diff --git a/.github/deploy-on-aws.png b/.github/deploy-on-aws.png
new file mode 100644
index 00000000000..06d41f2a5e0
Binary files /dev/null and b/.github/deploy-on-aws.png differ
diff --git a/.github/deploy-on-gcp.png b/.github/deploy-on-gcp.png
new file mode 100644
index 00000000000..e831a8c2e4e
Binary files /dev/null and b/.github/deploy-on-gcp.png differ
diff --git a/.github/workflows/osv-scan.yml b/.github/workflows/osv-scan.yml
index 9dd321f88db..0cd94fdd9e2 100644
--- a/.github/workflows/osv-scan.yml
+++ b/.github/workflows/osv-scan.yml
@@ -7,11 +7,6 @@ on:
- litellm_internal_staging
- litellm_oss_branch
- "litellm_**"
- paths:
- - uv.lock
- - ui/litellm-dashboard/package-lock.json
- - osv-scanner.toml
- - .github/workflows/osv-scan.yml
schedule:
- cron: "23 6 * * *"
workflow_dispatch:
diff --git a/.github/workflows/test-linting.yml b/.github/workflows/test-linting.yml
index de7e1b68346..950d6ca31a6 100644
--- a/.github/workflows/test-linting.yml
+++ b/.github/workflows/test-linting.yml
@@ -14,7 +14,7 @@ permissions:
jobs:
lint:
runs-on: ubuntu-latest
- timeout-minutes: 10
+ timeout-minutes: 15
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
@@ -87,9 +87,11 @@ jobs:
run: |
uv run --no-sync python -c "import openai; print(f'OpenAI version: {openai.__version__}')"
- - name: Run basedpyright type checking
+ - name: Check basedpyright budget (delta vs base)
+ env:
+ BASE_SHA: ${{ github.event.pull_request.base.sha }}
run: |
- (uv run --no-sync basedpyright --outputjson || true) | uv run --no-sync python scripts/type_check_gate.py
+ (uv run --no-sync basedpyright --outputjson || true) | uv run --no-sync python scripts/type_check_gate.py --base "$BASE_SHA"
- name: Check for circular imports
run: |
diff --git a/.github/workflows/test-rust.yml b/.github/workflows/test-rust.yml
new file mode 100644
index 00000000000..3d0a159cdc7
--- /dev/null
+++ b/.github/workflows/test-rust.yml
@@ -0,0 +1,65 @@
+name: LiteLLM Rust
+
+on:
+ push:
+ paths:
+ - "litellm-rust/**"
+ - ".github/workflows/test-rust.yml"
+ pull_request:
+ branches:
+ - main
+ - litellm_internal_staging
+ - litellm_oss_branch
+ - "litellm_**"
+ paths:
+ - "litellm-rust/**"
+ - ".github/workflows/test-rust.yml"
+
+permissions:
+ contents: read
+
+concurrency:
+ group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
+ cancel-in-progress: true
+
+jobs:
+ rust-checks:
+ name: rustfmt, clippy, test
+ runs-on: ubuntu-latest
+ timeout-minutes: 10
+ defaults:
+ run:
+ working-directory: litellm-rust
+ env:
+ CARGO_TERM_COLOR: always
+
+ steps:
+ - name: Checkout repository
+ uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
+ with:
+ persist-credentials: false
+
+ - name: Set up Rust
+ run: |
+ rustup toolchain install stable --profile minimal --component clippy,rustfmt
+ rustup default stable
+
+ - name: Cache Cargo registry and target
+ uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
+ with:
+ path: |
+ ~/.cargo/registry
+ ~/.cargo/git
+ litellm-rust/target
+ key: ${{ runner.os }}-cargo-${{ hashFiles('litellm-rust/Cargo.lock') }}
+ restore-keys: |
+ ${{ runner.os }}-cargo-
+
+ - name: Check Rust formatting
+ run: cargo fmt --check
+
+ - name: Run Clippy
+ run: cargo clippy --workspace --all-targets --locked -- -D warnings
+
+ - name: Run Rust tests
+ run: cargo test --workspace --locked
diff --git a/.github/workflows/test-unit-misc.yml b/.github/workflows/test-unit-misc.yml
index a7363ac3b43..2226d519331 100644
--- a/.github/workflows/test-unit-misc.yml
+++ b/.github/workflows/test-unit-misc.yml
@@ -32,7 +32,9 @@ jobs:
tests/test_litellm/repositories
tests/test_litellm/images
tests/test_litellm/interactions
+ tests/test_litellm/ocr
tests/test_litellm/passthrough
+ tests/test_litellm/sandbox
tests/test_litellm/vector_stores
tests/test_litellm/test_*.py
workers: 2
diff --git a/CLAUDE.md b/CLAUDE.md
index 2070b6fcdd6..b721064aaa7 100644
--- a/CLAUDE.md
+++ b/CLAUDE.md
@@ -29,7 +29,7 @@ If you ever make public-facing PR descriptions, comments, issues, commit message
- don't use "—". Instead, reach for ";", ".", etc.
- don't use the pattern "It's not X, it's Y", "You're not X, you're Y", etc.
- don't use bulleted or numbered lists unless it would be nonsensical not to. Instead, prefer prose
-- don't add a trailing "." at the end of paragraphs (just like this file)
+- don't add a trailing "." at the end of paragraphs (just like this file). That means every paragraph, not just the last one (of the markdown file, PR description, GitHub comment, etc.). Rule of thumb: unless there's a sentence immediately after, don't add a "."
- don't use →. Instead, prefer not to use arrows, and if need be, use -> instead
Don't hesitate to use values in .env to get needed API keys and other secrets, as long as you never add them to conversation history, commit them, or include them in GitHub issues / PRs
diff --git a/Dockerfile b/Dockerfile
index 4d55148ff89..af49dc8d8cf 100644
--- a/Dockerfile
+++ b/Dockerfile
@@ -1,8 +1,8 @@
# Base image for building
-ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:31da6565f35af6401031c1d7aa91dc84ac76c5c48edd17fb90f0ed9e3173c7a9
+ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:c61ac6919b811ea53c4782d69f1fe05218ba3c25d53f01b6ab7892e621bd4370
# Runtime image
-ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:31da6565f35af6401031c1d7aa91dc84ac76c5c48edd17fb90f0ed9e3173c7a9
+ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:c61ac6919b811ea53c4782d69f1fe05218ba3c25d53f01b6ab7892e621bd4370
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
FROM $UV_IMAGE AS uvbin
diff --git a/Makefile b/Makefile
index 27150aec938..076eac0f4a7 100644
--- a/Makefile
+++ b/Makefile
@@ -125,7 +125,8 @@ lint-ruff-FULL-dev: install-dev
else echo "No changed .py files to check."; fi
lint-basedpyright: install-dev
- ($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py
+ git fetch origin litellm_internal_staging
+ ($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py --base origin/litellm_internal_staging
lint-basedpyright-budget-update: install-dev
($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py --update
diff --git a/README.md b/README.md
index b26ad39eada..3d0f7282d7c 100644
--- a/README.md
+++ b/README.md
@@ -6,10 +6,10 @@
Open Source AI Gateway for 100+ LLMs. Self-hosted. Enterprise-ready. Call any LLM in OpenAI format.
-
-
-
-
+
+
+
+
@@ -406,6 +406,140 @@ You can use LiteLLM through either the Proxy Server or Python SDK. Both give you
Support for more providers. Missing a provider or LLM Platform, raise a [feature request](https://github.com/BerriAI/litellm/issues/new?assignees=&labels=enhancement&projects=&template=feature_request.yml&title=%5BFeature%5D%3A+).
+### Deploy on AWS or GCP with Terraform
+
+Run the LiteLLM proxy as a production-ready componentized stack (gateway, backend, UI on separate services; managed Postgres + Redis + object store) using the published Terraform modules. Both modules are on the [public Terraform Registry](https://registry.terraform.io/namespaces/BerriAI) — no auth needed.
+
+#### AWS — ECS Fargate + Aurora + ElastiCache + ALB
+
+[](https://console.aws.amazon.com/cloudshell/home) — opens an in-browser shell, already authenticated to your AWS account. Once inside, run:
+
+```bash
+git clone https://github.com/BerriAI/litellm.git
+cd litellm/terraform/litellm/aws/examples/default
+cp terraform.tfvars.example terraform.tfvars # edit region/tenant/env
+terraform init && terraform apply
+```
+
+[Module page →](https://registry.terraform.io/modules/BerriAI/litellm/aws/latest)
+
+Or call the module from your own root config:
+
+```hcl
+# main.tf
+terraform {
+ required_version = ">= 1.6.0"
+ required_providers {
+ aws = { source = "hashicorp/aws", version = "~> 5.60" }
+ }
+}
+
+provider "aws" {
+ region = "us-west-2"
+}
+
+module "litellm" {
+ source = "BerriAI/litellm/aws"
+ version = "~> 1.89"
+
+ region = "us-west-2"
+ azs = ["us-west-2a", "us-west-2b"]
+ tenant = "acme"
+ env = "prod"
+
+ # Production: provide an ACM cert. Without one, set allow_plaintext_alb = true
+ # (dev/trial only).
+ # acm_certificate_arn = "arn:aws:acm:us-west-2:111122223333:certificate/..."
+ allow_plaintext_alb = true
+}
+
+output "litellm_url" {
+ value = module.litellm.alb_dns_name
+}
+```
+
+```bash
+terraform init
+terraform apply
+```
+
+Provider API keys live in AWS Secrets Manager; reference ARNs via `gateway_extra_secrets`. Full input list and architecture diagram on the [registry page](https://registry.terraform.io/modules/BerriAI/litellm/aws/latest?tab=inputs).
+
+#### GCP — Cloud Run + Cloud SQL + Memorystore + HTTPS LB
+
+[](https://ssh.cloud.google.com/cloudshell/editor?cloudshell_git_repo=https%3A%2F%2Fgithub.com%2FBerriAI%2Flitellm&cloudshell_workspace=terraform%2Flitellm%2Fgcp%2Fexamples%2Fdefault&cloudshell_tutorial=TUTORIAL.md&cloudshell_image=gcr.io/ds-artifacts-cloudshell/deploystack_custom_image&shellonly=true)
+
+Real 1-click. Opens Cloud Shell, clones this repo, and walks you through `terraform apply` via a built-in [DeployStack tutorial](./terraform/litellm/gcp/examples/default/TUTORIAL.md) — pick the project, the tutorial sets up the Artifact Registry remote repo, writes `terraform.tfvars` from your answers, and runs apply.
+
+[Module page →](https://registry.terraform.io/modules/BerriAI/litellm/google/latest)
+
+To call the module from your own config instead, Cloud Run can't pull from `ghcr.io` directly, so first set up a one-time Artifact Registry remote repo backed by GHCR:
+
+```bash
+gcloud artifacts repositories create litellm \
+ --location=us-central1 \
+ --repository-format=docker \
+ --mode=remote-repository \
+ --remote-docker-repo=https://ghcr.io \
+ --project=my-gcp-project
+```
+
+Then:
+
+```hcl
+# main.tf
+terraform {
+ required_version = ">= 1.6.0"
+ required_providers {
+ google = { source = "hashicorp/google", version = "~> 6.10" }
+ google-beta = { source = "hashicorp/google-beta", version = "~> 6.10" }
+ }
+}
+
+provider "google" { project = "my-gcp-project"; region = "us-central1" }
+provider "google-beta" { project = "my-gcp-project"; region = "us-central1" }
+
+module "litellm" {
+ source = "BerriAI/litellm/google"
+ version = "~> 1.89"
+
+ project_id = "my-gcp-project"
+ region = "us-central1"
+ tenant = "acme"
+ env = "prod"
+
+ # Replace my-gcp-project with your GCP project ID (same value as project_id above).
+ image_registry = "us-central1-docker.pkg.dev/my-gcp-project/litellm/berriai"
+
+ # Production: provide DNS already pointing at the LB IP for Google-managed certs.
+ # Without one, set allow_plaintext_lb = true (dev/trial only).
+ # lb_domains = ["proxy.example.com"]
+ allow_plaintext_lb = true
+}
+
+output "litellm_url" {
+ value = module.litellm.load_balancer_url
+}
+```
+
+```bash
+terraform init
+terraform apply
+```
+
+Provider API keys live in Secret Manager; reference resource IDs (e.g. `projects/my-gcp-project/secrets/openai-api-key`) via `gateway_extra_secrets`. Full input list and architecture diagram on the [registry page](https://registry.terraform.io/modules/BerriAI/litellm/google/latest?tab=inputs).
+
+#### Both stacks include
+
+- The full componentized split (gateway / backend / UI as independent services)
+- Managed Postgres (writer + reader) and Redis
+- Versioned object store for proxy state + file uploads
+- An auto-generated `LITELLM_MASTER_KEY` in your cloud's secret manager
+- A one-off migration job that runs `prisma migrate deploy` before the proxy starts
+- The same `proxy_config` surface as the [Helm chart](./helm/litellm/) — pass YAML as a typed map
+
+The Terraform modules live at [`terraform/litellm/aws/`](./terraform/litellm/aws/) and [`terraform/litellm/gcp/`](./terraform/litellm/gcp/) in this repo; the registry entries are read-only mirrors updated on each release.
+
### Run in Developer Mode
#### Services
1. Setup .env file in root
diff --git a/backend/Dockerfile b/backend/Dockerfile
index 2cfdde8a517..667bdb073eb 100644
--- a/backend/Dockerfile
+++ b/backend/Dockerfile
@@ -1,5 +1,5 @@
-ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:31da6565f35af6401031c1d7aa91dc84ac76c5c48edd17fb90f0ed9e3173c7a9
-ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:31da6565f35af6401031c1d7aa91dc84ac76c5c48edd17fb90f0ed9e3173c7a9
+ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:c61ac6919b811ea53c4782d69f1fe05218ba3c25d53f01b6ab7892e621bd4370
+ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:c61ac6919b811ea53c4782d69f1fe05218ba3c25d53f01b6ab7892e621bd4370
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
FROM $UV_IMAGE AS uvbin
diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json
index 7ba7656e407..f5b0a9aaf81 100644
--- a/basedpyright-code-budget.json
+++ b/basedpyright-code-budget.json
@@ -121,7 +121,7 @@
},
"reportReturnType": {
"baseline": 126,
- "slack": 13
+ "slack": 100
},
"reportTypedDictNotRequiredAccess": {
"baseline": 20,
@@ -157,7 +157,7 @@
},
"reportUnnecessaryComparison": {
"baseline": 683,
- "slack": 10
+ "slack": 100
},
"reportUnnecessaryContains": {
"baseline": 4,
diff --git a/docker/Dockerfile.database b/docker/Dockerfile.database
index e591a4a2adb..50ef55e3261 100644
--- a/docker/Dockerfile.database
+++ b/docker/Dockerfile.database
@@ -1,8 +1,8 @@
# Base image for building
-ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:3258be472764337fd13095bcbb3182da170243b5819fd67ad4c0754590588b31
+ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:c61ac6919b811ea53c4782d69f1fe05218ba3c25d53f01b6ab7892e621bd4370
# Runtime image
-ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:3258be472764337fd13095bcbb3182da170243b5819fd67ad4c0754590588b31
+ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:c61ac6919b811ea53c4782d69f1fe05218ba3c25d53f01b6ab7892e621bd4370
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
FROM $UV_IMAGE AS uvbin
diff --git a/docker/Dockerfile.non_root b/docker/Dockerfile.non_root
index eafbd23fd90..ab02b43d0f9 100644
--- a/docker/Dockerfile.non_root
+++ b/docker/Dockerfile.non_root
@@ -1,6 +1,6 @@
# Base images
-ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:3258be472764337fd13095bcbb3182da170243b5819fd67ad4c0754590588b31
-ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:3258be472764337fd13095bcbb3182da170243b5819fd67ad4c0754590588b31
+ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:c61ac6919b811ea53c4782d69f1fe05218ba3c25d53f01b6ab7892e621bd4370
+ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:c61ac6919b811ea53c4782d69f1fe05218ba3c25d53f01b6ab7892e621bd4370
ARG PROXY_EXTRAS_SOURCE=published
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
diff --git a/gateway/Dockerfile b/gateway/Dockerfile
index 19c8a10fdfe..716b2fa09d1 100644
--- a/gateway/Dockerfile
+++ b/gateway/Dockerfile
@@ -1,5 +1,5 @@
-ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:31da6565f35af6401031c1d7aa91dc84ac76c5c48edd17fb90f0ed9e3173c7a9
-ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:31da6565f35af6401031c1d7aa91dc84ac76c5c48edd17fb90f0ed9e3173c7a9
+ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:c61ac6919b811ea53c4782d69f1fe05218ba3c25d53f01b6ab7892e621bd4370
+ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:c61ac6919b811ea53c4782d69f1fe05218ba3c25d53f01b6ab7892e621bd4370
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
FROM $UV_IMAGE AS uvbin
diff --git a/litellm-rust/.gitignore b/litellm-rust/.gitignore
new file mode 100644
index 00000000000..b83d22266ac
--- /dev/null
+++ b/litellm-rust/.gitignore
@@ -0,0 +1 @@
+/target/
diff --git a/litellm-rust/ADDING_A_PROVIDER.md b/litellm-rust/ADDING_A_PROVIDER.md
new file mode 100644
index 00000000000..2fa81798605
--- /dev/null
+++ b/litellm-rust/ADDING_A_PROVIDER.md
@@ -0,0 +1,9 @@
+# Adding a provider / route to litellm-rust
+
+Three layers, same for every route (see `ocr` and `realtime` as references):
+
+1. **Transform contract (pure)** — `crates/core/src//transformation.rs`: a `…ProviderConfig` trait (URL build + request/response transforms) + types in `types.rs`. No network, env, or auth.
+2. **Provider config (pure)** — `crates/providers/src///transformation.rs`: implement that trait as a `const __CONFIG`, mirroring the Python provider tree. Add parity unit tests.
+3. **HTTP / transport (the host)** — `crates/providers/src/.rs` (e.g. `ocr.rs`, `realtime.rs`): the callable fn (`run_ocr`, `realtime`). It resolves the key, builds the auth header, builds URL + transforms via the config, then does the network call. This is the only layer allowed to do I/O.
+
+**Calling:** the host invokes the route fn — the Python bridge calls `run_ocr`; the `ai-gateway` server calls `realtime`. Register new modules in `lib.rs` / `mod.rs`, then run `cargo fmt && cargo clippy --workspace -- -D warnings && cargo test --workspace`.
diff --git a/litellm-rust/CLAUDE.md b/litellm-rust/CLAUDE.md
new file mode 100644
index 00000000000..1d2987e0a1a
--- /dev/null
+++ b/litellm-rust/CLAUDE.md
@@ -0,0 +1,88 @@
+# CLAUDE.md
+
+This file defines the rules for Rust work in LiteLLM.
+
+## Core Boundary
+
+The `core` and `providers` crates describe work; hosts execute work.
+
+Route-level Rust structure mirrors LiteLLM's Python responsibilities:
+- `core/src//` owns the route contract, shared types, and provider
+ template traits. For OCR, this means `core/src/ocr`.
+- `providers/src///transformation.rs` owns the
+ provider-specific transform. For Mistral OCR, this means
+ `providers/src/mistral/ocr/transformation.rs`.
+- Future network execution belongs in a host/transport layer such as
+ `llm_http_handler`, not inside `core` or `providers`.
+
+Allowed in `core` and `providers`:
+- Pure request transforms
+- Pure response transforms
+- Pure stream chunk normalization
+- Shared data types and validation errors
+- Deterministic token/cost helper logic
+
+Not allowed in `core` or `providers`:
+- Network calls
+- Environment variable or secret reads
+- Filesystem access
+- Database or cache access
+- Provider SDK signing or auth flows
+- Logging callbacks, spend writes, or custom callbacks
+- Global mutable runtime state
+
+Python owns rollout state and fallback while Rust is being introduced. Rust
+paths must be off by default until parity tests prove equivalence with Python.
+
+## Production Bar
+
+Rust code in this workspace is held to a strict parity and robustness bar from
+the first PR:
+
+- Correctness parity is proven with tests. Do not rely on README claims or
+ manual inspection for a port that mirrors Python behavior.
+- Every provider transform must have unit tests for supported-parameter
+ filtering, request body shape, response normalization, missing/null fields,
+ and bad-input errors.
+- When Rust is exposed through Python, add Python tests that prove disabled,
+ enabled, and unavailable-bridge fallback behavior.
+- Avoid panics on user/provider input. Return typed errors and let the host map
+ them to Python exceptions or HTTP responses.
+- OCR handles documents that often contain personal data. Do not log document
+ contents, base64 payloads, provider response bodies, or secrets.
+- Error messages must be useful but data-minimized. Truncate or sanitize any
+ upstream body before it crosses a host boundary.
+- Treat empty or whitespace-only credentials, URLs, and config values as absent
+ at the host/config resolution layer.
+- Preserve Python output shape intentionally. If a field is always serialized as
+ `null` for Python parity, leave a short comment explaining that parity choice.
+
+## Host I/O Rules
+
+These rules apply when adding future crates or modules that execute network I/O,
+such as `ai-gateway`, router hosts, or standalone servers:
+
+- Set connect and full-request timeouts. No unbounded waits.
+- Reuse HTTP clients; do not construct clients per request.
+- Prefer rustls TLS for portable Python wheels and Linux images unless there is
+ a documented reason not to.
+- Add request IDs and structured tracing at the host layer, without logging OCR
+ document contents or secrets.
+- Do not echo raw upstream response bodies to callers. Sanitize and bound them.
+- Avoid `expect`/`unwrap` in server startup and request paths unless the panic is
+ impossible by construction and documented.
+
+## Checks
+
+Run these before pushing Rust changes. The same checks run in GitHub Actions
+for changes under `litellm-rust/`.
+
+```bash
+cd litellm-rust
+cargo fmt --check
+cargo clippy --workspace --all-targets -- -D warnings
+cargo test --workspace
+```
+
+When a Rust path is exposed through Python, add Python parity tests that compare
+the existing Python output with the Rust-backed output.
diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock
new file mode 100644
index 00000000000..a269a224d97
--- /dev/null
+++ b/litellm-rust/Cargo.lock
@@ -0,0 +1,1872 @@
+# This file is automatically @generated by Cargo.
+# It is not intended for manual editing.
+version = 4
+
+[[package]]
+name = "async-trait"
+version = "0.1.89"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "9035ad2d096bed7955a320ee7e2230574d28fd3c3a0f186cbea1ff3c7eed5dbb"
+dependencies = [
+ "proc-macro2",
+ "quote",
+ "syn",
+]
+
+[[package]]
+name = "atomic-waker"
+version = "1.1.2"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "1505bd5d3d116872e7271a6d4e16d81d0c8570876c8de68093a09ac269d8aac0"
+
+[[package]]
+name = "autocfg"
+version = "1.5.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "f2032f911046de80f0a198e0901378627c33f59ea0ac00e363d481118bd70a53"
+
+[[package]]
+name = "axum"
+version = "0.7.9"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "edca88bc138befd0323b20752846e6587272d3b03b0343c8ea28a6f819e6e71f"
+dependencies = [
+ "async-trait",
+ "axum-core",
+ "base64",
+ "bytes",
+ "futures-util",
+ "http",
+ "http-body",
+ "http-body-util",
+ "hyper",
+ "hyper-util",
+ "itoa",
+ "matchit",
+ "memchr",
+ "mime",
+ "percent-encoding",
+ "pin-project-lite",
+ "rustversion",
+ "serde",
+ "serde_json",
+ "serde_path_to_error",
+ "serde_urlencoded",
+ "sha1",
+ "sync_wrapper",
+ "tokio",
+ "tokio-tungstenite",
+ "tower",
+ "tower-layer",
+ "tower-service",
+ "tracing",
+]
+
+[[package]]
+name = "axum-core"
+version = "0.4.5"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "09f2bd6146b97ae3359fa0cc6d6b376d9539582c7b4220f041a33ec24c226199"
+dependencies = [
+ "async-trait",
+ "bytes",
+ "futures-util",
+ "http",
+ "http-body",
+ "http-body-util",
+ "mime",
+ "pin-project-lite",
+ "rustversion",
+ "sync_wrapper",
+ "tower-layer",
+ "tower-service",
+ "tracing",
+]
+
+[[package]]
+name = "base64"
+version = "0.22.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "72b3254f16251a8381aa12e40e3c4d2f0199f8c6508fbecb9d91f575e0fbb8c6"
+
+[[package]]
+name = "bitflags"
+version = "2.13.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "b4388bee8683e3d04af747c73422af53102d2bd24d9eadb6cbc100baef4b43f8"
+
+[[package]]
+name = "block-buffer"
+version = "0.10.4"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "3078c7629b62d3f0439517fa394996acacc5cbc91c5a20d8c658e77abd503a71"
+dependencies = [
+ "generic-array",
+]
+
+[[package]]
+name = "bumpalo"
+version = "3.20.3"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "72f5acc6cb2ba439de613abc23857ec3d78374d8ed5ac84e9d11336e87da8649"
+
+[[package]]
+name = "byteorder"
+version = "1.5.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "1fd0f2584146f6f2ef48085050886acf353beff7305ebd1ae69500e27c67f64b"
+
+[[package]]
+name = "bytes"
+version = "1.12.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "8ae3f5d315924270530207e2a68396c3cc547f6dca3fbdca317cfb1a51edb593"
+
+[[package]]
+name = "cc"
+version = "1.2.65"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "e228eec9be7c17ccb640b59b36a5cd805ea2a564a4c5e162c2f659fea30d3b96"
+dependencies = [
+ "find-msvc-tools",
+ "shlex",
+]
+
+[[package]]
+name = "cfg-if"
+version = "1.0.4"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "9330f8b2ff13f34540b44e946ef35111825727b38d33286ef986142615121801"
+
+[[package]]
+name = "cfg_aliases"
+version = "0.2.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724"
+
+[[package]]
+name = "core-foundation"
+version = "0.10.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "b2a6cd9ae233e7f62ba4e9353e81a88df7fc8a5987b8d445b4d90c879bd156f6"
+dependencies = [
+ "core-foundation-sys",
+ "libc",
+]
+
+[[package]]
+name = "core-foundation-sys"
+version = "0.8.7"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "773648b94d0e5d620f64f280777445740e61fe701025087ec8b57f45c791888b"
+
+[[package]]
+name = "cpufeatures"
+version = "0.2.17"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "59ed5838eebb26a2bb2e58f6d5b5316989ae9d08bab10e0e6d103e656d1b0280"
+dependencies = [
+ "libc",
+]
+
+[[package]]
+name = "crypto-common"
+version = "0.1.7"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "78c8292055d1c1df0cce5d180393dc8cce0abec0a7102adb6c7b1eef6016d60a"
+dependencies = [
+ "generic-array",
+ "typenum",
+]
+
+[[package]]
+name = "data-encoding"
+version = "2.11.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "a4ae5f15dda3c708c0ade84bfee31ccab44a3da4f88015ed22f63732abe300c8"
+
+[[package]]
+name = "digest"
+version = "0.10.7"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292"
+dependencies = [
+ "block-buffer",
+ "crypto-common",
+]
+
+[[package]]
+name = "displaydoc"
+version = "0.2.6"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "1ac70aa55017e108007fbaf5aa0f54b021c98f92ff8af59d42eda9da96e3dd4f"
+dependencies = [
+ "proc-macro2",
+ "quote",
+ "syn",
+]
+
+[[package]]
+name = "find-msvc-tools"
+version = "0.1.9"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "5baebc0774151f905a1a2cc41989300b1e6fbb29aff0ceffa1064fdd3088d582"
+
+[[package]]
+name = "form_urlencoded"
+version = "1.2.2"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "cb4cb245038516f5f85277875cdaa4f7d2c9a0fa0468de06ed190163b1581fcf"
+dependencies = [
+ "percent-encoding",
+]
+
+[[package]]
+name = "futures-channel"
+version = "0.3.32"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "07bbe89c50d7a535e539b8c17bc0b49bdb77747034daa8087407d655f3f7cc1d"
+dependencies = [
+ "futures-core",
+ "futures-sink",
+]
+
+[[package]]
+name = "futures-core"
+version = "0.3.32"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "7e3450815272ef58cec6d564423f6e755e25379b217b0bc688e295ba24df6b1d"
+
+[[package]]
+name = "futures-io"
+version = "0.3.32"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "cecba35d7ad927e23624b22ad55235f2239cfa44fd10428eecbeba6d6a717718"
+
+[[package]]
+name = "futures-sink"
+version = "0.3.32"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "c39754e157331b013978ec91992bde1ac089843443c49cbc7f46150b0fad0893"
+
+[[package]]
+name = "futures-task"
+version = "0.3.32"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "037711b3d59c33004d3856fbdc83b99d4ff37a24768fa1be9ce3538a1cde4393"
+
+[[package]]
+name = "futures-util"
+version = "0.3.32"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "389ca41296e6190b48053de0321d02a77f32f8a5d2461dd38762c0593805c6d6"
+dependencies = [
+ "futures-core",
+ "futures-io",
+ "futures-sink",
+ "futures-task",
+ "memchr",
+ "pin-project-lite",
+ "slab",
+]
+
+[[package]]
+name = "generic-array"
+version = "0.14.7"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "85649ca51fd72272d7821adaf274ad91c288277713d9c18820d8499a7ff69e9a"
+dependencies = [
+ "typenum",
+ "version_check",
+]
+
+[[package]]
+name = "getrandom"
+version = "0.2.17"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "ff2abc00be7fca6ebc474524697ae276ad847ad0a6b3faa4bcb027e9a4614ad0"
+dependencies = [
+ "cfg-if",
+ "js-sys",
+ "libc",
+ "wasi",
+ "wasm-bindgen",
+]
+
+[[package]]
+name = "getrandom"
+version = "0.3.4"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "899def5c37c4fd7b2664648c28120ecec138e4d395b459e5ca34f9cce2dd77fd"
+dependencies = [
+ "cfg-if",
+ "js-sys",
+ "libc",
+ "r-efi",
+ "wasip2",
+ "wasm-bindgen",
+]
+
+[[package]]
+name = "heck"
+version = "0.5.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea"
+
+[[package]]
+name = "http"
+version = "1.4.2"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "6970f50e31d6fc17d3fa27329444bfa74e196cf62e95052a3f6fee181dba6425"
+dependencies = [
+ "bytes",
+ "itoa",
+]
+
+[[package]]
+name = "http-body"
+version = "1.0.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "1efedce1fb8e6913f23e0c92de8e62cd5b772a67e7b3946df930a62566c93184"
+dependencies = [
+ "bytes",
+ "http",
+]
+
+[[package]]
+name = "http-body-util"
+version = "0.1.3"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "b021d93e26becf5dc7e1b75b1bed1fd93124b374ceb73f43d4d4eafec896a64a"
+dependencies = [
+ "bytes",
+ "futures-core",
+ "http",
+ "http-body",
+ "pin-project-lite",
+]
+
+[[package]]
+name = "httparse"
+version = "1.10.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "6dbf3de79e51f3d586ab4cb9d5c3e2c14aa28ed23d180cf89b4df0454a69cc87"
+
+[[package]]
+name = "httpdate"
+version = "1.0.3"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "df3b46402a9d5adb4c86a0cf463f42e19994e3ee891101b1841f30a545cb49a9"
+
+[[package]]
+name = "hyper"
+version = "1.10.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "55281c53a1894c864990125767da440a4e630446785086f52523b20033b74498"
+dependencies = [
+ "atomic-waker",
+ "bytes",
+ "futures-channel",
+ "futures-core",
+ "http",
+ "http-body",
+ "httparse",
+ "httpdate",
+ "itoa",
+ "pin-project-lite",
+ "smallvec",
+ "tokio",
+ "want",
+]
+
+[[package]]
+name = "hyper-rustls"
+version = "0.27.9"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "33ca68d021ef39cf6463ab54c1d0f5daf03377b70561305bb89a8f83aab66e0f"
+dependencies = [
+ "http",
+ "hyper",
+ "hyper-util",
+ "rustls",
+ "tokio",
+ "tokio-rustls",
+ "tower-service",
+ "webpki-roots",
+]
+
+[[package]]
+name = "hyper-util"
+version = "0.1.20"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "96547c2556ec9d12fb1578c4eaf448b04993e7fb79cbaad930a656880a6bdfa0"
+dependencies = [
+ "base64",
+ "bytes",
+ "futures-channel",
+ "futures-util",
+ "http",
+ "http-body",
+ "hyper",
+ "ipnet",
+ "libc",
+ "percent-encoding",
+ "pin-project-lite",
+ "socket2",
+ "tokio",
+ "tower-service",
+ "tracing",
+]
+
+[[package]]
+name = "icu_collections"
+version = "2.2.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "2984d1cd16c883d7935b9e07e44071dca8d917fd52ecc02c04d5fa0b5a3f191c"
+dependencies = [
+ "displaydoc",
+ "potential_utf",
+ "utf8_iter",
+ "yoke",
+ "zerofrom",
+ "zerovec",
+]
+
+[[package]]
+name = "icu_locale_core"
+version = "2.2.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "92219b62b3e2b4d88ac5119f8904c10f8f61bf7e95b640d25ba3075e6cac2c29"
+dependencies = [
+ "displaydoc",
+ "litemap",
+ "tinystr",
+ "writeable",
+ "zerovec",
+]
+
+[[package]]
+name = "icu_normalizer"
+version = "2.2.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "c56e5ee99d6e3d33bd91c5d85458b6005a22140021cc324cea84dd0e72cff3b4"
+dependencies = [
+ "icu_collections",
+ "icu_normalizer_data",
+ "icu_properties",
+ "icu_provider",
+ "smallvec",
+ "zerovec",
+]
+
+[[package]]
+name = "icu_normalizer_data"
+version = "2.2.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "da3be0ae77ea334f4da67c12f149704f19f81d1adf7c51cf482943e84a2bad38"
+
+[[package]]
+name = "icu_properties"
+version = "2.2.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "bee3b67d0ea5c2cca5003417989af8996f8604e34fb9ddf96208a033901e70de"
+dependencies = [
+ "icu_collections",
+ "icu_locale_core",
+ "icu_properties_data",
+ "icu_provider",
+ "zerotrie",
+ "zerovec",
+]
+
+[[package]]
+name = "icu_properties_data"
+version = "2.2.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "8e2bbb201e0c04f7b4b3e14382af113e17ba4f63e2c9d2ee626b720cbce54a14"
+
+[[package]]
+name = "icu_provider"
+version = "2.2.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "139c4cf31c8b5f33d7e199446eff9c1e02decfc2f0eec2c8d71f65befa45b421"
+dependencies = [
+ "displaydoc",
+ "icu_locale_core",
+ "writeable",
+ "yoke",
+ "zerofrom",
+ "zerotrie",
+ "zerovec",
+]
+
+[[package]]
+name = "idna"
+version = "1.1.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "3b0875f23caa03898994f6ddc501886a45c7d3d62d04d2d90788d47be1b1e4de"
+dependencies = [
+ "idna_adapter",
+ "smallvec",
+ "utf8_iter",
+]
+
+[[package]]
+name = "idna_adapter"
+version = "1.2.2"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "cb68373c0d6620ef8105e855e7745e18b0d00d3bdb07fb532e434244cdb9a714"
+dependencies = [
+ "icu_normalizer",
+ "icu_properties",
+]
+
+[[package]]
+name = "indoc"
+version = "2.0.7"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "79cf5c93f93228cf8efb3ba362535fb11199ac548a09ce117c9b1adc3030d706"
+dependencies = [
+ "rustversion",
+]
+
+[[package]]
+name = "ipnet"
+version = "2.12.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "d98f6fed1fde3f8c21bc40a1abb88dd75e67924f9cffc3ef95607bad8017f8e2"
+
+[[package]]
+name = "itoa"
+version = "1.0.18"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "8f42a60cbdf9a97f5d2305f08a87dc4e09308d1276d28c869c684d7777685682"
+
+[[package]]
+name = "js-sys"
+version = "0.3.102"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "03d04c30968dffe80775bd4d7fb676131cd04a1fb46d2686dbffbaec2d9dfd31"
+dependencies = [
+ "cfg-if",
+ "futures-util",
+ "wasm-bindgen",
+]
+
+[[package]]
+name = "libc"
+version = "0.2.186"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "68ab91017fe16c622486840e4c83c9a37afeff978bd239b5293d61ece587de66"
+
+[[package]]
+name = "litellm-ai-gateway"
+version = "0.1.0"
+dependencies = [
+ "axum",
+ "futures-util",
+ "litellm-core",
+ "litellm-providers",
+ "pyo3",
+ "serde",
+ "serde_json",
+ "subtle",
+ "tokio",
+]
+
+[[package]]
+name = "litellm-core"
+version = "0.1.0"
+dependencies = [
+ "rand 0.8.6",
+ "serde",
+ "serde_json",
+ "thiserror 2.0.18",
+]
+
+[[package]]
+name = "litellm-providers"
+version = "0.1.0"
+dependencies = [
+ "futures-channel",
+ "futures-util",
+ "litellm-core",
+ "reqwest",
+ "serde_json",
+ "tokio",
+ "tokio-tungstenite",
+]
+
+[[package]]
+name = "litellm-python-bridge"
+version = "0.1.0"
+dependencies = [
+ "litellm-core",
+ "litellm-providers",
+ "pyo3",
+ "serde_json",
+]
+
+[[package]]
+name = "litemap"
+version = "0.8.2"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "92daf443525c4cce67b150400bc2316076100ce0b3686209eb8cf3c31612e6f0"
+
+[[package]]
+name = "log"
+version = "0.4.33"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "0ceec5bc11778974d1bcb055b18002eba7f4b3518b6a0081b3af5f21666da9ad"
+
+[[package]]
+name = "lru-slab"
+version = "0.1.2"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "112b39cec0b298b6c1999fee3e31427f74f676e4cb9879ed1a121b43661a4154"
+
+[[package]]
+name = "matchit"
+version = "0.7.3"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "0e7465ac9959cc2b1404e8e2367b43684a6d13790fe23056cc8c6c5a6b7bcb94"
+
+[[package]]
+name = "memchr"
+version = "2.8.2"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "88904434abc2901f197fe8cc55f0445e7ded921dba5911dad2e2b39b48e663c4"
+
+[[package]]
+name = "memoffset"
+version = "0.9.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "488016bfae457b036d996092f6cb448677611ce4449e970ceaf42695203f218a"
+dependencies = [
+ "autocfg",
+]
+
+[[package]]
+name = "mime"
+version = "0.3.17"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "6877bb514081ee2a7ff5ef9de3281f14a4dd4bceac4c09388074a6b5df8a139a"
+
+[[package]]
+name = "mio"
+version = "1.2.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "02bd0af71c67b473010cbbc60715ee815645a4dc942899111f494b4b737d6fda"
+dependencies = [
+ "libc",
+ "wasi",
+ "windows-sys 0.61.2",
+]
+
+[[package]]
+name = "once_cell"
+version = "1.21.4"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50"
+
+[[package]]
+name = "openssl-probe"
+version = "0.2.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "7c87def4c32ab89d880effc9e097653c8da5d6ef28e6b539d313baaacfbafcbe"
+
+[[package]]
+name = "percent-encoding"
+version = "2.3.2"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "9b4f627cb1b25917193a259e49bdad08f671f8d9708acfd5fe0a8c1455d87220"
+
+[[package]]
+name = "pin-project-lite"
+version = "0.2.17"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "a89322df9ebe1c1578d689c92318e070967d1042b512afbe49518723f4e6d5cd"
+
+[[package]]
+name = "portable-atomic"
+version = "1.13.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "c33a9471896f1c69cecef8d20cbe2f7accd12527ce60845ff44c153bb2a21b49"
+
+[[package]]
+name = "potential_utf"
+version = "0.1.5"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "0103b1cef7ec0cf76490e969665504990193874ea05c85ff9bab8b911d0a0564"
+dependencies = [
+ "zerovec",
+]
+
+[[package]]
+name = "ppv-lite86"
+version = "0.2.21"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "85eae3c4ed2f50dcfe72643da4befc30deadb458a9b590d720cde2f2b1e97da9"
+dependencies = [
+ "zerocopy",
+]
+
+[[package]]
+name = "proc-macro2"
+version = "1.0.106"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "8fd00f0bb2e90d81d1044c2b32617f68fcb9fa3bb7640c23e9c748e53fb30934"
+dependencies = [
+ "unicode-ident",
+]
+
+[[package]]
+name = "pyo3"
+version = "0.23.5"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "7778bffd85cf38175ac1f545509665d0b9b92a198ca7941f131f85f7a4f9a872"
+dependencies = [
+ "cfg-if",
+ "indoc",
+ "libc",
+ "memoffset",
+ "once_cell",
+ "portable-atomic",
+ "pyo3-build-config",
+ "pyo3-ffi",
+ "pyo3-macros",
+ "unindent",
+]
+
+[[package]]
+name = "pyo3-build-config"
+version = "0.23.5"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "94f6cbe86ef3bf18998d9df6e0f3fc1050a8c5efa409bf712e661a4366e010fb"
+dependencies = [
+ "once_cell",
+ "target-lexicon",
+]
+
+[[package]]
+name = "pyo3-ffi"
+version = "0.23.5"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "e9f1b4c431c0bb1c8fb0a338709859eed0d030ff6daa34368d3b152a63dfdd8d"
+dependencies = [
+ "libc",
+ "pyo3-build-config",
+]
+
+[[package]]
+name = "pyo3-macros"
+version = "0.23.5"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "fbc2201328f63c4710f68abdf653c89d8dbc2858b88c5d88b0ff38a75288a9da"
+dependencies = [
+ "proc-macro2",
+ "pyo3-macros-backend",
+ "quote",
+ "syn",
+]
+
+[[package]]
+name = "pyo3-macros-backend"
+version = "0.23.5"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "fca6726ad0f3da9c9de093d6f116a93c1a38e417ed73bf138472cf4064f72028"
+dependencies = [
+ "heck",
+ "proc-macro2",
+ "pyo3-build-config",
+ "quote",
+ "syn",
+]
+
+[[package]]
+name = "quinn"
+version = "0.11.11"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "0c1a41e437b6bbd489372cd4971de128e85c855f56c57f283d20ff016cf7c0a8"
+dependencies = [
+ "bytes",
+ "cfg_aliases",
+ "pin-project-lite",
+ "quinn-proto",
+ "quinn-udp",
+ "rustc-hash",
+ "rustls",
+ "socket2",
+ "thiserror 2.0.18",
+ "tokio",
+ "tracing",
+ "web-time",
+]
+
+[[package]]
+name = "quinn-proto"
+version = "0.11.15"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "4fcb935c5bec503c2f0e306bdd3e58bb9029dcb14fa8d9ac76e3a5256ac0763e"
+dependencies = [
+ "bytes",
+ "getrandom 0.3.4",
+ "lru-slab",
+ "rand 0.9.4",
+ "ring",
+ "rustc-hash",
+ "rustls",
+ "rustls-pki-types",
+ "slab",
+ "thiserror 2.0.18",
+ "tinyvec",
+ "tracing",
+ "web-time",
+]
+
+[[package]]
+name = "quinn-udp"
+version = "0.5.14"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "addec6a0dcad8a8d96a771f815f0eaf55f9d1805756410b39f5fa81332574cbd"
+dependencies = [
+ "cfg_aliases",
+ "libc",
+ "once_cell",
+ "socket2",
+ "tracing",
+ "windows-sys 0.60.2",
+]
+
+[[package]]
+name = "quote"
+version = "1.0.46"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "dfbc457d0c7a0759a614551b11a6409e5951f6c7537be1f1b7682b9ae9230368"
+dependencies = [
+ "proc-macro2",
+]
+
+[[package]]
+name = "r-efi"
+version = "5.3.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "69cdb34c158ceb288df11e18b4bd39de994f6657d83847bdffdbd7f346754b0f"
+
+[[package]]
+name = "rand"
+version = "0.8.6"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "5ca0ecfa931c29007047d1bc58e623ab12e5590e8c7cc53200d5202b69266d8a"
+dependencies = [
+ "libc",
+ "rand_chacha 0.3.1",
+ "rand_core 0.6.4",
+]
+
+[[package]]
+name = "rand"
+version = "0.9.4"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "44c5af06bb1b7d3216d91932aed5265164bf384dc89cd6ba05cf59a35f5f76ea"
+dependencies = [
+ "rand_chacha 0.9.0",
+ "rand_core 0.9.5",
+]
+
+[[package]]
+name = "rand_chacha"
+version = "0.3.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "e6c10a63a0fa32252be49d21e7709d4d4baf8d231c2dbce1eaa8141b9b127d88"
+dependencies = [
+ "ppv-lite86",
+ "rand_core 0.6.4",
+]
+
+[[package]]
+name = "rand_chacha"
+version = "0.9.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "d3022b5f1df60f26e1ffddd6c66e8aa15de382ae63b3a0c1bfc0e4d3e3f325cb"
+dependencies = [
+ "ppv-lite86",
+ "rand_core 0.9.5",
+]
+
+[[package]]
+name = "rand_core"
+version = "0.6.4"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "ec0be4795e2f6a28069bec0b5ff3e2ac9bafc99e6a9a7dc3547996c5c816922c"
+dependencies = [
+ "getrandom 0.2.17",
+]
+
+[[package]]
+name = "rand_core"
+version = "0.9.5"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "76afc826de14238e6e8c374ddcc1fa19e374fd8dd986b0d2af0d02377261d83c"
+dependencies = [
+ "getrandom 0.3.4",
+]
+
+[[package]]
+name = "reqwest"
+version = "0.12.28"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "eddd3ca559203180a307f12d114c268abf583f59b03cb906fd0b3ff8646c1147"
+dependencies = [
+ "base64",
+ "bytes",
+ "futures-channel",
+ "futures-core",
+ "futures-util",
+ "http",
+ "http-body",
+ "http-body-util",
+ "hyper",
+ "hyper-rustls",
+ "hyper-util",
+ "js-sys",
+ "log",
+ "percent-encoding",
+ "pin-project-lite",
+ "quinn",
+ "rustls",
+ "rustls-pki-types",
+ "serde",
+ "serde_json",
+ "serde_urlencoded",
+ "sync_wrapper",
+ "tokio",
+ "tokio-rustls",
+ "tower",
+ "tower-http",
+ "tower-service",
+ "url",
+ "wasm-bindgen",
+ "wasm-bindgen-futures",
+ "web-sys",
+ "webpki-roots",
+]
+
+[[package]]
+name = "ring"
+version = "0.17.14"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "a4689e6c2294d81e88dc6261c768b63bc4fcdb852be6d1352498b114f61383b7"
+dependencies = [
+ "cc",
+ "cfg-if",
+ "getrandom 0.2.17",
+ "libc",
+ "untrusted",
+ "windows-sys 0.52.0",
+]
+
+[[package]]
+name = "rustc-hash"
+version = "2.1.2"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "94300abf3f1ae2e2b8ffb7b58043de3d399c73fa6f4b73826402a5c457614dbe"
+
+[[package]]
+name = "rustls"
+version = "0.23.41"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "6b92b125634d9b795e7beca796cc790df15a7fb38323bf3196fda83292d06b1f"
+dependencies = [
+ "once_cell",
+ "ring",
+ "rustls-pki-types",
+ "rustls-webpki",
+ "subtle",
+ "zeroize",
+]
+
+[[package]]
+name = "rustls-native-certs"
+version = "0.8.4"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "dab5152771c58876a2146916e53e35057e1a4dfa2b9df0f0305b07f611fdea4d"
+dependencies = [
+ "openssl-probe",
+ "rustls-pki-types",
+ "schannel",
+ "security-framework",
+]
+
+[[package]]
+name = "rustls-pki-types"
+version = "1.14.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "30a7197ae7eb376e574fe940d068c30fe0462554a3ddbe4eca7838e049c937a9"
+dependencies = [
+ "web-time",
+ "zeroize",
+]
+
+[[package]]
+name = "rustls-webpki"
+version = "0.103.13"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "61c429a8649f110dddef65e2a5ad240f747e85f7758a6bccc7e5777bd33f756e"
+dependencies = [
+ "ring",
+ "rustls-pki-types",
+ "untrusted",
+]
+
+[[package]]
+name = "rustversion"
+version = "1.0.22"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "b39cdef0fa800fc44525c84ccb54a029961a8215f9619753635a9c0d2538d46d"
+
+[[package]]
+name = "ryu"
+version = "1.0.23"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "9774ba4a74de5f7b1c1451ed6cd5285a32eddb5cccb8cc655a4e50009e06477f"
+
+[[package]]
+name = "schannel"
+version = "0.1.29"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "91c1b7e4904c873ef0710c1f407dde2e6287de2bebc1bbbf7d430bb7cbffd939"
+dependencies = [
+ "windows-sys 0.61.2",
+]
+
+[[package]]
+name = "security-framework"
+version = "3.7.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "b7f4bc775c73d9a02cde8bf7b2ec4c9d12743edf609006c7facc23998404cd1d"
+dependencies = [
+ "bitflags",
+ "core-foundation",
+ "core-foundation-sys",
+ "libc",
+ "security-framework-sys",
+]
+
+[[package]]
+name = "security-framework-sys"
+version = "2.17.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "6ce2691df843ecc5d231c0b14ece2acc3efb62c0a398c7e1d875f3983ce020e3"
+dependencies = [
+ "core-foundation-sys",
+ "libc",
+]
+
+[[package]]
+name = "serde"
+version = "1.0.228"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "9a8e94ea7f378bd32cbbd37198a4a91436180c5bb472411e48b5ec2e2124ae9e"
+dependencies = [
+ "serde_core",
+ "serde_derive",
+]
+
+[[package]]
+name = "serde_core"
+version = "1.0.228"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "41d385c7d4ca58e59fc732af25c3983b67ac852c1a25000afe1175de458b67ad"
+dependencies = [
+ "serde_derive",
+]
+
+[[package]]
+name = "serde_derive"
+version = "1.0.228"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "d540f220d3187173da220f885ab66608367b6574e925011a9353e4badda91d79"
+dependencies = [
+ "proc-macro2",
+ "quote",
+ "syn",
+]
+
+[[package]]
+name = "serde_json"
+version = "1.0.150"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "e8014e44b4736ed0538adeecded0fce2a272f22dc9578a7eb6b2d9993c74cfb9"
+dependencies = [
+ "itoa",
+ "memchr",
+ "serde",
+ "serde_core",
+ "zmij",
+]
+
+[[package]]
+name = "serde_path_to_error"
+version = "0.1.20"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "10a9ff822e371bb5403e391ecd83e182e0e77ba7f6fe0160b795797109d1b457"
+dependencies = [
+ "itoa",
+ "serde",
+ "serde_core",
+]
+
+[[package]]
+name = "serde_urlencoded"
+version = "0.7.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "d3491c14715ca2294c4d6a88f15e84739788c1d030eed8c110436aafdaa2f3fd"
+dependencies = [
+ "form_urlencoded",
+ "itoa",
+ "ryu",
+ "serde",
+]
+
+[[package]]
+name = "sha1"
+version = "0.10.6"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "e3bf829a2d51ab4a5ddf1352d8470c140cadc8301b2ae1789db023f01cedd6ba"
+dependencies = [
+ "cfg-if",
+ "cpufeatures",
+ "digest",
+]
+
+[[package]]
+name = "shlex"
+version = "2.0.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "f8fadd59c855ef2080decdef8ff161eb6661b86933c9d82e5ba29dc602a55aba"
+
+[[package]]
+name = "slab"
+version = "0.4.12"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "0c790de23124f9ab44544d7ac05d60440adc586479ce501c1d6d7da3cd8c9cf5"
+
+[[package]]
+name = "smallvec"
+version = "1.15.2"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "8ed6a63f02c8539c91a8685a86f4099661ba3da017932f6ebbea6de3f0fa7c90"
+
+[[package]]
+name = "socket2"
+version = "0.6.4"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "52d1cfed4120b4d927bf7c0f86d2087a4a7d6027c906d9f9d525a80573b9be51"
+dependencies = [
+ "libc",
+ "windows-sys 0.61.2",
+]
+
+[[package]]
+name = "stable_deref_trait"
+version = "1.2.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "6ce2be8dc25455e1f91df71bfa12ad37d7af1092ae736f3a6cd0e37bc7810596"
+
+[[package]]
+name = "subtle"
+version = "2.6.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "13c2bddecc57b384dee18652358fb23172facb8a2c51ccc10d74c157bdea3292"
+
+[[package]]
+name = "syn"
+version = "2.0.118"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "1b9ae57f904213ebb649ce6895b8a66c66f0203b9319718f69a5612a065b1422"
+dependencies = [
+ "proc-macro2",
+ "quote",
+ "unicode-ident",
+]
+
+[[package]]
+name = "sync_wrapper"
+version = "1.0.2"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "0bf256ce5efdfa370213c1dabab5935a12e49f2c58d15e9eac2870d3b4f27263"
+dependencies = [
+ "futures-core",
+]
+
+[[package]]
+name = "synstructure"
+version = "0.13.2"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "728a70f3dbaf5bab7f0c4b1ac8d7ae5ea60a4b5549c8a5914361c99147a709d2"
+dependencies = [
+ "proc-macro2",
+ "quote",
+ "syn",
+]
+
+[[package]]
+name = "target-lexicon"
+version = "0.12.16"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "61c41af27dd6d1e27b1b16b489db798443478cef1f06a660c96db617ba5de3b1"
+
+[[package]]
+name = "thiserror"
+version = "1.0.69"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "b6aaf5339b578ea85b50e080feb250a3e8ae8cfcdff9a461c9ec2904bc923f52"
+dependencies = [
+ "thiserror-impl 1.0.69",
+]
+
+[[package]]
+name = "thiserror"
+version = "2.0.18"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "4288b5bcbc7920c07a1149a35cf9590a2aa808e0bc1eafaade0b80947865fbc4"
+dependencies = [
+ "thiserror-impl 2.0.18",
+]
+
+[[package]]
+name = "thiserror-impl"
+version = "1.0.69"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "4fee6c4efc90059e10f81e6d42c60a18f76588c3d74cb83a0b242a2b6c7504c1"
+dependencies = [
+ "proc-macro2",
+ "quote",
+ "syn",
+]
+
+[[package]]
+name = "thiserror-impl"
+version = "2.0.18"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "ebc4ee7f67670e9b64d05fa4253e753e016c6c95ff35b89b7941d6b856dec1d5"
+dependencies = [
+ "proc-macro2",
+ "quote",
+ "syn",
+]
+
+[[package]]
+name = "tinystr"
+version = "0.8.3"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "c8323304221c2a851516f22236c5722a72eaa19749016521d6dff0824447d96d"
+dependencies = [
+ "displaydoc",
+ "zerovec",
+]
+
+[[package]]
+name = "tinyvec"
+version = "1.11.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "3e61e67053d25a4e82c844e8424039d9745781b3fc4f32b8d55ed50f5f667ef3"
+dependencies = [
+ "tinyvec_macros",
+]
+
+[[package]]
+name = "tinyvec_macros"
+version = "0.1.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20"
+
+[[package]]
+name = "tokio"
+version = "1.52.3"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "8fc7f01b389ac15039e4dc9531aa973a135d7a4135281b12d7c1bc79fd57fffe"
+dependencies = [
+ "bytes",
+ "libc",
+ "mio",
+ "pin-project-lite",
+ "socket2",
+ "tokio-macros",
+ "windows-sys 0.61.2",
+]
+
+[[package]]
+name = "tokio-macros"
+version = "2.7.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "385a6cb71ab9ab790c5fe8d67f1645e6c450a7ce006a33de03daa956cf70a496"
+dependencies = [
+ "proc-macro2",
+ "quote",
+ "syn",
+]
+
+[[package]]
+name = "tokio-rustls"
+version = "0.26.4"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "1729aa945f29d91ba541258c8df89027d5792d85a8841fb65e8bf0f4ede4ef61"
+dependencies = [
+ "rustls",
+ "tokio",
+]
+
+[[package]]
+name = "tokio-tungstenite"
+version = "0.24.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "edc5f74e248dc973e0dbb7b74c7e0d6fcc301c694ff50049504004ef4d0cdcd9"
+dependencies = [
+ "futures-util",
+ "log",
+ "rustls",
+ "rustls-native-certs",
+ "rustls-pki-types",
+ "tokio",
+ "tokio-rustls",
+ "tungstenite",
+]
+
+[[package]]
+name = "tower"
+version = "0.5.3"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "ebe5ef63511595f1344e2d5cfa636d973292adc0eec1f0ad45fae9f0851ab1d4"
+dependencies = [
+ "futures-core",
+ "futures-util",
+ "pin-project-lite",
+ "sync_wrapper",
+ "tokio",
+ "tower-layer",
+ "tower-service",
+ "tracing",
+]
+
+[[package]]
+name = "tower-http"
+version = "0.6.11"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "4cfcf7e2740e6fc6d4d688b4ef00650406bb94adf4731e43c096c3a19fe40840"
+dependencies = [
+ "bitflags",
+ "bytes",
+ "futures-util",
+ "http",
+ "http-body",
+ "pin-project-lite",
+ "tower",
+ "tower-layer",
+ "tower-service",
+ "url",
+]
+
+[[package]]
+name = "tower-layer"
+version = "0.3.3"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "121c2a6cda46980bb0fcd1647ffaf6cd3fc79a013de288782836f6df9c48780e"
+
+[[package]]
+name = "tower-service"
+version = "0.3.3"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "8df9b6e13f2d32c91b9bd719c00d1958837bc7dec474d94952798cc8e69eeec3"
+
+[[package]]
+name = "tracing"
+version = "0.1.44"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "63e71662fa4b2a2c3a26f570f037eb95bb1f85397f3cd8076caed2f026a6d100"
+dependencies = [
+ "log",
+ "pin-project-lite",
+ "tracing-core",
+]
+
+[[package]]
+name = "tracing-core"
+version = "0.1.36"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "db97caf9d906fbde555dd62fa95ddba9eecfd14cb388e4f491a66d74cd5fb79a"
+dependencies = [
+ "once_cell",
+]
+
+[[package]]
+name = "try-lock"
+version = "0.2.5"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "e421abadd41a4225275504ea4d6566923418b7f05506fbc9c0fe86ba7396114b"
+
+[[package]]
+name = "tungstenite"
+version = "0.24.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "18e5b8366ee7a95b16d32197d0b2604b43a0be89dc5fac9f8e96ccafbaedda8a"
+dependencies = [
+ "byteorder",
+ "bytes",
+ "data-encoding",
+ "http",
+ "httparse",
+ "log",
+ "rand 0.8.6",
+ "rustls",
+ "rustls-pki-types",
+ "sha1",
+ "thiserror 1.0.69",
+ "utf-8",
+]
+
+[[package]]
+name = "typenum"
+version = "1.20.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "b6f5e870be6c3b371b77fe0ee0bafb859fa4964b4404c27de1d380043c4dda20"
+
+[[package]]
+name = "unicode-ident"
+version = "1.0.24"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75"
+
+[[package]]
+name = "unindent"
+version = "0.2.4"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "7264e107f553ccae879d21fbea1d6724ac785e8c3bfc762137959b5802826ef3"
+
+[[package]]
+name = "untrusted"
+version = "0.9.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "8ecb6da28b8a351d773b68d5825ac39017e680750f980f3a1a85cd8dd28a47c1"
+
+[[package]]
+name = "url"
+version = "2.5.8"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "ff67a8a4397373c3ef660812acab3268222035010ab8680ec4215f38ba3d0eed"
+dependencies = [
+ "form_urlencoded",
+ "idna",
+ "percent-encoding",
+ "serde",
+]
+
+[[package]]
+name = "utf-8"
+version = "0.7.6"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "09cc8ee72d2a9becf2f2febe0205bbed8fc6615b7cb429ad062dc7b7ddd036a9"
+
+[[package]]
+name = "utf8_iter"
+version = "1.0.4"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "b6c140620e7ffbb22c2dee59cafe6084a59b5ffc27a8859a5f0d494b5d52b6be"
+
+[[package]]
+name = "version_check"
+version = "0.9.5"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "0b928f33d975fc6ad9f86c8f283853ad26bdd5b10b7f1542aa2fa15e2289105a"
+
+[[package]]
+name = "want"
+version = "0.3.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "bfa7760aed19e106de2c7c0b581b509f2f25d3dacaf737cb82ac61bc6d760b0e"
+dependencies = [
+ "try-lock",
+]
+
+[[package]]
+name = "wasi"
+version = "0.11.1+wasi-snapshot-preview1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "ccf3ec651a847eb01de73ccad15eb7d99f80485de043efb2f370cd654f4ea44b"
+
+[[package]]
+name = "wasip2"
+version = "1.0.4+wasi-0.2.12"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "b67efb37e106e55ce722a510d6b5f9c17f083e5fc79afc2badeb12cc313d9487"
+dependencies = [
+ "wit-bindgen",
+]
+
+[[package]]
+name = "wasm-bindgen"
+version = "0.2.125"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "8ddb3f79143bced6de84270411622a2699cee572fc0875aeaf1e7867cf9fca1a"
+dependencies = [
+ "cfg-if",
+ "once_cell",
+ "rustversion",
+ "wasm-bindgen-macro",
+ "wasm-bindgen-shared",
+]
+
+[[package]]
+name = "wasm-bindgen-futures"
+version = "0.4.75"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "503b14d284f2c8dac03b819967e155ea753f573586193b2b2c95990cb5d69280"
+dependencies = [
+ "js-sys",
+ "wasm-bindgen",
+]
+
+[[package]]
+name = "wasm-bindgen-macro"
+version = "0.2.125"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "4e21a184b13fb19e157296e2c46056aec9092264fab83e4ba59e68c61b323c3d"
+dependencies = [
+ "quote",
+ "wasm-bindgen-macro-support",
+]
+
+[[package]]
+name = "wasm-bindgen-macro-support"
+version = "0.2.125"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "fecefd9c35bd935a20fc3fc344b5f29138961e4f47fb03297d88f2587afb5ebd"
+dependencies = [
+ "bumpalo",
+ "proc-macro2",
+ "quote",
+ "syn",
+ "wasm-bindgen-shared",
+]
+
+[[package]]
+name = "wasm-bindgen-shared"
+version = "0.2.125"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "23939e44bb9a5d7576fa2b563dc2e136628f1224e88a8deed09e04858b77871f"
+dependencies = [
+ "unicode-ident",
+]
+
+[[package]]
+name = "web-sys"
+version = "0.3.102"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "a6430a72df5eb332242960fe84b3002a241163998241eb596d4f739b9757061d"
+dependencies = [
+ "js-sys",
+ "wasm-bindgen",
+]
+
+[[package]]
+name = "web-time"
+version = "1.1.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "5a6580f308b1fad9207618087a65c04e7a10bc77e02c8e84e9b00dd4b12fa0bb"
+dependencies = [
+ "js-sys",
+ "wasm-bindgen",
+]
+
+[[package]]
+name = "webpki-roots"
+version = "1.0.8"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "bf85cb06032201fa7c6f829d7db5a7e5aa45bcc0655327713065f6f0576731bf"
+dependencies = [
+ "rustls-pki-types",
+]
+
+[[package]]
+name = "windows-link"
+version = "0.2.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5"
+
+[[package]]
+name = "windows-sys"
+version = "0.52.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "282be5f36a8ce781fad8c8ae18fa3f9beff57ec1b52cb3de0789201425d9a33d"
+dependencies = [
+ "windows-targets 0.52.6",
+]
+
+[[package]]
+name = "windows-sys"
+version = "0.60.2"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "f2f500e4d28234f72040990ec9d39e3a6b950f9f22d3dba18416c35882612bcb"
+dependencies = [
+ "windows-targets 0.53.5",
+]
+
+[[package]]
+name = "windows-sys"
+version = "0.61.2"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "ae137229bcbd6cdf0f7b80a31df61766145077ddf49416a728b02cb3921ff3fc"
+dependencies = [
+ "windows-link",
+]
+
+[[package]]
+name = "windows-targets"
+version = "0.52.6"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "9b724f72796e036ab90c1021d4780d4d3d648aca59e491e6b98e725b84e99973"
+dependencies = [
+ "windows_aarch64_gnullvm 0.52.6",
+ "windows_aarch64_msvc 0.52.6",
+ "windows_i686_gnu 0.52.6",
+ "windows_i686_gnullvm 0.52.6",
+ "windows_i686_msvc 0.52.6",
+ "windows_x86_64_gnu 0.52.6",
+ "windows_x86_64_gnullvm 0.52.6",
+ "windows_x86_64_msvc 0.52.6",
+]
+
+[[package]]
+name = "windows-targets"
+version = "0.53.5"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "4945f9f551b88e0d65f3db0bc25c33b8acea4d9e41163edf90dcd0b19f9069f3"
+dependencies = [
+ "windows-link",
+ "windows_aarch64_gnullvm 0.53.1",
+ "windows_aarch64_msvc 0.53.1",
+ "windows_i686_gnu 0.53.1",
+ "windows_i686_gnullvm 0.53.1",
+ "windows_i686_msvc 0.53.1",
+ "windows_x86_64_gnu 0.53.1",
+ "windows_x86_64_gnullvm 0.53.1",
+ "windows_x86_64_msvc 0.53.1",
+]
+
+[[package]]
+name = "windows_aarch64_gnullvm"
+version = "0.52.6"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "32a4622180e7a0ec044bb555404c800bc9fd9ec262ec147edd5989ccd0c02cd3"
+
+[[package]]
+name = "windows_aarch64_gnullvm"
+version = "0.53.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "a9d8416fa8b42f5c947f8482c43e7d89e73a173cead56d044f6a56104a6d1b53"
+
+[[package]]
+name = "windows_aarch64_msvc"
+version = "0.52.6"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "09ec2a7bb152e2252b53fa7803150007879548bc709c039df7627cabbd05d469"
+
+[[package]]
+name = "windows_aarch64_msvc"
+version = "0.53.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "b9d782e804c2f632e395708e99a94275910eb9100b2114651e04744e9b125006"
+
+[[package]]
+name = "windows_i686_gnu"
+version = "0.52.6"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "8e9b5ad5ab802e97eb8e295ac6720e509ee4c243f69d781394014ebfe8bbfa0b"
+
+[[package]]
+name = "windows_i686_gnu"
+version = "0.53.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "960e6da069d81e09becb0ca57a65220ddff016ff2d6af6a223cf372a506593a3"
+
+[[package]]
+name = "windows_i686_gnullvm"
+version = "0.52.6"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "0eee52d38c090b3caa76c563b86c3a4bd71ef1a819287c19d586d7334ae8ed66"
+
+[[package]]
+name = "windows_i686_gnullvm"
+version = "0.53.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "fa7359d10048f68ab8b09fa71c3daccfb0e9b559aed648a8f95469c27057180c"
+
+[[package]]
+name = "windows_i686_msvc"
+version = "0.52.6"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "240948bc05c5e7c6dabba28bf89d89ffce3e303022809e73deaefe4f6ec56c66"
+
+[[package]]
+name = "windows_i686_msvc"
+version = "0.53.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "1e7ac75179f18232fe9c285163565a57ef8d3c89254a30685b57d83a38d326c2"
+
+[[package]]
+name = "windows_x86_64_gnu"
+version = "0.52.6"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "147a5c80aabfbf0c7d901cb5895d1de30ef2907eb21fbbab29ca94c5b08b1a78"
+
+[[package]]
+name = "windows_x86_64_gnu"
+version = "0.53.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "9c3842cdd74a865a8066ab39c8a7a473c0778a3f29370b5fd6b4b9aa7df4a499"
+
+[[package]]
+name = "windows_x86_64_gnullvm"
+version = "0.52.6"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "24d5b23dc417412679681396f2b49f3de8c1473deb516bd34410872eff51ed0d"
+
+[[package]]
+name = "windows_x86_64_gnullvm"
+version = "0.53.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "0ffa179e2d07eee8ad8f57493436566c7cc30ac536a3379fdf008f47f6bb7ae1"
+
+[[package]]
+name = "windows_x86_64_msvc"
+version = "0.52.6"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec"
+
+[[package]]
+name = "windows_x86_64_msvc"
+version = "0.53.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "d6bbff5f0aada427a1e5a6da5f1f98158182f26556f345ac9e04d36d0ebed650"
+
+[[package]]
+name = "wit-bindgen"
+version = "0.57.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "1ebf944e87a7c253233ad6766e082e3cd714b5d03812acc24c318f549614536e"
+
+[[package]]
+name = "writeable"
+version = "0.6.3"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "1ffae5123b2d3fc086436f8834ae3ab053a283cfac8fe0a0b8eaae044768a4c4"
+
+[[package]]
+name = "yoke"
+version = "0.8.3"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "709fe23a0424b6a435d82152b1bd3fdfb0833487d5fa90d05d42762a9891fef5"
+dependencies = [
+ "stable_deref_trait",
+ "yoke-derive",
+ "zerofrom",
+]
+
+[[package]]
+name = "yoke-derive"
+version = "0.8.2"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "de844c262c8848816172cef550288e7dc6c7b7814b4ee56b3e1553f275f1858e"
+dependencies = [
+ "proc-macro2",
+ "quote",
+ "syn",
+ "synstructure",
+]
+
+[[package]]
+name = "zerocopy"
+version = "0.8.52"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "ce1022995ff5ff5d841ad7d994facc23098cd40152f2c1d11cd607c6f530653f"
+dependencies = [
+ "zerocopy-derive",
+]
+
+[[package]]
+name = "zerocopy-derive"
+version = "0.8.52"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "1ae7f38b72ec2a254e2b87ef277cf2cd4fb97cbebf944faa6f33354da0867930"
+dependencies = [
+ "proc-macro2",
+ "quote",
+ "syn",
+]
+
+[[package]]
+name = "zerofrom"
+version = "0.1.8"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "0ec05a11813ea801ff6d75110ad09cd0824ddba17dfe17128ea0d5f68e6c5272"
+dependencies = [
+ "zerofrom-derive",
+]
+
+[[package]]
+name = "zerofrom-derive"
+version = "0.1.7"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "11532158c46691caf0f2593ea8358fed6bbf68a0315e80aae9bd41fbade684a1"
+dependencies = [
+ "proc-macro2",
+ "quote",
+ "syn",
+ "synstructure",
+]
+
+[[package]]
+name = "zeroize"
+version = "1.9.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "e13c156562582aa81c60cb29407084cdb54c4164760106ab78e6c5b0858cf64e"
+
+[[package]]
+name = "zerotrie"
+version = "0.2.4"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "0f9152d31db0792fa83f70fb2f83148effb5c1f5b8c7686c3459e361d9bc20bf"
+dependencies = [
+ "displaydoc",
+ "yoke",
+ "zerofrom",
+]
+
+[[package]]
+name = "zerovec"
+version = "0.11.6"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "90f911cbc359ab6af17377d242225f4d75119aec87ea711a880987b18cd7b239"
+dependencies = [
+ "yoke",
+ "zerofrom",
+ "zerovec-derive",
+]
+
+[[package]]
+name = "zerovec-derive"
+version = "0.11.3"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "625dc425cab0dca6dc3c3319506e6593dcb08a9f387ea3b284dbd52a92c40555"
+dependencies = [
+ "proc-macro2",
+ "quote",
+ "syn",
+]
+
+[[package]]
+name = "zmij"
+version = "1.0.21"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "b8848ee67ecc8aedbaf3e4122217aff892639231befc6a1b58d29fff4c2cabaa"
diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml
new file mode 100644
index 00000000000..06289e5a46f
--- /dev/null
+++ b/litellm-rust/Cargo.toml
@@ -0,0 +1,28 @@
+[workspace]
+members = [
+ "crates/core",
+ "crates/providers",
+ "crates/python-bridge",
+ "crates/ai-gateway",
+]
+resolver = "2"
+
+[workspace.package]
+edition = "2021"
+license = "MIT"
+repository = "https://github.com/BerriAI/litellm"
+
+[workspace.dependencies]
+litellm-core = { path = "crates/core" }
+litellm-providers = { path = "crates/providers" }
+axum = "0.7"
+pyo3 = "0.23.5"
+rand = "0.8"
+reqwest = { version = "0.12", default-features = false, features = ["blocking", "json", "rustls-tls"] }
+serde = { version = "1.0", features = ["derive"] }
+serde_json = "1.0"
+subtle = "2"
+thiserror = "2.0"
+tokio = { version = "1", features = ["rt-multi-thread", "macros", "time"] }
+tokio-tungstenite = { version = "0.24", default-features = false, features = ["connect", "rustls-tls-native-roots"] }
+futures-util = { version = "0.3", default-features = false, features = ["sink", "std"] }
diff --git a/litellm-rust/README.md b/litellm-rust/README.md
new file mode 100644
index 00000000000..15ad1855420
--- /dev/null
+++ b/litellm-rust/README.md
@@ -0,0 +1,34 @@
+# LiteLLM Rust
+
+This workspace contains the staged Rust implementation for LiteLLM.
+
+Rust starts as a pure transform core used by the existing Python host. Python
+continues to own auth, configuration, network I/O, retries, routing, logging,
+callbacks, spend tracking, and customer plugins until each Rust path has parity
+coverage and production evidence.
+
+## Layout
+
+```text
+crates/
+ core/ Route contracts, shared pure types, errors, and templates.
+ src/ocr/
+ providers/ Provider-specific pure transforms.
+ src/mistral/ocr/transformation.rs
+ python-bridge/ PyO3 bridge for Python LiteLLM.
+```
+
+The folder shape should follow the Python provider tree:
+`providers/src///transformation.rs`. The bridge should expose
+one function per top-level route, starting with `ocr(payload)`.
+
+## Checks
+
+Run these before pushing Rust changes. GitHub Actions runs the same checks for
+changes under `litellm-rust/`.
+
+```bash
+cargo fmt --check
+cargo clippy --workspace --all-targets -- -D warnings
+cargo test --workspace
+```
diff --git a/litellm-rust/crates/ai-gateway/AGENTS.md b/litellm-rust/crates/ai-gateway/AGENTS.md
new file mode 100644
index 00000000000..d9e6e1adde5
--- /dev/null
+++ b/litellm-rust/crates/ai-gateway/AGENTS.md
@@ -0,0 +1,50 @@
+# ai-gateway — folder architecture
+
+The Axum server that fronts the Rust gateway. It owns transport + config + auth
+only; deployment selection lives in `core::router`, transforms in `core`/`providers`.
+
+```
+src/
+ main.rs # entrypoint: build AppState (router + master key), bind, serve
+ state.rs # AppState — shared Arc + master_key
+ gil.rs # GIL-activity tracker (records Python acquisitions)
+ auth/ # authentication as an axum extractor — added to handler args
+ mod.rs # RequireMasterKey: FromRequestParts, single master key (LITELLM_MASTER_KEY)
+ routes/ # one module per route, all matching the same template
+ AGENTS.md # ← the route template (read this before adding a route)
+ mod.rs # app(): merges every module's router()
+ health.rs # simple route (one file): router() + liveness/readiness
+ gil.rs # simple route (one file): router() + GET /health/gil
+ realtime/ # route with logic → axum surface + a no-axum service:
+ mod.rs # router() + handler + WS<->events adapter (the axum surface)
+ service.rs # business logic (select deployment, call provider) — no axum, testable
+ python/ # Python interop (feature: python-config) — load-time only
+ mod.rs, config.rs, AGENTS.md
+```
+
+## Rules
+
+- **Routes follow one template.** Each route module exposes
+ `pub fn router() -> Router`; `routes/mod.rs` only merges them. Simple
+ routes are one file; non-trivial routes are a folder (`handler`/`service`/
+ `transport`). See `routes/AGENTS.md`.
+- **Auth is an extractor.** Add `crate::auth::RequireMasterKey` to a handler's
+ args; it runs during extraction. Never re-implement the check per route.
+- **Handlers are thin.** A handler validates and delegates to its `service`. No
+ business logic, no provider calls, no transforms in handlers.
+- **State is shared and cheap to clone.** Long-lived handles live behind `Arc` in
+ `state.rs`; read env/config only in `main.rs` when building state.
+
+## Auth (interim)
+
+A single **master key** (`LITELLM_MASTER_KEY`), enforced by the
+`auth::RequireMasterKey` extractor: any caller presenting it as
+`Authorization: Bearer ` may invoke the gateway. Fails closed (500) when
+unset; constant-time compare. The server binds `127.0.0.1` by default (`HOST` to
+override). Full per-key auth + budgets/rate-limits are delegated to the Python
+proxy in a later phase. Health routes don't add the extractor (unauthenticated).
+
+## Python interop
+
+Anything that calls into Python lives in `python/` and is **load-time only** — see
+`python/AGENTS.md`. The realtime data path never takes the GIL.
diff --git a/litellm-rust/crates/ai-gateway/Cargo.toml b/litellm-rust/crates/ai-gateway/Cargo.toml
new file mode 100644
index 00000000000..79bdc4bdb26
--- /dev/null
+++ b/litellm-rust/crates/ai-gateway/Cargo.toml
@@ -0,0 +1,26 @@
+[package]
+name = "litellm-ai-gateway"
+version = "0.1.0"
+edition.workspace = true
+license.workspace = true
+repository.workspace = true
+
+[[bin]]
+name = "litellm-ai-gateway"
+path = "src/main.rs"
+
+[dependencies]
+litellm-core.workspace = true
+litellm-providers.workspace = true
+axum = { workspace = true, features = ["ws"] }
+futures-util.workspace = true
+tokio = { workspace = true, features = ["rt-multi-thread", "macros", "net", "time"] }
+serde.workspace = true
+serde_json.workspace = true
+subtle.workspace = true
+pyo3 = { workspace = true, features = ["auto-initialize"], optional = true }
+
+[features]
+# Build the gateway's config from the proxy YAML via an embedded Python
+# interpreter (links libpython; requires `litellm` importable at runtime).
+python-config = ["dep:pyo3"]
diff --git a/litellm-rust/crates/ai-gateway/Dockerfile b/litellm-rust/crates/ai-gateway/Dockerfile
new file mode 100644
index 00000000000..adf6fca0741
--- /dev/null
+++ b/litellm-rust/crates/ai-gateway/Dockerfile
@@ -0,0 +1,86 @@
+# Multi-stage build for the LiteLLM Rust AI Gateway (realtime WebSocket proxy).
+#
+# Build context is the **repo root** so we can install `litellm` from this repo's
+# source (the gateway loads its model_list via litellm.proxy.read_model_list,
+# which is not in any PyPI release yet) AND build the rust workspace under
+# litellm-rust/.
+#
+# docker build -f litellm-rust/crates/ai-gateway/Dockerfile -t litellm-ai-gateway .
+#
+# No secrets live in this file. Runtime config (LITELLM_MASTER_KEY,
+# OPENAI_API_KEY referenced by config.yaml, etc.) is injected as environment
+# variables at deploy time.
+
+# ---- Chef -------------------------------------------------------------------
+# cargo-chef caches the dependency build so only the gateway crate recompiles on
+# a source-only change. python3-dev is present in every rust stage because the
+# `python-config` feature links libpython via pyo3 (even in the cook step).
+FROM rust:1.90-slim-bookworm AS chef
+ENV PYO3_PYTHON=python3.11
+RUN apt-get update \
+ && apt-get install -y --no-install-recommends \
+ python3 python3-dev pkg-config libssl-dev clang \
+ && rm -rf /var/lib/apt/lists/* \
+ && cargo install cargo-chef --locked --version 0.1.77
+WORKDIR /build/litellm-rust
+
+# ---- Planner ----------------------------------------------------------------
+# Produce the dependency recipe from the rust workspace manifests + Cargo.lock.
+FROM chef AS planner
+COPY litellm-rust/ .
+RUN cargo chef prepare --recipe-path recipe.json
+
+# ---- Builder ----------------------------------------------------------------
+FROM chef AS builder
+# Cook (compile) just the dependencies first — this layer is cached and reused
+# whenever only gateway source changes.
+COPY --from=planner /build/litellm-rust/recipe.json recipe.json
+RUN cargo chef cook --locked --release \
+ -p litellm-ai-gateway --features python-config \
+ --recipe-path recipe.json
+# Now copy the real sources and build the gateway binary. Deps are already cooked
+# above, so this step only recompiles the gateway crate.
+COPY litellm-rust/ .
+RUN cargo build --locked --release -p litellm-ai-gateway --features python-config
+
+# ---- Runtime ----------------------------------------------------------------
+# python:3.11-slim-bookworm ships libpython3.11, matching the builder's PyO3
+# 3.11 ABI so the embedded interpreter links and imports cleanly.
+FROM python:3.11-slim-bookworm AS runtime
+
+# CA certificates for outbound TLS to the OpenAI realtime endpoint.
+RUN apt-get update \
+ && apt-get install -y --no-install-recommends ca-certificates \
+ && rm -rf /var/lib/apt/lists/*
+
+WORKDIR /app
+
+# Install litellm (with proxy extras) FROM THIS REPO'S SOURCE so
+# `import litellm.proxy.read_model_list` works — it is not on PyPI yet. Copy the
+# package + packaging metadata, then pip install the proxy extra.
+COPY pyproject.toml README.md LICENSE ./
+COPY litellm/ ./litellm/
+RUN pip install --no-cache-dir ".[proxy]"
+
+# The compiled gateway binary (pure-Rust realtime hot path; Python is load-time
+# only).
+COPY --from=builder /build/litellm-rust/target/release/litellm-ai-gateway /usr/local/bin/litellm-ai-gateway
+
+# Default config.yaml. A real deploy can override this (e.g. mount a Render
+# secret file at the same path) — never bake secrets into the image.
+COPY litellm-rust/crates/ai-gateway/config.yaml /app/config.yaml
+
+# Bind to all interfaces (Render routes to 0.0.0.0:$PORT) and load the model_list
+# from config.yaml via the embedded python config reader.
+ENV HOST=0.0.0.0 \
+ LITELLM_CONFIG_PATH=/app/config.yaml
+
+# Drop to a non-root user. The realtime hot path needs no root privileges, so
+# running unprivileged limits blast radius if the process is ever compromised.
+# The binary in /usr/local/bin is world-executable (COPY default mode 755); we
+# only need /app (and the config.yaml it reads) owned by the unprivileged user.
+RUN useradd --system --no-create-home --uid 10001 appuser \
+ && chown -R appuser:appuser /app
+USER appuser
+
+ENTRYPOINT ["/usr/local/bin/litellm-ai-gateway"]
diff --git a/litellm-rust/crates/ai-gateway/Dockerfile.dockerignore b/litellm-rust/crates/ai-gateway/Dockerfile.dockerignore
new file mode 100644
index 00000000000..030ee6a37c5
--- /dev/null
+++ b/litellm-rust/crates/ai-gateway/Dockerfile.dockerignore
@@ -0,0 +1,45 @@
+# Dockerfile-specific ignore-file for the Rust AI Gateway build.
+#
+# The build context is the repo root (so the image can pip install litellm from
+# source AND build the rust workspace). BuildKit honors `.dockerignore`
+# next to the Dockerfile and it takes precedence over the repo-root `.dockerignore`,
+# so this file shrinks the (large) repo-root context for THIS build only without
+# touching the root `.dockerignore` used by the main litellm images.
+#
+# Strategy: ignore everything, then re-include only what the build needs:
+# - litellm/ (pip install . needs the full package + proxy reader)
+# - litellm-rust/ (the rust workspace; Cargo.lock + crate sources)
+# - pyproject.toml / README.md / LICENSE (packaging metadata for pip install)
+*
+
+# --- re-include the build inputs ---
+!litellm/
+!litellm-rust/
+!pyproject.toml
+!README.md
+!LICENSE
+
+# --- prune heavy / irrelevant subpaths back out of the re-included trees ---
+# Rust build artifacts (huge; regenerated in the builder).
+**/target/
+# Python caches and compiled bytecode.
+**/__pycache__/
+**/*.pyc
+**/*.pyo
+**/.pytest_cache/
+**/.ruff_cache/
+**/.mypy_cache/
+# Node / UI build output bundled under the python package (not needed to import
+# litellm.proxy.read_model_list).
+**/node_modules/
+litellm/proxy/_experimental/out/
+# Tests, logs, and local scratch.
+**/tests/
+**/test/
+*.log
+log.txt
+*.tgz
+# VCS / editor / CI metadata that may live under re-included trees.
+**/.git/
+.git/
+**/.DS_Store
diff --git a/litellm-rust/crates/ai-gateway/README.md b/litellm-rust/crates/ai-gateway/README.md
new file mode 100644
index 00000000000..3662ce2584a
--- /dev/null
+++ b/litellm-rust/crates/ai-gateway/README.md
@@ -0,0 +1,172 @@
+# LiteLLM Rust AI Gateway
+
+A minimal Axum service that fronts OpenAI's realtime API. Clients open a
+WebSocket to `GET /v1/realtime`; the gateway authenticates, selects a deployment,
+dials OpenAI upstream, and splices the two sockets frame-by-frame.
+
+- **Client endpoint:** `wss:///v1/realtime?model=` (WebSocket)
+- **Auth:** `Authorization: Bearer $LITELLM_MASTER_KEY` (fails closed if unset)
+- **Health:** `GET /health/readiness`, `GET /health/liveness`, `GET /health/gil`
+
+> **Realtime serving is pure Rust.** Python is used at **load time only** — to
+> read the config once at boot. The realtime hot path never touches Python.
+
+## Configuration (config.yaml)
+
+The gateway loads its `model_list` from a **config.yaml**, the same as the
+LiteLLM proxy. Point `LITELLM_CONFIG_PATH` at the file:
+
+```yaml
+# config.yaml
+model_list:
+ - model_name: gpt-realtime
+ litellm_params:
+ model: openai/gpt-realtime
+ api_key: os.environ/OPENAI_API_KEY
+```
+
+```bash
+LITELLM_CONFIG_PATH=./config.yaml ./litellm-ai-gateway
+```
+
+At boot the gateway calls into `litellm.proxy.read_model_list`, which reuses the
+**real proxy config reader** (`ProxyConfig.get_config`). That means everything
+the proxy supports in config.yaml works here too:
+
+- `include:` to merge in other config files,
+- `os.environ/VAR` secret references (resolved via the secret manager, never
+ inlined),
+- DB-stored models (when a database is configured).
+
+Secrets stay out of the config — reference them with `os.environ/...` and set
+the env var at deploy time. The shipped Docker image is built with the
+`python-config` feature and **bundles litellm**, so config loading works out of
+the box; the default baked config lives at `/app/config.yaml` and can be
+overridden at deploy time (e.g. a Render secret file mounted at the same path).
+
+### Environment variables
+
+| Var | Required | Default | Purpose |
+|---|---|---|---|
+| `LITELLM_CONFIG_PATH` | yes (config mode) | — | Path to the config.yaml the gateway loads its `model_list` from. The Docker image defaults this to `/app/config.yaml`. |
+| `LITELLM_MASTER_KEY` | yes | — | Bearer token clients must send. Unset ⇒ all `/v1/realtime` requests are rejected (fail closed). |
+| `OPENAI_API_KEY` | yes | — | Upstream OpenAI key. Referenced by config.yaml as `os.environ/OPENAI_API_KEY` for the gateway→OpenAI dial. |
+| `HOST` | no | `127.0.0.1` | **Set to `0.0.0.0` in any container/deploy** or external traffic is refused. |
+| `PORT` | no | `4001` | Listen port. Render and most PaaS inject this automatically. |
+
+> Secrets (`LITELLM_MASTER_KEY`, `OPENAI_API_KEY`) are never baked into the image
+> or `render.yaml` — inject them at deploy time only.
+
+### Lean env stand-in (fallback)
+
+If the binary is built **without** `python-config` (default features), or
+`LITELLM_CONFIG_PATH` is unset, the gateway falls back to a single-deployment
+stand-in built from the environment:
+
+| Var | Default | Purpose |
+|---|---|---|
+| `OPENAI_REALTIME_MODEL` | `gpt-realtime` | The single deployment's model name (also the `?model=` clients pass). |
+
+This mode links no libpython and needs no config file, but it only supports one
+hard-coded OpenAI deployment. **config.yaml is the recommended path** — use the
+stand-in only for the leanest possible build.
+
+## Build & run with Docker
+
+The image is built `--features python-config` and installs litellm **from this
+repo's source** (the config reader is newer than any PyPI release), so the build
+**context is the repo root**:
+
+```bash
+# from the repo root
+docker build -f litellm-rust/crates/ai-gateway/Dockerfile -t litellm-ai-gateway .
+
+docker run --rm -p 4001:4001 \
+ -e HOST=0.0.0.0 -e PORT=4001 \
+ -e LITELLM_MASTER_KEY=sk-local \
+ -e OPENAI_API_KEY=$OPENAI_API_KEY \
+ litellm-ai-gateway # LITELLM_CONFIG_PATH defaults to /app/config.yaml
+
+# smoke test
+curl -s -o /dev/null -w '%{http_code}\n' localhost:4001/health/readiness # -> 200
+curl -s -o /dev/null -w '%{http_code}\n' localhost:4001/v1/realtime # -> 401 (auth fails closed)
+```
+
+On boot you should see `loaded model_list from /app/config.yaml via python
+config reader` — that confirms the config path (not the env stand-in fallback).
+To use your own config, mount it over the default:
+
+```bash
+docker run --rm -p 4001:4001 \
+ -e HOST=0.0.0.0 -e LITELLM_MASTER_KEY=sk-local -e OPENAI_API_KEY=$OPENAI_API_KEY \
+ -v $(pwd)/my-config.yaml:/app/config.yaml:ro \
+ litellm-ai-gateway
+```
+
+### Cargo-only (no Docker)
+
+```bash
+# config.yaml mode — needs litellm importable in the active python env
+LITELLM_CONFIG_PATH=./crates/ai-gateway/config.yaml \
+ cargo run --release -p litellm-ai-gateway --features python-config
+
+# env stand-in mode — no python, no config
+cargo run --release -p litellm-ai-gateway
+```
+
+## Deploy on Render
+
+The service is a Docker **web service**; Render terminates TLS and supports
+WebSockets, so the public endpoint is `wss://.onrender.com/v1/realtime`.
+
+### Option A — Blueprint (`render.yaml`)
+
+`crates/ai-gateway/render.yaml` describes the service (Docker runtime,
+`healthCheckPath: /health/readiness`, repo-root `dockerContext: .`,
+`dockerfilePath: ./litellm-rust/crates/ai-gateway/Dockerfile`,
+`LITELLM_CONFIG_PATH: /app/config.yaml`). `LITELLM_MASTER_KEY` and
+`OPENAI_API_KEY` are `sync: false` — set them in the dashboard after the first
+deploy. To use a non-default model_list, mount a **Render Secret File** at
+`/app/config.yaml`. Point a Render Blueprint at this repo/branch and apply.
+
+### Option B — Render API
+
+```bash
+# create a Docker web service from this repo+branch, then set env vars:
+curl -X POST https://api.render.com/v1/services \
+ -H "Authorization: Bearer $RENDER_API_KEY" -H "Content-Type: application/json" \
+ -d '{
+ "type": "web_service", "name": "litellm-rust-ai-gateway",
+ "ownerId": "", "repo": "https://github.com/BerriAI/litellm",
+ "branch": "",
+ "serviceDetails": {
+ "env": "docker",
+ "envSpecificDetails": {
+ "dockerfilePath": "./litellm-rust/crates/ai-gateway/Dockerfile",
+ "dockerContext": "."
+ },
+ "healthCheckPath": "/health/readiness"
+ }
+ }'
+# then set env vars LITELLM_MASTER_KEY, OPENAI_API_KEY, HOST=0.0.0.0,
+# LITELLM_CONFIG_PATH=/app/config.yaml
+```
+
+Health check path **must** be `/health/readiness`. `autoDeploy` is off by default
+in the blueprint — trigger deploys manually (or flip it on) to pick up new commits.
+
+## Scaling
+
+Concurrency is what matters, not total connections: each in-flight session holds
+one client socket + one upstream socket. To scale, raise the instance count /
+enable autoscaling on the Render service (e.g. baseline 10, max 100). Each
+instance needs file descriptors for `2 × peak_concurrent_sessions` — raise
+`ulimit -n` if you push very high concurrency.
+
+## Latency note
+
+The gateway adds the cost of one extra hop: client→gateway, then a fresh
+gateway→OpenAI realtime handshake (TLS + WS upgrade + `session.created`). In
+benchmarks this is ~100–150 ms of added session-establishment time; first-audio
+and steady-state streaming add no measurable overhead. To minimize it, deploy the
+gateway in the Render region with the lowest RTT to OpenAI's realtime endpoint.
diff --git a/litellm-rust/crates/ai-gateway/benchmarks/realtime/README.md b/litellm-rust/crates/ai-gateway/benchmarks/realtime/README.md
new file mode 100644
index 00000000000..84e926af243
--- /dev/null
+++ b/litellm-rust/crates/ai-gateway/benchmarks/realtime/README.md
@@ -0,0 +1,55 @@
+# Realtime gateway benchmark — pool on/off
+
+Measures what the gateway adds over talking to OpenAI's realtime WebSocket
+directly, and what the pre-warmed connection pool removes. See
+`../../src/routes/realtime/README.md` for how the pool works.
+
+## Results
+
+5000 calls / 500 concurrency, gateway at 10 instances, pool ON
+(`REALTIME_POOL_SIZE=64`), upstream OpenAI `gpt-realtime`. Each leg run twice.
+Times in **ms**. Phases per connection: **dial** = TCP+TLS+WS upgrade,
+**session** = upgrade → `session.created` (the phase the pool removes),
+**1st-audio** = `response.create` → first audio delta (OpenAI inference),
+**total** = full wall-clock.
+
+| metric | Direct OpenAI | Gateway (pool ON) | Overhead (ms) | vs OpenAI |
+| ------------------ | ------------- | ----------------- | ------------- | ---------- |
+| success rate (%) | 99.8 | 99.8 | — | — |
+| dial p50 (ms) | 276 | 158 | −118 | **faster** |
+| session p50 (ms) | 7 | 0 | −7 | **faster** |
+| 1st-audio p50 (ms) | 440 | 664 | +224 | slower¹ |
+| total p50 (ms) | 816 | 1010 | +194 | slower¹ |
+| total p95 (ms) | 2152 | 1970 | −182 | **faster** |
+| total p99 (ms) | 2692 | 2610 | −82 | **faster** |
+
+The gateway is **faster than direct on 4 of 6 metrics**. The warm pool makes the
+**session phase sub-millisecond** at the median — ~76% of connects hit the pool,
+~70% had session < 1 ms. ¹ The two "slower" rows are not gateway overhead:
+`1st-audio` is OpenAI's own inference time (the gateway only relays it), which ran
+slower during the gateway legs and drags `total p50` with it.
+
+**Pool OFF** (control, `REALTIME_POOL_SIZE=0`): session p50 was **367 ms** — the
+fresh-dial overhead the pool removes.
+
+## Reproduce
+
+The load generator lives in a separate repo:
+**https://github.com/ishaan-berri/litellm-realtime-bench**
+
+```bash
+git clone https://github.com/ishaan-berri/litellm-realtime-bench
+cd litellm-realtime-bench && go build -o wsbench .
+
+# Direct to OpenAI (baseline)
+./wsbench -host api.openai.com -key "$OPENAI_API_KEY" -m gpt-realtime -n 5000 -c 500 -t 60
+
+# Through the gateway — run once with pool ON, once with REALTIME_POOL_SIZE=0
+./wsbench -host -key "$LITELLM_MASTER_KEY" -m gpt-realtime -n 5000 -c 500 -t 60
+```
+
+Run the gateway with the env stand-in (`OPENAI_REALTIME_MODEL=gpt-realtime`,
+`OPENAI_API_KEY`, `LITELLM_MASTER_KEY`, `REALTIME_POOL_SIZE`, `HOST=0.0.0.0`). At
+500 concurrency over N instances, size the pool to `≈ 500 / N` per instance (64 was
+used here for 10 instances). The bench repo's README covers running 500-concurrency
+legs from a hosted multi-vCPU runner. **Never commit keys — pass them via `-key`.**
diff --git a/litellm-rust/crates/ai-gateway/config.yaml b/litellm-rust/crates/ai-gateway/config.yaml
new file mode 100644
index 00000000000..ac598c220dd
--- /dev/null
+++ b/litellm-rust/crates/ai-gateway/config.yaml
@@ -0,0 +1,13 @@
+# Sample realtime config for the LiteLLM Rust AI Gateway.
+#
+# The gateway loads this model_list at boot via the embedded python config
+# reader (litellm.proxy.read_model_list), which reuses the proxy's own reader —
+# so include:, os.environ/ secrets, and DB-stored models all work here too.
+#
+# Secrets are referenced (never inlined) via os.environ/. A real deploy can
+# override this file (e.g. mount a Render secret file at LITELLM_CONFIG_PATH).
+model_list:
+ - model_name: gpt-realtime
+ litellm_params:
+ model: openai/gpt-realtime
+ api_key: os.environ/OPENAI_API_KEY
diff --git a/litellm-rust/crates/ai-gateway/render.yaml b/litellm-rust/crates/ai-gateway/render.yaml
new file mode 100644
index 00000000000..4170849f65d
--- /dev/null
+++ b/litellm-rust/crates/ai-gateway/render.yaml
@@ -0,0 +1,35 @@
+# Render blueprint for the LiteLLM Rust AI Gateway (realtime WebSocket proxy).
+#
+# Single instance for now (no autoscaling). The public endpoint is a
+# WebSocket served over TLS: wss://.onrender.com/v1/realtime
+#
+# Paths are relative to the **repo root** (Render's convention). The build
+# context is the repo root so the image can install litellm from source — the
+# gateway loads its model_list via litellm.proxy.read_model_list at boot.
+#
+# Secrets (LITELLM_MASTER_KEY, OPENAI_API_KEY) are marked sync: false — set
+# them in the Render dashboard or via the API, never inline here.
+services:
+ - type: web
+ name: litellm-rust-ai-gateway
+ runtime: docker
+ plan: standard
+ dockerfilePath: ./litellm-rust/crates/ai-gateway/Dockerfile
+ dockerContext: .
+ healthCheckPath: /health/readiness
+ numInstances: 1
+ envVars:
+ # The gateway loads its model_list from this config.yaml via the embedded
+ # python config reader. The image bakes a default config at /app/config.yaml;
+ # a real deploy can override it by mounting a Render secret file at this
+ # same path (Dashboard → Environment → Secret Files) — never inline secrets.
+ - key: LITELLM_CONFIG_PATH
+ value: /app/config.yaml
+ - key: HOST
+ value: 0.0.0.0
+ # Bearer token clients must send on /v1/realtime (fail closed if unset).
+ - key: LITELLM_MASTER_KEY
+ sync: false
+ # Referenced by config.yaml as os.environ/OPENAI_API_KEY for the upstream dial.
+ - key: OPENAI_API_KEY
+ sync: false
diff --git a/litellm-rust/crates/ai-gateway/src/auth/mod.rs b/litellm-rust/crates/ai-gateway/src/auth/mod.rs
new file mode 100644
index 00000000000..e2dd51f656d
--- /dev/null
+++ b/litellm-rust/crates/ai-gateway/src/auth/mod.rs
@@ -0,0 +1,54 @@
+//! Gateway authentication, as an axum **extractor** (the idiomatic pattern —
+//! keeps handlers clean and auth testable).
+//!
+//! For now this is a single **master key**: any caller presenting it as
+//! `Authorization: Bearer ` may invoke the gateway. Per-key auth, budgets,
+//! and rate limits are delegated to the Python proxy in a later phase.
+//!
+//! A handler opts in by adding [`RequireMasterKey`] to its arguments; auth then
+//! runs during extraction, before the handler body. Routes never re-implement it.
+
+use axum::extract::FromRequestParts;
+use axum::http::header::AUTHORIZATION;
+use axum::http::request::Parts;
+use axum::http::StatusCode;
+use subtle::ConstantTimeEq;
+
+use crate::state::AppState;
+
+/// Extractor that requires the configured master key as a bearer token.
+///
+/// Rejections: `500` when no master key is configured (permanent
+/// misconfiguration, not a transient outage); `401` on a missing/incorrect
+/// token. The comparison is constant-time.
+pub struct RequireMasterKey;
+
+#[axum::async_trait]
+impl FromRequestParts for RequireMasterKey {
+ type Rejection = (StatusCode, String);
+
+ async fn from_request_parts(
+ parts: &mut Parts,
+ state: &AppState,
+ ) -> Result {
+ let Some(expected) = state.master_key.as_deref() else {
+ return Err((
+ StatusCode::INTERNAL_SERVER_ERROR,
+ "gateway auth not configured (set LITELLM_MASTER_KEY)".to_string(),
+ ));
+ };
+ let provided = parts
+ .headers
+ .get(AUTHORIZATION)
+ .and_then(|value| value.to_str().ok())
+ .and_then(|value| value.strip_prefix("Bearer "))
+ .map(str::trim);
+ match provided {
+ Some(token) if bool::from(token.as_bytes().ct_eq(expected.as_bytes())) => Ok(Self),
+ _ => Err((
+ StatusCode::UNAUTHORIZED,
+ "missing or invalid bearer token".to_string(),
+ )),
+ }
+ }
+}
diff --git a/litellm-rust/crates/ai-gateway/src/gil.rs b/litellm-rust/crates/ai-gateway/src/gil.rs
new file mode 100644
index 00000000000..c749f722c73
--- /dev/null
+++ b/litellm-rust/crates/ai-gateway/src/gil.rs
@@ -0,0 +1,58 @@
+//! GIL-activity tracking.
+//!
+//! Every acquisition of the Python GIL is recorded here so the `/health/gil`
+//! endpoint can report whether Python was touched recently. The design goal is
+//! that the GIL is acquired **only at load time** (config read) and never on the
+//! realtime hot path — polling this endpoint during traffic should show the
+//! count holding steady and `acquired_last_30s` falling to `false`.
+
+use std::sync::atomic::{AtomicU64, Ordering};
+use std::time::{SystemTime, UNIX_EPOCH};
+
+/// Window (seconds) for the "recently acquired" signal.
+pub const RECENT_WINDOW_SECS: u64 = 30;
+
+static GIL_ACQUISITIONS: AtomicU64 = AtomicU64::new(0);
+/// Unix seconds of the last acquisition; `0` means "never".
+static LAST_GIL_UNIX_SECS: AtomicU64 = AtomicU64::new(0);
+
+fn now_unix_secs() -> u64 {
+ SystemTime::now()
+ .duration_since(UNIX_EPOCH)
+ .map(|d| d.as_secs())
+ .unwrap_or(0)
+}
+
+/// Record that the GIL was just acquired. Call immediately before taking the GIL.
+///
+/// Only invoked under the `python-config` feature; without it the gateway never
+/// touches Python, so the recorder is unused (and the endpoint reports zero).
+#[cfg_attr(not(feature = "python-config"), allow(dead_code))]
+pub fn record_acquisition() {
+ GIL_ACQUISITIONS.fetch_add(1, Ordering::Relaxed);
+ LAST_GIL_UNIX_SECS.store(now_unix_secs(), Ordering::Relaxed);
+}
+
+/// Point-in-time view of GIL activity.
+pub struct GilSnapshot {
+ pub total_acquisitions: u64,
+ pub seconds_since_last: Option,
+ pub acquired_last_30s: bool,
+}
+
+/// Read the current GIL-activity snapshot.
+pub fn snapshot() -> GilSnapshot {
+ let total = GIL_ACQUISITIONS.load(Ordering::Relaxed);
+ let last = LAST_GIL_UNIX_SECS.load(Ordering::Relaxed);
+ let seconds_since_last = if last == 0 {
+ None
+ } else {
+ Some(now_unix_secs().saturating_sub(last))
+ };
+ let acquired_last_30s = seconds_since_last.is_some_and(|secs| secs <= RECENT_WINDOW_SECS);
+ GilSnapshot {
+ total_acquisitions: total,
+ seconds_since_last,
+ acquired_last_30s,
+ }
+}
diff --git a/litellm-rust/crates/ai-gateway/src/main.rs b/litellm-rust/crates/ai-gateway/src/main.rs
new file mode 100644
index 00000000000..71e4a6836ad
--- /dev/null
+++ b/litellm-rust/crates/ai-gateway/src/main.rs
@@ -0,0 +1,152 @@
+//! LiteLLM AI Gateway — a minimal Axum server fronting the Rust router.
+//!
+//! Flow: client → `POST /v1/realtime` → `router.realtime()` selects a deployment
+//! (simple-shuffle) → `providers::realtime::realtime()` invokes OpenAI. The
+//! server owns transport + config; routing lives in the `router` crate.
+
+mod auth;
+mod gil;
+#[cfg(feature = "python-config")]
+mod python;
+mod routes;
+mod state;
+
+use std::sync::Arc;
+
+use litellm_core::router::{Deployment, LiteLLMParams, Router};
+use litellm_providers::realtime_pool::{upstream_key, PoolConfig, RealtimePool};
+
+use crate::state::AppState;
+
+/// Bind to localhost by default so the gateway is not a public, unauthenticated
+/// provider proxy out of the box. Override with `HOST` (e.g. `0.0.0.0`).
+const DEFAULT_HOST: &str = "127.0.0.1";
+const DEFAULT_PORT: u16 = 4001;
+
+#[tokio::main]
+async fn main() {
+ // Trim before storing so it matches the trimmed bearer token in `auth`
+ // (avoids a silent auth failure when the env var has surrounding whitespace).
+ let master_key: Option> = std::env::var("LITELLM_MASTER_KEY")
+ .ok()
+ .map(|key| key.trim().to_string())
+ .filter(|key| !key.is_empty())
+ .map(Arc::from);
+ if master_key.is_none() {
+ eprintln!(
+ "warning: LITELLM_MASTER_KEY is not set; /v1/realtime will reject all requests (fail closed)"
+ );
+ }
+
+ let router = Arc::new(build_router());
+
+ // Build the pre-warmed realtime pool and register each deployment's upstream
+ // so the background replenisher starts warming it. `REALTIME_POOL_SIZE=0`
+ // yields a disabled pool → every connect fresh-dials (original behavior).
+ let pool_config = PoolConfig::from_env();
+ let realtime_pool = RealtimePool::spawn(pool_config);
+ if pool_config.enabled() {
+ register_deployments(&router, &realtime_pool);
+ eprintln!(
+ "realtime connection pool enabled: target {} warm sockets/key, max idle {}s",
+ pool_config.target_size,
+ pool_config.max_idle.as_secs()
+ );
+ } else {
+ eprintln!(
+ "realtime connection pool disabled (REALTIME_POOL_SIZE=0); fresh-dialing each connect"
+ );
+ }
+
+ let state = AppState {
+ router,
+ master_key,
+ realtime_pool,
+ };
+
+ let host = std::env::var("HOST").unwrap_or_else(|_| DEFAULT_HOST.to_string());
+ let port = resolve_port();
+
+ let listener = tokio::net::TcpListener::bind((host.as_str(), port))
+ .await
+ .expect("failed to bind listener");
+ eprintln!("litellm-ai-gateway listening on {host}:{port}");
+ axum::serve(listener, routes::app(state))
+ .await
+ .expect("server error");
+}
+
+/// Register every deployment's upstream key with the pool so the replenisher
+/// pre-warms it. Mirrors `service::run`'s key derivation (strip `openai/`, resolve
+/// api_key); deployments whose key can't be resolved are skipped (they fresh-dial
+/// and surface the auth error on the request path, as before).
+fn register_deployments(router: &Router, pool: &RealtimePool) {
+ for deployment in router.deployments() {
+ let params = &deployment.litellm_params;
+ let provider_model = params
+ .model
+ .strip_prefix("openai/")
+ .unwrap_or(¶ms.model);
+ if let Some(key) = upstream_key(
+ provider_model,
+ params.api_key.as_deref(),
+ params.api_base.as_deref(),
+ ) {
+ pool.register(key);
+ }
+ }
+}
+
+/// Resolve `PORT`, warning (rather than silently defaulting) on an invalid value.
+fn resolve_port() -> u16 {
+ match std::env::var("PORT") {
+ Ok(raw) => raw.parse().unwrap_or_else(|_| {
+ eprintln!("warning: PORT={raw:?} is not a valid port; using {DEFAULT_PORT}");
+ DEFAULT_PORT
+ }),
+ Err(_) => DEFAULT_PORT,
+ }
+}
+
+/// Build the router. With the `python-config` feature and `LITELLM_CONFIG_PATH`
+/// set, load the resolved `model_list` from the proxy config via the embedded
+/// Python reader (load time only). Otherwise fall back to the env stand-in.
+fn build_router() -> Router {
+ #[cfg(feature = "python-config")]
+ if let Ok(config_path) = std::env::var("LITELLM_CONFIG_PATH") {
+ match python::config::load_router_from_config(&config_path) {
+ Ok(router) => {
+ eprintln!("loaded model_list from {config_path} via python config reader");
+ return router;
+ }
+ Err(err) => {
+ eprintln!("config load failed ({err}); falling back to env deployment");
+ }
+ }
+ }
+ build_router_from_env()
+}
+
+/// Build a minimal single-deployment `model_list` from the environment.
+///
+/// A real deployment loads `model_list` from config; this is the minimal stand-in
+/// so the gateway has one OpenAI deployment to route to.
+fn build_router_from_env() -> Router {
+ let model =
+ std::env::var("OPENAI_REALTIME_MODEL").unwrap_or_else(|_| "gpt-realtime".to_string());
+ let api_key = std::env::var("OPENAI_API_KEY").ok();
+ if api_key.is_none() {
+ eprintln!(
+ "warning: OPENAI_API_KEY is not set; realtime requests will fail with auth errors"
+ );
+ }
+ let deployment = Deployment {
+ model_name: model.clone(),
+ litellm_params: LiteLLMParams {
+ model,
+ api_key,
+ api_base: None,
+ },
+ };
+ Router::new(vec![deployment])
+}
diff --git a/litellm-rust/crates/ai-gateway/src/python/AGENTS.md b/litellm-rust/crates/ai-gateway/src/python/AGENTS.md
new file mode 100644
index 00000000000..47aa117e0b9
--- /dev/null
+++ b/litellm-rust/crates/ai-gateway/src/python/AGENTS.md
@@ -0,0 +1,27 @@
+# ai-gateway/src/python — Python interop (load-time only)
+
+Functions here embed the Python interpreter (pyo3) and take the GIL to call into
+`litellm` (e.g. read the proxy `model_list`). Compiled only under the
+`python-config` feature.
+
+## Hard rule: non-hot-path functions only
+
+Everything in this folder MUST run **at most once per process lifetime — at
+startup / load time** (config read, warm-up). NEVER call into Python on the
+request path:
+
+- No GIL acquisition per request, per connection, or per realtime event.
+- No Python call inside a route handler, the router's hot path, or any loop that
+ scales with traffic.
+
+**Why:** the GIL serializes execution and would cap throughput; the realtime data
+path must stay pure Rust. Every acquisition is recorded by `crate::gil` — poll
+`GET /health/gil`, and `total_acquisitions` MUST stay flat under load.
+
+## How to add one
+
+Resolve whatever Python-derived data you need **once at boot** and hand the rest
+of the gateway an owned, plain-Rust value (e.g. build a `Router` from the
+resolved `model_list`). Record the acquisition via `crate::gil::record_acquisition()`
+immediately before taking the GIL. If a function would need to run per request,
+it does not belong here — move the work to Rust, or pre-resolve it at startup.
diff --git a/litellm-rust/crates/ai-gateway/src/python/config.rs b/litellm-rust/crates/ai-gateway/src/python/config.rs
new file mode 100644
index 00000000000..6ec9595469d
--- /dev/null
+++ b/litellm-rust/crates/ai-gateway/src/python/config.rs
@@ -0,0 +1,39 @@
+//! Build the router by calling the Python proxy config reader (load time only).
+//!
+//! Embeds the interpreter via pyo3 and calls
+//! `litellm.proxy.read_model_list.read_model_list`, which reuses the proxy's
+//! `os.environ/` + secret-manager resolution. The GIL is taken **once at boot**
+//! (and recorded in [`crate::gil`]); the realtime hot path never touches Python.
+//!
+//! Compiled only under the `python-config` feature.
+
+use litellm_core::error::CoreError;
+use litellm_core::router::{Deployment, Router};
+use litellm_core::CoreResult;
+use pyo3::prelude::*;
+
+use crate::gil;
+
+/// Load the router's `model_list` from `config_path` via the Python reader.
+pub fn load_router_from_config(config_path: &str) -> CoreResult {
+ gil::record_acquisition();
+ Python::with_gil(|py| {
+ let model_list = py
+ .import("litellm.proxy.read_model_list")
+ .and_then(|module| module.getattr("read_model_list"))
+ .and_then(|reader| reader.call1((config_path,)))
+ .map_err(|err| CoreError::Routing(format!("read_model_list failed: {err}")))?;
+
+ let model_list_json: String = py
+ .import("json")
+ .and_then(|json| json.getattr("dumps"))
+ .and_then(|dumps| dumps.call1((model_list,)))
+ .and_then(|encoded| encoded.extract())
+ .map_err(|err| CoreError::Routing(format!("serializing model_list failed: {err}")))?;
+
+ let deployments: Vec = serde_json::from_str(&model_list_json)
+ .map_err(|err| CoreError::Routing(format!("parsing model_list failed: {err}")))?;
+
+ Ok(Router::new(deployments))
+ })
+}
diff --git a/litellm-rust/crates/ai-gateway/src/python/mod.rs b/litellm-rust/crates/ai-gateway/src/python/mod.rs
new file mode 100644
index 00000000000..a677bade676
--- /dev/null
+++ b/litellm-rust/crates/ai-gateway/src/python/mod.rs
@@ -0,0 +1,4 @@
+//! Python interop for the gateway. See `AGENTS.md`: **load-time / non-hot-path
+//! only.** Compiled only under the `python-config` feature.
+
+pub mod config;
diff --git a/litellm-rust/crates/ai-gateway/src/routes/AGENTS.md b/litellm-rust/crates/ai-gateway/src/routes/AGENTS.md
new file mode 100644
index 00000000000..02c5f18c4f3
--- /dev/null
+++ b/litellm-rust/crates/ai-gateway/src/routes/AGENTS.md
@@ -0,0 +1,38 @@
+# routes/ — the route template
+
+Every route follows the **same shape** so the layout is predictable. The rule:
+
+> **Each route module exposes `pub fn router() -> Router`.**
+> `routes/mod.rs::app` merges them all and applies state once. Adding a route is:
+> create the module, then add one `.merge(::router())` line.
+
+## Default: one file
+A route is a single file containing `router()` + its handler(s) (handlers stay
+private). This is the norm — don't split until it hurts.
+```
+pub fn router() -> Router { Router::new().route(PATH, get(handle)) }
+async fn handle(...) -> impl IntoResponse { ... }
+```
+`health.rs` and `gil.rs` are examples.
+
+## Split out `service` when there's real logic
+When a route has business logic worth testing without axum, put it in a sibling
+`service` (a file, or a folder if the route grows). The route file stays the
+**axum surface** (router + handler + any socket/SSE adapter); `service` is plain
+Rust with **no axum types**. `realtime/` is the example:
+```
+realtime/
+ mod.rs # axum surface: router() + handler + the WS<->events adapter
+ service.rs # pure logic: select deployment + call provider (no axum) — testable
+```
+Split `service` further (or add `transport`, `repo`, …) only once a single file
+genuinely gets hard to read.
+
+## Invariants
+- **Auth is an extractor, not a manual call.** A handler requires auth by adding
+ `crate::auth::RequireMasterKey` to its arguments; it runs during extraction.
+ Never re-implement the check per route.
+- **Handlers contain no business logic; `service` contains no axum types.**
+- A route owns its paths in its own `router()`; `mod.rs` only merges.
+- Cross-cutting concerns (logging, CORS, timeouts) → Tower layers in `mod.rs`,
+ not duplicated in handlers.
diff --git a/litellm-rust/crates/ai-gateway/src/routes/gil.rs b/litellm-rust/crates/ai-gateway/src/routes/gil.rs
new file mode 100644
index 00000000000..0db0c6f0b14
--- /dev/null
+++ b/litellm-rust/crates/ai-gateway/src/routes/gil.rs
@@ -0,0 +1,30 @@
+//! `GET /health/gil` — poll to confirm Python is only touched at load time.
+//! Simple-route template: a `router()` plus its handler, in one file.
+
+use axum::routing::get;
+use axum::{Json, Router};
+use serde::Serialize;
+
+use crate::gil;
+use crate::state::AppState;
+
+/// This route's contribution to the app router.
+pub fn router() -> Router {
+ Router::new().route("/health/gil", get(status))
+}
+
+#[derive(Debug, Serialize)]
+struct GilStatusResponse {
+ gil_acquired_last_30s: bool,
+ total_acquisitions: u64,
+ seconds_since_last: Option,
+}
+
+async fn status() -> Json {
+ let snapshot = gil::snapshot();
+ Json(GilStatusResponse {
+ gil_acquired_last_30s: snapshot.acquired_last_30s,
+ total_acquisitions: snapshot.total_acquisitions,
+ seconds_since_last: snapshot.seconds_since_last,
+ })
+}
diff --git a/litellm-rust/crates/ai-gateway/src/routes/health.rs b/litellm-rust/crates/ai-gateway/src/routes/health.rs
new file mode 100644
index 00000000000..15c67fea325
--- /dev/null
+++ b/litellm-rust/crates/ai-gateway/src/routes/health.rs
@@ -0,0 +1,24 @@
+//! Health probes. Simple-route template: a `router()` plus its handlers, in one file.
+
+use axum::http::StatusCode;
+use axum::routing::get;
+use axum::Router;
+
+use crate::state::AppState;
+
+/// This route's contribution to the app router.
+pub fn router() -> Router {
+ Router::new()
+ .route("/health/liveness", get(liveness))
+ .route("/health/readiness", get(readiness))
+}
+
+/// The process is up.
+async fn liveness() -> StatusCode {
+ StatusCode::OK
+}
+
+/// The server is ready to accept traffic.
+async fn readiness() -> StatusCode {
+ StatusCode::OK
+}
diff --git a/litellm-rust/crates/ai-gateway/src/routes/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/mod.rs
new file mode 100644
index 00000000000..c6b9573781a
--- /dev/null
+++ b/litellm-rust/crates/ai-gateway/src/routes/mod.rs
@@ -0,0 +1,23 @@
+//! HTTP routes.
+//!
+//! **Template:** every route module exposes `pub fn router() -> Router`
+//! that mounts its own paths; [`app`] merges them. A trivial route is a single
+//! file (`health.rs`, `gil.rs`); a non-trivial one is a folder (`realtime/`) with
+//! `handler` (entry) + `service` (logic) + `transport` (adapters). See AGENTS.md.
+
+pub mod gil;
+pub mod health;
+pub mod realtime;
+
+use axum::Router;
+
+use crate::state::AppState;
+
+/// Assemble the application router by merging every route module's `router()`.
+pub fn app(state: AppState) -> Router {
+ Router::new()
+ .merge(health::router())
+ .merge(gil::router())
+ .merge(realtime::router())
+ .with_state(state)
+}
diff --git a/litellm-rust/crates/ai-gateway/src/routes/realtime/README.md b/litellm-rust/crates/ai-gateway/src/routes/realtime/README.md
new file mode 100644
index 00000000000..3301576bb85
--- /dev/null
+++ b/litellm-rust/crates/ai-gateway/src/routes/realtime/README.md
@@ -0,0 +1,87 @@
+# Realtime route (`GET /v1/realtime`)
+
+Proxies OpenAI's realtime WebSocket. `mod.rs` is the axum surface (handler +
+socket↔events adapter); `service.rs` is the pure logic (select a deployment, then
+splice client ↔ upstream). The pool itself lives in
+`crates/providers/src/realtime_pool.rs`.
+
+## Connection pooling
+
+### The problem
+
+The gateway's realtime overhead lives **entirely in session establishment**. On each
+client connect it dials a *fresh* upstream WS to OpenAI and waits for
+`session.created` before it can serve. Measured at 5000 calls / 500 concurrency, the
+fresh-dial session phase is **~360 ms** vs **~7 ms** direct; dial, first-audio, and
+streaming add ~0. So the one lever is removing that per-connect handshake from the
+critical path.
+
+### The idea
+
+Keep a few upstream OpenAI sockets **already connected and already past
+`session.created`** (buffered). On a client connect, hand off a warm socket — relay
+its buffered `session.created` instantly (a local `Vec::pop`, sub-millisecond) and
+splice exactly as a fresh dial would. A background task keeps the pool topped up. On
+a miss or dead socket we fall back to fresh-dial: the pool is a latency optimization,
+never a correctness dependency.
+
+```
+ ┌───────────────────────────────────────┐
+ client connect ──────► │ routes/realtime → service::run │
+ │ pool.take(key) │
+ │ hit → relay buffered │
+ │ session.created, then splice │
+ │ miss → fresh dial (original path) │
+ └───────────────┬───────────────────────┘
+ │ replenish (async, concurrent)
+ ┌───────────────▼───────────────────────┐
+ background task ─────► │ RealtimePool: per-key warm sockets │
+ │ each = { ws, buffered session.created}│
+ │ liveness-checked before handoff │
+ └─────────────────────────────────────────┘
+```
+
+A warm session is indistinguishable from a fresh one: OpenAI sends `session.created`
+unprompted on connect, we pre-read exactly that one frame and relay it on handoff,
+and we send nothing else on the socket before a client exists — so the client's first
+`session.update` behaves identically either way.
+
+### Sizing
+
+Each warm socket serves **exactly one** session (realtime isn't multiplexed), so the
+pool is sized to the **peak concurrent connects per instance**, not total live
+connections:
+
+```
+REALTIME_POOL_SIZE ≈ peak_concurrency / instance_count
+```
+
+e.g. 500 concurrency over 10 instances → ~50–64 per instance. The replenisher dials
+the missing sockets **concurrently**, so a drained pool refills in ~one handshake
+window and keeps supply close to the connect rate. Over-provisioning just burns idle
+upstream sockets, which is why warm sockets are short-lived
+(`REALTIME_POOL_MAX_IDLE_SECS`).
+
+### Config
+
+| env | default | meaning |
+| ----------------------------- | ------- | --------------------------------------------------------------- |
+| `REALTIME_POOL_SIZE` | `4` | target warm sockets per key. `0` disables pooling (fresh-dial). |
+| `REALTIME_POOL_MAX_IDLE_SECS` | `30` | max time a warm socket sits before it's closed and replaced. |
+
+### Notes
+
+- **Miss / dead socket → fresh dial.** Burst beyond warm supply, or a socket that
+ died, never blocks or fails — it falls back to the original path. The pool can only
+ make a connect faster, never slower or more fragile.
+- **Auth scope.** The pool key includes `api_key`, so a warm socket is only handed to
+ a request resolving to the same key — no cross-tenant reuse.
+- **Idle billing.** Warm sockets are liveness-checked at handoff and capped at
+ `REALTIME_POOL_MAX_IDLE_SECS` to bound idle billing and dodge OpenAI's idle timeout.
+- **Replenish backoff.** If a key's warm-up dials all fail (invalid credentials, an
+ unreachable upstream), the replenisher puts that key into exponential backoff
+ (500 ms → 30 s cap) instead of re-dialing it every tick. This bounds connection
+ attempts against a broken key so it can't exhaust upstream rate limits and degrade
+ valid cold-path traffic; the backoff resets the moment a dial succeeds.
+
+Benchmarks and repro: `../../benchmarks/realtime/README.md`.
diff --git a/litellm-rust/crates/ai-gateway/src/routes/realtime/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/realtime/mod.rs
new file mode 100644
index 00000000000..695e0c6bb39
--- /dev/null
+++ b/litellm-rust/crates/ai-gateway/src/routes/realtime/mod.rs
@@ -0,0 +1,88 @@
+//! `GET /v1/realtime` (WebSocket).
+//!
+//! This file is the **axum surface**: `router()`, the handler, and the small
+//! socket↔events adapter. The pure logic (no axum) lives in [`service`]. Auth is
+//! the `RequireMasterKey` extractor, so the handler stays thin.
+
+mod service;
+
+use std::sync::Arc;
+
+use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
+use axum::extract::{Query, State};
+use axum::http::StatusCode;
+use axum::response::Response;
+use axum::routing::get;
+use axum::Router;
+use futures_util::{SinkExt, StreamExt};
+use litellm_core::realtime::types::RealtimeEvent;
+use litellm_core::router::Router as ModelRouter;
+use litellm_providers::realtime_pool::RealtimePool;
+use serde::Deserialize;
+
+use crate::auth::RequireMasterKey;
+use crate::state::AppState;
+
+/// This route's contribution to the app router.
+pub fn router() -> Router {
+ Router::new().route("/v1/realtime", get(handle))
+}
+
+#[derive(Debug, Deserialize)]
+struct RealtimeQuery {
+ model: String,
+}
+
+/// Auth runs via the `RequireMasterKey` extractor. We validate the model BEFORE
+/// the upgrade so failures are clean HTTP (400/404), not a socket that opens then
+/// closes, then hand the socket to `bridge`.
+async fn handle(
+ _auth: RequireMasterKey,
+ ws: WebSocketUpgrade,
+ State(state): State,
+ Query(query): Query,
+) -> Result {
+ if query.model.trim().is_empty() {
+ return Err((
+ StatusCode::BAD_REQUEST,
+ "missing 'model' query param".to_string(),
+ ));
+ }
+ if !state.router.has_deployment(&query.model) {
+ return Err((
+ StatusCode::NOT_FOUND,
+ format!("no deployment for model '{}'", query.model),
+ ));
+ }
+
+ let router = state.router.clone();
+ let pool = state.realtime_pool.clone();
+ let model = query.model;
+ Ok(ws.on_upgrade(move |socket| bridge(socket, router, pool, model)))
+}
+
+/// Adapt the axum socket (text frames) to the typed-event `Stream`/`Sink` the
+/// service wants, keeping axum types out of `service`.
+async fn bridge(
+ socket: WebSocket,
+ router: Arc,
+ pool: Arc,
+ model: String,
+) {
+ let (ws_sink, ws_stream) = socket.split();
+
+ let client_in = ws_stream.filter_map(|message| async move {
+ match message {
+ Ok(Message::Text(text)) => serde_json::from_str::(&text).ok(),
+ _ => None,
+ }
+ });
+ let client_out = ws_sink.with(|event: RealtimeEvent| async move {
+ Ok::(Message::Text(
+ serde_json::to_string(&event).unwrap_or_default(),
+ ))
+ });
+
+ futures_util::pin_mut!(client_in, client_out);
+ let _ = service::run(&router, &pool, &model, None, client_in, client_out).await;
+}
diff --git a/litellm-rust/crates/ai-gateway/src/routes/realtime/service.rs b/litellm-rust/crates/ai-gateway/src/routes/realtime/service.rs
new file mode 100644
index 00000000000..0cbd00d664f
--- /dev/null
+++ b/litellm-rust/crates/ai-gateway/src/routes/realtime/service.rs
@@ -0,0 +1,76 @@
+//! Business logic: select a deployment with the (pure) core router, then call the
+//! provider splice. The seam between `core::router` (selection only) and
+//! `providers` (the actual WebSocket I/O).
+//!
+//! On connect we try a pre-warmed upstream from the pool (handshake already paid,
+//! `session.created` buffered) and relay it instantly. On a pool miss or dead warm
+//! socket we fresh-dial exactly as before — the pool is never on the critical path
+//! for correctness, only latency.
+
+use std::time::Duration;
+
+use futures_util::{Sink, Stream};
+use litellm_core::error::CoreError;
+use litellm_core::realtime::types::RealtimeEvent;
+use litellm_core::router::Router;
+use litellm_core::CoreResult;
+use litellm_providers::realtime_pool::{upstream_key, RealtimePool};
+
+/// Select a deployment for `model` and splice the client stream to the provider.
+///
+/// `pool` supplies a pre-warmed upstream when one is available; otherwise we
+/// fresh-dial. A disabled pool always misses, so this collapses to the original
+/// fresh-dial behavior.
+pub async fn run(
+ router: &Router,
+ pool: &RealtimePool,
+ model: &str,
+ idle_timeout: Option,
+ client_in: In,
+ client_out: Out,
+) -> CoreResult<()>
+where
+ In: Stream- + Unpin + Send,
+ Out: Sink + Unpin + Send,
+ >::Error: std::fmt::Display,
+{
+ let deployment = router.get_available_deployment(model).ok_or_else(|| {
+ CoreError::Routing(format!("no deployment available for model '{model}'"))
+ })?;
+ let params = &deployment.litellm_params;
+ // Strip a leading `openai/` so the OpenAI-only realtime fn gets the bare model.
+ let provider_model = params
+ .model
+ .strip_prefix("openai/")
+ .unwrap_or(¶ms.model);
+
+ // Warm path: take a pooled upstream (handshake already paid) and relay its
+ // buffered session.created immediately. On miss/dead socket fall through.
+ if let Some(key) = upstream_key(
+ provider_model,
+ params.api_key.as_deref(),
+ params.api_base.as_deref(),
+ ) {
+ if let Some(handoff) = pool.take(&key) {
+ return litellm_providers::realtime::realtime_warm(
+ provider_model,
+ handoff,
+ idle_timeout,
+ client_in,
+ client_out,
+ )
+ .await;
+ }
+ }
+
+ // Cold path: fresh dial (the original behavior).
+ litellm_providers::realtime::realtime(
+ provider_model,
+ params.api_key.as_deref(),
+ params.api_base.as_deref(),
+ idle_timeout,
+ client_in,
+ client_out,
+ )
+ .await
+}
diff --git a/litellm-rust/crates/ai-gateway/src/state.rs b/litellm-rust/crates/ai-gateway/src/state.rs
new file mode 100644
index 00000000000..ef96037d477
--- /dev/null
+++ b/litellm-rust/crates/ai-gateway/src/state.rs
@@ -0,0 +1,17 @@
+use std::sync::Arc;
+
+use litellm_core::router::Router;
+use litellm_providers::realtime_pool::RealtimePool;
+
+/// Shared application state handed to every route handler.
+#[derive(Clone)]
+pub struct AppState {
+ pub router: Arc,
+ /// The gateway master key. Any caller presenting it as a bearer token may
+ /// invoke the gateway. `None` → auth not configured (routes fail closed).
+ pub master_key: Option>,
+ /// Pre-warmed upstream realtime connection pool. Disabled
+ /// (`RealtimePool::disabled()`) when `REALTIME_POOL_SIZE=0`, in which case
+ /// every realtime connect fresh-dials exactly as before.
+ pub realtime_pool: Arc,
+}
diff --git a/litellm-rust/crates/core/CLAUDE.md b/litellm-rust/crates/core/CLAUDE.md
new file mode 100644
index 00000000000..20873878967
--- /dev/null
+++ b/litellm-rust/crates/core/CLAUDE.md
@@ -0,0 +1,47 @@
+# CLAUDE.md
+
+Rules for `litellm-rust/crates/core`.
+
+## Responsibility
+
+`core` owns shared data types, typed errors, and deterministic helper contracts.
+It must stay pure and host-independent.
+
+Allowed:
+- Shared request/response structs.
+- Typed errors with stable, non-sensitive messages.
+- Deterministic validation helpers.
+- Serialization helpers that intentionally mirror Python output shape.
+- Route templates that match Python base config responsibilities, such as
+ `ocr::transformation::OcrProviderConfig`.
+
+Not allowed:
+- Network, filesystem, database, cache, or environment access.
+- Secret reads or auth/header construction.
+- Logging callbacks, tracing spans, spend writes, or customer callbacks.
+- Provider-specific branching that belongs in `providers`.
+- Panics for user/provider-controlled input.
+
+## Typed Contracts (core rule)
+
+Trait and function boundaries MUST be strongly typed. No stringly-typed JSON
+(`&str` / `String` / `Vec` / bare `serde_json::Value`) as a transform
+input or output. Parse wire bytes into typed structs/enums at the host edge;
+`core` and `providers` operate only on those types (e.g. `RealtimeEvent`,
+`RealtimeTransformResult`, `OcrRequestData`). A `type`-style discriminator is a
+typed field on a struct, not a raw string threaded through the API.
+
+## Structure
+
+Use route names directly under `src/`: `ocr`, future `messages`,
+`chat_completions`, `embeddings`, and similar top-level LiteLLM calls. Do not
+invent broad names like `engine` for route contracts.
+
+## Parity Rules
+
+- Every shared type used by a provider transform needs unit tests for
+ serialization shape.
+- If Python parity requires always emitting a `null` field instead of omitting
+ it, document that in code and pin it with a test.
+- Error enums should preserve enough detail for Python/HTTP hosts to map errors
+ consistently without exposing document contents or upstream bodies.
diff --git a/litellm-rust/crates/core/Cargo.toml b/litellm-rust/crates/core/Cargo.toml
new file mode 100644
index 00000000000..1881bcfa602
--- /dev/null
+++ b/litellm-rust/crates/core/Cargo.toml
@@ -0,0 +1,12 @@
+[package]
+name = "litellm-core"
+version = "0.1.0"
+edition.workspace = true
+license.workspace = true
+repository.workspace = true
+
+[dependencies]
+rand.workspace = true
+serde.workspace = true
+serde_json.workspace = true
+thiserror.workspace = true
diff --git a/litellm-rust/crates/core/src/error.rs b/litellm-rust/crates/core/src/error.rs
new file mode 100644
index 00000000000..9b29260cca4
--- /dev/null
+++ b/litellm-rust/crates/core/src/error.rs
@@ -0,0 +1,35 @@
+use thiserror::Error;
+
+pub type CoreResult = Result;
+
+#[derive(Debug, Error, PartialEq, Eq)]
+pub enum CoreError {
+ #[error("expected {expected}, got {actual}")]
+ InvalidType {
+ expected: &'static str,
+ actual: &'static str,
+ },
+ #[error("missing required field: {0}")]
+ MissingField(&'static str),
+ #[error("invalid response: {0}")]
+ InvalidResponse(String),
+ #[error("{0}")]
+ Auth(String),
+ #[error("OCR request failed with status {status}: {body}")]
+ Http { status: u16, body: String },
+ #[error("OCR network error: {0}")]
+ Network(String),
+ #[error("routing error: {0}")]
+ Routing(String),
+}
+
+pub fn json_type_name(value: &serde_json::Value) -> &'static str {
+ match value {
+ serde_json::Value::Null => "null",
+ serde_json::Value::Bool(_) => "bool",
+ serde_json::Value::Number(_) => "number",
+ serde_json::Value::String(_) => "string",
+ serde_json::Value::Array(_) => "array",
+ serde_json::Value::Object(_) => "object",
+ }
+}
diff --git a/litellm-rust/crates/core/src/lib.rs b/litellm-rust/crates/core/src/lib.rs
new file mode 100644
index 00000000000..9d686626edc
--- /dev/null
+++ b/litellm-rust/crates/core/src/lib.rs
@@ -0,0 +1,6 @@
+pub mod error;
+pub mod ocr;
+pub mod realtime;
+pub mod router;
+
+pub use error::{CoreError, CoreResult};
diff --git a/litellm-rust/crates/core/src/ocr/mod.rs b/litellm-rust/crates/core/src/ocr/mod.rs
new file mode 100644
index 00000000000..ec2fbb969a6
--- /dev/null
+++ b/litellm-rust/crates/core/src/ocr/mod.rs
@@ -0,0 +1,2 @@
+pub mod transformation;
+pub mod types;
diff --git a/litellm-rust/crates/core/src/ocr/transformation.rs b/litellm-rust/crates/core/src/ocr/transformation.rs
new file mode 100644
index 00000000000..7353d9d22c4
--- /dev/null
+++ b/litellm-rust/crates/core/src/ocr/transformation.rs
@@ -0,0 +1,32 @@
+use serde_json::{Map, Value};
+
+use crate::CoreResult;
+
+use super::types::{OcrRequestData, OcrResponseData};
+
+pub trait OcrProviderConfig {
+ fn supported_ocr_params(&self) -> &'static [&'static str];
+
+ fn map_ocr_params(&self, non_default_params: &Map) -> Map {
+ let mut mapped_params = Map::new();
+ for (param, value) in non_default_params {
+ if self.supported_ocr_params().contains(¶m.as_str()) {
+ mapped_params.insert(param.clone(), value.clone());
+ }
+ }
+ mapped_params
+ }
+
+ fn transform_ocr_request(
+ &self,
+ model: &str,
+ document: Value,
+ optional_params: Map,
+ ) -> CoreResult;
+
+ fn transform_ocr_response(
+ &self,
+ model: &str,
+ response_json: Value,
+ ) -> CoreResult;
+}
diff --git a/litellm-rust/crates/core/src/ocr/types.rs b/litellm-rust/crates/core/src/ocr/types.rs
new file mode 100644
index 00000000000..1a72b8f1d66
--- /dev/null
+++ b/litellm-rust/crates/core/src/ocr/types.rs
@@ -0,0 +1,29 @@
+use serde::{Deserialize, Serialize};
+use serde_json::Value;
+
+#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
+pub struct OcrRequestData {
+ pub data: Value,
+ pub files: Option,
+}
+
+#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
+pub struct OcrResponseData {
+ pub pages: Vec,
+ pub model: String,
+ pub document_annotation: Option,
+ pub usage_info: Option,
+ pub object: String,
+}
+
+impl OcrResponseData {
+ pub fn into_json(self) -> Value {
+ serde_json::json!({
+ "pages": self.pages,
+ "model": self.model,
+ "document_annotation": self.document_annotation,
+ "usage_info": self.usage_info,
+ "object": self.object,
+ })
+ }
+}
diff --git a/litellm-rust/crates/core/src/realtime/mod.rs b/litellm-rust/crates/core/src/realtime/mod.rs
new file mode 100644
index 00000000000..ec2fbb969a6
--- /dev/null
+++ b/litellm-rust/crates/core/src/realtime/mod.rs
@@ -0,0 +1,2 @@
+pub mod transformation;
+pub mod types;
diff --git a/litellm-rust/crates/core/src/realtime/transformation.rs b/litellm-rust/crates/core/src/realtime/transformation.rs
new file mode 100644
index 00000000000..a4baa27a6c2
--- /dev/null
+++ b/litellm-rust/crates/core/src/realtime/transformation.rs
@@ -0,0 +1,22 @@
+use crate::realtime::types::{RealtimeEvent, RealtimeTransformResult};
+use crate::CoreResult;
+
+pub trait RealtimeProviderConfig {
+ /// Build the upstream WebSocket URL (e.g. `wss://api.openai.com/v1/realtime?model=…`).
+ /// Pure string construction only — no network, no env.
+ fn complete_url(&self, api_base: Option<&str>, model: &str) -> String;
+
+ /// Transform a client → backend event before it is forwarded upstream.
+ fn transform_realtime_request(
+ &self,
+ event: &RealtimeEvent,
+ model: &str,
+ ) -> CoreResult;
+
+ /// Transform a backend → client event before it is forwarded downstream.
+ fn transform_realtime_response(
+ &self,
+ event: &RealtimeEvent,
+ model: &str,
+ ) -> CoreResult;
+}
diff --git a/litellm-rust/crates/core/src/realtime/types.rs b/litellm-rust/crates/core/src/realtime/types.rs
new file mode 100644
index 00000000000..3b59224b6e9
--- /dev/null
+++ b/litellm-rust/crates/core/src/realtime/types.rs
@@ -0,0 +1,60 @@
+use serde::{Deserialize, Serialize};
+use serde_json::{Map, Value};
+
+/// A single realtime event exchanged over the WebSocket.
+///
+/// The `type` discriminator is a typed field; the remaining fields are
+/// preserved losslessly in `data` so a transform can pass an event through, or
+/// inspect/modify specific fields, without enumerating every event variant.
+/// Wire (de)serialization happens at the host edge — `core`/`providers` operate
+/// only on this typed form.
+#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
+pub struct RealtimeEvent {
+ #[serde(rename = "type")]
+ pub event_type: String,
+ #[serde(flatten)]
+ pub data: Map,
+}
+
+/// One or more typed events produced by a realtime transform.
+#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
+pub struct RealtimeTransformResult {
+ pub events: Vec,
+}
+
+impl RealtimeTransformResult {
+ /// Forward a single event unchanged (the OpenAI baseline).
+ pub fn passthrough(event: RealtimeEvent) -> Self {
+ Self {
+ events: vec![event],
+ }
+ }
+}
+
+#[cfg(test)]
+mod tests {
+ use super::*;
+
+ fn event(raw: &str) -> RealtimeEvent {
+ serde_json::from_str(raw).expect("valid event json")
+ }
+
+ #[test]
+ fn realtime_event_round_trips_type_and_extra_fields() {
+ let raw = r#"{"type":"response.output_text.delta","delta":"hi","response_id":"r1"}"#;
+ let parsed = event(raw);
+ assert_eq!(parsed.event_type, "response.output_text.delta");
+ assert_eq!(parsed.data.get("delta"), Some(&Value::String("hi".into())));
+ // Re-serializing yields a semantically-equal event (key order may differ).
+ let reparsed: RealtimeEvent =
+ serde_json::from_str(&serde_json::to_string(&parsed).unwrap()).unwrap();
+ assert_eq!(parsed, reparsed);
+ }
+
+ #[test]
+ fn passthrough_produces_single_element_vec() {
+ let parsed = event(r#"{"type":"session.update"}"#);
+ let result = RealtimeTransformResult::passthrough(parsed.clone());
+ assert_eq!(result.events, vec![parsed]);
+ }
+}
diff --git a/litellm-rust/crates/core/src/router/deployment.rs b/litellm-rust/crates/core/src/router/deployment.rs
new file mode 100644
index 00000000000..1ee88e682a3
--- /dev/null
+++ b/litellm-rust/crates/core/src/router/deployment.rs
@@ -0,0 +1,44 @@
+//! `model_list` data types, mirroring Python's deployment dict. Deserialize-ready
+//! so a deployment can be loaded straight from the proxy config's `model_list`.
+
+use serde::Deserialize;
+
+/// Per-deployment call parameters, mirroring Python's `litellm_params`.
+#[derive(Clone, Debug, Deserialize)]
+pub struct LiteLLMParams {
+ /// Provider model, e.g. `gpt-realtime` or `openai/gpt-realtime`.
+ pub model: String,
+ #[serde(default)]
+ pub api_key: Option,
+ #[serde(default)]
+ pub api_base: Option,
+}
+
+/// One entry of the `model_list`, mirroring Python's deployment dict.
+#[derive(Clone, Debug, Deserialize)]
+pub struct Deployment {
+ /// Public alias clients request, e.g. `gpt-realtime`.
+ pub model_name: String,
+ pub litellm_params: LiteLLMParams,
+}
+
+#[cfg(test)]
+mod tests {
+ use super::*;
+
+ #[test]
+ fn deserializes_from_model_list_entry() {
+ let entry = r#"{
+ "model_name": "gpt-realtime",
+ "litellm_params": {"model": "openai/gpt-realtime", "api_base": "https://x"}
+ }"#;
+ let deployment: Deployment = serde_json::from_str(entry).expect("valid entry");
+ assert_eq!(deployment.model_name, "gpt-realtime");
+ assert_eq!(deployment.litellm_params.model, "openai/gpt-realtime");
+ assert_eq!(deployment.litellm_params.api_key, None);
+ assert_eq!(
+ deployment.litellm_params.api_base.as_deref(),
+ Some("https://x")
+ );
+ }
+}
diff --git a/litellm-rust/crates/core/src/router/mod.rs b/litellm-rust/crates/core/src/router/mod.rs
new file mode 100644
index 00000000000..96bc91bc6b5
--- /dev/null
+++ b/litellm-rust/crates/core/src/router/mod.rs
@@ -0,0 +1,93 @@
+//! Minimal Rust port of LiteLLM's `router.py` deployment selection.
+//!
+//! A [`Router`] is built from a `model_list` of [`Deployment`]s
+//! (`{ model_name, litellm_params: { model, api_key, api_base } }`) and selects
+//! one per request via a [`RoutingStrategy`]. For now the only strategy is
+//! `simple-shuffle` — a uniform random pick within a `model_name` group.
+//!
+//! This stays pure (no I/O): it only *chooses* a deployment. The host (the
+//! gateway) takes the chosen deployment and performs the actual provider call.
+//!
+//! - [`deployment`] — the `model_list` data types.
+//! - [`strategy`] — how a deployment is chosen.
+
+mod deployment;
+mod strategy;
+
+pub use deployment::{Deployment, LiteLLMParams};
+pub use strategy::RoutingStrategy;
+
+/// Load-balancing router over a `model_list`.
+#[derive(Clone, Debug, Default)]
+pub struct Router {
+ model_list: Vec,
+ routing_strategy: RoutingStrategy,
+}
+
+impl Router {
+ /// Build a router from a `model_list` using the default `simple-shuffle` strategy.
+ pub fn new(model_list: Vec) -> Self {
+ Self {
+ model_list,
+ routing_strategy: RoutingStrategy::SimpleShuffle,
+ }
+ }
+
+ /// All deployments in the `model_list`. Read-only; used by the host to
+ /// enumerate upstreams (e.g. to pre-warm a connection pool per deployment).
+ pub fn deployments(&self) -> &[Deployment] {
+ &self.model_list
+ }
+
+ /// Whether any deployment is registered under `model`.
+ pub fn has_deployment(&self, model: &str) -> bool {
+ self.model_list
+ .iter()
+ .any(|deployment| deployment.model_name == model)
+ }
+
+ /// Pick a deployment for `model` per the routing strategy. Returns `None`
+ /// when no deployment is registered under that `model_name`.
+ pub fn get_available_deployment(&self, model: &str) -> Option<&Deployment> {
+ let candidates: Vec<&Deployment> = self
+ .model_list
+ .iter()
+ .filter(|deployment| deployment.model_name == model)
+ .collect();
+ self.routing_strategy.select(&candidates)
+ }
+}
+
+#[cfg(test)]
+mod tests {
+ use super::*;
+
+ fn deployment(name: &str, model: &str) -> Deployment {
+ Deployment {
+ model_name: name.to_string(),
+ litellm_params: LiteLLMParams {
+ model: model.to_string(),
+ api_key: None,
+ api_base: None,
+ },
+ }
+ }
+
+ #[test]
+ fn selects_a_matching_deployment() {
+ let router = Router::new(vec![
+ deployment("gpt-realtime", "gpt-realtime"),
+ deployment("other", "other-model"),
+ ]);
+ let chosen = router
+ .get_available_deployment("gpt-realtime")
+ .expect("a deployment should match");
+ assert_eq!(chosen.model_name, "gpt-realtime");
+ }
+
+ #[test]
+ fn unknown_model_returns_none() {
+ let router = Router::new(vec![deployment("gpt-realtime", "gpt-realtime")]);
+ assert!(router.get_available_deployment("missing").is_none());
+ }
+}
diff --git a/litellm-rust/crates/core/src/router/strategy/mod.rs b/litellm-rust/crates/core/src/router/strategy/mod.rs
new file mode 100644
index 00000000000..7e8ac217db3
--- /dev/null
+++ b/litellm-rust/crates/core/src/router/strategy/mod.rs
@@ -0,0 +1,26 @@
+//! Routing policy: how the router picks one deployment from a model group.
+//!
+//! One module per strategy; [`RoutingStrategy::select`] dispatches to it. New
+//! strategies (least-busy, latency-based, …) get their own file here.
+
+mod simple_shuffle;
+
+use super::Deployment;
+
+/// How the router chooses among the deployments sharing a `model_name`.
+#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
+pub enum RoutingStrategy {
+ /// Uniform random pick among the matching deployments.
+ #[default]
+ SimpleShuffle,
+}
+
+impl RoutingStrategy {
+ /// Choose one deployment from `candidates` (all sharing the requested
+ /// `model_name`). Returns `None` when there are no candidates.
+ pub fn select<'a>(&self, candidates: &[&'a Deployment]) -> Option<&'a Deployment> {
+ match self {
+ RoutingStrategy::SimpleShuffle => simple_shuffle::select(candidates),
+ }
+ }
+}
diff --git a/litellm-rust/crates/core/src/router/strategy/simple_shuffle.rs b/litellm-rust/crates/core/src/router/strategy/simple_shuffle.rs
new file mode 100644
index 00000000000..74ce0c21e80
--- /dev/null
+++ b/litellm-rust/crates/core/src/router/strategy/simple_shuffle.rs
@@ -0,0 +1,47 @@
+//! `simple-shuffle`: a uniform random pick among the candidate deployments.
+
+use rand::seq::SliceRandom;
+
+use crate::router::Deployment;
+
+/// Uniform random choice among `candidates` (all sharing the requested
+/// `model_name`). Returns `None` when there are no candidates.
+pub fn select<'a>(candidates: &[&'a Deployment]) -> Option<&'a Deployment> {
+ candidates.choose(&mut rand::thread_rng()).copied()
+}
+
+#[cfg(test)]
+mod tests {
+ use super::*;
+ use crate::router::{Deployment, LiteLLMParams};
+
+ fn deployment(model: &str) -> Deployment {
+ Deployment {
+ model_name: "gpt-realtime".to_string(),
+ litellm_params: LiteLLMParams {
+ model: model.to_string(),
+ api_key: None,
+ api_base: None,
+ },
+ }
+ }
+
+ #[test]
+ fn picks_from_candidates() {
+ let a = deployment("key-a");
+ let b = deployment("key-b");
+ let candidates = vec![&a, &b];
+ for _ in 0..20 {
+ let chosen = select(&candidates).expect("non-empty");
+ assert!(matches!(
+ chosen.litellm_params.model.as_str(),
+ "key-a" | "key-b"
+ ));
+ }
+ }
+
+ #[test]
+ fn empty_candidates_select_none() {
+ assert!(select(&[]).is_none());
+ }
+}
diff --git a/litellm-rust/crates/providers/CLAUDE.md b/litellm-rust/crates/providers/CLAUDE.md
new file mode 100644
index 00000000000..0f7fdcda2aa
--- /dev/null
+++ b/litellm-rust/crates/providers/CLAUDE.md
@@ -0,0 +1,53 @@
+# CLAUDE.md
+
+Rules for `litellm-rust/crates/providers`.
+
+## Responsibility
+
+`providers` owns provider-specific pure transforms. It mirrors the existing
+Python provider modules closely enough that parity review is mechanical.
+
+Provider files should map to the Python provider tree:
+
+```text
+providers/src///transformation.rs
+```
+
+For example, Mistral OCR lives at
+`providers/src/mistral/ocr/transformation.rs`, matching
+`litellm/llms/mistral/ocr/transformation.py`.
+
+Allowed:
+- Provider request transforms.
+- Provider response normalization.
+- Supported-parameter filtering.
+- Provider-specific validation that does not require I/O or secrets.
+
+Not allowed:
+- HTTP clients or provider SDK calls.
+- Environment variable reads.
+- API key resolution or auth header construction.
+- Logging, callbacks, spend tracking, retries, routing, cooldowns, or fallbacks.
+- Panics on bad user/provider input.
+
+## Required Tests
+
+Every provider transform must include focused unit tests for:
+- Supported params matching the Python provider config.
+- Unknown params being dropped or transformed the same way as Python.
+- Request body shape matching Python output.
+- Response normalization with complete, missing, null, and extra fields.
+- Bad input returning typed errors.
+
+For OCR specifically, assume documents can contain personal data. Tests should
+prove transforms do not copy document contents into error messages.
+
+## Implementation Rules
+
+- Prefer static supported-parameter lists over allocating strings on every call.
+- Keep transforms deterministic and allocation-conscious, but choose clarity over
+ premature micro-optimization for tiny parameter lists.
+- Use typed errors from `core`; avoid stringly-typed error plumbing.
+- Add comments only when they explain Python-parity decisions or provider quirks.
+- Put route-level provider dispatch in a route file such as `providers/src/ocr.rs`.
+ Do not move provider-specific transform logic into the Python bridge.
diff --git a/litellm-rust/crates/providers/Cargo.toml b/litellm-rust/crates/providers/Cargo.toml
new file mode 100644
index 00000000000..c5b41424d66
--- /dev/null
+++ b/litellm-rust/crates/providers/Cargo.toml
@@ -0,0 +1,18 @@
+[package]
+name = "litellm-providers"
+version = "0.1.0"
+edition.workspace = true
+license.workspace = true
+repository.workspace = true
+
+[dependencies]
+litellm-core.workspace = true
+reqwest.workspace = true
+serde_json.workspace = true
+tokio.workspace = true
+tokio-tungstenite.workspace = true
+futures-util.workspace = true
+
+[dev-dependencies]
+serde_json.workspace = true
+futures-channel = "0.3"
diff --git a/litellm-rust/crates/providers/src/lib.rs b/litellm-rust/crates/providers/src/lib.rs
new file mode 100644
index 00000000000..40e18961f43
--- /dev/null
+++ b/litellm-rust/crates/providers/src/lib.rs
@@ -0,0 +1,5 @@
+pub mod mistral;
+pub mod ocr;
+pub mod openai;
+pub mod realtime;
+pub mod realtime_pool;
diff --git a/litellm-rust/crates/providers/src/mistral/mod.rs b/litellm-rust/crates/providers/src/mistral/mod.rs
new file mode 100644
index 00000000000..3621ff6a2fd
--- /dev/null
+++ b/litellm-rust/crates/providers/src/mistral/mod.rs
@@ -0,0 +1 @@
+pub mod ocr;
diff --git a/litellm-rust/crates/providers/src/mistral/ocr/mod.rs b/litellm-rust/crates/providers/src/mistral/ocr/mod.rs
new file mode 100644
index 00000000000..f239b6921fa
--- /dev/null
+++ b/litellm-rust/crates/providers/src/mistral/ocr/mod.rs
@@ -0,0 +1 @@
+pub mod transformation;
diff --git a/litellm-rust/crates/providers/src/mistral/ocr/transformation.rs b/litellm-rust/crates/providers/src/mistral/ocr/transformation.rs
new file mode 100644
index 00000000000..fd691177783
--- /dev/null
+++ b/litellm-rust/crates/providers/src/mistral/ocr/transformation.rs
@@ -0,0 +1,292 @@
+use litellm_core::error::{json_type_name, CoreError, CoreResult};
+use litellm_core::ocr::transformation::OcrProviderConfig;
+use litellm_core::ocr::types::{OcrRequestData, OcrResponseData};
+use serde_json::{Map, Value};
+
+const SUPPORTED_OCR_PARAMS: &[&str] = &[
+ "pages",
+ "include_image_base64",
+ "image_limit",
+ "image_min_size",
+ "bbox_annotation_format",
+ "document_annotation_format",
+ "document_annotation_prompt",
+ "extract_header",
+ "extract_footer",
+ "table_format",
+ "confidence_scores_granularity",
+ "id",
+];
+
+/// Default Mistral API base, used when the caller does not override `api_base`.
+pub const MISTRAL_DEFAULT_API_BASE: &str = "https://api.mistral.ai/v1";
+
+/// Environment variable holding the Mistral API key.
+pub const MISTRAL_API_KEY_ENV: &str = "MISTRAL_API_KEY";
+
+/// Error message raised when no Mistral API key can be resolved.
+pub const MISSING_KEY_MESSAGE: &str = "Missing Mistral API Key - A call is being made to Mistral but no key is set either in the environment variables or via params";
+
+/// Build the complete OCR endpoint URL, de-duplicating a trailing `/v1`.
+///
+/// Blank/whitespace `api_base` is treated as absent (guard at resolution time).
+pub fn complete_url(api_base: Option<&str>) -> String {
+ let base = api_base
+ .map(str::trim)
+ .filter(|base| !base.is_empty())
+ .unwrap_or(MISTRAL_DEFAULT_API_BASE)
+ .trim_end_matches('/');
+
+ if base.ends_with("/v1") {
+ format!("{base}/ocr")
+ } else {
+ format!("{base}/v1/ocr")
+ }
+}
+
+/// Resolve the Mistral API key from the explicit param or the environment.
+///
+/// Blank/whitespace values are treated as absent. Returns `CoreError::Auth`
+/// when no usable key is available.
+///
+/// Note: the env fallback only reads the process environment. Secret-manager
+/// backends (AWS/Azure/GCP/Vault) are resolved on the Python side and passed in
+/// via `api_key`; this fallback is a last resort for direct/standalone use.
+pub fn resolve_api_key(
+ api_key: Option<&str>,
+ env_lookup: &dyn Fn(&str) -> Option,
+) -> CoreResult {
+ api_key
+ .map(str::trim)
+ .filter(|key| !key.is_empty())
+ .map(str::to_string)
+ .or_else(|| env_lookup(MISTRAL_API_KEY_ENV).filter(|key| !key.trim().is_empty()))
+ .ok_or_else(|| CoreError::Auth(MISSING_KEY_MESSAGE.to_string()))
+}
+
+pub struct MistralOcrConfig;
+
+pub const MISTRAL_OCR_CONFIG: MistralOcrConfig = MistralOcrConfig;
+
+impl OcrProviderConfig for MistralOcrConfig {
+ fn supported_ocr_params(&self) -> &'static [&'static str] {
+ SUPPORTED_OCR_PARAMS
+ }
+
+ fn transform_ocr_request(
+ &self,
+ model: &str,
+ document: Value,
+ optional_params: Map,
+ ) -> CoreResult {
+ if !document.is_object() {
+ return Err(CoreError::InvalidType {
+ expected: "object",
+ actual: json_type_name(&document),
+ });
+ }
+
+ let mut data = Map::new();
+ data.insert("model".to_string(), Value::String(model.to_string()));
+ data.insert("document".to_string(), document);
+ for (param, value) in optional_params {
+ data.insert(param, value);
+ }
+
+ Ok(OcrRequestData {
+ data: Value::Object(data),
+ files: None,
+ })
+ }
+
+ fn transform_ocr_response(
+ &self,
+ model: &str,
+ response_json: Value,
+ ) -> CoreResult {
+ let response_object = response_json
+ .as_object()
+ .ok_or_else(|| CoreError::InvalidType {
+ expected: "object",
+ actual: json_type_name(&response_json),
+ })?;
+
+ let pages = response_object
+ .get("pages")
+ .and_then(Value::as_array)
+ .cloned()
+ .unwrap_or_default();
+ let model = response_object
+ .get("model")
+ .and_then(Value::as_str)
+ .unwrap_or(model)
+ .to_string();
+ let document_annotation = response_object.get("document_annotation").cloned();
+ let usage_info = response_object.get("usage_info").cloned();
+
+ Ok(OcrResponseData {
+ pages,
+ model,
+ document_annotation,
+ usage_info,
+ object: "ocr".to_string(),
+ })
+ }
+}
+
+pub fn supported_ocr_params() -> &'static [&'static str] {
+ MISTRAL_OCR_CONFIG.supported_ocr_params()
+}
+
+pub fn map_ocr_params(non_default_params: &Map) -> Map {
+ MISTRAL_OCR_CONFIG.map_ocr_params(non_default_params)
+}
+
+pub fn transform_ocr_request(
+ model: &str,
+ document: Value,
+ optional_params: Map,
+) -> CoreResult {
+ MISTRAL_OCR_CONFIG.transform_ocr_request(model, document, optional_params)
+}
+
+pub fn transform_ocr_response(model: &str, response_json: Value) -> CoreResult {
+ MISTRAL_OCR_CONFIG.transform_ocr_response(model, response_json)
+}
+
+#[cfg(test)]
+mod tests {
+ use super::*;
+ use serde_json::json;
+
+ #[test]
+ fn supported_params_match_python_mistral_ocr_config() {
+ assert_eq!(
+ supported_ocr_params(),
+ &[
+ "pages",
+ "include_image_base64",
+ "image_limit",
+ "image_min_size",
+ "bbox_annotation_format",
+ "document_annotation_format",
+ "document_annotation_prompt",
+ "extract_header",
+ "extract_footer",
+ "table_format",
+ "confidence_scores_granularity",
+ "id",
+ ]
+ );
+ }
+
+ #[test]
+ fn map_ocr_params_drops_unknown_params() {
+ let params = json!({
+ "extract_header": true,
+ "unsupported_param": "value",
+ "pages": [0, 1]
+ });
+ let mapped = map_ocr_params(params.as_object().unwrap());
+
+ assert_eq!(mapped.get("extract_header"), Some(&json!(true)));
+ assert_eq!(mapped.get("pages"), Some(&json!([0, 1])));
+ assert!(!mapped.contains_key("unsupported_param"));
+ }
+
+ #[test]
+ fn transform_ocr_request_builds_mistral_body() {
+ let document = json!({
+ "type": "document_url",
+ "document_url": "https://example.com/doc.pdf"
+ });
+ let optional_params = json!({
+ "include_image_base64": true,
+ "table_format": "html"
+ })
+ .as_object()
+ .unwrap()
+ .clone();
+
+ let result = transform_ocr_request("mistral-ocr-latest", document.clone(), optional_params)
+ .expect("request should transform");
+
+ assert_eq!(
+ result.data,
+ json!({
+ "model": "mistral-ocr-latest",
+ "document": document,
+ "include_image_base64": true,
+ "table_format": "html"
+ })
+ );
+ assert_eq!(result.files, None);
+ }
+
+ #[test]
+ fn transform_ocr_request_rejects_non_object_document() {
+ let err = transform_ocr_request("mistral-ocr-latest", json!("bad"), Map::new())
+ .expect_err("string document should be rejected");
+
+ assert_eq!(
+ err,
+ CoreError::InvalidType {
+ expected: "object",
+ actual: "string",
+ }
+ );
+ }
+
+ #[test]
+ fn transform_ocr_response_normalizes_mistral_json() {
+ let response = json!({
+ "pages": [{"index": 0, "markdown": "hello"}],
+ "model": "mistral-ocr-2505-completion",
+ "document_annotation": null,
+ "usage_info": {"pages_processed": 1}
+ });
+
+ let result = transform_ocr_response("mistral-ocr-latest", response)
+ .expect("response should transform");
+
+ assert_eq!(result.pages, vec![json!({"index": 0, "markdown": "hello"})]);
+ assert_eq!(result.model, "mistral-ocr-2505-completion");
+ assert_eq!(result.document_annotation, Some(Value::Null));
+ assert_eq!(result.usage_info, Some(json!({"pages_processed": 1})));
+ assert_eq!(result.object, "ocr");
+ }
+
+ #[test]
+ fn complete_url_defaults_and_dedupes_v1() {
+ assert_eq!(complete_url(None), "https://api.mistral.ai/v1/ocr");
+ assert_eq!(complete_url(Some(" ")), "https://api.mistral.ai/v1/ocr");
+ assert_eq!(
+ complete_url(Some("https://proxy.internal")),
+ "https://proxy.internal/v1/ocr"
+ );
+ assert_eq!(
+ complete_url(Some("https://proxy.internal/v1/")),
+ "https://proxy.internal/v1/ocr"
+ );
+ }
+
+ #[test]
+ fn resolve_api_key_prefers_param_then_env() {
+ let no_env = |_: &str| None;
+ assert_eq!(
+ resolve_api_key(Some("sk-param"), &no_env).unwrap(),
+ "sk-param"
+ );
+
+ let with_env = |key: &str| (key == MISTRAL_API_KEY_ENV).then(|| "sk-env".to_string());
+ assert_eq!(resolve_api_key(None, &with_env).unwrap(), "sk-env");
+ // Blank param falls through to the environment.
+ assert_eq!(resolve_api_key(Some(" "), &with_env).unwrap(), "sk-env");
+ }
+
+ #[test]
+ fn resolve_api_key_errors_when_absent() {
+ let err = resolve_api_key(None, &|_| None).expect_err("missing key should error");
+ assert_eq!(err, CoreError::Auth(MISSING_KEY_MESSAGE.to_string()));
+ }
+}
diff --git a/litellm-rust/crates/providers/src/ocr.rs b/litellm-rust/crates/providers/src/ocr.rs
new file mode 100644
index 00000000000..dcd56a5f0b4
--- /dev/null
+++ b/litellm-rust/crates/providers/src/ocr.rs
@@ -0,0 +1,127 @@
+//! End-to-end OCR orchestration.
+//!
+//! Owns the whole Mistral OCR call so the Python side stays a thin bridge:
+//! resolve the API key, build the URL + body via the pure transforms, POST it,
+//! and normalize the response. The HTTP client is built once and reused.
+
+use std::sync::OnceLock;
+use std::time::Duration;
+
+use litellm_core::error::CoreError;
+use litellm_core::ocr::transformation::OcrProviderConfig;
+use litellm_core::CoreResult;
+use serde_json::{Map, Value};
+
+use crate::mistral::ocr::transformation as mistral;
+use crate::mistral::ocr::transformation::MISTRAL_OCR_CONFIG;
+
+/// OCR over large documents can take a while; bound it generously rather than
+/// hanging forever on an unresponsive upstream. The client-level limit is the
+/// outer ceiling; callers can tighten it per request via ``run_ocr``'s ``timeout``.
+const OCR_TIMEOUT_SECS: u64 = 600;
+
+/// Maximum upstream body characters retained in error messages. OCR responses
+/// can echo document contents and prompts; keep enough for debugging without
+/// forwarding sensitive payloads across the host boundary.
+const ERROR_BODY_MAX_CHARS: usize = 256;
+
+/// Process-wide blocking HTTP client (connection pool + TLS reused across calls).
+fn http_client() -> &'static reqwest::blocking::Client {
+ static CLIENT: OnceLock = OnceLock::new();
+ CLIENT.get_or_init(|| {
+ reqwest::blocking::Client::builder()
+ .timeout(Duration::from_secs(OCR_TIMEOUT_SECS))
+ .build()
+ .expect("failed to build reqwest client")
+ })
+}
+
+fn truncate_error_body(body: &str) -> String {
+ if body.chars().count() <= ERROR_BODY_MAX_CHARS {
+ return body.to_string();
+ }
+ let truncated: String = body.chars().take(ERROR_BODY_MAX_CHARS).collect();
+ format!("{truncated}... (truncated)")
+}
+
+/// Perform a Mistral OCR call end to end and return the normalized response as
+/// JSON (the shape the Python `OCRResponse` model expects).
+///
+/// Blocking: intended to be called with the GIL released from the Python bridge.
+pub fn run_ocr(
+ model: &str,
+ document: Value,
+ api_key: Option<&str>,
+ api_base: Option<&str>,
+ optional_params: Map,
+ timeout: Option,
+) -> CoreResult {
+ let config = &MISTRAL_OCR_CONFIG;
+
+ let api_key = mistral::resolve_api_key(api_key, &|key| std::env::var(key).ok())?;
+ let url = mistral::complete_url(api_base);
+ let filtered_params = config.map_ocr_params(&optional_params);
+ let body = config
+ .transform_ocr_request(model, document, filtered_params)?
+ .data;
+
+ let mut request = http_client().post(&url).bearer_auth(&api_key).json(&body);
+ if let Some(duration) = timeout {
+ request = request.timeout(duration);
+ }
+
+ let response = request
+ .send()
+ .map_err(|err| CoreError::Network(err.to_string()))?;
+
+ let status = response.status();
+ let text = response
+ .text()
+ .map_err(|err| CoreError::Network(err.to_string()))?;
+
+ if !status.is_success() {
+ return Err(CoreError::Http {
+ status: status.as_u16(),
+ body: truncate_error_body(&text),
+ });
+ }
+
+ let response_json: Value = serde_json::from_str(&text)
+ .map_err(|err| CoreError::InvalidResponse(format!("invalid OCR response JSON: {err}")))?;
+
+ Ok(config
+ .transform_ocr_response(model, response_json)?
+ .into_json())
+}
+
+#[cfg(test)]
+mod tests {
+ use super::*;
+
+ #[test]
+ fn truncate_error_body_passes_short_strings_through() {
+ let body = "Unauthorized";
+ assert_eq!(truncate_error_body(body), "Unauthorized");
+ }
+
+ #[test]
+ fn truncate_error_body_caps_long_payloads() {
+ let body = "x".repeat(ERROR_BODY_MAX_CHARS + 50);
+ let truncated = truncate_error_body(&body);
+
+ assert!(truncated.ends_with("... (truncated)"));
+ let prefix_chars = truncated
+ .strip_suffix("... (truncated)")
+ .expect("truncated marker present")
+ .chars()
+ .count();
+ assert_eq!(prefix_chars, ERROR_BODY_MAX_CHARS);
+ }
+
+ #[test]
+ fn truncate_error_body_does_not_split_multibyte_chars() {
+ let body = "é".repeat(ERROR_BODY_MAX_CHARS + 10);
+ let truncated = truncate_error_body(&body);
+ assert!(truncated.is_char_boundary(truncated.len()));
+ }
+}
diff --git a/litellm-rust/crates/providers/src/openai/mod.rs b/litellm-rust/crates/providers/src/openai/mod.rs
new file mode 100644
index 00000000000..403e32975cf
--- /dev/null
+++ b/litellm-rust/crates/providers/src/openai/mod.rs
@@ -0,0 +1 @@
+pub mod realtime;
diff --git a/litellm-rust/crates/providers/src/openai/realtime/mod.rs b/litellm-rust/crates/providers/src/openai/realtime/mod.rs
new file mode 100644
index 00000000000..f239b6921fa
--- /dev/null
+++ b/litellm-rust/crates/providers/src/openai/realtime/mod.rs
@@ -0,0 +1 @@
+pub mod transformation;
diff --git a/litellm-rust/crates/providers/src/openai/realtime/transformation.rs b/litellm-rust/crates/providers/src/openai/realtime/transformation.rs
new file mode 100644
index 00000000000..2e127c699e0
--- /dev/null
+++ b/litellm-rust/crates/providers/src/openai/realtime/transformation.rs
@@ -0,0 +1,189 @@
+use litellm_core::realtime::transformation::RealtimeProviderConfig;
+use litellm_core::realtime::types::{RealtimeEvent, RealtimeTransformResult};
+use litellm_core::CoreResult;
+
+/// Default OpenAI API base, used when the caller does not override `api_base`.
+pub const OPENAI_REALTIME_DEFAULT_API_BASE: &str = "https://api.openai.com";
+
+/// Path appended to the resolved host base to reach the realtime endpoint.
+pub const OPENAI_REALTIME_PATH: &str = "/v1/realtime";
+
+/// Percent-encode a query value, escaping any char outside the RFC 3986
+/// unreserved set (`A-Za-z0-9-._~`). Keeps us dependency-free; common realtime
+/// model slugs have no special chars, but this stays correct for the rest.
+fn percent_encode(value: &str) -> String {
+ let mut encoded = String::with_capacity(value.len());
+ for byte in value.bytes() {
+ let unreserved = byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'.' | b'_' | b'~');
+ if unreserved {
+ encoded.push(byte as char);
+ } else {
+ encoded.push('%');
+ encoded.push_str(&format!("{byte:02X}"));
+ }
+ }
+ encoded
+}
+
+/// Build the realtime WebSocket URL, porting Python's `OpenAIRealtime._construct_url`.
+///
+/// Blank/whitespace `api_base` is treated as absent (guard at resolution time),
+/// falling back to the default. The scheme is swapped to its WebSocket
+/// equivalent (`https://`→`wss://`, `http://`→`ws://`); bases already using
+/// `ws`/`wss` are left untouched. A bare host or unrecognized scheme defaults to
+/// secure `wss://` so we never hand a scheme-less URL to the connector (this is
+/// a deliberate hardening over Python's `_construct_url`, which would emit a
+/// scheme-less URL here). A trailing `/` is trimmed before the path and
+/// `?model=` are appended.
+pub fn complete_url(api_base: Option<&str>, model: &str) -> String {
+ let base = api_base
+ .map(str::trim)
+ .filter(|base| !base.is_empty())
+ .unwrap_or(OPENAI_REALTIME_DEFAULT_API_BASE);
+
+ let base = if let Some(rest) = base.strip_prefix("https://") {
+ format!("wss://{rest}")
+ } else if let Some(rest) = base.strip_prefix("http://") {
+ format!("ws://{rest}")
+ } else if base.starts_with("wss://") || base.starts_with("ws://") {
+ base.to_string()
+ } else {
+ format!("wss://{base}")
+ };
+
+ let base = base.trim_end_matches('/');
+
+ format!(
+ "{base}{OPENAI_REALTIME_PATH}?model={}",
+ percent_encode(model)
+ )
+}
+
+pub struct OpenAiRealtimeConfig;
+
+pub const OPENAI_REALTIME_CONFIG: OpenAiRealtimeConfig = OpenAiRealtimeConfig;
+
+impl RealtimeProviderConfig for OpenAiRealtimeConfig {
+ fn complete_url(&self, api_base: Option<&str>, model: &str) -> String {
+ complete_url(api_base, model)
+ }
+
+ fn transform_realtime_request(
+ &self,
+ event: &RealtimeEvent,
+ _model: &str,
+ ) -> CoreResult {
+ Ok(RealtimeTransformResult::passthrough(event.clone()))
+ }
+
+ fn transform_realtime_response(
+ &self,
+ event: &RealtimeEvent,
+ _model: &str,
+ ) -> CoreResult {
+ Ok(RealtimeTransformResult::passthrough(event.clone()))
+ }
+}
+
+pub fn transform_realtime_request(
+ event: &RealtimeEvent,
+ model: &str,
+) -> CoreResult {
+ OPENAI_REALTIME_CONFIG.transform_realtime_request(event, model)
+}
+
+pub fn transform_realtime_response(
+ event: &RealtimeEvent,
+ model: &str,
+) -> CoreResult {
+ OPENAI_REALTIME_CONFIG.transform_realtime_response(event, model)
+}
+
+#[cfg(test)]
+mod tests {
+ use super::*;
+
+ #[test]
+ fn complete_url_defaults_to_openai_wss() {
+ assert_eq!(
+ complete_url(None, "gpt-4o-realtime-preview"),
+ "wss://api.openai.com/v1/realtime?model=gpt-4o-realtime-preview"
+ );
+ }
+
+ #[test]
+ fn complete_url_blank_base_uses_default() {
+ assert_eq!(
+ complete_url(Some(" "), "gpt-4o-realtime-preview"),
+ "wss://api.openai.com/v1/realtime?model=gpt-4o-realtime-preview"
+ );
+ }
+
+ #[test]
+ fn complete_url_swaps_http_to_ws() {
+ assert_eq!(
+ complete_url(Some("http://localhost:8080"), "gpt-4o-realtime-preview"),
+ "ws://localhost:8080/v1/realtime?model=gpt-4o-realtime-preview"
+ );
+ }
+
+ #[test]
+ fn complete_url_dedupes_trailing_slash() {
+ assert_eq!(
+ complete_url(Some("https://api.openai.com/"), "gpt-4o-realtime-preview"),
+ "wss://api.openai.com/v1/realtime?model=gpt-4o-realtime-preview"
+ );
+ }
+
+ #[test]
+ fn complete_url_custom_base() {
+ assert_eq!(
+ complete_url(Some("https://oai.azure.example"), "gpt-4o-realtime-preview"),
+ "wss://oai.azure.example/v1/realtime?model=gpt-4o-realtime-preview"
+ );
+ }
+
+ #[test]
+ fn complete_url_preserves_existing_wss_scheme() {
+ assert_eq!(
+ complete_url(Some("wss://api.openai.com"), "gpt-realtime"),
+ "wss://api.openai.com/v1/realtime?model=gpt-realtime"
+ );
+ }
+
+ #[test]
+ fn complete_url_bare_host_defaults_to_wss() {
+ assert_eq!(
+ complete_url(Some("api.openai.com"), "gpt-realtime"),
+ "wss://api.openai.com/v1/realtime?model=gpt-realtime"
+ );
+ }
+
+ #[test]
+ fn complete_url_percent_encodes_model_space() {
+ assert_eq!(
+ complete_url(None, "gpt 4o"),
+ "wss://api.openai.com/v1/realtime?model=gpt%204o"
+ );
+ }
+
+ #[test]
+ fn transform_realtime_request_passthrough_preserves_event() {
+ let event: RealtimeEvent =
+ serde_json::from_str(r#"{"type":"session.update","session":{"voice":"alloy"}}"#)
+ .expect("valid event");
+ let result =
+ transform_realtime_request(&event, "gpt-realtime").expect("passthrough is infallible");
+ assert_eq!(result.events, vec![event]);
+ }
+
+ #[test]
+ fn transform_realtime_response_passthrough_preserves_event() {
+ let event: RealtimeEvent =
+ serde_json::from_str(r#"{"type":"response.output_audio.delta","delta":"abc=="}"#)
+ .expect("valid event");
+ let result =
+ transform_realtime_response(&event, "gpt-realtime").expect("passthrough is infallible");
+ assert_eq!(result.events, vec![event]);
+ }
+}
diff --git a/litellm-rust/crates/providers/src/realtime.rs b/litellm-rust/crates/providers/src/realtime.rs
new file mode 100644
index 00000000000..398158f6dba
--- /dev/null
+++ b/litellm-rust/crates/providers/src/realtime.rs
@@ -0,0 +1,374 @@
+//! End-to-end OpenAI realtime invocation.
+//!
+//! The host-facing entry point, mirroring `providers::ocr::run_ocr`: open the
+//! WebSocket to OpenAI, then splice a client realtime stream to the upstream,
+//! driving typed events through the pure `OPENAI_REALTIME_CONFIG` transforms.
+//! Network, auth header, key resolution, and wire (de)serialization live here so
+//! the `transformation` module stays pure and typed.
+//!
+//! The dial and splice steps are factored out ([`dial_upstream`], [`splice`]) so
+//! the connection pool ([`crate::realtime_pool`]) can pre-establish an upstream,
+//! buffer its `session.created`, and later hand the live socket to the same
+//! splice loop a fresh dial uses.
+
+use std::time::Duration;
+
+use futures_util::stream::{SplitSink, SplitStream};
+use futures_util::{Sink, SinkExt, Stream, StreamExt};
+use litellm_core::error::CoreError;
+use litellm_core::realtime::transformation::RealtimeProviderConfig;
+use litellm_core::realtime::types::RealtimeEvent;
+use litellm_core::CoreResult;
+use tokio::net::TcpStream;
+use tokio_tungstenite::tungstenite::client::IntoClientRequest;
+use tokio_tungstenite::tungstenite::http::header::AUTHORIZATION;
+use tokio_tungstenite::tungstenite::http::HeaderValue;
+use tokio_tungstenite::tungstenite::Message;
+use tokio_tungstenite::{connect_async, MaybeTlsStream, WebSocketStream};
+
+use crate::openai::realtime::transformation::OPENAI_REALTIME_CONFIG;
+
+/// Environment variable holding the OpenAI API key (last-resort fallback).
+const OPENAI_API_KEY_ENV: &str = "OPENAI_API_KEY";
+
+const MISSING_KEY_MESSAGE: &str = "Missing OpenAI API Key - a realtime call is being made but no key was passed via params or the OPENAI_API_KEY environment variable";
+
+/// Default **idle** timeout: if neither side sends a frame for this long, the
+/// session is reaped. It resets on any activity, so it does not cap a healthy
+/// (continuously streaming) session — it only frees a stalled one (e.g. a
+/// half-open upstream that keeps the socket open but stops sending).
+const DEFAULT_IDLE_TIMEOUT_SECS: u64 = 300;
+
+/// The concrete upstream WebSocket type (TLS or plain). Shared by the dial path
+/// and the pool so warm sockets and fresh sockets are the exact same type.
+pub type UpstreamWs = WebSocketStream>;
+pub(crate) type UpstreamTx = SplitSink;
+pub(crate) type UpstreamRx = SplitStream;
+
+/// Resolve the OpenAI API key from the explicit param or the environment.
+///
+/// Blank/whitespace values are treated as absent (guard at resolution time).
+pub(crate) fn resolve_api_key(api_key: Option<&str>) -> CoreResult {
+ api_key
+ .map(str::trim)
+ .filter(|key| !key.is_empty())
+ .map(str::to_string)
+ .or_else(|| {
+ std::env::var(OPENAI_API_KEY_ENV)
+ .ok()
+ .filter(|key| !key.trim().is_empty())
+ })
+ .ok_or_else(|| CoreError::Auth(MISSING_KEY_MESSAGE.to_string()))
+}
+
+/// Open the upstream WebSocket to OpenAI for `(model, api_key, api_base)`.
+///
+/// This is the dial half of [`realtime`], factored out so the pool can
+/// pre-establish sockets ahead of any client. `api_key` here is already resolved
+/// (non-blank) — the pool resolves it once when it is created.
+pub(crate) async fn dial_upstream(
+ model: &str,
+ api_key: &str,
+ api_base: Option<&str>,
+) -> CoreResult {
+ let url = OPENAI_REALTIME_CONFIG.complete_url(api_base, model);
+
+ let mut request = url
+ .as_str()
+ .into_client_request()
+ .map_err(|err| CoreError::Network(err.to_string()))?;
+ // GA realtime: only Authorization. The legacy OpenAI-Beta header triggers
+ // beta_api_shape_disabled, so we do not send it.
+ request.headers_mut().insert(
+ AUTHORIZATION,
+ HeaderValue::from_str(&format!("Bearer {api_key}"))
+ .map_err(|err| CoreError::Auth(err.to_string()))?,
+ );
+
+ let (upstream, _response) = connect_async(request)
+ .await
+ .map_err(|err| CoreError::Network(err.to_string()))?;
+ Ok(upstream)
+}
+
+/// Read the next text frame from the upstream and decode it as a typed event.
+///
+/// Used by the pool to pre-read OpenAI's unprompted `session.created`. Returns an
+/// error on a non-text frame, a closed socket, or undecodable JSON so the pool can
+/// discard a misbehaving socket rather than warm it.
+pub(crate) async fn read_event(upstream_rx: &mut UpstreamRx) -> CoreResult {
+ loop {
+ let message = upstream_rx
+ .next()
+ .await
+ .ok_or_else(|| CoreError::Network("upstream closed before first event".to_string()))?
+ .map_err(|err| CoreError::Network(err.to_string()))?;
+ match message {
+ Message::Text(text) => {
+ return serde_json::from_str(&text)
+ .map_err(|err| CoreError::InvalidResponse(err.to_string()));
+ }
+ // Ignore protocol frames (ping/pong) while waiting for the first event.
+ Message::Ping(_) | Message::Pong(_) => continue,
+ Message::Close(_) => {
+ return Err(CoreError::Network(
+ "upstream closed before first event".to_string(),
+ ))
+ }
+ _ => continue,
+ }
+ }
+}
+
+/// Splice an already-connected upstream to the client streams.
+///
+/// `prelude` is relayed to the client first (the pool passes the buffered
+/// `session.created` here; the fresh-dial path passes `None` and lets the upstream
+/// deliver it). Then a single select loop forwards both directions through the
+/// transforms until either side closes or the idle timeout fires.
+#[allow(clippy::too_many_arguments)]
+pub(crate) async fn splice(
+ model: &str,
+ mut upstream_tx: UpstreamTx,
+ mut upstream_rx: UpstreamRx,
+ prelude: Option,
+ idle_timeout: Option,
+ mut client_in: In,
+ mut client_out: Out,
+) -> CoreResult<()>
+where
+ In: Stream
- + Unpin + Send,
+ Out: Sink + Unpin + Send,
+ >::Error: std::fmt::Display,
+{
+ let config = &OPENAI_REALTIME_CONFIG;
+
+ // Relay a buffered backend event (warm handoff's session.created) first, so a
+ // warm session looks identical to a fresh one from the client's view.
+ if let Some(event) = prelude {
+ for outbound in config.transform_realtime_response(&event, model)?.events {
+ client_out
+ .send(outbound)
+ .await
+ .map_err(|err| CoreError::Network(err.to_string()))?;
+ }
+ }
+
+ let idle = idle_timeout.unwrap_or(Duration::from_secs(DEFAULT_IDLE_TIMEOUT_SECS));
+
+ // One loop forwarding both directions. The `sleep(idle)` arm is rebuilt every
+ // iteration, so any frame (either way) resets it — it fires only when the
+ // session has been fully idle for `idle`, reaping a stalled connection
+ // (task + upstream TCP socket) instead of leaking it.
+ loop {
+ tokio::select! {
+ // client -> upstream
+ client_event = client_in.next() => {
+ let Some(event) = client_event else { break }; // client disconnected
+ for outbound in config.transform_realtime_request(&event, model)?.events {
+ let payload = serde_json::to_string(&outbound)
+ .map_err(|err| CoreError::InvalidResponse(err.to_string()))?;
+ upstream_tx
+ .send(Message::Text(payload))
+ .await
+ .map_err(|err| CoreError::Network(err.to_string()))?;
+ }
+ }
+ // upstream -> client
+ upstream_message = upstream_rx.next() => {
+ let Some(message) = upstream_message else { break }; // upstream closed
+ match message.map_err(|err| CoreError::Network(err.to_string()))? {
+ Message::Text(text) => {
+ let event: RealtimeEvent = serde_json::from_str(&text)
+ .map_err(|err| CoreError::InvalidResponse(err.to_string()))?;
+ for outbound in config.transform_realtime_response(&event, model)?.events {
+ client_out
+ .send(outbound)
+ .await
+ .map_err(|err| CoreError::Network(err.to_string()))?;
+ }
+ }
+ Message::Close(_) => break,
+ _ => {}
+ }
+ }
+ // idle timeout: no activity from either side within `idle`
+ _ = tokio::time::sleep(idle) => break,
+ }
+ }
+ Ok(())
+}
+
+/// Splice a client realtime stream to OpenAI: forward client events upstream
+/// (via `transform_realtime_request`) and backend events downstream (via
+/// `transform_realtime_response`). Returns when either side closes.
+///
+/// Generic over the client transport (typed events) so this crate stays
+/// framework-agnostic; the gateway adapts its axum socket to these. This is the
+/// fresh-dial path: dial, then splice. The pool's warm-handoff path skips the dial
+/// and calls [`splice`] directly with a buffered `session.created`.
+pub async fn realtime(
+ model: &str,
+ api_key: Option<&str>,
+ api_base: Option<&str>,
+ idle_timeout: Option,
+ client_in: In,
+ client_out: Out,
+) -> CoreResult<()>
+where
+ In: Stream
- + Unpin + Send,
+ Out: Sink + Unpin + Send,
+ >::Error: std::fmt::Display,
+{
+ let api_key = resolve_api_key(api_key)?;
+ let upstream = dial_upstream(model, &api_key, api_base).await?;
+ let (upstream_tx, upstream_rx) = upstream.split();
+ splice(
+ model,
+ upstream_tx,
+ upstream_rx,
+ None,
+ idle_timeout,
+ client_in,
+ client_out,
+ )
+ .await
+}
+
+/// Splice a pre-warmed upstream (taken from [`crate::realtime_pool`]) to the
+/// client. Relays the buffered `session.created` first, then splices exactly like
+/// the fresh-dial path — so a warm session is indistinguishable from a fresh one.
+pub async fn realtime_warm(
+ model: &str,
+ handoff: crate::realtime_pool::WarmHandoff,
+ idle_timeout: Option,
+ client_in: In,
+ client_out: Out,
+) -> CoreResult<()>
+where
+ In: Stream
- + Unpin + Send,
+ Out: Sink + Unpin + Send,
+ >::Error: std::fmt::Display,
+{
+ splice(
+ model,
+ handoff.tx,
+ handoff.rx,
+ Some(handoff.session_created),
+ idle_timeout,
+ client_in,
+ client_out,
+ )
+ .await
+}
+
+#[cfg(test)]
+mod tests {
+ use super::*;
+
+ fn event(raw: &str) -> RealtimeEvent {
+ serde_json::from_str(raw).expect("valid event json")
+ }
+
+ #[test]
+ fn resolve_api_key_prefers_param_then_blank_falls_through() {
+ assert_eq!(resolve_api_key(Some("sk-test")).unwrap(), "sk-test");
+ // A blank param with no env set should error.
+ if std::env::var(OPENAI_API_KEY_ENV).is_err() {
+ assert!(resolve_api_key(Some(" ")).is_err());
+ }
+ }
+
+ /// Live end-to-end check against OpenAI. Ignored by default (CI never runs
+ /// it); run explicitly with `OPENAI_API_KEY` set:
+ /// `cargo test -p litellm-providers realtime_invokes_openai -- --ignored --nocapture`
+ #[tokio::test]
+ #[ignore = "hits the live OpenAI realtime API; needs OPENAI_API_KEY"]
+ async fn realtime_invokes_openai_and_responds() {
+ use futures_channel::mpsc;
+
+ let key =
+ std::env::var(OPENAI_API_KEY_ENV).expect("set OPENAI_API_KEY to run this ignored test");
+
+ // client -> provider (we hold `client_tx` to push events upstream)
+ let (mut client_tx, client_in) = mpsc::unbounded::();
+ // provider -> client (we hold `backend_rx` to read backend events)
+ let (client_out, mut backend_rx) = mpsc::unbounded::();
+
+ // Clone the key so the spawned task owns its `String` (no borrow across await).
+ let key_owned = key.clone();
+ let call = tokio::spawn(async move {
+ realtime(
+ "gpt-realtime",
+ Some(&key_owned),
+ None,
+ None,
+ client_in,
+ client_out,
+ )
+ .await
+ });
+
+ // 1. First backend event should be session.created.
+ let first = tokio::time::timeout(Duration::from_secs(30), backend_rx.next())
+ .await
+ .expect("timed out waiting for session.created")
+ .expect("backend stream closed before session.created");
+ assert_eq!(
+ first.event_type, "session.created",
+ "expected session.created, got: {}",
+ first.event_type
+ );
+
+ // 2. Ask for a short audio response.
+ client_tx
+ .send(event(
+ r#"{"type":"conversation.item.create","item":{"type":"message","role":"user","content":[{"type":"input_text","text":"Say hi."}]}}"#,
+ ))
+ .await
+ .expect("send conversation.item.create");
+ client_tx
+ .send(event(r#"{"type":"response.create"}"#))
+ .await
+ .expect("send response.create");
+
+ // 3. Read backend events; require a non-empty audio delta, then response.done.
+ let mut saw_audio_delta = false;
+ let mut saw_done = false;
+ for _ in 0..500 {
+ let next = tokio::time::timeout(Duration::from_secs(30), backend_rx.next()).await;
+ let event = match next {
+ Ok(Some(event)) => event,
+ Ok(None) => break,
+ Err(_) => panic!("timed out waiting for backend events"),
+ };
+ match event.event_type.as_str() {
+ "response.output_audio.delta" => {
+ let delta = event
+ .data
+ .get("delta")
+ .and_then(|value| value.as_str())
+ .unwrap_or("");
+ if !delta.is_empty() {
+ saw_audio_delta = true;
+ }
+ }
+ "response.done" => {
+ saw_done = true;
+ break;
+ }
+ _ => {}
+ }
+ }
+
+ assert!(
+ saw_audio_delta,
+ "expected a response.output_audio.delta with non-empty delta"
+ );
+ assert!(saw_done, "expected a response.done event");
+
+ // Drop the client sender so the provider's to_upstream side finishes.
+ drop(client_tx);
+ let _ = call.await;
+ }
+}
diff --git a/litellm-rust/crates/providers/src/realtime_pool.rs b/litellm-rust/crates/providers/src/realtime_pool.rs
new file mode 100644
index 00000000000..1b1fc8112c5
--- /dev/null
+++ b/litellm-rust/crates/providers/src/realtime_pool.rs
@@ -0,0 +1,712 @@
+//! Pre-warmed upstream realtime connection pool.
+//!
+//! The gateway's realtime overhead lives entirely in session establishment: on
+//! every client connect it dials a fresh upstream WS to OpenAI and waits for
+//! `session.created` before it can serve. This pool keeps a small set of upstream
+//! sockets **already connected and already past `session.created`** so a connect
+//! can be served from a warm socket and the handshake is off the critical path.
+//!
+//! Layering: this stays in `providers` (axum-free) next to the dial/splice it
+//! reuses. The gateway holds an `Arc` in its state and asks for a
+//! warm socket per connect; on a miss it fresh-dials exactly as before. The pool
+//! is a latency optimization, never a correctness dependency — see the gateway's
+//! `src/routes/realtime/README.md`.
+//!
+//! ## Caveats (enforced here)
+//! - One warm socket serves exactly one session (realtime isn't multiplexed), so
+//! the pool is sized to the connect *rate*, not concurrent connections.
+//! - `session.created` is pre-read once and buffered; nothing else is read from a
+//! warm socket before handoff, so a warm session starts at OpenAI defaults just
+//! like a fresh one (`session.update` semantics unchanged).
+//! - Warm sockets are short-lived (`max_idle`) and liveness-checked at handoff to
+//! bound idle billing / dodge OpenAI's idle timeout.
+//! - On miss or dead socket the caller fresh-dials; the pool never blocks or fails
+//! a connect because it is empty.
+
+use std::collections::HashMap;
+use std::sync::{Arc, Mutex};
+use std::time::{Duration, Instant};
+
+use futures_util::StreamExt;
+use litellm_core::realtime::types::RealtimeEvent;
+use litellm_core::CoreResult;
+
+use crate::realtime::{
+ dial_upstream, read_event, resolve_api_key, UpstreamRx, UpstreamTx, UpstreamWs,
+};
+
+/// Default target warm sockets per key when pooling is enabled.
+pub const DEFAULT_POOL_SIZE: usize = 4;
+
+/// Default max time a warm socket may sit before it is closed and replaced.
+pub const DEFAULT_MAX_IDLE: Duration = Duration::from_secs(30);
+
+/// Env var: target warm sockets per key. `0` disables pooling (fresh-dial only).
+pub const POOL_SIZE_ENV: &str = "REALTIME_POOL_SIZE";
+
+/// Env var: max warm-socket idle lifetime, in seconds.
+pub const MAX_IDLE_ENV: &str = "REALTIME_POOL_MAX_IDLE_SECS";
+
+/// How often the background replenisher wakes to top up and reap stale sockets.
+const REPLENISH_TICK: Duration = Duration::from_millis(250);
+
+/// Backoff floor after a key's warm-up dials all fail. The first failed pass
+/// waits this long before retrying that key.
+const BACKOFF_BASE: Duration = Duration::from_millis(500);
+
+/// Backoff ceiling. A key that keeps failing (invalid credentials, an
+/// unreachable upstream) is retried at most once per this interval — instead of
+/// firing `needed` concurrent TLS dials every 250 ms tick, which would hammer
+/// the upstream and risk rate-limit exhaustion that degrades valid cold-path
+/// traffic. Backoff resets the moment a dial for the key succeeds.
+const BACKOFF_MAX: Duration = Duration::from_secs(30);
+
+/// Identifies an upstream connection: the tuple that fully determines the dial.
+/// `api_key` is included so a warm socket is only ever reused for the same key
+/// (no cross-tenant reuse).
+#[derive(Clone, PartialEq, Eq, Hash)]
+pub struct UpstreamKey {
+ pub model: String,
+ pub api_key: String,
+ pub api_base: Option,
+}
+
+impl std::fmt::Debug for UpstreamKey {
+ fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
+ f.debug_struct("UpstreamKey")
+ .field("model", &self.model)
+ .field("api_key", &"[REDACTED]")
+ .field("api_base", &self.api_base)
+ .finish()
+ }
+}
+
+/// A warm upstream: split halves + the buffered `session.created` + when it was
+/// warmed (for `max_idle` expiry).
+struct WarmConnection {
+ tx: UpstreamTx,
+ rx: UpstreamRx,
+ session_created: RealtimeEvent,
+ warmed_at: Instant,
+}
+
+/// A live upstream taken from the pool, ready to splice. The caller relays
+/// `session_created` to the client first, then splices `(tx, rx)` as usual.
+pub struct WarmHandoff {
+ pub tx: UpstreamTx,
+ pub rx: UpstreamRx,
+ pub session_created: RealtimeEvent,
+}
+
+/// Pool configuration, resolved once at startup from the environment.
+#[derive(Clone, Copy, Debug)]
+pub struct PoolConfig {
+ /// Target warm sockets per key. `0` disables pooling.
+ pub target_size: usize,
+ /// Max time a warm socket may sit before it is closed and replaced.
+ pub max_idle: Duration,
+}
+
+impl Default for PoolConfig {
+ fn default() -> Self {
+ Self {
+ target_size: DEFAULT_POOL_SIZE,
+ max_idle: DEFAULT_MAX_IDLE,
+ }
+ }
+}
+
+impl PoolConfig {
+ /// Read config from the environment, falling back to defaults. An invalid
+ /// value warns and uses the default rather than failing startup.
+ pub fn from_env() -> Self {
+ let target_size = match std::env::var(POOL_SIZE_ENV) {
+ Ok(raw) => raw.trim().parse().unwrap_or_else(|_| {
+ eprintln!("warning: {POOL_SIZE_ENV}={raw:?} is not a valid size; using {DEFAULT_POOL_SIZE}");
+ DEFAULT_POOL_SIZE
+ }),
+ Err(_) => DEFAULT_POOL_SIZE,
+ };
+ let max_idle = match std::env::var(MAX_IDLE_ENV) {
+ Ok(raw) => raw
+ .trim()
+ .parse()
+ .map(Duration::from_secs)
+ .unwrap_or_else(|_| {
+ eprintln!(
+ "warning: {MAX_IDLE_ENV}={raw:?} is not a valid number of seconds; using {}s",
+ DEFAULT_MAX_IDLE.as_secs()
+ );
+ DEFAULT_MAX_IDLE
+ }),
+ Err(_) => DEFAULT_MAX_IDLE,
+ };
+ Self {
+ target_size,
+ max_idle,
+ }
+ }
+
+ /// Whether pooling is on (`target_size > 0`).
+ pub fn enabled(&self) -> bool {
+ self.target_size > 0
+ }
+}
+
+/// Per-key warm sockets, behind a single `Mutex`. Realtime warm sockets are few
+/// (the pool is small), so a plain mutex over a `VecDeque`-ish `Vec` is simpler
+/// and faster than sharding; contention is negligible at this scale.
+type Warm = HashMap>;
+
+/// Per-key replenish backoff. Absent (or `consecutive_failures == 0`) means the
+/// key is healthy and replenished every tick. After a pass whose dials all fail,
+/// `retry_after` is pushed out with exponential backoff so a broken key (invalid
+/// credentials, unreachable upstream) is not re-dialed on every 250 ms tick.
+#[derive(Default)]
+struct Backoff {
+ /// Don't attempt warm-up dials for this key until this instant. `None` =
+ /// eligible now.
+ retry_after: Option,
+ consecutive_failures: u32,
+}
+
+type Backoffs = HashMap;
+
+/// Pre-warmed upstream realtime connection pool.
+///
+/// Cheap to clone-via-`Arc`. The background replenisher is spawned by
+/// [`RealtimePool::spawn`]; a pool built with [`RealtimePool::disabled`] never
+/// warms anything and every `take` misses (callers fresh-dial).
+pub struct RealtimePool {
+ config: PoolConfig,
+ warm: Mutex,
+ /// Per-key replenish backoff so a broken key doesn't trigger unbounded
+ /// concurrent dials every tick. Separate lock from `warm` so the request
+ /// hot path (`take`) never contends on it.
+ backoff: Mutex,
+}
+
+impl RealtimePool {
+ /// A disabled pool: no background task, every `take` returns `None`.
+ pub fn disabled() -> Arc {
+ Arc::new(Self {
+ config: PoolConfig {
+ target_size: 0,
+ ..PoolConfig::default()
+ },
+ warm: Mutex::new(HashMap::new()),
+ backoff: Mutex::new(HashMap::new()),
+ })
+ }
+
+ /// Build a pool from config **without** the background replenisher. The pool
+ /// only warms when [`RealtimePool::warm_now`] is called. Used by deterministic
+ /// unit tests; production uses [`RealtimePool::spawn`].
+ #[cfg(test)]
+ fn new_unspawned(config: PoolConfig) -> Arc {
+ Arc::new(Self {
+ config,
+ warm: Mutex::new(HashMap::new()),
+ backoff: Mutex::new(HashMap::new()),
+ })
+ }
+
+ /// Build a pool from config and, if enabled, spawn the background replenisher.
+ /// Returns the shared handle the gateway stores in its state.
+ pub fn spawn(config: PoolConfig) -> Arc {
+ let pool = Arc::new(Self {
+ config,
+ warm: Mutex::new(HashMap::new()),
+ backoff: Mutex::new(HashMap::new()),
+ });
+ if config.enabled() {
+ let weak = Arc::downgrade(&pool);
+ tokio::spawn(async move {
+ let mut tick = tokio::time::interval(REPLENISH_TICK);
+ tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
+ loop {
+ tick.tick().await;
+ // Stop once the gateway has dropped its handle.
+ let Some(pool) = weak.upgrade() else { break };
+ pool.replenish_all().await;
+ }
+ });
+ }
+ pool
+ }
+
+ /// Resolved config (test/inspection).
+ pub fn config(&self) -> PoolConfig {
+ self.config
+ }
+
+ /// Register a key so the replenisher starts warming it. Idempotent. The
+ /// gateway calls this once per known deployment at startup; the pool only
+ /// warms keys it has seen, so it never dials a model nobody asked for.
+ pub fn register(&self, key: UpstreamKey) {
+ if !self.config.enabled() {
+ return;
+ }
+ self.warm.lock().unwrap().entry(key).or_default();
+ }
+
+ /// Take a warm, live socket for `key`, or `None` on miss / dead socket.
+ ///
+ /// Pops the freshest non-expired socket and liveness-checks it; a socket that
+ /// is too old or already dead is dropped (closing it) and the next candidate
+ /// tried. Never blocks: if nothing warm is live, returns `None` so the caller
+ /// fresh-dials.
+ pub fn take(&self, key: &UpstreamKey) -> Option {
+ if !self.config.enabled() {
+ return None;
+ }
+ loop {
+ let mut candidate = {
+ let mut warm = self.warm.lock().unwrap();
+ let bucket = warm.get_mut(key)?;
+ bucket.pop()?
+ };
+ // Discard sockets past their warm lifetime (idle-billing guard).
+ if candidate.warmed_at.elapsed() > self.config.max_idle {
+ continue; // drops `candidate`, closing the socket
+ }
+ // Liveness: a non-blocking check that the socket hasn't already
+ // delivered a Close/Err. A warm socket should be silent after
+ // session.created, so anything pending means it is unhealthy.
+ if is_dead(&mut candidate.rx) {
+ continue;
+ }
+ return Some(WarmHandoff {
+ tx: candidate.tx,
+ rx: candidate.rx,
+ session_created: candidate.session_created,
+ });
+ }
+ }
+
+ /// One replenish pass over every registered key: reap stale sockets, then
+ /// dial up to `target_size`. Dials run concurrently; failures are swallowed
+ /// (a key that can't be warmed just keeps fresh-dialing on the request path)
+ /// and put the key into exponential backoff so a broken key isn't re-dialed
+ /// on every tick.
+ async fn replenish_all(&self) {
+ let keys: Vec = { self.warm.lock().unwrap().keys().cloned().collect() };
+ for key in keys {
+ self.reap_stale(&key);
+ // Skip keys still in backoff from a prior all-failed pass — this is
+ // what bounds dials against an invalid/unreachable key to once per
+ // `BACKOFF_MAX` instead of `needed` dials every 250 ms tick.
+ if self.in_backoff(&key) {
+ continue;
+ }
+ let needed = {
+ let warm = self.warm.lock().unwrap();
+ let have = warm.get(&key).map(Vec::len).unwrap_or(0);
+ self.config.target_size.saturating_sub(have)
+ };
+ if needed == 0 {
+ continue;
+ }
+ // Dial the missing sockets CONCURRENTLY. A sequential loop here makes
+ // a full refill cost `needed × handshake` (~needed × 350 ms), which
+ // can't keep up with a high connect rate — the pool drains faster
+ // than it refills and most connects miss. Firing the dials together
+ // refills in ~one handshake window, keeping warm supply ≈ peak
+ // concurrent connects so the sub-ms warm handoff becomes the median,
+ // not the lucky-hit tail.
+ let dials = (0..needed).map(|_| warm_one(&key));
+ let results = futures_util::future::join_all(dials).await;
+ let mut any_ok = false;
+ // `.flatten()` keeps only the successful dials; a key that can't be
+ // warmed just keeps fresh-dialing on the request path.
+ for conn in results.into_iter().flatten() {
+ any_ok = true;
+ self.warm
+ .lock()
+ .unwrap()
+ .entry(key.clone())
+ .or_default()
+ .push(conn);
+ }
+ // Reset backoff on any success; otherwise grow it. We only ever enter
+ // backoff when a pass that *attempted* dials produced none — a `needed
+ // == 0` pass is handled by the `continue` above and never touches it.
+ self.record_replenish_outcome(&key, any_ok);
+ }
+ }
+
+ /// Whether `key` is currently in a backoff window (a prior pass failed and
+ /// the retry time hasn't arrived). Eligible keys are pruned from the backoff
+ /// map so it doesn't grow unbounded for healthy keys.
+ fn in_backoff(&self, key: &UpstreamKey) -> bool {
+ let mut backoff = self.backoff.lock().unwrap();
+ match backoff.get(key).and_then(|b| b.retry_after) {
+ Some(retry_after) if Instant::now() < retry_after => true,
+ Some(_) => {
+ // Window elapsed — allow the attempt. Keep the failure count so a
+ // still-broken key backs off further, but clear the gate so this
+ // tick proceeds.
+ if let Some(b) = backoff.get_mut(key) {
+ b.retry_after = None;
+ }
+ false
+ }
+ None => false,
+ }
+ }
+
+ /// Update a key's backoff after a replenish attempt. Success clears it;
+ /// failure grows the retry delay exponentially up to `BACKOFF_MAX`.
+ fn record_replenish_outcome(&self, key: &UpstreamKey, any_ok: bool) {
+ let mut backoff = self.backoff.lock().unwrap();
+ if any_ok {
+ backoff.remove(key);
+ return;
+ }
+ let entry = backoff.entry(key.clone()).or_default();
+ entry.consecutive_failures = entry.consecutive_failures.saturating_add(1);
+ // Exponential: BASE * 2^(failures-1), saturating at MAX. `min` of the
+ // shift exponent keeps the doubling from overflowing.
+ let shift = (entry.consecutive_failures - 1).min(16);
+ let delay = BACKOFF_BASE.saturating_mul(1u32 << shift).min(BACKOFF_MAX);
+ entry.retry_after = Some(Instant::now() + delay);
+ }
+
+ /// Drop sockets past `max_idle` or already dead for a key.
+ fn reap_stale(&self, key: &UpstreamKey) {
+ let mut warm = self.warm.lock().unwrap();
+ if let Some(bucket) = warm.get_mut(key) {
+ bucket.retain_mut(|conn| {
+ conn.warmed_at.elapsed() <= self.config.max_idle && !is_dead(&mut conn.rx)
+ });
+ }
+ }
+
+ /// Test/inspection: number of warm sockets currently held for `key`.
+ #[cfg(test)]
+ pub fn warm_len(&self, key: &UpstreamKey) -> usize {
+ self.warm
+ .lock()
+ .unwrap()
+ .get(key)
+ .map(Vec::len)
+ .unwrap_or(0)
+ }
+
+ /// Test/inspection: consecutive replenish failures recorded for `key` (0 if
+ /// the key is healthy / has no backoff entry).
+ #[cfg(test)]
+ pub fn backoff_failures(&self, key: &UpstreamKey) -> u32 {
+ self.backoff
+ .lock()
+ .unwrap()
+ .get(key)
+ .map(|b| b.consecutive_failures)
+ .unwrap_or(0)
+ }
+
+ /// Test helper: synchronously warm `target_size` sockets for `key` (no
+ /// background task). Lets tests assert handoff behavior deterministically.
+ #[cfg(test)]
+ pub async fn warm_now(&self, key: &UpstreamKey) {
+ let needed = {
+ let warm = self.warm.lock().unwrap();
+ let have = warm.get(key).map(Vec::len).unwrap_or(0);
+ self.config.target_size.saturating_sub(have)
+ };
+ for _ in 0..needed {
+ if let Ok(conn) = warm_one(key).await {
+ self.warm
+ .lock()
+ .unwrap()
+ .entry(key.clone())
+ .or_default()
+ .push(conn);
+ }
+ }
+ }
+
+ /// Test helper: insert an already-built warm connection (used to inject a
+ /// dead socket and assert it is discarded at handoff).
+ #[cfg(test)]
+ fn insert_warm(&self, key: UpstreamKey, conn: WarmConnection) {
+ self.warm.lock().unwrap().entry(key).or_default().push(conn);
+ }
+}
+
+/// Dial one upstream and pre-read its `session.created` into a [`WarmConnection`].
+///
+/// `key.api_key` is already resolved (non-blank). The first frame OpenAI sends
+/// unprompted is `session.created`; we buffer exactly that and read nothing more.
+async fn warm_one(key: &UpstreamKey) -> CoreResult {
+ let upstream: UpstreamWs =
+ dial_upstream(&key.model, &key.api_key, key.api_base.as_deref()).await?;
+ let (tx, mut rx) = upstream.split();
+ let session_created = read_event(&mut rx).await?;
+ Ok(WarmConnection {
+ tx,
+ rx,
+ session_created,
+ warmed_at: Instant::now(),
+ })
+}
+
+/// Resolve a deployment's API key into the pool key, returning `None` when no key
+/// can be resolved (those deployments simply aren't pooled — the request path
+/// still fresh-dials and surfaces the auth error there).
+pub fn upstream_key(
+ model: &str,
+ api_key: Option<&str>,
+ api_base: Option<&str>,
+) -> Option {
+ let api_key = resolve_api_key(api_key).ok()?;
+ Some(UpstreamKey {
+ model: model.to_string(),
+ api_key,
+ api_base: api_base.map(str::to_string),
+ })
+}
+
+/// Non-blocking liveness check: poll the upstream once. A warm socket is silent
+/// after `session.created`, so a pending `Close`/`Err`/`None` means it is dead.
+/// A pending data frame (shouldn't happen pre-handoff) is also treated as
+/// unhealthy — we'd rather discard and fresh-dial than hand over a socket in an
+/// unexpected state. `Pending` (the healthy case) returns `false`.
+fn is_dead(rx: &mut UpstreamRx) -> bool {
+ use futures_util::task::noop_waker_ref;
+ use futures_util::Stream;
+ use std::pin::Pin;
+ use std::task::{Context, Poll};
+
+ let mut cx = Context::from_waker(noop_waker_ref());
+ match Pin::new(rx).poll_next(&mut cx) {
+ Poll::Pending => false,
+ Poll::Ready(None) => true,
+ Poll::Ready(Some(Err(_))) => true,
+ // Any frame arriving before handoff is unexpected for a silent warm
+ // socket; treat it as unhealthy.
+ Poll::Ready(Some(Ok(_))) => true,
+ }
+}
+
+#[cfg(test)]
+mod tests {
+ use super::*;
+ use futures_util::SinkExt;
+ use std::net::SocketAddr;
+ use tokio::net::TcpListener;
+ use tokio_tungstenite::tungstenite::Message;
+
+ /// An in-process fake OpenAI realtime WS server. On connect it sends
+ /// `session.created`; on `response.create` it sends `response.created` +
+ /// `response.output_audio.delta` + `response.done`. Returns its `ws://` base.
+ async fn spawn_fake_openai() -> String {
+ let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
+ let addr: SocketAddr = listener.local_addr().unwrap();
+ tokio::spawn(async move {
+ while let Ok((stream, _)) = listener.accept().await {
+ tokio::spawn(handle_fake_conn(stream));
+ }
+ });
+ format!("ws://{addr}")
+ }
+
+ async fn handle_fake_conn(stream: tokio::net::TcpStream) {
+ let mut ws = match tokio_tungstenite::accept_async(stream).await {
+ Ok(ws) => ws,
+ Err(_) => return,
+ };
+ // Unprompted session.created, exactly like OpenAI.
+ let _ = ws
+ .send(Message::Text(
+ r#"{"type":"session.created","session":{"id":"sess_fake"}}"#.to_string(),
+ ))
+ .await;
+ while let Some(Ok(msg)) = ws.next().await {
+ if let Message::Text(text) = msg {
+ if text.contains("response.create") {
+ for frame in [
+ r#"{"type":"response.created"}"#,
+ r#"{"type":"response.output_audio.delta","delta":"AAAA"}"#,
+ r#"{"type":"response.done"}"#,
+ ] {
+ let _ = ws.send(Message::Text(frame.to_string())).await;
+ }
+ }
+ }
+ }
+ }
+
+ fn test_config() -> PoolConfig {
+ PoolConfig {
+ target_size: 2,
+ max_idle: Duration::from_secs(30),
+ }
+ }
+
+ fn key_for(base: &str) -> UpstreamKey {
+ UpstreamKey {
+ model: "gpt-realtime".to_string(),
+ api_key: "sk-test".to_string(),
+ api_base: Some(base.to_string()),
+ }
+ }
+
+ #[tokio::test]
+ async fn warm_handoff_relays_buffered_session_created() {
+ let base = spawn_fake_openai().await;
+ let pool = RealtimePool::new_unspawned(test_config());
+ let key = key_for(&base);
+ pool.register(key.clone());
+ pool.warm_now(&key).await;
+ assert_eq!(pool.warm_len(&key), 2);
+
+ let handoff = pool.take(&key).expect("a warm socket should be available");
+ assert_eq!(handoff.session_created.event_type, "session.created");
+ assert_eq!(
+ handoff
+ .session_created
+ .data
+ .get("session")
+ .and_then(|s| s.get("id"))
+ .and_then(|v| v.as_str()),
+ Some("sess_fake")
+ );
+ // Taking one leaves one.
+ assert_eq!(pool.warm_len(&key), 1);
+ }
+
+ #[tokio::test]
+ async fn pool_miss_returns_none_for_fresh_dial_fallback() {
+ let base = spawn_fake_openai().await;
+ let pool = RealtimePool::new_unspawned(test_config());
+ let key = key_for(&base);
+ // Registered but never warmed → empty bucket → miss.
+ pool.register(key.clone());
+ assert!(pool.take(&key).is_none());
+
+ // Unknown key → miss.
+ let other = key_for("ws://127.0.0.1:1");
+ assert!(pool.take(&other).is_none());
+ }
+
+ #[tokio::test]
+ async fn disabled_pool_never_hands_off() {
+ let pool = RealtimePool::disabled();
+ let key = key_for("ws://127.0.0.1:1");
+ pool.register(key.clone());
+ assert_eq!(pool.warm_len(&key), 0);
+ assert!(pool.take(&key).is_none());
+ }
+
+ #[tokio::test]
+ async fn dead_warm_socket_is_discarded() {
+ let base = spawn_fake_openai().await;
+ let pool = RealtimePool::new_unspawned(test_config());
+ let key = key_for(&base);
+ pool.register(key.clone());
+
+ // Build one real warm connection, then kill the upstream by dropping the
+ // server side: easiest is to dial, read session.created, then close our
+ // own rx's peer. Instead we forge "dead" via an already-closed socket:
+ // dial a connection and immediately send a Close from the client side so
+ // the server closes back, then warm it. Simpler: warm normally, then
+ // mark it stale by backdating warmed_at past max_idle and confirm it's
+ // dropped — that exercises the same discard path.
+ let mut conn = warm_one(&key).await.expect("warm one");
+ conn.warmed_at = Instant::now() - Duration::from_secs(3600); // past max_idle
+ pool.insert_warm(key.clone(), conn);
+ assert_eq!(pool.warm_len(&key), 1);
+
+ // take() must discard the stale socket and report a miss.
+ assert!(pool.take(&key).is_none());
+ assert_eq!(pool.warm_len(&key), 0);
+ }
+
+ #[tokio::test]
+ async fn background_replenisher_tops_up_registered_key() {
+ let base = spawn_fake_openai().await;
+ let pool = RealtimePool::spawn(test_config());
+ let key = key_for(&base);
+ pool.register(key.clone());
+
+ // Wait (bounded) for the background task to reach the target size.
+ let mut warmed = 0;
+ for _ in 0..40 {
+ tokio::time::sleep(Duration::from_millis(50)).await;
+ warmed = pool.warm_len(&key);
+ if warmed >= test_config().target_size {
+ break;
+ }
+ }
+ assert_eq!(
+ warmed,
+ test_config().target_size,
+ "background replenisher should warm up to target_size"
+ );
+ let handoff = pool.take(&key).expect("a warm socket should be available");
+ assert_eq!(handoff.session_created.event_type, "session.created");
+ }
+
+ #[tokio::test]
+ async fn closed_upstream_socket_is_detected_dead() {
+ // A genuinely dead socket: dial the fake, read session.created, then drop
+ // the server by closing from our side and waiting for the close to land.
+ let base = spawn_fake_openai().await;
+ let pool = RealtimePool::new_unspawned(test_config());
+ let key = key_for(&base);
+ pool.register(key.clone());
+
+ let mut conn = warm_one(&key).await.expect("warm one");
+ // Close the upstream from the client side; the server echoes a close.
+ let _ = conn.tx.send(Message::Close(None)).await;
+ // Give the close a moment to arrive on rx.
+ tokio::time::sleep(Duration::from_millis(50)).await;
+ pool.insert_warm(key.clone(), conn);
+
+ // Liveness check at take() should detect the close and discard it.
+ assert!(pool.take(&key).is_none());
+ assert_eq!(pool.warm_len(&key), 0);
+ }
+
+ #[tokio::test]
+ async fn broken_key_backs_off_instead_of_dialing_every_tick() {
+ // A key whose upstream is unreachable: every warm-up dial fails.
+ let pool = RealtimePool::new_unspawned(test_config());
+ let key = key_for("ws://127.0.0.1:1"); // nothing listens here
+ pool.register(key.clone());
+
+ // First pass attempts dials, they all fail → key enters backoff, no warm
+ // sockets, one recorded failure.
+ pool.replenish_all().await;
+ assert_eq!(pool.warm_len(&key), 0);
+ assert_eq!(pool.backoff_failures(&key), 1);
+ assert!(
+ pool.in_backoff(&key),
+ "a key whose dials all failed must be in backoff"
+ );
+
+ // An immediate next pass must be SKIPPED (still in the backoff window), so
+ // it does NOT fire another round of dials — the failure count is unchanged.
+ pool.replenish_all().await;
+ assert_eq!(
+ pool.backoff_failures(&key),
+ 1,
+ "replenish during the backoff window must not re-dial the broken key"
+ );
+ }
+
+ #[tokio::test]
+ async fn healthy_key_never_enters_backoff_and_clears_after_recovery() {
+ let base = spawn_fake_openai().await;
+ let pool = RealtimePool::new_unspawned(test_config());
+ let key = key_for(&base);
+ pool.register(key.clone());
+
+ // A reachable upstream: the pass succeeds, so the key is never backed off.
+ pool.replenish_all().await;
+ assert_eq!(pool.warm_len(&key), test_config().target_size);
+ assert_eq!(pool.backoff_failures(&key), 0);
+ assert!(!pool.in_backoff(&key));
+ }
+}
diff --git a/litellm-rust/crates/python-bridge/CLAUDE.md b/litellm-rust/crates/python-bridge/CLAUDE.md
new file mode 100644
index 00000000000..efa1a554c9c
--- /dev/null
+++ b/litellm-rust/crates/python-bridge/CLAUDE.md
@@ -0,0 +1,36 @@
+# CLAUDE.md
+
+Rules for `litellm-rust/crates/python-bridge`.
+
+## Responsibility
+
+`python-bridge` is the PyO3 boundary between Python LiteLLM and Rust transforms.
+Keep this crate thin. It adapts Python objects to Rust payloads and returns
+Python-compatible dictionaries.
+
+## Bridge Shape
+
+- Prefer one stable method per top-level LiteLLM route, for example
+ `ocr(payload)`.
+- Do not add one exported PyO3 function per provider helper unless there is a
+ measured reason.
+- Provider dispatch belongs in Rust route modules such as
+ `litellm_providers::ocr`, not in this PyO3 crate.
+- Python owns rollout state and fallback. Rust should return errors; Python
+ decides whether to raise or fall back.
+
+## Data Handling
+
+- OCR payloads can contain personal data and large base64 images. Do not log
+ payloads or provider responses.
+- Avoid copying large payloads more than needed. The current JSON round-trip is
+ acceptable for the first scaffold, but future performance work should evaluate
+ direct PyO3 conversion before expanding Rust coverage to image-heavy paths.
+- Do not expose raw Rust errors that include document contents or upstream
+ bodies.
+
+## Tests
+
+- `cargo test --workspace` must compile this crate.
+- Python tests must cover bridge disabled, bridge enabled, and module-missing
+ fallback behavior for every exposed route.
diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml
new file mode 100644
index 00000000000..80b6478daac
--- /dev/null
+++ b/litellm-rust/crates/python-bridge/Cargo.toml
@@ -0,0 +1,16 @@
+[package]
+name = "litellm-python-bridge"
+version = "0.1.0"
+edition.workspace = true
+license.workspace = true
+repository.workspace = true
+
+[lib]
+name = "litellm_python_bridge"
+crate-type = ["cdylib"]
+
+[dependencies]
+litellm-core.workspace = true
+litellm-providers.workspace = true
+pyo3 = { workspace = true, features = ["extension-module"] }
+serde_json.workspace = true
diff --git a/litellm-rust/crates/python-bridge/src/gil.rs b/litellm-rust/crates/python-bridge/src/gil.rs
new file mode 100644
index 00000000000..dc1b591735c
--- /dev/null
+++ b/litellm-rust/crates/python-bridge/src/gil.rs
@@ -0,0 +1,32 @@
+//! GIL accounting.
+//!
+//! A single chokepoint for releasing the GIL around blocking work. Every
+//! blocking call in the bridge goes through [`release_gil`] instead of calling
+//! `Python::allow_threads` directly, so the release count stays accurate and we
+//! have one place to extend later (timing histograms, per-call labels, etc.).
+
+use std::sync::atomic::{AtomicU64, Ordering};
+
+use pyo3::prelude::*;
+
+/// Number of times the bridge has released the GIL since process start.
+static GIL_RELEASES: AtomicU64 = AtomicU64::new(0);
+
+/// Release the GIL around `f`, recording the release.
+///
+/// `f` must not touch any Python state — that is what makes releasing the GIL
+/// safe. Returning the value back to Python re-acquires the GIL at the call
+/// site, after `f` has finished.
+pub fn release_gil(py: Python<'_>, f: F) -> T
+where
+ F: FnOnce() -> T + Send,
+ T: Send,
+{
+ GIL_RELEASES.fetch_add(1, Ordering::Relaxed);
+ py.allow_threads(f)
+}
+
+/// Total GIL releases performed by the bridge so far.
+pub fn release_count() -> u64 {
+ GIL_RELEASES.load(Ordering::Relaxed)
+}
diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs
new file mode 100644
index 00000000000..15e93f7b00c
--- /dev/null
+++ b/litellm-rust/crates/python-bridge/src/lib.rs
@@ -0,0 +1,100 @@
+use std::time::Duration;
+
+use litellm_core::error::CoreError;
+use litellm_providers::ocr::run_ocr;
+use pyo3::exceptions::{PyRuntimeError, PyValueError};
+use pyo3::prelude::*;
+use pyo3::types::{PyAny, PyDict};
+use serde_json::{Map, Value};
+
+mod gil;
+
+fn py_to_json(py: Python<'_>, value: &Bound<'_, PyAny>) -> PyResult {
+ let json = py.import("json")?;
+ let encoded: String = json.call_method1("dumps", (value,))?.extract()?;
+ serde_json::from_str(&encoded).map_err(|err| PyValueError::new_err(err.to_string()))
+}
+
+fn json_to_py(py: Python<'_>, value: Value) -> PyResult> {
+ let json = py.import("json")?;
+ let encoded =
+ serde_json::to_string(&value).map_err(|err| PyValueError::new_err(err.to_string()))?;
+ Ok(json.call_method1("loads", (encoded,))?.unbind())
+}
+
+/// Map a core error to the closest Python exception. Caller-input problems
+/// (auth, bad types, missing fields) -> `ValueError`; everything else
+/// (network, upstream status, parse failures) -> `RuntimeError`.
+fn core_error_to_pyerr(err: CoreError) -> PyErr {
+ match err {
+ CoreError::Auth(message) => PyValueError::new_err(message),
+ CoreError::InvalidType { .. } | CoreError::MissingField(_) => {
+ PyValueError::new_err(err.to_string())
+ }
+ other => PyRuntimeError::new_err(other.to_string()),
+ }
+}
+
+/// Perform a Mistral OCR call end to end and return the response as a dict.
+#[pyfunction]
+#[pyo3(signature = (model, document, api_key=None, api_base=None, optional_params=None, timeout_seconds=None))]
+fn ocr(
+ py: Python<'_>,
+ model: String,
+ document: Py,
+ api_key: Option,
+ api_base: Option,
+ optional_params: Option>,
+ timeout_seconds: Option,
+) -> PyResult> {
+ let document = py_to_json(py, document.bind(py))?;
+
+ let optional_params = match optional_params {
+ Some(params) => match py_to_json(py, params.bind(py))? {
+ Value::Object(map) => map,
+ _ => return Err(PyValueError::new_err("optional_params must be a dict")),
+ },
+ None => Map::new(),
+ };
+
+ let timeout = timeout_seconds.and_then(|secs| {
+ if secs.is_finite() && secs > 0.0 {
+ Some(Duration::from_secs_f64(secs))
+ } else {
+ None
+ }
+ });
+
+ // Release the GIL during the blocking HTTP call (counted for observability).
+ let result = gil::release_gil(py, || {
+ run_ocr(
+ &model,
+ document,
+ api_key.as_deref(),
+ api_base.as_deref(),
+ optional_params,
+ timeout,
+ )
+ });
+
+ match result {
+ Ok(value) => json_to_py(py, value),
+ Err(err) => Err(core_error_to_pyerr(err)),
+ }
+}
+
+/// Bridge GIL accounting, e.g. `{"releases": 12}`. Lets the Python side observe
+/// how often the bridge has dropped the GIL for blocking work.
+#[pyfunction]
+fn gil_stats(py: Python<'_>) -> PyResult> {
+ let stats = PyDict::new(py);
+ stats.set_item("releases", gil::release_count())?;
+ Ok(stats.into_any().unbind())
+}
+
+#[pymodule]
+fn litellm_python_bridge(module: &Bound<'_, PyModule>) -> PyResult<()> {
+ module.add_function(wrap_pyfunction!(ocr, module)?)?;
+ module.add_function(wrap_pyfunction!(gil_stats, module)?)?;
+ Ok(())
+}
diff --git a/litellm/__init__.py b/litellm/__init__.py
index d21234d2a81..d0513f77b35 100644
--- a/litellm/__init__.py
+++ b/litellm/__init__.py
@@ -80,6 +80,7 @@ from litellm.constants import (
WANDB_MODELS,
REPEATED_STREAMING_CHUNK_LIMIT,
request_timeout,
+ request_timeout_explicitly_set as request_timeout_explicitly_set,
open_ai_embedding_models,
cohere_embedding_models,
bedrock_embedding_models,
@@ -673,6 +674,7 @@ elevenlabs_models: Set = set()
dashscope_models: Set = set()
moonshot_models: Set = set()
publicai_models: Set = set()
+darkbloom_models: Set = set()
v0_models: Set = set()
morph_models: Set = set()
lambda_ai_models: Set = set()
@@ -927,6 +929,8 @@ def add_known_models(model_cost_map: Optional[Dict] = None):
moonshot_models.add(key)
elif value.get("litellm_provider") == "publicai":
publicai_models.add(key)
+ elif value.get("litellm_provider") == "darkbloom":
+ darkbloom_models.add(key)
elif value.get("litellm_provider") == "v0":
v0_models.add(key)
elif value.get("litellm_provider") == "morph":
@@ -1075,6 +1079,7 @@ model_list = list(
| dashscope_models
| moonshot_models
| publicai_models
+ | darkbloom_models
| v0_models
| morph_models
| lambda_ai_models
@@ -1179,6 +1184,7 @@ models_by_provider: dict = {
"modelscope": modelscope_models,
"moonshot": moonshot_models,
"publicai": publicai_models,
+ "darkbloom": darkbloom_models,
"v0": v0_models,
"morph": morph_models,
"lambda_ai": lambda_ai_models,
@@ -1400,6 +1406,7 @@ from .skills.main import (
)
from .containers.main import *
from .ocr.main import *
+from .ocr.rust_bridge import use_litellm_rust
from .rag.main import *
from .sandbox.main import *
from .search.main import *
@@ -1922,9 +1929,6 @@ if TYPE_CHECKING:
from .llms.fireworks_ai.completion.transformation import (
FireworksAITextCompletionConfig as FireworksAITextCompletionConfig,
)
- from .llms.fireworks_ai.audio_transcription.transformation import (
- FireworksAIAudioTranscriptionConfig as FireworksAIAudioTranscriptionConfig,
- )
from .llms.fireworks_ai.embed.fireworks_ai_transformation import (
FireworksAIEmbeddingConfig as FireworksAIEmbeddingConfig,
)
diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py
index e653b40fd04..4f131354d2e 100644
--- a/litellm/_lazy_imports_registry.py
+++ b/litellm/_lazy_imports_registry.py
@@ -260,7 +260,6 @@ LLM_CONFIG_NAMES = (
"SambaNovaEmbeddingConfig",
"FireworksAIConfig",
"FireworksAITextCompletionConfig",
- "FireworksAIAudioTranscriptionConfig",
"FireworksAIEmbeddingConfig",
"FriendliaiChatConfig",
"JinaAIEmbeddingConfig",
@@ -1027,10 +1026,6 @@ _LLM_CONFIGS_IMPORT_MAP = {
".llms.fireworks_ai.completion.transformation",
"FireworksAITextCompletionConfig",
),
- "FireworksAIAudioTranscriptionConfig": (
- ".llms.fireworks_ai.audio_transcription.transformation",
- "FireworksAIAudioTranscriptionConfig",
- ),
"FireworksAIEmbeddingConfig": (
".llms.fireworks_ai.embed.fireworks_ai_transformation",
"FireworksAIEmbeddingConfig",
diff --git a/litellm/constants.py b/litellm/constants.py
index c0e265c0e4a..212d34357f8 100644
--- a/litellm/constants.py
+++ b/litellm/constants.py
@@ -201,6 +201,18 @@ DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET = int(
# Provider-specific API base URLs
XAI_API_BASE = "https://api.x.ai/v1"
+OPEN_SANDBOX_API_BASE_ENV_VAR = "OPEN_SANDBOX_API_BASE"
+OPEN_SANDBOX_API_KEY_ENV_VAR = "OPEN_SANDBOX_API_KEY"
+OPEN_SANDBOX_DEFAULT_TEMPLATE = "opensandbox/code-interpreter:v1.1.0"
+_OPEN_SANDBOX_FALLBACK_ENTRYPOINT = "/opt/code-interpreter/code-interpreter.sh"
+OPEN_SANDBOX_DEFAULT_ENTRYPOINT = (_OPEN_SANDBOX_FALLBACK_ENTRYPOINT,)
+OPEN_SANDBOX_DEFAULT_LANGUAGE = "python"
+OPEN_SANDBOX_DEFAULT_CPU_LIMIT = "1"
+OPEN_SANDBOX_DEFAULT_MEMORY_LIMIT = "2Gi"
+OPEN_SANDBOX_EXECD_PORT = 44772
+OPEN_SANDBOX_DEFAULT_TIMEOUT = 300
+OPEN_SANDBOX_READY_TIMEOUT = 30.0
+OPEN_SANDBOX_POLL_INTERVAL = 0.2
DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET = int(
os.getenv("DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET", 1024)
@@ -456,6 +468,7 @@ HTTP_HANDLER_CONNECT_TIMEOUT_SECONDS: float = 5.0
request_timeout: float = float(
os.getenv("REQUEST_TIMEOUT", str(int(DEFAULT_REQUEST_TIMEOUT_SECONDS)))
)
+request_timeout_explicitly_set: bool = "REQUEST_TIMEOUT" in os.environ
DEFAULT_A2A_AGENT_TIMEOUT: float = float(
os.getenv("DEFAULT_A2A_AGENT_TIMEOUT", 6000)
) # 10 minutes
@@ -867,6 +880,7 @@ openai_compatible_providers: List = [
"docker_model_runner",
"ragflow",
"pinstripes", # Pinstripes - JSON-configured provider
+ "darkbloom",
]
openai_text_completion_compatible_providers: List = (
[ # providers that support `/v1/completions`
diff --git a/litellm/integrations/code_interpreter_interception/handler.py b/litellm/integrations/code_interpreter_interception/handler.py
index da8149eab9b..362581937d7 100644
--- a/litellm/integrations/code_interpreter_interception/handler.py
+++ b/litellm/integrations/code_interpreter_interception/handler.py
@@ -9,9 +9,11 @@ captured stdout back through the typed agentic loop plan.
import json
import time
import uuid
-from typing import Any, cast
+from typing import Any, Literal, TypedDict, cast
import litellm
+from pydantic import ValidationError
+
from litellm._logging import verbose_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.integrations.code_interpreter_interception import (
@@ -20,15 +22,93 @@ from litellm.types.integrations.code_interpreter_interception import (
from litellm.types.integrations.custom_logger import (
AgenticLoopPlan,
AgenticLoopRequestPatch,
+ CHAT_COMPLETION_AGENTIC_SURFACE,
+ NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES,
+ is_interception_internal_key,
+)
+from litellm.types.llms.openai import (
+ ChatCompletionAssistantMessage,
+ ChatCompletionAssistantToolCall,
+ ChatCompletionToolMessage,
+)
+from litellm.types.utils import (
+ CallTypes,
+ ChatCompletionMessageToolCall,
+ ModelResponse,
)
-from litellm.types.utils import CallTypes
LITELLM_CODE_EXECUTION_TOOL_NAME = "litellm_code_execution"
_INTERCEPTION_ACTIVE_KEY = "_code_interpreter_interception_active"
_SANDBOX_KEY = "_code_interpreter_interception_sandbox_key"
+_CONVERTED_STREAM_KEY = "_code_interpreter_interception_converted_stream"
+_LITELLM_METADATA_KEY = "litellm_metadata"
_CACHE_TTL_SECONDS = 15 * 60
+class CodeExecutionToolCall(TypedDict, total=False):
+ id: str | None
+ call_id: str | None
+ type: Literal["function"]
+ name: str
+ arguments: str
+
+
+class CodeInterpreterLogOutput(TypedDict):
+ type: Literal["logs"]
+ logs: str
+
+
+class CodeInterpreterCall(TypedDict):
+ id: str
+ type: Literal["code_interpreter_call"]
+ status: Literal["completed"]
+ code: str
+ container_id: str | None
+ outputs: list[CodeInterpreterLogOutput]
+
+
+class CodeExecutionFunctionParameters(TypedDict):
+ type: Literal["object"]
+ properties: dict[str, dict[str, str]]
+ required: list[str]
+
+
+class ResponsesFunctionTool(TypedDict):
+ type: Literal["function"]
+ name: str
+ description: str
+ parameters: CodeExecutionFunctionParameters
+
+
+class ChatCompletionFunctionDefinition(TypedDict):
+ name: str
+ description: str
+ parameters: CodeExecutionFunctionParameters
+
+
+class ChatCompletionFunctionTool(TypedDict):
+ type: Literal["function"]
+ function: ChatCompletionFunctionDefinition
+
+
+CodeExecutionFunctionTool = ResponsesFunctionTool | ChatCompletionFunctionTool
+
+
+class ResponsesFunctionToolChoice(TypedDict):
+ type: Literal["function"]
+ name: str
+
+
+class ChatCompletionFunctionToolChoice(TypedDict):
+ type: Literal["function"]
+ function: dict[str, str]
+
+
+CodeExecutionFunctionToolChoice = (
+ ResponsesFunctionToolChoice | ChatCompletionFunctionToolChoice
+)
+
+
def _resolve_sandbox_tool(sandbox_tool_name: str | None) -> dict[str, Any] | None:
try:
from litellm.sandbox.sandbox_tools import resolve_sandbox_tool
@@ -97,9 +177,15 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
if not kwargs.get("_agentic_loop_depth"):
kwargs.pop(_INTERCEPTION_ACTIVE_KEY, None)
kwargs.pop(_SANDBOX_KEY, None)
+ self._strip_interception_metadata(kwargs)
if not self.enabled:
return None
- if call_type not in (CallTypes.responses, CallTypes.aresponses):
+ if call_type not in (
+ CallTypes.responses,
+ CallTypes.aresponses,
+ CallTypes.completion,
+ CallTypes.acompletion,
+ ):
return None
if (
self.enabled_providers is not None
@@ -120,18 +206,10 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
kwargs[_SANDBOX_KEY] = uuid.uuid4().hex
if kwargs.get("stream"):
kwargs["stream"] = False
- kwargs["_code_interpreter_interception_converted_stream"] = True
+ kwargs[_CONVERTED_STREAM_KEY] = True
+ self._write_interception_metadata(kwargs)
- function_tool = {
- "type": "function",
- "name": LITELLM_CODE_EXECUTION_TOOL_NAME,
- "description": "Execute python code in a sandbox and return stdout.",
- "parameters": {
- "type": "object",
- "properties": {"code": {"type": "string"}},
- "required": ["code"],
- },
- }
+ function_tool = self._get_function_tool(call_type=call_type)
kwargs["tools"] = [
(
function_tool
@@ -141,19 +219,90 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
for tool in tools
]
if self._tool_choice_targets_code_interpreter(kwargs.get("tool_choice")):
- kwargs["tool_choice"] = {
- "type": "function",
- "name": LITELLM_CODE_EXECUTION_TOOL_NAME,
- }
+ kwargs["tool_choice"] = self._get_function_tool_choice(call_type=call_type)
return kwargs
+ @staticmethod
+ def _strip_interception_metadata(kwargs: dict[str, Any]) -> None:
+ metadata = kwargs.get(_LITELLM_METADATA_KEY)
+ if not isinstance(metadata, dict):
+ return
+ filtered_metadata = {
+ key: value
+ for key, value in metadata.items()
+ if not is_interception_internal_key(key)
+ and not key.startswith("_agentic_loop")
+ and key != "max_agentic_loops"
+ }
+ if filtered_metadata:
+ kwargs[_LITELLM_METADATA_KEY] = filtered_metadata
+ else:
+ kwargs.pop(_LITELLM_METADATA_KEY, None)
+
+ @staticmethod
+ def _write_interception_metadata(kwargs: dict[str, Any]) -> None:
+ metadata = kwargs.get(_LITELLM_METADATA_KEY)
+ metadata = dict(metadata) if isinstance(metadata, dict) else {}
+ for key in (_INTERCEPTION_ACTIVE_KEY, _SANDBOX_KEY, _CONVERTED_STREAM_KEY):
+ if key in kwargs:
+ metadata[key] = kwargs[key]
+ kwargs[_LITELLM_METADATA_KEY] = metadata
+
+ @staticmethod
+ def _get_function_parameters() -> CodeExecutionFunctionParameters:
+ return {
+ "type": "object",
+ "properties": {"code": {"type": "string"}},
+ "required": ["code"],
+ }
+
+ def _get_function_tool(
+ self, call_type: CallTypes | None
+ ) -> CodeExecutionFunctionTool:
+ description = "Execute python code in a sandbox and return stdout."
+ if call_type in (CallTypes.completion, CallTypes.acompletion):
+ return {
+ "type": "function",
+ "function": {
+ "name": LITELLM_CODE_EXECUTION_TOOL_NAME,
+ "description": description,
+ "parameters": self._get_function_parameters(),
+ },
+ }
+ return {
+ "type": "function",
+ "name": LITELLM_CODE_EXECUTION_TOOL_NAME,
+ "description": description,
+ "parameters": self._get_function_parameters(),
+ }
+
+ @staticmethod
+ def _get_function_tool_choice(
+ call_type: CallTypes | None,
+ ) -> CodeExecutionFunctionToolChoice:
+ if call_type in (CallTypes.completion, CallTypes.acompletion):
+ return {
+ "type": "function",
+ "function": {"name": LITELLM_CODE_EXECUTION_TOOL_NAME},
+ }
+ return {
+ "type": "function",
+ "name": LITELLM_CODE_EXECUTION_TOOL_NAME,
+ }
+
@staticmethod
def _tool_choice_targets_code_interpreter(tool_choice: Any) -> bool:
if not isinstance(tool_choice, dict):
return False
+ function = tool_choice.get("function")
return (
tool_choice.get("type") == "code_interpreter"
or tool_choice.get("name") == "code_interpreter"
+ or tool_choice.get("name") == LITELLM_CODE_EXECUTION_TOOL_NAME
+ or (
+ isinstance(function, dict)
+ and function.get("name") == LITELLM_CODE_EXECUTION_TOOL_NAME
+ )
)
def _resolve_provider(self, kwargs: dict[str, Any]) -> str | None:
@@ -188,7 +337,12 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
):
return False, {}
- tool_calls = self._extract_code_execution_tool_calls(response=response)
+ tool_calls = (
+ self._extract_chat_completion_code_execution_tool_calls(response=response)
+ if kwargs.get("_agentic_loop_api_surface")
+ == CHAT_COMPLETION_AGENTIC_SURFACE
+ else self._extract_code_execution_tool_calls(response=response)
+ )
if not tool_calls:
return False, {}
@@ -206,15 +360,24 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
stream: bool,
kwargs: dict,
) -> AgenticLoopPlan:
+ if kwargs.get("_agentic_loop_api_surface") == CHAT_COMPLETION_AGENTIC_SURFACE:
+ return await self._build_chat_completion_agentic_loop_plan(
+ tools=tools,
+ model=model,
+ messages=messages,
+ optional_params=anthropic_messages_optional_request_params,
+ kwargs=kwargs,
+ )
+
await self._prune_expired_cache()
- tool_calls = cast(list[dict[str, Any]], tools.get("tool_calls", []))
+ tool_calls = cast(list[CodeExecutionToolCall], tools.get("tool_calls", []))
sandbox_key = kwargs.get(_SANDBOX_KEY)
container, params = await self._get_or_create_container(cache_key=sandbox_key)
try:
- container_id = getattr(container, "id", None)
+ container_id = cast(str | None, getattr(container, "id", None))
input_list = self._normalize_messages(messages)
- code_interpreter_calls = []
+ code_interpreter_calls: list[CodeInterpreterCall] = []
for tool_call in tool_calls:
arguments = tool_call.get("arguments", "")
code = self._parse_code(arguments)
@@ -256,9 +419,12 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
request_patch = AgenticLoopRequestPatch(
model=model,
messages=input_list,
- tools=optional_params.get("tools"),
- optional_params={k: v for k, v in optional_params.items() if k != "tools"},
- kwargs={k: v for k, v in kwargs.items() if k != "litellm_logging_obj"},
+ tools=self._get_followup_tools(
+ tools=optional_params.get("tools"),
+ call_type=CallTypes.responses,
+ ),
+ optional_params=self._get_followup_optional_params(optional_params),
+ kwargs=self._filter_agentic_loop_kwargs(kwargs),
)
return AgenticLoopPlan(
@@ -271,12 +437,134 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
},
)
+ async def _build_chat_completion_agentic_loop_plan(
+ self,
+ tools: dict[str, object],
+ model: str,
+ messages: list[dict],
+ optional_params: dict[str, object],
+ kwargs: dict[str, object],
+ ) -> AgenticLoopPlan:
+ await self._prune_expired_cache()
+ tool_calls = cast(list[CodeExecutionToolCall], tools.get("tool_calls", []))
+ sandbox_key = cast(str | None, kwargs.get(_SANDBOX_KEY))
+ container, params = await self._get_or_create_container(cache_key=sandbox_key)
+
+ try:
+ container_id = cast(str | None, getattr(container, "id", None))
+ tool_results = [
+ await self._build_chat_completion_tool_result(
+ container=container,
+ params=params,
+ tool_call=tool_call,
+ container_id=container_id,
+ )
+ for tool_call in tool_calls
+ ]
+ except Exception:
+ await self._delete_container_for_cache_key(sandbox_key)
+ raise
+ tool_messages = [result[0] for result in tool_results]
+ code_interpreter_calls = [result[1] for result in tool_results]
+
+ request_patch = AgenticLoopRequestPatch(
+ model=model,
+ messages=list(messages)
+ + [self._build_chat_completion_assistant_message(tool_calls)]
+ + tool_messages,
+ tools=self._get_followup_tools(
+ tools=optional_params.get("tools"),
+ call_type=CallTypes.completion,
+ ),
+ optional_params=self._get_followup_optional_params(optional_params),
+ kwargs=self._filter_agentic_loop_kwargs(kwargs),
+ )
+
+ return AgenticLoopPlan(
+ run_agentic_loop=True,
+ request_patch=request_patch,
+ metadata={
+ "tool_type": "code_interpreter",
+ "sandbox_key": sandbox_key or "",
+ "code_interpreter_calls": code_interpreter_calls,
+ "response_format": "openai",
+ },
+ )
+
+ async def _build_chat_completion_tool_result(
+ self,
+ container: object,
+ params: dict[str, Any] | None,
+ tool_call: CodeExecutionToolCall,
+ container_id: str | None,
+ ) -> tuple[ChatCompletionToolMessage, CodeInterpreterCall]:
+ arguments = tool_call.get("arguments", "")
+ code = self._parse_code(arguments)
+ stdout = await self._run_tool_call(
+ container=container, params=params, arguments=arguments
+ )
+ tool_call_id = (
+ tool_call.get("id") or tool_call.get("call_id") or uuid.uuid4().hex
+ )
+ return (
+ {
+ "role": "tool",
+ "tool_call_id": tool_call_id,
+ "content": stdout,
+ },
+ {
+ "id": f"ci_{uuid.uuid4().hex}",
+ "type": "code_interpreter_call",
+ "status": "completed",
+ "code": code,
+ "container_id": container_id,
+ "outputs": [{"type": "logs", "logs": stdout}] if stdout else [],
+ },
+ )
+
async def async_agentic_loop_cleanup_hook(
self, plan: AgenticLoopPlan, kwargs: dict
) -> None:
metadata = plan.metadata or {} if plan else {}
await self._delete_container_for_cache_key(metadata.get("sandbox_key"))
+ @staticmethod
+ def _filter_agentic_loop_kwargs(kwargs: dict[str, object]) -> dict[str, object]:
+ return {
+ k: v
+ for k, v in kwargs.items()
+ if k not in {"litellm_logging_obj", "acompletion"}
+ and not is_interception_internal_key(
+ k, prefixes=NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES
+ )
+ }
+
+ def _get_followup_tools(
+ self, tools: object, call_type: CallTypes | None
+ ) -> list[dict[str, Any]] | None:
+ if not isinstance(tools, list):
+ return None
+ return [
+ (
+ self._get_function_tool(call_type=call_type)
+ if isinstance(tool, dict) and tool.get("type") == "code_interpreter"
+ else tool
+ )
+ for tool in tools
+ ]
+
+ def _get_followup_optional_params(
+ self, optional_params: dict[str, object]
+ ) -> dict[str, object]:
+ drop_tool_choice = self._tool_choice_targets_code_interpreter(
+ optional_params.get("tool_choice")
+ )
+ return {
+ k: v
+ for k, v in optional_params.items()
+ if k != "tools" and not (k == "tool_choice" and drop_tool_choice)
+ }
+
async def async_post_agentic_loop_response_hook(
self, response: Any, plan: AgenticLoopPlan, kwargs: dict
) -> Any:
@@ -420,7 +708,9 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
return list(messages)
return []
- def _extract_code_execution_tool_calls(self, response: Any) -> list[dict[str, Any]]:
+ def _extract_code_execution_tool_calls(
+ self, response: object
+ ) -> list[CodeExecutionToolCall]:
if isinstance(response, dict):
output = response.get("output", [])
else:
@@ -446,6 +736,82 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
if self._is_code_execution_call(item)
]
+ def _extract_chat_completion_code_execution_tool_calls(
+ self, response: ModelResponse | dict[str, Any]
+ ) -> list[CodeExecutionToolCall]:
+ model_response = self._to_model_response(response)
+ if model_response is None:
+ return []
+ choices = model_response.choices or []
+ if not choices:
+ return []
+ message = choices[0].message
+ tool_calls = message.tool_calls or []
+
+ return [
+ normalized
+ for tool_call in tool_calls
+ if (normalized := self._normalize_chat_completion_tool_call(tool_call))
+ is not None
+ ]
+
+ @staticmethod
+ def _normalize_chat_completion_tool_call(
+ tool_call: ChatCompletionMessageToolCall,
+ ) -> CodeExecutionToolCall | None:
+ if (
+ tool_call.type != "function"
+ or tool_call.function.name != LITELLM_CODE_EXECUTION_TOOL_NAME
+ ):
+ return None
+
+ arguments = tool_call.function.arguments
+ if isinstance(arguments, dict):
+ arguments = json.dumps(arguments)
+ elif not isinstance(arguments, str):
+ arguments = "" if arguments is None else str(arguments)
+
+ return {
+ "id": tool_call.id,
+ "call_id": tool_call.id,
+ "type": "function",
+ "name": LITELLM_CODE_EXECUTION_TOOL_NAME,
+ "arguments": arguments,
+ }
+
+ @staticmethod
+ def _build_chat_completion_assistant_message(
+ tool_calls: list[CodeExecutionToolCall],
+ ) -> ChatCompletionAssistantMessage:
+ return {
+ "role": "assistant",
+ "tool_calls": [
+ cast(
+ ChatCompletionAssistantToolCall,
+ {
+ "id": tool_call.get("id"),
+ "type": "function",
+ "function": {
+ "name": LITELLM_CODE_EXECUTION_TOOL_NAME,
+ "arguments": tool_call.get("arguments", ""),
+ },
+ },
+ )
+ for tool_call in tool_calls
+ ],
+ }
+
+ @staticmethod
+ def _to_model_response(
+ response: ModelResponse | dict[str, Any],
+ ) -> ModelResponse | None:
+ if isinstance(response, ModelResponse):
+ return response
+ try:
+ return ModelResponse(**response)
+ except (TypeError, ValidationError):
+ return None
+
def _is_code_execution_call(self, item: Any) -> bool:
if isinstance(item, dict):
return (
diff --git a/litellm/interactions/litellm_responses_transformation/transformation.py b/litellm/interactions/litellm_responses_transformation/transformation.py
index 173d4ca8764..0ff1a97cd0b 100644
--- a/litellm/interactions/litellm_responses_transformation/transformation.py
+++ b/litellm/interactions/litellm_responses_transformation/transformation.py
@@ -300,9 +300,6 @@ class LiteLLMResponsesInteractionsConfig:
"total_output_tokens": getattr(usage, "output_tokens", 0),
}
- # Add role
- interactions_response_dict["role"] = "model"
-
# Add updated (same as created for now)
interactions_response_dict["updated"] = created
diff --git a/litellm/litellm_core_utils/chat_completion_agentic_loop.py b/litellm/litellm_core_utils/chat_completion_agentic_loop.py
new file mode 100644
index 00000000000..938e892bd50
--- /dev/null
+++ b/litellm/litellm_core_utils/chat_completion_agentic_loop.py
@@ -0,0 +1,332 @@
+# this is a patch to allow for agentic loops covering llm_http_handler.py and openai sdk based calling flows for the .completion() api
+
+import json
+from typing import cast
+
+from litellm._logging import verbose_logger
+from litellm.integrations.custom_logger import CustomLogger
+from litellm.types.integrations.custom_logger import (
+ CHAT_COMPLETION_AGENTIC_SURFACE,
+ NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES,
+ AgenticLoopPlan,
+ AgenticLoopRequestPatch,
+ is_interception_internal_key,
+)
+from litellm.types.utils import ModelResponse
+from litellm.utils import CustomStreamWrapper
+
+_FOLLOWUP_INTERNAL_PARAMS = frozenset(
+ (
+ "acompletion",
+ "litellm_logging_obj",
+ "custom_llm_provider",
+ "model_alias_map",
+ "stream_response",
+ "custom_prompt_dict",
+ "_agentic_loop_api_surface",
+ )
+)
+
+
+def _gate_overridden(callback: CustomLogger) -> bool:
+ base = CustomLogger.async_should_run_agentic_loop
+ func = type(callback).async_should_run_agentic_loop
+ return getattr(func, "__func__", func) is not getattr(base, "__func__", base)
+
+
+def _build_plan_overridden(callback: CustomLogger) -> bool:
+ base = CustomLogger.async_build_agentic_loop_plan
+ func = type(callback).async_build_agentic_loop_plan
+ return getattr(func, "__func__", func) is not getattr(base, "__func__", base)
+
+
+def _post_hook_overridden(callback: CustomLogger) -> bool:
+ base = CustomLogger.async_post_agentic_loop_response_hook
+ func = type(callback).async_post_agentic_loop_response_hook
+ return getattr(func, "__func__", func) is not getattr(base, "__func__", base)
+
+
+def _coerce_int(value: object, default: int) -> int:
+ return int(value) if isinstance(value, (int, str)) else default
+
+
+def _agentic_loop_settings(kwargs: dict[str, object]) -> tuple[int, int, list[str]]:
+ depth = _coerce_int(kwargs.get("_agentic_loop_depth"), 0)
+ max_loops = max(_coerce_int(kwargs.get("max_agentic_loops"), 3), 1)
+ raw_fingerprints = kwargs.get("_agentic_loop_fingerprints")
+ fingerprints = (
+ [str(fp) for fp in raw_fingerprints]
+ if isinstance(raw_fingerprints, list)
+ else []
+ )
+ return depth, max_loops, fingerprints
+
+
+def _fingerprint_tools(tool_calls: object) -> str:
+ try:
+ return json.dumps(tool_calls, sort_keys=True, default=str)
+ except Exception:
+ return str(tool_calls)
+
+
+def _check_agentic_loop_safety(
+ tool_calls: object,
+ fingerprints: list[str],
+ depth: int,
+ max_loops: int,
+ model: str,
+) -> str:
+ fingerprint = _fingerprint_tools(tool_calls)
+ if fingerprint in fingerprints:
+ raise ValueError(
+ "Agentic loop detected repeated tool-call fingerprint; aborting rerun"
+ )
+ if depth >= max_loops:
+ raise ValueError(f"Exceeded max_agentic_loops={max_loops} for model={model}")
+ return fingerprint
+
+
+def _wrap_response_as_fake_stream(response: object) -> object:
+ if getattr(response, "object", None) == "chat.completion.chunk":
+ return response
+ if not hasattr(response, "choices"):
+ return response
+ from litellm.llms.base_llm.base_model_iterator import (
+ convert_model_response_to_streaming,
+ )
+
+ return convert_model_response_to_streaming(cast(ModelResponse, response))
+
+
+def _add_agentic_loop_metadata(kwargs_for_followup: dict[str, object]) -> None:
+ metadata = kwargs_for_followup.get("litellm_metadata")
+ metadata = dict(metadata) if isinstance(metadata, dict) else {}
+ for key, value in kwargs_for_followup.items():
+ if (
+ key.startswith("_agentic_loop")
+ or key == "max_agentic_loops"
+ or is_interception_internal_key(key)
+ ):
+ metadata[key] = value
+ kwargs_for_followup["litellm_metadata"] = metadata
+
+
+def _filter_followup_kwargs(source: dict[str, object]) -> dict[str, object]:
+ return {
+ k: v
+ for k, v in source.items()
+ if not is_interception_internal_key(
+ k, prefixes=NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES
+ )
+ and k not in _FOLLOWUP_INTERNAL_PARAMS
+ }
+
+
+async def _execute_chat_completion_agentic_plan(
+ *,
+ plan: AgenticLoopPlan,
+ callback: CustomLogger,
+ model: str,
+ optional_params: dict[str, object],
+ kwargs: dict[str, object],
+ logging_obj: object,
+ custom_llm_provider: str,
+ depth: int,
+ max_loops: int,
+ fingerprints: list[str],
+ fingerprint: str,
+) -> object:
+ import litellm
+
+ patch = plan.request_patch or AgenticLoopRequestPatch()
+ if patch.messages is None:
+ raise ValueError("Agentic loop plan missing patched messages")
+
+ full_model_name = patch.model or model
+ if "/" not in full_model_name:
+ full_model_name = f"{custom_llm_provider}/{full_model_name}"
+
+ optional_params_for_followup = {**optional_params, **patch.optional_params}
+ if patch.tools is not None:
+ optional_params_for_followup["tools"] = patch.tools
+ if "tool_choice" not in patch.optional_params:
+ optional_params_for_followup.pop("tool_choice", None)
+
+ kwargs_for_followup = _filter_followup_kwargs(kwargs)
+ kwargs_for_followup.update(
+ {
+ k: v
+ for k, v in _filter_followup_kwargs(patch.kwargs).items()
+ if k not in optional_params_for_followup
+ }
+ )
+ kwargs_for_followup["_agentic_loop_depth"] = depth + 1
+ kwargs_for_followup["max_agentic_loops"] = max_loops
+ kwargs_for_followup["_agentic_loop_fingerprints"] = fingerprints + [fingerprint]
+ _add_agentic_loop_metadata(kwargs_for_followup)
+
+ try:
+ response_followup = await litellm.acompletion(
+ model=full_model_name,
+ messages=patch.messages,
+ **optional_params_for_followup,
+ **kwargs_for_followup,
+ )
+ if _post_hook_overridden(callback):
+ try:
+ response_followup = (
+ await callback.async_post_agentic_loop_response_hook(
+ response=response_followup, plan=plan, kwargs=kwargs
+ )
+ )
+ except Exception as e:
+ _call_id = getattr(logging_obj, "litellm_call_id", "unknown")
+ verbose_logger.exception(
+ "LiteLLM.AgenticHookError: Exception in "
+ "async_post_agentic_loop_response_hook [call_id=%s model=%s]: %s",
+ _call_id,
+ model,
+ str(e),
+ )
+ if kwargs.get("_code_interpreter_interception_converted_stream") and not depth:
+ return _wrap_response_as_fake_stream(response_followup)
+ return response_followup
+ finally:
+ try:
+ await callback.async_agentic_loop_cleanup_hook(plan=plan, kwargs=kwargs)
+ except Exception as e:
+ _call_id = getattr(logging_obj, "litellm_call_id", "unknown")
+ verbose_logger.exception(
+ "LiteLLM.AgenticHookError: Exception in "
+ "async_agentic_loop_cleanup_hook [call_id=%s model=%s]: %s",
+ _call_id,
+ model,
+ str(e),
+ )
+
+
+async def maybe_run_chat_completion_agentic_loop(
+ *,
+ response: ModelResponse,
+ model: str,
+ messages: list,
+ optional_params: dict,
+ kwargs: dict,
+ logging_obj: object,
+ custom_llm_provider: str,
+ stream: bool,
+) -> ModelResponse | CustomStreamWrapper | None:
+ import litellm
+
+ callbacks = litellm.callbacks + (
+ getattr(logging_obj, "dynamic_success_callbacks", None) or []
+ )
+ depth, max_loops, fingerprints = _agentic_loop_settings(kwargs)
+ tools = optional_params.get("tools", [])
+
+ for callback in callbacks:
+ if not isinstance(callback, CustomLogger):
+ continue
+ if not _gate_overridden(callback):
+ continue
+
+ gate_kwargs = {
+ **kwargs,
+ "_agentic_loop_api_surface": CHAT_COMPLETION_AGENTIC_SURFACE,
+ "custom_llm_provider": custom_llm_provider,
+ }
+ try:
+ should_run, tool_calls = await callback.async_should_run_agentic_loop(
+ response=response,
+ model=model,
+ messages=messages,
+ tools=tools,
+ stream=stream,
+ custom_llm_provider=custom_llm_provider,
+ kwargs=gate_kwargs,
+ )
+ except Exception as e:
+ verbose_logger.exception(
+ "LiteLLM.AgenticHookError: Exception in chat completion agentic gate: %s",
+ str(e),
+ )
+ continue
+
+ if not should_run:
+ continue
+
+ fingerprint = _check_agentic_loop_safety(
+ tool_calls=tool_calls,
+ fingerprints=fingerprints,
+ depth=depth,
+ max_loops=max_loops,
+ model=model,
+ )
+
+ try:
+ plan_kwargs = {
+ **kwargs,
+ "_agentic_loop_api_surface": CHAT_COMPLETION_AGENTIC_SURFACE,
+ "custom_llm_provider": custom_llm_provider,
+ }
+ if not _build_plan_overridden(callback):
+ return await callback.async_run_agentic_loop(
+ tools=tool_calls,
+ model=model,
+ messages=messages,
+ response=response,
+ anthropic_messages_provider_config=None,
+ anthropic_messages_optional_request_params=optional_params,
+ logging_obj=logging_obj,
+ stream=stream,
+ kwargs=plan_kwargs,
+ )
+
+ plan = await callback.async_build_agentic_loop_plan(
+ tools=tool_calls,
+ model=model,
+ messages=messages,
+ response=response,
+ anthropic_messages_provider_config=None,
+ anthropic_messages_optional_request_params=optional_params,
+ logging_obj=logging_obj,
+ stream=stream,
+ kwargs=plan_kwargs,
+ )
+
+ if plan.response_override is not None:
+ return plan.response_override
+ if plan.terminate:
+ return response
+ if not plan.run_agentic_loop:
+ continue
+
+ return await _execute_chat_completion_agentic_plan(
+ plan=plan,
+ callback=callback,
+ model=model,
+ optional_params=optional_params,
+ kwargs=kwargs,
+ logging_obj=logging_obj,
+ custom_llm_provider=custom_llm_provider,
+ depth=depth,
+ max_loops=max_loops,
+ fingerprints=fingerprints,
+ fingerprint=fingerprint,
+ )
+ except Exception as e:
+ verbose_logger.exception(
+ "LiteLLM.AgenticHookError: Exception in chat completion agentic hooks: %s",
+ str(e),
+ )
+
+ if (
+ kwargs.get("_code_interpreter_interception_converted_stream")
+ and not depth
+ and hasattr(response, "choices")
+ ):
+ return cast(
+ "ModelResponse | CustomStreamWrapper",
+ _wrap_response_as_fake_stream(response),
+ )
+ return None
diff --git a/litellm/litellm_core_utils/completion_timeout.py b/litellm/litellm_core_utils/completion_timeout.py
index 5350d88e593..70c6896a323 100644
--- a/litellm/litellm_core_utils/completion_timeout.py
+++ b/litellm/litellm_core_utils/completion_timeout.py
@@ -6,10 +6,7 @@ from typing import Callable, Optional, Union
import httpx
-from litellm.constants import (
- COMPLETION_HTTP_FALLBACK_SECONDS,
- DEFAULT_REQUEST_TIMEOUT_SECONDS,
-)
+from litellm.constants import COMPLETION_HTTP_FALLBACK_SECONDS
class CompletionTimeout:
@@ -22,17 +19,13 @@ class CompletionTimeout:
"""
Used when ``model_timeout`` and kwargs timeouts are all unset.
- ``global_timeout`` is :attr:`litellm.request_timeout` (numeric / string), not
- :class:`httpx.Timeout`.
-
- If it equals :data:`~litellm.constants.DEFAULT_REQUEST_TIMEOUT_SECONDS` (6000),
- return :data:`~litellm.constants.COMPLETION_HTTP_FALLBACK_SECONDS`. Same if
- ``None``. Otherwise return ``float(global_timeout)``.
+ ``global_timeout`` is the explicitly-configured ``litellm.request_timeout``
+ (numeric / string) or ``None`` when it was never set. ``None`` falls back to
+ :data:`~litellm.constants.COMPLETION_HTTP_FALLBACK_SECONDS`; any explicit value
+ (including ``6000``) is honored.
"""
if global_timeout is None:
return COMPLETION_HTTP_FALLBACK_SECONDS
- if float(global_timeout) == float(DEFAULT_REQUEST_TIMEOUT_SECONDS):
- return COMPLETION_HTTP_FALLBACK_SECONDS
return float(global_timeout)
@staticmethod
@@ -50,11 +43,10 @@ class CompletionTimeout:
1. ``model_timeout`` (call argument / merged ``litellm_params``)
2. ``kwargs["timeout"]``
3. ``kwargs["request_timeout"]``
- 4. Fallback from ``global_timeout`` (:attr:`litellm.request_timeout`) — if it is
- the package default (6000), use 600 instead.
+ 4. ``global_timeout`` (the explicitly-configured ``litellm.request_timeout``),
+ or 600 when nothing was configured.
Coerce :class:`httpx.Timeout` when the provider does not support it.
- Explicit ``6000`` on the model or in kwargs is kept as ``6000``.
"""
resolved: Union[float, str, httpx.Timeout]
if model_timeout is not None:
diff --git a/litellm/litellm_core_utils/get_supported_openai_params.py b/litellm/litellm_core_utils/get_supported_openai_params.py
index e87042b9101..c22d3b99705 100644
--- a/litellm/litellm_core_utils/get_supported_openai_params.py
+++ b/litellm/litellm_core_utils/get_supported_openai_params.py
@@ -86,9 +86,7 @@ def get_supported_openai_params(
model=model
)
elif request_type == "transcription":
- return litellm.FireworksAIAudioTranscriptionConfig().get_supported_openai_params(
- model=model
- )
+ return None
else:
return litellm.FireworksAIConfig().get_supported_openai_params(model=model)
elif custom_llm_provider == "nvidia_nim":
@@ -191,7 +189,9 @@ def get_supported_openai_params(
)
elif custom_llm_provider == "sambanova":
if request_type == "embeddings":
- litellm.SambaNovaEmbeddingConfig().get_supported_openai_params(model=model)
+ return litellm.SambaNovaEmbeddingConfig().get_supported_openai_params(
+ model=model
+ )
else:
return litellm.SambanovaConfig().get_supported_openai_params(model=model)
elif custom_llm_provider == "nebius":
diff --git a/litellm/litellm_core_utils/request_timeout_resolver.py b/litellm/litellm_core_utils/request_timeout_resolver.py
new file mode 100644
index 00000000000..146c39ce9f3
--- /dev/null
+++ b/litellm/litellm_core_utils/request_timeout_resolver.py
@@ -0,0 +1,29 @@
+"""Single source of truth for whether ``litellm.request_timeout`` was configured.
+
+``litellm.request_timeout`` always holds a value (the package default,
+:data:`~litellm.constants.DEFAULT_REQUEST_TIMEOUT_SECONDS`), so a bare read can't
+tell "user asked for this" from "nobody set it". This resolver answers that:
+
+* ``request_timeout_explicitly_set`` is the authoritative signal, set when the
+ value comes from the ``REQUEST_TIMEOUT`` env var or ``litellm_settings``.
+* A runtime value that differs from the package default (e.g. ``litellm.request_timeout
+ = 300`` in SDK code) is also treated as explicit, for backwards compatibility.
+"""
+
+from __future__ import annotations
+
+from typing import Optional
+
+from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS
+
+
+def get_configured_request_timeout() -> Optional[float]:
+ """Return the explicitly-configured ``litellm.request_timeout``, else ``None``."""
+ import litellm
+
+ timeout = float(litellm.request_timeout)
+ if litellm.request_timeout_explicitly_set:
+ return timeout
+ if timeout != float(DEFAULT_REQUEST_TIMEOUT_SECONDS):
+ return timeout
+ return None
diff --git a/litellm/litellm_core_utils/sensitive_data_masker.py b/litellm/litellm_core_utils/sensitive_data_masker.py
index 4928dd08386..b14e12de7cd 100644
--- a/litellm/litellm_core_utils/sensitive_data_masker.py
+++ b/litellm/litellm_core_utils/sensitive_data_masker.py
@@ -12,6 +12,7 @@ class SensitiveDataMasker:
visible_prefix: int = 4,
visible_suffix: int = 4,
mask_char: str = "*",
+ mask_short_values: bool = True,
):
self.sensitive_patterns = sensitive_patterns or {
"password",
@@ -38,12 +39,17 @@ class SensitiveDataMasker:
self.visible_prefix = visible_prefix
self.visible_suffix = visible_suffix
self.mask_char = mask_char
+ self.mask_short_values = mask_short_values
def _mask_value(self, value: str) -> str:
- if not value or len(str(value)) < (self.visible_prefix + self.visible_suffix):
- return value
-
value_str = str(value)
+ if not value_str:
+ return value
+ if len(value_str) <= (self.visible_prefix + self.visible_suffix):
+ return (
+ self.mask_char * len(value_str) if self.mask_short_values else value_str
+ )
+
masked_length = len(value_str) - (self.visible_prefix + self.visible_suffix)
# Handle the case where visible_suffix is 0 to avoid showing the entire string
diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py
index c24c990f356..822b75b37f4 100644
--- a/litellm/llms/anthropic/chat/transformation.py
+++ b/litellm/llms/anthropic/chat/transformation.py
@@ -229,6 +229,11 @@ DROP_UNSUPPORTED_OUTPUT_CONFIG_WARNING = (
"Sonnet 4.6+, and Mythos Preview."
)
+DROP_UNSUPPORTED_SPEED_WARNING = (
+ "Dropping unsupported `speed` for model=%s "
+ "(drop_params=True). Fast mode is only supported on select Opus models."
+)
+
class AnthropicConfig(AnthropicModelInfo, BaseConfig):
"""
@@ -374,6 +379,51 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
for level in ("low", "minimal", "medium", "high", "xhigh", "max")
)
+ @staticmethod
+ def _model_supports_speed_param(
+ model: str, custom_llm_provider: Optional[str] = None
+ ) -> bool:
+ """Whether the model accepts Anthropic's ``speed`` parameter (fast mode).
+
+ Fast mode is direct Anthropic API-only (not Bedrock, Vertex, or Azure).
+ Those providers strip their prefix before this shared transform runs, so a
+ bare ``claude-opus-4-8`` would otherwise resolve to the direct-API entry;
+ the routed provider is checked explicitly to keep them out.
+ """
+ if custom_llm_provider is not None and custom_llm_provider != "anthropic":
+ return False
+ return (
+ AnthropicModelInfo._get_exact_model_capability(model, "supports_speed")
+ is True
+ )
+
+ @staticmethod
+ def _maybe_drop_speed_param(
+ model: str,
+ optional_params: dict,
+ drop_params: bool,
+ custom_llm_provider: Optional[str] = None,
+ ) -> None:
+ if "speed" not in optional_params:
+ return
+ if AnthropicConfig._model_supports_speed_param(model, custom_llm_provider):
+ return
+ if not (litellm.drop_params or drop_params):
+ speed_value = optional_params.get("speed")
+ raise litellm.utils.UnsupportedParamsError(
+ message=(
+ f"{model} does not support speed={speed_value!r}. "
+ "To drop unsupported params, set "
+ "`litellm.drop_params = True`."
+ ),
+ status_code=400,
+ )
+ litellm.verbose_logger.warning(
+ DROP_UNSUPPORTED_SPEED_WARNING,
+ model,
+ )
+ optional_params.pop("speed", None)
+
@staticmethod
def _raise_invalid_reasoning_effort(
model: str, value: Any, llm_provider: str
@@ -1569,8 +1619,13 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
anthropic_context_management
)
elif param == "speed" and isinstance(value, str):
- # Pass through Anthropic-specific speed parameter for fast mode
optional_params["speed"] = value
+ AnthropicConfig._maybe_drop_speed_param(
+ model=model,
+ optional_params=optional_params,
+ drop_params=drop_params,
+ custom_llm_provider=self.custom_llm_provider,
+ )
elif param == "cache_control" and isinstance(value, dict):
# Pass through top-level cache_control for automatic prompt caching
optional_params["cache_control"] = value
@@ -1875,6 +1930,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
"has no thinking_blocks. The model won't use extended thinking for this turn."
)
+ AnthropicConfig._maybe_drop_speed_param(
+ model=model,
+ optional_params=optional_params,
+ drop_params=litellm.drop_params
+ or litellm_params.get("drop_params") is True,
+ custom_llm_provider=self.custom_llm_provider,
+ )
+
headers = self.update_headers_with_optional_anthropic_beta(
headers=headers, optional_params=optional_params
)
diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py
index 5741513903c..0e41ef619ba 100644
--- a/litellm/llms/anthropic/common_utils.py
+++ b/litellm/llms/anthropic/common_utils.py
@@ -367,6 +367,16 @@ class AnthropicModelInfo(BaseLLMModelInfo):
pass
return None
+ @staticmethod
+ def _get_exact_model_capability(model: str, key: str) -> Optional[bool]:
+ """Read boolean capability ``key`` from the exact model-map entry only.
+
+ Unlike ``_get_model_capability``, does not walk stripped provider aliases.
+ Use when a feature is tied to a specific host (e.g. Anthropic API fast mode).
+ """
+ value = litellm.model_cost.get(model, {}).get(key)
+ return value if isinstance(value, bool) else None
+
@staticmethod
def _supports_model_capability(model: str, key: str) -> bool:
"""Check a boolean capability ``key`` in the model map.
diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py
index a3ac465c463..7b10a447bc8 100644
--- a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py
+++ b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py
@@ -507,7 +507,10 @@ def anthropic_messages_handler(
local_vars.update(kwargs)
anthropic_messages_optional_request_params = (
AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param(
- params=local_vars
+ params=local_vars,
+ model=model,
+ drop_params=litellm_params.get("drop_params") is True,
+ custom_llm_provider=custom_llm_provider,
)
)
if is_reasoning_auto_summary_enabled():
diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/utils.py b/litellm/llms/anthropic/experimental_pass_through/messages/utils.py
index 88832fb3f63..42167e0fdaa 100644
--- a/litellm/llms/anthropic/experimental_pass_through/messages/utils.py
+++ b/litellm/llms/anthropic/experimental_pass_through/messages/utils.py
@@ -23,12 +23,19 @@ class AnthropicMessagesRequestUtils:
@staticmethod
def get_requested_anthropic_messages_optional_param(
params: Dict[str, Any],
+ *,
+ model: str | None = None,
+ drop_params: bool = False,
+ custom_llm_provider: str | None = None,
) -> AnthropicMessagesRequestOptionalParams:
"""
Filter parameters to only include those defined in AnthropicMessagesRequestOptionalParams.
Args:
params: Dictionary of parameters to filter
+ model: Resolved model id; when set, unsupported params may be dropped
+ drop_params: Per-request drop_params flag (also respects litellm.drop_params)
+ custom_llm_provider: Routed provider; fast mode is gated to direct Anthropic
Returns:
AnthropicMessagesRequestOptionalParams instance with only the valid parameters
@@ -37,6 +44,15 @@ class AnthropicMessagesRequestUtils:
filtered_params = {
k: v for k, v in params.items() if k in valid_keys and v is not None
}
+ if model is not None:
+ from litellm.llms.anthropic.chat.transformation import AnthropicConfig
+
+ AnthropicConfig._maybe_drop_speed_param(
+ model=model,
+ optional_params=filtered_params,
+ drop_params=drop_params,
+ custom_llm_provider=custom_llm_provider,
+ )
return cast(AnthropicMessagesRequestOptionalParams, filtered_params)
diff --git a/litellm/llms/apiserpent/search/transformation.py b/litellm/llms/apiserpent/search/transformation.py
index 1eb7d34c875..bc11875ba12 100644
--- a/litellm/llms/apiserpent/search/transformation.py
+++ b/litellm/llms/apiserpent/search/transformation.py
@@ -53,7 +53,13 @@ class APISerpentSearchConfig(BaseSearchConfig):
api_base: Optional[str] = None,
**kwargs,
) -> Dict:
- api_key = api_key or get_secret_str("APISERPENT_API_KEY")
+ api_key = self.resolve_server_api_key(
+ caller_api_key=api_key,
+ caller_api_base=api_base,
+ key_env_vars=("APISERPENT_API_KEY",),
+ base_env_var="APISERPENT_API_BASE",
+ default_api_base=APISERPENT_BASE,
+ )
if not api_key:
raise ValueError(
"APISERPENT_API_KEY is not set. Set `APISERPENT_API_KEY` environment variable."
diff --git a/litellm/llms/base_llm/sandbox/transformation.py b/litellm/llms/base_llm/sandbox/transformation.py
index 6ad945f47a3..1c012a15fdb 100644
--- a/litellm/llms/base_llm/sandbox/transformation.py
+++ b/litellm/llms/base_llm/sandbox/transformation.py
@@ -8,10 +8,14 @@ run code -> delete container; `code_interpreter_tool` combines all three.
from typing import Any, Union
+import httpx
+
from pydantic import Field, PrivateAttr
from litellm.types.llms.base import LiteLLMPydanticObjectBase
+SANDBOX_MAX_OUTPUT_BYTES = 10 * 1024 * 1024
+
class ContainerHandle(LiteLLMPydanticObjectBase):
"""A live sandbox container. Carries everything needed to reach it again."""
@@ -53,7 +57,7 @@ class BaseSandboxConfig:
*,
template: str | None = None,
timeout: int | None = None,
- allow_internet_access: bool = True,
+ allow_internet_access: bool | None = None,
api_key: str | None = None,
**kwargs,
) -> ContainerHandle:
@@ -77,3 +81,16 @@ class BaseSandboxConfig:
**kwargs,
) -> bool:
raise NotImplementedError("adelete_sandbox must be implemented by provider")
+
+ async def _read_capped_lines(self, response: httpx.Response) -> list[str]:
+ lines: list[str] = []
+ total = 0
+ async for line in response.aiter_lines():
+ total += len(line.encode("utf-8"))
+ if total > SANDBOX_MAX_OUTPUT_BYTES:
+ raise ValueError(
+ f"Sandbox output exceeded {SANDBOX_MAX_OUTPUT_BYTES} bytes; aborting "
+ "to avoid unbounded memory use."
+ )
+ lines.append(line)
+ return lines
diff --git a/litellm/llms/base_llm/search/transformation.py b/litellm/llms/base_llm/search/transformation.py
index 4dfe86685fb..1581d8bb064 100644
--- a/litellm/llms/base_llm/search/transformation.py
+++ b/litellm/llms/base_llm/search/transformation.py
@@ -3,11 +3,13 @@ Base Search transformation configuration.
"""
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union
+from urllib.parse import urlsplit
import httpx
from pydantic import PrivateAttr
from litellm.llms.base_llm.chat.transformation import BaseLLMException
+from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.base import LiteLLMPydanticObjectBase
if TYPE_CHECKING:
@@ -16,6 +18,29 @@ else:
LiteLLMLoggingObj = Any
+def _search_host(url: str) -> str:
+ return urlsplit(url).netloc.lower()
+
+
+def _is_trusted_search_api_base(
+ caller_api_base: str,
+ default_api_base: str | None,
+ base_env_var: str | None,
+) -> bool:
+ candidate = _search_host(caller_api_base)
+ if not candidate:
+ return False
+ trusted = {
+ _search_host(base)
+ for base in (
+ default_api_base,
+ get_secret_str(base_env_var) if base_env_var else None,
+ )
+ if base
+ }
+ return candidate in trusted
+
+
class SearchResult(LiteLLMPydanticObjectBase):
"""Single search result."""
@@ -86,6 +111,60 @@ class BaseSearchConfig:
"max_tokens_per_page",
}
+ def _assert_trusted_api_base_for_server_credential(
+ self,
+ caller_api_base: str | None,
+ default_api_base: str | None,
+ base_env_var: str | None,
+ credential_name: str,
+ ) -> None:
+ """
+ Block sending a server-managed credential to a caller-chosen host.
+
+ A caller-supplied api_base is honored when constructing the request URL, so
+ falling back to a server-configured secret while the caller controls the host
+ leaks that secret. The provider default and the operator's own api_base
+ override are the only trusted destinations for a server-managed credential.
+ """
+ if not caller_api_base:
+ return
+ if _is_trusted_search_api_base(caller_api_base, default_api_base, base_env_var):
+ return
+ raise ValueError(
+ f"Refusing to send the server-configured {credential_name} to the "
+ f"caller-supplied api_base '{caller_api_base}'. Pass an explicit api_key "
+ f"when overriding api_base for this search provider."
+ )
+
+ def resolve_server_api_key(
+ self,
+ *,
+ caller_api_key: str | None,
+ caller_api_base: str | None,
+ key_env_vars: tuple[str, ...],
+ base_env_var: str | None,
+ default_api_base: str | None,
+ ) -> str | None:
+ """
+ Resolve a single-secret search API key, falling back to a server-managed
+ secret only when the request targets a trusted host.
+
+ Returns the caller's key when provided, otherwise the first set
+ server-managed secret (or None when none is set, for keyless providers).
+ """
+ if caller_api_key:
+ return caller_api_key
+ server_key = next(
+ (key for key in (get_secret_str(var) for var in key_env_vars) if key),
+ None,
+ )
+ if server_key is None:
+ return None
+ self._assert_trusted_api_base_for_server_credential(
+ caller_api_base, default_api_base, base_env_var, key_env_vars[0]
+ )
+ return server_key
+
def validate_environment(
self,
headers: Dict,
diff --git a/litellm/llms/bedrock/chat/invoke_handler.py b/litellm/llms/bedrock/chat/invoke_handler.py
index 75b560b4d6d..9fca7bc61af 100644
--- a/litellm/llms/bedrock/chat/invoke_handler.py
+++ b/litellm/llms/bedrock/chat/invoke_handler.py
@@ -70,6 +70,7 @@ from ..base_aws_llm import BaseAWSLLM
from ..common_utils import (
BedrockError,
ModelResponseIterator,
+ build_bedrock_stream_error,
get_bedrock_response_stream_shape,
get_bedrock_tool_name,
)
@@ -1841,23 +1842,7 @@ class AWSEventStreamDecoder:
parsed_response = self.parser.parse(response_dict, response_stream_shape)
if response_dict["status_code"] != 200:
- decoded_body = response_dict["body"].decode()
- if isinstance(decoded_body, dict):
- error_message = decoded_body.get("message")
- elif isinstance(decoded_body, str):
- error_message = decoded_body
- else:
- error_message = ""
- exception_status = response_dict["headers"].get(":exception-type")
- error_message = exception_status + " " + error_message
- raise BedrockError(
- status_code=response_dict["status_code"],
- message=(
- json.dumps(error_message)
- if isinstance(error_message, dict)
- else error_message
- ),
- )
+ raise build_bedrock_stream_error(response_dict, response_stream_shape)
if "chunk" in parsed_response:
chunk = parsed_response.get("chunk")
if not chunk:
diff --git a/litellm/llms/bedrock/chat/mantle/transformation.py b/litellm/llms/bedrock/chat/mantle/transformation.py
index cbed2232be5..93306025b02 100644
--- a/litellm/llms/bedrock/chat/mantle/transformation.py
+++ b/litellm/llms/bedrock/chat/mantle/transformation.py
@@ -12,6 +12,7 @@ from typing import TYPE_CHECKING, Any, List, Optional
from litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transformation import (
AmazonAnthropicClaudeConfig,
)
+from litellm.llms.bedrock.common_utils import build_mantle_messages_url
from litellm.types.llms.openai import AllMessageValues
if TYPE_CHECKING:
@@ -21,10 +22,6 @@ if TYPE_CHECKING:
else:
LiteLLMLoggingObj = Any
-MANTLE_ENDPOINT_TEMPLATE = (
- "https://bedrock-mantle.{region}.api.aws/anthropic/v1/messages"
-)
-
class AmazonMantleConfig(AmazonAnthropicClaudeConfig):
"""
@@ -46,7 +43,13 @@ class AmazonMantleConfig(AmazonAnthropicClaudeConfig):
stream: Optional[bool] = None,
) -> str:
region = self._get_aws_region_name(optional_params=optional_params, model=model)
- return MANTLE_ENDPOINT_TEMPLATE.format(region=region)
+ return build_mantle_messages_url(
+ api_base=api_base,
+ aws_bedrock_runtime_endpoint=optional_params.get(
+ "aws_bedrock_runtime_endpoint"
+ ),
+ region=region,
+ )
def validate_environment(
self,
diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py
index bdc5da321c6..5e97394f459 100644
--- a/litellm/llms/bedrock/common_utils.py
+++ b/litellm/llms/bedrock/common_utils.py
@@ -7,9 +7,21 @@ Common utilities used across bedrock chat/embedding/image generation
import functools
import json
import os
-from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union
+from typing import (
+ TYPE_CHECKING,
+ Any,
+ Dict,
+ List,
+ Literal,
+ Mapping,
+ Optional,
+ TypedDict,
+ Union,
+)
if TYPE_CHECKING:
+ from botocore.model import Shape
+
from litellm.types.llms.bedrock import BedrockCreateBatchRequest
import httpx
@@ -610,6 +622,31 @@ def strip_bedrock_throughput_suffix(model: str) -> str:
return model
+MANTLE_MESSAGES_PATH = "/anthropic/v1/messages"
+
+
+def build_mantle_messages_url(
+ api_base: Optional[str],
+ aws_bedrock_runtime_endpoint: Optional[str],
+ region: str,
+) -> str:
+ """Build the bedrock-mantle Anthropic /messages URL.
+
+ Honors an explicit endpoint override (``api_base``, then
+ ``aws_bedrock_runtime_endpoint``) so private VPC / VPCE / GovCloud Mantle
+ endpoints are reachable; otherwise falls back to the public regional host.
+ The mantle messages path is appended unless the override already carries it,
+ so callers can pass either the host or the full messages URL.
+ """
+ override = api_base or aws_bedrock_runtime_endpoint
+ if override:
+ base = override.rstrip("/")
+ if base.endswith(MANTLE_MESSAGES_PATH):
+ return base
+ return f"{base}{MANTLE_MESSAGES_PATH}"
+ return f"https://bedrock-mantle.{region}.api.aws{MANTLE_MESSAGES_PATH}"
+
+
def get_bedrock_base_model(model: str) -> str:
"""
Get the base model from the given model name.
@@ -1132,6 +1169,39 @@ def get_bedrock_response_stream_shape():
return _load_bedrock_response_stream_shape()
+class BedrockEventStreamResponseDict(TypedDict):
+ status_code: int
+ headers: Mapping[str, str]
+ body: bytes
+
+
+def build_bedrock_stream_error(
+ response_dict: BedrockEventStreamResponseDict,
+ response_stream_shape: Shape | None,
+) -> BedrockError:
+ """Build a BedrockError for a non-200 event-stream error event.
+
+ botocore hard-codes HTTP 400 on every mid-stream error event, so the modeled
+ ResponseStream member's httpStatusCode is the real status. Resolve it from the
+ shape and fall back to the raw status when the type is not modeled.
+ """
+ exception_type = response_dict["headers"].get(":exception-type")
+ decoded_body = response_dict["body"].decode()
+ message = f"{exception_type} {decoded_body}" if exception_type else decoded_body
+
+ status_code = response_dict["status_code"]
+ if exception_type is not None and response_stream_shape is not None:
+ member = response_stream_shape.members.get(exception_type)
+ if member is not None:
+ modeled_status = (
+ (member.metadata or {}).get("error", {}).get("httpStatusCode")
+ )
+ if modeled_status is not None:
+ status_code = int(modeled_status)
+
+ return BedrockError(status_code=status_code, message=message)
+
+
class BedrockEventStreamDecoderBase:
"""
Base class for event stream decoding for Bedrock
@@ -1156,23 +1226,7 @@ class BedrockEventStreamDecoderBase:
parsed_response = self.parser.parse(response_dict, response_stream_shape)
if response_dict["status_code"] != 200:
- decoded_body = response_dict["body"].decode()
- if isinstance(decoded_body, dict):
- error_message = decoded_body.get("message")
- elif isinstance(decoded_body, str):
- error_message = decoded_body
- else:
- error_message = ""
- exception_status = response_dict["headers"].get(":exception-type")
- error_message = exception_status + " " + error_message
- raise BedrockError(
- status_code=response_dict["status_code"],
- message=(
- json.dumps(error_message)
- if isinstance(error_message, dict)
- else error_message
- ),
- )
+ raise build_bedrock_stream_error(response_dict, response_stream_shape)
if "chunk" in parsed_response:
chunk = parsed_response.get("chunk")
if not chunk:
diff --git a/litellm/llms/bedrock/messages/mantle_transformation.py b/litellm/llms/bedrock/messages/mantle_transformation.py
index 900d9aa97d8..94e7f90b719 100644
--- a/litellm/llms/bedrock/messages/mantle_transformation.py
+++ b/litellm/llms/bedrock/messages/mantle_transformation.py
@@ -8,6 +8,7 @@ stripping that are specific to the bedrock-mantle endpoint.
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple
+from litellm.llms.bedrock.common_utils import build_mantle_messages_url
from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation import (
AmazonAnthropicClaudeMessagesConfig,
)
@@ -20,10 +21,6 @@ if TYPE_CHECKING:
else:
LiteLLMLoggingObj = Any
-MANTLE_ENDPOINT_TEMPLATE = (
- "https://bedrock-mantle.{region}.api.aws/anthropic/v1/messages"
-)
-
class AmazonMantleMessagesConfig(AmazonAnthropicClaudeMessagesConfig):
"""
@@ -43,7 +40,13 @@ class AmazonMantleMessagesConfig(AmazonAnthropicClaudeMessagesConfig):
stream: Optional[bool] = None,
) -> str:
region = self._get_aws_region_name(optional_params=optional_params, model=model)
- return MANTLE_ENDPOINT_TEMPLATE.format(region=region)
+ return build_mantle_messages_url(
+ api_base=api_base,
+ aws_bedrock_runtime_endpoint=optional_params.get(
+ "aws_bedrock_runtime_endpoint"
+ ),
+ region=region,
+ )
def validate_anthropic_messages_environment(
self,
diff --git a/litellm/llms/brave/search/transformation.py b/litellm/llms/brave/search/transformation.py
index 9dfcd6bc75a..8ffe7dcb126 100644
--- a/litellm/llms/brave/search/transformation.py
+++ b/litellm/llms/brave/search/transformation.py
@@ -115,7 +115,13 @@ class BraveSearchConfig(BaseSearchConfig):
"""
Validate environment and return headers.
"""
- api_key = api_key or get_secret_str("BRAVE_API_KEY")
+ api_key = self.resolve_server_api_key(
+ caller_api_key=api_key,
+ caller_api_base=api_base,
+ key_env_vars=("BRAVE_API_KEY",),
+ base_env_var="BRAVE_API_BASE",
+ default_api_base=self.BRAVE_API_BASE,
+ )
if not api_key:
raise ValueError(
diff --git a/litellm/llms/cloudflare/chat/transformation.py b/litellm/llms/cloudflare/chat/transformation.py
index 66e253f304d..68f08741cc5 100644
--- a/litellm/llms/cloudflare/chat/transformation.py
+++ b/litellm/llms/cloudflare/chat/transformation.py
@@ -1,26 +1,15 @@
-import json
-import time
-from typing import AsyncIterator, Iterator, List, Optional, Union
+from typing import List, Optional, Union
import httpx
-import litellm
-from litellm.litellm_core_utils.url_utils import encode_url_path_segments
-from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
-from litellm.llms.base_llm.chat.transformation import (
- BaseConfig,
- BaseLLMException,
- LiteLLMLoggingObj,
+from litellm._logging import verbose_logger
+from litellm.llms.base_llm.chat.transformation import BaseLLMException
+from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
+from litellm.secret_managers.main import (
+ get_secret_str,
+ normalize_nonempty_secret_str,
)
-from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllMessageValues
-from litellm.types.utils import (
- ChatCompletionToolCallChunk,
- ChatCompletionUsageBlock,
- GenericStreamingChunk,
- ModelResponse,
- Usage,
-)
class CloudflareError(BaseLLMException):
@@ -34,26 +23,46 @@ class CloudflareError(BaseLLMException):
message=message,
request=self.request,
response=self.response,
- ) # Call the base class constructor with the parameters it needs
+ )
-class CloudflareChatConfig(BaseConfig):
- max_tokens: Optional[int] = None
- stream: Optional[bool] = None
-
- def __init__(
+class CloudflareChatConfig(OpenAIGPTConfig):
+ def get_complete_url(
self,
- max_tokens: Optional[int] = None,
+ api_base: Optional[str],
+ api_key: Optional[str],
+ model: str,
+ optional_params: dict,
+ litellm_params: dict,
stream: Optional[bool] = None,
- ) -> None:
- locals_ = locals().copy()
- for key, value in locals_.items():
- if key != "self" and value is not None:
- setattr(self.__class__, key, value)
+ ) -> str:
+ return super().get_complete_url(
+ api_base=self._resolve_api_base(api_base),
+ api_key=api_key,
+ model=model,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ stream=stream,
+ )
- @classmethod
- def get_config(cls):
- return super().get_config()
+ @staticmethod
+ def _resolve_api_base(api_base: Optional[str]) -> str:
+ if not api_base:
+ account_id = normalize_nonempty_secret_str(
+ get_secret_str("CLOUDFLARE_ACCOUNT_ID")
+ )
+ if account_id is None:
+ raise ValueError(
+ "Missing CLOUDFLARE_ACCOUNT_ID - set CLOUDFLARE_ACCOUNT_ID in the environment or pass api_base explicitly"
+ )
+ return f"https://api.cloudflare.com/client/v4/accounts/{account_id}/ai/v1"
+ trimmed = api_base.rstrip("/")
+ if trimmed.endswith("/ai/run"):
+ verbose_logger.warning(
+ "Cloudflare api_base ending in '/ai/run' is the legacy Workers AI path and no longer serves OpenAI-compatible requests; rewriting to the '/ai/v1' endpoint"
+ )
+ return f"{trimmed[: -len('/ai/run')]}/ai/v1"
+ return api_base
def validate_environment(
self,
@@ -67,107 +76,18 @@ class CloudflareChatConfig(BaseConfig):
) -> dict:
if api_key is None:
raise ValueError(
- "Missing CloudflareError API Key - A call is being made to cloudflare but no key is set either in the environment variables or via params"
+ "Missing Cloudflare API Key - A call is being made to cloudflare but no key is set either in the environment variables or via params"
)
- headers = {
- "accept": "application/json",
- "content-type": "apbplication/json",
- "Authorization": "Bearer " + api_key,
- }
- return headers
-
- def get_complete_url(
- self,
- api_base: Optional[str],
- api_key: Optional[str],
- model: str,
- optional_params: dict,
- litellm_params: dict,
- stream: Optional[bool] = None,
- ) -> str:
- if api_base is None:
- account_id = get_secret_str("CLOUDFLARE_ACCOUNT_ID")
- api_base = (
- f"https://api.cloudflare.com/client/v4/accounts/{account_id}/ai/run/"
- )
- encoded_model = encode_url_path_segments(model, field_name="model")
- return api_base + encoded_model
-
- def get_supported_openai_params(self, model: str) -> List[str]:
- return [
- "stream",
- "max_tokens",
- ]
-
- def map_openai_params(
- self,
- non_default_params: dict,
- optional_params: dict,
- model: str,
- drop_params: bool,
- ) -> dict:
- supported_openai_params = self.get_supported_openai_params(model=model)
- for param, value in non_default_params.items():
- if param == "max_completion_tokens":
- optional_params["max_tokens"] = value
- elif param in supported_openai_params:
- optional_params[param] = value
- return optional_params
-
- def transform_request(
- self,
- model: str,
- messages: List[AllMessageValues],
- optional_params: dict,
- litellm_params: dict,
- headers: dict,
- ) -> dict:
- config = litellm.CloudflareChatConfig.get_config()
- for k, v in config.items():
- if k not in optional_params:
- optional_params[k] = v
-
- data = {
- "messages": messages,
- **optional_params,
- }
- return data
-
- def transform_response(
- self,
- model: str,
- raw_response: httpx.Response,
- model_response: ModelResponse,
- logging_obj: LiteLLMLoggingObj,
- request_data: dict,
- messages: List[AllMessageValues],
- optional_params: dict,
- litellm_params: dict,
- encoding: str,
- api_key: Optional[str] = None,
- json_mode: Optional[bool] = None,
- ) -> ModelResponse:
- completion_response = raw_response.json()
-
- # Support both "response" and "response_text" keys (newer models like Nemotron use "response_text")
- result = completion_response["result"]
- model_response.choices[0].message.content = result.get("response") if result.get("response") is not None else result.get("response_text", "") # type: ignore
-
- prompt_tokens = litellm.utils.get_token_count(messages=messages, model=model)
- completion_tokens = len(
- encoding.encode(model_response["choices"][0]["message"].get("content", ""))
+ return super().validate_environment(
+ headers=headers,
+ model=model,
+ messages=messages,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ api_key=api_key,
+ api_base=api_base,
)
- model_response.created = int(time.time())
- model_response.model = "cloudflare/" + model
- usage = Usage(
- prompt_tokens=prompt_tokens,
- completion_tokens=completion_tokens,
- total_tokens=prompt_tokens + completion_tokens,
- )
- setattr(model_response, "usage", usage)
- return model_response
-
def get_error_class(
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
) -> BaseLLMException:
@@ -175,48 +95,3 @@ class CloudflareChatConfig(BaseConfig):
status_code=status_code,
message=error_message,
)
-
- def get_model_response_iterator(
- self,
- streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse],
- sync_stream: bool,
- json_mode: Optional[bool] = False,
- ):
- return CloudflareChatResponseIterator(
- streaming_response=streaming_response,
- sync_stream=sync_stream,
- json_mode=json_mode,
- )
-
-
-class CloudflareChatResponseIterator(BaseModelResponseIterator):
- def chunk_parser(self, chunk: dict) -> GenericStreamingChunk:
- try:
- text = ""
- tool_use: Optional[ChatCompletionToolCallChunk] = None
- is_finished = False
- finish_reason = ""
- usage: Optional[ChatCompletionUsageBlock] = None
- provider_specific_fields = None
-
- index = int(chunk.get("index", 0))
-
- if "response" in chunk and chunk["response"] is not None:
- text = chunk["response"]
- elif "response_text" in chunk and chunk["response_text"] is not None:
- text = chunk["response_text"]
-
- returned_chunk = GenericStreamingChunk(
- text=text,
- tool_use=tool_use,
- is_finished=is_finished,
- finish_reason=finish_reason,
- usage=usage,
- index=index,
- provider_specific_fields=provider_specific_fields,
- )
-
- return returned_chunk
-
- except json.JSONDecodeError:
- raise ValueError(f"Failed to decode JSON from chunk: {chunk}")
diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py
index 01c94476431..1000ab12803 100644
--- a/litellm/llms/custom_httpx/http_handler.py
+++ b/litellm/llms/custom_httpx/http_handler.py
@@ -42,6 +42,9 @@ from litellm.constants import (
HTTP_HANDLER_CONNECT_TIMEOUT_SECONDS,
)
from litellm.litellm_core_utils.logging_utils import track_llm_api_timing
+from litellm.litellm_core_utils.request_timeout_resolver import (
+ get_configured_request_timeout,
+)
from litellm.types.llms.custom_http import *
if TYPE_CHECKING:
@@ -134,6 +137,18 @@ _DEFAULT_TIMEOUT = httpx.Timeout(
timeout=COMPLETION_HTTP_FALLBACK_SECONDS,
connect=HTTP_HANDLER_CONNECT_TIMEOUT_SECONDS,
)
+
+
+def _default_cached_client_timeout() -> httpx.Timeout:
+ """Timeout for cached default httpx clients; honors an explicit litellm.request_timeout."""
+ configured = get_configured_request_timeout()
+ if configured is None:
+ return _DEFAULT_TIMEOUT
+ return httpx.Timeout(
+ timeout=configured, connect=HTTP_HANDLER_CONNECT_TIMEOUT_SECONDS
+ )
+
+
_STREAMING_ERROR_BODY_READ_TIMEOUT_SECONDS = 5.0
_STREAMING_ERROR_BODY_READ_EXECUTOR = concurrent.futures.ThreadPoolExecutor(
max_workers=50,
@@ -1379,7 +1394,7 @@ def get_async_httpx_client(
_new_client = AsyncHTTPHandler(**handler_params)
else:
_new_client = AsyncHTTPHandler(
- timeout=_DEFAULT_TIMEOUT,
+ timeout=_default_cached_client_timeout(),
shared_session=shared_session,
)
@@ -1428,7 +1443,7 @@ def _get_httpx_client(params: Optional[dict] = None) -> HTTPHandler:
}
_new_client = HTTPHandler(**handler_params)
else:
- _new_client = HTTPHandler(timeout=_DEFAULT_TIMEOUT)
+ _new_client = HTTPHandler(timeout=_default_cached_client_timeout())
cache.set_cache(
key=_cache_key_name,
diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py
index 790bd0519d7..948c90f9f99 100644
--- a/litellm/llms/custom_httpx/llm_http_handler.py
+++ b/litellm/llms/custom_httpx/llm_http_handler.py
@@ -1,5 +1,6 @@
import json
import ssl
+from functools import lru_cache
from urllib.parse import parse_qs, urlencode, urlparse, urlunparse
from typing import (
TYPE_CHECKING,
@@ -13,6 +14,7 @@ from typing import (
Tuple,
Union,
cast,
+ get_type_hints,
)
import httpx # type: ignore
@@ -26,6 +28,7 @@ from litellm._logging import _redact_string, verbose_logger
from litellm.anthropic_beta_headers_manager import update_headers_with_filtered_beta
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming
+from litellm.litellm_core_utils.asyncify import run_async_function
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
from litellm.llms.base_llm.anthropic_messages.transformation import (
BaseAnthropicMessagesConfig,
@@ -101,6 +104,7 @@ from litellm.types.llms.openai import (
HttpxBinaryResponseContent,
OpenAIFileObject,
ResponseInputParam,
+ ResponsesAPIOptionalRequestParams,
ResponsesAPIResponse,
)
from litellm.types.rerank import RerankResponse
@@ -135,6 +139,7 @@ from litellm.utils import (
ImageResponse,
ModelResponse,
ProviderConfigManager,
+ async_pre_call_deployment_hook,
)
from .http_handler import get_shared_realtime_ssl_context
@@ -184,6 +189,47 @@ def _google_genai_streaming_hidden_params(
}
+@lru_cache(maxsize=None)
+def _responses_api_optional_request_param_names() -> frozenset[str]:
+ return frozenset(get_type_hints(ResponsesAPIOptionalRequestParams).keys())
+
+
+def _custom_logger_callbacks(logging_obj: Any) -> list[Any]:
+ from litellm.integrations.custom_logger import CustomLogger
+ from litellm.litellm_core_utils.litellm_logging import (
+ get_custom_logger_compatible_class,
+ )
+
+ dynamic_success_callbacks = getattr(logging_obj, "dynamic_success_callbacks", None)
+ callbacks = list(litellm.callbacks)
+ if isinstance(dynamic_success_callbacks, (list, tuple)):
+ callbacks.extend(dynamic_success_callbacks)
+
+ custom_loggers: list[Any] = []
+ for cb in callbacks:
+ if isinstance(cb, str):
+ resolved = get_custom_logger_compatible_class(cb) # type: ignore[arg-type]
+ if resolved is None:
+ continue
+ cb = resolved
+ if isinstance(cb, CustomLogger):
+ custom_loggers.append(cb)
+ return custom_loggers
+
+
+def _has_pre_call_deployment_hook(logging_obj: Any) -> bool:
+ from litellm.integrations.custom_logger import CustomLogger
+
+ base_func = CustomLogger.async_pre_call_deployment_hook
+ for cb in _custom_logger_callbacks(logging_obj):
+ cb_func = getattr(type(cb), "async_pre_call_deployment_hook", base_func)
+ if getattr(cb_func, "__func__", cb_func) is not getattr(
+ base_func, "__func__", base_func
+ ):
+ return True
+ return False
+
+
class BaseLLMHTTPHandler:
async def _make_common_async_call(
self,
@@ -1833,6 +1879,9 @@ class BaseLLMHTTPHandler:
data = provider_config.transform_search_request(
query=query,
optional_params=optional_params,
+ api_key=api_key,
+ api_base=api_base,
+ headers=headers or {},
)
# Get complete URL (pass data for providers that need request body for URL construction)
@@ -2224,12 +2273,92 @@ class BaseLLMHTTPHandler:
)
raise ValueError("anthropic_messages_handler is not implemented for sync calls")
+ def _run_sync_responses_pre_call_deployment_hook(
+ self,
+ *,
+ model: str,
+ input: Union[str, ResponseInputParam],
+ custom_llm_provider: str,
+ response_api_optional_request_params: dict[str, Any],
+ litellm_params: GenericLiteLLMParams,
+ logging_obj: LiteLLMLoggingObj,
+ ) -> tuple[
+ str,
+ Union[str, ResponseInputParam],
+ str,
+ dict[str, Any],
+ GenericLiteLLMParams,
+ ]:
+ if not _has_pre_call_deployment_hook(logging_obj):
+ return (
+ model,
+ input,
+ custom_llm_provider,
+ response_api_optional_request_params,
+ litellm_params,
+ )
+
+ modified_kwargs = run_async_function(
+ async_pre_call_deployment_hook,
+ {
+ **dict(litellm_params),
+ **response_api_optional_request_params,
+ "model": model,
+ "input": input,
+ "custom_llm_provider": custom_llm_provider,
+ },
+ CallTypes.responses.value,
+ )
+ if modified_kwargs is None:
+ return (
+ model,
+ input,
+ custom_llm_provider,
+ response_api_optional_request_params,
+ litellm_params,
+ )
+
+ optional_param_names = _responses_api_optional_request_param_names()
+ updated_response_params = {
+ **response_api_optional_request_params,
+ **{
+ key: value
+ for key, value in modified_kwargs.items()
+ if key in optional_param_names
+ },
+ }
+ updated_litellm_params = GenericLiteLLMParams(
+ **{
+ **dict(litellm_params),
+ **{
+ key: value
+ for key, value in modified_kwargs.items()
+ if key not in optional_param_names
+ and key not in {"model", "input", "custom_llm_provider"}
+ },
+ }
+ )
+ return (
+ str(modified_kwargs["model"]) if "model" in modified_kwargs else model,
+ cast(
+ Union[str, ResponseInputParam],
+ modified_kwargs["input"] if "input" in modified_kwargs else input,
+ ),
+ (
+ str(modified_kwargs["custom_llm_provider"])
+ if "custom_llm_provider" in modified_kwargs
+ else custom_llm_provider
+ ),
+ updated_response_params,
+ updated_litellm_params,
+ )
+
def response_api_handler(
self,
model: str,
input: Union[str, ResponseInputParam],
responses_api_provider_config: BaseResponsesAPIConfig,
- response_api_optional_request_params: Dict,
+ response_api_optional_request_params: dict[str, Any],
custom_llm_provider: str,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
@@ -2276,6 +2405,21 @@ class BaseLLMHTTPHandler:
shared_session=shared_session,
)
+ (
+ model,
+ input,
+ custom_llm_provider,
+ response_api_optional_request_params,
+ litellm_params,
+ ) = self._run_sync_responses_pre_call_deployment_hook(
+ model=model,
+ input=input,
+ custom_llm_provider=custom_llm_provider,
+ response_api_optional_request_params=response_api_optional_request_params,
+ litellm_params=litellm_params,
+ logging_obj=logging_obj,
+ )
+
if client is None or not isinstance(client, HTTPHandler):
sync_httpx_client = _get_httpx_client(
params={"ssl_verify": litellm_params.get("ssl_verify", None)}
@@ -2414,9 +2558,27 @@ class BaseLLMHTTPHandler:
logging_obj=logging_obj,
)
)
- # Responses agentic interception (e.g. code interpreter) runs the follow-up
- # loop via the async hook, so it is async-only for now; the sync path returns
- # the initial response unchanged.
+
+ if self._has_agentic_completion_hook(logging_obj):
+ final_response = run_async_function(
+ self._call_agentic_completion_hooks,
+ response=initial_response,
+ model=model,
+ messages=(
+ input
+ if isinstance(input, list)
+ else [{"role": "user", "content": input}]
+ ),
+ anthropic_messages_provider_config=responses_api_provider_config,
+ anthropic_messages_optional_request_params=response_api_optional_request_params,
+ logging_obj=logging_obj,
+ stream=False,
+ custom_llm_provider=custom_llm_provider,
+ kwargs=dict(litellm_params),
+ api_surface="responses",
+ )
+ return final_response if final_response is not None else initial_response
+
return initial_response
async def async_response_api_handler(
@@ -4772,22 +4934,9 @@ class BaseLLMHTTPHandler:
agentic callback is detected too.
"""
from litellm.integrations.custom_logger import CustomLogger
- from litellm.litellm_core_utils.litellm_logging import (
- get_custom_logger_compatible_class,
- )
base_func = CustomLogger.async_should_run_agentic_loop
- callbacks = litellm.callbacks + (
- getattr(logging_obj, "dynamic_success_callbacks", None) or []
- )
- for cb in callbacks:
- if isinstance(cb, str):
- resolved = get_custom_logger_compatible_class(cb) # type: ignore[arg-type]
- if resolved is None:
- continue
- cb = resolved
- if not isinstance(cb, CustomLogger):
- continue
+ for cb in _custom_logger_callbacks(logging_obj):
cb_func = getattr(type(cb), "async_should_run_agentic_loop", base_func)
if getattr(cb_func, "__func__", cb_func) is not getattr(
base_func, "__func__", base_func
@@ -5537,9 +5686,7 @@ class BaseLLMHTTPHandler:
import websockets
from websockets.asyncio.client import ClientConnection
- url = self._append_query_params(
- provider_config.get_complete_url(api_base, model, api_key), query_params
- )
+ url = provider_config.get_complete_url(api_base, model, api_key)
headers = provider_config.validate_environment(
headers=headers,
model=model,
diff --git a/litellm/llms/dataforseo/search/transformation.py b/litellm/llms/dataforseo/search/transformation.py
index 27c10d740b5..701db586b72 100644
--- a/litellm/llms/dataforseo/search/transformation.py
+++ b/litellm/llms/dataforseo/search/transformation.py
@@ -61,9 +61,18 @@ class DataForSEOSearchConfig(BaseSearchConfig):
password = get_secret_str("DATAFORSEO_PASSWORD")
# If api_key is provided in "login:password" format, use it
+ caller_supplied_credentials = bool(api_key and ":" in api_key)
if api_key and ":" in api_key:
login, password = api_key.split(":", 1)
+ if not caller_supplied_credentials and login and password:
+ self._assert_trusted_api_base_for_server_credential(
+ api_base,
+ self.DATAFORSEO_API_BASE,
+ "DATAFORSEO_API_BASE",
+ "DATAFORSEO_LOGIN",
+ )
+
if not login:
raise ValueError(
"DATAFORSEO_LOGIN is not set. Set `DATAFORSEO_LOGIN` environment variable or pass credentials in api_key parameter."
diff --git a/litellm/llms/e2b/sandbox/transformation.py b/litellm/llms/e2b/sandbox/transformation.py
index c279fab22ab..ecfc1642c97 100644
--- a/litellm/llms/e2b/sandbox/transformation.py
+++ b/litellm/llms/e2b/sandbox/transformation.py
@@ -16,6 +16,7 @@ from litellm.llms.base_llm.sandbox.transformation import (
BaseSandboxConfig,
CodeExecutionResult,
ContainerHandle,
+ SANDBOX_MAX_OUTPUT_BYTES,
)
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
@@ -29,7 +30,7 @@ E2B_DEFAULT_TEMPLATE = "code-interpreter-v1"
E2B_DEFAULT_DOMAIN = "e2b.app"
JUPYTER_PORT = 49999
DEFAULT_SANDBOX_TIMEOUT = 300
-MAX_OUTPUT_BYTES = 10 * 1024 * 1024
+MAX_OUTPUT_BYTES = SANDBOX_MAX_OUTPUT_BYTES
class E2BSandboxConfig(BaseSandboxConfig):
@@ -49,7 +50,7 @@ class E2BSandboxConfig(BaseSandboxConfig):
*,
template: str | None = None,
timeout: int | None = None,
- allow_internet_access: bool = True,
+ allow_internet_access: bool | None = None,
api_key: str | None = None,
api_base: str | None = None,
metadata: dict | None = None,
@@ -62,7 +63,9 @@ class E2BSandboxConfig(BaseSandboxConfig):
"templateID": template or E2B_DEFAULT_TEMPLATE,
"timeout": timeout if timeout is not None else DEFAULT_SANDBOX_TIMEOUT,
"secure": True,
- "allow_internet_access": allow_internet_access,
+ "allow_internet_access": (
+ True if allow_internet_access is None else allow_internet_access
+ ),
}
if metadata:
body["metadata"] = metadata
@@ -168,20 +171,6 @@ class E2BSandboxConfig(BaseSandboxConfig):
handle._hidden_params = {}
return handle
- @staticmethod
- async def _read_capped_lines(response: httpx.Response) -> list[str]:
- lines: list[str] = []
- total = 0
- async for line in response.aiter_lines():
- total += len(line.encode("utf-8"))
- if total > MAX_OUTPUT_BYTES:
- raise ValueError(
- f"Sandbox output exceeded {MAX_OUTPUT_BYTES} bytes; aborting to "
- "avoid unbounded memory use."
- )
- lines.append(line)
- return lines
-
@staticmethod
def _parse_lines(lines: list[str]) -> CodeExecutionResult:
def _try_parse(stripped: str):
@@ -192,10 +181,9 @@ class E2BSandboxConfig(BaseSandboxConfig):
messages = tuple(
parsed
- for stripped in (line.strip() for line in lines)
- if stripped
- for parsed in (_try_parse(stripped),)
- if parsed is not None
+ for line in lines
+ if (stripped := line.strip())
+ if (parsed := _try_parse(stripped)) is not None
)
def of_type(message_type: str):
diff --git a/litellm/llms/exa_ai/search/transformation.py b/litellm/llms/exa_ai/search/transformation.py
index 7a34ededa6b..5cfd14aeaa9 100644
--- a/litellm/llms/exa_ai/search/transformation.py
+++ b/litellm/llms/exa_ai/search/transformation.py
@@ -65,7 +65,13 @@ class ExaAISearchConfig(BaseSearchConfig):
"""
Validate environment and return headers.
"""
- api_key = api_key or get_secret_str("EXA_API_KEY")
+ api_key = self.resolve_server_api_key(
+ caller_api_key=api_key,
+ caller_api_base=api_base,
+ key_env_vars=("EXA_API_KEY",),
+ base_env_var="EXA_API_BASE",
+ default_api_base=self.EXA_AI_API_BASE,
+ )
if not api_key:
raise ValueError(
"EXA_API_KEY is not set. Set `EXA_API_KEY` environment variable."
diff --git a/litellm/llms/fastcrw/search/transformation.py b/litellm/llms/fastcrw/search/transformation.py
index ce702266e7b..b571a659cac 100644
--- a/litellm/llms/fastcrw/search/transformation.py
+++ b/litellm/llms/fastcrw/search/transformation.py
@@ -57,7 +57,13 @@ class FastCRWSearchConfig(BaseSearchConfig):
"""
Validate environment and return headers.
"""
- api_key = api_key or get_secret_str("CRW_API_KEY")
+ api_key = self.resolve_server_api_key(
+ caller_api_key=api_key,
+ caller_api_base=api_base,
+ key_env_vars=("CRW_API_KEY",),
+ base_env_var="CRW_API_BASE",
+ default_api_base=self.FASTCRW_API_BASE,
+ )
if not api_key:
raise ValueError(
"CRW_API_KEY is not set. Set `CRW_API_KEY` environment variable."
diff --git a/litellm/llms/firecrawl/search/transformation.py b/litellm/llms/firecrawl/search/transformation.py
index 18cf1d28c4d..7e01ba58706 100644
--- a/litellm/llms/firecrawl/search/transformation.py
+++ b/litellm/llms/firecrawl/search/transformation.py
@@ -61,7 +61,13 @@ class FirecrawlSearchConfig(BaseSearchConfig):
"""
Validate environment and return headers.
"""
- api_key = api_key or get_secret_str("FIRECRAWL_API_KEY")
+ api_key = self.resolve_server_api_key(
+ caller_api_key=api_key,
+ caller_api_base=api_base,
+ key_env_vars=("FIRECRAWL_API_KEY",),
+ base_env_var="FIRECRAWL_API_BASE",
+ default_api_base=self.FIRECRAWL_API_BASE,
+ )
if not api_key:
raise ValueError(
"FIRECRAWL_API_KEY is not set. Set `FIRECRAWL_API_KEY` environment variable."
diff --git a/litellm/llms/fireworks_ai/audio_transcription/transformation.py b/litellm/llms/fireworks_ai/audio_transcription/transformation.py
deleted file mode 100644
index 00bb5f26797..00000000000
--- a/litellm/llms/fireworks_ai/audio_transcription/transformation.py
+++ /dev/null
@@ -1,17 +0,0 @@
-from typing import List
-
-from litellm.types.llms.openai import OpenAIAudioTranscriptionOptionalParams
-
-from ...openai.transcriptions.whisper_transformation import (
- OpenAIWhisperAudioTranscriptionConfig,
-)
-from ..common_utils import FireworksAIMixin
-
-
-class FireworksAIAudioTranscriptionConfig(
- FireworksAIMixin, OpenAIWhisperAudioTranscriptionConfig
-):
- def get_supported_openai_params(
- self, model: str
- ) -> List[OpenAIAudioTranscriptionOptionalParams]:
- return ["language", "prompt", "response_format", "timestamp_granularities"]
diff --git a/litellm/llms/gemini/realtime/transformation.py b/litellm/llms/gemini/realtime/transformation.py
index 74f6cd4d831..e153d00e6ab 100644
--- a/litellm/llms/gemini/realtime/transformation.py
+++ b/litellm/llms/gemini/realtime/transformation.py
@@ -103,6 +103,10 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
# bypassing spend and budget accounting.
self._pending_usage_metadata: Optional[dict] = None
+ def _include_function_response_id(self) -> bool:
+ """Google AI Studio Gemini 3.5+ accepts ``id`` on functionResponses; Vertex AI rejects it."""
+ return True
+
@staticmethod
def _usage_detail_alias(details: Any, defaults: Dict[str, int]) -> Dict[str, Any]:
if not isinstance(details, dict):
@@ -604,10 +608,9 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
)
# Build Gemini toolResponse format
- function_response = {
- "id": call_id,
- "response": output_dict,
- }
+ function_response: dict[str, Any] = {"response": output_dict}
+ if self._include_function_response_id() and call_id:
+ function_response["id"] = call_id
if function_name:
function_response["name"] = function_name
diff --git a/litellm/llms/google_pse/search/transformation.py b/litellm/llms/google_pse/search/transformation.py
index a8aa109cbf0..5cd3f2085a8 100644
--- a/litellm/llms/google_pse/search/transformation.py
+++ b/litellm/llms/google_pse/search/transformation.py
@@ -85,7 +85,13 @@ class GooglePSESearchConfig(BaseSearchConfig):
Google PSE uses API key as a query parameter, not in headers.
This method is called but headers are not used for authentication.
"""
- api_key = api_key or get_secret_str("GOOGLE_PSE_API_KEY")
+ api_key = self.resolve_server_api_key(
+ caller_api_key=api_key,
+ caller_api_base=api_base,
+ key_env_vars=("GOOGLE_PSE_API_KEY",),
+ base_env_var="GOOGLE_PSE_API_BASE",
+ default_api_base=self.GOOGLE_PSE_API_BASE,
+ )
if not api_key:
raise ValueError(
"GOOGLE_PSE_API_KEY is not set. Set `GOOGLE_PSE_API_KEY` environment variable."
@@ -137,6 +143,7 @@ class GooglePSESearchConfig(BaseSearchConfig):
query: Union[str, List[str]],
optional_params: dict,
api_key: Optional[str] = None,
+ api_base: str | None = None,
search_engine_id: Optional[str] = None,
**kwargs,
) -> Dict:
@@ -165,8 +172,16 @@ class GooglePSESearchConfig(BaseSearchConfig):
# Google PSE only supports single string queries
query = " ".join(query)
- # Get API credentials
- api_key = api_key or get_secret_str("GOOGLE_PSE_API_KEY")
+ # Get API credentials. The key is sent as a query param to api_base, so
+ # resolve it host-aware to avoid leaking a server-managed key to a
+ # caller-supplied host.
+ api_key = self.resolve_server_api_key(
+ caller_api_key=api_key,
+ caller_api_base=api_base,
+ key_env_vars=("GOOGLE_PSE_API_KEY",),
+ base_env_var="GOOGLE_PSE_API_BASE",
+ default_api_base=self.GOOGLE_PSE_API_BASE,
+ )
search_engine_id = search_engine_id or get_secret_str("GOOGLE_PSE_ENGINE_ID")
if not api_key:
diff --git a/litellm/llms/linkup/search/transformation.py b/litellm/llms/linkup/search/transformation.py
index 2b17d5642ac..d27ae038f9e 100644
--- a/litellm/llms/linkup/search/transformation.py
+++ b/litellm/llms/linkup/search/transformation.py
@@ -61,7 +61,13 @@ class LinkupSearchConfig(BaseSearchConfig):
"""
Validate environment and return headers.
"""
- api_key = api_key or get_secret_str("LINKUP_API_KEY")
+ api_key = self.resolve_server_api_key(
+ caller_api_key=api_key,
+ caller_api_base=api_base,
+ key_env_vars=("LINKUP_API_KEY",),
+ base_env_var="LINKUP_API_BASE",
+ default_api_base=self.LINKUP_API_BASE,
+ )
if not api_key:
raise ValueError(
"LINKUP_API_KEY is not set. Set `LINKUP_API_KEY` environment variable."
diff --git a/litellm/llms/mistral/chat/transformation.py b/litellm/llms/mistral/chat/transformation.py
index f1ad3708236..8d0cf993814 100644
--- a/litellm/llms/mistral/chat/transformation.py
+++ b/litellm/llms/mistral/chat/transformation.py
@@ -247,6 +247,8 @@ class MistralConfig(OpenAIGPTConfig):
The above statement is not valid now. Need to plan to remove all the #1,2,3
Mistral API supports content as a list.
"""
+ messages = [self._strip_output_only_fields(m) for m in messages]
+
## 1. If 'image_url' or 'file' in content, then transform with base class and mistral-specific handling
for m in messages:
_content_block = m.get("content")
@@ -409,6 +411,25 @@ class MistralConfig(OpenAIGPTConfig):
return cleaned_tools
+ @classmethod
+ def _strip_output_only_fields(cls, message: AllMessageValues) -> AllMessageValues:
+ """
+ ``reasoning_content`` and ``thinking_blocks`` are output-only fields that
+ LiteLLM attaches to assistant responses. Mistral's input schema forbids
+ unknown fields, so replaying them verbatim in a follow-up turn triggers a
+ 422 ``extra_forbidden``. Drop them before the request is sent.
+ """
+ if message["role"] != "assistant":
+ return message
+ return cast(
+ AllMessageValues,
+ {
+ k: v
+ for k, v in message.items()
+ if k not in ("reasoning_content", "thinking_blocks")
+ },
+ )
+
@classmethod
def _handle_name_in_message(cls, message: AllMessageValues) -> AllMessageValues:
"""
diff --git a/litellm/llms/openai_like/providers.json b/litellm/llms/openai_like/providers.json
index 24943563937..d87346fea70 100644
--- a/litellm/llms/openai_like/providers.json
+++ b/litellm/llms/openai_like/providers.json
@@ -115,6 +115,14 @@
"max_completion_tokens": "max_tokens"
}
},
+ "darkbloom": {
+ "base_url": "https://api.darkbloom.dev/v1",
+ "api_key_env": "DARKBLOOM_API_KEY",
+ "api_base_env": "DARKBLOOM_API_BASE",
+ "param_mappings": {
+ "max_completion_tokens": "max_tokens"
+ }
+ },
"neosantara": {
"base_url": "https://api.neosantara.xyz/v1",
"api_key_env": "NEOSANTARA_API_KEY",
diff --git a/litellm/llms/opensandbox/__init__.py b/litellm/llms/opensandbox/__init__.py
new file mode 100644
index 00000000000..8b137891791
--- /dev/null
+++ b/litellm/llms/opensandbox/__init__.py
@@ -0,0 +1 @@
+
diff --git a/litellm/llms/opensandbox/sandbox/__init__.py b/litellm/llms/opensandbox/sandbox/__init__.py
new file mode 100644
index 00000000000..8b137891791
--- /dev/null
+++ b/litellm/llms/opensandbox/sandbox/__init__.py
@@ -0,0 +1 @@
+
diff --git a/litellm/llms/opensandbox/sandbox/transformation.py b/litellm/llms/opensandbox/sandbox/transformation.py
new file mode 100644
index 00000000000..dc9f8440d30
--- /dev/null
+++ b/litellm/llms/opensandbox/sandbox/transformation.py
@@ -0,0 +1,598 @@
+import asyncio
+import json
+import time
+from typing import Union, cast
+
+import httpx
+
+from litellm.constants import (
+ OPEN_SANDBOX_API_BASE_ENV_VAR,
+ OPEN_SANDBOX_API_KEY_ENV_VAR,
+ OPEN_SANDBOX_DEFAULT_CPU_LIMIT,
+ OPEN_SANDBOX_DEFAULT_ENTRYPOINT,
+ OPEN_SANDBOX_DEFAULT_LANGUAGE,
+ OPEN_SANDBOX_DEFAULT_MEMORY_LIMIT,
+ OPEN_SANDBOX_DEFAULT_TEMPLATE,
+ OPEN_SANDBOX_DEFAULT_TIMEOUT,
+ OPEN_SANDBOX_EXECD_PORT,
+ OPEN_SANDBOX_POLL_INTERVAL,
+ OPEN_SANDBOX_READY_TIMEOUT,
+)
+from litellm.llms.base_llm.sandbox.transformation import (
+ BaseSandboxConfig,
+ CodeExecutionResult,
+ ContainerHandle,
+ SANDBOX_MAX_OUTPUT_BYTES,
+)
+from litellm.llms.custom_httpx.http_handler import (
+ AsyncHTTPHandler,
+ get_async_httpx_client,
+)
+from litellm.secret_managers.main import get_secret_str
+from litellm.types.llms.custom_http import httpxSpecialProvider
+
+DEFAULT_SANDBOX_TIMEOUT = OPEN_SANDBOX_DEFAULT_TIMEOUT
+DEFAULT_READY_TIMEOUT = OPEN_SANDBOX_READY_TIMEOUT
+DEFAULT_POLL_INTERVAL = OPEN_SANDBOX_POLL_INTERVAL
+MAX_OUTPUT_BYTES = SANDBOX_MAX_OUTPUT_BYTES
+
+
+class OpenSandboxSandboxConfig(BaseSandboxConfig):
+ def _http(self, client: AsyncHTTPHandler | None) -> AsyncHTTPHandler:
+ if client is not None:
+ return client
+ return get_async_httpx_client(llm_provider=httpxSpecialProvider.Sandbox)
+
+ def validate_environment(self, api_key: str | None = None, **kwargs) -> str:
+ if api_key is not None:
+ return api_key
+ return get_secret_str(OPEN_SANDBOX_API_KEY_ENV_VAR) or ""
+
+ async def acreate_sandbox(
+ self,
+ *,
+ template: str | None = None,
+ timeout: int | None = None,
+ allow_internet_access: bool | None = None,
+ api_key: str | None = None,
+ api_base: str | None = None,
+ metadata: dict[str, str] | None = None,
+ env_vars: dict[str, str] | None = None,
+ resource_limits: dict[str, str] | None = None,
+ resource_requests: dict[str, str] | None = None,
+ entrypoint: list[str] | tuple[str, ...] | None = None,
+ network_policy: dict[str, object] | None = None,
+ secure_access: bool = False,
+ use_server_proxy: bool = False,
+ ready_timeout: float | None = None,
+ poll_interval: float | None = None,
+ client: AsyncHTTPHandler | None = None,
+ **kwargs,
+ ) -> ContainerHandle:
+ key = self.validate_environment(api_key=api_key)
+ base = self._api_base(api_base)
+ ready_timeout_seconds = (
+ float(ready_timeout) if ready_timeout is not None else DEFAULT_READY_TIMEOUT
+ )
+ poll_interval_seconds = (
+ float(poll_interval) if poll_interval is not None else DEFAULT_POLL_INTERVAL
+ )
+ body = self._create_body(
+ template=template,
+ timeout=timeout,
+ allow_internet_access=allow_internet_access,
+ metadata=metadata,
+ env_vars=env_vars,
+ resource_limits=resource_limits,
+ resource_requests=resource_requests,
+ entrypoint=entrypoint,
+ network_policy=network_policy,
+ secure_access=secure_access,
+ )
+
+ response = cast(
+ httpx.Response,
+ await self._http(client).post(
+ url=f"{base}/sandboxes",
+ headers=self._lifecycle_headers(key),
+ json=body,
+ ),
+ )
+ data = response.json()
+ sandbox_id = str(data["id"])
+
+ if self._sandbox_state(data) != "Running":
+ await self._wait_until_running(
+ sandbox_id=sandbox_id,
+ api_base=base,
+ headers=self._lifecycle_headers(key),
+ client=client,
+ ready_timeout=ready_timeout_seconds,
+ poll_interval=poll_interval_seconds,
+ )
+
+ endpoint, endpoint_headers = await self._wait_for_execd_endpoint(
+ sandbox_id=sandbox_id,
+ api_base=base,
+ headers=self._lifecycle_headers(key),
+ use_server_proxy=use_server_proxy,
+ client=client,
+ ready_timeout=ready_timeout_seconds,
+ poll_interval=poll_interval_seconds,
+ )
+
+ handle = ContainerHandle(id=sandbox_id, provider="opensandbox", domain=base)
+ handle._hidden_params = {
+ "api_base": base,
+ "api_key": key,
+ "execd_endpoint": endpoint,
+ "execd_headers": endpoint_headers,
+ "use_server_proxy": use_server_proxy,
+ }
+ return handle
+
+ async def arun_code(
+ self,
+ *,
+ container: Union[ContainerHandle, str],
+ code: str,
+ api_key: str | None = None,
+ api_base: str | None = None,
+ language: str = OPEN_SANDBOX_DEFAULT_LANGUAGE,
+ use_server_proxy: bool = False,
+ ready_timeout: float | None = None,
+ poll_interval: float | None = None,
+ client: AsyncHTTPHandler | None = None,
+ **kwargs,
+ ) -> CodeExecutionResult:
+ handle = await self._ensure_handle(
+ container=container,
+ api_key=api_key,
+ api_base=api_base,
+ use_server_proxy=use_server_proxy,
+ ready_timeout=(
+ float(ready_timeout)
+ if ready_timeout is not None
+ else DEFAULT_READY_TIMEOUT
+ ),
+ poll_interval=(
+ float(poll_interval)
+ if poll_interval is not None
+ else DEFAULT_POLL_INTERVAL
+ ),
+ client=client,
+ )
+ endpoint = str(handle._hidden_params["execd_endpoint"])
+ endpoint_headers = self._as_str_dict(handle._hidden_params.get("execd_headers"))
+ base = str(
+ handle._hidden_params.get("api_base")
+ or handle.domain
+ or self._api_base(api_base)
+ )
+ lines = await self._post_code(
+ url=f"{self._endpoint_base_url(endpoint, base)}/code",
+ headers={
+ "Content-Type": "application/json",
+ "Accept": "text/event-stream",
+ "Cache-Control": "no-cache",
+ **endpoint_headers,
+ },
+ body={
+ "code": code,
+ "context": {"language": language},
+ },
+ client=client,
+ )
+ return self._parse_lines(lines)
+
+ async def adelete_sandbox(
+ self,
+ *,
+ container: Union[ContainerHandle, str],
+ api_key: str | None = None,
+ api_base: str | None = None,
+ client: AsyncHTTPHandler | None = None,
+ **kwargs,
+ ) -> bool:
+ handle = self._as_handle(container, api_base=api_base)
+ base = str(handle._hidden_params.get("api_base") or self._api_base(api_base))
+ key = self._api_key(api_key=api_key, handle=handle)
+ try:
+ response = cast(
+ httpx.Response,
+ await self._http(client).delete(
+ url=f"{base}/sandboxes/{handle.id}",
+ headers=self._lifecycle_headers(key),
+ ),
+ )
+ except httpx.HTTPStatusError as e:
+ if e.response.status_code == 404:
+ return False
+ raise
+ return 200 <= response.status_code < 300
+
+ async def _ensure_handle(
+ self,
+ *,
+ container: Union[ContainerHandle, str],
+ api_key: str | None,
+ api_base: str | None,
+ use_server_proxy: bool,
+ ready_timeout: float,
+ poll_interval: float,
+ client: AsyncHTTPHandler | None,
+ ) -> ContainerHandle:
+ handle = self._as_handle(container, api_base=api_base)
+ if handle._hidden_params.get("execd_endpoint"):
+ return handle
+
+ base = str(handle._hidden_params.get("api_base") or self._api_base(api_base))
+ key = self._api_key(api_key=api_key, handle=handle)
+ resolved_use_server_proxy = bool(
+ handle._hidden_params.get("use_server_proxy", use_server_proxy)
+ )
+ endpoint, endpoint_headers = await self._wait_for_execd_endpoint(
+ sandbox_id=handle.id,
+ api_base=base,
+ headers=self._lifecycle_headers(key),
+ use_server_proxy=resolved_use_server_proxy,
+ client=client,
+ ready_timeout=ready_timeout,
+ poll_interval=poll_interval,
+ )
+ handle.domain = base
+ handle._hidden_params = {
+ **handle._hidden_params,
+ "api_base": base,
+ "api_key": key,
+ "execd_endpoint": endpoint,
+ "execd_headers": endpoint_headers,
+ "use_server_proxy": resolved_use_server_proxy,
+ }
+ return handle
+
+ async def _wait_until_running(
+ self,
+ *,
+ sandbox_id: str,
+ api_base: str,
+ headers: dict[str, str],
+ client: AsyncHTTPHandler | None,
+ ready_timeout: float,
+ poll_interval: float,
+ ) -> None:
+ deadline = time.monotonic() + ready_timeout
+ while True:
+ response = cast(
+ httpx.Response,
+ await self._http(client).get(
+ url=f"{api_base}/sandboxes/{sandbox_id}",
+ headers=headers,
+ ),
+ )
+ data = response.json()
+ state = self._sandbox_state(data)
+ if state == "Running":
+ return
+ if state in {"Failed", "Stopping", "Terminated"}:
+ raise ValueError(f"OpenSandbox sandbox {sandbox_id} entered {state}")
+ if time.monotonic() >= deadline:
+ raise TimeoutError(
+ f"OpenSandbox sandbox {sandbox_id} was not Running within "
+ f"{ready_timeout} seconds"
+ )
+ await asyncio.sleep(poll_interval)
+
+ async def _wait_for_execd_endpoint(
+ self,
+ *,
+ sandbox_id: str,
+ api_base: str,
+ headers: dict[str, str],
+ use_server_proxy: bool,
+ client: AsyncHTTPHandler | None,
+ ready_timeout: float,
+ poll_interval: float,
+ ) -> tuple[str, dict[str, str]]:
+ deadline = time.monotonic() + ready_timeout
+ last_error: Exception | None = None
+ while True:
+ try:
+ return await self._get_execd_endpoint(
+ sandbox_id=sandbox_id,
+ api_base=api_base,
+ headers=headers,
+ use_server_proxy=use_server_proxy,
+ client=client,
+ )
+ except httpx.HTTPStatusError as e:
+ if e.response.status_code != 404:
+ raise
+ last_error = e
+ except ValueError as e:
+ last_error = e
+
+ if time.monotonic() >= deadline:
+ raise TimeoutError(
+ f"OpenSandbox execd endpoint for {sandbox_id} was not ready within "
+ f"{ready_timeout} seconds"
+ ) from last_error
+ await asyncio.sleep(poll_interval)
+
+ async def _get_execd_endpoint(
+ self,
+ *,
+ sandbox_id: str,
+ api_base: str,
+ headers: dict[str, str],
+ use_server_proxy: bool,
+ client: AsyncHTTPHandler | None,
+ ) -> tuple[str, dict[str, str]]:
+ response = cast(
+ httpx.Response,
+ await self._http(client).get(
+ url=f"{api_base}/sandboxes/{sandbox_id}/endpoints/{OPEN_SANDBOX_EXECD_PORT}",
+ headers=headers,
+ params={"use_server_proxy": use_server_proxy},
+ ),
+ )
+ data = response.json()
+ endpoint = data.get("endpoint")
+ if not endpoint:
+ raise ValueError(
+ f"OpenSandbox did not return an execd endpoint for {sandbox_id}"
+ )
+ return str(endpoint), self._as_str_dict(data.get("headers"))
+
+ async def _post_code(
+ self,
+ *,
+ url: str,
+ headers: dict[str, str],
+ body: dict[str, object],
+ client: AsyncHTTPHandler | None,
+ ) -> list[str]:
+ timeout = httpx.Timeout(connect=30.0, read=None, write=30.0, pool=None)
+ response = cast(
+ httpx.Response,
+ await self._http(client).post(
+ url=url,
+ headers=headers,
+ timeout=timeout,
+ json=body,
+ stream=True,
+ ),
+ )
+ return await self._read_capped_lines(response)
+
+ def _api_key(self, *, api_key: str | None, handle: ContainerHandle) -> str:
+ if api_key is not None:
+ return api_key
+ if "api_key" in handle._hidden_params:
+ return str(handle._hidden_params["api_key"])
+ return self.validate_environment()
+
+ @staticmethod
+ def _create_body(
+ *,
+ template: str | None,
+ timeout: int | None,
+ allow_internet_access: bool | None,
+ metadata: dict[str, str] | None,
+ env_vars: dict[str, str] | None,
+ resource_limits: dict[str, str] | None,
+ resource_requests: dict[str, str] | None,
+ entrypoint: list[str] | tuple[str, ...] | None,
+ network_policy: dict[str, object] | None,
+ secure_access: bool,
+ ) -> dict[str, object]:
+ body: dict[str, object] = {
+ "image": {"uri": template or OPEN_SANDBOX_DEFAULT_TEMPLATE},
+ "entrypoint": list(entrypoint or OPEN_SANDBOX_DEFAULT_ENTRYPOINT),
+ "timeout": timeout if timeout is not None else DEFAULT_SANDBOX_TIMEOUT,
+ "resourceLimits": resource_limits
+ or OpenSandboxSandboxConfig._default_resource_limits(),
+ }
+ if metadata:
+ body["metadata"] = metadata
+ if env_vars:
+ body["env"] = env_vars
+ if resource_requests:
+ body["resourceRequests"] = resource_requests
+ if network_policy is not None:
+ body["networkPolicy"] = network_policy
+ elif allow_internet_access is not True:
+ body["networkPolicy"] = {"defaultAction": "deny", "egress": []}
+ if secure_access:
+ body["secureAccess"] = True
+ return body
+
+ @staticmethod
+ def _default_resource_limits() -> dict[str, str]:
+ return {
+ "cpu": OPEN_SANDBOX_DEFAULT_CPU_LIMIT,
+ "memory": OPEN_SANDBOX_DEFAULT_MEMORY_LIMIT,
+ }
+
+ @staticmethod
+ def _sandbox_state(data: object) -> str | None:
+ if not isinstance(data, dict):
+ return None
+ status = data.get("status")
+ if not isinstance(status, dict):
+ return None
+ state = status.get("state")
+ return str(state) if state is not None else None
+
+ @staticmethod
+ def _as_str_dict(value: object) -> dict[str, str]:
+ if not isinstance(value, dict):
+ return {}
+ return {str(k): str(v) for k, v in value.items()}
+
+ @staticmethod
+ def _api_base(api_base: str | None) -> str:
+ base = api_base or get_secret_str(OPEN_SANDBOX_API_BASE_ENV_VAR)
+ if not base:
+ raise ValueError(
+ "OpenSandbox api_base is required. Pass api_base or set "
+ f"{OPEN_SANDBOX_API_BASE_ENV_VAR}."
+ )
+ return str(base).rstrip("/")
+
+ @staticmethod
+ def _lifecycle_headers(api_key: str) -> dict[str, str]:
+ headers = {"Content-Type": "application/json"}
+ if api_key:
+ headers["OPEN-SANDBOX-API-KEY"] = api_key
+ return headers
+
+ @staticmethod
+ def _endpoint_base_url(endpoint: str, api_base: str) -> str:
+ normalized_endpoint = endpoint.rstrip("/")
+ if normalized_endpoint.startswith(("http://", "https://")):
+ return normalized_endpoint
+ protocol = api_base.split("://", 1)[0] if "://" in api_base else "http"
+ return f"{protocol}://{normalized_endpoint}"
+
+ @staticmethod
+ def _as_handle(
+ container: Union[ContainerHandle, str], *, api_base: str | None
+ ) -> ContainerHandle:
+ if isinstance(container, ContainerHandle):
+ return container
+ handle = ContainerHandle(
+ id=str(container),
+ provider="opensandbox",
+ domain=OpenSandboxSandboxConfig._api_base(api_base),
+ )
+ handle._hidden_params = {}
+ return handle
+
+ @staticmethod
+ def _parse_lines(lines: list[str]) -> CodeExecutionResult:
+ messages = tuple(
+ event
+ for line in lines
+ if (event := OpenSandboxSandboxConfig._parse_sse_line(line)) is not None
+ )
+
+ def of_type(message_type: str):
+ return (m for m in messages if m.get("type") == message_type)
+
+ error = next(
+ (OpenSandboxSandboxConfig._normalize_error(m) for m in of_type("error")),
+ None,
+ )
+ execution_count = next(
+ (
+ OpenSandboxSandboxConfig._as_int(m.get("execution_count"))
+ for m in of_type("execution_count")
+ if OpenSandboxSandboxConfig._as_int(m.get("execution_count"))
+ is not None
+ ),
+ None,
+ )
+
+ return CodeExecutionResult(
+ stdout="".join(str(m.get("text", "")) for m in of_type("stdout")),
+ stderr="".join(str(m.get("text", "")) for m in of_type("stderr")),
+ results=[
+ OpenSandboxSandboxConfig._normalize_result(m) for m in of_type("result")
+ ],
+ error=error,
+ execution_count=execution_count,
+ )
+
+ @staticmethod
+ def _parse_sse_line(line: str) -> dict[str, object] | None:
+ stripped = line.strip()
+ if not stripped or stripped.startswith(
+ (
+ ":",
+ "event:",
+ "id:",
+ "retry:",
+ )
+ ):
+ return None
+ data = stripped[5:].strip() if stripped.startswith("data:") else stripped
+ if not data:
+ return None
+ try:
+ parsed = json.loads(data)
+ except json.JSONDecodeError:
+ return None
+ if not isinstance(parsed, dict):
+ return None
+ if "type" not in parsed and "code" in parsed and "message" in parsed:
+ return {
+ "type": "error",
+ "error": {
+ "ename": str(parsed["code"]),
+ "evalue": str(parsed["message"]),
+ "traceback": [],
+ },
+ }
+ return parsed
+
+ @staticmethod
+ def _normalize_result(message: dict[str, object]) -> dict[str, object]:
+ results = message.get("results")
+ if isinstance(results, dict):
+ return {str(k): v for k, v in results.items()}
+ return {
+ str(k): v
+ for k, v in message.items()
+ if k not in {"type", "timestamp", "execution_count"}
+ }
+
+ @staticmethod
+ def _normalize_error(message: dict[str, object]) -> dict[str, object]:
+ raw_error = message.get("error")
+ if isinstance(raw_error, dict):
+ name = OpenSandboxSandboxConfig._first_non_none_value(
+ raw_error, "ename", "name", default=""
+ )
+ value = OpenSandboxSandboxConfig._first_non_none_value(
+ raw_error, "evalue", "value", default=""
+ )
+ traceback = OpenSandboxSandboxConfig._first_non_none_value(
+ raw_error, "traceback", default=[]
+ )
+ return {
+ "name": name,
+ "value": value,
+ "traceback": traceback,
+ }
+ return {
+ "name": OpenSandboxSandboxConfig._first_non_none_value(
+ message, "name", default=""
+ ),
+ "value": OpenSandboxSandboxConfig._first_non_none_value(
+ message, "value", "text", default=""
+ ),
+ "traceback": OpenSandboxSandboxConfig._first_non_none_value(
+ message, "traceback", default=[]
+ ),
+ }
+
+ @staticmethod
+ def _as_int(value: object) -> int | None:
+ if isinstance(value, int):
+ return value
+ if isinstance(value, str):
+ try:
+ return int(value)
+ except ValueError:
+ return None
+ return None
+
+ @staticmethod
+ def _first_non_none_value(
+ values: dict[str, object], *keys: str, default: object
+ ) -> object:
+ return next(
+ (values[key] for key in keys if key in values and values[key] is not None),
+ default,
+ )
diff --git a/litellm/llms/parallel_ai/search/transformation.py b/litellm/llms/parallel_ai/search/transformation.py
index 85602bf1d86..35a0d84df40 100644
--- a/litellm/llms/parallel_ai/search/transformation.py
+++ b/litellm/llms/parallel_ai/search/transformation.py
@@ -67,10 +67,12 @@ class ParallelAISearchConfig(BaseSearchConfig):
api_base: Optional[str] = None,
**kwargs,
) -> Dict:
- api_key = (
- api_key
- or get_secret_str("PARALLEL_AI_API_KEY")
- or get_secret_str("PARALLEL_API_KEY")
+ api_key = self.resolve_server_api_key(
+ caller_api_key=api_key,
+ caller_api_base=api_base,
+ key_env_vars=("PARALLEL_AI_API_KEY", "PARALLEL_API_KEY"),
+ base_env_var="PARALLEL_AI_API_BASE",
+ default_api_base=self.PARALLEL_AI_API_BASE,
)
if not api_key:
raise ValueError(
diff --git a/litellm/llms/perplexity/cost_calculator.py b/litellm/llms/perplexity/cost_calculator.py
index bf055f91aa0..ec7ec397ea6 100644
--- a/litellm/llms/perplexity/cost_calculator.py
+++ b/litellm/llms/perplexity/cost_calculator.py
@@ -98,10 +98,11 @@ def cost_per_token(model: str, usage: Usage) -> Tuple[float, float]:
if num_search_queries > 0 and search_cost_value is not None:
# Handle both dict and float formats
if isinstance(search_cost_value, dict):
- # Use the "low" size as default - tests expect 0.005 / 1000
- search_cost_per_query = (
- _safe_float_cast(search_cost_value.get("search_context_size_low", 0))
- / 1000
+ # search_context_cost_per_query stores the per-request price in USD
+ # (e.g. sonar low = $0.005/request). Use it directly, matching the
+ # gemini cost calculator which reads the same field per request.
+ search_cost_per_query = _safe_float_cast(
+ search_cost_value.get("search_context_size_low", 0)
)
else:
search_cost_per_query = _safe_float_cast(search_cost_value)
diff --git a/litellm/llms/perplexity/search/transformation.py b/litellm/llms/perplexity/search/transformation.py
index ea96f87957c..55de52c5384 100644
--- a/litellm/llms/perplexity/search/transformation.py
+++ b/litellm/llms/perplexity/search/transformation.py
@@ -50,7 +50,13 @@ class PerplexitySearchConfig(BaseSearchConfig):
"""
Validate environment and return headers.
"""
- api_key = api_key or get_secret_str("PERPLEXITYAI_API_KEY")
+ api_key = self.resolve_server_api_key(
+ caller_api_key=api_key,
+ caller_api_base=api_base,
+ key_env_vars=("PERPLEXITYAI_API_KEY",),
+ base_env_var="PERPLEXITY_API_BASE",
+ default_api_base=self.PERPLEXITY_API_BASE,
+ )
if not api_key:
raise ValueError(
"PERPLEXITYAI_API_KEY is not set. Set `PERPLEXITYAI_API_KEY` environment variable."
diff --git a/litellm/llms/searchapi/search/transformation.py b/litellm/llms/searchapi/search/transformation.py
index c04e1377f9c..ae8413684cc 100644
--- a/litellm/llms/searchapi/search/transformation.py
+++ b/litellm/llms/searchapi/search/transformation.py
@@ -74,7 +74,13 @@ class SearchAPIConfig(BaseSearchConfig):
"""
Validate environment and return headers.
"""
- api_key = api_key or get_secret_str("SEARCHAPI_API_KEY")
+ api_key = self.resolve_server_api_key(
+ caller_api_key=api_key,
+ caller_api_base=api_base,
+ key_env_vars=("SEARCHAPI_API_KEY",),
+ base_env_var="SEARCHAPI_API_BASE",
+ default_api_base=self.SEARCHAPI_API_BASE,
+ )
if not api_key:
raise ValueError(
@@ -114,6 +120,7 @@ class SearchAPIConfig(BaseSearchConfig):
query: Union[str, List[str]],
optional_params: dict,
api_key: Optional[str] = None,
+ api_base: str | None = None,
search_engine_id: Optional[str] = None,
**kwargs,
) -> Dict:
@@ -137,8 +144,16 @@ class SearchAPIConfig(BaseSearchConfig):
if isinstance(query, list):
query = " ".join(query)
- # Get API key from parameter or environment
- api_key = api_key or get_secret_str("SEARCHAPI_API_KEY")
+ # Get API key from parameter or environment. The key is sent as a query
+ # param to api_base, so resolve it host-aware to avoid leaking a
+ # server-managed key to a caller-supplied host.
+ api_key = self.resolve_server_api_key(
+ caller_api_key=api_key,
+ caller_api_base=api_base,
+ key_env_vars=("SEARCHAPI_API_KEY",),
+ base_env_var="SEARCHAPI_API_BASE",
+ default_api_base=self.SEARCHAPI_API_BASE,
+ )
if not api_key:
raise ValueError(
"SEARCHAPI_API_KEY is not set. Set `SEARCHAPI_API_KEY` environment variable."
diff --git a/litellm/llms/searxng/search/transformation.py b/litellm/llms/searxng/search/transformation.py
index ee6f3895721..ff68be5709e 100644
--- a/litellm/llms/searxng/search/transformation.py
+++ b/litellm/llms/searxng/search/transformation.py
@@ -61,7 +61,13 @@ class SearXNGSearchConfig(BaseSearchConfig):
Some instances may require authentication via headers.
"""
# SearXNG typically doesn't require API keys, but support optional auth
- api_key = api_key or get_secret_str("SEARXNG_API_KEY")
+ api_key = self.resolve_server_api_key(
+ caller_api_key=api_key,
+ caller_api_base=api_base,
+ key_env_vars=("SEARXNG_API_KEY",),
+ base_env_var="SEARXNG_API_BASE",
+ default_api_base=None,
+ )
if api_key:
headers["Authorization"] = f"Bearer {api_key}"
headers["Content-Type"] = "application/json"
diff --git a/litellm/llms/serper/search/transformation.py b/litellm/llms/serper/search/transformation.py
index 0daccbe652b..dd43f2d2dc9 100644
--- a/litellm/llms/serper/search/transformation.py
+++ b/litellm/llms/serper/search/transformation.py
@@ -55,7 +55,13 @@ class SerperSearchConfig(BaseSearchConfig):
"""
Validate environment and return headers.
"""
- api_key = api_key or get_secret_str("SERPER_API_KEY")
+ api_key = self.resolve_server_api_key(
+ caller_api_key=api_key,
+ caller_api_base=api_base,
+ key_env_vars=("SERPER_API_KEY",),
+ base_env_var="SERPER_API_BASE",
+ default_api_base=self.SERPER_API_BASE,
+ )
if not api_key:
raise ValueError(
"SERPER_API_KEY is not set. Set `SERPER_API_KEY` environment variable."
diff --git a/litellm/llms/tavily/search/transformation.py b/litellm/llms/tavily/search/transformation.py
index ec96db96f36..647cfb5fa84 100644
--- a/litellm/llms/tavily/search/transformation.py
+++ b/litellm/llms/tavily/search/transformation.py
@@ -64,7 +64,13 @@ class TavilySearchConfig(BaseSearchConfig):
"""
Validate environment and return headers.
"""
- api_key = api_key or get_secret_str("TAVILY_API_KEY")
+ api_key = self.resolve_server_api_key(
+ caller_api_key=api_key,
+ caller_api_base=api_base,
+ key_env_vars=("TAVILY_API_KEY",),
+ base_env_var="TAVILY_API_BASE",
+ default_api_base=self.TAVILY_API_BASE,
+ )
if not api_key:
raise ValueError(
"TAVILY_API_KEY is not set. Set `TAVILY_API_KEY` environment variable."
diff --git a/litellm/llms/tinyfish/search/transformation.py b/litellm/llms/tinyfish/search/transformation.py
index c4949380e3a..b92f7ca1aff 100644
--- a/litellm/llms/tinyfish/search/transformation.py
+++ b/litellm/llms/tinyfish/search/transformation.py
@@ -67,7 +67,13 @@ class TinyfishSearchConfig(BaseSearchConfig):
api_base: str | None = None,
**kwargs: object,
) -> dict[str, str]:
- resolved_key = api_key or get_secret_str("TINYFISH_API_KEY")
+ resolved_key = self.resolve_server_api_key(
+ caller_api_key=api_key,
+ caller_api_base=api_base,
+ key_env_vars=("TINYFISH_API_KEY",),
+ base_env_var="TINYFISH_API_BASE",
+ default_api_base=self.TINYFISH_API_BASE,
+ )
if not resolved_key:
raise ValueError(
"TINYFISH_API_KEY is not set. Set `TINYFISH_API_KEY` environment variable."
diff --git a/litellm/llms/vertex_ai/realtime/transformation.py b/litellm/llms/vertex_ai/realtime/transformation.py
index d6441db7856..1fe9f15c9f0 100644
--- a/litellm/llms/vertex_ai/realtime/transformation.py
+++ b/litellm/llms/vertex_ai/realtime/transformation.py
@@ -32,6 +32,9 @@ class VertexAIRealtimeConfig(GeminiRealtimeConfig):
self._project = project
self._location = location
+ def _include_function_response_id(self) -> bool:
+ return False
+
# ------------------------------------------------------------------
# URL
# ------------------------------------------------------------------
diff --git a/litellm/llms/you_com/search/transformation.py b/litellm/llms/you_com/search/transformation.py
index 3c94b991735..0c7916e4c05 100644
--- a/litellm/llms/you_com/search/transformation.py
+++ b/litellm/llms/you_com/search/transformation.py
@@ -64,7 +64,13 @@ class YouComSearchConfig(BaseSearchConfig):
endpoint with the `X-API-Key` header. Otherwise fall through to the
keyless free tier; no auth header is required.
"""
- api_key = api_key or get_secret_str("YOUCOM_API_KEY")
+ api_key = self.resolve_server_api_key(
+ caller_api_key=api_key,
+ caller_api_base=api_base,
+ key_env_vars=("YOUCOM_API_KEY",),
+ base_env_var="YOUCOM_API_BASE",
+ default_api_base=self.YOU_COM_API_BASE,
+ )
headers["Content-Type"] = "application/json"
# Pin Accept-Encoding to identity: the keyless `api.you.com/v1/agents/search`
# endpoint advertises gzip content-encoding but returns body bytes the
diff --git a/litellm/main.py b/litellm/main.py
index 63c5798e70a..c3d7ca28c49 100644
--- a/litellm/main.py
+++ b/litellm/main.py
@@ -81,11 +81,17 @@ from litellm.constants import (
from litellm.exceptions import LiteLLMUnknownProvider
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.asyncify import run_async_function
+from litellm.litellm_core_utils.chat_completion_agentic_loop import (
+ maybe_run_chat_completion_agentic_loop,
+)
from litellm.litellm_core_utils.audio_utils.utils import (
calculate_request_duration,
get_audio_file_for_health_check,
)
from litellm.litellm_core_utils.completion_timeout import CompletionTimeout
+from litellm.litellm_core_utils.request_timeout_resolver import (
+ get_configured_request_timeout,
+)
from litellm.litellm_core_utils.get_litellm_params import OPTIONAL_KWARGS_KEYS
from litellm.litellm_core_utils.dd_tracing import tracer
from litellm.litellm_core_utils.get_provider_specific_headers import (
@@ -118,6 +124,10 @@ from litellm.llms.vertex_ai.common_utils import (
)
from litellm.realtime_api.main import _realtime_health_check
from litellm.secret_managers.main import get_secret_bool, get_secret_str
+from litellm.types.completion import (
+ _CompletionDispatchContext,
+ _CompletionDispatchResult,
+)
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import (
CustomPricingLiteLLMParams,
@@ -650,6 +660,39 @@ async def acompletion(
response_object=response,
model_response_object=litellm.ModelResponse(),
)
+ # Provider-agnostic dispatch point for the chat-completions agentic loop
+ # (code-interpreter interception, etc). Chat routing forks per provider
+ # before this (OpenAI goes through the OpenAI SDK in openai.py, others
+ # through the shared httpx handler), so a dispatch inside any single
+ # provider handler would miss the others. Here is where every fork
+ # reconverges, so the loop runs once for all providers. Responses needs
+ # no equivalent: every provider already funnels through one shared
+ # handler where the loop is dispatched.
+ if isinstance(response, litellm.ModelResponse):
+ looped = await maybe_run_chat_completion_agentic_loop(
+ response=response,
+ model=model,
+ messages=messages,
+ optional_params={
+ k: v
+ for k, v in completion_kwargs.items()
+ if v is not None
+ and k
+ not in (
+ "model",
+ "messages",
+ "stream",
+ "acompletion",
+ "deployment_id",
+ )
+ },
+ kwargs=kwargs,
+ logging_obj=kwargs.get("litellm_logging_obj"),
+ custom_llm_provider=custom_llm_provider,
+ stream=bool(stream),
+ )
+ if looped is not None:
+ response = looped
if isinstance(response, CustomStreamWrapper):
response.set_logging_event_loop(
loop=loop
@@ -1084,6 +1127,3825 @@ def _build_custom_pricing_entry(
return entry
+def _complete_azure(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ _azure_detection_model = ctx._azure_detection_model
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ api_version = ctx.api_version
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ extra_headers = ctx.extra_headers
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logger_fn = ctx.logger_fn
+ logging = ctx.logging
+ max_retries = ctx.max_retries
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ timeout = ctx.timeout
+
+ dynamic_params = False
+ if client is not None and (
+ isinstance(client, openai.AzureOpenAI)
+ or isinstance(client, openai.AsyncAzureOpenAI)
+ ):
+ dynamic_params = _check_dynamic_azure_params(
+ azure_client_params={"api_version": api_version},
+ azure_client=client,
+ )
+
+ api_type = get_secret("AZURE_API_TYPE") or "azure"
+
+ api_base = api_base or litellm.api_base or get_secret("AZURE_API_BASE")
+
+ api_version = (
+ api_version
+ or litellm.api_version
+ or get_secret_str("AZURE_API_VERSION")
+ or litellm.AZURE_DEFAULT_API_VERSION
+ )
+
+ api_key = (
+ api_key
+ or litellm.api_key
+ or litellm.azure_key
+ or get_secret_str("AZURE_OPENAI_API_KEY")
+ or get_secret_str("AZURE_API_KEY")
+ )
+
+ azure_ad_token = optional_params.get("extra_body", {}).pop(
+ "azure_ad_token", None
+ ) or get_secret_str("AZURE_AD_TOKEN")
+
+ azure_ad_token_provider = litellm_params.get("azure_ad_token_provider", None)
+
+ headers = headers or litellm.headers
+
+ if extra_headers is not None:
+ optional_params["extra_headers"] = extra_headers
+ if max_retries is not None:
+ optional_params["max_retries"] = max_retries
+
+ if litellm.AzureOpenAIO1Config().is_o_series_model(model=_azure_detection_model):
+ ## LOAD CONFIG - if set
+ config = litellm.AzureOpenAIO1Config.get_config()
+ for k, v in config.items():
+ if (
+ k not in optional_params
+ ): # completion(top_k=3) > azure_config(top_k=3) <- allows for dynamic variables to be passed in
+ optional_params[k] = v
+
+ response = azure_o1_chat_completions.completion(
+ model=model,
+ messages=messages,
+ headers=headers,
+ api_key=api_key,
+ api_base=api_base,
+ api_version=api_version,
+ dynamic_params=dynamic_params,
+ azure_ad_token=azure_ad_token,
+ model_response=model_response,
+ print_verbose=print_verbose,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ logger_fn=logger_fn,
+ logging_obj=logging,
+ acompletion=acompletion,
+ timeout=timeout, # type: ignore
+ client=client, # pass AsyncAzureOpenAI, AzureOpenAI client
+ custom_llm_provider=custom_llm_provider,
+ )
+ else:
+ ## LOAD CONFIG - if set
+ config = litellm.AzureOpenAIConfig.get_config()
+ for k, v in config.items():
+ if (
+ k not in optional_params
+ ): # completion(top_k=3) > azure_config(top_k=3) <- allows for dynamic variables to be passed in
+ optional_params[k] = v
+
+ ## COMPLETION CALL
+ response = azure_chat_completions.completion(
+ model=model,
+ messages=messages,
+ headers=headers,
+ api_key=api_key,
+ api_base=api_base,
+ api_version=api_version,
+ api_type=api_type,
+ dynamic_params=dynamic_params,
+ azure_ad_token=azure_ad_token,
+ azure_ad_token_provider=azure_ad_token_provider,
+ model_response=model_response,
+ print_verbose=print_verbose,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ logger_fn=logger_fn,
+ logging_obj=logging,
+ acompletion=acompletion,
+ timeout=timeout, # type: ignore
+ client=client, # pass AsyncAzureOpenAI, AzureOpenAI client
+ )
+
+ if optional_params.get("stream", False):
+ ## LOGGING
+ logging.post_call(
+ input=messages,
+ api_key=api_key,
+ original_response=response,
+ additional_args={
+ "headers": headers,
+ "api_version": api_version,
+ "api_base": api_base,
+ },
+ )
+
+ return response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract
+
+
+def _complete_azure_text(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ api_version = ctx.api_version
+ client = ctx.client
+ extra_headers = ctx.extra_headers
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logger_fn = ctx.logger_fn
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ timeout = ctx.timeout
+
+ api_type = get_secret_str("AZURE_API_TYPE") or "azure"
+
+ api_base = api_base or litellm.api_base or get_secret_str("AZURE_API_BASE")
+
+ if api_base is None:
+ raise ValueError(
+ "api_base is required for Azure OpenAI LLM provider. Either set it dynamically or set the AZURE_API_BASE environment variable."
+ )
+
+ api_version = (
+ api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION")
+ )
+
+ api_key = (
+ api_key
+ or litellm.api_key
+ or litellm.azure_key
+ or get_secret_str("AZURE_OPENAI_API_KEY")
+ or get_secret_str("AZURE_API_KEY")
+ )
+
+ azure_ad_token = optional_params.get("extra_body", {}).pop(
+ "azure_ad_token", None
+ ) or get_secret_str("AZURE_AD_TOKEN")
+
+ azure_ad_token_provider = litellm_params.get("azure_ad_token_provider", None)
+
+ headers = headers or litellm.headers
+
+ if extra_headers is not None:
+ optional_params["extra_headers"] = extra_headers
+
+ ## LOAD CONFIG - if set
+ config = litellm.AzureOpenAIConfig.get_config()
+ for k, v in config.items():
+ if (
+ k not in optional_params
+ ): # completion(top_k=3) > azure_config(top_k=3) <- allows for dynamic variables to be passed in
+ optional_params[k] = v
+
+ ## COMPLETION CALL
+ response = azure_text_completions.completion(
+ model=model,
+ messages=messages,
+ headers=headers,
+ api_key=api_key,
+ api_base=api_base,
+ api_version=cast(str, api_version),
+ api_type=api_type,
+ azure_ad_token=azure_ad_token,
+ azure_ad_token_provider=azure_ad_token_provider,
+ model_response=model_response,
+ print_verbose=print_verbose,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ logger_fn=logger_fn,
+ logging_obj=logging,
+ acompletion=acompletion,
+ timeout=timeout,
+ client=client, # pass AsyncAzureOpenAI, AzureOpenAI client
+ )
+
+ if optional_params.get("stream", False) or acompletion is True:
+ ## LOGGING
+ logging.post_call(
+ input=messages,
+ api_key=api_key,
+ original_response=response,
+ additional_args={
+ "headers": headers,
+ "api_version": api_version,
+ "api_base": api_base,
+ },
+ )
+
+ return response
+
+
+def _complete_deepseek(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ provider_config = ctx.provider_config
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ try:
+ response = base_llm_http_handler.completion(
+ model=model,
+ messages=messages,
+ headers=headers,
+ model_response=model_response,
+ api_key=api_key,
+ api_base=api_base,
+ acompletion=acompletion,
+ logging_obj=logging,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ timeout=timeout, # type: ignore
+ client=client,
+ custom_llm_provider=custom_llm_provider,
+ encoding=_get_encoding(),
+ stream=stream,
+ provider_config=provider_config,
+ )
+ except Exception as e:
+ ## LOGGING - log the original exception returned
+ logging.post_call(
+ input=messages,
+ api_key=api_key,
+ original_response=str(e),
+ additional_args={"headers": headers},
+ )
+ raise e
+
+ return response
+
+
+def _complete_azure_ai(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ extra_headers = ctx.extra_headers
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logger_fn = ctx.logger_fn
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ from litellm.llms.azure_ai.common_utils import AzureFoundryModelInfo
+
+ azure_ai_route = AzureFoundryModelInfo.get_azure_ai_route(model)
+
+ # Check if this is an agents route - model format: azure_ai/agents/
+ if azure_ai_route == "agents":
+ from litellm.llms.azure_ai.agents import AzureAIAgentsConfig
+
+ api_base = AzureFoundryModelInfo.get_api_base(api_base)
+ if api_base is None:
+ raise ValueError(
+ "Azure AI Agents requests require an api_base. "
+ "Set `api_base` or the AZURE_AI_API_BASE env var."
+ )
+ api_key = AzureFoundryModelInfo.get_api_key(api_key)
+
+ response = AzureAIAgentsConfig.completion(
+ model=model,
+ messages=messages,
+ api_base=api_base,
+ api_key=api_key,
+ model_response=model_response,
+ logging_obj=logging,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ timeout=timeout,
+ acompletion=acompletion,
+ stream=stream,
+ headers=headers or litellm.headers,
+ )
+
+ # Check if this is a Claude model - route to Azure Anthropic handler
+ elif "claude" in model.lower():
+ # Use Azure Anthropic handler for Claude models
+ api_base = AzureFoundryModelInfo.get_api_base(api_base)
+ if api_base is None:
+ raise ValueError(
+ "Azure Anthropic requests require an api_base. "
+ "Set `api_base` or the AZURE_AI_API_BASE env var."
+ )
+ api_key = AzureFoundryModelInfo.get_api_key(api_key)
+
+ # Ensure the URL ends with /v1/messages for Anthropic
+ if api_base:
+ api_base = api_base.rstrip("/")
+ if not api_base.endswith("/v1/messages"):
+ if "/anthropic" in api_base:
+ parts = api_base.split("/anthropic", 1)
+ api_base = parts[0] + "/anthropic"
+ else:
+ api_base = api_base + "/anthropic"
+ api_base = api_base + "/v1/messages"
+
+ response = azure_anthropic_chat_completions.completion(
+ model=model,
+ messages=messages,
+ api_base=api_base,
+ acompletion=acompletion,
+ custom_prompt_dict=litellm.custom_prompt_dict,
+ model_response=model_response,
+ print_verbose=print_verbose,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ logger_fn=logger_fn,
+ encoding=_get_encoding(),
+ api_key=api_key,
+ logging_obj=logging,
+ headers=headers,
+ timeout=timeout,
+ client=client,
+ custom_llm_provider=custom_llm_provider,
+ )
+ if optional_params.get("stream", False) or acompletion is True:
+ ## LOGGING
+ logging.post_call(
+ input=messages,
+ api_key=api_key,
+ original_response=response,
+ )
+ response = response
+ else:
+ # Non-Claude models use standard Azure AI flow
+ api_base = AzureFoundryModelInfo.get_api_base(api_base)
+ # set API KEY
+ api_key = AzureFoundryModelInfo.get_api_key(api_key)
+
+ headers = headers or litellm.headers
+
+ if extra_headers is not None:
+ optional_params["extra_headers"] = extra_headers
+
+ ## FOR COHERE
+ if "command-r" in model: # make sure tool call in messages are str
+ messages = stringify_json_tool_call_content(messages=messages)
+
+ ## COMPLETION CALL
+ try:
+ response = base_llm_http_handler.completion(
+ model=model,
+ messages=messages,
+ headers=headers,
+ model_response=model_response,
+ api_key=api_key,
+ api_base=api_base,
+ acompletion=acompletion,
+ logging_obj=logging,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ timeout=timeout, # type: ignore
+ client=client, # pass AsyncOpenAI, OpenAI client
+ custom_llm_provider=custom_llm_provider,
+ encoding=_get_encoding(),
+ stream=stream,
+ )
+ except Exception as e:
+ ## LOGGING - log the original exception returned
+ logging.post_call(
+ input=messages,
+ api_key=api_key,
+ original_response=str(e),
+ additional_args={"headers": headers},
+ )
+ raise e
+
+ if optional_params.get("stream", False):
+ ## LOGGING
+ logging.post_call(
+ input=messages,
+ api_key=api_key,
+ original_response=response,
+ additional_args={"headers": headers},
+ )
+
+ return response
+
+
+def _complete_text_completion_openai(
+ ctx: _CompletionDispatchContext,
+) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logger_fn = ctx.logger_fn
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ text_completion = ctx.text_completion
+ timeout = ctx.timeout
+
+ openai.api_type = "openai"
+
+ api_base = (
+ api_base
+ or litellm.api_base
+ or get_secret("OPENAI_BASE_URL")
+ or get_secret("OPENAI_API_BASE")
+ or "https://api.openai.com/v1"
+ )
+
+ openai.api_version = None
+ # set API KEY
+
+ api_key = (
+ api_key or litellm.api_key or litellm.openai_key or get_secret("OPENAI_API_KEY")
+ )
+
+ headers = headers or litellm.headers
+
+ ## LOAD CONFIG - if set
+ config = litellm.OpenAITextCompletionConfig.get_config()
+ for k, v in config.items():
+ if (
+ k not in optional_params
+ ): # completion(top_k=3) > openai_text_config(top_k=3) <- allows for dynamic variables to be passed in
+ optional_params[k] = v
+ if litellm.organization:
+ openai.organization = litellm.organization
+
+ ## COMPLETION CALL
+ _response = openai_text_completions.completion(
+ model=model,
+ messages=messages,
+ headers=headers,
+ model_response=model_response,
+ print_verbose=print_verbose,
+ api_key=api_key,
+ custom_llm_provider=custom_llm_provider,
+ api_base=api_base,
+ acompletion=acompletion,
+ client=client, # pass AsyncOpenAI, OpenAI client
+ logging_obj=logging,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ logger_fn=logger_fn,
+ timeout=timeout, # type: ignore
+ )
+
+ if (
+ optional_params.get("stream", False) is False
+ and acompletion is False
+ and text_completion is False
+ ):
+ # convert to chat completion response
+ _response = (
+ litellm.OpenAITextCompletionConfig().convert_to_chat_model_response_object(
+ response_object=_response, model_response_object=model_response
+ )
+ )
+
+ if optional_params.get("stream", False) or acompletion is True:
+ ## LOGGING
+ logging.post_call(
+ input=messages,
+ api_key=api_key,
+ original_response=_response,
+ additional_args={"headers": headers},
+ )
+ return _response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract
+
+
+def _complete_fireworks_ai(
+ ctx: _CompletionDispatchContext,
+) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ provider_config = ctx.provider_config
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ try:
+ response = base_llm_http_handler.completion(
+ model=model,
+ messages=messages,
+ headers=headers,
+ model_response=model_response,
+ api_key=api_key,
+ api_base=api_base,
+ acompletion=acompletion,
+ logging_obj=logging,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ timeout=timeout, # type: ignore
+ client=client,
+ custom_llm_provider=custom_llm_provider,
+ encoding=_get_encoding(),
+ stream=stream,
+ provider_config=provider_config,
+ )
+ except Exception as e:
+ ## LOGGING - log the original exception returned
+ logging.post_call(
+ input=messages,
+ api_key=api_key,
+ original_response=str(e),
+ additional_args={"headers": headers},
+ )
+ raise e
+
+ return response
+
+
+def _complete_heroku(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ provider_config = ctx.provider_config
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ try:
+ response = base_llm_http_handler.completion(
+ model=model,
+ messages=messages,
+ headers=headers,
+ model_response=model_response,
+ api_key=api_key,
+ api_base=api_base,
+ acompletion=acompletion,
+ logging_obj=logging,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ timeout=timeout,
+ client=client,
+ custom_llm_provider=custom_llm_provider,
+ encoding=_get_encoding(),
+ stream=stream,
+ provider_config=provider_config,
+ )
+ except Exception as e:
+ logging.post_call(
+ input=messages,
+ api_key=api_key,
+ original_response=str(e),
+ additional_args={"headers": headers},
+ )
+ raise e
+
+ return response
+
+
+def _complete_ragflow(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ provider_config = ctx.provider_config
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ try:
+ response = base_llm_http_handler.completion(
+ model=model,
+ messages=messages,
+ headers=headers,
+ model_response=model_response,
+ api_key=api_key,
+ api_base=api_base,
+ acompletion=acompletion,
+ logging_obj=logging,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ timeout=timeout,
+ client=client,
+ custom_llm_provider=custom_llm_provider,
+ encoding=_get_encoding(),
+ stream=stream,
+ provider_config=provider_config,
+ )
+ except Exception as e:
+ logging.post_call(
+ input=messages,
+ api_key=api_key,
+ original_response=str(e),
+ additional_args={"headers": headers},
+ )
+ raise e
+
+ return response
+
+
+def _complete_xai(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ provider_config = ctx.provider_config
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ try:
+ response = base_llm_http_handler.completion(
+ model=model,
+ messages=messages,
+ headers=headers,
+ model_response=model_response,
+ api_key=api_key,
+ api_base=api_base,
+ acompletion=acompletion,
+ logging_obj=logging,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ timeout=timeout, # type: ignore
+ client=client,
+ custom_llm_provider=custom_llm_provider,
+ encoding=_get_encoding(),
+ stream=stream,
+ provider_config=provider_config,
+ )
+ except Exception as e:
+ ## LOGGING - log the original exception returned
+ logging.post_call(
+ input=messages,
+ api_key=api_key,
+ original_response=str(e),
+ additional_args={"headers": headers},
+ )
+ raise e
+
+ return response
+
+
+def _complete_groq(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ api_base = (
+ api_base # for deepinfra/perplexity/anyscale/groq/friendliai we check in get_llm_provider and pass in the api base from there
+ or litellm.api_base
+ or get_secret("GROQ_API_BASE")
+ or "https://api.groq.com/openai/v1"
+ )
+
+ # set API KEY
+ api_key = (
+ api_key
+ or litellm.api_key # for deepinfra/perplexity/anyscale/friendliai we check in get_llm_provider and pass in the api key from there
+ or litellm.groq_key
+ or get_secret("GROQ_API_KEY")
+ )
+
+ headers = headers or litellm.headers
+
+ ## LOAD CONFIG - if set
+ config = litellm.GroqChatConfig.get_config()
+ for k, v in config.items():
+ if (
+ k not in optional_params
+ ): # completion(top_k=3) > openai_config(top_k=3) <- allows for dynamic variables to be passed in
+ optional_params[k] = v
+
+ return base_llm_http_handler.completion(
+ model=model,
+ stream=stream,
+ messages=messages,
+ acompletion=acompletion,
+ api_base=api_base,
+ model_response=model_response,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ custom_llm_provider=custom_llm_provider,
+ timeout=timeout,
+ headers=headers,
+ encoding=_get_encoding(),
+ api_key=api_key,
+ logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
+ client=client,
+ )
+
+
+def _complete_bedrock_mantle(
+ ctx: _CompletionDispatchContext,
+) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ api_base = api_base or litellm.api_base or get_secret("BEDROCK_MANTLE_API_BASE")
+ api_key = api_key or litellm.api_key or get_secret("BEDROCK_MANTLE_API_KEY")
+ headers = headers or litellm.headers
+ config = litellm.BedrockMantleChatConfig.get_config()
+ for k, v in config.items():
+ if k not in optional_params:
+ optional_params[k] = v
+ return base_llm_http_handler.completion(
+ model=model,
+ stream=stream,
+ messages=messages,
+ acompletion=acompletion,
+ api_base=api_base,
+ model_response=model_response,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ custom_llm_provider=custom_llm_provider,
+ timeout=timeout,
+ headers=headers,
+ encoding=_get_encoding(),
+ api_key=api_key,
+ logging_obj=logging,
+ client=client,
+ )
+
+
+def _complete_a2a(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ provider_config = ctx.provider_config
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ (
+ api_base,
+ api_key,
+ headers,
+ ) = litellm.A2AConfig.resolve_agent_config_from_registry(
+ model=model,
+ api_base=api_base,
+ api_key=api_key,
+ headers=headers,
+ optional_params=optional_params,
+ )
+
+ # Fall back to environment variables and defaults
+ api_base = api_base or litellm.api_base or get_secret_str("A2A_API_BASE")
+
+ if api_base is None:
+ raise Exception(
+ "api_base is required for A2A provider. "
+ "Either provide api_base parameter, set A2A_API_BASE environment variable, "
+ "or register the agent in the proxy with model='a2a/'."
+ )
+
+ headers = headers or litellm.headers
+
+ return base_llm_http_handler.completion(
+ model=model,
+ stream=stream,
+ messages=messages,
+ acompletion=acompletion,
+ api_base=api_base,
+ model_response=model_response,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ custom_llm_provider=custom_llm_provider,
+ timeout=timeout,
+ headers=headers,
+ encoding=_get_encoding(),
+ api_key=api_key,
+ logging_obj=logging,
+ client=client,
+ provider_config=provider_config,
+ )
+
+
+def _complete_gigachat(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ provider_config = ctx.provider_config
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ api_key = (
+ api_key
+ or litellm.api_key
+ or litellm.gigachat_key
+ or get_secret("GIGACHAT_API_KEY")
+ or get_secret("GIGACHAT_CREDENTIALS")
+ )
+
+ headers = headers or litellm.headers or {}
+
+ ## COMPLETION CALL
+ try:
+ response = base_llm_http_handler.completion(
+ model=model,
+ messages=messages,
+ headers=headers,
+ model_response=model_response,
+ api_key=api_key,
+ api_base=api_base,
+ acompletion=acompletion,
+ logging_obj=logging,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ timeout=timeout,
+ client=client,
+ custom_llm_provider=custom_llm_provider,
+ encoding=_get_encoding(),
+ stream=stream,
+ provider_config=provider_config,
+ )
+ except Exception as e:
+ ## LOGGING - log the original exception returned
+ logging.post_call(
+ input=messages,
+ api_key=api_key,
+ original_response=str(e),
+ additional_args={"headers": headers},
+ )
+ raise e
+
+ return response
+
+
+def _complete_sap(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ headers = headers or litellm.headers
+ ## LOAD CONFIG - if set
+ config = litellm.GenAIHubOrchestrationConfig.get_config()
+ for k, v in config.items():
+ if (
+ k not in optional_params
+ ): # completion(top_k=3) > openai_config(top_k=3) <- allows for dynamic variables to be passed in
+ optional_params[k] = v
+
+ return sap_gen_ai_hub_chat_completions.completion(
+ model=model,
+ messages=messages,
+ headers=headers,
+ model_response=model_response,
+ acompletion=acompletion,
+ logging_obj=logging,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ timeout=timeout, # type: ignore
+ shared_session=shared_session,
+ client=client,
+ custom_llm_provider=custom_llm_provider,
+ encoding=_get_encoding(),
+ api_key=api_key,
+ api_base=api_base,
+ stream=stream,
+ )
+
+
+def _complete_aiohttp_openai(
+ ctx: _CompletionDispatchContext,
+) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ extra_headers = ctx.extra_headers
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ api_base = (
+ api_base # for deepinfra/perplexity/anyscale/groq/friendliai we check in get_llm_provider and pass in the api base from there
+ or litellm.api_base
+ or get_secret("OPENAI_BASE_URL")
+ or get_secret("OPENAI_API_BASE")
+ or "https://api.openai.com/v1"
+ )
+ # set API KEY
+ api_key = (
+ api_key
+ or litellm.api_key # for deepinfra/perplexity/anyscale/friendliai we check in get_llm_provider and pass in the api key from there
+ or litellm.openai_key
+ or get_secret("OPENAI_API_KEY")
+ )
+
+ headers = headers or litellm.headers
+
+ if extra_headers is not None:
+ optional_params["extra_headers"] = extra_headers
+ return base_llm_aiohttp_handler.completion(
+ model=model,
+ messages=messages,
+ headers=headers,
+ model_response=model_response,
+ api_key=api_key,
+ api_base=api_base,
+ acompletion=acompletion,
+ logging_obj=logging,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ timeout=timeout,
+ client=client,
+ custom_llm_provider=custom_llm_provider,
+ encoding=_get_encoding(),
+ stream=stream,
+ )
+
+
+def _complete_cometapi(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ provider_config = ctx.provider_config
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ api_key = (
+ api_key
+ or litellm.cometapi_key
+ or get_secret_str("COMETAPI_KEY")
+ or litellm.api_key
+ )
+
+ api_base = (
+ api_base
+ or litellm.api_base
+ or get_secret_str("COMETAPI_API_BASE")
+ or "https://api.cometapi.com/v1"
+ )
+
+ ## COMPLETION CALL
+ response = base_llm_http_handler.completion(
+ model=model,
+ messages=messages,
+ headers=headers,
+ model_response=model_response,
+ api_key=api_key,
+ api_base=api_base,
+ acompletion=acompletion,
+ logging_obj=logging,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ timeout=timeout,
+ client=client,
+ custom_llm_provider=custom_llm_provider,
+ encoding=_get_encoding(),
+ stream=stream,
+ provider_config=provider_config,
+ )
+
+ ## LOGGING
+ logging.post_call(input=messages, api_key=api_key, original_response=response)
+
+ return response
+
+
+def _complete_minimax(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ provider_config = ctx.provider_config
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ api_key = api_key or get_secret_str("MINIMAX_API_KEY") or litellm.api_key
+
+ api_base = (
+ api_base
+ or litellm.api_base
+ or get_secret_str("MINIMAX_API_BASE")
+ or "https://api.minimax.io/v1"
+ )
+
+ response = base_llm_http_handler.completion(
+ model=model,
+ messages=messages,
+ api_base=api_base,
+ custom_llm_provider=custom_llm_provider,
+ model_response=model_response,
+ encoding=_get_encoding(),
+ logging_obj=logging,
+ optional_params=optional_params,
+ timeout=timeout,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ acompletion=acompletion,
+ stream=stream,
+ api_key=api_key,
+ headers=headers,
+ client=client,
+ provider_config=provider_config,
+ )
+ logging.post_call(input=messages, api_key=api_key, original_response=response)
+
+ return response
+
+
+def _complete_hosted_vllm(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ provider_config = ctx.provider_config
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ api_base = api_base or litellm.api_base or get_secret_str("HOSTED_VLLM_API_BASE")
+
+ response = base_llm_http_handler.completion(
+ model=model,
+ messages=messages,
+ api_base=api_base,
+ custom_llm_provider=custom_llm_provider,
+ model_response=model_response,
+ encoding=_get_encoding(),
+ logging_obj=logging,
+ optional_params=optional_params,
+ timeout=timeout,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ acompletion=acompletion,
+ stream=stream,
+ api_key=api_key,
+ headers=headers,
+ client=client,
+ provider_config=provider_config,
+ )
+ logging.post_call(input=messages, api_key=api_key, original_response=response)
+
+ return response
+
+
+def _complete_custom_openai(
+ ctx: _CompletionDispatchContext,
+) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ custom_prompt_dict = ctx.custom_prompt_dict
+ extra_headers = ctx.extra_headers
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logger_fn = ctx.logger_fn
+ logging = ctx.logging
+ messages = ctx.messages
+ metadata = ctx.metadata
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ organization = ctx.organization
+ provider_config = ctx.provider_config
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ api_base = (
+ api_base # for deepinfra/perplexity/anyscale/groq/friendliai we check in get_llm_provider and pass in the api base from there
+ or litellm.api_base
+ or get_secret("OPENAI_BASE_URL")
+ or get_secret("OPENAI_API_BASE")
+ or "https://api.openai.com/v1"
+ )
+ organization = (
+ organization
+ or litellm.organization
+ or get_secret("OPENAI_ORGANIZATION")
+ or None # default - https://github.com/openai/openai-python/blob/284c1799070c723c6a553337134148a7ab088dd8/openai/util.py#L105
+ )
+ openai.organization = organization
+ # set API KEY
+ api_key = (
+ api_key
+ or litellm.api_key # for deepinfra/perplexity/anyscale/friendliai we check in get_llm_provider and pass in the api key from there
+ or litellm.openai_key
+ or get_secret("OPENAI_API_KEY")
+ )
+
+ headers = headers or litellm.headers
+
+ # Add GitHub Copilot headers (same as /responses endpoint does)
+ if custom_llm_provider == "github_copilot":
+ from litellm.llms.github_copilot.authenticator import Authenticator
+ from litellm.llms.github_copilot.common_utils import (
+ get_copilot_default_headers,
+ )
+
+ copilot_auth = Authenticator()
+ copilot_api_key = copilot_auth.get_api_key()
+ copilot_headers = get_copilot_default_headers(copilot_api_key)
+ if extra_headers:
+ copilot_headers.update(extra_headers)
+ extra_headers = copilot_headers
+
+ if extra_headers is not None:
+ optional_params["extra_headers"] = extra_headers
+
+ if (
+ litellm.enable_preview_features and metadata is not None
+ ): # [PREVIEW] allow metadata to be passed to OPENAI
+ openai_metadata = get_requester_metadata(metadata)
+ if openai_metadata is not None:
+ optional_params["metadata"] = openai_metadata
+
+ ## LOAD CONFIG - if set
+ config = litellm.OpenAIConfig.get_config()
+ for k, v in config.items():
+ if (
+ k not in optional_params
+ ): # completion(top_k=3) > openai_config(top_k=3) <- allows for dynamic variables to be passed in
+ optional_params[k] = v
+
+ ## COMPLETION CALL
+ use_base_llm_http_handler = get_secret_bool(
+ "EXPERIMENTAL_OPENAI_BASE_LLM_HTTP_HANDLER"
+ )
+
+ try:
+ if use_base_llm_http_handler:
+ response = base_llm_http_handler.completion(
+ model=model,
+ messages=messages,
+ api_base=api_base,
+ custom_llm_provider=custom_llm_provider,
+ model_response=model_response,
+ encoding=_get_encoding(),
+ logging_obj=logging,
+ optional_params=optional_params,
+ timeout=timeout,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ acompletion=acompletion,
+ stream=stream,
+ api_key=api_key,
+ headers=headers,
+ client=client,
+ provider_config=provider_config,
+ )
+ else:
+ response = openai_chat_completions.completion(
+ model=model,
+ messages=messages,
+ headers=headers,
+ model_response=model_response,
+ print_verbose=print_verbose,
+ api_key=api_key,
+ api_base=api_base,
+ acompletion=acompletion,
+ logging_obj=logging,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ logger_fn=logger_fn,
+ timeout=timeout, # type: ignore
+ custom_prompt_dict=custom_prompt_dict,
+ client=client, # pass AsyncOpenAI, OpenAI client
+ organization=organization,
+ custom_llm_provider=custom_llm_provider,
+ shared_session=shared_session,
+ )
+ except Exception as e:
+ ## LOGGING - log the original exception returned
+ logging.post_call(
+ input=messages,
+ api_key=api_key,
+ original_response=str(e),
+ additional_args={"headers": headers},
+ )
+ raise e
+
+ if optional_params.get("stream", False):
+ ## LOGGING
+ logging.post_call(
+ input=messages,
+ api_key=api_key,
+ original_response=response,
+ additional_args={"headers": headers},
+ )
+
+ return response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract
+
+
+def _complete_mistral(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ provider_config = ctx.provider_config
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ api_key = api_key or litellm.api_key or get_secret("MISTRAL_API_KEY")
+ api_base = (
+ api_base
+ or litellm.api_base
+ or get_secret("MISTRAL_API_BASE")
+ or "https://api.mistral.ai/v1"
+ )
+
+ return base_llm_http_handler.completion(
+ model=model,
+ messages=messages,
+ api_base=api_base,
+ custom_llm_provider=custom_llm_provider,
+ model_response=model_response,
+ encoding=_get_encoding(),
+ logging_obj=logging,
+ optional_params=optional_params,
+ timeout=timeout,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ acompletion=acompletion,
+ stream=stream,
+ api_key=api_key,
+ headers=headers,
+ client=client,
+ provider_config=provider_config,
+ )
+
+
+def _complete_replicate(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ custom_prompt_dict = ctx.custom_prompt_dict
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logger_fn = ctx.logger_fn
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+
+ replicate_key = (
+ api_key
+ or litellm.replicate_key
+ or litellm.api_key
+ or get_secret("REPLICATE_API_KEY")
+ or get_secret("REPLICATE_API_TOKEN")
+ )
+
+ api_base = (
+ api_base
+ or litellm.api_base
+ or get_secret("REPLICATE_API_BASE")
+ or "https://api.replicate.com/v1"
+ )
+
+ custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict
+
+ model_response = replicate_chat_completion( # type: ignore
+ model=model,
+ messages=messages,
+ api_base=api_base,
+ model_response=model_response,
+ print_verbose=print_verbose,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ logger_fn=logger_fn,
+ encoding=_get_encoding(), # for calculating input/output tokens
+ api_key=replicate_key,
+ logging_obj=logging,
+ custom_prompt_dict=custom_prompt_dict,
+ acompletion=acompletion,
+ headers=headers,
+ )
+
+ if optional_params.get("stream", False) is True:
+ ## LOGGING
+ logging.post_call(
+ input=messages,
+ api_key=replicate_key,
+ original_response=model_response,
+ )
+
+ return model_response
+
+
+def _complete_anthropic_text(
+ ctx: _CompletionDispatchContext,
+) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ custom_prompt_dict = ctx.custom_prompt_dict
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ api_key = (
+ api_key
+ or litellm.anthropic_key
+ or litellm.api_key
+ or os.environ.get("ANTHROPIC_API_KEY")
+ )
+ custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict
+ api_base = cast(
+ Optional[str],
+ api_base
+ or litellm.api_base
+ or get_secret("ANTHROPIC_API_BASE")
+ or get_secret("ANTHROPIC_BASE_URL")
+ or "https://api.anthropic.com/v1/complete",
+ )
+
+ # Check if we should disable automatic URL suffix appending
+ disable_url_suffix = get_secret_bool("LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX")
+ if (
+ api_base is not None
+ and not disable_url_suffix
+ and not api_base.endswith("/v1/complete")
+ ):
+ api_base += "/v1/complete"
+ elif disable_url_suffix:
+ verbose_logger.debug(
+ "LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX is set, skipping /v1/complete suffix"
+ )
+
+ return base_llm_http_handler.completion(
+ model=model,
+ stream=stream,
+ messages=messages,
+ acompletion=acompletion,
+ api_base=api_base,
+ model_response=model_response,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ custom_llm_provider="anthropic_text",
+ timeout=timeout,
+ headers=headers,
+ encoding=_get_encoding(),
+ api_key=api_key,
+ logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
+ )
+
+
+def _complete_anthropic(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ custom_prompt_dict = ctx.custom_prompt_dict
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logger_fn = ctx.logger_fn
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ timeout = ctx.timeout
+
+ api_key = (
+ api_key
+ or litellm.anthropic_key
+ or litellm.api_key
+ or os.environ.get("ANTHROPIC_API_KEY")
+ )
+ custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict
+ # call /messages
+ # default route for all anthropic models
+ api_base = cast(
+ Optional[str],
+ api_base
+ or litellm.api_base
+ or get_secret("ANTHROPIC_API_BASE")
+ or get_secret("ANTHROPIC_BASE_URL")
+ or "https://api.anthropic.com/v1/messages",
+ )
+
+ # Check if we should disable automatic URL suffix appending
+ disable_url_suffix = get_secret_bool("LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX")
+ if (
+ api_base is not None
+ and not disable_url_suffix
+ and not api_base.endswith("/v1/messages")
+ ):
+ api_base += "/v1/messages"
+ elif disable_url_suffix:
+ verbose_logger.debug(
+ "LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX is set, skipping /v1/messages suffix"
+ )
+
+ response = anthropic_chat_completions.completion(
+ model=model,
+ messages=messages,
+ api_base=api_base,
+ acompletion=acompletion,
+ custom_prompt_dict=litellm.custom_prompt_dict,
+ model_response=model_response,
+ print_verbose=print_verbose,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ logger_fn=logger_fn,
+ encoding=_get_encoding(), # for calculating input/output tokens
+ api_key=api_key,
+ logging_obj=logging,
+ headers=headers,
+ timeout=timeout,
+ client=client,
+ custom_llm_provider=custom_llm_provider,
+ )
+ if optional_params.get("stream", False) or acompletion is True:
+ ## LOGGING
+ logging.post_call(
+ input=messages,
+ api_key=api_key,
+ original_response=response,
+ )
+ return response
+
+
+def _complete_nlp_cloud(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ litellm_params = ctx.litellm_params
+ logger_fn = ctx.logger_fn
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+
+ nlp_cloud_key = (
+ api_key
+ or litellm.nlp_cloud_key
+ or get_secret("NLP_CLOUD_API_KEY")
+ or litellm.api_key
+ )
+
+ api_base = (
+ api_base
+ or litellm.api_base
+ or get_secret("NLP_CLOUD_API_BASE")
+ or "https://api.nlpcloud.io/v1/gpu/"
+ )
+
+ response = nlp_cloud_chat_completion(
+ model=model,
+ messages=messages,
+ api_base=api_base,
+ model_response=model_response,
+ print_verbose=print_verbose,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ logger_fn=logger_fn,
+ encoding=_get_encoding(),
+ api_key=nlp_cloud_key,
+ logging_obj=logging,
+ )
+
+ if "stream" in optional_params and optional_params["stream"] is True:
+ # don't try to access stream object,
+ response = CustomStreamWrapper(
+ response,
+ model,
+ custom_llm_provider="nlp_cloud",
+ logging_obj=logging,
+ )
+
+ if optional_params.get("stream", False) or acompletion is True:
+ ## LOGGING
+ logging.post_call(
+ input=messages,
+ api_key=api_key,
+ original_response=response,
+ )
+
+ return response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract
+
+
+def _complete_aleph_alpha(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ litellm_params = ctx.litellm_params
+ logger_fn = ctx.logger_fn
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+
+ aleph_alpha_key = (
+ api_key
+ or litellm.aleph_alpha_key
+ or get_secret("ALEPH_ALPHA_API_KEY")
+ or get_secret("ALEPHALPHA_API_KEY")
+ or litellm.api_key
+ )
+
+ api_base = (
+ api_base
+ or litellm.api_base
+ or get_secret("ALEPH_ALPHA_API_BASE")
+ or "https://api.aleph-alpha.com/complete"
+ )
+
+ model_response = aleph_alpha.completion(
+ model=model,
+ messages=messages,
+ api_base=api_base,
+ model_response=model_response,
+ print_verbose=print_verbose,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ logger_fn=logger_fn,
+ encoding=_get_encoding(),
+ default_max_tokens_to_sample=litellm.max_tokens,
+ api_key=aleph_alpha_key,
+ logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
+ )
+
+ if "stream" in optional_params and optional_params["stream"] is True:
+ # don't try to access stream object,
+ return CustomStreamWrapper(
+ model_response,
+ model,
+ custom_llm_provider="aleph_alpha",
+ logging_obj=logging,
+ )
+ return model_response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract
+
+
+def _complete_cohere_chat(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ extra_headers = ctx.extra_headers
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ provider_config = ctx.provider_config
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ cohere_key = (
+ api_key
+ or litellm.cohere_key
+ or get_secret_str("COHERE_API_KEY")
+ or get_secret_str("CO_API_KEY")
+ or litellm.api_key
+ )
+
+ cohere_route = CohereModelInfo.get_cohere_route(model)
+ verbose_logger.debug(f"Cohere route: {cohere_route}")
+ # Set API base based on route
+ if cohere_route == "v2":
+ api_base = (
+ api_base
+ or litellm.api_base
+ or get_secret_str("COHERE_API_BASE")
+ or "https://api.cohere.com/v2/chat"
+ )
+ # Remove v2/ prefix from model name for the actual API call
+ if "v2/" in model:
+ model = model.replace("v2/", "")
+ else:
+ api_base = (
+ api_base
+ or litellm.api_base
+ or get_secret_str("COHERE_API_BASE")
+ or "https://api.cohere.ai/v1/chat"
+ )
+
+ headers = headers or litellm.headers or {}
+ if headers is None:
+ headers = {}
+
+ if extra_headers is not None:
+ headers.update(extra_headers)
+
+ verbose_logger.debug(f"Model: {model}, API Base: {api_base}")
+ verbose_logger.debug(f"Provider Config: {provider_config}")
+ return base_llm_http_handler.completion(
+ model=model,
+ stream=stream,
+ messages=messages,
+ acompletion=acompletion,
+ api_base=api_base,
+ model_response=model_response,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ custom_llm_provider="cohere_chat",
+ timeout=timeout,
+ headers=headers,
+ encoding=_get_encoding(),
+ api_key=cohere_key,
+ provider_config=provider_config,
+ logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
+ )
+
+
+def _complete_maritalk(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ custom_prompt_dict = ctx.custom_prompt_dict
+ litellm_params = ctx.litellm_params
+ logger_fn = ctx.logger_fn
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+
+ maritalk_key = (
+ api_key
+ or litellm.maritalk_key
+ or get_secret("MARITALK_API_KEY")
+ or litellm.api_key
+ )
+
+ api_base = (
+ api_base
+ or litellm.api_base
+ or get_secret("MARITALK_API_BASE")
+ or "https://chat.maritaca.ai/api"
+ )
+
+ return openai_like_chat_completion.completion(
+ model=model,
+ messages=messages,
+ api_base=api_base,
+ model_response=model_response,
+ print_verbose=print_verbose,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ logger_fn=logger_fn,
+ encoding=_get_encoding(),
+ api_key=maritalk_key,
+ logging_obj=logging,
+ custom_llm_provider="maritalk",
+ custom_prompt_dict=custom_prompt_dict,
+ )
+
+
+def _complete_amazon_nova(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ custom_llm_provider = ctx.custom_llm_provider
+ custom_prompt_dict = ctx.custom_prompt_dict
+ litellm_params = ctx.litellm_params
+ logger_fn = ctx.logger_fn
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ timeout = ctx.timeout
+
+ api_key = (
+ api_key
+ or litellm.amazon_nova_api_key
+ or get_secret_str("AMAZON_NOVA_API_KEY")
+ or litellm.api_key
+ )
+ api_base = (
+ api_base
+ or litellm.api_base
+ or get_secret_str("AMAZON_NOVA_API_BASE")
+ or "https://api.nova.amazon.com/v1"
+ )
+ return openai_like_chat_completion.completion(
+ model=model,
+ messages=messages,
+ api_base=api_base,
+ model_response=model_response,
+ print_verbose=print_verbose,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ logger_fn=logger_fn,
+ encoding=_get_encoding(),
+ api_key=api_key,
+ logging_obj=logging,
+ timeout=timeout,
+ custom_llm_provider=custom_llm_provider,
+ custom_prompt_dict=custom_prompt_dict,
+ )
+
+
+def _complete_huggingface(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ huggingface_key = (
+ api_key
+ or litellm.huggingface_key
+ or os.environ.get("HF_TOKEN")
+ or os.environ.get("HUGGINGFACE_API_KEY")
+ or litellm.api_key
+ )
+ hf_headers = headers or litellm.headers
+ return base_llm_http_handler.completion(
+ model=model,
+ messages=messages,
+ headers=hf_headers,
+ model_response=model_response,
+ api_key=huggingface_key,
+ api_base=api_base,
+ acompletion=acompletion,
+ logging_obj=logging,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ timeout=timeout, # type: ignore
+ client=client,
+ custom_llm_provider=custom_llm_provider,
+ encoding=_get_encoding(),
+ stream=stream,
+ )
+
+
+def _complete_oci(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ return base_llm_http_handler.completion(
+ model=model,
+ messages=messages,
+ headers=headers,
+ model_response=model_response,
+ api_key=api_key,
+ api_base=api_base,
+ acompletion=acompletion,
+ logging_obj=logging,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ timeout=timeout, # type: ignore
+ client=client,
+ custom_llm_provider=custom_llm_provider,
+ encoding=_get_encoding(),
+ stream=stream,
+ )
+
+
+def _complete_compactifai(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ provider_config = ctx.provider_config
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ api_key = api_key or get_secret_str("COMPACTIFAI_API_KEY") or litellm.api_key
+
+ api_base = api_base or "https://api.compactif.ai/v1"
+
+ ## COMPLETION CALL
+ return base_llm_http_handler.completion(
+ model=model,
+ messages=messages,
+ headers=headers,
+ model_response=model_response,
+ api_key=api_key,
+ api_base=api_base,
+ acompletion=acompletion,
+ logging_obj=logging,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ timeout=timeout,
+ client=client,
+ custom_llm_provider=custom_llm_provider,
+ encoding=_get_encoding(),
+ stream=stream,
+ provider_config=provider_config,
+ )
+
+
+def _complete_oobabooga(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ api_base = ctx.api_base
+ litellm_params = ctx.litellm_params
+ logger_fn = ctx.logger_fn
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+
+ model_response = oobabooga.completion(
+ model=model,
+ messages=messages,
+ model_response=model_response,
+ api_base=api_base, # type: ignore
+ print_verbose=print_verbose,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ api_key=None,
+ logger_fn=logger_fn,
+ encoding=_get_encoding(),
+ logging_obj=logging,
+ )
+ if "stream" in optional_params and optional_params["stream"] is True:
+ # don't try to access stream object,
+ return CustomStreamWrapper(
+ model_response,
+ model,
+ custom_llm_provider="oobabooga",
+ logging_obj=logging,
+ )
+ return model_response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract
+
+
+def _complete_databricks(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ api_base = (
+ api_base # for databricks we check in get_llm_provider and pass in the api base from there
+ or litellm.api_base
+ or os.getenv("DATABRICKS_API_BASE")
+ )
+
+ # set API KEY
+ api_key = (
+ api_key
+ or litellm.api_key # for databricks we check in get_llm_provider and pass in the api key from there
+ or litellm.databricks_key
+ or get_secret("DATABRICKS_API_KEY")
+ )
+
+ headers = headers or litellm.headers
+
+ ## COMPLETION CALL
+ try:
+ response = base_llm_http_handler.completion(
+ model=model,
+ stream=stream,
+ messages=messages,
+ acompletion=acompletion,
+ api_base=api_base,
+ model_response=model_response,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ custom_llm_provider="databricks",
+ timeout=timeout,
+ headers=headers,
+ encoding=_get_encoding(),
+ api_key=api_key,
+ logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
+ client=client,
+ )
+ except Exception as e:
+ ## LOGGING - log the original exception returned
+ logging.post_call(
+ input=messages,
+ api_key=api_key,
+ original_response=str(e),
+ additional_args={"headers": headers},
+ )
+ raise e
+
+ if optional_params.get("stream", False):
+ ## LOGGING
+ logging.post_call(
+ input=messages,
+ api_key=api_key,
+ original_response=response,
+ additional_args={"headers": headers},
+ )
+
+ return response
+
+
+def _complete_datarobot(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ provider_config = ctx.provider_config
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ return base_llm_http_handler.completion(
+ model=model,
+ messages=messages,
+ headers=headers,
+ model_response=model_response,
+ api_key=api_key,
+ api_base=api_base,
+ acompletion=acompletion,
+ logging_obj=logging,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ timeout=timeout, # type: ignore
+ client=client,
+ custom_llm_provider=custom_llm_provider,
+ encoding=_get_encoding(),
+ stream=stream,
+ provider_config=provider_config,
+ )
+
+
+def _complete_openrouter(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ api_base = (
+ api_base
+ or litellm.api_base
+ or get_secret_str("OPENROUTER_API_BASE")
+ or "https://openrouter.ai/api/v1"
+ )
+
+ api_key = (
+ api_key
+ or litellm.api_key
+ or litellm.openrouter_key
+ or get_secret_str("OPENROUTER_API_KEY")
+ or get_secret_str("OR_API_KEY")
+ )
+
+ openrouter_site_url = get_secret("OR_SITE_URL") or "https://litellm.ai"
+ openrouter_app_name = get_secret("OR_APP_NAME") or "liteLLM"
+
+ openrouter_headers = {
+ "HTTP-Referer": openrouter_site_url,
+ "X-Title": openrouter_app_name,
+ }
+
+ _headers = headers or litellm.headers
+ if _headers:
+ openrouter_headers.update(_headers)
+
+ headers = openrouter_headers
+
+ ## Load Config
+ config = litellm.OpenrouterConfig.get_config()
+ for k, v in config.items():
+ if k == "extra_body":
+ # we use openai 'extra_body' to pass openrouter specific params - transforms, route, models
+ if "extra_body" in optional_params:
+ optional_params[k].update(v)
+ else:
+ optional_params[k] = v
+ elif k not in optional_params:
+ optional_params[k] = v
+
+ ## COMPLETION CALL
+ response = base_llm_http_handler.completion(
+ model=model,
+ stream=stream,
+ messages=messages,
+ acompletion=acompletion,
+ api_base=api_base,
+ model_response=model_response,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ custom_llm_provider="openrouter",
+ timeout=timeout,
+ headers=headers,
+ encoding=_get_encoding(),
+ api_key=api_key,
+ logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
+ client=client,
+ )
+ ## LOGGING
+ logging.post_call(
+ input=messages, api_key=openai.api_key, original_response=response
+ )
+
+ return response
+
+
+def _complete_vercel_ai_gateway(
+ ctx: _CompletionDispatchContext,
+) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ api_base = (
+ api_base
+ or litellm.api_base
+ or get_secret_str("VERCEL_AI_GATEWAY_API_BASE")
+ or "https://ai-gateway.vercel.sh/v1"
+ )
+
+ api_key = api_key or litellm.api_key or get_secret("VERCEL_AI_GATEWAY_API_KEY")
+
+ vercel_site_url = get_secret("VERCEL_SITE_URL") or "https://litellm.ai"
+ vercel_app_name = get_secret("VERCEL_APP_NAME") or "liteLLM"
+
+ vercel_headers = {
+ "http-referer": vercel_site_url,
+ "x-title": vercel_app_name,
+ }
+
+ _headers = headers or litellm.headers
+ if _headers:
+ vercel_headers.update(_headers)
+
+ headers = vercel_headers
+
+ ## Load Config
+ config = litellm.VercelAIGatewayConfig.get_config()
+ for k, v in config.items():
+ if k == "extra_body":
+ # we use openai 'extra_body' to pass vercel specific params - providerOptions
+ if "extra_body" in optional_params:
+ optional_params[k].update(v)
+ else:
+ optional_params[k] = v
+ elif k not in optional_params:
+ optional_params[k] = v
+
+ ## COMPLETION CALL
+ response = base_llm_http_handler.completion(
+ model=model,
+ stream=stream,
+ messages=messages,
+ acompletion=acompletion,
+ api_base=api_base,
+ model_response=model_response,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ custom_llm_provider="vercel_ai_gateway",
+ timeout=timeout,
+ headers=headers,
+ encoding=_get_encoding(),
+ api_key=api_key,
+ logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
+ client=client,
+ )
+ ## LOGGING
+ logging.post_call(
+ input=messages, api_key=openai.api_key, original_response=response
+ )
+
+ return response
+
+
+def _complete_vertex_ai_beta(
+ ctx: _CompletionDispatchContext,
+) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logger_fn = ctx.logger_fn
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ timeout = ctx.timeout
+
+ vertex_ai_project = (
+ optional_params.pop("vertex_project", None)
+ or optional_params.pop("vertex_ai_project", None)
+ or litellm.vertex_project
+ or get_secret("VERTEXAI_PROJECT")
+ )
+ vertex_ai_location = (
+ optional_params.pop("vertex_location", None)
+ or optional_params.pop("vertex_ai_location", None)
+ or litellm.vertex_location
+ or get_secret("VERTEXAI_LOCATION")
+ )
+ vertex_credentials = (
+ optional_params.pop("vertex_credentials", None)
+ or optional_params.pop("vertex_ai_credentials", None)
+ or get_secret("VERTEXAI_CREDENTIALS")
+ )
+
+ gemini_api_key = (
+ api_key
+ or get_api_key_from_env()
+ or get_secret("PALM_API_KEY") # older palm api key should also work
+ or litellm.api_key
+ )
+
+ api_base = api_base or litellm.api_base or get_secret("GEMINI_API_BASE")
+ new_params = safe_deep_copy(optional_params or {})
+ return vertex_chat_completion.completion( # type: ignore
+ model=model,
+ messages=messages,
+ model_response=model_response,
+ print_verbose=print_verbose,
+ optional_params=new_params,
+ litellm_params=litellm_params, # type: ignore
+ logger_fn=logger_fn,
+ encoding=_get_encoding(),
+ vertex_location=vertex_ai_location,
+ vertex_project=vertex_ai_project,
+ vertex_credentials=vertex_credentials,
+ gemini_api_key=gemini_api_key,
+ logging_obj=logging,
+ acompletion=acompletion,
+ timeout=timeout,
+ custom_llm_provider=custom_llm_provider, # type: ignore
+ client=client,
+ api_base=api_base,
+ extra_headers=headers,
+ )
+
+
+def _complete_vertex_ai(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ custom_prompt_dict = ctx.custom_prompt_dict
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logger_fn = ctx.logger_fn
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ vertex_ai_project = (
+ optional_params.pop("vertex_project", None)
+ or optional_params.pop("vertex_ai_project", None)
+ or litellm.vertex_project
+ or get_secret("VERTEXAI_PROJECT")
+ )
+ vertex_ai_location = (
+ optional_params.pop("vertex_location", None)
+ or optional_params.pop("vertex_ai_location", None)
+ or litellm.vertex_location
+ or get_secret("VERTEXAI_LOCATION")
+ )
+ vertex_credentials = (
+ optional_params.pop("vertex_credentials", None)
+ or optional_params.pop("vertex_ai_credentials", None)
+ or get_secret("VERTEXAI_CREDENTIALS")
+ )
+
+ api_base = api_base or litellm.api_base or get_secret("VERTEXAI_API_BASE")
+
+ new_params = safe_deep_copy(optional_params or {})
+ model_route = get_vertex_ai_model_route(model=model, litellm_params=litellm_params)
+
+ if model_route == VertexAIModelRoute.PARTNER_MODELS:
+ model_response = vertex_partner_models_chat_completion.completion(
+ model=model,
+ messages=messages,
+ model_response=model_response,
+ print_verbose=print_verbose,
+ optional_params=new_params,
+ litellm_params=litellm_params, # type: ignore
+ logger_fn=logger_fn,
+ encoding=_get_encoding(),
+ api_base=api_base,
+ vertex_location=vertex_ai_location,
+ vertex_project=vertex_ai_project,
+ vertex_credentials=vertex_credentials,
+ logging_obj=logging,
+ acompletion=acompletion,
+ headers=headers,
+ custom_prompt_dict=custom_prompt_dict,
+ timeout=timeout,
+ client=client,
+ )
+ elif model_route == VertexAIModelRoute.GEMINI:
+ model_response = vertex_chat_completion.completion( # type: ignore
+ model=model,
+ messages=messages,
+ model_response=model_response,
+ print_verbose=print_verbose,
+ optional_params=new_params,
+ litellm_params=litellm_params, # type: ignore
+ logger_fn=logger_fn,
+ encoding=_get_encoding(),
+ vertex_location=vertex_ai_location,
+ vertex_project=vertex_ai_project,
+ vertex_credentials=vertex_credentials,
+ gemini_api_key=None,
+ logging_obj=logging,
+ acompletion=acompletion,
+ timeout=timeout,
+ custom_llm_provider=custom_llm_provider, # type: ignore
+ client=client,
+ api_base=api_base,
+ extra_headers=headers,
+ )
+ elif model_route == VertexAIModelRoute.GEMMA:
+ # Vertex Gemma Models with custom prediction endpoint
+ model_response = vertex_gemma_chat_completion.completion(
+ model=model,
+ messages=messages,
+ model_response=model_response,
+ print_verbose=print_verbose,
+ optional_params=new_params,
+ litellm_params=litellm_params, # type: ignore
+ logger_fn=logger_fn,
+ encoding=_get_encoding(),
+ api_base=api_base,
+ vertex_location=vertex_ai_location,
+ vertex_project=vertex_ai_project,
+ vertex_credentials=vertex_credentials,
+ logging_obj=logging,
+ acompletion=acompletion,
+ headers=headers,
+ custom_prompt_dict=custom_prompt_dict,
+ timeout=timeout,
+ client=client,
+ )
+ elif model_route == VertexAIModelRoute.MODEL_GARDEN:
+ # Vertex Model Garden - OpenAI compatible models
+ model_response = vertex_model_garden_chat_completion.completion(
+ model=model,
+ messages=messages,
+ model_response=model_response,
+ print_verbose=print_verbose,
+ optional_params=new_params,
+ litellm_params=litellm_params, # type: ignore
+ logger_fn=logger_fn,
+ encoding=_get_encoding(),
+ api_base=api_base,
+ vertex_location=vertex_ai_location,
+ vertex_project=vertex_ai_project,
+ vertex_credentials=vertex_credentials,
+ logging_obj=logging,
+ acompletion=acompletion,
+ headers=headers,
+ custom_prompt_dict=custom_prompt_dict,
+ timeout=timeout,
+ client=client,
+ )
+ elif model_route == VertexAIModelRoute.AGENT_ENGINE:
+ # Vertex AI Agent Engine (Reasoning Engines)
+ from litellm.llms.vertex_ai.agent_engine.transformation import (
+ VertexAgentEngineConfig,
+ )
+
+ vertex_agent_engine_config = VertexAgentEngineConfig()
+
+ # Update litellm_params with vertex credentials
+ litellm_params["vertex_project"] = vertex_ai_project
+ litellm_params["vertex_location"] = vertex_ai_location
+ litellm_params["vertex_credentials"] = vertex_credentials
+
+ model_response = base_llm_http_handler.completion(
+ model=model,
+ stream=stream,
+ messages=messages,
+ model_response=model_response,
+ optional_params=new_params,
+ litellm_params=litellm_params, # type: ignore
+ encoding=_get_encoding(),
+ api_key=None,
+ api_base=api_base,
+ logging_obj=logging,
+ acompletion=acompletion,
+ timeout=timeout,
+ client=client,
+ custom_llm_provider="vertex_ai",
+ provider_config=vertex_agent_engine_config,
+ headers=headers or {},
+ )
+ else: # VertexAIModelRoute.NON_GEMINI
+ model_response = vertex_ai_non_gemini.completion(
+ model=model,
+ messages=messages,
+ model_response=model_response,
+ print_verbose=print_verbose,
+ optional_params=new_params,
+ litellm_params=litellm_params,
+ logger_fn=logger_fn,
+ encoding=_get_encoding(),
+ vertex_location=vertex_ai_location,
+ vertex_project=vertex_ai_project,
+ vertex_credentials=vertex_credentials,
+ logging_obj=logging,
+ acompletion=acompletion,
+ )
+
+ if (
+ "stream" in optional_params
+ and optional_params["stream"] is True
+ and acompletion is False
+ ):
+ return CustomStreamWrapper(
+ model_response,
+ model,
+ custom_llm_provider="vertex_ai",
+ logging_obj=logging,
+ )
+ return model_response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract
+
+
+def _complete_predibase(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ custom_prompt_dict = ctx.custom_prompt_dict
+ litellm_params = ctx.litellm_params
+ logger_fn = ctx.logger_fn
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ timeout = ctx.timeout
+
+ tenant_id = (
+ optional_params.pop("tenant_id", None)
+ or optional_params.pop("predibase_tenant_id", None)
+ or litellm.predibase_tenant_id
+ or get_secret("PREDIBASE_TENANT_ID")
+ )
+
+ if tenant_id is None:
+ raise ValueError(
+ "Missing Predibase Tenant ID - Required for making the request. Set dynamically (e.g. `completion(..tenant_id=)`) or in env - `PREDIBASE_TENANT_ID`."
+ )
+
+ api_base = (
+ api_base
+ or optional_params.pop("api_base", None)
+ or optional_params.pop("base_url", None)
+ or litellm.api_base
+ or get_secret("PREDIBASE_API_BASE")
+ )
+
+ api_key = (
+ api_key
+ or litellm.api_key
+ or litellm.predibase_key
+ or get_secret("PREDIBASE_API_KEY")
+ )
+
+ _model_response = predibase_chat_completions.completion(
+ model=model,
+ messages=messages,
+ model_response=model_response,
+ print_verbose=print_verbose,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ logger_fn=logger_fn,
+ encoding=_get_encoding(),
+ logging_obj=logging,
+ acompletion=acompletion,
+ api_base=api_base,
+ custom_prompt_dict=custom_prompt_dict,
+ api_key=api_key,
+ tenant_id=tenant_id,
+ timeout=timeout,
+ )
+
+ if (
+ "stream" in optional_params
+ and optional_params["stream"] is True
+ and acompletion is False
+ ):
+ return _model_response
+ return _model_response
+
+
+def _complete_text_completion_codestral(
+ ctx: _CompletionDispatchContext,
+) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ custom_prompt_dict = ctx.custom_prompt_dict
+ litellm_params = ctx.litellm_params
+ logger_fn = ctx.logger_fn
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ optional_params = ctx.optional_params
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ api_base = (
+ api_base
+ or optional_params.pop("api_base", None)
+ or optional_params.pop("base_url", None)
+ or litellm.api_base
+ or "https://codestral.mistral.ai/v1/fim/completions"
+ )
+
+ api_key = api_key or litellm.api_key or get_secret("CODESTRAL_API_KEY")
+
+ text_completion_model_response = litellm.TextCompletionResponse(stream=stream)
+
+ _model_response = codestral_text_completions.completion( # type: ignore
+ model=model,
+ messages=messages,
+ model_response=text_completion_model_response,
+ print_verbose=print_verbose,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ logger_fn=logger_fn,
+ encoding=_get_encoding(),
+ logging_obj=logging,
+ acompletion=acompletion,
+ api_base=api_base,
+ custom_prompt_dict=custom_prompt_dict,
+ api_key=api_key,
+ timeout=timeout,
+ )
+
+ if (
+ "stream" in optional_params
+ and optional_params["stream"] is True
+ and acompletion is False
+ ):
+ return _model_response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract
+ return _model_response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract
+
+
+def _complete_text_completion_inception(
+ ctx: _CompletionDispatchContext,
+) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logger_fn = ctx.logger_fn
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ text_completion = ctx.text_completion
+ timeout = ctx.timeout
+
+ passed_api_base = (
+ api_base
+ or optional_params.pop("api_base", None)
+ or optional_params.pop("base_url", None)
+ )
+ api_base = (
+ passed_api_base
+ or get_secret_str("INCEPTION_API_BASE")
+ or "https://api.inceptionlabs.ai/v1"
+ )
+ # FIM is served at `/v1/fim/completions`; the OpenAI client appends
+ # `/completions`, so point it at the `/v1/fim` base.
+ api_base = api_base.rstrip("/")
+ if not api_base.endswith("/fim"):
+ api_base += "/fim"
+
+ # Don't forward the server-managed Inception key to a caller-supplied
+ # api_base; only resolve it for the default/server base, or when the
+ # caller passes their own key.
+ if passed_api_base is None or api_key:
+ api_key = (
+ api_key or litellm.inception_key or get_secret_str("INCEPTION_API_KEY")
+ )
+
+ _response = openai_text_completions.completion(
+ model=model,
+ messages=messages,
+ model_response=model_response,
+ print_verbose=print_verbose,
+ api_key=api_key, # type: ignore[arg-type]
+ custom_llm_provider="text-completion-inception",
+ api_base=api_base,
+ acompletion=acompletion,
+ client=client,
+ logging_obj=logging,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ logger_fn=logger_fn,
+ timeout=timeout, # type: ignore
+ )
+
+ if (
+ optional_params.get("stream", False) is False
+ and acompletion is False
+ and text_completion is False
+ ):
+ _response = (
+ litellm.OpenAITextCompletionConfig().convert_to_chat_model_response_object(
+ response_object=_response, model_response_object=model_response
+ )
+ )
+
+ if optional_params.get("stream", False) or acompletion is True:
+ logging.post_call(
+ input=messages,
+ api_key=api_key,
+ original_response=_response,
+ additional_args={"headers": headers},
+ )
+ return _response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract
+
+
+def _complete_sagemaker_chat(
+ ctx: _CompletionDispatchContext,
+) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ return base_llm_http_handler.completion(
+ model=model,
+ stream=stream,
+ messages=messages,
+ acompletion=acompletion,
+ api_base=api_base,
+ model_response=model_response,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ custom_llm_provider=custom_llm_provider,
+ timeout=timeout,
+ headers=headers,
+ encoding=_get_encoding(),
+ api_key=api_key,
+ logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
+ client=client,
+ )
+
+
+def _complete_sagemaker(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ custom_prompt_dict = ctx.custom_prompt_dict
+ hf_model_name = ctx.hf_model_name
+ litellm_params = ctx.litellm_params
+ logger_fn = ctx.logger_fn
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+
+ return sagemaker_llm.completion(
+ model=model,
+ messages=messages,
+ model_response=model_response,
+ print_verbose=print_verbose,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ custom_prompt_dict=custom_prompt_dict,
+ hf_model_name=hf_model_name,
+ logger_fn=logger_fn,
+ encoding=_get_encoding(),
+ logging_obj=logging,
+ acompletion=acompletion,
+ )
+
+
+def _complete_bedrock(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_prompt_dict = ctx.custom_prompt_dict
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logger_fn = ctx.logger_fn
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ provider_config = ctx.provider_config
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict
+
+ if "aws_bedrock_client" in optional_params:
+ verbose_logger.warning(
+ "'aws_bedrock_client' is a deprecated param. Please move to another auth method - https://docs.litellm.ai/docs/providers/bedrock#boto3---authentication."
+ )
+ # Extract credentials for legacy boto3 client and pass thru to httpx
+ aws_bedrock_client = optional_params.pop("aws_bedrock_client")
+ creds = aws_bedrock_client._get_credentials().get_frozen_credentials()
+
+ if creds.access_key:
+ optional_params["aws_access_key_id"] = creds.access_key
+ if creds.secret_key:
+ optional_params["aws_secret_access_key"] = creds.secret_key
+ if creds.token:
+ optional_params["aws_session_token"] = creds.token
+ if (
+ "aws_region_name" not in optional_params
+ or optional_params["aws_region_name"] is None
+ ):
+ optional_params["aws_region_name"] = aws_bedrock_client.meta.region_name
+
+ bedrock_route = BedrockModelInfo.get_bedrock_route(model)
+ if bedrock_route == "claude_platform":
+ provider_config = ProviderConfigManager.get_provider_chat_config(
+ model=model,
+ provider=LlmProviders.BEDROCK,
+ )
+ model = BedrockModelInfo.get_claude_platform_model(model)
+ return base_llm_http_handler.completion(
+ model=model,
+ stream=stream,
+ messages=messages,
+ acompletion=acompletion,
+ api_base=api_base,
+ model_response=model_response,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ custom_llm_provider="bedrock",
+ timeout=timeout,
+ headers=headers,
+ encoding=_get_encoding(),
+ api_key=api_key,
+ logging_obj=logging,
+ client=client,
+ provider_config=provider_config,
+ )
+ elif bedrock_route == "converse":
+ model = model.replace("converse/", "")
+ response = bedrock_converse_chat_completion.completion(
+ model=model,
+ messages=messages,
+ custom_prompt_dict=custom_prompt_dict,
+ model_response=model_response,
+ optional_params=optional_params,
+ litellm_params=litellm_params, # type: ignore
+ logger_fn=logger_fn,
+ encoding=_get_encoding(),
+ logging_obj=logging,
+ extra_headers=headers, # Use merged headers instead of original extra_headers
+ timeout=timeout,
+ acompletion=acompletion,
+ client=client,
+ api_base=api_base,
+ api_key=api_key,
+ )
+ elif bedrock_route == "converse_like":
+ model = model.replace("converse_like/", "")
+ response = base_llm_http_handler.completion(
+ model=model,
+ stream=stream,
+ messages=messages,
+ acompletion=acompletion,
+ api_base=api_base,
+ model_response=model_response,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ custom_llm_provider="bedrock",
+ timeout=timeout,
+ headers=headers,
+ encoding=_get_encoding(),
+ api_key=api_key,
+ logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
+ client=client,
+ )
+ else:
+ response = base_llm_http_handler.completion(
+ model=model,
+ stream=stream,
+ messages=messages,
+ acompletion=acompletion,
+ api_base=api_base,
+ model_response=model_response,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ custom_llm_provider="bedrock",
+ timeout=timeout,
+ headers=headers,
+ encoding=_get_encoding(),
+ api_key=api_key,
+ logging_obj=logging,
+ client=client,
+ )
+
+ return response
+
+
+def _complete_watsonx(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_prompt_dict = ctx.custom_prompt_dict
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logger_fn = ctx.logger_fn
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ timeout = ctx.timeout
+
+ return watsonx_chat_completion.completion(
+ model=model,
+ messages=messages,
+ headers=headers,
+ model_response=model_response,
+ print_verbose=print_verbose,
+ api_key=api_key,
+ api_base=api_base,
+ acompletion=acompletion,
+ logging_obj=logging,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ logger_fn=logger_fn,
+ timeout=timeout, # type: ignore
+ custom_prompt_dict=custom_prompt_dict,
+ client=client, # pass AsyncOpenAI, OpenAI client
+ encoding=_get_encoding(),
+ custom_llm_provider="watsonx",
+ )
+
+
+def _complete_watsonx_text(
+ ctx: _CompletionDispatchContext,
+) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ api_key = (
+ api_key
+ or optional_params.pop("apikey", None)
+ or get_secret_str("WATSONX_APIKEY")
+ or get_secret_str("WATSONX_API_KEY")
+ or get_secret_str("WX_API_KEY")
+ )
+
+ api_base = (
+ api_base
+ or optional_params.pop(
+ "url",
+ optional_params.pop("api_base", optional_params.pop("base_url", None)),
+ )
+ or get_secret_str("WATSONX_API_BASE")
+ or get_secret_str("WATSONX_URL")
+ or get_secret_str("WX_URL")
+ or get_secret_str("WML_URL")
+ )
+
+ wx_credentials = optional_params.pop(
+ "wx_credentials",
+ optional_params.pop(
+ "watsonx_credentials", None
+ ), # follow {provider}_credentials, same as vertex ai
+ )
+
+ token: Optional[str] = None
+ if wx_credentials is not None:
+ api_base = wx_credentials.get("url", api_base)
+ api_key = wx_credentials.get("apikey", wx_credentials.get("api_key", api_key))
+ token = wx_credentials.get(
+ "token",
+ wx_credentials.get(
+ "watsonx_token", None
+ ), # follow format of {provider}_token, same as azure - e.g. 'azure_ad_token=..'
+ )
+
+ if token is not None:
+ optional_params["token"] = token
+
+ return base_llm_http_handler.completion(
+ model=model,
+ stream=stream,
+ messages=messages,
+ acompletion=acompletion,
+ api_base=api_base,
+ model_response=model_response,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ custom_llm_provider="watsonx_text",
+ timeout=timeout,
+ headers=headers,
+ encoding=_get_encoding(),
+ api_key=api_key,
+ logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
+ client=client,
+ )
+
+
+def _complete_vllm(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ custom_prompt_dict = ctx.custom_prompt_dict
+ litellm_params = ctx.litellm_params
+ logger_fn = ctx.logger_fn
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+
+ custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict
+ model_response = vllm_handler.completion(
+ model=model,
+ messages=messages,
+ custom_prompt_dict=custom_prompt_dict,
+ model_response=model_response,
+ print_verbose=print_verbose,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ logger_fn=logger_fn,
+ encoding=_get_encoding(),
+ logging_obj=logging,
+ )
+
+ if "stream" in optional_params and optional_params["stream"] is True: ## [BETA]
+ # don't try to access stream object,
+ return CustomStreamWrapper(
+ model_response,
+ model,
+ custom_llm_provider="vllm",
+ logging_obj=logging,
+ )
+
+ ## RESPONSE OBJECT
+ return model_response
+
+
+def _complete_ollama(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ api_base = (
+ litellm.api_base
+ or api_base
+ or get_secret("OLLAMA_API_BASE")
+ or "http://localhost:11434"
+ )
+ if api_key is not None and "Authorization" not in headers:
+ headers["Authorization"] = f"Bearer {api_key}"
+
+ return base_llm_http_handler.completion(
+ model=model,
+ stream=stream,
+ messages=messages,
+ acompletion=acompletion,
+ api_base=api_base,
+ model_response=model_response,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ custom_llm_provider="ollama",
+ timeout=timeout,
+ headers=headers,
+ encoding=_get_encoding(),
+ api_key=api_key,
+ logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
+ client=client,
+ )
+
+
+def _complete_ollama_chat(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ api_base = (
+ litellm.api_base
+ or api_base
+ or get_secret("OLLAMA_API_BASE")
+ or "http://localhost:11434"
+ )
+
+ api_key = (
+ api_key
+ or litellm.ollama_key
+ or os.environ.get("OLLAMA_API_KEY")
+ or litellm.api_key
+ )
+ if api_key is not None and "Authorization" not in headers:
+ headers["Authorization"] = f"Bearer {api_key}"
+
+ return base_llm_http_handler.completion(
+ model=model,
+ stream=stream,
+ messages=messages,
+ acompletion=acompletion,
+ api_base=api_base,
+ model_response=model_response,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ custom_llm_provider="ollama_chat",
+ timeout=timeout,
+ headers=headers,
+ encoding=_get_encoding(),
+ api_key=api_key,
+ logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
+ client=client,
+ )
+
+
+def _complete_triton(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ api_base = litellm.api_base or api_base
+ return base_llm_http_handler.completion(
+ model=model,
+ stream=stream,
+ messages=messages,
+ acompletion=acompletion,
+ api_base=api_base,
+ model_response=model_response,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ custom_llm_provider=custom_llm_provider,
+ timeout=timeout,
+ headers=headers,
+ encoding=_get_encoding(),
+ api_key=api_key,
+ logging_obj=logging,
+ )
+
+
+def _complete_cloudflare(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ custom_prompt_dict = ctx.custom_prompt_dict
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ api_key = (
+ api_key
+ or litellm.cloudflare_api_key
+ or litellm.api_key
+ or get_secret("CLOUDFLARE_API_KEY")
+ )
+ api_base = api_base or litellm.api_base or get_secret("CLOUDFLARE_API_BASE")
+
+ custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict
+ return base_llm_http_handler.completion(
+ model=model,
+ stream=stream,
+ messages=messages,
+ acompletion=acompletion,
+ api_base=api_base,
+ model_response=model_response,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ custom_llm_provider="cloudflare",
+ timeout=timeout,
+ headers=headers,
+ encoding=_get_encoding(),
+ api_key=api_key,
+ logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
+ )
+
+
+def _complete_petals(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ api_base = ctx.api_base
+ client = ctx.client
+ litellm_params = ctx.litellm_params
+ logger_fn = ctx.logger_fn
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ stream = ctx.stream
+
+ api_base = api_base or litellm.api_base
+
+ stream = optional_params.pop("stream", False)
+ model_response = petals_handler.completion(
+ model=model,
+ messages=messages,
+ api_base=api_base,
+ model_response=model_response,
+ print_verbose=print_verbose,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ logger_fn=logger_fn,
+ encoding=_get_encoding(),
+ logging_obj=logging,
+ client=client,
+ )
+ if stream is True: ## [BETA]
+ # Fake streaming for petals
+ resp_string = model_response["choices"][0]["message"]["content"]
+ return CustomStreamWrapper(
+ resp_string,
+ model,
+ custom_llm_provider="petals",
+ logging_obj=logging,
+ )
+ return model_response
+
+
+def _complete_snowflake(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ try:
+ client = (
+ HTTPHandler(timeout=timeout) if stream is False else None
+ ) # Keep this here, otherwise, the httpx.client closes and streaming is impossible
+ response = base_llm_http_handler.completion(
+ model=model,
+ messages=messages,
+ headers=headers,
+ model_response=model_response,
+ api_key=api_key,
+ api_base=api_base,
+ acompletion=acompletion,
+ logging_obj=logging,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ timeout=timeout, # type: ignore
+ client=client,
+ custom_llm_provider=custom_llm_provider,
+ encoding=_get_encoding(),
+ stream=stream,
+ )
+
+ except Exception as e:
+ ## LOGGING - log the original exception returned
+ logging.post_call(
+ input=messages,
+ api_key=api_key,
+ original_response=str(e),
+ additional_args={"headers": headers},
+ )
+ raise e
+
+ return response
+
+
+def _complete_gradient_ai(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ api_base = litellm.api_base or api_base
+ return base_llm_http_handler.completion(
+ model=model,
+ stream=stream,
+ messages=messages,
+ acompletion=acompletion,
+ api_base=api_base,
+ model_response=model_response,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ custom_llm_provider="gradient_ai",
+ timeout=timeout,
+ headers=headers,
+ encoding=_get_encoding(),
+ api_key=api_key,
+ logging_obj=logging,
+ )
+
+
+def _complete_bytez(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ api_key = (
+ api_key
+ or litellm.bytez_key
+ or get_secret_str("BYTEZ_API_KEY")
+ or litellm.api_key
+ )
+
+ response = base_llm_http_handler.completion(
+ model=model,
+ messages=messages,
+ headers=headers,
+ model_response=model_response,
+ api_key=api_key,
+ api_base=api_base,
+ acompletion=acompletion,
+ logging_obj=logging,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ timeout=timeout, # type: ignore
+ client=client,
+ custom_llm_provider=custom_llm_provider,
+ encoding=_get_encoding(),
+ stream=stream,
+ provider_config=bytez_transformation,
+ )
+
+ pass
+
+ return response
+
+
+def _complete_lemonade(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ api_key = (
+ api_key
+ or litellm.lemonade_key
+ or get_secret_str("LEMONADE_API_KEY")
+ or litellm.api_key
+ )
+
+ response = base_llm_http_handler.completion(
+ model=model,
+ messages=messages,
+ headers=headers,
+ model_response=model_response,
+ api_key=api_key,
+ api_base=api_base,
+ acompletion=acompletion,
+ logging_obj=logging,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ timeout=timeout, # type: ignore
+ client=client,
+ custom_llm_provider=custom_llm_provider,
+ encoding=_get_encoding(),
+ stream=stream,
+ provider_config=lemonade_transformation,
+ )
+
+ pass
+
+ return response
+
+
+def _complete_ovhcloud(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ api_key = (
+ api_key
+ or litellm.ovhcloud_key
+ or get_secret_str("OVHCLOUD_API_KEY")
+ or litellm.api_key
+ )
+
+ api_base = (
+ api_base
+ or litellm.api_base
+ or get_secret_str("OVHCLOUD_API_BASE")
+ or "https://oai.endpoints.kepler.ai.cloud.ovh.net/v1"
+ )
+
+ response = base_llm_http_handler.completion(
+ model=model,
+ messages=messages,
+ headers=headers,
+ model_response=model_response,
+ api_key=api_key,
+ api_base=api_base,
+ acompletion=acompletion,
+ logging_obj=logging,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ timeout=timeout, # type: ignore
+ client=client,
+ custom_llm_provider=custom_llm_provider,
+ encoding=_get_encoding(),
+ stream=stream,
+ provider_config=ovhcloud_transformation,
+ )
+
+ pass
+
+ return response
+
+
+def _complete_custom(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ api_base = ctx.api_base
+ headers = ctx.headers
+ kwargs = ctx.kwargs
+ max_tokens = ctx.max_tokens
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ temperature = ctx.temperature
+ top_p = ctx.top_p
+
+ url = litellm.api_base or api_base or ""
+ if url is None or url == "":
+ raise ValueError(
+ "api_base not set. Set api_base or litellm.api_base for custom endpoints"
+ )
+
+ """
+ assume input to custom LLM api bases follow this format:
+ resp = litellm.module_level_client.post(
+ api_base,
+ json={
+ 'model': 'meta-llama/Llama-2-13b-hf', # model name
+ 'params': {
+ 'prompt': ["The capital of France is P"],
+ 'max_tokens': 32,
+ 'temperature': 0.7,
+ 'top_p': 1.0,
+ 'top_k': 40,
+ }
+ }
+ )
+
+ """
+ prompt = " ".join([message["content"] for message in messages]) # type: ignore
+ resp = litellm.module_level_client.post(
+ url,
+ headers=headers,
+ json={
+ "model": model,
+ "params": {
+ "prompt": [prompt],
+ "max_tokens": max_tokens,
+ "temperature": temperature,
+ "top_p": top_p,
+ "top_k": kwargs.get("top_k"),
+ },
+ **kwargs.get("extra_body", {}),
+ },
+ )
+ response_json = resp.json()
+ """
+ assume all responses from custom api_bases of this format:
+ {
+ 'data': [
+ {
+ 'prompt': 'The capital of France is P',
+ 'output': ['The capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France'],
+ 'params': {'temperature': 0.7, 'top_k': 40, 'top_p': 1}}],
+ 'message': 'ok'
+ }
+ ]
+ }
+ """
+ string_response = response_json["data"][0]["output"][0]
+ ## RESPONSE OBJECT
+ model_response.choices[0].message.content = string_response # type: ignore
+ model_response.created = int(time.time())
+ model_response.model = model
+ return model_response
+
+
+def _complete_custom_providers(
+ ctx: _CompletionDispatchContext,
+) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ custom_prompt_dict = ctx.custom_prompt_dict
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logger_fn = ctx.logger_fn
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ custom_handler: Optional[CustomLLM] = None
+ for item in litellm.custom_provider_map:
+ if item["provider"] == custom_llm_provider:
+ custom_handler = item["custom_handler"]
+
+ if custom_handler is None:
+ raise LiteLLMUnknownProvider(
+ model=model, custom_llm_provider=custom_llm_provider
+ )
+
+ ## ROUTE LLM CALL ##
+ handler_fn = custom_chat_llm_router(
+ async_fn=acompletion, stream=stream, custom_llm=custom_handler
+ )
+
+ headers = headers or litellm.headers or {}
+
+ ## CALL FUNCTION
+ response = handler_fn(
+ model=model,
+ messages=messages,
+ headers=headers,
+ model_response=model_response,
+ print_verbose=print_verbose,
+ api_key=api_key,
+ api_base=api_base,
+ acompletion=acompletion,
+ logging_obj=logging,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ logger_fn=logger_fn,
+ timeout=timeout, # type: ignore
+ custom_prompt_dict=custom_prompt_dict,
+ client=client, # pass AsyncOpenAI, OpenAI client
+ encoding=_get_encoding(),
+ )
+ if stream is True:
+ return CustomStreamWrapper(
+ completion_stream=response,
+ model=model,
+ custom_llm_provider=custom_llm_provider,
+ logging_obj=logging,
+ )
+
+ return response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract
+
+
+def _complete_langgraph(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ from litellm.llms.langgraph.chat.transformation import LangGraphConfig
+
+ (
+ api_base,
+ api_key,
+ ) = LangGraphConfig()._get_openai_compatible_provider_info(
+ api_base=api_base or litellm.api_base,
+ api_key=api_key or litellm.api_key,
+ )
+
+ headers = headers or litellm.headers
+
+ return base_llm_http_handler.completion(
+ model=model,
+ stream=stream,
+ messages=messages,
+ acompletion=acompletion,
+ api_base=api_base,
+ model_response=model_response,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ custom_llm_provider=custom_llm_provider,
+ timeout=timeout,
+ headers=headers,
+ encoding=_get_encoding(),
+ api_key=api_key,
+ logging_obj=logging,
+ client=client,
+ )
+
+
+def _complete_langflow(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
+ acompletion = ctx.acompletion
+ api_base = ctx.api_base
+ api_key = ctx.api_key
+ client = ctx.client
+ custom_llm_provider = ctx.custom_llm_provider
+ headers = ctx.headers
+ litellm_params = ctx.litellm_params
+ logging = ctx.logging
+ messages = ctx.messages
+ model = ctx.model
+ model_response = ctx.model_response
+ optional_params = ctx.optional_params
+ shared_session = ctx.shared_session
+ stream = ctx.stream
+ timeout = ctx.timeout
+
+ from litellm.llms.langflow.chat.transformation import LangFlowConfig
+
+ (
+ api_base,
+ api_key,
+ ) = LangFlowConfig()._get_openai_compatible_provider_info(
+ api_base=api_base or litellm.api_base,
+ api_key=api_key or litellm.api_key,
+ )
+
+ headers = headers or litellm.headers
+
+ return base_llm_http_handler.completion(
+ model=model,
+ stream=stream,
+ messages=messages,
+ acompletion=acompletion,
+ api_base=api_base,
+ model_response=model_response,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ custom_llm_provider=custom_llm_provider,
+ timeout=timeout,
+ headers=headers,
+ encoding=_get_encoding(),
+ api_key=api_key,
+ logging_obj=logging,
+ client=client,
+ )
+
+
@tracer.wrap()
@client
def completion( # type: ignore
@@ -1215,9 +5077,7 @@ def completion( # type: ignore
if LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway(
tools=tools_for_mcp
):
- # Return coroutine - acompletion will await it
- # completion() can return a coroutine when MCP tools are present, which acompletion() awaits
- return acompletion_with_mcp( # type: ignore[return-value]
+ return acompletion_with_mcp( # pyright: ignore[reportReturnType] # MCP path returns a coroutine that acompletion() awaits; completion()'s sync return type omits it
model=model,
messages=messages,
functions=functions,
@@ -1389,12 +5249,16 @@ def completion( # type: ignore
logging: LiteLLMLoggingObj = cast(LiteLLMLoggingObj, litellm_logging_obj)
fallbacks = fallbacks or litellm.model_fallbacks
if fallbacks is not None:
- return completion_with_fallbacks(**args)
+ return completion_with_fallbacks( # pyright: ignore[reportReturnType] # fallback runner is untyped; resolves to ModelResponse|CustomStreamWrapper at runtime
+ **args
+ )
if model_list is not None:
deployments = [
m["litellm_params"] for m in model_list if m["model_name"] == model
]
- return litellm.batch_completion_models(deployments=deployments, **args)
+ return litellm.batch_completion_models( # pyright: ignore[reportReturnType] # batch path returns a list of responses, outside completion()'s single-response return type
+ deployments=deployments, **args
+ )
if litellm.model_alias_map and model in litellm.model_alias_map:
model = litellm.model_alias_map[
model
@@ -1454,7 +5318,7 @@ def completion( # type: ignore
timeout,
kwargs,
custom_llm_provider,
- global_timeout=getattr(litellm, "request_timeout", None),
+ global_timeout=get_configured_request_timeout(),
supports_httpx_timeout=supports_httpx_timeout,
)
@@ -1716,7 +5580,7 @@ def completion( # type: ignore
else:
optional_params["reasoning_effort"] = {"summary": rs_val}
- return responses_api_bridge.completion(
+ return responses_api_bridge.completion( # pyright: ignore[reportReturnType] # bridge returns a coroutine on the acompletion path; awaited by the async caller
model=model,
messages=messages,
headers=headers,
@@ -1746,375 +5610,52 @@ def completion( # type: ignore
optional_params
)
+ _dispatch_ctx = _CompletionDispatchContext(
+ _azure_detection_model=_azure_detection_model,
+ acompletion=acompletion,
+ api_base=api_base,
+ api_key=api_key,
+ api_version=api_version,
+ client=client,
+ custom_llm_provider=custom_llm_provider,
+ custom_prompt_dict=custom_prompt_dict,
+ extra_headers=extra_headers,
+ headers=headers,
+ hf_model_name=hf_model_name,
+ kwargs=kwargs,
+ litellm_params=litellm_params,
+ logger_fn=logger_fn,
+ logging=logging,
+ max_retries=max_retries,
+ max_tokens=max_tokens,
+ messages=messages,
+ metadata=metadata,
+ model=model,
+ model_response=model_response,
+ optional_params=optional_params,
+ organization=organization,
+ provider_config=provider_config,
+ shared_session=shared_session,
+ stream=stream,
+ temperature=temperature,
+ text_completion=text_completion,
+ timeout=timeout,
+ top_p=top_p,
+ )
if custom_llm_provider == "azure":
# azure configs
## check dynamic params ##
- dynamic_params = False
- if client is not None and (
- isinstance(client, openai.AzureOpenAI)
- or isinstance(client, openai.AsyncAzureOpenAI)
- ):
- dynamic_params = _check_dynamic_azure_params(
- azure_client_params={"api_version": api_version},
- azure_client=client,
- )
-
- api_type = get_secret("AZURE_API_TYPE") or "azure"
-
- api_base = api_base or litellm.api_base or get_secret("AZURE_API_BASE")
-
- api_version = (
- api_version
- or litellm.api_version
- or get_secret_str("AZURE_API_VERSION")
- or litellm.AZURE_DEFAULT_API_VERSION
- )
-
- api_key = (
- api_key
- or litellm.api_key
- or litellm.azure_key
- or get_secret_str("AZURE_OPENAI_API_KEY")
- or get_secret_str("AZURE_API_KEY")
- )
-
- azure_ad_token = optional_params.get("extra_body", {}).pop(
- "azure_ad_token", None
- ) or get_secret_str("AZURE_AD_TOKEN")
-
- azure_ad_token_provider = litellm_params.get(
- "azure_ad_token_provider", None
- )
-
- headers = headers or litellm.headers
-
- if extra_headers is not None:
- optional_params["extra_headers"] = extra_headers
- if max_retries is not None:
- optional_params["max_retries"] = max_retries
-
- if litellm.AzureOpenAIO1Config().is_o_series_model(
- model=_azure_detection_model
- ):
- ## LOAD CONFIG - if set
- config = litellm.AzureOpenAIO1Config.get_config()
- for k, v in config.items():
- if (
- k not in optional_params
- ): # completion(top_k=3) > azure_config(top_k=3) <- allows for dynamic variables to be passed in
- optional_params[k] = v
-
- response = azure_o1_chat_completions.completion(
- model=model,
- messages=messages,
- headers=headers,
- api_key=api_key,
- api_base=api_base,
- api_version=api_version,
- dynamic_params=dynamic_params,
- azure_ad_token=azure_ad_token,
- model_response=model_response,
- print_verbose=print_verbose,
- optional_params=optional_params,
- litellm_params=litellm_params,
- logger_fn=logger_fn,
- logging_obj=logging,
- acompletion=acompletion,
- timeout=timeout, # type: ignore
- client=client, # pass AsyncAzureOpenAI, AzureOpenAI client
- custom_llm_provider=custom_llm_provider,
- )
- else:
- ## LOAD CONFIG - if set
- config = litellm.AzureOpenAIConfig.get_config()
- for k, v in config.items():
- if (
- k not in optional_params
- ): # completion(top_k=3) > azure_config(top_k=3) <- allows for dynamic variables to be passed in
- optional_params[k] = v
-
- ## COMPLETION CALL
- response = azure_chat_completions.completion(
- model=model,
- messages=messages,
- headers=headers,
- api_key=api_key,
- api_base=api_base,
- api_version=api_version,
- api_type=api_type,
- dynamic_params=dynamic_params,
- azure_ad_token=azure_ad_token,
- azure_ad_token_provider=azure_ad_token_provider,
- model_response=model_response,
- print_verbose=print_verbose,
- optional_params=optional_params,
- litellm_params=litellm_params,
- logger_fn=logger_fn,
- logging_obj=logging,
- acompletion=acompletion,
- timeout=timeout, # type: ignore
- client=client, # pass AsyncAzureOpenAI, AzureOpenAI client
- )
-
- if optional_params.get("stream", False):
- ## LOGGING
- logging.post_call(
- input=messages,
- api_key=api_key,
- original_response=response,
- additional_args={
- "headers": headers,
- "api_version": api_version,
- "api_base": api_base,
- },
- )
+ response = _complete_azure(_dispatch_ctx)
elif custom_llm_provider == "azure_text":
# azure configs
- api_type = get_secret_str("AZURE_API_TYPE") or "azure"
-
- api_base = api_base or litellm.api_base or get_secret_str("AZURE_API_BASE")
-
- if api_base is None:
- raise ValueError(
- "api_base is required for Azure OpenAI LLM provider. Either set it dynamically or set the AZURE_API_BASE environment variable."
- )
-
- api_version = (
- api_version
- or litellm.api_version
- or get_secret_str("AZURE_API_VERSION")
- )
-
- api_key = (
- api_key
- or litellm.api_key
- or litellm.azure_key
- or get_secret_str("AZURE_OPENAI_API_KEY")
- or get_secret_str("AZURE_API_KEY")
- )
-
- azure_ad_token = optional_params.get("extra_body", {}).pop(
- "azure_ad_token", None
- ) or get_secret_str("AZURE_AD_TOKEN")
-
- azure_ad_token_provider = litellm_params.get(
- "azure_ad_token_provider", None
- )
-
- headers = headers or litellm.headers
-
- if extra_headers is not None:
- optional_params["extra_headers"] = extra_headers
-
- ## LOAD CONFIG - if set
- config = litellm.AzureOpenAIConfig.get_config()
- for k, v in config.items():
- if (
- k not in optional_params
- ): # completion(top_k=3) > azure_config(top_k=3) <- allows for dynamic variables to be passed in
- optional_params[k] = v
-
- ## COMPLETION CALL
- response = azure_text_completions.completion(
- model=model,
- messages=messages,
- headers=headers,
- api_key=api_key,
- api_base=api_base,
- api_version=cast(str, api_version),
- api_type=api_type,
- azure_ad_token=azure_ad_token,
- azure_ad_token_provider=azure_ad_token_provider,
- model_response=model_response,
- print_verbose=print_verbose,
- optional_params=optional_params,
- litellm_params=litellm_params,
- logger_fn=logger_fn,
- logging_obj=logging,
- acompletion=acompletion,
- timeout=timeout,
- client=client, # pass AsyncAzureOpenAI, AzureOpenAI client
- )
-
- if optional_params.get("stream", False) or acompletion is True:
- ## LOGGING
- logging.post_call(
- input=messages,
- api_key=api_key,
- original_response=response,
- additional_args={
- "headers": headers,
- "api_version": api_version,
- "api_base": api_base,
- },
- )
+ response = _complete_azure_text(_dispatch_ctx)
elif custom_llm_provider == "deepseek":
## COMPLETION CALL
- try:
- response = base_llm_http_handler.completion(
- model=model,
- messages=messages,
- headers=headers,
- model_response=model_response,
- api_key=api_key,
- api_base=api_base,
- acompletion=acompletion,
- logging_obj=logging,
- optional_params=optional_params,
- litellm_params=litellm_params,
- shared_session=shared_session,
- timeout=timeout, # type: ignore
- client=client,
- custom_llm_provider=custom_llm_provider,
- encoding=_get_encoding(),
- stream=stream,
- provider_config=provider_config,
- )
- except Exception as e:
- ## LOGGING - log the original exception returned
- logging.post_call(
- input=messages,
- api_key=api_key,
- original_response=str(e),
- additional_args={"headers": headers},
- )
- raise e
+ response = _complete_deepseek(_dispatch_ctx)
elif custom_llm_provider == "azure_ai":
- from litellm.llms.azure_ai.common_utils import AzureFoundryModelInfo
-
- azure_ai_route = AzureFoundryModelInfo.get_azure_ai_route(model)
-
- # Check if this is an agents route - model format: azure_ai/agents/
- if azure_ai_route == "agents":
- from litellm.llms.azure_ai.agents import AzureAIAgentsConfig
-
- api_base = AzureFoundryModelInfo.get_api_base(api_base)
- if api_base is None:
- raise ValueError(
- "Azure AI Agents requests require an api_base. "
- "Set `api_base` or the AZURE_AI_API_BASE env var."
- )
- api_key = AzureFoundryModelInfo.get_api_key(api_key)
-
- response = AzureAIAgentsConfig.completion(
- model=model,
- messages=messages,
- api_base=api_base,
- api_key=api_key,
- model_response=model_response,
- logging_obj=logging,
- optional_params=optional_params,
- litellm_params=litellm_params,
- timeout=timeout,
- acompletion=acompletion,
- stream=stream,
- headers=headers or litellm.headers,
- )
-
- # Check if this is a Claude model - route to Azure Anthropic handler
- elif "claude" in model.lower():
- # Use Azure Anthropic handler for Claude models
- api_base = AzureFoundryModelInfo.get_api_base(api_base)
- if api_base is None:
- raise ValueError(
- "Azure Anthropic requests require an api_base. "
- "Set `api_base` or the AZURE_AI_API_BASE env var."
- )
- api_key = AzureFoundryModelInfo.get_api_key(api_key)
-
- # Ensure the URL ends with /v1/messages for Anthropic
- if api_base:
- api_base = api_base.rstrip("/")
- if not api_base.endswith("/v1/messages"):
- if "/anthropic" in api_base:
- parts = api_base.split("/anthropic", 1)
- api_base = parts[0] + "/anthropic"
- else:
- api_base = api_base + "/anthropic"
- api_base = api_base + "/v1/messages"
-
- response = azure_anthropic_chat_completions.completion(
- model=model,
- messages=messages,
- api_base=api_base,
- acompletion=acompletion,
- custom_prompt_dict=litellm.custom_prompt_dict,
- model_response=model_response,
- print_verbose=print_verbose,
- optional_params=optional_params,
- litellm_params=litellm_params,
- logger_fn=logger_fn,
- encoding=_get_encoding(),
- api_key=api_key,
- logging_obj=logging,
- headers=headers,
- timeout=timeout,
- client=client,
- custom_llm_provider=custom_llm_provider,
- )
- if optional_params.get("stream", False) or acompletion is True:
- ## LOGGING
- logging.post_call(
- input=messages,
- api_key=api_key,
- original_response=response,
- )
- response = response
- else:
- # Non-Claude models use standard Azure AI flow
- api_base = AzureFoundryModelInfo.get_api_base(api_base)
- # set API KEY
- api_key = AzureFoundryModelInfo.get_api_key(api_key)
-
- headers = headers or litellm.headers
-
- if extra_headers is not None:
- optional_params["extra_headers"] = extra_headers
-
- ## FOR COHERE
- if "command-r" in model: # make sure tool call in messages are str
- messages = stringify_json_tool_call_content(messages=messages)
-
- ## COMPLETION CALL
- try:
- response = base_llm_http_handler.completion(
- model=model,
- messages=messages,
- headers=headers,
- model_response=model_response,
- api_key=api_key,
- api_base=api_base,
- acompletion=acompletion,
- logging_obj=logging,
- optional_params=optional_params,
- litellm_params=litellm_params,
- shared_session=shared_session,
- timeout=timeout, # type: ignore
- client=client, # pass AsyncOpenAI, OpenAI client
- custom_llm_provider=custom_llm_provider,
- encoding=_get_encoding(),
- stream=stream,
- )
- except Exception as e:
- ## LOGGING - log the original exception returned
- logging.post_call(
- input=messages,
- api_key=api_key,
- original_response=str(e),
- additional_args={"headers": headers},
- )
- raise e
-
- if optional_params.get("stream", False):
- ## LOGGING
- logging.post_call(
- input=messages,
- api_key=api_key,
- original_response=response,
- additional_args={"headers": headers},
- )
+ response = _complete_azure_ai(_dispatch_ctx)
elif (
custom_llm_provider == "text-completion-openai"
or "ft:babbage-002" in model
@@ -2123,535 +5664,42 @@ def completion( # type: ignore
in litellm.openai_text_completion_compatible_providers
and kwargs.get("text_completion") is True
):
- openai.api_type = "openai"
-
- api_base = (
- api_base
- or litellm.api_base
- or get_secret("OPENAI_BASE_URL")
- or get_secret("OPENAI_API_BASE")
- or "https://api.openai.com/v1"
- )
-
- openai.api_version = None
- # set API KEY
-
- api_key = (
- api_key
- or litellm.api_key
- or litellm.openai_key
- or get_secret("OPENAI_API_KEY")
- )
-
- headers = headers or litellm.headers
-
- ## LOAD CONFIG - if set
- config = litellm.OpenAITextCompletionConfig.get_config()
- for k, v in config.items():
- if (
- k not in optional_params
- ): # completion(top_k=3) > openai_text_config(top_k=3) <- allows for dynamic variables to be passed in
- optional_params[k] = v
- if litellm.organization:
- openai.organization = litellm.organization
-
- if (
- len(messages) > 0
- and "content" in messages[0]
- and isinstance(messages[0]["content"], list)
- ):
- # text-davinci-003 can accept a string or array, if it's an array, assume the array is set in messages[0]['content']
- # https://platform.openai.com/docs/api-reference/completions/create
- prompt = messages[0]["content"]
- else:
- prompt = " ".join([message["content"] for message in messages]) # type: ignore
-
- ## COMPLETION CALL
- _response = openai_text_completions.completion(
- model=model,
- messages=messages,
- headers=headers,
- model_response=model_response,
- print_verbose=print_verbose,
- api_key=api_key,
- custom_llm_provider=custom_llm_provider,
- api_base=api_base,
- acompletion=acompletion,
- client=client, # pass AsyncOpenAI, OpenAI client
- logging_obj=logging,
- optional_params=optional_params,
- litellm_params=litellm_params,
- logger_fn=logger_fn,
- timeout=timeout, # type: ignore
- )
-
- if (
- optional_params.get("stream", False) is False
- and acompletion is False
- and text_completion is False
- ):
- # convert to chat completion response
- _response = litellm.OpenAITextCompletionConfig().convert_to_chat_model_response_object(
- response_object=_response, model_response_object=model_response
- )
-
- if optional_params.get("stream", False) or acompletion is True:
- ## LOGGING
- logging.post_call(
- input=messages,
- api_key=api_key,
- original_response=_response,
- additional_args={"headers": headers},
- )
- response = _response
+ response = _complete_text_completion_openai(_dispatch_ctx)
elif custom_llm_provider == "fireworks_ai":
## COMPLETION CALL
- try:
- response = base_llm_http_handler.completion(
- model=model,
- messages=messages,
- headers=headers,
- model_response=model_response,
- api_key=api_key,
- api_base=api_base,
- acompletion=acompletion,
- logging_obj=logging,
- optional_params=optional_params,
- litellm_params=litellm_params,
- shared_session=shared_session,
- timeout=timeout, # type: ignore
- client=client,
- custom_llm_provider=custom_llm_provider,
- encoding=_get_encoding(),
- stream=stream,
- provider_config=provider_config,
- )
- except Exception as e:
- ## LOGGING - log the original exception returned
- logging.post_call(
- input=messages,
- api_key=api_key,
- original_response=str(e),
- additional_args={"headers": headers},
- )
- raise e
+ response = _complete_fireworks_ai(_dispatch_ctx)
elif custom_llm_provider == "heroku":
- try:
- response = base_llm_http_handler.completion(
- model=model,
- messages=messages,
- headers=headers,
- model_response=model_response,
- api_key=api_key,
- api_base=api_base,
- acompletion=acompletion,
- logging_obj=logging,
- optional_params=optional_params,
- litellm_params=litellm_params,
- shared_session=shared_session,
- timeout=timeout,
- client=client,
- custom_llm_provider=custom_llm_provider,
- encoding=_get_encoding(),
- stream=stream,
- provider_config=provider_config,
- )
- except Exception as e:
- logging.post_call(
- input=messages,
- api_key=api_key,
- original_response=str(e),
- additional_args={"headers": headers},
- )
- raise e
+ response = _complete_heroku(_dispatch_ctx)
elif custom_llm_provider == "ragflow":
## COMPLETION CALL - RAGFlow uses HTTP handler to support custom URL paths
- try:
- response = base_llm_http_handler.completion(
- model=model,
- messages=messages,
- headers=headers,
- model_response=model_response,
- api_key=api_key,
- api_base=api_base,
- acompletion=acompletion,
- logging_obj=logging,
- optional_params=optional_params,
- litellm_params=litellm_params,
- shared_session=shared_session,
- timeout=timeout,
- client=client,
- custom_llm_provider=custom_llm_provider,
- encoding=_get_encoding(),
- stream=stream,
- provider_config=provider_config,
- )
- except Exception as e:
- logging.post_call(
- input=messages,
- api_key=api_key,
- original_response=str(e),
- additional_args={"headers": headers},
- )
- raise e
+ response = _complete_ragflow(_dispatch_ctx)
elif custom_llm_provider == "xai":
## COMPLETION CALL
- try:
- response = base_llm_http_handler.completion(
- model=model,
- messages=messages,
- headers=headers,
- model_response=model_response,
- api_key=api_key,
- api_base=api_base,
- acompletion=acompletion,
- logging_obj=logging,
- optional_params=optional_params,
- litellm_params=litellm_params,
- shared_session=shared_session,
- timeout=timeout, # type: ignore
- client=client,
- custom_llm_provider=custom_llm_provider,
- encoding=_get_encoding(),
- stream=stream,
- provider_config=provider_config,
- )
- except Exception as e:
- ## LOGGING - log the original exception returned
- logging.post_call(
- input=messages,
- api_key=api_key,
- original_response=str(e),
- additional_args={"headers": headers},
- )
- raise e
+ response = _complete_xai(_dispatch_ctx)
elif custom_llm_provider == "groq":
- api_base = (
- api_base # for deepinfra/perplexity/anyscale/groq/friendliai we check in get_llm_provider and pass in the api base from there
- or litellm.api_base
- or get_secret("GROQ_API_BASE")
- or "https://api.groq.com/openai/v1"
- )
-
- # set API KEY
- api_key = (
- api_key
- or litellm.api_key # for deepinfra/perplexity/anyscale/friendliai we check in get_llm_provider and pass in the api key from there
- or litellm.groq_key
- or get_secret("GROQ_API_KEY")
- )
-
- headers = headers or litellm.headers
-
- ## LOAD CONFIG - if set
- config = litellm.GroqChatConfig.get_config()
- for k, v in config.items():
- if (
- k not in optional_params
- ): # completion(top_k=3) > openai_config(top_k=3) <- allows for dynamic variables to be passed in
- optional_params[k] = v
-
- response = base_llm_http_handler.completion(
- model=model,
- stream=stream,
- messages=messages,
- acompletion=acompletion,
- api_base=api_base,
- model_response=model_response,
- optional_params=optional_params,
- litellm_params=litellm_params,
- shared_session=shared_session,
- custom_llm_provider=custom_llm_provider,
- timeout=timeout,
- headers=headers,
- encoding=_get_encoding(),
- api_key=api_key,
- logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
- client=client,
- )
+ response = _complete_groq(_dispatch_ctx)
elif custom_llm_provider == "bedrock_mantle":
- api_base = (
- api_base or litellm.api_base or get_secret("BEDROCK_MANTLE_API_BASE")
- )
- api_key = api_key or litellm.api_key or get_secret("BEDROCK_MANTLE_API_KEY")
- headers = headers or litellm.headers
- config = litellm.BedrockMantleChatConfig.get_config()
- for k, v in config.items():
- if k not in optional_params:
- optional_params[k] = v
- response = base_llm_http_handler.completion(
- model=model,
- stream=stream,
- messages=messages,
- acompletion=acompletion,
- api_base=api_base,
- model_response=model_response,
- optional_params=optional_params,
- litellm_params=litellm_params,
- shared_session=shared_session,
- custom_llm_provider=custom_llm_provider,
- timeout=timeout,
- headers=headers,
- encoding=_get_encoding(),
- api_key=api_key,
- logging_obj=logging,
- client=client,
- )
+ response = _complete_bedrock_mantle(_dispatch_ctx)
elif custom_llm_provider == "a2a":
# A2A (Agent-to-Agent) Protocol
# Resolve agent configuration from registry if model format is "a2a/"
- (
- api_base,
- api_key,
- headers,
- ) = litellm.A2AConfig.resolve_agent_config_from_registry(
- model=model,
- api_base=api_base,
- api_key=api_key,
- headers=headers,
- optional_params=optional_params,
- )
-
- # Fall back to environment variables and defaults
- api_base = api_base or litellm.api_base or get_secret_str("A2A_API_BASE")
-
- if api_base is None:
- raise Exception(
- "api_base is required for A2A provider. "
- "Either provide api_base parameter, set A2A_API_BASE environment variable, "
- "or register the agent in the proxy with model='a2a/'."
- )
-
- headers = headers or litellm.headers
-
- response = base_llm_http_handler.completion(
- model=model,
- stream=stream,
- messages=messages,
- acompletion=acompletion,
- api_base=api_base,
- model_response=model_response,
- optional_params=optional_params,
- litellm_params=litellm_params,
- shared_session=shared_session,
- custom_llm_provider=custom_llm_provider,
- timeout=timeout,
- headers=headers,
- encoding=_get_encoding(),
- api_key=api_key,
- logging_obj=logging,
- client=client,
- provider_config=provider_config,
- )
+ response = _complete_a2a(_dispatch_ctx)
elif custom_llm_provider == "gigachat":
# GigaChat - Sber AI's LLM (Russia)
- api_key = (
- api_key
- or litellm.api_key
- or litellm.gigachat_key
- or get_secret("GIGACHAT_API_KEY")
- or get_secret("GIGACHAT_CREDENTIALS")
- )
-
- headers = headers or litellm.headers or {}
-
- ## COMPLETION CALL
- try:
- response = base_llm_http_handler.completion(
- model=model,
- messages=messages,
- headers=headers,
- model_response=model_response,
- api_key=api_key,
- api_base=api_base,
- acompletion=acompletion,
- logging_obj=logging,
- optional_params=optional_params,
- litellm_params=litellm_params,
- shared_session=shared_session,
- timeout=timeout,
- client=client,
- custom_llm_provider=custom_llm_provider,
- encoding=_get_encoding(),
- stream=stream,
- provider_config=provider_config,
- )
- except Exception as e:
- ## LOGGING - log the original exception returned
- logging.post_call(
- input=messages,
- api_key=api_key,
- original_response=str(e),
- additional_args={"headers": headers},
- )
- raise e
+ response = _complete_gigachat(_dispatch_ctx)
elif custom_llm_provider == "sap":
- headers = headers or litellm.headers
- ## LOAD CONFIG - if set
- config = litellm.GenAIHubOrchestrationConfig.get_config()
- for k, v in config.items():
- if (
- k not in optional_params
- ): # completion(top_k=3) > openai_config(top_k=3) <- allows for dynamic variables to be passed in
- optional_params[k] = v
-
- response = sap_gen_ai_hub_chat_completions.completion(
- model=model,
- messages=messages,
- headers=headers,
- model_response=model_response,
- acompletion=acompletion,
- logging_obj=logging,
- optional_params=optional_params,
- litellm_params=litellm_params,
- timeout=timeout, # type: ignore
- shared_session=shared_session,
- client=client,
- custom_llm_provider=custom_llm_provider,
- encoding=_get_encoding(),
- api_key=api_key,
- api_base=api_base,
- stream=stream,
- )
+ response = _complete_sap(_dispatch_ctx)
elif custom_llm_provider == "aiohttp_openai":
# NEW aiohttp provider for 10-100x higher RPS
- api_base = (
- api_base # for deepinfra/perplexity/anyscale/groq/friendliai we check in get_llm_provider and pass in the api base from there
- or litellm.api_base
- or get_secret("OPENAI_BASE_URL")
- or get_secret("OPENAI_API_BASE")
- or "https://api.openai.com/v1"
- )
- # set API KEY
- api_key = (
- api_key
- or litellm.api_key # for deepinfra/perplexity/anyscale/friendliai we check in get_llm_provider and pass in the api key from there
- or litellm.openai_key
- or get_secret("OPENAI_API_KEY")
- )
-
- headers = headers or litellm.headers
-
- if extra_headers is not None:
- optional_params["extra_headers"] = extra_headers
- response = base_llm_aiohttp_handler.completion(
- model=model,
- messages=messages,
- headers=headers,
- model_response=model_response,
- api_key=api_key,
- api_base=api_base,
- acompletion=acompletion,
- logging_obj=logging,
- optional_params=optional_params,
- litellm_params=litellm_params,
- timeout=timeout,
- client=client,
- custom_llm_provider=custom_llm_provider,
- encoding=_get_encoding(),
- stream=stream,
- )
+ response = _complete_aiohttp_openai(_dispatch_ctx)
elif custom_llm_provider == "cometapi":
- api_key = (
- api_key
- or litellm.cometapi_key
- or get_secret_str("COMETAPI_KEY")
- or litellm.api_key
- )
-
- api_base = (
- api_base
- or litellm.api_base
- or get_secret_str("COMETAPI_API_BASE")
- or "https://api.cometapi.com/v1"
- )
-
- ## COMPLETION CALL
- response = base_llm_http_handler.completion(
- model=model,
- messages=messages,
- headers=headers,
- model_response=model_response,
- api_key=api_key,
- api_base=api_base,
- acompletion=acompletion,
- logging_obj=logging,
- optional_params=optional_params,
- litellm_params=litellm_params,
- shared_session=shared_session,
- timeout=timeout,
- client=client,
- custom_llm_provider=custom_llm_provider,
- encoding=_get_encoding(),
- stream=stream,
- provider_config=provider_config,
- )
-
- ## LOGGING
- logging.post_call(
- input=messages, api_key=api_key, original_response=response
- )
+ response = _complete_cometapi(_dispatch_ctx)
elif custom_llm_provider == "minimax":
- api_key = api_key or get_secret_str("MINIMAX_API_KEY") or litellm.api_key
-
- api_base = (
- api_base
- or litellm.api_base
- or get_secret_str("MINIMAX_API_BASE")
- or "https://api.minimax.io/v1"
- )
-
- response = base_llm_http_handler.completion(
- model=model,
- messages=messages,
- api_base=api_base,
- custom_llm_provider=custom_llm_provider,
- model_response=model_response,
- encoding=_get_encoding(),
- logging_obj=logging,
- optional_params=optional_params,
- timeout=timeout,
- litellm_params=litellm_params,
- shared_session=shared_session,
- acompletion=acompletion,
- stream=stream,
- api_key=api_key,
- headers=headers,
- client=client,
- provider_config=provider_config,
- )
- logging.post_call(
- input=messages, api_key=api_key, original_response=response
- )
+ response = _complete_minimax(_dispatch_ctx)
elif custom_llm_provider == "hosted_vllm":
- api_base = (
- api_base or litellm.api_base or get_secret_str("HOSTED_VLLM_API_BASE")
- )
-
- response = base_llm_http_handler.completion(
- model=model,
- messages=messages,
- api_base=api_base,
- custom_llm_provider=custom_llm_provider,
- model_response=model_response,
- encoding=_get_encoding(),
- logging_obj=logging,
- optional_params=optional_params,
- timeout=timeout,
- litellm_params=litellm_params,
- shared_session=shared_session,
- acompletion=acompletion,
- stream=stream,
- api_key=api_key,
- headers=headers,
- client=client,
- provider_config=provider_config,
- )
- logging.post_call(
- input=messages, api_key=api_key, original_response=response
- )
+ response = _complete_hosted_vllm(_dispatch_ctx)
elif (
model in litellm.open_ai_chat_completion_models
or custom_llm_provider == "custom_openai"
@@ -2676,205 +5724,17 @@ def completion( # type: ignore
): # allow user to make an openai call with a custom base
# note: if a user sets a custom base - we should ensure this works
# allow for the setting of dynamic and stateful api-bases
- api_base = (
- api_base # for deepinfra/perplexity/anyscale/groq/friendliai we check in get_llm_provider and pass in the api base from there
- or litellm.api_base
- or get_secret("OPENAI_BASE_URL")
- or get_secret("OPENAI_API_BASE")
- or "https://api.openai.com/v1"
- )
- organization = (
- organization
- or litellm.organization
- or get_secret("OPENAI_ORGANIZATION")
- or None # default - https://github.com/openai/openai-python/blob/284c1799070c723c6a553337134148a7ab088dd8/openai/util.py#L105
- )
- openai.organization = organization
- # set API KEY
- api_key = (
- api_key
- or litellm.api_key # for deepinfra/perplexity/anyscale/friendliai we check in get_llm_provider and pass in the api key from there
- or litellm.openai_key
- or get_secret("OPENAI_API_KEY")
- )
-
- headers = headers or litellm.headers
-
- # Add GitHub Copilot headers (same as /responses endpoint does)
- if custom_llm_provider == "github_copilot":
- from litellm.llms.github_copilot.authenticator import Authenticator
- from litellm.llms.github_copilot.common_utils import (
- get_copilot_default_headers,
- )
-
- copilot_auth = Authenticator()
- copilot_api_key = copilot_auth.get_api_key()
- copilot_headers = get_copilot_default_headers(copilot_api_key)
- if extra_headers:
- copilot_headers.update(extra_headers)
- extra_headers = copilot_headers
-
- if extra_headers is not None:
- optional_params["extra_headers"] = extra_headers
-
- if (
- litellm.enable_preview_features and metadata is not None
- ): # [PREVIEW] allow metadata to be passed to OPENAI
- openai_metadata = get_requester_metadata(metadata)
- if openai_metadata is not None:
- optional_params["metadata"] = openai_metadata
-
- ## LOAD CONFIG - if set
- config = litellm.OpenAIConfig.get_config()
- for k, v in config.items():
- if (
- k not in optional_params
- ): # completion(top_k=3) > openai_config(top_k=3) <- allows for dynamic variables to be passed in
- optional_params[k] = v
-
- ## COMPLETION CALL
- use_base_llm_http_handler = get_secret_bool(
- "EXPERIMENTAL_OPENAI_BASE_LLM_HTTP_HANDLER"
- )
-
- try:
- if use_base_llm_http_handler:
- response = base_llm_http_handler.completion(
- model=model,
- messages=messages,
- api_base=api_base,
- custom_llm_provider=custom_llm_provider,
- model_response=model_response,
- encoding=_get_encoding(),
- logging_obj=logging,
- optional_params=optional_params,
- timeout=timeout,
- litellm_params=litellm_params,
- shared_session=shared_session,
- acompletion=acompletion,
- stream=stream,
- api_key=api_key,
- headers=headers,
- client=client,
- provider_config=provider_config,
- )
- else:
- response = openai_chat_completions.completion(
- model=model,
- messages=messages,
- headers=headers,
- model_response=model_response,
- print_verbose=print_verbose,
- api_key=api_key,
- api_base=api_base,
- acompletion=acompletion,
- logging_obj=logging,
- optional_params=optional_params,
- litellm_params=litellm_params,
- logger_fn=logger_fn,
- timeout=timeout, # type: ignore
- custom_prompt_dict=custom_prompt_dict,
- client=client, # pass AsyncOpenAI, OpenAI client
- organization=organization,
- custom_llm_provider=custom_llm_provider,
- shared_session=shared_session,
- )
- except Exception as e:
- ## LOGGING - log the original exception returned
- logging.post_call(
- input=messages,
- api_key=api_key,
- original_response=str(e),
- additional_args={"headers": headers},
- )
- raise e
-
- if optional_params.get("stream", False):
- ## LOGGING
- logging.post_call(
- input=messages,
- api_key=api_key,
- original_response=response,
- additional_args={"headers": headers},
- )
+ response = _complete_custom_openai(_dispatch_ctx)
elif custom_llm_provider == "mistral":
- api_key = api_key or litellm.api_key or get_secret("MISTRAL_API_KEY")
- api_base = (
- api_base
- or litellm.api_base
- or get_secret("MISTRAL_API_BASE")
- or "https://api.mistral.ai/v1"
- )
-
- response = base_llm_http_handler.completion(
- model=model,
- messages=messages,
- api_base=api_base,
- custom_llm_provider=custom_llm_provider,
- model_response=model_response,
- encoding=_get_encoding(),
- logging_obj=logging,
- optional_params=optional_params,
- timeout=timeout,
- litellm_params=litellm_params,
- shared_session=shared_session,
- acompletion=acompletion,
- stream=stream,
- api_key=api_key,
- headers=headers,
- client=client,
- provider_config=provider_config,
- )
+ response = _complete_mistral(_dispatch_ctx)
elif (
"replicate" in model
or custom_llm_provider == "replicate"
or model in litellm.replicate_models
):
# Setting the relevant API KEY for replicate, replicate defaults to using os.environ.get("REPLICATE_API_TOKEN")
- replicate_key = (
- api_key
- or litellm.replicate_key
- or litellm.api_key
- or get_secret("REPLICATE_API_KEY")
- or get_secret("REPLICATE_API_TOKEN")
- )
-
- api_base = (
- api_base
- or litellm.api_base
- or get_secret("REPLICATE_API_BASE")
- or "https://api.replicate.com/v1"
- )
-
- custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict
-
- model_response = replicate_chat_completion( # type: ignore
- model=model,
- messages=messages,
- api_base=api_base,
- model_response=model_response,
- print_verbose=print_verbose,
- optional_params=optional_params,
- litellm_params=litellm_params,
- logger_fn=logger_fn,
- encoding=_get_encoding(), # for calculating input/output tokens
- api_key=replicate_key,
- logging_obj=logging,
- custom_prompt_dict=custom_prompt_dict,
- acompletion=acompletion,
- headers=headers,
- )
-
- if optional_params.get("stream", False) is True:
- ## LOGGING
- logging.post_call(
- input=messages,
- api_key=replicate_key,
- original_response=model_response,
- )
-
- response = model_response
+ response = _complete_replicate(_dispatch_ctx)
elif (
"clarifai" in model
or custom_llm_provider == "clarifai"
@@ -2882,614 +5742,36 @@ def completion( # type: ignore
):
pass # Deprecated - handled in the openai compatible provider section above
elif custom_llm_provider == "anthropic_text":
- api_key = (
- api_key
- or litellm.anthropic_key
- or litellm.api_key
- or os.environ.get("ANTHROPIC_API_KEY")
- )
- custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict
- api_base = (
- api_base
- or litellm.api_base
- or get_secret("ANTHROPIC_API_BASE")
- or get_secret("ANTHROPIC_BASE_URL")
- or "https://api.anthropic.com/v1/complete"
- )
-
- # Check if we should disable automatic URL suffix appending
- disable_url_suffix = get_secret_bool("LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX")
- if (
- api_base is not None
- and not disable_url_suffix
- and not api_base.endswith("/v1/complete")
- ):
- api_base += "/v1/complete"
- elif disable_url_suffix:
- verbose_logger.debug(
- "LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX is set, skipping /v1/complete suffix"
- )
-
- response = base_llm_http_handler.completion(
- model=model,
- stream=stream,
- messages=messages,
- acompletion=acompletion,
- api_base=api_base,
- model_response=model_response,
- optional_params=optional_params,
- litellm_params=litellm_params,
- shared_session=shared_session,
- custom_llm_provider="anthropic_text",
- timeout=timeout,
- headers=headers,
- encoding=_get_encoding(),
- api_key=api_key,
- logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
- )
+ response = _complete_anthropic_text(_dispatch_ctx)
elif custom_llm_provider == "anthropic":
- api_key = (
- api_key
- or litellm.anthropic_key
- or litellm.api_key
- or os.environ.get("ANTHROPIC_API_KEY")
- )
- custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict
- # call /messages
- # default route for all anthropic models
- api_base = (
- api_base
- or litellm.api_base
- or get_secret("ANTHROPIC_API_BASE")
- or get_secret("ANTHROPIC_BASE_URL")
- or "https://api.anthropic.com/v1/messages"
- )
-
- # Check if we should disable automatic URL suffix appending
- disable_url_suffix = get_secret_bool("LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX")
- if (
- api_base is not None
- and not disable_url_suffix
- and not api_base.endswith("/v1/messages")
- ):
- api_base += "/v1/messages"
- elif disable_url_suffix:
- verbose_logger.debug(
- "LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX is set, skipping /v1/messages suffix"
- )
-
- response = anthropic_chat_completions.completion(
- model=model,
- messages=messages,
- api_base=api_base,
- acompletion=acompletion,
- custom_prompt_dict=litellm.custom_prompt_dict,
- model_response=model_response,
- print_verbose=print_verbose,
- optional_params=optional_params,
- litellm_params=litellm_params,
- logger_fn=logger_fn,
- encoding=_get_encoding(), # for calculating input/output tokens
- api_key=api_key,
- logging_obj=logging,
- headers=headers,
- timeout=timeout,
- client=client,
- custom_llm_provider=custom_llm_provider,
- )
- if optional_params.get("stream", False) or acompletion is True:
- ## LOGGING
- logging.post_call(
- input=messages,
- api_key=api_key,
- original_response=response,
- )
- response = response
+ response = _complete_anthropic(_dispatch_ctx)
elif custom_llm_provider == "nlp_cloud":
- nlp_cloud_key = (
- api_key
- or litellm.nlp_cloud_key
- or get_secret("NLP_CLOUD_API_KEY")
- or litellm.api_key
- )
-
- api_base = (
- api_base
- or litellm.api_base
- or get_secret("NLP_CLOUD_API_BASE")
- or "https://api.nlpcloud.io/v1/gpu/"
- )
-
- response = nlp_cloud_chat_completion(
- model=model,
- messages=messages,
- api_base=api_base,
- model_response=model_response,
- print_verbose=print_verbose,
- optional_params=optional_params,
- litellm_params=litellm_params,
- logger_fn=logger_fn,
- encoding=_get_encoding(),
- api_key=nlp_cloud_key,
- logging_obj=logging,
- )
-
- if "stream" in optional_params and optional_params["stream"] is True:
- # don't try to access stream object,
- response = CustomStreamWrapper(
- response,
- model,
- custom_llm_provider="nlp_cloud",
- logging_obj=logging,
- )
-
- if optional_params.get("stream", False) or acompletion is True:
- ## LOGGING
- logging.post_call(
- input=messages,
- api_key=api_key,
- original_response=response,
- )
-
- response = response
+ response = _complete_nlp_cloud(_dispatch_ctx)
elif custom_llm_provider == "aleph_alpha":
- aleph_alpha_key = (
- api_key
- or litellm.aleph_alpha_key
- or get_secret("ALEPH_ALPHA_API_KEY")
- or get_secret("ALEPHALPHA_API_KEY")
- or litellm.api_key
- )
-
- api_base = (
- api_base
- or litellm.api_base
- or get_secret("ALEPH_ALPHA_API_BASE")
- or "https://api.aleph-alpha.com/complete"
- )
-
- model_response = aleph_alpha.completion(
- model=model,
- messages=messages,
- api_base=api_base,
- model_response=model_response,
- print_verbose=print_verbose,
- optional_params=optional_params,
- litellm_params=litellm_params,
- logger_fn=logger_fn,
- encoding=_get_encoding(),
- default_max_tokens_to_sample=litellm.max_tokens,
- api_key=aleph_alpha_key,
- logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
- )
-
- if "stream" in optional_params and optional_params["stream"] is True:
- # don't try to access stream object,
- response = CustomStreamWrapper(
- model_response,
- model,
- custom_llm_provider="aleph_alpha",
- logging_obj=logging,
- )
- return response
- response = model_response
+ response = _complete_aleph_alpha(_dispatch_ctx)
elif custom_llm_provider == "cohere_chat" or custom_llm_provider == "cohere":
- cohere_key = (
- api_key
- or litellm.cohere_key
- or get_secret_str("COHERE_API_KEY")
- or get_secret_str("CO_API_KEY")
- or litellm.api_key
- )
-
- cohere_route = CohereModelInfo.get_cohere_route(model)
- verbose_logger.debug(f"Cohere route: {cohere_route}")
- # Set API base based on route
- if cohere_route == "v2":
- api_base = (
- api_base
- or litellm.api_base
- or get_secret_str("COHERE_API_BASE")
- or "https://api.cohere.com/v2/chat"
- )
- # Remove v2/ prefix from model name for the actual API call
- if "v2/" in model:
- model = model.replace("v2/", "")
- else:
- api_base = (
- api_base
- or litellm.api_base
- or get_secret_str("COHERE_API_BASE")
- or "https://api.cohere.ai/v1/chat"
- )
-
- headers = headers or litellm.headers or {}
- if headers is None:
- headers = {}
-
- if extra_headers is not None:
- headers.update(extra_headers)
-
- verbose_logger.debug(f"Model: {model}, API Base: {api_base}")
- verbose_logger.debug(f"Provider Config: {provider_config}")
- response = base_llm_http_handler.completion(
- model=model,
- stream=stream,
- messages=messages,
- acompletion=acompletion,
- api_base=api_base,
- model_response=model_response,
- optional_params=optional_params,
- litellm_params=litellm_params,
- shared_session=shared_session,
- custom_llm_provider="cohere_chat",
- timeout=timeout,
- headers=headers,
- encoding=_get_encoding(),
- api_key=cohere_key,
- provider_config=provider_config,
- logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
- )
+ response = _complete_cohere_chat(_dispatch_ctx)
elif custom_llm_provider == "maritalk":
- maritalk_key = (
- api_key
- or litellm.maritalk_key
- or get_secret("MARITALK_API_KEY")
- or litellm.api_key
- )
-
- api_base = (
- api_base
- or litellm.api_base
- or get_secret("MARITALK_API_BASE")
- or "https://chat.maritaca.ai/api"
- )
-
- model_response = openai_like_chat_completion.completion(
- model=model,
- messages=messages,
- api_base=api_base,
- model_response=model_response,
- print_verbose=print_verbose,
- optional_params=optional_params,
- litellm_params=litellm_params,
- logger_fn=logger_fn,
- encoding=_get_encoding(),
- api_key=maritalk_key,
- logging_obj=logging,
- custom_llm_provider="maritalk",
- custom_prompt_dict=custom_prompt_dict,
- )
-
- response = model_response
+ response = _complete_maritalk(_dispatch_ctx)
elif custom_llm_provider == "amazon_nova":
- api_key = (
- api_key
- or litellm.amazon_nova_api_key
- or get_secret_str("AMAZON_NOVA_API_KEY")
- or litellm.api_key
- )
- api_base = (
- api_base
- or litellm.api_base
- or get_secret_str("AMAZON_NOVA_API_BASE")
- or "https://api.nova.amazon.com/v1"
- )
- response = openai_like_chat_completion.completion(
- model=model,
- messages=messages,
- api_base=api_base,
- model_response=model_response,
- print_verbose=print_verbose,
- optional_params=optional_params,
- litellm_params=litellm_params,
- logger_fn=logger_fn,
- encoding=_get_encoding(),
- api_key=api_key,
- logging_obj=logging,
- timeout=timeout,
- custom_llm_provider=custom_llm_provider,
- custom_prompt_dict=custom_prompt_dict,
- )
+ response = _complete_amazon_nova(_dispatch_ctx)
elif custom_llm_provider == "huggingface":
- huggingface_key = (
- api_key
- or litellm.huggingface_key
- or os.environ.get("HF_TOKEN")
- or os.environ.get("HUGGINGFACE_API_KEY")
- or litellm.api_key
- )
- hf_headers = headers or litellm.headers
- response = base_llm_http_handler.completion(
- model=model,
- messages=messages,
- headers=hf_headers,
- model_response=model_response,
- api_key=huggingface_key,
- api_base=api_base,
- acompletion=acompletion,
- logging_obj=logging,
- optional_params=optional_params,
- litellm_params=litellm_params,
- timeout=timeout, # type: ignore
- client=client,
- custom_llm_provider=custom_llm_provider,
- encoding=_get_encoding(),
- stream=stream,
- )
+ response = _complete_huggingface(_dispatch_ctx)
elif custom_llm_provider == "oci":
- response = base_llm_http_handler.completion(
- model=model,
- messages=messages,
- headers=headers,
- model_response=model_response,
- api_key=api_key,
- api_base=api_base,
- acompletion=acompletion,
- logging_obj=logging,
- optional_params=optional_params,
- litellm_params=litellm_params,
- timeout=timeout, # type: ignore
- client=client,
- custom_llm_provider=custom_llm_provider,
- encoding=_get_encoding(),
- stream=stream,
- )
+ response = _complete_oci(_dispatch_ctx)
elif custom_llm_provider == "compactifai":
- api_key = (
- api_key or get_secret_str("COMPACTIFAI_API_KEY") or litellm.api_key
- )
-
- api_base = api_base or "https://api.compactif.ai/v1"
-
- ## COMPLETION CALL
- response = base_llm_http_handler.completion(
- model=model,
- messages=messages,
- headers=headers,
- model_response=model_response,
- api_key=api_key,
- api_base=api_base,
- acompletion=acompletion,
- logging_obj=logging,
- optional_params=optional_params,
- litellm_params=litellm_params,
- timeout=timeout,
- client=client,
- custom_llm_provider=custom_llm_provider,
- encoding=_get_encoding(),
- stream=stream,
- provider_config=provider_config,
- )
+ response = _complete_compactifai(_dispatch_ctx)
elif custom_llm_provider == "oobabooga":
- custom_llm_provider = "oobabooga"
- model_response = oobabooga.completion(
- model=model,
- messages=messages,
- model_response=model_response,
- api_base=api_base, # type: ignore
- print_verbose=print_verbose,
- optional_params=optional_params,
- litellm_params=litellm_params,
- api_key=None,
- logger_fn=logger_fn,
- encoding=_get_encoding(),
- logging_obj=logging,
- )
- if "stream" in optional_params and optional_params["stream"] is True:
- # don't try to access stream object,
- response = CustomStreamWrapper(
- model_response,
- model,
- custom_llm_provider="oobabooga",
- logging_obj=logging,
- )
- return response
- response = model_response
+ response = _complete_oobabooga(_dispatch_ctx)
elif custom_llm_provider == "databricks":
- api_base = (
- api_base # for databricks we check in get_llm_provider and pass in the api base from there
- or litellm.api_base
- or os.getenv("DATABRICKS_API_BASE")
- )
-
- # set API KEY
- api_key = (
- api_key
- or litellm.api_key # for databricks we check in get_llm_provider and pass in the api key from there
- or litellm.databricks_key
- or get_secret("DATABRICKS_API_KEY")
- )
-
- headers = headers or litellm.headers
-
- ## COMPLETION CALL
- try:
- response = base_llm_http_handler.completion(
- model=model,
- stream=stream,
- messages=messages,
- acompletion=acompletion,
- api_base=api_base,
- model_response=model_response,
- optional_params=optional_params,
- litellm_params=litellm_params,
- custom_llm_provider="databricks",
- timeout=timeout,
- headers=headers,
- encoding=_get_encoding(),
- api_key=api_key,
- logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
- client=client,
- )
- except Exception as e:
- ## LOGGING - log the original exception returned
- logging.post_call(
- input=messages,
- api_key=api_key,
- original_response=str(e),
- additional_args={"headers": headers},
- )
- raise e
-
- if optional_params.get("stream", False):
- ## LOGGING
- logging.post_call(
- input=messages,
- api_key=api_key,
- original_response=response,
- additional_args={"headers": headers},
- )
+ response = _complete_databricks(_dispatch_ctx)
elif custom_llm_provider == "datarobot":
- response = base_llm_http_handler.completion(
- model=model,
- messages=messages,
- headers=headers,
- model_response=model_response,
- api_key=api_key,
- api_base=api_base,
- acompletion=acompletion,
- logging_obj=logging,
- optional_params=optional_params,
- litellm_params=litellm_params,
- timeout=timeout, # type: ignore
- client=client,
- custom_llm_provider=custom_llm_provider,
- encoding=_get_encoding(),
- stream=stream,
- provider_config=provider_config,
- )
+ response = _complete_datarobot(_dispatch_ctx)
elif custom_llm_provider == "openrouter":
- api_base = (
- api_base
- or litellm.api_base
- or get_secret_str("OPENROUTER_API_BASE")
- or "https://openrouter.ai/api/v1"
- )
-
- api_key = (
- api_key
- or litellm.api_key
- or litellm.openrouter_key
- or get_secret_str("OPENROUTER_API_KEY")
- or get_secret_str("OR_API_KEY")
- )
-
- openrouter_site_url = get_secret("OR_SITE_URL") or "https://litellm.ai"
- openrouter_app_name = get_secret("OR_APP_NAME") or "liteLLM"
-
- openrouter_headers = {
- "HTTP-Referer": openrouter_site_url,
- "X-Title": openrouter_app_name,
- }
-
- _headers = headers or litellm.headers
- if _headers:
- openrouter_headers.update(_headers)
-
- headers = openrouter_headers
-
- ## Load Config
- config = litellm.OpenrouterConfig.get_config()
- for k, v in config.items():
- if k == "extra_body":
- # we use openai 'extra_body' to pass openrouter specific params - transforms, route, models
- if "extra_body" in optional_params:
- optional_params[k].update(v)
- else:
- optional_params[k] = v
- elif k not in optional_params:
- optional_params[k] = v
-
- data = {"model": model, "messages": messages, **optional_params}
-
- ## COMPLETION CALL
- response = base_llm_http_handler.completion(
- model=model,
- stream=stream,
- messages=messages,
- acompletion=acompletion,
- api_base=api_base,
- model_response=model_response,
- optional_params=optional_params,
- litellm_params=litellm_params,
- shared_session=shared_session,
- custom_llm_provider="openrouter",
- timeout=timeout,
- headers=headers,
- encoding=_get_encoding(),
- api_key=api_key,
- logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
- client=client,
- )
- ## LOGGING
- logging.post_call(
- input=messages, api_key=openai.api_key, original_response=response
- )
+ response = _complete_openrouter(_dispatch_ctx)
elif custom_llm_provider == "vercel_ai_gateway":
- api_base = (
- api_base
- or litellm.api_base
- or get_secret_str("VERCEL_AI_GATEWAY_API_BASE")
- or "https://ai-gateway.vercel.sh/v1"
- )
-
- api_key = (
- api_key or litellm.api_key or get_secret("VERCEL_AI_GATEWAY_API_KEY")
- )
-
- vercel_site_url = get_secret("VERCEL_SITE_URL") or "https://litellm.ai"
- vercel_app_name = get_secret("VERCEL_APP_NAME") or "liteLLM"
-
- vercel_headers = {
- "http-referer": vercel_site_url,
- "x-title": vercel_app_name,
- }
-
- _headers = headers or litellm.headers
- if _headers:
- vercel_headers.update(_headers)
-
- headers = vercel_headers
-
- ## Load Config
- config = litellm.VercelAIGatewayConfig.get_config()
- for k, v in config.items():
- if k == "extra_body":
- # we use openai 'extra_body' to pass vercel specific params - providerOptions
- if "extra_body" in optional_params:
- optional_params[k].update(v)
- else:
- optional_params[k] = v
- elif k not in optional_params:
- optional_params[k] = v
-
- data = {"model": model, "messages": messages, **optional_params}
-
- ## COMPLETION CALL
- response = base_llm_http_handler.completion(
- model=model,
- stream=stream,
- messages=messages,
- acompletion=acompletion,
- api_base=api_base,
- model_response=model_response,
- optional_params=optional_params,
- litellm_params=litellm_params,
- shared_session=shared_session,
- custom_llm_provider="vercel_ai_gateway",
- timeout=timeout,
- headers=headers,
- encoding=_get_encoding(),
- api_key=api_key,
- logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
- client=client,
- )
- ## LOGGING
- logging.post_call(
- input=messages, api_key=openai.api_key, original_response=response
- )
+ response = _complete_vercel_ai_gateway(_dispatch_ctx)
elif (
custom_llm_provider == "together_ai"
or ("togethercomputer" in model)
@@ -3504,1114 +5786,75 @@ def completion( # type: ignore
"Palm was decommisioned on October 2024. Please use the `gemini/` route for Gemini Google AI Studio Models. Announcement: https://ai.google.dev/palm_docs/palm?hl=en"
)
elif custom_llm_provider == "vertex_ai_beta" or custom_llm_provider == "gemini":
- vertex_ai_project = (
- optional_params.pop("vertex_project", None)
- or optional_params.pop("vertex_ai_project", None)
- or litellm.vertex_project
- or get_secret("VERTEXAI_PROJECT")
- )
- vertex_ai_location = (
- optional_params.pop("vertex_location", None)
- or optional_params.pop("vertex_ai_location", None)
- or litellm.vertex_location
- or get_secret("VERTEXAI_LOCATION")
- )
- vertex_credentials = (
- optional_params.pop("vertex_credentials", None)
- or optional_params.pop("vertex_ai_credentials", None)
- or get_secret("VERTEXAI_CREDENTIALS")
- )
-
- gemini_api_key = (
- api_key
- or get_api_key_from_env()
- or get_secret("PALM_API_KEY") # older palm api key should also work
- or litellm.api_key
- )
-
- api_base = api_base or litellm.api_base or get_secret("GEMINI_API_BASE")
- new_params = safe_deep_copy(optional_params or {})
- response = vertex_chat_completion.completion( # type: ignore
- model=model,
- messages=messages,
- model_response=model_response,
- print_verbose=print_verbose,
- optional_params=new_params,
- litellm_params=litellm_params, # type: ignore
- logger_fn=logger_fn,
- encoding=_get_encoding(),
- vertex_location=vertex_ai_location,
- vertex_project=vertex_ai_project,
- vertex_credentials=vertex_credentials,
- gemini_api_key=gemini_api_key,
- logging_obj=logging,
- acompletion=acompletion,
- timeout=timeout,
- custom_llm_provider=custom_llm_provider, # type: ignore
- client=client,
- api_base=api_base,
- extra_headers=headers,
- )
+ response = _complete_vertex_ai_beta(_dispatch_ctx)
elif custom_llm_provider == "vertex_ai":
- vertex_ai_project = (
- optional_params.pop("vertex_project", None)
- or optional_params.pop("vertex_ai_project", None)
- or litellm.vertex_project
- or get_secret("VERTEXAI_PROJECT")
- )
- vertex_ai_location = (
- optional_params.pop("vertex_location", None)
- or optional_params.pop("vertex_ai_location", None)
- or litellm.vertex_location
- or get_secret("VERTEXAI_LOCATION")
- )
- vertex_credentials = (
- optional_params.pop("vertex_credentials", None)
- or optional_params.pop("vertex_ai_credentials", None)
- or get_secret("VERTEXAI_CREDENTIALS")
- )
-
- api_base = api_base or litellm.api_base or get_secret("VERTEXAI_API_BASE")
-
- new_params = safe_deep_copy(optional_params or {})
- model_route = get_vertex_ai_model_route(
- model=model, litellm_params=litellm_params
- )
-
- if model_route == VertexAIModelRoute.PARTNER_MODELS:
- model_response = vertex_partner_models_chat_completion.completion(
- model=model,
- messages=messages,
- model_response=model_response,
- print_verbose=print_verbose,
- optional_params=new_params,
- litellm_params=litellm_params, # type: ignore
- logger_fn=logger_fn,
- encoding=_get_encoding(),
- api_base=api_base,
- vertex_location=vertex_ai_location,
- vertex_project=vertex_ai_project,
- vertex_credentials=vertex_credentials,
- logging_obj=logging,
- acompletion=acompletion,
- headers=headers,
- custom_prompt_dict=custom_prompt_dict,
- timeout=timeout,
- client=client,
- )
- elif model_route == VertexAIModelRoute.GEMINI:
- model_response = vertex_chat_completion.completion( # type: ignore
- model=model,
- messages=messages,
- model_response=model_response,
- print_verbose=print_verbose,
- optional_params=new_params,
- litellm_params=litellm_params, # type: ignore
- logger_fn=logger_fn,
- encoding=_get_encoding(),
- vertex_location=vertex_ai_location,
- vertex_project=vertex_ai_project,
- vertex_credentials=vertex_credentials,
- gemini_api_key=None,
- logging_obj=logging,
- acompletion=acompletion,
- timeout=timeout,
- custom_llm_provider=custom_llm_provider, # type: ignore
- client=client,
- api_base=api_base,
- extra_headers=headers,
- )
- elif model_route == VertexAIModelRoute.GEMMA:
- # Vertex Gemma Models with custom prediction endpoint
- model_response = vertex_gemma_chat_completion.completion(
- model=model,
- messages=messages,
- model_response=model_response,
- print_verbose=print_verbose,
- optional_params=new_params,
- litellm_params=litellm_params, # type: ignore
- logger_fn=logger_fn,
- encoding=_get_encoding(),
- api_base=api_base,
- vertex_location=vertex_ai_location,
- vertex_project=vertex_ai_project,
- vertex_credentials=vertex_credentials,
- logging_obj=logging,
- acompletion=acompletion,
- headers=headers,
- custom_prompt_dict=custom_prompt_dict,
- timeout=timeout,
- client=client,
- )
- elif model_route == VertexAIModelRoute.MODEL_GARDEN:
- # Vertex Model Garden - OpenAI compatible models
- model_response = vertex_model_garden_chat_completion.completion(
- model=model,
- messages=messages,
- model_response=model_response,
- print_verbose=print_verbose,
- optional_params=new_params,
- litellm_params=litellm_params, # type: ignore
- logger_fn=logger_fn,
- encoding=_get_encoding(),
- api_base=api_base,
- vertex_location=vertex_ai_location,
- vertex_project=vertex_ai_project,
- vertex_credentials=vertex_credentials,
- logging_obj=logging,
- acompletion=acompletion,
- headers=headers,
- custom_prompt_dict=custom_prompt_dict,
- timeout=timeout,
- client=client,
- )
- elif model_route == VertexAIModelRoute.AGENT_ENGINE:
- # Vertex AI Agent Engine (Reasoning Engines)
- from litellm.llms.vertex_ai.agent_engine.transformation import (
- VertexAgentEngineConfig,
- )
-
- vertex_agent_engine_config = VertexAgentEngineConfig()
-
- # Update litellm_params with vertex credentials
- litellm_params["vertex_project"] = vertex_ai_project
- litellm_params["vertex_location"] = vertex_ai_location
- litellm_params["vertex_credentials"] = vertex_credentials
-
- model_response = base_llm_http_handler.completion(
- model=model,
- stream=stream,
- messages=messages,
- model_response=model_response,
- optional_params=new_params,
- litellm_params=litellm_params, # type: ignore
- encoding=_get_encoding(),
- api_key=None,
- api_base=api_base,
- logging_obj=logging,
- acompletion=acompletion,
- timeout=timeout,
- client=client,
- custom_llm_provider="vertex_ai",
- provider_config=vertex_agent_engine_config,
- headers=headers or {},
- )
- else: # VertexAIModelRoute.NON_GEMINI
- model_response = vertex_ai_non_gemini.completion(
- model=model,
- messages=messages,
- model_response=model_response,
- print_verbose=print_verbose,
- optional_params=new_params,
- litellm_params=litellm_params,
- logger_fn=logger_fn,
- encoding=_get_encoding(),
- vertex_location=vertex_ai_location,
- vertex_project=vertex_ai_project,
- vertex_credentials=vertex_credentials,
- logging_obj=logging,
- acompletion=acompletion,
- )
-
- if (
- "stream" in optional_params
- and optional_params["stream"] is True
- and acompletion is False
- ):
- response = CustomStreamWrapper(
- model_response,
- model,
- custom_llm_provider="vertex_ai",
- logging_obj=logging,
- )
- return response
- response = model_response
+ response = _complete_vertex_ai(_dispatch_ctx)
elif custom_llm_provider == "predibase":
- tenant_id = (
- optional_params.pop("tenant_id", None)
- or optional_params.pop("predibase_tenant_id", None)
- or litellm.predibase_tenant_id
- or get_secret("PREDIBASE_TENANT_ID")
- )
-
- if tenant_id is None:
- raise ValueError(
- "Missing Predibase Tenant ID - Required for making the request. Set dynamically (e.g. `completion(..tenant_id=)`) or in env - `PREDIBASE_TENANT_ID`."
- )
-
- api_base = (
- api_base
- or optional_params.pop("api_base", None)
- or optional_params.pop("base_url", None)
- or litellm.api_base
- or get_secret("PREDIBASE_API_BASE")
- )
-
- api_key = (
- api_key
- or litellm.api_key
- or litellm.predibase_key
- or get_secret("PREDIBASE_API_KEY")
- )
-
- _model_response = predibase_chat_completions.completion(
- model=model,
- messages=messages,
- model_response=model_response,
- print_verbose=print_verbose,
- optional_params=optional_params,
- litellm_params=litellm_params,
- logger_fn=logger_fn,
- encoding=_get_encoding(),
- logging_obj=logging,
- acompletion=acompletion,
- api_base=api_base,
- custom_prompt_dict=custom_prompt_dict,
- api_key=api_key,
- tenant_id=tenant_id,
- timeout=timeout,
- )
-
- if (
- "stream" in optional_params
- and optional_params["stream"] is True
- and acompletion is False
- ):
- return _model_response
- response = _model_response
+ response = _complete_predibase(_dispatch_ctx)
elif custom_llm_provider == "text-completion-codestral":
- api_base = (
- api_base
- or optional_params.pop("api_base", None)
- or optional_params.pop("base_url", None)
- or litellm.api_base
- or "https://codestral.mistral.ai/v1/fim/completions"
- )
-
- api_key = api_key or litellm.api_key or get_secret("CODESTRAL_API_KEY")
-
- text_completion_model_response = litellm.TextCompletionResponse(
- stream=stream
- )
-
- _model_response = codestral_text_completions.completion( # type: ignore
- model=model,
- messages=messages,
- model_response=text_completion_model_response,
- print_verbose=print_verbose,
- optional_params=optional_params,
- litellm_params=litellm_params,
- logger_fn=logger_fn,
- encoding=_get_encoding(),
- logging_obj=logging,
- acompletion=acompletion,
- api_base=api_base,
- custom_prompt_dict=custom_prompt_dict,
- api_key=api_key,
- timeout=timeout,
- )
-
- if (
- "stream" in optional_params
- and optional_params["stream"] is True
- and acompletion is False
- ):
- return _model_response
- response = _model_response
+ response = _complete_text_completion_codestral(_dispatch_ctx)
elif custom_llm_provider == "text-completion-inception":
- passed_api_base = (
- api_base
- or optional_params.pop("api_base", None)
- or optional_params.pop("base_url", None)
- )
- api_base = (
- passed_api_base
- or get_secret_str("INCEPTION_API_BASE")
- or "https://api.inceptionlabs.ai/v1"
- )
- # FIM is served at `/v1/fim/completions`; the OpenAI client appends
- # `/completions`, so point it at the `/v1/fim` base.
- api_base = api_base.rstrip("/")
- if not api_base.endswith("/fim"):
- api_base += "/fim"
-
- # Don't forward the server-managed Inception key to a caller-supplied
- # api_base; only resolve it for the default/server base, or when the
- # caller passes their own key.
- if passed_api_base is None or api_key:
- api_key = (
- api_key
- or litellm.inception_key
- or get_secret_str("INCEPTION_API_KEY")
- )
-
- _response = openai_text_completions.completion(
- model=model,
- messages=messages,
- model_response=model_response,
- print_verbose=print_verbose,
- api_key=api_key, # type: ignore[arg-type]
- custom_llm_provider="text-completion-inception",
- api_base=api_base,
- acompletion=acompletion,
- client=client,
- logging_obj=logging,
- optional_params=optional_params,
- litellm_params=litellm_params,
- logger_fn=logger_fn,
- timeout=timeout, # type: ignore
- )
-
- if (
- optional_params.get("stream", False) is False
- and acompletion is False
- and text_completion is False
- ):
- _response = litellm.OpenAITextCompletionConfig().convert_to_chat_model_response_object(
- response_object=_response, model_response_object=model_response
- )
-
- if optional_params.get("stream", False) or acompletion is True:
- logging.post_call(
- input=messages,
- api_key=api_key,
- original_response=_response,
- additional_args={"headers": headers},
- )
- response = _response
+ response = _complete_text_completion_inception(_dispatch_ctx)
elif custom_llm_provider in ("sagemaker_chat", "sagemaker_nova"):
# boto3 reads keys from .env
# sagemaker_chat: HF Messages API endpoints
# sagemaker_nova: Nova models on SageMaker (OpenAI-compatible)
- model_response = base_llm_http_handler.completion(
- model=model,
- stream=stream,
- messages=messages,
- acompletion=acompletion,
- api_base=api_base,
- model_response=model_response,
- optional_params=optional_params,
- litellm_params=litellm_params,
- custom_llm_provider=custom_llm_provider,
- timeout=timeout,
- headers=headers,
- encoding=_get_encoding(),
- api_key=api_key,
- logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
- client=client,
- )
-
- ## RESPONSE OBJECT
- response = model_response
+ response = _complete_sagemaker_chat(_dispatch_ctx)
elif custom_llm_provider == "sagemaker":
# boto3 reads keys from .env
- model_response = sagemaker_llm.completion(
- model=model,
- messages=messages,
- model_response=model_response,
- print_verbose=print_verbose,
- optional_params=optional_params,
- litellm_params=litellm_params,
- custom_prompt_dict=custom_prompt_dict,
- hf_model_name=hf_model_name,
- logger_fn=logger_fn,
- encoding=_get_encoding(),
- logging_obj=logging,
- acompletion=acompletion,
- )
-
- ## RESPONSE OBJECT
- response = model_response
+ response = _complete_sagemaker(_dispatch_ctx)
elif custom_llm_provider == "bedrock":
# boto3 reads keys from .env
- custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict
-
- if "aws_bedrock_client" in optional_params:
- verbose_logger.warning(
- "'aws_bedrock_client' is a deprecated param. Please move to another auth method - https://docs.litellm.ai/docs/providers/bedrock#boto3---authentication."
- )
- # Extract credentials for legacy boto3 client and pass thru to httpx
- aws_bedrock_client = optional_params.pop("aws_bedrock_client")
- creds = aws_bedrock_client._get_credentials().get_frozen_credentials()
-
- if creds.access_key:
- optional_params["aws_access_key_id"] = creds.access_key
- if creds.secret_key:
- optional_params["aws_secret_access_key"] = creds.secret_key
- if creds.token:
- optional_params["aws_session_token"] = creds.token
- if (
- "aws_region_name" not in optional_params
- or optional_params["aws_region_name"] is None
- ):
- optional_params["aws_region_name"] = (
- aws_bedrock_client.meta.region_name
- )
-
- bedrock_route = BedrockModelInfo.get_bedrock_route(model)
- if bedrock_route == "claude_platform":
- provider_config = ProviderConfigManager.get_provider_chat_config(
- model=model,
- provider=LlmProviders.BEDROCK,
- )
- model = BedrockModelInfo.get_claude_platform_model(model)
- response = base_llm_http_handler.completion(
- model=model,
- stream=stream,
- messages=messages,
- acompletion=acompletion,
- api_base=api_base,
- model_response=model_response,
- optional_params=optional_params,
- litellm_params=litellm_params,
- shared_session=shared_session,
- custom_llm_provider="bedrock",
- timeout=timeout,
- headers=headers,
- encoding=_get_encoding(),
- api_key=api_key,
- logging_obj=logging,
- client=client,
- provider_config=provider_config,
- )
- return response
- elif bedrock_route == "converse":
- model = model.replace("converse/", "")
- response = bedrock_converse_chat_completion.completion(
- model=model,
- messages=messages,
- custom_prompt_dict=custom_prompt_dict,
- model_response=model_response,
- optional_params=optional_params,
- litellm_params=litellm_params, # type: ignore
- logger_fn=logger_fn,
- encoding=_get_encoding(),
- logging_obj=logging,
- extra_headers=headers, # Use merged headers instead of original extra_headers
- timeout=timeout,
- acompletion=acompletion,
- client=client,
- api_base=api_base,
- api_key=api_key,
- )
- elif bedrock_route == "converse_like":
- model = model.replace("converse_like/", "")
- response = base_llm_http_handler.completion(
- model=model,
- stream=stream,
- messages=messages,
- acompletion=acompletion,
- api_base=api_base,
- model_response=model_response,
- optional_params=optional_params,
- litellm_params=litellm_params,
- custom_llm_provider="bedrock",
- timeout=timeout,
- headers=headers,
- encoding=_get_encoding(),
- api_key=api_key,
- logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
- client=client,
- )
- else:
- response = base_llm_http_handler.completion(
- model=model,
- stream=stream,
- messages=messages,
- acompletion=acompletion,
- api_base=api_base,
- model_response=model_response,
- optional_params=optional_params,
- litellm_params=litellm_params,
- custom_llm_provider="bedrock",
- timeout=timeout,
- headers=headers,
- encoding=_get_encoding(),
- api_key=api_key,
- logging_obj=logging,
- client=client,
- )
+ response = _complete_bedrock(_dispatch_ctx)
elif custom_llm_provider == "watsonx":
- response = watsonx_chat_completion.completion(
- model=model,
- messages=messages,
- headers=headers,
- model_response=model_response,
- print_verbose=print_verbose,
- api_key=api_key,
- api_base=api_base,
- acompletion=acompletion,
- logging_obj=logging,
- optional_params=optional_params,
- litellm_params=litellm_params,
- logger_fn=logger_fn,
- timeout=timeout, # type: ignore
- custom_prompt_dict=custom_prompt_dict,
- client=client, # pass AsyncOpenAI, OpenAI client
- encoding=_get_encoding(),
- custom_llm_provider="watsonx",
- )
+ response = _complete_watsonx(_dispatch_ctx)
elif custom_llm_provider == "watsonx_text":
- api_key = (
- api_key
- or optional_params.pop("apikey", None)
- or get_secret_str("WATSONX_APIKEY")
- or get_secret_str("WATSONX_API_KEY")
- or get_secret_str("WX_API_KEY")
- )
-
- api_base = (
- api_base
- or optional_params.pop(
- "url",
- optional_params.pop(
- "api_base", optional_params.pop("base_url", None)
- ),
- )
- or get_secret_str("WATSONX_API_BASE")
- or get_secret_str("WATSONX_URL")
- or get_secret_str("WX_URL")
- or get_secret_str("WML_URL")
- )
-
- wx_credentials = optional_params.pop(
- "wx_credentials",
- optional_params.pop(
- "watsonx_credentials", None
- ), # follow {provider}_credentials, same as vertex ai
- )
-
- token: Optional[str] = None
- if wx_credentials is not None:
- api_base = wx_credentials.get("url", api_base)
- api_key = wx_credentials.get(
- "apikey", wx_credentials.get("api_key", api_key)
- )
- token = wx_credentials.get(
- "token",
- wx_credentials.get(
- "watsonx_token", None
- ), # follow format of {provider}_token, same as azure - e.g. 'azure_ad_token=..'
- )
-
- if token is not None:
- optional_params["token"] = token
-
- response = base_llm_http_handler.completion(
- model=model,
- stream=stream,
- messages=messages,
- acompletion=acompletion,
- api_base=api_base,
- model_response=model_response,
- optional_params=optional_params,
- litellm_params=litellm_params,
- shared_session=shared_session,
- custom_llm_provider="watsonx_text",
- timeout=timeout,
- headers=headers,
- encoding=_get_encoding(),
- api_key=api_key,
- logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
- client=client,
- )
+ response = _complete_watsonx_text(_dispatch_ctx)
elif custom_llm_provider == "vllm":
- custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict
- model_response = vllm_handler.completion(
- model=model,
- messages=messages,
- custom_prompt_dict=custom_prompt_dict,
- model_response=model_response,
- print_verbose=print_verbose,
- optional_params=optional_params,
- litellm_params=litellm_params,
- logger_fn=logger_fn,
- encoding=_get_encoding(),
- logging_obj=logging,
- )
-
- if (
- "stream" in optional_params and optional_params["stream"] is True
- ): ## [BETA]
- # don't try to access stream object,
- response = CustomStreamWrapper(
- model_response,
- model,
- custom_llm_provider="vllm",
- logging_obj=logging,
- )
- return response
-
- ## RESPONSE OBJECT
- response = model_response
+ response = _complete_vllm(_dispatch_ctx)
elif custom_llm_provider == "ollama":
- api_base = (
- litellm.api_base
- or api_base
- or get_secret("OLLAMA_API_BASE")
- or "http://localhost:11434"
- )
- if api_key is not None and "Authorization" not in headers:
- headers["Authorization"] = f"Bearer {api_key}"
-
- response = base_llm_http_handler.completion(
- model=model,
- stream=stream,
- messages=messages,
- acompletion=acompletion,
- api_base=api_base,
- model_response=model_response,
- optional_params=optional_params,
- litellm_params=litellm_params,
- shared_session=shared_session,
- custom_llm_provider="ollama",
- timeout=timeout,
- headers=headers,
- encoding=_get_encoding(),
- api_key=api_key,
- logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
- client=client,
- )
+ response = _complete_ollama(_dispatch_ctx)
elif custom_llm_provider == "ollama_chat":
- api_base = (
- litellm.api_base
- or api_base
- or get_secret("OLLAMA_API_BASE")
- or "http://localhost:11434"
- )
-
- api_key = (
- api_key
- or litellm.ollama_key
- or os.environ.get("OLLAMA_API_KEY")
- or litellm.api_key
- )
- if api_key is not None and "Authorization" not in headers:
- headers["Authorization"] = f"Bearer {api_key}"
-
- response = base_llm_http_handler.completion(
- model=model,
- stream=stream,
- messages=messages,
- acompletion=acompletion,
- api_base=api_base,
- model_response=model_response,
- optional_params=optional_params,
- litellm_params=litellm_params,
- shared_session=shared_session,
- custom_llm_provider="ollama_chat",
- timeout=timeout,
- headers=headers,
- encoding=_get_encoding(),
- api_key=api_key,
- logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
- client=client,
- )
+ response = _complete_ollama_chat(_dispatch_ctx)
elif custom_llm_provider == "triton":
- api_base = litellm.api_base or api_base
- response = base_llm_http_handler.completion(
- model=model,
- stream=stream,
- messages=messages,
- acompletion=acompletion,
- api_base=api_base,
- model_response=model_response,
- optional_params=optional_params,
- litellm_params=litellm_params,
- shared_session=shared_session,
- custom_llm_provider=custom_llm_provider,
- timeout=timeout,
- headers=headers,
- encoding=_get_encoding(),
- api_key=api_key,
- logging_obj=logging,
- )
+ response = _complete_triton(_dispatch_ctx)
elif custom_llm_provider == "cloudflare":
- api_key = (
- api_key
- or litellm.cloudflare_api_key
- or litellm.api_key
- or get_secret("CLOUDFLARE_API_KEY")
- )
- account_id = get_secret("CLOUDFLARE_ACCOUNT_ID")
- api_base = (
- api_base
- or litellm.api_base
- or get_secret("CLOUDFLARE_API_BASE")
- or f"https://api.cloudflare.com/client/v4/accounts/{account_id}/ai/run/"
- )
-
- custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict
- response = base_llm_http_handler.completion(
- model=model,
- stream=stream,
- messages=messages,
- acompletion=acompletion,
- api_base=api_base,
- model_response=model_response,
- optional_params=optional_params,
- litellm_params=litellm_params,
- shared_session=shared_session,
- custom_llm_provider="cloudflare",
- timeout=timeout,
- headers=headers,
- encoding=_get_encoding(),
- api_key=api_key,
- logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
- )
+ response = _complete_cloudflare(_dispatch_ctx)
elif custom_llm_provider == "petals" or model in litellm.petals_models:
- api_base = api_base or litellm.api_base
-
- custom_llm_provider = "petals"
- stream = optional_params.pop("stream", False)
- model_response = petals_handler.completion(
- model=model,
- messages=messages,
- api_base=api_base,
- model_response=model_response,
- print_verbose=print_verbose,
- optional_params=optional_params,
- litellm_params=litellm_params,
- logger_fn=logger_fn,
- encoding=_get_encoding(),
- logging_obj=logging,
- client=client,
- )
- if stream is True: ## [BETA]
- # Fake streaming for petals
- resp_string = model_response["choices"][0]["message"]["content"]
- response = CustomStreamWrapper(
- resp_string,
- model,
- custom_llm_provider="petals",
- logging_obj=logging,
- )
- return response
- response = model_response
+ response = _complete_petals(_dispatch_ctx)
elif custom_llm_provider == "snowflake" or model in litellm.snowflake_models:
- try:
- client = (
- HTTPHandler(timeout=timeout) if stream is False else None
- ) # Keep this here, otherwise, the httpx.client closes and streaming is impossible
- response = base_llm_http_handler.completion(
- model=model,
- messages=messages,
- headers=headers,
- model_response=model_response,
- api_key=api_key,
- api_base=api_base,
- acompletion=acompletion,
- logging_obj=logging,
- optional_params=optional_params,
- litellm_params=litellm_params,
- shared_session=shared_session,
- timeout=timeout, # type: ignore
- client=client,
- custom_llm_provider=custom_llm_provider,
- encoding=_get_encoding(),
- stream=stream,
- )
-
- except Exception as e:
- ## LOGGING - log the original exception returned
- logging.post_call(
- input=messages,
- api_key=api_key,
- original_response=str(e),
- additional_args={"headers": headers},
- )
- raise e
+ response = _complete_snowflake(_dispatch_ctx)
elif custom_llm_provider == "gradient_ai":
- api_base = litellm.api_base or api_base
- response = base_llm_http_handler.completion(
- model=model,
- stream=stream,
- messages=messages,
- acompletion=acompletion,
- api_base=api_base,
- model_response=model_response,
- optional_params=optional_params,
- litellm_params=litellm_params,
- shared_session=shared_session,
- custom_llm_provider="gradient_ai",
- timeout=timeout,
- headers=headers,
- encoding=_get_encoding(),
- api_key=api_key,
- logging_obj=logging,
- )
+ response = _complete_gradient_ai(_dispatch_ctx)
elif custom_llm_provider == "bytez":
- api_key = (
- api_key
- or litellm.bytez_key
- or get_secret_str("BYTEZ_API_KEY")
- or litellm.api_key
- )
-
- response = base_llm_http_handler.completion(
- model=model,
- messages=messages,
- headers=headers,
- model_response=model_response,
- api_key=api_key,
- api_base=api_base,
- acompletion=acompletion,
- logging_obj=logging,
- optional_params=optional_params,
- litellm_params=litellm_params,
- timeout=timeout, # type: ignore
- client=client,
- custom_llm_provider=custom_llm_provider,
- encoding=_get_encoding(),
- stream=stream,
- provider_config=bytez_transformation,
- )
-
- pass
+ response = _complete_bytez(_dispatch_ctx)
elif custom_llm_provider == "lemonade":
- api_key = (
- api_key
- or litellm.lemonade_key
- or get_secret_str("LEMONADE_API_KEY")
- or litellm.api_key
- )
-
- response = base_llm_http_handler.completion(
- model=model,
- messages=messages,
- headers=headers,
- model_response=model_response,
- api_key=api_key,
- api_base=api_base,
- acompletion=acompletion,
- logging_obj=logging,
- optional_params=optional_params,
- litellm_params=litellm_params,
- timeout=timeout, # type: ignore
- client=client,
- custom_llm_provider=custom_llm_provider,
- encoding=_get_encoding(),
- stream=stream,
- provider_config=lemonade_transformation,
- )
-
- pass
+ response = _complete_lemonade(_dispatch_ctx)
elif custom_llm_provider == "ovhcloud" or model in litellm.ovhcloud_models:
- api_key = (
- api_key
- or litellm.ovhcloud_key
- or get_secret_str("OVHCLOUD_API_KEY")
- or litellm.api_key
- )
-
- api_base = (
- api_base
- or litellm.api_base
- or get_secret_str("OVHCLOUD_API_BASE")
- or "https://oai.endpoints.kepler.ai.cloud.ovh.net/v1"
- )
-
- response = base_llm_http_handler.completion(
- model=model,
- messages=messages,
- headers=headers,
- model_response=model_response,
- api_key=api_key,
- api_base=api_base,
- acompletion=acompletion,
- logging_obj=logging,
- optional_params=optional_params,
- litellm_params=litellm_params,
- timeout=timeout, # type: ignore
- client=client,
- custom_llm_provider=custom_llm_provider,
- encoding=_get_encoding(),
- stream=stream,
- provider_config=ovhcloud_transformation,
- )
-
- pass
+ response = _complete_ovhcloud(_dispatch_ctx)
elif custom_llm_provider == "custom":
- url = litellm.api_base or api_base or ""
- if url is None or url == "":
- raise ValueError(
- "api_base not set. Set api_base or litellm.api_base for custom endpoints"
- )
-
- """
- assume input to custom LLM api bases follow this format:
- resp = litellm.module_level_client.post(
- api_base,
- json={
- 'model': 'meta-llama/Llama-2-13b-hf', # model name
- 'params': {
- 'prompt': ["The capital of France is P"],
- 'max_tokens': 32,
- 'temperature': 0.7,
- 'top_p': 1.0,
- 'top_k': 40,
- }
- }
- )
-
- """
- prompt = " ".join([message["content"] for message in messages]) # type: ignore
- resp = litellm.module_level_client.post(
- url,
- headers=headers,
- json={
- "model": model,
- "params": {
- "prompt": [prompt],
- "max_tokens": max_tokens,
- "temperature": temperature,
- "top_p": top_p,
- "top_k": kwargs.get("top_k"),
- },
- **kwargs.get("extra_body", {}),
- },
- )
- response_json = resp.json()
- """
- assume all responses from custom api_bases of this format:
- {
- 'data': [
- {
- 'prompt': 'The capital of France is P',
- 'output': ['The capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France'],
- 'params': {'temperature': 0.7, 'top_k': 40, 'top_p': 1}}],
- 'message': 'ok'
- }
- ]
- }
- """
- string_response = response_json["data"][0]["output"][0]
- ## RESPONSE OBJECT
- model_response.choices[0].message.content = string_response # type: ignore
- model_response.created = int(time.time())
- model_response.model = model
- response = model_response
+ response = _complete_custom(_dispatch_ctx)
elif (
custom_llm_provider in litellm._custom_providers
): # Assume custom LLM provider
# Get the Custom Handler
- custom_handler: Optional[CustomLLM] = None
- for item in litellm.custom_provider_map:
- if item["provider"] == custom_llm_provider:
- custom_handler = item["custom_handler"]
-
- if custom_handler is None:
- raise LiteLLMUnknownProvider(
- model=model, custom_llm_provider=custom_llm_provider
- )
-
- ## ROUTE LLM CALL ##
- handler_fn = custom_chat_llm_router(
- async_fn=acompletion, stream=stream, custom_llm=custom_handler
- )
-
- headers = headers or litellm.headers or {}
-
- ## CALL FUNCTION
- response = handler_fn(
- model=model,
- messages=messages,
- headers=headers,
- model_response=model_response,
- print_verbose=print_verbose,
- api_key=api_key,
- api_base=api_base,
- acompletion=acompletion,
- logging_obj=logging,
- optional_params=optional_params,
- litellm_params=litellm_params,
- logger_fn=logger_fn,
- timeout=timeout, # type: ignore
- custom_prompt_dict=custom_prompt_dict,
- client=client, # pass AsyncOpenAI, OpenAI client
- encoding=_get_encoding(),
- )
- if stream is True:
- return CustomStreamWrapper(
- completion_stream=response,
- model=model,
- custom_llm_provider=custom_llm_provider,
- logging_obj=logging,
- )
+ response = _complete_custom_providers(_dispatch_ctx)
elif custom_llm_provider == "langgraph":
# LangGraph - Agent Runtime Provider
- from litellm.llms.langgraph.chat.transformation import LangGraphConfig
-
- (
- api_base,
- api_key,
- ) = LangGraphConfig()._get_openai_compatible_provider_info(
- api_base=api_base or litellm.api_base,
- api_key=api_key or litellm.api_key,
- )
-
- headers = headers or litellm.headers
-
- response = base_llm_http_handler.completion(
- model=model,
- stream=stream,
- messages=messages,
- acompletion=acompletion,
- api_base=api_base,
- model_response=model_response,
- optional_params=optional_params,
- litellm_params=litellm_params,
- shared_session=shared_session,
- custom_llm_provider=custom_llm_provider,
- timeout=timeout,
- headers=headers,
- encoding=_get_encoding(),
- api_key=api_key,
- logging_obj=logging,
- client=client,
- )
+ response = _complete_langgraph(_dispatch_ctx)
elif custom_llm_provider == "langflow":
# LangFlow - Visual AI Agent Platform
- from litellm.llms.langflow.chat.transformation import LangFlowConfig
-
- (
- api_base,
- api_key,
- ) = LangFlowConfig()._get_openai_compatible_provider_info(
- api_base=api_base or litellm.api_base,
- api_key=api_key or litellm.api_key,
- )
-
- headers = headers or litellm.headers
-
- response = base_llm_http_handler.completion(
- model=model,
- stream=stream,
- messages=messages,
- acompletion=acompletion,
- api_base=api_base,
- model_response=model_response,
- optional_params=optional_params,
- litellm_params=litellm_params,
- shared_session=shared_session,
- custom_llm_provider=custom_llm_provider,
- timeout=timeout,
- headers=headers,
- encoding=_get_encoding(),
- api_key=api_key,
- logging_obj=logging,
- client=client,
- )
+ response = _complete_langflow(_dispatch_ctx)
else:
raise LiteLLMUnknownProvider(
diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json
index 5c962cf8440..6ebac7efc8d 100644
--- a/litellm/model_prices_and_context_window_backup.json
+++ b/litellm/model_prices_and_context_window_backup.json
@@ -571,7 +571,7 @@
"output_vector_size": 1536
},
"amazon.titan-embed-text-v2:0": {
- "input_cost_per_token": 2e-07,
+ "input_cost_per_token": 2e-08,
"litellm_provider": "bedrock",
"max_input_tokens": 8192,
"max_tokens": 8192,
@@ -10443,7 +10443,8 @@
"fast": 6.0
},
"supports_output_config": true,
- "supports_max_reasoning_effort": true
+ "supports_max_reasoning_effort": true,
+ "supports_speed": true
},
"claude-opus-4-6-20260205": {
"cache_creation_input_token_cost": 6.25e-06,
@@ -10476,7 +10477,8 @@
"fast": 6.0
},
"supports_max_reasoning_effort": true,
- "supports_output_config": true
+ "supports_output_config": true,
+ "supports_speed": true
},
"claude-opus-4-7": {
"cache_creation_input_token_cost": 6.25e-06,
@@ -10511,7 +10513,8 @@
"us": 1.1,
"fast": 6.0
},
- "supports_output_config": true
+ "supports_output_config": true,
+ "supports_speed": true
},
"claude-opus-4-7-20260416": {
"cache_creation_input_token_cost": 6.25e-06,
@@ -10546,7 +10549,8 @@
"us": 1.1,
"fast": 6.0
},
- "supports_output_config": true
+ "supports_output_config": true,
+ "supports_speed": true
},
"claude-fable-5": {
"cache_creation_input_token_cost": 1.25e-05,
@@ -10615,7 +10619,8 @@
"us": 1.1,
"fast": 2.0
},
- "supports_output_config": true
+ "supports_output_config": true,
+ "supports_speed": true
},
"claude-sonnet-4-20250514": {
"deprecation_date": "2026-05-14",
@@ -10684,6 +10689,268 @@
"mode": "chat",
"output_cost_per_token": 1.923e-06
},
+ "cloudflare/@cf/openai/gpt-oss-120b": {
+ "input_cost_per_token": 3.5e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 128000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 7.5e-07,
+ "supports_function_calling": true,
+ "supports_reasoning": true
+ },
+ "cloudflare/@cf/google/gemma-2b-it-lora": {
+ "input_cost_per_token": 0.0,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 8192,
+ "max_output_tokens": 8192,
+ "max_tokens": 8192,
+ "mode": "chat",
+ "output_cost_per_token": 0.0
+ },
+ "cloudflare/@cf/meta/llama-3.2-3b-instruct": {
+ "input_cost_per_token": 5.09e-08,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 80000,
+ "max_output_tokens": 80000,
+ "max_tokens": 80000,
+ "mode": "chat",
+ "output_cost_per_token": 3.35e-07
+ },
+ "cloudflare/@cf/meta/llama-guard-3-8b": {
+ "input_cost_per_token": 4.84e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 131072,
+ "max_output_tokens": 131072,
+ "max_tokens": 131072,
+ "mode": "chat",
+ "output_cost_per_token": 3e-08
+ },
+ "cloudflare/@cf/mistral/mistral-7b-instruct-v0.2-lora": {
+ "input_cost_per_token": 0.0,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 15000,
+ "max_output_tokens": 15000,
+ "max_tokens": 15000,
+ "mode": "chat",
+ "output_cost_per_token": 0.0
+ },
+ "cloudflare/@cf/moonshotai/kimi-k2.7-code": {
+ "cache_read_input_token_cost": 1.9e-07,
+ "input_cost_per_token": 9.5e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 262144,
+ "max_output_tokens": 262144,
+ "max_tokens": 262144,
+ "mode": "chat",
+ "output_cost_per_token": 4e-06,
+ "supports_function_calling": true,
+ "supports_reasoning": true
+ },
+ "cloudflare/@cf/deepseek-ai/deepseek-r1-distill-qwen-32b": {
+ "input_cost_per_token": 4.97e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 80000,
+ "max_output_tokens": 80000,
+ "max_tokens": 80000,
+ "mode": "chat",
+ "output_cost_per_token": 4.881e-06,
+ "supports_reasoning": true
+ },
+ "cloudflare/@cf/meta/llama-3.1-8b-instruct-fp8": {
+ "input_cost_per_token": 1.52e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 32000,
+ "max_output_tokens": 32000,
+ "max_tokens": 32000,
+ "mode": "chat",
+ "output_cost_per_token": 2.87e-07
+ },
+ "cloudflare/@cf/meta/llama-3.2-1b-instruct": {
+ "input_cost_per_token": 2.7e-08,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 60000,
+ "max_output_tokens": 60000,
+ "max_tokens": 60000,
+ "mode": "chat",
+ "output_cost_per_token": 2.01e-07
+ },
+ "cloudflare/@cf/moonshotai/kimi-k2.6": {
+ "cache_read_input_token_cost": 1.6e-07,
+ "input_cost_per_token": 9.5e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 262144,
+ "max_output_tokens": 262144,
+ "max_tokens": 262144,
+ "mode": "chat",
+ "output_cost_per_token": 4e-06,
+ "supports_function_calling": true,
+ "supports_reasoning": true
+ },
+ "cloudflare/@cf/zai-org/glm-4.7-flash": {
+ "input_cost_per_token": 6.05e-08,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 131072,
+ "max_output_tokens": 131072,
+ "max_tokens": 131072,
+ "mode": "chat",
+ "output_cost_per_token": 4e-07,
+ "supports_function_calling": true,
+ "supports_reasoning": true
+ },
+ "cloudflare/@cf/meta-llama/llama-2-7b-chat-hf-lora": {
+ "input_cost_per_token": 0.0,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 8192,
+ "max_output_tokens": 8192,
+ "max_tokens": 8192,
+ "mode": "chat",
+ "output_cost_per_token": 0.0
+ },
+ "cloudflare/@cf/meta/llama-3.3-70b-instruct-fp8-fast": {
+ "input_cost_per_token": 2.93e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 24000,
+ "max_output_tokens": 24000,
+ "max_tokens": 24000,
+ "mode": "chat",
+ "output_cost_per_token": 2.253e-06,
+ "supports_function_calling": true
+ },
+ "cloudflare/@cf/ibm-granite/granite-4.0-h-micro": {
+ "input_cost_per_token": 1.7e-08,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 131000,
+ "max_output_tokens": 131000,
+ "max_tokens": 131000,
+ "mode": "chat",
+ "output_cost_per_token": 1.12e-07,
+ "supports_function_calling": true
+ },
+ "cloudflare/@cf/qwen/qwen2.5-coder-32b-instruct": {
+ "input_cost_per_token": 6.6e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 32768,
+ "max_output_tokens": 32768,
+ "max_tokens": 32768,
+ "mode": "chat",
+ "output_cost_per_token": 1e-06
+ },
+ "cloudflare/@cf/zai-org/glm-5.2": {
+ "cache_read_input_token_cost": 2.6e-07,
+ "input_cost_per_token": 1.4e-06,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 262144,
+ "max_output_tokens": 262144,
+ "max_tokens": 262144,
+ "mode": "chat",
+ "output_cost_per_token": 4.4e-06,
+ "supports_function_calling": true,
+ "supports_reasoning": true
+ },
+ "cloudflare/@cf/nvidia/nemotron-3-120b-a12b": {
+ "input_cost_per_token": 5e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 256000,
+ "max_output_tokens": 256000,
+ "max_tokens": 256000,
+ "mode": "chat",
+ "output_cost_per_token": 1.5e-06,
+ "supports_function_calling": true,
+ "supports_reasoning": true
+ },
+ "cloudflare/@cf/aisingapore/gemma-sea-lion-v4-27b-it": {
+ "input_cost_per_token": 3.51e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 128000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 5.55e-07
+ },
+ "cloudflare/@cf/qwen/qwen3-30b-a3b-fp8": {
+ "input_cost_per_token": 5.09e-08,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 32768,
+ "max_output_tokens": 32768,
+ "max_tokens": 32768,
+ "mode": "chat",
+ "output_cost_per_token": 3.35e-07,
+ "supports_function_calling": true,
+ "supports_reasoning": true
+ },
+ "cloudflare/@cf/google/gemma-7b-it-lora": {
+ "input_cost_per_token": 0.0,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 3500,
+ "max_output_tokens": 3500,
+ "max_tokens": 3500,
+ "mode": "chat",
+ "output_cost_per_token": 0.0
+ },
+ "cloudflare/@cf/google/gemma-4-26b-a4b-it": {
+ "input_cost_per_token": 1e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 256000,
+ "max_output_tokens": 256000,
+ "max_tokens": 256000,
+ "mode": "chat",
+ "output_cost_per_token": 3e-07,
+ "supports_function_calling": true,
+ "supports_reasoning": true
+ },
+ "cloudflare/@cf/mistralai/mistral-small-3.1-24b-instruct": {
+ "input_cost_per_token": 3.51e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 128000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 5.55e-07,
+ "supports_function_calling": true
+ },
+ "cloudflare/@cf/meta/llama-3.2-11b-vision-instruct": {
+ "input_cost_per_token": 4.85e-08,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 128000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 6.76e-07,
+ "supports_vision": true
+ },
+ "cloudflare/@cf/openai/gpt-oss-20b": {
+ "input_cost_per_token": 2e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 128000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 3e-07,
+ "supports_function_calling": true,
+ "supports_reasoning": true
+ },
+ "cloudflare/@cf/meta/llama-4-scout-17b-16e-instruct": {
+ "input_cost_per_token": 2.7e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 131000,
+ "max_output_tokens": 131000,
+ "max_tokens": 131000,
+ "mode": "chat",
+ "output_cost_per_token": 8.5e-07,
+ "supports_function_calling": true
+ },
+ "cloudflare/@cf/qwen/qwq-32b": {
+ "input_cost_per_token": 6.6e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 24000,
+ "max_output_tokens": 24000,
+ "max_tokens": 24000,
+ "mode": "chat",
+ "output_cost_per_token": 1e-06,
+ "supports_reasoning": true
+ },
"codestral/codestral-2405": {
"input_cost_per_token": 0.0,
"litellm_provider": "codestral",
@@ -20088,8 +20355,6 @@
"output_cost_per_token": 8e-06,
"output_cost_per_token_batches": 4e-06,
"output_cost_per_token_priority": 1.4e-05,
- "regional_processing_uplift_multiplier_eu": 1.10,
- "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -20163,8 +20428,6 @@
"output_cost_per_token": 1.6e-06,
"output_cost_per_token_batches": 8e-07,
"output_cost_per_token_priority": 2.8e-06,
- "regional_processing_uplift_multiplier_eu": 1.10,
- "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -20238,8 +20501,6 @@
"output_cost_per_token": 4e-07,
"output_cost_per_token_batches": 2e-07,
"output_cost_per_token_priority": 8e-07,
- "regional_processing_uplift_multiplier_eu": 1.10,
- "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -20311,8 +20572,6 @@
"output_cost_per_token": 1e-05,
"output_cost_per_token_batches": 5e-06,
"output_cost_per_token_priority": 1.7e-05,
- "regional_processing_uplift_multiplier_eu": 1.10,
- "regional_processing_uplift_multiplier_us": 1.10,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
@@ -20354,8 +20613,6 @@
"mode": "chat",
"output_cost_per_token": 1e-05,
"output_cost_per_token_batches": 5e-06,
- "regional_processing_uplift_multiplier_eu": 1.10,
- "regional_processing_uplift_multiplier_us": 1.10,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
@@ -20377,8 +20634,6 @@
"mode": "chat",
"output_cost_per_token": 1e-05,
"output_cost_per_token_batches": 5e-06,
- "regional_processing_uplift_multiplier_eu": 1.10,
- "regional_processing_uplift_multiplier_us": 1.10,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
@@ -20667,8 +20922,6 @@
"output_cost_per_token": 6e-07,
"output_cost_per_token_batches": 3e-07,
"output_cost_per_token_priority": 1e-06,
- "regional_processing_uplift_multiplier_eu": 1.10,
- "regional_processing_uplift_multiplier_us": 1.10,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
@@ -21372,8 +21625,6 @@
"output_cost_per_token": 1e-05,
"output_cost_per_token_flex": 5e-06,
"output_cost_per_token_priority": 2e-05,
- "regional_processing_uplift_multiplier_eu": 1.10,
- "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -21767,6 +22018,8 @@
"output_cost_per_token_flex": 1.5e-05,
"output_cost_per_token_batches": 1.5e-05,
"output_cost_per_token_priority": 6e-05,
+ "regional_processing_uplift_multiplier_eu": 1.10,
+ "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -21815,6 +22068,8 @@
"output_cost_per_token_flex": 1.5e-05,
"output_cost_per_token_batches": 1.5e-05,
"output_cost_per_token_priority": 6e-05,
+ "regional_processing_uplift_multiplier_eu": 1.10,
+ "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -21859,6 +22114,8 @@
"output_cost_per_token_above_272k_tokens": 0.00027,
"output_cost_per_token_flex": 9e-05,
"output_cost_per_token_batches": 9e-05,
+ "regional_processing_uplift_multiplier_eu": 1.10,
+ "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/responses",
"/v1/batch"
@@ -21903,6 +22160,8 @@
"output_cost_per_token_above_272k_tokens": 0.00027,
"output_cost_per_token_flex": 9e-05,
"output_cost_per_token_batches": 9e-05,
+ "regional_processing_uplift_multiplier_eu": 1.10,
+ "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/responses",
"/v1/batch"
@@ -21951,6 +22210,8 @@
"output_cost_per_token_flex": 7.5e-06,
"output_cost_per_token_batches": 7.5e-06,
"output_cost_per_token_priority": 3e-05,
+ "regional_processing_uplift_multiplier_eu": 1.10,
+ "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -21998,6 +22259,8 @@
"output_cost_per_token_flex": 7.5e-06,
"output_cost_per_token_batches": 7.5e-06,
"output_cost_per_token_priority": 3e-05,
+ "regional_processing_uplift_multiplier_eu": 1.10,
+ "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -22038,6 +22301,8 @@
"output_cost_per_token_above_272k_tokens": 0.00027,
"output_cost_per_token_flex": 9e-05,
"output_cost_per_token_batches": 9e-05,
+ "regional_processing_uplift_multiplier_eu": 1.10,
+ "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/responses",
"/v1/batch"
@@ -22081,6 +22346,8 @@
"output_cost_per_token_above_272k_tokens": 0.00027,
"output_cost_per_token_flex": 9e-05,
"output_cost_per_token_batches": 9e-05,
+ "regional_processing_uplift_multiplier_eu": 1.10,
+ "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/responses",
"/v1/batch"
@@ -22126,6 +22393,8 @@
"output_cost_per_token_flex": 2.25e-06,
"output_cost_per_token_batches": 2.25e-06,
"output_cost_per_token_priority": 9e-06,
+ "regional_processing_uplift_multiplier_eu": 1.10,
+ "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -22172,6 +22441,8 @@
"output_cost_per_token_flex": 2.25e-06,
"output_cost_per_token_batches": 2.25e-06,
"output_cost_per_token_priority": 9e-06,
+ "regional_processing_uplift_multiplier_eu": 1.10,
+ "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -22215,6 +22486,8 @@
"output_cost_per_token": 1.25e-06,
"output_cost_per_token_flex": 6.25e-07,
"output_cost_per_token_batches": 6.25e-07,
+ "regional_processing_uplift_multiplier_eu": 1.10,
+ "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -22258,6 +22531,8 @@
"output_cost_per_token": 1.25e-06,
"output_cost_per_token_flex": 6.25e-07,
"output_cost_per_token_batches": 6.25e-07,
+ "regional_processing_uplift_multiplier_eu": 1.10,
+ "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -22296,8 +22571,6 @@
"mode": "responses",
"output_cost_per_token": 0.00012,
"output_cost_per_token_batches": 6e-05,
- "regional_processing_uplift_multiplier_eu": 1.10,
- "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/batch",
"/v1/responses"
@@ -22704,8 +22977,6 @@
"output_cost_per_token": 2e-06,
"output_cost_per_token_flex": 1e-06,
"output_cost_per_token_priority": 3.6e-06,
- "regional_processing_uplift_multiplier_eu": 1.10,
- "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -22787,8 +23058,6 @@
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
- "regional_processing_uplift_multiplier_eu": 1.10,
- "regional_processing_uplift_multiplier_us": 1.10,
"mode": "chat",
"output_cost_per_token": 4e-07,
"output_cost_per_token_flex": 2e-07,
@@ -39908,24 +40177,6 @@
"litellm_provider": "fireworks_ai",
"mode": "chat"
},
- "fireworks_ai/accounts/fireworks/models/whisper-v3": {
- "max_tokens": 4096,
- "max_input_tokens": 4096,
- "max_output_tokens": 4096,
- "input_cost_per_token": 0.0,
- "output_cost_per_token": 0.0,
- "litellm_provider": "fireworks_ai",
- "mode": "audio_transcription"
- },
- "fireworks_ai/accounts/fireworks/models/whisper-v3-turbo": {
- "max_tokens": 4096,
- "max_input_tokens": 4096,
- "max_output_tokens": 4096,
- "input_cost_per_token": 0.0,
- "output_cost_per_token": 0.0,
- "litellm_provider": "fireworks_ai",
- "mode": "audio_transcription"
- },
"fireworks_ai/accounts/fireworks/models/yi-34b": {
"max_tokens": 4096,
"max_input_tokens": 4096,
@@ -43061,6 +43312,40 @@
"supports_tool_choice": true,
"supports_vision": false
},
+ "darkbloom/gemma-4-26b": {
+ "input_cost_per_token": 3e-08,
+ "litellm_provider": "darkbloom",
+ "max_input_tokens": 131072,
+ "max_output_tokens": 32768,
+ "max_tokens": 32768,
+ "mode": "chat",
+ "output_cost_per_token": 1.65e-07,
+ "source": "https://www.darkbloom.dev/",
+ "supported_endpoints": [
+ "/v1/chat/completions"
+ ],
+ "supports_function_calling": true,
+ "supports_native_streaming": true,
+ "supports_system_messages": true,
+ "supports_tool_choice": true
+ },
+ "darkbloom/gpt-oss-20b": {
+ "input_cost_per_token": 1.45e-08,
+ "litellm_provider": "darkbloom",
+ "max_input_tokens": 131072,
+ "max_output_tokens": 32768,
+ "max_tokens": 32768,
+ "mode": "chat",
+ "output_cost_per_token": 7e-08,
+ "source": "https://www.darkbloom.dev/",
+ "supported_endpoints": [
+ "/v1/chat/completions"
+ ],
+ "supports_function_calling": true,
+ "supports_native_streaming": true,
+ "supports_system_messages": true,
+ "supports_tool_choice": true
+ },
"deepseek/deepseek-v4-pro": {
"cache_creation_input_token_cost": 0.0,
"cache_read_input_token_cost": 3.625e-09,
diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py
index b27082c361a..3a9ef8db804 100644
--- a/litellm/ocr/main.py
+++ b/litellm/ocr/main.py
@@ -10,7 +10,7 @@ import os
import re
from functools import partial
from io import IOBase
-from typing import Any, Coroutine, Dict, Optional, Union
+from typing import Any, Callable, Coroutine, Dict, Optional, Union, cast
import httpx
@@ -20,6 +20,7 @@ from litellm.constants import request_timeout
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig, OCRResponse
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
+from litellm.ocr.rust_bridge import RustOcr, load_rust_ocr, rust_ocr_enabled
from litellm.types.router import GenericLiteLLMParams
from litellm.utils import ProviderConfigManager, client
@@ -28,6 +29,82 @@ base_llm_http_handler = BaseLLMHTTPHandler()
#################################################
+def _timeout_to_seconds(
+ timeout: Optional[Union[float, httpx.Timeout]],
+) -> Optional[float]:
+ """Convert the Python OCR timeout to a single seconds value for the Rust bridge.
+
+ The Rust HTTP client takes one duration; ``httpx.Timeout`` carries separate
+ connect/read/write/pool values, so pick the read deadline as the closest
+ analog to a total-request timeout.
+ """
+ if timeout is None:
+ return None
+ if isinstance(timeout, httpx.Timeout):
+ return timeout.read
+ return float(timeout)
+
+
+def _run_rust_ocr(
+ rust_ocr: RustOcr,
+ logging_obj: LiteLLMLoggingObj,
+ provider_config: BaseOCRConfig,
+ resolve_api_key: Callable[[str], Optional[str]],
+ model: str,
+ document: dict[str, object],
+ api_key: Optional[str],
+ api_base: Optional[str],
+ optional_params: dict[str, object],
+ litellm_params: dict[str, object],
+ timeout_seconds: Optional[float],
+) -> OCRResponse:
+ """Run the Mistral OCR call through the Rust bridge and wrap the result.
+
+ Resolves the key the same way the Python path does so secret-manager backends
+ (AWS/Azure/GCP/Vault) work; the Rust bridge's own fallback only reads the
+ process environment. The request that Rust actually sends (resolved URL and
+ headers) is mirrored into pre_call so logs match the wire. Dependencies are
+ injected so this stays unit-testable without patching module globals.
+ """
+ resolved_api_key = api_key or resolve_api_key("MISTRAL_API_KEY")
+ resolved_headers = provider_config.validate_environment(
+ headers={},
+ model=model,
+ api_key=resolved_api_key,
+ api_base=api_base,
+ litellm_params=litellm_params,
+ )
+ resolved_complete_url = provider_config.get_complete_url(
+ api_base=api_base,
+ model=model,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ )
+ logging_obj.pre_call(
+ input="OCR document processing",
+ api_key=resolved_api_key,
+ additional_args={
+ "complete_input_dict": {
+ "model": model,
+ "document": document,
+ **optional_params,
+ },
+ "api_base": resolved_complete_url,
+ "headers": resolved_headers,
+ },
+ )
+ return OCRResponse.model_validate(
+ rust_ocr(
+ model=model,
+ document=document,
+ api_key=resolved_api_key,
+ api_base=api_base,
+ optional_params=optional_params,
+ timeout_seconds=timeout_seconds,
+ )
+ )
+
+
@client
async def aocr(
model: str,
@@ -220,7 +297,7 @@ def ocr(
"""
local_vars = locals()
try:
- litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore
+ litellm_logging_obj = cast(LiteLLMLoggingObj, kwargs.pop("litellm_logging_obj"))
litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None)
_is_async = kwargs.pop("aocr", False) is True
@@ -261,7 +338,6 @@ def ocr(
if dynamic_api_base:
api_base = dynamic_api_base
- # Get provider config
ocr_provider_config: Optional[BaseOCRConfig] = (
ProviderConfigManager.get_provider_ocr_config(
model=model,
@@ -278,17 +354,14 @@ def ocr(
f"OCR call - model: {model}, provider: {custom_llm_provider}"
)
- # Get litellm params using GenericLiteLLMParams (same as responses API)
litellm_params = GenericLiteLLMParams(**kwargs)
- # Extract OCR-specific parameters from kwargs
supported_params = ocr_provider_config.get_supported_ocr_params(model=model)
non_default_params = {}
for param in supported_params:
if param in kwargs:
non_default_params[param] = kwargs.pop(param)
- # Map parameters to provider-specific format
optional_params = ocr_provider_config.map_ocr_params(
non_default_params=non_default_params,
optional_params={},
@@ -297,7 +370,8 @@ def ocr(
verbose_logger.debug(f"OCR optional_params after mapping: {optional_params}")
- # Pre Call logging
+ effective_timeout = timeout or request_timeout
+
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model,
@@ -309,12 +383,35 @@ def ocr(
custom_llm_provider=custom_llm_provider,
)
- # Call the handler - pass document dict directly
+ # Optional Rust path: hand the whole Mistral OCR call to the Rust bridge.
+ if custom_llm_provider == "mistral" and rust_ocr_enabled():
+ rust_ocr = load_rust_ocr()
+ if rust_ocr is None:
+ verbose_logger.debug(
+ "Rust OCR bridge unavailable; falling back to Python path"
+ )
+ else:
+ from litellm.secret_managers.main import get_secret_str
+
+ return _run_rust_ocr(
+ rust_ocr=rust_ocr,
+ logging_obj=litellm_logging_obj,
+ provider_config=ocr_provider_config,
+ resolve_api_key=get_secret_str,
+ model=model,
+ document=document,
+ api_key=api_key,
+ api_base=api_base,
+ optional_params=optional_params,
+ litellm_params=dict(litellm_params),
+ timeout_seconds=_timeout_to_seconds(effective_timeout),
+ )
+
response = base_llm_http_handler.ocr(
model=model,
- document=document, # Pass the entire document dict
+ document=document,
optional_params=optional_params,
- timeout=timeout or request_timeout,
+ timeout=effective_timeout,
logging_obj=litellm_logging_obj,
api_key=api_key,
api_base=api_base,
diff --git a/litellm/ocr/rust_bridge.py b/litellm/ocr/rust_bridge.py
new file mode 100644
index 00000000000..61f9e8ca69a
--- /dev/null
+++ b/litellm/ocr/rust_bridge.py
@@ -0,0 +1,74 @@
+"""
+Optional Rust-backed OCR path.
+
+Enable with ``litellm.use_litellm_rust()``; the sync ``litellm.ocr()`` entrypoint
+then routes supported Mistral calls through the compiled ``litellm_python_bridge``
+extension, which performs the whole OCR call (URL, headers, HTTP, parse) in Rust.
+
+No module-level ``litellm`` imports keep this a leaf so ``litellm/ocr/main.py``
+can import it statically without forming an import cycle.
+"""
+
+from __future__ import annotations
+
+from typing import Final, Protocol, cast
+
+
+class RustOcr(Protocol):
+ """Signature of the compiled ``litellm_python_bridge.ocr`` entrypoint."""
+
+ def __call__(
+ self,
+ model: str,
+ document: dict[str, object],
+ api_key: str | None,
+ api_base: str | None,
+ optional_params: dict[str, object],
+ timeout_seconds: float | None,
+ ) -> dict[str, object]: ...
+
+
+class _Unset:
+ """Sentinel type so ``ocr=None`` can clear a prior injection while omission preserves it."""
+
+
+_UNSET: Final[_Unset] = _Unset()
+
+_rust_ocr_enabled = False
+_rust_ocr_impl: RustOcr | None = None
+
+
+def use_litellm_rust(
+ enabled: bool = True, *, ocr: RustOcr | None | _Unset = _UNSET
+) -> None:
+ """Route supported OCR calls through the Rust ``litellm_python_bridge`` extension.
+
+ ``ocr`` injects the bridge callable; when omitted the compiled extension is
+ loaded on demand and any previously injected bridge is preserved. Pass
+ ``ocr=None`` explicitly to clear a prior injection.
+ """
+ global _rust_ocr_enabled, _rust_ocr_impl
+ _rust_ocr_enabled = enabled
+ if not isinstance(ocr, _Unset):
+ _rust_ocr_impl = ocr
+
+
+def rust_ocr_enabled() -> bool:
+ """Whether the Rust OCR path has been turned on via ``use_litellm_rust()``."""
+ return _rust_ocr_enabled
+
+
+def load_rust_ocr() -> RustOcr | None:
+ """Return the Rust OCR callable, or ``None`` when no bridge is available.
+
+ Prefers an injected implementation, otherwise loads the compiled
+ ``litellm_python_bridge`` extension; a missing extension yields ``None`` so
+ the caller can fall back to the Python path instead of hard-failing.
+ """
+ if _rust_ocr_impl is not None:
+ return _rust_ocr_impl
+ try:
+ import litellm_python_bridge
+ except ImportError:
+ return None
+ return cast(RustOcr, litellm_python_bridge.ocr)
diff --git a/litellm/provider_endpoints_support_backup.json b/litellm/provider_endpoints_support_backup.json
index db6183edaa0..dd7712aabca 100644
--- a/litellm/provider_endpoints_support_backup.json
+++ b/litellm/provider_endpoints_support_backup.json
@@ -1835,6 +1835,23 @@
"interactions": true
}
},
+ "darkbloom": {
+ "display_name": "Darkbloom (`darkbloom`)",
+ "url": "https://docs.litellm.ai/docs/providers/darkbloom",
+ "endpoints": {
+ "chat_completions": true,
+ "messages": false,
+ "responses": false,
+ "embeddings": false,
+ "image_generations": false,
+ "audio_transcriptions": false,
+ "audio_speech": false,
+ "moderations": false,
+ "batches": false,
+ "rerank": false,
+ "a2a": false
+ }
+ },
"predibase": {
"display_name": "Predibase (`predibase`)",
"url": "https://docs.litellm.ai/docs/providers/predibase",
diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/__init__.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/__init__.py
index b357504979d..73166a45d6e 100644
--- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/__init__.py
+++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/__init__.py
@@ -1,16 +1,20 @@
"""Typed upstream-credential resolution for MCP servers.
-This subpackage houses the typed credential vocabulary and (in a later PR) the
-``resolve_credentials`` dispatch. A server declares one per-mode config from the
-``AuthConfig`` discriminated union; failures are modeled as values via :mod:`.result`
-(``Result[T, CredError]``) rather than raised, so every seam is total. Nothing here is
-wired onto a live request path yet.
+This subpackage houses the typed credential vocabulary and the ``resolve_credentials``
+dispatch. A server declares one per-mode config from the ``AuthConfig`` discriminated union;
+``UpstreamCredentialProvider.resolve_credentials`` selects one arm and returns an ``httpx.Auth``
+or a typed ``CredError``. Failures are modeled as values via :mod:`.result` (``Result[T,
+CredError]``) rather than raised, so every seam is total. Nothing here is wired onto a live
+request path yet.
"""
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
NoOpAuth,
StaticHeaderAuth,
)
+from litellm.proxy._experimental.mcp_server.outbound_credentials.resolver import (
+ UpstreamCredentialProvider,
+)
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import (
Error,
Ok,
@@ -45,6 +49,7 @@ __all__ = [
"Result",
"NoOpAuth",
"StaticHeaderAuth",
+ "UpstreamCredentialProvider",
"AuthSpecKind",
"CredError",
"Subject",
diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py
new file mode 100644
index 00000000000..7bcdb3e6529
--- /dev/null
+++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py
@@ -0,0 +1,70 @@
+"""The one credential resolver: dispatch on the declared mode, fail closed.
+
+`resolve_credentials` selects exactly one arm off the server's typed `config` and either
+produces an `httpx.Auth` or returns a typed `CredError`. The `match` is over the `AuthConfig`
+variant, so each arm receives its own fully-typed config with no field-presence inference and
+no precedence cascade. It is wildcard-free with an `assert_never` tail, so adding a mode without
+an arm fails the type gate (basedpyright `reportMatchNotExhaustive`); a bypassed gate fails loudly
+at runtime instead of returning `None`.
+
+This skeleton ships every arm as a `not_implemented` stub. Each mode's real body, with its
+injected seam, lands in its own follow-up PR; until then the arm returns a typed error rather
+than silently producing no credential. Pure v2: no imports from v1.
+"""
+
+from __future__ import annotations
+
+import httpx
+from typing_extensions import assert_never
+
+from litellm.proxy._experimental.mcp_server.outbound_credentials.result import (
+ Error,
+ Result,
+)
+from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
+ ApiKeyConfig,
+ AuthorizationCodeConfig,
+ AuthSpecKind,
+ AwsSigV4Config,
+ ClientCredentialsConfig,
+ CredError,
+ NoneConfig,
+ PassthroughConfig,
+ ServerSpec,
+ Subject,
+ TokenExchangeConfig,
+)
+
+
+class UpstreamCredentialProvider:
+ """Produces the one `httpx.Auth` for a `(subject, upstream)` pair, per declared mode.
+
+ Collaborators (the per-mode credential stores and token fetchers) are injected as each arm
+ is built; the skeleton needs none, since every arm is a stub.
+ """
+
+ async def resolve_credentials(
+ self, subject: Subject, server: ServerSpec
+ ) -> Result[httpx.Auth, CredError]:
+ match server.config:
+ case NoneConfig():
+ return _not_implemented(AuthSpecKind.none)
+ case ApiKeyConfig():
+ return _not_implemented(AuthSpecKind.api_key)
+ case PassthroughConfig():
+ return _not_implemented(AuthSpecKind.passthrough)
+ case ClientCredentialsConfig():
+ return _not_implemented(AuthSpecKind.client_credentials)
+ case TokenExchangeConfig():
+ return _not_implemented(AuthSpecKind.token_exchange)
+ case AuthorizationCodeConfig():
+ return _not_implemented(AuthSpecKind.authorization_code)
+ case AwsSigV4Config():
+ return _not_implemented(AuthSpecKind.aws_sigv4)
+ assert_never(server.config)
+
+
+def _not_implemented(kind: AuthSpecKind) -> Result[httpx.Auth, CredError]:
+ return Error(
+ CredError.of_not_implemented(f"{kind.value}: resolver arm not implemented yet")
+ )
diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py
index ac90302fdaa..5bba842c7eb 100644
--- a/litellm/proxy/_types.py
+++ b/litellm/proxy/_types.py
@@ -3358,7 +3358,9 @@ class ProxyException(Exception):
class CommonProxyErrors(str, enum.Enum):
db_not_connected_error = (
- "DB not connected. See https://docs.litellm.ai/docs/proxy/virtual_keys"
+ "DB not connected. This endpoint needs a database; set DATABASE_URL to a "
+ "PostgreSQL connection string (postgresql://...) to enable it. "
+ "See https://docs.litellm.ai/docs/proxy/virtual_keys"
)
no_llm_router = "No models configured on proxy"
not_allowed_access = "Admin-only endpoint. Not allowed to access this."
diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py
index 52788ed9238..88db2a2b7ea 100644
--- a/litellm/proxy/auth/auth_checks.py
+++ b/litellm/proxy/auth/auth_checks.py
@@ -2954,6 +2954,26 @@ async def _get_agent_ids_from_access_groups(
)
+def _resolve_all_team_model_sentinel_for_auth_check(
+ models: List[str],
+ llm_router: Optional[Router],
+ team_id: Optional[str],
+) -> List[str]:
+ if (
+ SpecialModelNames.all_team_models.value not in models
+ or team_id is None
+ or llm_router is None
+ ):
+ return models
+ proxy_models = llm_router.get_model_names()
+ non_sentinel_models = [
+ model for model in models if model != SpecialModelNames.all_team_models.value
+ ]
+ if not proxy_models:
+ return non_sentinel_models or models
+ return list(dict.fromkeys(non_sentinel_models + proxy_models))
+
+
def _check_model_access_helper(
model: str,
llm_router: Optional[Router],
@@ -2971,6 +2991,12 @@ def _check_model_access_helper(
model_name=model, team_id=team_id
)
+ models = _resolve_all_team_model_sentinel_for_auth_check(
+ models=models,
+ llm_router=llm_router,
+ team_id=team_id,
+ )
+
if (
len(access_groups) > 0 and llm_router is not None
): # check if token contains any model access groups
diff --git a/litellm/proxy/auth/model_checks.py b/litellm/proxy/auth/model_checks.py
index b89db51c6f1..aa53954da8f 100644
--- a/litellm/proxy/auth/model_checks.py
+++ b/litellm/proxy/auth/model_checks.py
@@ -122,9 +122,16 @@ def get_key_models(
SpecialModelNames.all_team_models.value in all_models
and user_api_key_dict.team_id is not None
):
- all_models = list(
- user_api_key_dict.team_models
- ) # copy to avoid mutating cached objects
+ all_models = list(user_api_key_dict.team_models)
+ if SpecialModelNames.all_team_models.value in all_models:
+ all_models = [
+ model
+ for model in all_models
+ if model != SpecialModelNames.all_team_models.value
+ ]
+ all_models.extend(proxy_model_list)
+ if include_model_access_groups:
+ all_models.extend(model_access_groups.keys())
if SpecialModelNames.all_proxy_models.value in all_models:
all_models = list(proxy_model_list) # copy to avoid mutating caller's list
if include_model_access_groups:
@@ -160,6 +167,12 @@ def get_team_models(
all_models_set.update(team_models)
if SpecialModelNames.all_team_models.value in all_models_set:
all_models_set.update(team_models)
+ # GH#30619: expand all-team-models sentinel
+ # to the actual proxy model list
+ all_models_set.discard(SpecialModelNames.all_team_models.value)
+ all_models_set.update(proxy_model_list)
+ if include_model_access_groups:
+ all_models_set.update(model_access_groups.keys())
if SpecialModelNames.all_proxy_models.value in all_models_set:
all_models_set.update(proxy_model_list)
if include_model_access_groups:
diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py
index 8ef931e8d25..8dec08460b4 100644
--- a/litellm/proxy/common_request_processing.py
+++ b/litellm/proxy/common_request_processing.py
@@ -1037,6 +1037,8 @@ class ProxyBaseLLMRequestProcessing:
version=version,
proxy_config=proxy_config,
)
+ if not general_settings.get("expose_fallback_errors_to_caller"):
+ self.data.pop("include_fallback_errors", None)
if route_type in {"aresponses", "_aresponses_websocket"}:
await _authorize_response_file_search_vector_stores(
data=self.data,
diff --git a/litellm/proxy/db/db_url_settings.py b/litellm/proxy/db/db_url_settings.py
index 58478db5e2e..ae2307658dd 100644
--- a/litellm/proxy/db/db_url_settings.py
+++ b/litellm/proxy/db/db_url_settings.py
@@ -32,7 +32,7 @@ password when their ``*_READ_REPLICA`` counterpart is unset.
import os
import urllib.parse
-from typing import Optional, cast
+from typing import Final, cast
from pydantic import AliasChoices, Field
from pydantic_settings import BaseSettings, SettingsConfigDict
@@ -44,6 +44,41 @@ from litellm.proxy.auth import rds_iam_token
_IAM_ENV_KEY = "IAM_TOKEN_DB_AUTH"
_DEFAULT_PG_PORT = "5432"
+# schema.prisma pins `provider = "postgresql"`, so these are the only schemes
+# Prisma can actually connect with.
+SUPPORTED_DB_SCHEMES: Final[frozenset[str]] = frozenset({"postgresql", "postgres"})
+_MISSING_SCHEME = ""
+
+
+def unsupported_db_scheme(database_url: str) -> str | None:
+ """Return the connection URL scheme when it is not PostgreSQL, else None.
+
+ A `sqlite://` / `mysql://` URL can never connect against the
+ postgresql-only datasource, but the resulting Prisma failure is opaque and
+ version-dependent (a confusing migration error, or a startup that never
+ binds). Callers use this to reject the URL up front with an actionable
+ error instead.
+
+ A schemeless value (e.g. a malformed DSN like ``user:pass@host/db``) yields
+ the ``_MISSING_SCHEME`` placeholder rather than the raw URL, so callers that
+ log the return value never echo embedded credentials.
+ """
+ scheme = urllib.parse.urlsplit(database_url).scheme.lower()
+ if scheme in SUPPORTED_DB_SCHEMES:
+ return None
+ return scheme or _MISSING_SCHEME
+
+
+def unsupported_db_scheme_message(env_var: str, scheme: str) -> str:
+ """Operator-facing message naming the offending env var and scheme."""
+ return (
+ f"{env_var} uses unsupported scheme '{scheme}'. LiteLLM's database "
+ "features (virtual keys, store_model_in_db, spend tracking) require "
+ "PostgreSQL; use a 'postgresql://' connection string. SQLite and other "
+ "engines are not supported. "
+ "See https://docs.litellm.ai/docs/proxy/virtual_keys"
+ )
+
class DatabaseURLSettings(BaseSettings):
"""Discrete ``DATABASE_*`` env vars, loaded once at process start.
@@ -58,46 +93,47 @@ class DatabaseURLSettings(BaseSettings):
iam_token_db_auth: bool = Field(default=False, validation_alias=_IAM_ENV_KEY)
# Writer
- database_url: Optional[str] = Field(default=None, validation_alias="DATABASE_URL")
- database_host: Optional[str] = Field(default=None, validation_alias="DATABASE_HOST")
+ database_url: str | None = Field(default=None, validation_alias="DATABASE_URL")
+ direct_url: str | None = Field(default=None, validation_alias="DIRECT_URL")
+ database_host: str | None = Field(default=None, validation_alias="DATABASE_HOST")
database_port: str = Field(
default=_DEFAULT_PG_PORT, validation_alias="DATABASE_PORT"
)
- database_user: Optional[str] = Field(
+ database_user: str | None = Field(
default=None,
validation_alias=AliasChoices("DATABASE_USER", "DATABASE_USERNAME"),
)
- database_name: Optional[str] = Field(default=None, validation_alias="DATABASE_NAME")
- database_schema: Optional[str] = Field(
+ database_name: str | None = Field(default=None, validation_alias="DATABASE_NAME")
+ database_schema: str | None = Field(
default=None, validation_alias="DATABASE_SCHEMA"
)
- database_password: Optional[str] = Field(
+ database_password: str | None = Field(
default=None, validation_alias="DATABASE_PASSWORD"
)
# Read replica
- database_url_read_replica: Optional[str] = Field(
+ database_url_read_replica: str | None = Field(
default=None, validation_alias="DATABASE_URL_READ_REPLICA"
)
- database_host_read_replica: Optional[str] = Field(
+ database_host_read_replica: str | None = Field(
default=None, validation_alias="DATABASE_HOST_READ_REPLICA"
)
- database_port_read_replica: Optional[str] = Field(
+ database_port_read_replica: str | None = Field(
default=None, validation_alias="DATABASE_PORT_READ_REPLICA"
)
- database_user_read_replica: Optional[str] = Field(
+ database_user_read_replica: str | None = Field(
default=None,
validation_alias=AliasChoices(
"DATABASE_USER_READ_REPLICA", "DATABASE_USERNAME_READ_REPLICA"
),
)
- database_name_read_replica: Optional[str] = Field(
+ database_name_read_replica: str | None = Field(
default=None, validation_alias="DATABASE_NAME_READ_REPLICA"
)
- database_schema_read_replica: Optional[str] = Field(
+ database_schema_read_replica: str | None = Field(
default=None, validation_alias="DATABASE_SCHEMA_READ_REPLICA"
)
- database_password_read_replica: Optional[str] = Field(
+ database_password_read_replica: str | None = Field(
default=None, validation_alias="DATABASE_PASSWORD_READ_REPLICA"
)
@@ -106,7 +142,7 @@ class DatabaseURLSettings(BaseSettings):
"""Load the settings from ``os.environ`` (read at call time)."""
return cls()
- def build_writer_url(self) -> Optional[str]:
+ def build_writer_url(self) -> str | None:
"""Return the writer URL to set, or ``None`` to leave it as-is.
Raises ``RuntimeError`` (naming the offending vars) when IAM auth is
@@ -156,7 +192,7 @@ class DatabaseURLSettings(BaseSettings):
)
return None
- def build_reader_url(self) -> Optional[str]:
+ def build_reader_url(self) -> str | None:
"""Return the read-replica URL to set, or ``None`` to leave it as-is.
Opt-in via ``DATABASE_HOST_READ_REPLICA``; never clobbers a
@@ -217,11 +253,11 @@ class DatabaseURLSettings(BaseSettings):
def _password_url(
*,
user: str,
- password: Optional[str],
+ password: str | None,
host: str,
port: str,
name: str,
- schema: Optional[str],
+ schema: str | None,
) -> str:
"""Percent-encode credentials into a ``postgresql://`` URL.
@@ -239,6 +275,26 @@ class DatabaseURLSettings(BaseSettings):
url += f"?schema={schema}"
return url
+ def _raise_for_unsupported_scheme(self) -> None:
+ """Reject an operator-pinned non-PostgreSQL writer / direct / reader URL.
+
+ The componentized entrypoints (gateway / backend / migrations) call
+ ``apply_to_env`` and then hand the URL straight to Prisma, bypassing
+ the CLI's own guard. A pinned URL flows through untouched, so validate
+ the same three vars the CLI guard checks (DATABASE_URL, DIRECT_URL, and
+ the read replica) rather than letting Prisma stall on an unusable scheme.
+ """
+ for env_var, url in (
+ ("DATABASE_URL", self.database_url),
+ ("DIRECT_URL", self.direct_url),
+ ("DATABASE_URL_READ_REPLICA", self.database_url_read_replica),
+ ):
+ if not url:
+ continue
+ bad_scheme = unsupported_db_scheme(url)
+ if bad_scheme is not None:
+ raise RuntimeError(unsupported_db_scheme_message(env_var, bad_scheme))
+
def apply_to_env(self) -> bool:
"""Write the assembled URL(s) into ``os.environ``.
@@ -246,6 +302,7 @@ class DatabaseURLSettings(BaseSettings):
password auth that assembled a fresh URL). False means there was
nothing to do — an operator-pinned URL, or no discrete fields.
"""
+ self._raise_for_unsupported_scheme()
wrote_writer = False
writer_url = self.build_writer_url()
if writer_url is not None:
diff --git a/litellm/proxy/hooks/mcp_semantic_filter/hook.py b/litellm/proxy/hooks/mcp_semantic_filter/hook.py
index 6343faaa965..9888baf897e 100644
--- a/litellm/proxy/hooks/mcp_semantic_filter/hook.py
+++ b/litellm/proxy/hooks/mcp_semantic_filter/hook.py
@@ -123,11 +123,70 @@ class SemanticToolFilterHook(CustomLogger):
return openai_tools_as_dicts
+ def _is_mcp_tool(self, tool: object) -> bool:
+ """
+ Check whether *tool* is registered in the MCP semantic router.
+
+ Classification strategy (shape-first, lookup-second):
+ 1. Chat Completions format dicts are always native.
+ 2. Responses API function tools are always native.
+ 3. Everything else is looked up by name in the MCP registry.
+ """
+ if (
+ isinstance(tool, dict)
+ and tool.get("type") == "function"
+ and isinstance(tool.get("function"), dict)
+ ):
+ 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)
+ return bool(name) and name in self.filter._tool_map
+
def _get_metadata_variable_name(self, data: dict) -> str:
if "litellm_metadata" in data:
return "litellm_metadata"
return "metadata"
+ def _emit_filter_metadata(
+ self,
+ data: dict,
+ mcp_tools: list[object],
+ filtered_mcp_tools: list[object],
+ native_tools: list[object],
+ filtered_tools: list[object],
+ ) -> None:
+ """
+ Emit response-header metadata when MCP tools were filtered.
+
+ Stats report MCP-only counts so downstream consumers see accurate
+ semantic filter metrics. Skips metadata entirely for purely-native
+ requests to avoid spurious headers.
+ """
+ if mcp_tools:
+ filter_stats = f"{len(mcp_tools)}->{len(filtered_mcp_tools)}"
+ tool_names_csv = self._get_tool_names_csv(filtered_mcp_tools)
+
+ _metadata_variable_name = self._get_metadata_variable_name(data)
+ metadata = data.setdefault(_metadata_variable_name, {})
+ metadata["litellm_semantic_filter_stats"] = filter_stats
+ metadata["litellm_semantic_filter_tools"] = tool_names_csv
+
+ verbose_proxy_logger.info(
+ f"Semantic tool filter: {filter_stats} MCP tools "
+ f"({len(native_tools)} native preserved, "
+ f"{len(filtered_tools)} total)"
+ )
+ else:
+ verbose_proxy_logger.info(
+ f"Semantic tool filter: all {len(native_tools)} tools "
+ f"are native, no MCP filtering applied"
+ )
+
async def async_pre_call_hook(
self,
user_api_key_dict: "UserAPIKeyAuth",
@@ -140,53 +199,55 @@ class SemanticToolFilterHook(CustomLogger):
This hook is called before the LLM request is made. It filters the
tools list to only include semantically relevant tools.
-
- Args:
- user_api_key_dict: User authentication
- cache: Cache instance
- data: Request data containing messages and tools
- call_type: Type of call (completion, acompletion, etc.)
-
- Returns:
- Modified data dict with filtered tools, or None if no changes
"""
- # Only filter endpoints that support tools
if call_type not in ("completion", "acompletion", "aresponses"):
verbose_proxy_logger.debug(
f"Skipping semantic filter for call_type={call_type}"
)
return None
- # Check if tools are present
tools = data.get("tools")
if not tools:
verbose_proxy_logger.debug("No tools in request, skipping semantic filter")
return None
- original_tool_count = len(tools)
-
- # Check for MCP references (server_url="litellm_proxy") and expand them
+ # Expanded MCP tools are in OpenAI nested format which
+ # filter_tools/_extract_tool_info cannot name-match, so we skip
+ # semantic filtering and return early.
if self._should_expand_mcp_tools(tools):
verbose_proxy_logger.debug(
"Detected litellm_proxy MCP references, expanding before semantic filtering"
)
try:
+ native_tools_before_expand = [
+ t
+ for t in tools
+ if not (isinstance(t, dict) and t.get("type") == "mcp")
+ ]
+
expanded_tools = await self._expand_mcp_tools(tools, user_api_key_dict)
if not expanded_tools:
+ if native_tools_before_expand:
+ data["tools"] = native_tools_before_expand
+ verbose_proxy_logger.warning(
+ "No MCP tools expanded, preserving "
+ f"{len(native_tools_before_expand)} native tools"
+ )
+ return data
verbose_proxy_logger.warning(
"No tools expanded from MCP references"
)
return None
+ data["tools"] = native_tools_before_expand + expanded_tools
verbose_proxy_logger.info(
- f"Expanded {len(tools)} MCP reference(s) to {len(expanded_tools)} tools"
+ f"Expanded MCP references to {len(expanded_tools)} tools "
+ f"({len(native_tools_before_expand)} native preserved), "
+ f"skipping semantic filter (OpenAI nested format)"
)
-
- # Update tools for filtering
- tools = expanded_tools
- original_tool_count = len(tools)
+ return data
except Exception as e:
verbose_proxy_logger.error(
@@ -194,7 +255,6 @@ class SemanticToolFilterHook(CustomLogger):
)
return None
- # Check if messages are present (try both "messages" and "input" for responses API)
messages = data.get("messages", [])
if not messages:
messages = data.get("input", [])
@@ -204,13 +264,11 @@ class SemanticToolFilterHook(CustomLogger):
)
return None
- # Check if filter is enabled
if not self.filter.enabled:
verbose_proxy_logger.debug("Semantic filter disabled, skipping")
return None
try:
- # Extract user query from messages
user_query = self.filter.extract_user_query(messages)
if not user_query:
verbose_proxy_logger.debug(
@@ -218,33 +276,60 @@ class SemanticToolFilterHook(CustomLogger):
)
return None
+ native_tools: list[object] = []
+ mcp_tools: list[object] = []
+ mcp_indices: set[int] = set()
+ for i, t in enumerate(tools):
+ if self._is_mcp_tool(t):
+ mcp_tools.append(t)
+ mcp_indices.add(i)
+ else:
+ native_tools.append(t)
+
verbose_proxy_logger.debug(
- f"Applying semantic filter to {len(tools)} tools "
- f"with query: '{user_query[:50]}...'"
+ f"Applying semantic filter: {len(mcp_tools)} MCP tools, "
+ f"{len(native_tools)} native tools, "
+ f"query: '{user_query[:50]}...'"
)
- # Filter tools semantically
- filtered_tools = await self.filter.filter_tools(
- query=user_query,
- available_tools=tools, # type: ignore
- )
+ if mcp_tools:
+ filtered_mcp_tools = await self.filter.filter_tools(
+ query=user_query,
+ available_tools=mcp_tools, # type: ignore
+ )
+ else:
+ filtered_mcp_tools = []
+
+ filtered_mcp_names: set[str] = set()
+ for t in filtered_mcp_tools:
+ name, _ = self.filter._extract_tool_info(t)
+ if name:
+ filtered_mcp_names.add(name)
+
+ filtered_tools: list[object] = []
+ for i, t in enumerate(tools):
+ if i in mcp_indices:
+ name, _ = self.filter._extract_tool_info(t)
+ if name in filtered_mcp_names:
+ filtered_tools.append(t)
+ else:
+ filtered_tools.append(t)
- # Always update tools and emit header (even if count unchanged)
data["tools"] = filtered_tools
- # Store filter stats and tool names for response header
- filter_stats = f"{original_tool_count}->{len(filtered_tools)}"
- tool_names_csv = self._get_tool_names_csv(filtered_tools)
-
- _metadata_variable_name = self._get_metadata_variable_name(data)
- data[_metadata_variable_name][
- "litellm_semantic_filter_stats"
- ] = filter_stats
- data[_metadata_variable_name][
- "litellm_semantic_filter_tools"
- ] = tool_names_csv
-
- verbose_proxy_logger.info(f"Semantic tool filter: {filter_stats} tools")
+ try:
+ self._emit_filter_metadata(
+ data=data,
+ mcp_tools=mcp_tools,
+ filtered_mcp_tools=filtered_mcp_tools,
+ native_tools=native_tools,
+ filtered_tools=filtered_tools,
+ )
+ except Exception as e:
+ verbose_proxy_logger.warning(
+ f"Failed to emit semantic filter metadata: {e}",
+ exc_info=True,
+ )
return data
@@ -266,7 +351,7 @@ class SemanticToolFilterHook(CustomLogger):
from litellm.constants import MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH
_metadata_variable_name = self._get_metadata_variable_name(data)
- metadata = data[_metadata_variable_name]
+ metadata = data.get(_metadata_variable_name, {})
filter_stats = metadata.get("litellm_semantic_filter_stats")
if not filter_stats:
diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py
index e86982307e7..f896047a219 100644
--- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py
+++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py
@@ -1016,15 +1016,10 @@ if MCP_AVAILABLE:
if is_restricted_virtual_key:
return _sanitize_mcp_server_list_for_virtual_key(redacted_mcp_servers)
- # Non-admin authenticated users may see the server inventory but
- # not credential-bearing fields like `url` (often contains bearer
- # tokens) or headers/env (often contain Authorization).
- if not _user_has_admin_view(user_api_key_dict):
- return _sanitize_mcp_server_list_for_non_admin(redacted_mcp_servers)
-
+ # only a full PROXY_ADMIN sees credential-bearing fields; everyone else
+ # goes through the non-admin sanitizer
if not _user_is_full_admin(user_api_key_dict):
- for server in redacted_mcp_servers:
- _redact_global_env_var_values(server)
+ return _sanitize_mcp_server_list_for_non_admin(redacted_mcp_servers)
return redacted_mcp_servers
@@ -1415,10 +1410,10 @@ if MCP_AVAILABLE:
redacted = _redact_mcp_credentials(mcp_server)
if is_restricted_virtual_key:
return _sanitize_mcp_server_for_virtual_key(redacted)
- if not _user_has_admin_view(user_api_key_dict):
- return _sanitize_mcp_server_for_non_admin(redacted)
+ # only a full PROXY_ADMIN sees credential-bearing fields; everyone else
+ # goes through the non-admin sanitizer
if not _user_is_full_admin(user_api_key_dict):
- _redact_global_env_var_values(redacted)
+ return _sanitize_mcp_server_for_non_admin(redacted)
return redacted
@router.post(
@@ -1935,12 +1930,9 @@ if MCP_AVAILABLE:
prisma_client = get_prisma_client_or_throw(
"Database not connected. Connect a database to your proxy"
)
- mcp_server = await get_mcp_server(prisma_client, server_id)
- if mcp_server is None:
- raise HTTPException(
- status_code=status.HTTP_404_NOT_FOUND,
- detail={"error": f"MCP Server {server_id} not found"},
- )
+ mcp_server = await _authorize_and_fetch_mcp_server(
+ prisma_client, user_api_key_dict, server_id
+ )
if not getattr(mcp_server, "is_byok", False):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
@@ -2015,12 +2007,9 @@ if MCP_AVAILABLE:
prisma_client = get_prisma_client_or_throw(
"Database not connected. Connect a database to your proxy"
)
- mcp_server = await get_mcp_server(prisma_client, server_id)
- if mcp_server is None:
- raise HTTPException(
- status_code=status.HTTP_404_NOT_FOUND,
- detail={"error": f"MCP Server {server_id} not found"},
- )
+ await _authorize_and_fetch_mcp_server(
+ prisma_client, user_api_key_dict, server_id
+ )
user_id = user_api_key_dict.user_id or ""
if not user_id:
raise HTTPException(
@@ -2182,37 +2171,47 @@ if MCP_AVAILABLE:
user_api_key_dict: UserAPIKeyAuth,
server_id: str,
) -> LiteLLM_MCPServerTable:
- """Return the MCP server the caller may manage env vars for.
+ """Resolve the MCP server a caller may manage their own per-user state for.
- Admins look the server up directly. Non-admins reuse the access-scoped
- listing that already loads every server they can see, so we don't issue
- a second per-server query just to re-fetch a record the authorization
- check produced. A non-admin who can't see the server gets 403 (never
- 404) so server ids can't be enumerated.
+ Looks the server up in the DB, then the in-memory registry, so a
+ config-defined server (which never gets a DB row) resolves too. Admins
+ may reach any server and get a 404 for an unknown id. A non-admin may
+ only reach a server in their allowed set and otherwise gets 403 (never
+ 404, so server ids can't be enumerated), using the same allowed-server
+ resolution the MCP gateway enforces on tool calls.
"""
+ server = await get_mcp_server(prisma_client, server_id)
+ if server is None:
+ registry_server = global_mcp_server_manager.get_mcp_server_by_id(server_id)
+ if registry_server is not None:
+ server = global_mcp_server_manager._build_mcp_server_table(
+ registry_server
+ )
+
if _user_has_admin_view(user_api_key_dict):
- server = await get_mcp_server(prisma_client, server_id)
if server is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail={"error": f"MCP Server {server_id} not found"},
)
return server
- accessible = await get_all_mcp_servers_for_user(
- prisma_client, user_api_key_dict
- )
- for server in accessible:
- if server.server_id == server_id:
- return server
- raise HTTPException(
- status_code=status.HTTP_403_FORBIDDEN,
- detail={
- "error": (
- f"User does not have permission to access mcp server with id {server_id}. "
- "You can only manage env vars for mcp servers that you have access to."
- )
- },
- )
+
+ allowed_server_ids: set[str] = set()
+ for auth_context in await build_effective_auth_contexts(user_api_key_dict):
+ allowed_server_ids.update(
+ await global_mcp_server_manager.get_allowed_mcp_servers(auth_context)
+ )
+ if server is None or server.server_id not in allowed_server_ids:
+ raise HTTPException(
+ status_code=status.HTTP_403_FORBIDDEN,
+ detail={
+ "error": (
+ f"User does not have permission to access mcp server with id {server_id}. "
+ "You can only manage mcp servers that you have access to."
+ )
+ },
+ )
+ return server
def _compute_user_env_var_status(
*,
diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py
index 9c4d7b1bb5d..d0281885482 100644
--- a/litellm/proxy/proxy_cli.py
+++ b/litellm/proxy/proxy_cli.py
@@ -1195,6 +1195,25 @@ def run_server(
os.getenv("DATABASE_URL", None) is not None
or os.getenv("DIRECT_URL", None) is not None
):
+ from litellm.proxy.db.db_url_settings import (
+ unsupported_db_scheme,
+ unsupported_db_scheme_message,
+ )
+
+ for _db_env in ("DATABASE_URL", "DIRECT_URL"):
+ _candidate_url = os.getenv(_db_env)
+ if _candidate_url is None:
+ continue
+ _bad_scheme = unsupported_db_scheme(_candidate_url)
+ if _bad_scheme is not None:
+ print(
+ f"\033[1;31mLiteLLM Proxy: "
+ f"{unsupported_db_scheme_message(_db_env, _bad_scheme)}"
+ "\033[0m",
+ file=sys.stderr,
+ flush=True,
+ )
+ sys.exit(1)
try:
from litellm.secret_managers.main import get_secret
diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py
index c5676fbb32f..4c36b42615e 100644
--- a/litellm/proxy/proxy_server.py
+++ b/litellm/proxy/proxy_server.py
@@ -38,7 +38,7 @@ from typing import (
import anyio
import websockets
import websockets.exceptions
-from pydantic import BaseModel, Json
+from pydantic import BaseModel, Json, JsonValue
from litellm._uuid import uuid
from litellm.constants import (
@@ -106,6 +106,10 @@ from litellm.proxy.common_utils.callback_utils import (
process_callback,
)
from litellm.proxy.common_utils.realtime_utils import _realtime_request_body
+from litellm.router_utils.add_retry_fallback_headers import (
+ get_fallback_errors_from_headers,
+ get_hidden_params_dict,
+)
from litellm.types.utils import (
ModelResponse,
ModelResponseStream,
@@ -4459,6 +4463,8 @@ class ProxyConfig:
f"{blue_color_code} setting litellm.{key}={value}{reset_color_code}"
)
setattr(litellm, key, value)
+ if key == "request_timeout":
+ litellm.request_timeout_explicitly_set = True
if key in {"s3_audit_callback_params", "s3_callback_params"}:
from litellm.integrations.s3_v2 import S3Logger as S3V2Logger
from litellm.litellm_core_utils.litellm_logging import (
@@ -7085,57 +7091,122 @@ def _get_client_requested_model_for_streaming(request_data: dict) -> str:
return requested_model if isinstance(requested_model, str) else ""
+def _is_positive_int_like(value: Any) -> bool:
+ try:
+ return int(value) > 0
+ except (TypeError, ValueError):
+ return False
+
+
+def _should_include_fallback_errors(request_data: dict[str, object]) -> bool:
+ if not general_settings.get("expose_fallback_errors_to_caller"):
+ return False
+ return request_data.get("include_fallback_errors") is True
+
+
+def _get_streaming_fallback_metadata(
+ response_obj: object,
+) -> tuple[bool, str | None, list[dict[str, object]]]:
+ additional_headers = get_hidden_params_dict(response_obj).get("additional_headers")
+ if not isinstance(additional_headers, dict):
+ return False, None, []
+
+ if not _is_positive_int_like(
+ additional_headers.get("x-litellm-attempted-fallbacks")
+ ):
+ return False, None, []
+
+ fallback_model = additional_headers.get("x-litellm-model-group")
+ fallback_errors = get_fallback_errors_from_headers(additional_headers)
+ if isinstance(fallback_model, str) and fallback_model:
+ return True, fallback_model, fallback_errors
+ return True, None, fallback_errors
+
+
+def _format_fallback_metadata_sse_event(
+ *,
+ fallback_model: str | None,
+ fallback_errors: list[dict[str, object]],
+) -> str:
+ import time
+
+ payload = {
+ "id": "litellm-fallback-metadata",
+ "object": "chat.completion.chunk",
+ "created": int(time.time()),
+ "model": fallback_model or "",
+ "choices": [],
+ "litellm_fallback": {
+ "fallback_model": fallback_model,
+ "errors": fallback_errors,
+ },
+ }
+ return f"data: {json.dumps(payload)}\n\n"
+
+
def _restamp_streaming_chunk_model(
*,
chunk: Any,
requested_model_from_client: str,
request_data: dict,
model_mismatch_logged: bool,
-) -> Tuple[Any, bool]:
+ fallback_was_attempted: bool = False,
+ fallback_model_from_metadata: str | None = None,
+) -> tuple[Any, bool]:
+ target_model = (
+ fallback_model_from_metadata
+ if fallback_was_attempted
+ else requested_model_from_client
+ )
# Always return the client-requested model name (not provider-prefixed internal identifiers)
# on streaming chunks.
+ # On fallback, use the public OpenAI-compatible model name. This keeps
+ # provider-prefixed internal identifiers from leaking into the public API.
#
# Note: This warning is intentionally verbose. A mismatch is a useful signal that an
# internal provider/deployment identifier is leaking into the public API, and helps
# maintainers/operators catch regressions while preserving OpenAI-compatible output.
- if not requested_model_from_client or not isinstance(chunk, (BaseModel, dict)):
+ if not target_model or not isinstance(chunk, (BaseModel, dict)):
return chunk, model_mismatch_logged
# For Azure Model Router, preserve the actual model used in each chunk
- if _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
# instead of stamping the comma-separated list the client sent.
- if request_data.get("fastest_response", False):
+ if not fallback_was_attempted and request_data.get("fastest_response", False):
return chunk, model_mismatch_logged
downstream_model = (
chunk.get("model") if isinstance(chunk, dict) else getattr(chunk, "model", None)
)
- if downstream_model == requested_model_from_client:
+ if downstream_model == target_model:
return chunk, model_mismatch_logged
- if not model_mismatch_logged and downstream_model != requested_model_from_client:
+ if not model_mismatch_logged and downstream_model != target_model:
verbose_proxy_logger.debug(
- "litellm_call_id=%s: streaming chunk model mismatch - requested=%r downstream=%r. Overriding model to requested.",
+ "litellm_call_id=%s: streaming chunk model mismatch - target=%r downstream=%r fallback_was_attempted=%s. Overriding chunk model to target.",
request_data.get("litellm_call_id"),
- requested_model_from_client,
+ target_model,
downstream_model,
+ fallback_was_attempted,
)
model_mismatch_logged = True
if isinstance(chunk, dict):
- chunk["model"] = requested_model_from_client
+ chunk["model"] = target_model
return chunk, model_mismatch_logged
try:
- setattr(chunk, "model", requested_model_from_client)
+ chunk.model = target_model
except Exception as e:
verbose_proxy_logger.error(
"litellm_call_id=%s: failed to override chunk.model=%r on chunk_type=%s. error=%s",
request_data.get("litellm_call_id"),
- requested_model_from_client,
+ target_model,
type(chunk),
str(e),
exc_info=True,
@@ -7294,7 +7365,14 @@ async def async_data_generator(
requested_model_from_client = _get_client_requested_model_for_streaming(
request_data=request_data
)
+ (
+ fallback_was_attempted,
+ fallback_model_from_metadata,
+ fallback_errors,
+ ) = _get_streaming_fallback_metadata(response)
model_mismatch_logged = False
+ fallback_metadata_event_sent = False
+ include_fallback_errors = _should_include_fallback_errors(request_data)
# Use a running string instead of list + join to avoid O(n^2) overhead.
# Previously "".join(str_so_far_parts) was called every chunk, re-joining
# the entire accumulated response. String += is O(n) amortized total.
@@ -7332,13 +7410,37 @@ async def async_data_generator(
str_so_far=_str_so_far,
)
+ # Mid-stream fallbacks surface metadata on individual chunks rather than
+ # the response wrapper. Keep scanning chunks until a fallback model is
+ # resolved, then latch it for the rest of the stream.
+ if fallback_model_from_metadata is None:
+ (
+ chunk_fallback_was_attempted,
+ chunk_fallback_model,
+ chunk_fallback_errors,
+ ) = _get_streaming_fallback_metadata(chunk)
+ if chunk_fallback_was_attempted:
+ fallback_was_attempted = True
+ fallback_model_from_metadata = chunk_fallback_model
+ fallback_errors = fallback_errors or chunk_fallback_errors
+
+ pending_fallback_event = (
+ include_fallback_errors
+ and fallback_was_attempted
+ and fallback_errors
+ and not fallback_metadata_event_sent
+ )
+
chunk, model_mismatch_logged = _restamp_streaming_chunk_model(
chunk=chunk,
requested_model_from_client=requested_model_from_client,
request_data=request_data,
model_mismatch_logged=model_mismatch_logged,
+ fallback_was_attempted=fallback_was_attempted,
+ fallback_model_from_metadata=fallback_model_from_metadata,
)
+ raw_passthrough = False
if isinstance(chunk, BaseModel):
chunk = _serialize_streaming_chunk(chunk)
elif isinstance(chunk, bytes):
@@ -7354,14 +7456,14 @@ async def async_data_generator(
raise ValueError(
"Raw SSE stream exceeded maximum buffered size without a frame delimiter"
)
- continue
- if chunk.startswith(("data:", "event:", ":")):
+ raw_passthrough = True
+ elif chunk.startswith(("data:", "event:", ":")):
yield (
chunk
if chunk.endswith(_SSE_FRAME_DELIMITERS)
else chunk + "\n\n"
)
- continue
+ raw_passthrough = True
elif isinstance(chunk, str) and is_raw_sse_stream:
raw_sse_buffer += chunk
while True:
@@ -7373,15 +7475,23 @@ async def async_data_generator(
raise ValueError(
"Raw SSE stream exceeded maximum buffered size without a frame delimiter"
)
- continue
+ raw_passthrough = True
elif isinstance(chunk, str) and chunk.startswith("data: "):
error_message = chunk
break
- try:
- yield _format_streaming_sse_chunk(chunk=chunk)
- except Exception as e:
- yield f"data: {str(e)}\n\n"
+ if not raw_passthrough:
+ try:
+ yield _format_streaming_sse_chunk(chunk=chunk)
+ except Exception as e:
+ yield f"data: {str(e)}\n\n"
+
+ if pending_fallback_event:
+ yield _format_fallback_metadata_sse_event(
+ fallback_model=fallback_model_from_metadata,
+ fallback_errors=fallback_errors,
+ )
+ fallback_metadata_event_sent = True
stream_completed = True
if not needs_iterator_wrap:
@@ -11390,18 +11500,20 @@ def get_direct_access_models(
llm_router: Router,
) -> List[str]:
"""
- Get all models that user has direct access to
- """
+ Get all models that user has direct access to.
- direct_access_models: List[str] = []
- for model in user_db_object.models:
- deployments = llm_router.get_model_list(model_name=model)
- if deployments is not None:
- for deployment in deployments:
- model_id = deployment.get("model_info", {}).get("id", None)
- if model_id is not None:
- direct_access_models.append(model_id)
- return direct_access_models
+ The 'all-proxy-models' sentinel grants direct access to every non-team
+ deployment, mirroring how get_key_models expands it for the key/team path.
+ """
+ if SpecialModelNames.all_proxy_models.value in user_db_object.models:
+ return llm_router.get_model_ids(exclude_team_models=True)
+
+ return [
+ model_id
+ for model in user_db_object.models
+ for deployment in (llm_router.get_model_list(model_name=model) or [])
+ if (model_id := deployment.get("model_info", {}).get("id", None)) is not None
+ ]
def _filter_models_to_user_accessible(all_models: List[Dict]) -> List[Dict]:
@@ -15073,6 +15185,68 @@ async def update_config_general_settings(
return response
+# Secret-bearing general_settings fields the segment masker does not match by
+# name: database_url and database_extra_connection_params embed DB credentials,
+# pass_through_endpoints carry upstream Authorization headers, and
+# alert_to_webhook_url is itself a webhook secret
+_EXTRA_SECRET_GENERAL_SETTINGS_FIELDS = frozenset(
+ {
+ "database_url",
+ "database_extra_connection_params",
+ "pass_through_endpoints",
+ "alert_to_webhook_url",
+ }
+)
+
+
+def _is_secret_general_setting_field(field_name: str) -> bool:
+ return (
+ field_name in _EXTRA_SECRET_GENERAL_SETTINGS_FIELDS
+ or SENSITIVE_DATA_MASKER.is_sensitive_key(field_name)
+ )
+
+
+# Matches the cap on _redact_sensitive_litellm_params (the closest analog in the
+# proxy). Past this depth we fail closed by returning "REDACTED" for the whole
+# subtree rather than recursing further — better to over-redact a pathological
+# config than to silently return a deeply-nested credential verbatim
+_REDACT_SECRET_MAX_DEPTH = 10
+
+
+def _redact_secret_values_in_obj(value: JsonValue, depth: int = 0) -> JsonValue:
+ """Recursively redact secret leaves inside a structured field so a nested
+ credential (e.g. aws_web_identity_token under database_args) is never
+ returned to a non-admin, while non-secret siblings stay visible. At
+ _REDACT_SECRET_MAX_DEPTH the whole subtree is replaced with "REDACTED"
+ so depth-overrun fails closed."""
+ if depth >= _REDACT_SECRET_MAX_DEPTH:
+ return "REDACTED"
+ if isinstance(value, dict):
+ return {
+ key: (
+ "REDACTED"
+ if _is_secret_general_setting_field(key)
+ else _redact_secret_values_in_obj(sub, depth + 1)
+ )
+ for key, sub in value.items()
+ }
+ if isinstance(value, list):
+ return [_redact_secret_values_in_obj(item, depth + 1) for item in value]
+ return value
+
+
+def _redact_general_setting_value(
+ field_name: str, value: JsonValue, is_full_admin: bool
+) -> JsonValue:
+ if is_full_admin:
+ return value
+ if _is_secret_general_setting_field(field_name):
+ return "REDACTED"
+ if isinstance(value, (dict, list)):
+ return _redact_secret_values_in_obj(value)
+ return value
+
+
@router.get(
"/config/field/info",
tags=["config.yaml"],
@@ -15125,9 +15299,11 @@ async def get_config_general_settings(
general_settings = dict(db_general_settings.param_value)
if field_name in general_settings:
- field_value = general_settings[field_name]
- # Redact plugin_key from plugin configs so the shared credential
- # is never returned even to admin-viewer callers.
+ field_value = _redact_general_setting_value(
+ field_name,
+ general_settings[field_name],
+ user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN,
+ )
if field_name == "plugins" and isinstance(field_value, list):
field_value = [
(
@@ -15183,6 +15359,8 @@ async def get_config_list(
},
)
+ is_full_admin = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
+
## get general settings from db
db_general_settings = await ConfigRepository(prisma_client).table.find_first(
where={"param_name": "general_settings"}
@@ -15231,7 +15409,11 @@ async def get_config_list(
field_name=sub_field,
field_type=sub_field_type.__name__,
field_description="", # Add custom logic if descriptions are available
- field_default_value=general_settings.get(sub_field, None),
+ field_default_value=_redact_general_setting_value(
+ sub_field,
+ general_settings.get(sub_field, None),
+ is_full_admin,
+ ),
stored_in_db=None,
)
for sub_field, sub_field_type in pydantic_class.__annotations__.items()
@@ -15261,7 +15443,11 @@ async def get_config_list(
field_name=field_name,
field_type=allowed_args[field_name]["type"],
field_description=field_info.description or "",
- field_value=general_settings.get(field_name, None),
+ field_value=_redact_general_setting_value(
+ field_name,
+ general_settings.get(field_name, None),
+ is_full_admin,
+ ),
stored_in_db=_stored_in_db,
field_default_value=field_info.default,
nested_fields=nested_fields,
@@ -15285,7 +15471,9 @@ async def get_config_list(
field_name=field_name,
field_type=allowed_args[field_name]["type"],
field_description=field_info.description or "",
- field_value=_field_value,
+ field_value=_redact_general_setting_value(
+ field_name, _field_value, is_full_admin
+ ),
stored_in_db=_stored_in_db,
field_default_value=field_info.default,
nested_fields=nested_fields,
diff --git a/litellm/proxy/read_model_list.py b/litellm/proxy/read_model_list.py
new file mode 100644
index 00000000000..2dff8eaf698
--- /dev/null
+++ b/litellm/proxy/read_model_list.py
@@ -0,0 +1,28 @@
+"""Resolve a proxy config's ``model_list`` for the Rust AI gateway.
+
+The Rust gateway calls this once at load time (via an embedded interpreter) and
+builds its own (Rust) router from the returned ``model_list``. We do NOT call
+``ProxyConfig.load_config`` here: that returns a *Python* ``litellm.Router`` (not
+usable from Rust) and boots the whole proxy (callbacks, cache, DB, auth) as side
+effects.
+
+Instead we reuse ``ProxyConfig.get_config`` — the actual config reader — so the
+gateway inherits the same heavy lifting the proxy does: ``include:`` merging,
+``os.environ/`` + secret-manager resolution, and DB-stored models (when a DB is
+configured). It has no proxy-setup side effects. Returns the resolved
+``model_list``; the Rust side deserializes each entry into its ``Deployment``.
+"""
+
+from __future__ import annotations
+
+import asyncio
+from typing import Any
+
+
+def read_model_list(config_path: str) -> list[dict[str, Any]]:
+ """Load ``config_path`` via the proxy's own reader and return its
+ resolved ``model_list``."""
+ from litellm.proxy.proxy_server import ProxyConfig
+
+ config = asyncio.run(ProxyConfig().get_config(config_file_path=config_path))
+ return config.get("model_list") or []
diff --git a/litellm/responses/main.py b/litellm/responses/main.py
index 34c9cdd3d1c..2c46baaada5 100644
--- a/litellm/responses/main.py
+++ b/litellm/responses/main.py
@@ -58,7 +58,10 @@ from litellm.llms.openai.data_residency import infer_openai_data_residency
from litellm.secret_managers.main import get_secret_str
from litellm.types.responses.main import *
from litellm.types.router import GenericLiteLLMParams
-from litellm.utils import ProviderConfigManager, client
+from litellm.utils import (
+ ProviderConfigManager,
+ client,
+)
if TYPE_CHECKING:
from mcp.types import Tool as MCPTool
diff --git a/litellm/router.py b/litellm/router.py
index e54eadfb872..6e7b9689415 100644
--- a/litellm/router.py
+++ b/litellm/router.py
@@ -40,7 +40,6 @@ import anyio
import httpx
import openai
from openai import AsyncOpenAI
-from pydantic import BaseModel
from typing_extensions import overload
import litellm
@@ -62,6 +61,9 @@ from litellm.constants import (
)
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.asyncify import run_async_function
+from litellm.litellm_core_utils.request_timeout_resolver import (
+ get_configured_request_timeout,
+)
from litellm.litellm_core_utils.core_helpers import (
_get_parent_otel_span_from_kwargs,
get_metadata_variable_name_from_kwargs,
@@ -81,8 +83,10 @@ from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2
from litellm.router_strategy.simple_shuffle import simple_shuffle
from litellm.router_strategy.tag_based_routing import get_deployments_for_tag
from litellm.router_utils.add_retry_fallback_headers import (
+ _HiddenParamsHost,
add_fallback_headers_to_response,
add_retry_headers_to_response,
+ get_hidden_params_dict,
)
from litellm.router_utils.batch_utils import (
_get_router_metadata_variable_name,
@@ -564,6 +568,12 @@ class Router:
self._explicit_timeout = timeout # None when user did not pass timeout
self.timeout = timeout or litellm.request_timeout
+ # Per-attempt request_timeout, independent of router_settings.timeout.
+ # Only stored when a router timeout is also set, since otherwise
+ # request_timeout already flows through self.timeout above.
+ self.request_timeout = (
+ get_configured_request_timeout() if timeout is not None else None
+ )
self.stream_timeout = stream_timeout
self.retry_after = retry_after
@@ -2165,6 +2175,36 @@ class Router:
)
setattr(fallback_item, "usage", combined_usage)
+ @staticmethod
+ def _prepare_fallback_hidden_params(
+ fallback_response: object,
+ ) -> tuple[dict[str, object], dict[str, object]]:
+ fallback_hidden_params = get_hidden_params_dict(fallback_response)
+ fallback_headers = fallback_hidden_params.get("additional_headers")
+ if not isinstance(fallback_headers, dict):
+ return fallback_hidden_params, {}
+ return fallback_hidden_params, cast("dict[str, object]", fallback_headers)
+
+ @staticmethod
+ def _apply_fallback_hidden_params_to_item(
+ fallback_item: object,
+ prepared_fallback_hidden_params: tuple[dict[str, object], dict[str, object]],
+ ) -> None:
+ if fallback_item is None or not hasattr(fallback_item, "_hidden_params"):
+ return
+
+ fallback_hidden_params, fallback_headers = prepared_fallback_hidden_params
+ item_hidden_params = get_hidden_params_dict(fallback_item)
+ item_headers = item_hidden_params.get("additional_headers")
+ if not isinstance(item_headers, dict):
+ item_headers = {}
+
+ cast(_HiddenParamsHost, fallback_item)._hidden_params = {
+ **item_hidden_params,
+ **fallback_hidden_params,
+ "additional_headers": {**item_headers, **fallback_headers},
+ }
+
async def _acompletion_streaming_iterator(
self,
model_response: CustomStreamWrapper,
@@ -2257,12 +2297,22 @@ class Router:
model_group=model_group,
args=(),
kwargs=initial_kwargs,
+ include_fallback_errors=initial_kwargs.get(
+ "include_fallback_errors", False
+ )
+ is True,
)
)
# If fallback returns a streaming response, iterate over it
if hasattr(fallback_response, "__aiter__"):
+ prepared_fallback_hidden_params = (
+ Router._prepare_fallback_hidden_params(fallback_response)
+ )
async for fallback_item in fallback_response: # type: ignore
+ Router._apply_fallback_hidden_params_to_item(
+ fallback_item, prepared_fallback_hidden_params
+ )
if (
fallback_item
and isinstance(fallback_item, ModelResponseStream)
@@ -2686,11 +2736,21 @@ class Router:
model_group=model_group,
args=(),
kwargs=initial_kwargs,
+ include_fallback_errors=initial_kwargs.get(
+ "include_fallback_errors", False
+ )
+ is True,
)
)
if hasattr(fallback_response, "__aiter__"):
+ prepared_fallback_hidden_params = (
+ Router._prepare_fallback_hidden_params(fallback_response)
+ )
async for fallback_item in fallback_response: # type: ignore
+ Router._apply_fallback_hidden_params_to_item(
+ fallback_item, prepared_fallback_hidden_params
+ )
if partial_usage is not None:
Router._combine_responses_fallback_usage(
fallback_item, partial_usage
@@ -2815,7 +2875,13 @@ class Router:
)
if hasattr(fallback_response, "__iter__"):
+ prepared_fallback_hidden_params = (
+ Router._prepare_fallback_hidden_params(fallback_response)
+ )
for fallback_item in fallback_response:
+ Router._apply_fallback_hidden_params_to_item(
+ fallback_item, prepared_fallback_hidden_params
+ )
if (
fallback_item
and isinstance(fallback_item, ModelResponseStream)
@@ -2972,6 +3038,7 @@ class Router:
**kwargs,
}
input_kwargs.pop("silent_model", None)
+ input_kwargs.pop("include_fallback_errors", None)
_response = litellm.acompletion(**input_kwargs)
@@ -3076,7 +3143,18 @@ class Router:
- litellm_trace_id
- metadata
"""
- kwargs["num_retries"] = kwargs.get("num_retries", self.num_retries)
+ # Normalise an explicit num_retries=None to the router default here (dict.get()
+ # only falls back when the key is absent, not when its value is None), then to 0
+ # if the router default is itself None - mirroring the guard in
+ # async_function_with_retries, which remains the safety net for paths that bypass
+ # this setter.
+ _req_num_retries = kwargs.get("num_retries")
+ if _req_num_retries is not None:
+ kwargs["num_retries"] = _req_num_retries
+ else:
+ kwargs["num_retries"] = (
+ self.num_retries if self.num_retries is not None else 0
+ )
kwargs.setdefault("litellm_trace_id", str(uuid.uuid4()))
model_group_alias: Optional[str] = None
if self._get_model_from_alias(model=model):
@@ -3322,6 +3400,7 @@ class Router:
"stream_timeout", None
) # timeout set on litellm_params for this deployment
or self.stream_timeout # timeout set on router
+ or self.request_timeout # litellm_settings.request_timeout (per-attempt)
or self.default_litellm_params.get("stream_timeout", None)
)
@@ -3338,7 +3417,8 @@ class Router:
or data.get(
"request_timeout", None
) # timeout set on litellm_params for this deployment
- or self.timeout # timeout set on router
+ or self.request_timeout # litellm_settings.request_timeout (per-attempt)
+ or self.timeout # timeout set on router (router_settings.timeout)
or self.default_litellm_params.get("timeout", None)
)
return timeout
@@ -6478,6 +6558,7 @@ class Router:
model_group: Optional[str],
args: tuple,
kwargs: dict,
+ include_fallback_errors: bool = False,
):
"""
Common utilities for async_function_with_fallbacks
@@ -6501,6 +6582,8 @@ class Router:
input_kwargs["max_fallbacks"] = self.max_fallbacks
if "fallback_depth" not in input_kwargs:
input_kwargs["fallback_depth"] = 0
+ if include_fallback_errors:
+ input_kwargs["include_fallback_errors"] = True
# ORDER-BASED FALLBACKS: prepend higher order levels to the fallback list
# Skip for error types that have their own dedicated fallback handlers
@@ -6759,6 +6842,7 @@ class Router:
If it fails after num_retries, fall back to another model group
"""
model_group: Optional[str] = kwargs.get("model")
+ include_fallback_errors = kwargs.get("include_fallback_errors", False) is True
disable_fallbacks: Optional[bool] = kwargs.pop("disable_fallbacks", False)
fallbacks: Optional[List] = kwargs.get("fallbacks", self.fallbacks)
context_window_fallbacks: Optional[List] = kwargs.get(
@@ -6802,6 +6886,7 @@ class Router:
model_group,
args,
kwargs,
+ include_fallback_errors=include_fallback_errors,
)
def _handle_mock_testing_fallbacks(
@@ -6868,7 +6953,11 @@ class Router:
"model_group_retry_policy", self.model_group_retry_policy
)
model_group: Optional[str] = kwargs.get("model")
- num_retries = kwargs.pop("num_retries")
+ num_retries = kwargs.pop("num_retries", None)
+ if num_retries is None:
+ # Fall back to the router setting (then 0) so the comparisons below never
+ # hit `None > int`, which would mask the real upstream error with a TypeError.
+ num_retries = self.num_retries if self.num_retries is not None else 0
## ADD MODEL GROUP SIZE TO METADATA - used for model_group_rate_limit_error tracking
_metadata: dict = kwargs.get("litellm_metadata", kwargs.get("metadata")) or {}
@@ -9725,17 +9814,19 @@ class Router:
# - if healthy_deployments > 1, return model group rate limit headers
# - else return the model's rate limit headers
"""
- if (
- isinstance(response, BaseModel)
- and hasattr(response, "_hidden_params")
- and isinstance(response._hidden_params, dict) # type: ignore
- ):
- response._hidden_params.setdefault("additional_headers", {}) # type: ignore
- response._hidden_params["additional_headers"][ # type: ignore
- "x-litellm-model-group"
- ] = model_group
+ if response is not None and hasattr(response, "_hidden_params"):
+ hidden_params = getattr(response, "_hidden_params", {}) or {}
+ if hasattr(hidden_params, "model_dump"):
+ hidden_params = hidden_params.model_dump()
+ if not isinstance(hidden_params, dict):
+ return response
+ response._hidden_params = hidden_params
- additional_headers = response._hidden_params["additional_headers"] # type: ignore
+ additional_headers = hidden_params.get("additional_headers")
+ if not isinstance(additional_headers, dict):
+ additional_headers = {}
+ hidden_params["additional_headers"] = additional_headers
+ additional_headers["x-litellm-model-group"] = model_group
# Lift QualityRouter routing decision into response headers for
# transparency. The decision is stashed in request_kwargs.metadata
diff --git a/litellm/router_utils/add_retry_fallback_headers.py b/litellm/router_utils/add_retry_fallback_headers.py
index 6b921a0db8a..0b927714ca9 100644
--- a/litellm/router_utils/add_retry_fallback_headers.py
+++ b/litellm/router_utils/add_retry_fallback_headers.py
@@ -1,44 +1,99 @@
-from typing import Any, Optional, Union
+import json
+from typing import Protocol, TypedDict, cast
from pydantic import BaseModel
-from litellm.types.utils import HiddenParams
+
+class FallbackErrorInfo(TypedDict):
+ message: str
+ type: str
+ param: str | None
+ code: str | None
-def _add_headers_to_response(response: Any, headers: dict) -> Any:
+class _HiddenParamsHost(Protocol):
+ _hidden_params: dict[str, object]
+
+
+def get_hidden_params_dict(response: object) -> dict[str, object]:
+ hidden_params: object = cast(object, getattr(response, "_hidden_params", None))
+ if isinstance(hidden_params, BaseModel):
+ return cast("dict[str, object]", hidden_params.model_dump())
+ if isinstance(hidden_params, dict):
+ return cast("dict[str, object]", hidden_params)
+ return {}
+
+
+def _ensure_additional_headers_dict(
+ hidden_params: dict[str, object],
+) -> dict[str, object]:
+ additional_headers = hidden_params.get("additional_headers")
+ if isinstance(additional_headers, dict):
+ return cast("dict[str, object]", additional_headers)
+ return {}
+
+
+def get_fallback_error_info(error: Exception) -> FallbackErrorInfo:
+ message = cast(object, getattr(error, "message", str(error)))
+ error_type = cast(object, getattr(error, "type", error.__class__.__name__))
+ param = cast(object, getattr(error, "param", None))
+ code = cast(object, getattr(error, "status_code", getattr(error, "code", None)))
+ return FallbackErrorInfo(
+ message=str(message),
+ type=str(error_type),
+ param=str(param) if param is not None else None,
+ code=str(code) if code is not None else None,
+ )
+
+
+def _coerce_error_dicts(items: list[object]) -> list[dict[str, object]]:
+ return [cast("dict[str, object]", item) for item in items if isinstance(item, dict)]
+
+
+def get_fallback_errors_from_headers(
+ additional_headers: dict[str, object],
+) -> list[dict[str, object]]:
+ existing_errors = additional_headers.get("x-litellm-fallback-errors")
+ if isinstance(existing_errors, list):
+ return _coerce_error_dicts(cast("list[object]", existing_errors))
+ if isinstance(existing_errors, str):
+ try:
+ parsed_errors: object = cast(object, json.loads(existing_errors))
+ except json.JSONDecodeError:
+ return []
+ if isinstance(parsed_errors, list):
+ return _coerce_error_dicts(cast("list[object]", parsed_errors))
+ return []
+
+
+def _add_headers_to_response(response: object, headers: dict[str, object]) -> object:
"""
Helper function to add headers to a response's hidden params
"""
- if response is None or not isinstance(response, BaseModel):
+ if response is None:
return response
- hidden_params: Optional[Union[dict, HiddenParams]] = getattr(
- response, "_hidden_params", {}
- )
+ if not isinstance(response, BaseModel) and not hasattr(response, "_hidden_params"):
+ return response
- if hidden_params is None:
- hidden_params_dict = {}
- elif isinstance(hidden_params, HiddenParams):
- hidden_params_dict = hidden_params.model_dump()
- else:
- hidden_params_dict = hidden_params
+ hidden_params = get_hidden_params_dict(response)
+ additional_headers = _ensure_additional_headers_dict(hidden_params)
+ additional_headers.update(headers)
+ hidden_params["additional_headers"] = additional_headers
- hidden_params_dict.setdefault("additional_headers", {})
- hidden_params_dict["additional_headers"].update(headers)
-
- setattr(response, "_hidden_params", hidden_params_dict)
+ cast(_HiddenParamsHost, response)._hidden_params = hidden_params
return response
def add_retry_headers_to_response(
- response: Any,
+ response: object,
attempted_retries: int,
- max_retries: Optional[int] = None,
-) -> Any:
+ max_retries: int | None = None,
+) -> object:
"""
Add retry headers to the request
"""
- retry_headers = {
+ retry_headers: dict[str, object] = {
"x-litellm-attempted-retries": attempted_retries,
}
if max_retries is not None:
@@ -48,9 +103,10 @@ def add_retry_headers_to_response(
def add_fallback_headers_to_response(
- response: Any,
+ response: object,
attempted_fallbacks: int,
-) -> Any:
+ fallback_errors: list[FallbackErrorInfo] | None = None,
+) -> object:
"""
Add fallback headers to the response
@@ -64,7 +120,19 @@ def add_fallback_headers_to_response(
Note: It's intentional that we don't add max_fallbacks in response headers
Want to avoid bloat in the response headers for performance.
"""
- fallback_headers = {
+ fallback_headers: dict[str, object] = {
"x-litellm-attempted-fallbacks": attempted_fallbacks,
}
- return _add_headers_to_response(response, fallback_headers)
+ response = _add_headers_to_response(response, fallback_headers)
+ if fallback_errors is None or response is None:
+ return response
+
+ hidden_params = get_hidden_params_dict(response)
+ additional_headers = _ensure_additional_headers_dict(hidden_params)
+ merged_errors = get_fallback_errors_from_headers(additional_headers) + [
+ cast("dict[str, object]", error) for error in fallback_errors
+ ]
+ additional_headers["x-litellm-fallback-errors"] = json.dumps(merged_errors)
+ hidden_params["additional_headers"] = additional_headers
+ cast(_HiddenParamsHost, response)._hidden_params = hidden_params
+ return response
diff --git a/litellm/router_utils/cooldown_cache.py b/litellm/router_utils/cooldown_cache.py
index b210ea44596..dcfa44381c1 100644
--- a/litellm/router_utils/cooldown_cache.py
+++ b/litellm/router_utils/cooldown_cache.py
@@ -38,6 +38,7 @@ class CooldownCache:
visible_prefix=50, # Show first 50 characters
visible_suffix=0, # Show last 0 characters
mask_char="*", # Use * for masking
+ mask_short_values=False, # Truncate long messages only; keep short ones readable
)
def _common_add_cooldown_logic(
diff --git a/litellm/router_utils/fallback_event_handlers.py b/litellm/router_utils/fallback_event_handlers.py
index eb756e3cf8b..f0edc7fc9db 100644
--- a/litellm/router_utils/fallback_event_handlers.py
+++ b/litellm/router_utils/fallback_event_handlers.py
@@ -6,6 +6,7 @@ from litellm._logging import verbose_router_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.router_utils.add_retry_fallback_headers import (
add_fallback_headers_to_response,
+ get_fallback_error_info,
)
from litellm.types.router import LiteLLMParamsTypedDict
@@ -90,6 +91,7 @@ async def run_async_fallback(
original_exception: Exception,
max_fallbacks: int,
fallback_depth: int,
+ include_fallback_errors: bool = False,
**kwargs,
) -> Any:
"""
@@ -118,6 +120,7 @@ async def run_async_fallback(
raise original_exception
error_from_fallbacks = original_exception
+ fallback_errors = (get_fallback_error_info(original_exception),)
for mg in fallback_model_group:
if mg == original_model_group:
@@ -136,6 +139,8 @@ async def run_async_fallback(
fallback_depth = fallback_depth + 1
kwargs["fallback_depth"] = fallback_depth
kwargs["max_fallbacks"] = max_fallbacks
+ if include_fallback_errors:
+ kwargs["include_fallback_errors"] = include_fallback_errors
response = await litellm_router.async_function_with_fallbacks(
*args, **kwargs
)
@@ -143,6 +148,9 @@ async def run_async_fallback(
response = add_fallback_headers_to_response(
response=response,
attempted_fallbacks=fallback_depth,
+ fallback_errors=(
+ list(fallback_errors) if include_fallback_errors else None
+ ),
)
# callback for successfull_fallback_event():
await log_success_fallback_event(
@@ -153,6 +161,7 @@ async def run_async_fallback(
return response
except Exception as e:
error_from_fallbacks = e
+ fallback_errors = fallback_errors + (get_fallback_error_info(e),)
await log_failure_fallback_event(
original_model_group=original_model_group,
kwargs=kwargs,
diff --git a/litellm/sandbox/main.py b/litellm/sandbox/main.py
index 45d3bffb4f9..76d9994c683 100644
--- a/litellm/sandbox/main.py
+++ b/litellm/sandbox/main.py
@@ -68,7 +68,7 @@ async def acreate_sandbox(
provider: str,
template: str | None = None,
timeout: int | None = None,
- allow_internet_access: bool = True,
+ allow_internet_access: bool | None = None,
api_key: str | None = None,
api_base: str | None = None,
**kwargs,
diff --git a/litellm/types/completion.py b/litellm/types/completion.py
index cb263914be8..a91f6234fad 100644
--- a/litellm/types/completion.py
+++ b/litellm/types/completion.py
@@ -1,8 +1,28 @@
-from typing import Iterable, List, Optional, Union
+from __future__ import annotations
+
+from dataclasses import dataclass
+from typing import (
+ TYPE_CHECKING,
+ Any,
+ Callable,
+ Coroutine,
+ Iterable,
+ List,
+ Optional,
+ Union,
+)
from pydantic import BaseModel, ConfigDict
from typing_extensions import Literal, Required, TypedDict
+if TYPE_CHECKING:
+ import httpx
+ from aiohttp import ClientSession
+
+ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
+ from litellm.llms.base_llm import BaseConfig
+ from litellm.utils import CustomStreamWrapper, ModelResponse
+
class ChatCompletionSystemMessageParam(TypedDict, total=False):
content: Required[str]
@@ -191,3 +211,44 @@ class CompletionRequest(BaseModel):
model_list: Optional[List[str]] = None
model_config = ConfigDict(protected_namespaces=(), extra="allow")
+
+
+@dataclass(frozen=True, slots=True)
+class _CompletionDispatchContext:
+ _azure_detection_model: str
+ acompletion: bool
+ api_base: Optional[str]
+ api_key: Optional[str]
+ api_version: Optional[str]
+ client: Any
+ custom_llm_provider: str
+ custom_prompt_dict: dict
+ extra_headers: Optional[dict]
+ headers: dict
+ hf_model_name: Optional[str]
+ kwargs: dict
+ litellm_params: dict
+ logger_fn: Optional[Callable]
+ logging: LiteLLMLoggingObj
+ max_retries: Optional[int]
+ max_tokens: Optional[int]
+ messages: list
+ metadata: Optional[dict]
+ model: str
+ model_response: ModelResponse
+ optional_params: dict
+ organization: Optional[str]
+ provider_config: Optional[BaseConfig]
+ shared_session: Optional[ClientSession]
+ stream: Optional[bool]
+ temperature: Optional[float]
+ text_completion: bool
+ timeout: Optional[Union[float, str, httpx.Timeout]]
+ top_p: Optional[float]
+
+
+_CompletionDispatchResult = Union[
+ Coroutine[Any, Any, Union["ModelResponse", "CustomStreamWrapper"]],
+ "ModelResponse",
+ "CustomStreamWrapper",
+]
diff --git a/litellm/types/integrations/custom_logger.py b/litellm/types/integrations/custom_logger.py
index b5726a11ca0..26a0be36ef4 100644
--- a/litellm/types/integrations/custom_logger.py
+++ b/litellm/types/integrations/custom_logger.py
@@ -2,6 +2,25 @@ from typing import Any, Dict, List, Optional
from pydantic import BaseModel, Field
+CHAT_COMPLETION_AGENTIC_SURFACE = "chat_completions"
+CODE_INTERPRETER_INTERCEPTION_PREFIX = "_code_interpreter_interception"
+NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES = frozenset(
+ ("_websearch_interception", "_compression_interception")
+)
+INTERCEPTION_INTERNAL_PREFIXES = frozenset(
+ (
+ *NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES,
+ CODE_INTERPRETER_INTERCEPTION_PREFIX,
+ )
+)
+
+
+def is_interception_internal_key(
+ key: str,
+ prefixes: frozenset[str] = INTERCEPTION_INTERNAL_PREFIXES,
+) -> bool:
+ return any(key.startswith(prefix) for prefix in prefixes)
+
class StandardCustomLoggerInitParams(BaseModel):
"""
diff --git a/litellm/types/interactions/generated.py b/litellm/types/interactions/generated.py
index b38cd8f58b9..a07642073af 100644
--- a/litellm/types/interactions/generated.py
+++ b/litellm/types/interactions/generated.py
@@ -954,9 +954,6 @@ class Interaction(BaseModel):
None,
description="Output only. The time at which the response was last updated in ISO 8601 format\n(YYYY-MM-DDThh:mm:ssZ).",
)
- role: Optional[str] = Field(
- None, description="Output only. The role of the interaction."
- )
outputs: Optional[List[Content]] = Field(
None, description="Output only. Responses from the model."
)
@@ -1031,9 +1028,6 @@ class CreateModelInteractionParams(BaseModel):
None,
description="Output only. The time at which the response was last updated in ISO 8601 format\n(YYYY-MM-DDThh:mm:ssZ).",
)
- role: Optional[str] = Field(
- None, description="Output only. The role of the interaction."
- )
outputs: Optional[List[Content]] = Field(
None, description="Output only. Responses from the model."
)
@@ -1101,9 +1095,6 @@ class CreateAgentInteractionParams(BaseModel):
None,
description="Output only. The time at which the response was last updated in ISO 8601 format\n(YYYY-MM-DDThh:mm:ssZ).",
)
- role: Optional[str] = Field(
- None, description="Output only. The role of the interaction."
- )
outputs: Optional[List[Content]] = Field(
None, description="Output only. Responses from the model."
)
@@ -1323,7 +1314,6 @@ class InteractionsAPIResponse(BaseLiteLLMOpenAIResponseObject):
status: Optional[str] = None
created: Optional[str] = None
updated: Optional[str] = None
- role: Optional[str] = None
# Legacy schema field (Api-Revision: 2026-05-07). Remove after June 8, 2026.
outputs: Optional[List[Dict[str, Any]]] = None
# New schema field (Api-Revision: 2026-05-20).
@@ -1356,7 +1346,6 @@ class InteractionsAPIStreamingResponse(BaseLiteLLMOpenAIResponseObject):
status: Optional[str] = None
created: Optional[str] = None
updated: Optional[str] = None
- role: Optional[str] = None
# Legacy schema field (Api-Revision: 2026-05-07). Remove after June 8, 2026.
outputs: Optional[List[Dict[str, Any]]] = None
# New schema field (Api-Revision: 2026-05-20).
diff --git a/litellm/types/utils.py b/litellm/types/utils.py
index 16e693c7d7c..24d6e84fba7 100644
--- a/litellm/types/utils.py
+++ b/litellm/types/utils.py
@@ -3163,8 +3163,26 @@ class CustomPricingLiteLLMParams(BaseModel):
regional_processing_uplift_multiplier_us: Optional[float] = None
+# Server-controlled fields that bound or drive an interceptor's agentic loop
+# (depth, cycle fingerprints, ceiling, code-interpreter sandbox state). Listed
+# in all_litellm_params so they are treated as LiteLLM-level and excluded from
+# get_non_default_completion_params; otherwise the OpenAI param builder sweeps
+# any unrecognized top-level key into extra_body and leaks them to the provider.
+# This is what lets the loop carry state across rerun calls without a provider
+# scrubber.
+agentic_loop_internal_litellm_params = [
+ "_agentic_loop_depth",
+ "_agentic_loop_fingerprints",
+ "_agentic_loop_api_surface",
+ "max_agentic_loops",
+ "_code_interpreter_interception_active",
+ "_code_interpreter_interception_sandbox_key",
+ "_code_interpreter_interception_converted_stream",
+]
+
all_litellm_params = (
- [
+ agentic_loop_internal_litellm_params
+ + [
"metadata",
"litellm_metadata",
"litellm_trace_id",
@@ -3463,6 +3481,7 @@ class LlmProviders(str, Enum):
TENSORMESH = "tensormesh"
LIBERTAI = "libertai"
PINSTRIPES = "pinstripes"
+ DARKBLOOM = "darkbloom"
LITELLM_AGENT = "litellm_agent"
CURSOR = "cursor"
BEDROCK_MANTLE = "bedrock_mantle"
@@ -3520,6 +3539,7 @@ class SandboxProviders(str, Enum):
"""
E2B = "e2b"
+ OPENSANDBOX = "opensandbox"
class LiteLLMLoggingBaseClass:
diff --git a/litellm/utils.py b/litellm/utils.py
index 29f703104da..5c3ab3e1490 100644
--- a/litellm/utils.py
+++ b/litellm/utils.py
@@ -3191,7 +3191,7 @@ def get_optional_params_transcription(
model=model,
drop_params=drop_params if drop_params is not None else False,
)
- elif provider_config is not None: # handles fireworks ai, and any future providers
+ elif provider_config is not None: # custom audio transcription config
supported_params = provider_config.get_supported_openai_params(model=model)
_check_valid_arg(supported_params=supported_params)
optional_params = provider_config.map_openai_params(
@@ -8915,8 +8915,6 @@ class ProviderConfigManager:
)
return AzureSpeechAudioTranscriptionConfig()
- if litellm.LlmProviders.FIREWORKS_AI == provider:
- return litellm.FireworksAIAudioTranscriptionConfig()
elif litellm.LlmProviders.DEEPGRAM == provider:
return litellm.DeepgramAudioTranscriptionConfig()
elif litellm.LlmProviders.ELEVENLABS == provider:
@@ -9733,9 +9731,14 @@ class ProviderConfigManager:
Get sandbox (code execution) configuration for a given provider.
"""
from litellm.llms.e2b.sandbox.transformation import E2BSandboxConfig
+ from litellm.llms.opensandbox.sandbox.transformation import (
+ OpenSandboxSandboxConfig,
+ )
if provider == SandboxProviders.E2B:
return E2BSandboxConfig()
+ if provider == SandboxProviders.OPENSANDBOX:
+ return OpenSandboxSandboxConfig()
return None
@staticmethod
diff --git a/migrations/Dockerfile b/migrations/Dockerfile
index a78a4e2225a..caca280cbfc 100644
--- a/migrations/Dockerfile
+++ b/migrations/Dockerfile
@@ -1,5 +1,5 @@
-ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:31da6565f35af6401031c1d7aa91dc84ac76c5c48edd17fb90f0ed9e3173c7a9
-ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:31da6565f35af6401031c1d7aa91dc84ac76c5c48edd17fb90f0ed9e3173c7a9
+ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:c61ac6919b811ea53c4782d69f1fe05218ba3c25d53f01b6ab7892e621bd4370
+ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:c61ac6919b811ea53c4782d69f1fe05218ba3c25d53f01b6ab7892e621bd4370
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
FROM $UV_IMAGE AS uvbin
diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json
index 3c50dde9277..5f3f2294147 100644
--- a/model_prices_and_context_window.json
+++ b/model_prices_and_context_window.json
@@ -571,7 +571,7 @@
"output_vector_size": 1536
},
"amazon.titan-embed-text-v2:0": {
- "input_cost_per_token": 2e-07,
+ "input_cost_per_token": 2e-08,
"litellm_provider": "bedrock",
"max_input_tokens": 8192,
"max_tokens": 8192,
@@ -10443,7 +10443,8 @@
"fast": 6.0
},
"supports_output_config": true,
- "supports_max_reasoning_effort": true
+ "supports_max_reasoning_effort": true,
+ "supports_speed": true
},
"claude-opus-4-6-20260205": {
"cache_creation_input_token_cost": 6.25e-06,
@@ -10476,7 +10477,8 @@
"fast": 6.0
},
"supports_max_reasoning_effort": true,
- "supports_output_config": true
+ "supports_output_config": true,
+ "supports_speed": true
},
"claude-opus-4-7": {
"cache_creation_input_token_cost": 6.25e-06,
@@ -10511,7 +10513,8 @@
"us": 1.1,
"fast": 6.0
},
- "supports_output_config": true
+ "supports_output_config": true,
+ "supports_speed": true
},
"claude-opus-4-7-20260416": {
"cache_creation_input_token_cost": 6.25e-06,
@@ -10546,7 +10549,8 @@
"us": 1.1,
"fast": 6.0
},
- "supports_output_config": true
+ "supports_output_config": true,
+ "supports_speed": true
},
"claude-fable-5": {
"cache_creation_input_token_cost": 1.25e-05,
@@ -10615,7 +10619,8 @@
"us": 1.1,
"fast": 2.0
},
- "supports_output_config": true
+ "supports_output_config": true,
+ "supports_speed": true
},
"claude-sonnet-4-20250514": {
"deprecation_date": "2026-05-14",
@@ -10684,6 +10689,268 @@
"mode": "chat",
"output_cost_per_token": 1.923e-06
},
+ "cloudflare/@cf/openai/gpt-oss-120b": {
+ "input_cost_per_token": 3.5e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 128000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 7.5e-07,
+ "supports_function_calling": true,
+ "supports_reasoning": true
+ },
+ "cloudflare/@cf/google/gemma-2b-it-lora": {
+ "input_cost_per_token": 0.0,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 8192,
+ "max_output_tokens": 8192,
+ "max_tokens": 8192,
+ "mode": "chat",
+ "output_cost_per_token": 0.0
+ },
+ "cloudflare/@cf/meta/llama-3.2-3b-instruct": {
+ "input_cost_per_token": 5.09e-08,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 80000,
+ "max_output_tokens": 80000,
+ "max_tokens": 80000,
+ "mode": "chat",
+ "output_cost_per_token": 3.35e-07
+ },
+ "cloudflare/@cf/meta/llama-guard-3-8b": {
+ "input_cost_per_token": 4.84e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 131072,
+ "max_output_tokens": 131072,
+ "max_tokens": 131072,
+ "mode": "chat",
+ "output_cost_per_token": 3e-08
+ },
+ "cloudflare/@cf/mistral/mistral-7b-instruct-v0.2-lora": {
+ "input_cost_per_token": 0.0,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 15000,
+ "max_output_tokens": 15000,
+ "max_tokens": 15000,
+ "mode": "chat",
+ "output_cost_per_token": 0.0
+ },
+ "cloudflare/@cf/moonshotai/kimi-k2.7-code": {
+ "cache_read_input_token_cost": 1.9e-07,
+ "input_cost_per_token": 9.5e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 262144,
+ "max_output_tokens": 262144,
+ "max_tokens": 262144,
+ "mode": "chat",
+ "output_cost_per_token": 4e-06,
+ "supports_function_calling": true,
+ "supports_reasoning": true
+ },
+ "cloudflare/@cf/deepseek-ai/deepseek-r1-distill-qwen-32b": {
+ "input_cost_per_token": 4.97e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 80000,
+ "max_output_tokens": 80000,
+ "max_tokens": 80000,
+ "mode": "chat",
+ "output_cost_per_token": 4.881e-06,
+ "supports_reasoning": true
+ },
+ "cloudflare/@cf/meta/llama-3.1-8b-instruct-fp8": {
+ "input_cost_per_token": 1.52e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 32000,
+ "max_output_tokens": 32000,
+ "max_tokens": 32000,
+ "mode": "chat",
+ "output_cost_per_token": 2.87e-07
+ },
+ "cloudflare/@cf/meta/llama-3.2-1b-instruct": {
+ "input_cost_per_token": 2.7e-08,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 60000,
+ "max_output_tokens": 60000,
+ "max_tokens": 60000,
+ "mode": "chat",
+ "output_cost_per_token": 2.01e-07
+ },
+ "cloudflare/@cf/moonshotai/kimi-k2.6": {
+ "cache_read_input_token_cost": 1.6e-07,
+ "input_cost_per_token": 9.5e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 262144,
+ "max_output_tokens": 262144,
+ "max_tokens": 262144,
+ "mode": "chat",
+ "output_cost_per_token": 4e-06,
+ "supports_function_calling": true,
+ "supports_reasoning": true
+ },
+ "cloudflare/@cf/zai-org/glm-4.7-flash": {
+ "input_cost_per_token": 6.05e-08,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 131072,
+ "max_output_tokens": 131072,
+ "max_tokens": 131072,
+ "mode": "chat",
+ "output_cost_per_token": 4e-07,
+ "supports_function_calling": true,
+ "supports_reasoning": true
+ },
+ "cloudflare/@cf/meta-llama/llama-2-7b-chat-hf-lora": {
+ "input_cost_per_token": 0.0,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 8192,
+ "max_output_tokens": 8192,
+ "max_tokens": 8192,
+ "mode": "chat",
+ "output_cost_per_token": 0.0
+ },
+ "cloudflare/@cf/meta/llama-3.3-70b-instruct-fp8-fast": {
+ "input_cost_per_token": 2.93e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 24000,
+ "max_output_tokens": 24000,
+ "max_tokens": 24000,
+ "mode": "chat",
+ "output_cost_per_token": 2.253e-06,
+ "supports_function_calling": true
+ },
+ "cloudflare/@cf/ibm-granite/granite-4.0-h-micro": {
+ "input_cost_per_token": 1.7e-08,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 131000,
+ "max_output_tokens": 131000,
+ "max_tokens": 131000,
+ "mode": "chat",
+ "output_cost_per_token": 1.12e-07,
+ "supports_function_calling": true
+ },
+ "cloudflare/@cf/qwen/qwen2.5-coder-32b-instruct": {
+ "input_cost_per_token": 6.6e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 32768,
+ "max_output_tokens": 32768,
+ "max_tokens": 32768,
+ "mode": "chat",
+ "output_cost_per_token": 1e-06
+ },
+ "cloudflare/@cf/zai-org/glm-5.2": {
+ "cache_read_input_token_cost": 2.6e-07,
+ "input_cost_per_token": 1.4e-06,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 262144,
+ "max_output_tokens": 262144,
+ "max_tokens": 262144,
+ "mode": "chat",
+ "output_cost_per_token": 4.4e-06,
+ "supports_function_calling": true,
+ "supports_reasoning": true
+ },
+ "cloudflare/@cf/nvidia/nemotron-3-120b-a12b": {
+ "input_cost_per_token": 5e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 256000,
+ "max_output_tokens": 256000,
+ "max_tokens": 256000,
+ "mode": "chat",
+ "output_cost_per_token": 1.5e-06,
+ "supports_function_calling": true,
+ "supports_reasoning": true
+ },
+ "cloudflare/@cf/aisingapore/gemma-sea-lion-v4-27b-it": {
+ "input_cost_per_token": 3.51e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 128000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 5.55e-07
+ },
+ "cloudflare/@cf/qwen/qwen3-30b-a3b-fp8": {
+ "input_cost_per_token": 5.09e-08,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 32768,
+ "max_output_tokens": 32768,
+ "max_tokens": 32768,
+ "mode": "chat",
+ "output_cost_per_token": 3.35e-07,
+ "supports_function_calling": true,
+ "supports_reasoning": true
+ },
+ "cloudflare/@cf/google/gemma-7b-it-lora": {
+ "input_cost_per_token": 0.0,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 3500,
+ "max_output_tokens": 3500,
+ "max_tokens": 3500,
+ "mode": "chat",
+ "output_cost_per_token": 0.0
+ },
+ "cloudflare/@cf/google/gemma-4-26b-a4b-it": {
+ "input_cost_per_token": 1e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 256000,
+ "max_output_tokens": 256000,
+ "max_tokens": 256000,
+ "mode": "chat",
+ "output_cost_per_token": 3e-07,
+ "supports_function_calling": true,
+ "supports_reasoning": true
+ },
+ "cloudflare/@cf/mistralai/mistral-small-3.1-24b-instruct": {
+ "input_cost_per_token": 3.51e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 128000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 5.55e-07,
+ "supports_function_calling": true
+ },
+ "cloudflare/@cf/meta/llama-3.2-11b-vision-instruct": {
+ "input_cost_per_token": 4.85e-08,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 128000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 6.76e-07,
+ "supports_vision": true
+ },
+ "cloudflare/@cf/openai/gpt-oss-20b": {
+ "input_cost_per_token": 2e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 128000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 3e-07,
+ "supports_function_calling": true,
+ "supports_reasoning": true
+ },
+ "cloudflare/@cf/meta/llama-4-scout-17b-16e-instruct": {
+ "input_cost_per_token": 2.7e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 131000,
+ "max_output_tokens": 131000,
+ "max_tokens": 131000,
+ "mode": "chat",
+ "output_cost_per_token": 8.5e-07,
+ "supports_function_calling": true
+ },
+ "cloudflare/@cf/qwen/qwq-32b": {
+ "input_cost_per_token": 6.6e-07,
+ "litellm_provider": "cloudflare",
+ "max_input_tokens": 24000,
+ "max_output_tokens": 24000,
+ "max_tokens": 24000,
+ "mode": "chat",
+ "output_cost_per_token": 1e-06,
+ "supports_reasoning": true
+ },
"codestral/codestral-2405": {
"input_cost_per_token": 0.0,
"litellm_provider": "codestral",
@@ -20096,8 +20363,6 @@
"output_cost_per_token": 8e-06,
"output_cost_per_token_batches": 4e-06,
"output_cost_per_token_priority": 1.4e-05,
- "regional_processing_uplift_multiplier_eu": 1.10,
- "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -20171,8 +20436,6 @@
"output_cost_per_token": 1.6e-06,
"output_cost_per_token_batches": 8e-07,
"output_cost_per_token_priority": 2.8e-06,
- "regional_processing_uplift_multiplier_eu": 1.10,
- "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -20246,8 +20509,6 @@
"output_cost_per_token": 4e-07,
"output_cost_per_token_batches": 2e-07,
"output_cost_per_token_priority": 8e-07,
- "regional_processing_uplift_multiplier_eu": 1.10,
- "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -20319,8 +20580,6 @@
"output_cost_per_token": 1e-05,
"output_cost_per_token_batches": 5e-06,
"output_cost_per_token_priority": 1.7e-05,
- "regional_processing_uplift_multiplier_eu": 1.10,
- "regional_processing_uplift_multiplier_us": 1.10,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
@@ -20362,8 +20621,6 @@
"mode": "chat",
"output_cost_per_token": 1e-05,
"output_cost_per_token_batches": 5e-06,
- "regional_processing_uplift_multiplier_eu": 1.10,
- "regional_processing_uplift_multiplier_us": 1.10,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
@@ -20385,8 +20642,6 @@
"mode": "chat",
"output_cost_per_token": 1e-05,
"output_cost_per_token_batches": 5e-06,
- "regional_processing_uplift_multiplier_eu": 1.10,
- "regional_processing_uplift_multiplier_us": 1.10,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
@@ -20675,8 +20930,6 @@
"output_cost_per_token": 6e-07,
"output_cost_per_token_batches": 3e-07,
"output_cost_per_token_priority": 1e-06,
- "regional_processing_uplift_multiplier_eu": 1.10,
- "regional_processing_uplift_multiplier_us": 1.10,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
@@ -21380,8 +21633,6 @@
"output_cost_per_token": 1e-05,
"output_cost_per_token_flex": 5e-06,
"output_cost_per_token_priority": 2e-05,
- "regional_processing_uplift_multiplier_eu": 1.10,
- "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -21775,6 +22026,8 @@
"output_cost_per_token_flex": 1.5e-05,
"output_cost_per_token_batches": 1.5e-05,
"output_cost_per_token_priority": 6e-05,
+ "regional_processing_uplift_multiplier_eu": 1.10,
+ "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -21823,6 +22076,8 @@
"output_cost_per_token_flex": 1.5e-05,
"output_cost_per_token_batches": 1.5e-05,
"output_cost_per_token_priority": 6e-05,
+ "regional_processing_uplift_multiplier_eu": 1.10,
+ "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -21867,6 +22122,8 @@
"output_cost_per_token_above_272k_tokens": 0.00027,
"output_cost_per_token_flex": 9e-05,
"output_cost_per_token_batches": 9e-05,
+ "regional_processing_uplift_multiplier_eu": 1.10,
+ "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/responses",
"/v1/batch"
@@ -21911,6 +22168,8 @@
"output_cost_per_token_above_272k_tokens": 0.00027,
"output_cost_per_token_flex": 9e-05,
"output_cost_per_token_batches": 9e-05,
+ "regional_processing_uplift_multiplier_eu": 1.10,
+ "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/responses",
"/v1/batch"
@@ -21959,6 +22218,8 @@
"output_cost_per_token_flex": 7.5e-06,
"output_cost_per_token_batches": 7.5e-06,
"output_cost_per_token_priority": 3e-05,
+ "regional_processing_uplift_multiplier_eu": 1.10,
+ "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -22006,6 +22267,8 @@
"output_cost_per_token_flex": 7.5e-06,
"output_cost_per_token_batches": 7.5e-06,
"output_cost_per_token_priority": 3e-05,
+ "regional_processing_uplift_multiplier_eu": 1.10,
+ "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -22046,6 +22309,8 @@
"output_cost_per_token_above_272k_tokens": 0.00027,
"output_cost_per_token_flex": 9e-05,
"output_cost_per_token_batches": 9e-05,
+ "regional_processing_uplift_multiplier_eu": 1.10,
+ "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/responses",
"/v1/batch"
@@ -22089,6 +22354,8 @@
"output_cost_per_token_above_272k_tokens": 0.00027,
"output_cost_per_token_flex": 9e-05,
"output_cost_per_token_batches": 9e-05,
+ "regional_processing_uplift_multiplier_eu": 1.10,
+ "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/responses",
"/v1/batch"
@@ -22134,6 +22401,8 @@
"output_cost_per_token_flex": 2.25e-06,
"output_cost_per_token_batches": 2.25e-06,
"output_cost_per_token_priority": 9e-06,
+ "regional_processing_uplift_multiplier_eu": 1.10,
+ "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -22180,6 +22449,8 @@
"output_cost_per_token_flex": 2.25e-06,
"output_cost_per_token_batches": 2.25e-06,
"output_cost_per_token_priority": 9e-06,
+ "regional_processing_uplift_multiplier_eu": 1.10,
+ "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -22223,6 +22494,8 @@
"output_cost_per_token": 1.25e-06,
"output_cost_per_token_flex": 6.25e-07,
"output_cost_per_token_batches": 6.25e-07,
+ "regional_processing_uplift_multiplier_eu": 1.10,
+ "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -22266,6 +22539,8 @@
"output_cost_per_token": 1.25e-06,
"output_cost_per_token_flex": 6.25e-07,
"output_cost_per_token_batches": 6.25e-07,
+ "regional_processing_uplift_multiplier_eu": 1.10,
+ "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -22304,8 +22579,6 @@
"mode": "responses",
"output_cost_per_token": 0.00012,
"output_cost_per_token_batches": 6e-05,
- "regional_processing_uplift_multiplier_eu": 1.10,
- "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/batch",
"/v1/responses"
@@ -22712,8 +22985,6 @@
"output_cost_per_token": 2e-06,
"output_cost_per_token_flex": 1e-06,
"output_cost_per_token_priority": 3.6e-06,
- "regional_processing_uplift_multiplier_eu": 1.10,
- "regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -22795,8 +23066,6 @@
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
- "regional_processing_uplift_multiplier_eu": 1.10,
- "regional_processing_uplift_multiplier_us": 1.10,
"mode": "chat",
"output_cost_per_token": 4e-07,
"output_cost_per_token_flex": 2e-07,
@@ -39946,24 +40215,6 @@
"litellm_provider": "fireworks_ai",
"mode": "chat"
},
- "fireworks_ai/accounts/fireworks/models/whisper-v3": {
- "max_tokens": 4096,
- "max_input_tokens": 4096,
- "max_output_tokens": 4096,
- "input_cost_per_token": 0.0,
- "output_cost_per_token": 0.0,
- "litellm_provider": "fireworks_ai",
- "mode": "audio_transcription"
- },
- "fireworks_ai/accounts/fireworks/models/whisper-v3-turbo": {
- "max_tokens": 4096,
- "max_input_tokens": 4096,
- "max_output_tokens": 4096,
- "input_cost_per_token": 0.0,
- "output_cost_per_token": 0.0,
- "litellm_provider": "fireworks_ai",
- "mode": "audio_transcription"
- },
"fireworks_ai/accounts/fireworks/models/yi-34b": {
"max_tokens": 4096,
"max_input_tokens": 4096,
@@ -43543,5 +43794,39 @@
"supports_assistant_prefill": true,
"supports_reasoning": false,
"source": "https://pinstripes.io/pricing"
+ },
+ "darkbloom/gemma-4-26b": {
+ "input_cost_per_token": 3e-08,
+ "litellm_provider": "darkbloom",
+ "max_input_tokens": 131072,
+ "max_output_tokens": 32768,
+ "max_tokens": 32768,
+ "mode": "chat",
+ "output_cost_per_token": 1.65e-07,
+ "source": "https://www.darkbloom.dev/",
+ "supported_endpoints": [
+ "/v1/chat/completions"
+ ],
+ "supports_function_calling": true,
+ "supports_native_streaming": true,
+ "supports_system_messages": true,
+ "supports_tool_choice": true
+ },
+ "darkbloom/gpt-oss-20b": {
+ "input_cost_per_token": 1.45e-08,
+ "litellm_provider": "darkbloom",
+ "max_input_tokens": 131072,
+ "max_output_tokens": 32768,
+ "max_tokens": 32768,
+ "mode": "chat",
+ "output_cost_per_token": 7e-08,
+ "source": "https://www.darkbloom.dev/",
+ "supported_endpoints": [
+ "/v1/chat/completions"
+ ],
+ "supports_function_calling": true,
+ "supports_native_streaming": true,
+ "supports_system_messages": true,
+ "supports_tool_choice": true
}
}
diff --git a/osv-scanner.toml b/osv-scanner.toml
index f0f5f045f1a..7ab450945f5 100644
--- a/osv-scanner.toml
+++ b/osv-scanner.toml
@@ -2,13 +2,3 @@
id = "GHSA-w8v5-vhqr-4h9v"
ignoreUntil = 2026-09-09
reason = "diskcache has no fixed release published; remove this entry once one exists"
-
-[[IgnoredVulns]]
-id = "GHSA-hg6j-4rv6-33pg"
-ignoreUntil = 2026-08-15
-reason = "aiohttp held at 3.13.5: vcrpy releases <= 8.1.1 cannot import aiohttp >= 3.14 and the merged upstream fix (vcrpy PR 996) is unreleased; bump aiohttp and drop this entry when a newer vcrpy ships"
-
-[[IgnoredVulns]]
-id = "GHSA-jg22-mg44-37j8"
-ignoreUntil = 2026-08-15
-reason = "aiohttp held at 3.13.5: vcrpy releases <= 8.1.1 cannot import aiohttp >= 3.14 and the merged upstream fix (vcrpy PR 996) is unreleased; bump aiohttp and drop this entry when a newer vcrpy ships"
diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json
index 7386ced3e6d..b137ec59a1f 100644
--- a/provider_endpoints_support.json
+++ b/provider_endpoints_support.json
@@ -1833,6 +1833,23 @@
"text_completion": true
}
},
+ "opensandbox": {
+ "display_name": "OpenSandbox (`opensandbox`)",
+ "url": "https://open-sandbox.ai/api/",
+ "endpoints": {
+ "chat_completions": false,
+ "messages": false,
+ "responses": false,
+ "embeddings": false,
+ "image_generations": false,
+ "audio_transcriptions": false,
+ "audio_speech": false,
+ "moderations": false,
+ "batches": false,
+ "rerank": false,
+ "sandbox": true
+ }
+ },
"openai_like": {
"display_name": "OpenAI-like (`openai_like`)",
"url": "https://docs.litellm.ai/docs/providers/openai_compatible",
@@ -2008,6 +2025,23 @@
"interactions": true
}
},
+ "darkbloom": {
+ "display_name": "Darkbloom (`darkbloom`)",
+ "url": "https://docs.litellm.ai/docs/providers/darkbloom",
+ "endpoints": {
+ "chat_completions": true,
+ "messages": false,
+ "responses": false,
+ "embeddings": false,
+ "image_generations": false,
+ "audio_transcriptions": false,
+ "audio_speech": false,
+ "moderations": false,
+ "batches": false,
+ "rerank": false,
+ "a2a": false
+ }
+ },
"predibase": {
"display_name": "Predibase (`predibase`)",
"url": "https://docs.litellm.ai/docs/providers/predibase",
diff --git a/pyproject.toml b/pyproject.toml
index b728bfcd514..1cb39153e83 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -1,6 +1,6 @@
[project]
name = "litellm"
-version = "1.90.0"
+version = "1.91.0"
description = "Library to easily interface with LLM API providers"
readme = "README.md"
requires-python = ">=3.10, <3.14"
@@ -55,7 +55,7 @@ proxy = [
"fastapi-sso>=0.19.0,<1.0",
"PyJWT>=2.13.0,<3.0",
"python-multipart>=0.0.27,<1.0",
- "cryptography>=46.0.7,<47.0",
+ "cryptography>=48.0.1,<49.0",
"pynacl>=1.6.2,<2.0",
"websockets>=15.0.1,<16.0",
"boto3>=1.43.1,<2.0",
@@ -182,7 +182,7 @@ dev = [
"parameterized==0.9.0",
"openapi-core==0.22.0; python_version < '3.14'",
"pytest-timeout==2.4.0",
- "vcrpy==8.1.1",
+ "vcrpy==8.2.1",
"pytest-recording==0.13.4",
]
proxy-dev = [
@@ -208,7 +208,7 @@ ci = [
"pytest-codspeed==4.3.0",
"pytest-retry==1.7.0",
"pyarrow==23.0.1",
- "langchain==1.2.10",
+ "langchain==1.3.9",
"lunary==1.4.36; python_version == '3.10'",
"lunary==1.4.37; python_version >= '3.11'",
"logfire==4.6.0",
@@ -225,11 +225,8 @@ ci = [
"pylint==4.0.5",
"langchain-mcp-adapters==0.2.1",
"langchain-openai==1.1.14",
- "langgraph==1.0.10",
- # langgraph-prebuilt 1.0.9 imports ExecutionInfo/ServerInfo from
- # langgraph.runtime, which is not exported until langgraph 1.1.0.
- # Pin to 1.0.8 so it pairs correctly with langgraph==1.0.10.
- "langgraph-prebuilt==1.0.8",
+ "langgraph>=1.2.4,<1.3.0",
+ "langgraph-prebuilt>=1.1.0,<1.3.0",
"claude-agent-sdk==0.1.44",
]
healthcheck = [
@@ -244,7 +241,7 @@ build-backend = "uv_build"
[tool.uv]
constraint-dependencies = [
"tornado>=6.5.6",
- "aiohttp>=3.13.5,<3.14",
+ "aiohttp>=3.14.1,<4.0",
]
default-groups = ["dev"]
required-version = ">=0.10.9"
@@ -273,7 +270,7 @@ source-exclude = [
profile = "black"
[tool.commitizen]
-version = "1.90.0"
+version = "1.91.0"
version_files = [
"pyproject.toml:^version",
]
diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json
index 62ebdb559fc..ae46f020de1 100644
--- a/ruff-strict-budget.json
+++ b/ruff-strict-budget.json
@@ -300,7 +300,7 @@
"slack": 3
},
"RET504": {
- "baseline": 709,
+ "baseline": 702,
"slack": 20
},
"RUF010": {
diff --git a/scripts/type_check_gate.py b/scripts/type_check_gate.py
index 0f9a44703f9..2ef332d91ea 100644
--- a/scripts/type_check_gate.py
+++ b/scripts/type_check_gate.py
@@ -1,21 +1,22 @@
#!/usr/bin/env python3
-"""Per-rule count gate for basedpyright.
+"""Delta-vs-base per-rule gate for basedpyright.
basedpyright's ``--outputjson`` is reduced to a count of errors per *rule*
(``reportAny``, ``reportArgumentType``, ...) and checked against a committed
budget of the form ``{rule: {baseline, slack}}``, the same shape as
-``ruff-strict-budget.json``. A rule fails when its codebase-wide total exceeds
-``baseline + slack``. Counts ignore file, line, and column, so a violation
-moving anywhere in the tree is invisible; only the per-rule total moves the
-needle.
+``ruff-strict-budget.json``. A rule fails only when its codebase-wide total is
+both over its ceiling (``baseline + slack``) *and* higher than the count on the
+base it merges into, so a change is blamed for the errors it adds, never for
+drift that already sits in the base. That ``> base`` guard is what stops an
+unrelated PR from inheriting a red once two PRs each land near the ceiling and
+their sum crosses it: the bystander's count equals its base, so it is spared,
+while any PR that actually grows the rule past the cap still fails.
-Unlike ``ruff_strict_gate.py`` this does *not* re-run the tool on the merge base
-to compute a delta: a second basedpyright pass is minutes and gigabytes, whereas
-ruff is milliseconds. The committed budget is the baseline instead -- exactly
-how the previous per-file gate worked -- so keep it fresh with ``--update``
-(ratchet), which re-captures every rule's count from the current tree while
-preserving each rule's slack. Tool output is read from stdin, so the caller
-decides how to invoke basedpyright (and from which cwd).
+Head counts are read from stdin (the caller runs basedpyright once and pipes
+``--outputjson`` in); the base count is a second basedpyright pass over a
+detached worktree at the merge-base, run under the same environment so import
+resolution matches. ``--update`` re-captures the absolute per-rule baselines for
+the ratchet, preserving each rule's slack.
``--outputjson`` is used rather than text diagnostics because the latter wrap
across lines, leaving the ``(reportRule)`` on a continuation line away from the
@@ -24,13 +25,21 @@ carries an unambiguous ``rule`` field.
"""
import argparse
+import contextlib
import json
+import shutil
+import subprocess
import sys
+import tempfile
from collections import Counter
+from collections.abc import Iterator, Mapping
from pathlib import Path
-from typing import Mapping, NamedTuple
+from typing import NamedTuple
REPO_ROOT = Path(__file__).resolve().parent.parent
+BUDGET_PATH = REPO_ROOT / "basedpyright-code-budget.json"
+PYRIGHT_CONFIG = REPO_ROOT / "pyrightconfig.json"
+DEFAULT_BASE = "origin/litellm_internal_staging"
# Bucket for a basedpyright diagnostic with no `rule`. Counted so it's gated.
UNCODED = ""
@@ -45,6 +54,7 @@ class Breach(NamedTuple):
code: str
total: int
cap: int
+ added: int
def _seed_slack(baseline: int) -> int:
@@ -54,18 +64,19 @@ def _seed_slack(baseline: int) -> int:
return 10 if baseline >= 50 else 3
-def _to_repo_relative(raw: str) -> str | None:
+def _to_relative(raw: str, root: Path) -> str | None:
path = Path(raw)
- absolute = path if path.is_absolute() else Path.cwd() / path
+ absolute = path if path.is_absolute() else root / path
try:
- return absolute.resolve().relative_to(REPO_ROOT).as_posix()
+ return absolute.resolve().relative_to(root).as_posix()
except ValueError:
return None
-def count_basedpyright(payload: str) -> dict[str, int]:
- """Count in-repo basedpyright errors per rule from `--outputjson`. Warnings
- and information are ignored; only `severity == "error"` is gated."""
+def count_basedpyright(payload: str, root: Path = REPO_ROOT) -> dict[str, int]:
+ """Count in-tree basedpyright errors per rule from `--outputjson`. Warnings
+ and information are ignored; only `severity == "error"` is gated. Files
+ outside `root` (the venv's site-packages, say) are dropped."""
try:
data = json.loads(payload or "{}")
except json.JSONDecodeError as exc:
@@ -79,21 +90,62 @@ def count_basedpyright(payload: str) -> dict[str, int]:
for diag in data.get("generalDiagnostics", []):
if diag.get("severity") != "error":
continue
- if _to_repo_relative(diag.get("file", "")) is None:
+ if _to_relative(diag.get("file", ""), root) is None:
continue
counts[diag.get("rule") or UNCODED] += 1
return dict(counts)
+def _run(cmd: list[str], cwd: Path = REPO_ROOT) -> str:
+ proc = subprocess.run(cmd, cwd=cwd, capture_output=True, text=True)
+ if proc.returncode not in (0, 1):
+ sys.stderr.write(proc.stderr)
+ raise SystemExit(f"{cmd[0]} exited {proc.returncode}")
+ return proc.stdout
+
+
+@contextlib.contextmanager
+def _temp_worktree(ref: str) -> Iterator[Path]:
+ parent = Path(tempfile.mkdtemp(prefix="bpr_base_"))
+ worktree = parent / "wt"
+ try:
+ _run(["git", "worktree", "add", "--detach", str(worktree), ref])
+ yield worktree
+ finally:
+ subprocess.run(
+ ["git", "worktree", "remove", "--force", str(worktree)],
+ cwd=REPO_ROOT,
+ capture_output=True,
+ text=True,
+ )
+ shutil.rmtree(parent, ignore_errors=True)
+
+
+def base_counts(ref: str) -> dict[str, int]:
+ """basedpyright error counts per rule for the merge-base tree. The head
+ config is copied in so the base is judged by today's rules, and the run uses
+ the head environment's basedpyright (on PATH) so imports resolve the same."""
+ exe = shutil.which("basedpyright") or "basedpyright"
+ with _temp_worktree(ref) as worktree:
+ shutil.copy(PYRIGHT_CONFIG, worktree / "pyrightconfig.json")
+ proc = subprocess.run(
+ [exe, "--outputjson"], cwd=worktree, capture_output=True, text=True
+ )
+ return count_basedpyright(proc.stdout, root=worktree)
+
+
def evaluate(
- counts: Mapping[str, int], budget: Mapping[str, Mapping[str, int]]
+ head: Mapping[str, int],
+ base: Mapping[str, int],
+ budget: Mapping[str, Mapping[str, int]],
) -> list[Breach]:
breaches = []
- for code, total in counts.items():
+ for code, total in head.items():
spec = budget.get(code)
cap = spec["baseline"] + spec["slack"] if spec else DEFAULT_SLACK
- if total > cap:
- breaches.append(Breach(code, total, cap))
+ prior = base.get(code, 0)
+ if total > cap and total > prior:
+ breaches.append(Breach(code, total, cap, total - prior))
return sorted(breaches)
@@ -107,9 +159,6 @@ def is_vacuous_run(
return not counts and any(spec["baseline"] for spec in budget.values())
-BUDGET_PATH = REPO_ROOT / "basedpyright-code-budget.json"
-
-
def cmd_update(counts: Mapping[str, int]) -> None:
existing = json.loads(BUDGET_PATH.read_text()) if BUDGET_PATH.exists() else {}
budget = {
@@ -127,9 +176,10 @@ def cmd_update(counts: Mapping[str, int]) -> None:
)
-def cmd_check(counts: Mapping[str, int]) -> None:
+def cmd_check(base_ref: str) -> None:
budget = json.loads(BUDGET_PATH.read_text())
- if is_vacuous_run(counts, budget):
+ head = count_basedpyright(sys.stdin.read())
+ if is_vacuous_run(head, budget):
expected = sum(spec["baseline"] for spec in budget.values())
print(
f"FAIL: basedpyright produced no errors, but {BUDGET_PATH.name} expects "
@@ -137,27 +187,44 @@ def cmd_check(counts: Mapping[str, int]) -> None:
f"nothing; refusing to certify a vacuous run."
)
raise SystemExit(1)
- breaches = evaluate(counts, budget)
+ base_point = _run(["git", "merge-base", base_ref, "HEAD"]).strip() or base_ref
+ base = base_counts(base_point)
+ if is_vacuous_run(base, budget):
+ print(
+ f"FAIL: basedpyright produced no errors for the base tree at "
+ f"{base_point[:12]}, so every rule would look freshly added. The base "
+ f"pass almost certainly crashed; refusing to blame this change for it."
+ )
+ raise SystemExit(1)
+ breaches = evaluate(head, base, budget)
if not breaches:
print(
- f"OK: every rule is within its basedpyright ceiling ({sum(counts.values())} errors total)"
+ f"OK: every rule is within its basedpyright ceiling or no higher than base ({sum(head.values())} errors total)"
)
return
print("FAIL: basedpyright errors exceed the per-rule ceiling:")
for breach in breaches:
- print(f" {breach.code}: {breach.total} errors over cap {breach.cap}")
+ print(
+ f" {breach.code}: total {breach.total} over cap {breach.cap} (this change added {breach.added})"
+ )
print(
- "Resolve the new errors, or run 'make lint-basedpyright-budget-update' if the ceiling should move."
+ "Reduce the new errors or remove an equal number elsewhere; the ceiling is "
+ "baseline + slack in basedpyright-code-budget.json."
)
+ summary = "; ".join(f"{b.code} {b.total}/{b.cap} (+{b.added})" for b in breaches)
+ print(f"BREACHED RULES: {summary}")
raise SystemExit(1)
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
+ parser.add_argument("--base", default=DEFAULT_BASE)
parser.add_argument("--update", action="store_true")
args = parser.parse_args()
- counts = count_basedpyright(sys.stdin.read())
- cmd_update(counts) if args.update else cmd_check(counts)
+ if args.update:
+ cmd_update(count_basedpyright(sys.stdin.read()))
+ else:
+ cmd_check(args.base)
if __name__ == "__main__":
diff --git a/tests/batches_tests/test_batch_custom_pricing.py b/tests/batches_tests/test_batch_custom_pricing.py
index cb2ca385ffc..3dc1d116e8d 100644
--- a/tests/batches_tests/test_batch_custom_pricing.py
+++ b/tests/batches_tests/test_batch_custom_pricing.py
@@ -159,12 +159,12 @@ def test_batch_cost_calculator_applies_data_residency_uplift(
base_prompt, base_completion = batch_cost_calculator(
usage=usage,
- model="gpt-5",
+ model="gpt-5.4",
custom_llm_provider="openai",
)
regional_prompt, regional_completion = batch_cost_calculator(
usage=usage,
- model="gpt-5",
+ model="gpt-5.4",
custom_llm_provider="openai",
data_residency=data_residency,
)
diff --git a/tests/code_coverage_tests/recursive_detector.py b/tests/code_coverage_tests/recursive_detector.py
index 254d700ee5a..1d11d676207 100644
--- a/tests/code_coverage_tests/recursive_detector.py
+++ b/tests/code_coverage_tests/recursive_detector.py
@@ -47,6 +47,7 @@ IGNORE_FUNCTIONS = [
"_read_image_bytes", # max depth set.
"_get_masked_values", # max depth set (default 20) to prevent infinite recursion while masking nested sensitive config dicts.
"_redact_sensitive_litellm_params", # max depth set (default 10).
+ "_redact_secret_values_in_obj", # max depth set (default 10, _REDACT_SECRET_MAX_DEPTH); fails closed by returning "REDACTED" at the cap.
"_resolve", # OCI: $ref resolver bounded by `resolving_stack` cycle guard.
"resolve_oci_schema_anyof", # OCI: bounded by JSON-schema tree depth (no cycles possible in well-formed input).
"sanitize_oci_schema", # OCI: bounded by JSON-schema tree depth.
diff --git a/tests/llm_translation/test_bedrock_completion.py b/tests/llm_translation/test_bedrock_completion.py
index fa22ff6b392..f4c307e9c8a 100644
--- a/tests/llm_translation/test_bedrock_completion.py
+++ b/tests/llm_translation/test_bedrock_completion.py
@@ -2502,19 +2502,34 @@ async def test_bedrock_image_url_sync_client():
mock_post.assert_called_once()
-def test_bedrock_error_handling_streaming():
+@pytest.mark.parametrize(
+ "exception_type, expected_status_code",
+ [
+ ("internalServerException", 500),
+ ("serviceUnavailableException", 503),
+ ("modelTimeoutException", 408),
+ ("modelStreamErrorException", 424),
+ ("validationException", 400),
+ ],
+)
+def test_bedrock_error_handling_streaming(exception_type, expected_status_code):
+ """Bedrock event-stream error events arrive with botocore's hard-coded
+ status_code=400; the decoder must surface the modeled HTTP status instead
+ (e.g. internalServerException -> 500). For 5xx this is what makes the error
+ retryable downstream; for all types it replaces the misleading 400 with the
+ true code. Regression for #24608."""
from litellm.llms.bedrock.chat.invoke_handler import (
AWSEventStreamDecoder,
BedrockError,
)
- from unittest.mock import patch, Mock
+ from unittest.mock import Mock
event = Mock()
event.to_response_dict = Mock(
return_value={
"status_code": 400,
"headers": {
- ":exception-type": "serviceUnavailableException",
+ ":exception-type": exception_type,
":content-type": "application/json",
":message-type": "exception",
},
@@ -2525,11 +2540,10 @@ def test_bedrock_error_handling_streaming():
decoder = AWSEventStreamDecoder(
model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0"
)
- with pytest.raises(Exception) as e:
+ with pytest.raises(BedrockError) as e:
decoder._parse_message_from_event(event)
- assert isinstance(e.value, BedrockError)
assert "Bedrock is unable to process your request." in e.value.message
- assert e.value.status_code == 400
+ assert e.value.status_code == expected_status_code
@pytest.mark.parametrize(
diff --git a/tests/llm_translation/test_bedrock_embedding_pricing.py b/tests/llm_translation/test_bedrock_embedding_pricing.py
new file mode 100644
index 00000000000..099d73fed87
--- /dev/null
+++ b/tests/llm_translation/test_bedrock_embedding_pricing.py
@@ -0,0 +1,34 @@
+"""
+Tests for AWS Bedrock embedding model pricing in the model cost map.
+
+Regression test for the Amazon Titan Text Embeddings V2 commercial price,
+which was previously set 10x too high (2e-07 instead of 2e-08).
+AWS lists Titan Text Embeddings V2 at $0.02 per 1M input tokens
+(= $0.00002 per 1K tokens = 2e-08 per token).
+"""
+
+import importlib
+
+
+class TestBedrockEmbeddingPricing:
+ """Test suite for Bedrock embedding model pricing in the cost map."""
+
+ def test_titan_embed_v2_commercial_input_cost(self, monkeypatch):
+ """Titan Text Embeddings V2 should be priced at $0.02 / 1M tokens (2e-08)."""
+ # Scope the local-cost-map flag to this test only, so it does not leak
+ # into sibling tests. monkeypatch restores the environment on teardown.
+ monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
+
+ import litellm.litellm_core_utils.get_model_cost_map
+ import litellm
+
+ # Reload so the cost map is re-read from the local file with the flag set.
+ importlib.reload(litellm.litellm_core_utils.get_model_cost_map)
+ importlib.reload(litellm)
+
+ model = litellm.model_cost["amazon.titan-embed-text-v2:0"]
+
+ assert model["input_cost_per_token"] == 2e-08
+ assert model["output_cost_per_token"] == 0.0
+ assert model["litellm_provider"] == "bedrock"
+ assert model["mode"] == "embedding"
diff --git a/tests/llm_translation/test_cloudflare.py b/tests/llm_translation/test_cloudflare.py
index 5a6a0008398..54c5d9e4e07 100644
--- a/tests/llm_translation/test_cloudflare.py
+++ b/tests/llm_translation/test_cloudflare.py
@@ -9,9 +9,7 @@ import pytest
from litellm import acompletion, completion
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
-FAKE_API_BASE = (
- "https://fake-cloudflare.example.com/client/v4/accounts/fake-acct/ai/run/"
-)
+FAKE_API_BASE = "https://fake-cloudflare.example.com/client/v4/accounts/fake-acct/ai/v1"
FAKE_API_KEY = "fake-cf-api-key"
@@ -26,28 +24,78 @@ def _make_mock_response(json_data: Dict[str, Any]) -> MagicMock:
def _chat_response() -> Dict[str, Any]:
return {
- "result": {
- "response": "I am a large language model created to assist you.",
- },
- "success": True,
- "errors": [],
- "messages": [],
+ "id": "chatcmpl-cf",
+ "object": "chat.completion",
+ "created": 1234567890,
+ "model": "@cf/meta/llama-2-7b-chat-int8",
+ "choices": [
+ {
+ "index": 0,
+ "message": {
+ "role": "assistant",
+ "content": "I am a large language model created to assist you.",
+ },
+ "finish_reason": "stop",
+ }
+ ],
+ "usage": {"prompt_tokens": 8, "completion_tokens": 11, "total_tokens": 19},
+ }
+
+
+def _tool_call_response() -> Dict[str, Any]:
+ return {
+ "id": "chatcmpl-cf-tools",
+ "object": "chat.completion",
+ "created": 1234567890,
+ "model": "@cf/meta/llama-2-7b-chat-int8",
+ "choices": [
+ {
+ "index": 0,
+ "message": {
+ "role": "assistant",
+ "content": None,
+ "tool_calls": [
+ {
+ "id": "call_1",
+ "type": "function",
+ "function": {
+ "name": "get_weather",
+ "arguments": '{"city": "New York"}',
+ },
+ }
+ ],
+ },
+ "finish_reason": "tool_calls",
+ }
+ ],
+ "usage": {"prompt_tokens": 20, "completion_tokens": 9, "total_tokens": 29},
}
def _streaming_chunks() -> list[str]:
+ base = {
+ "id": "chatcmpl-cf",
+ "object": "chat.completion.chunk",
+ "created": 1234567890,
+ "model": "@cf/meta/llama-2-7b-chat-int8",
+ }
return [
- json.dumps({"response": "I am"}),
- json.dumps({"response": " a language"}),
- json.dumps({"response": " model."}),
- ]
-
-
-def _streaming_chunks_response_text() -> list[str]:
- return [
- json.dumps({"response_text": "I am"}),
- json.dumps({"response_text": " a language"}),
- json.dumps({"response_text": " model."}),
+ json.dumps({**base, "choices": [{"index": 0, "delta": {"content": "I am"}}]}),
+ json.dumps(
+ {**base, "choices": [{"index": 0, "delta": {"content": " a language"}}]}
+ ),
+ json.dumps(
+ {
+ **base,
+ "choices": [
+ {
+ "index": 0,
+ "delta": {"content": " model."},
+ "finish_reason": "stop",
+ }
+ ],
+ }
+ ),
]
@@ -85,6 +133,48 @@ def test_completion_cloudflare(sync_mode):
assert response.choices[0].message.content is not None
assert "language model" in response.choices[0].message.content.lower()
+ called_url = mock_post.call_args.kwargs.get("url") or mock_post.call_args.args[0]
+ assert called_url.endswith("/ai/v1/chat/completions")
+ assert "/ai/run/" not in called_url
+
+
+def test_completion_cloudflare_tool_calls_sent_to_openai_endpoint():
+ messages = [{"role": "user", "content": "weather in New York?"}]
+ tools = [
+ {
+ "type": "function",
+ "function": {
+ "name": "get_weather",
+ "parameters": {
+ "type": "object",
+ "properties": {"city": {"type": "string"}},
+ "required": ["city"],
+ },
+ },
+ }
+ ]
+ mock_resp = _make_mock_response(_tool_call_response())
+
+ with patch.object(HTTPHandler, "post", return_value=mock_resp) as mock_post:
+ response = completion(
+ model="cloudflare/@cf/meta/llama-2-7b-chat-int8",
+ messages=messages,
+ tools=tools,
+ tool_choice="auto",
+ api_base=FAKE_API_BASE,
+ api_key=FAKE_API_KEY,
+ )
+ mock_post.assert_called_once()
+
+ sent_body = json.loads(mock_post.call_args.kwargs["data"])
+ assert sent_body["tools"] == tools
+ assert sent_body["tool_choice"] == "auto"
+
+ assert response.choices[0].finish_reason == "tool_calls"
+ tool_calls = response.choices[0].message.tool_calls
+ assert tool_calls is not None and len(tool_calls) == 1
+ assert tool_calls[0].function.name == "get_weather"
+
@pytest.mark.parametrize("sync_mode", [True, False])
def test_completion_cloudflare_stream(sync_mode):
@@ -153,76 +243,3 @@ def test_completion_cloudflare_stream(sync_mode):
if c.choices[0].delta.content
)
assert "language" in content.lower()
-
-
-@pytest.mark.parametrize("sync_mode", [True, False])
-def test_completion_cloudflare_stream_response_text(sync_mode):
- """Newer Cloudflare Workers AI models (e.g. Nemotron) emit `response_text`
- instead of `response` in streamed chunks. The iterator must surface that
- text so streaming output is not silently empty.
- """
- messages = [{"role": "user", "content": "what llm are you"}]
- raw_chunks = _streaming_chunks_response_text()
-
- if sync_mode:
-
- def _iter_lines():
- for chunk in raw_chunks:
- yield f"data: {chunk}"
- yield "data: [DONE]"
-
- mock_resp = MagicMock()
- mock_resp.iter_lines.return_value = _iter_lines()
- mock_resp.status_code = 200
- mock_resp.headers = {"content-type": "text/event-stream"}
-
- with patch.object(HTTPHandler, "post", return_value=mock_resp) as mock_post:
- response = completion(
- model="cloudflare/@cf/nvidia/nemotron-mini-4b-instruct",
- messages=messages,
- max_tokens=15,
- stream=True,
- api_base=FAKE_API_BASE,
- api_key=FAKE_API_KEY,
- )
- chunks_received = list(response)
- mock_post.assert_called_once()
- else:
-
- async def _aiter_lines():
- for chunk in raw_chunks:
- yield f"data: {chunk}"
- yield "data: [DONE]"
-
- mock_resp = MagicMock()
- mock_resp.aiter_lines.return_value = _aiter_lines()
- mock_resp.status_code = 200
- mock_resp.headers = {"content-type": "text/event-stream"}
-
- async def _run():
- with patch.object(
- AsyncHTTPHandler, "post", new_callable=AsyncMock, return_value=mock_resp
- ) as mock_post:
- resp = await acompletion(
- model="cloudflare/@cf/nvidia/nemotron-mini-4b-instruct",
- messages=messages,
- max_tokens=15,
- stream=True,
- api_base=FAKE_API_BASE,
- api_key=FAKE_API_KEY,
- )
- received = []
- async for chunk in resp:
- received.append(chunk)
- mock_post.assert_called_once()
- return received
-
- chunks_received = asyncio.run(_run())
-
- assert len(chunks_received) > 0
- content = "".join(
- c.choices[0].delta.content
- for c in chunks_received
- if c.choices[0].delta.content
- )
- assert "language" in content.lower()
diff --git a/tests/llm_translation/test_fireworks_ai_translation.py b/tests/llm_translation/test_fireworks_ai_translation.py
index 204f4d9e31b..27059581e4d 100644
--- a/tests/llm_translation/test_fireworks_ai_translation.py
+++ b/tests/llm_translation/test_fireworks_ai_translation.py
@@ -7,9 +7,10 @@ sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
import litellm
-from litellm import transcription
+from litellm.litellm_core_utils.get_supported_openai_params import (
+ get_supported_openai_params,
+)
from litellm.llms.fireworks_ai.chat.transformation import FireworksAIConfig
-from base_audio_transcription_unit_tests import BaseLLMAudioTranscriptionTest
fireworks = FireworksAIConfig()
@@ -69,74 +70,16 @@ def test_map_response_format():
assert result == {"response_format": response_format}
-_AUDIO_FILE_PATH = os.path.join(
- os.path.dirname(os.path.realpath(__file__)), "gettysburg.wav"
-)
-
-
-class TestFireworksAIAudioTranscription(BaseLLMAudioTranscriptionTest):
- def get_base_audio_transcription_call_args(self) -> dict:
- return {
- "model": "fireworks_ai/whisper-v3",
- "api_base": "https://audio-prod.api.fireworks.ai/v1",
- }
-
- def get_custom_llm_provider(self) -> litellm.LlmProviders:
- return litellm.LlmProviders.FIREWORKS_AI
-
- def test_audio_transcription(self):
- from unittest.mock import MagicMock
-
- from openai.types.audio import Transcription
-
- audio_file = open(_AUDIO_FILE_PATH, "rb")
- mock_client = MagicMock()
- mock_client.audio.transcriptions.create.return_value = Transcription(
- text="four score and seven years ago"
- )
-
- transcript = transcription(
- **self.get_base_audio_transcription_call_args(),
- file=audio_file,
- api_key="fw-test-key",
- client=mock_client,
- )
-
- assert transcript.text == "four score and seven years ago"
- sent = mock_client.audio.transcriptions.create.call_args.kwargs
- assert sent["model"] == "whisper-v3"
- assert sent["file"] is audio_file
-
- @pytest.mark.asyncio
- async def test_audio_transcription_async(self):
- from unittest.mock import AsyncMock, MagicMock
-
- from openai.types.audio import Transcription
-
- audio_file = open(_AUDIO_FILE_PATH, "rb")
- raw_response = MagicMock()
- raw_response.headers = {}
- raw_response.parse.return_value = Transcription(
- text="four score and seven years ago"
- )
- mock_client = MagicMock()
- mock_client.audio.transcriptions.with_raw_response.create = AsyncMock(
- return_value=raw_response
- )
-
- transcript = await litellm.atranscription(
- **self.get_base_audio_transcription_call_args(),
- file=audio_file,
- api_key="fw-test-key",
- client=mock_client,
- )
-
- assert transcript.text == "four score and seven years ago"
- sent = (
- mock_client.audio.transcriptions.with_raw_response.create.call_args.kwargs
- )
- assert sent["model"] == "whisper-v3"
- assert sent["file"] is audio_file
+def test_get_supported_openai_params_transcription_returns_none():
+ # Fireworks AI deprecated audio transcription on 2026-06-10; the endpoint
+ # is decommissioned. Returning None (not chat-completion params) signals
+ # to callers that transcription is unsupported for this provider.
+ result = get_supported_openai_params(
+ model="fireworks_ai/accounts/fireworks/models/whisper-v3",
+ custom_llm_provider="fireworks_ai",
+ request_type="transcription",
+ )
+ assert result is None
@pytest.mark.parametrize(
diff --git a/tests/llm_translation/test_prompt_factory.py b/tests/llm_translation/test_prompt_factory.py
index 36e47e3c2f4..05a58a135d2 100644
--- a/tests/llm_translation/test_prompt_factory.py
+++ b/tests/llm_translation/test_prompt_factory.py
@@ -605,8 +605,33 @@ def test_no_messages_yields_user_text():
assert contents == expected_output
-def test_convert_url():
- convert_url_to_base64("https://picsum.photos/id/237/200/300")
+def test_convert_url(monkeypatch):
+ import base64
+ from unittest.mock import MagicMock
+
+ import httpx
+
+ from litellm.litellm_core_utils.prompt_templates.image_handling import (
+ in_memory_cache,
+ )
+
+ url = "https://picsum.photos/id/237/200/300"
+ image_bytes = b"\x89PNG\r\n\x1a\nfake-png-bytes"
+
+ mock_client = MagicMock()
+ mock_client.get.return_value = httpx.Response(
+ 200, content=image_bytes, headers={"Content-Type": "image/png"}
+ )
+
+ monkeypatch.setattr(litellm, "user_url_validation", False, raising=False)
+ monkeypatch.setattr(litellm, "module_level_client", mock_client, raising=False)
+ in_memory_cache.flush_cache()
+
+ result = convert_url_to_base64(url)
+
+ expected = "data:image/png;base64," + base64.b64encode(image_bytes).decode("utf-8")
+ assert result == expected
+ mock_client.get.assert_called_once()
def test_azure_tool_call_invoke_helper():
diff --git a/tests/search_tests/test_searchapi_search.py b/tests/search_tests/test_searchapi_search.py
index 5ef9d922b89..d16868502a4 100644
--- a/tests/search_tests/test_searchapi_search.py
+++ b/tests/search_tests/test_searchapi_search.py
@@ -46,10 +46,9 @@ class TestSearchAPIConfig:
assert result["Content-Type"] == "application/json"
- @patch("litellm.llms.searchapi.search.transformation.get_secret_str")
- def test_validate_environment_without_api_key(self, mock_get_secret):
+ def test_validate_environment_without_api_key(self, monkeypatch):
"""Test environment validation without API key raises error."""
- mock_get_secret.return_value = None
+ monkeypatch.delenv("SEARCHAPI_API_KEY", raising=False)
config = SearchAPIConfig()
headers = {}
diff --git a/tests/search_tests/test_searxng_search.py b/tests/search_tests/test_searxng_search.py
index 45b0f3214d9..c12d44183b0 100644
--- a/tests/search_tests/test_searxng_search.py
+++ b/tests/search_tests/test_searxng_search.py
@@ -318,13 +318,11 @@ class TestSearXNGSearchHeaders:
assert headers["Content-Type"] == "application/json"
assert headers["Authorization"] == "Bearer test-key-123"
- def test_headers_with_env_api_key(self):
+ def test_headers_with_env_api_key(self, monkeypatch):
"""Test that headers use SEARXNG_API_KEY from env."""
- with patch(
- "litellm.llms.searxng.search.transformation.get_secret_str",
- return_value="env-key-456",
- ):
- headers = self.config.validate_environment(headers={})
+ monkeypatch.setenv("SEARXNG_API_KEY", "env-key-456")
+
+ headers = self.config.validate_environment(headers={})
assert headers["Authorization"] == "Bearer env-key-456"
diff --git a/tests/test_litellm/integrations/code_interpreter_interception/test_handler.py b/tests/test_litellm/integrations/code_interpreter_interception/test_handler.py
index f33814b86df..7ff58ba6324 100644
--- a/tests/test_litellm/integrations/code_interpreter_interception/test_handler.py
+++ b/tests/test_litellm/integrations/code_interpreter_interception/test_handler.py
@@ -1,8 +1,8 @@
"""
Unit tests for CodeInterpreterInterceptionLogger.
-All sandbox dependencies are injected (dependency injection, no monkeypatch):
-a FakeSandbox stands in for the real e2b config and records how it is called.
+All sandbox dependencies are injected: a FakeSandbox stands in for the real e2b
+config and records how it is called.
"""
import time
@@ -12,13 +12,17 @@ import pytest
from litellm.integrations.code_interpreter_interception.handler import (
CodeInterpreterInterceptionLogger,
LITELLM_CODE_EXECUTION_TOOL_NAME,
+ _INTERCEPTION_ACTIVE_KEY as _ACTIVE_KEY,
+ _SANDBOX_KEY,
+)
+from litellm.types.integrations.custom_logger import (
+ CHAT_COMPLETION_AGENTIC_SURFACE,
+ NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES,
+ is_interception_internal_key,
)
from litellm.llms.base_llm.sandbox.transformation import CodeExecutionResult
from litellm.types.utils import CallTypes
-_ACTIVE_KEY = "_code_interpreter_interception_active"
-_SANDBOX_KEY = "_code_interpreter_interception_sandbox_key"
-
class FakeHandle:
def __init__(self, sandbox_id="sbx_fake"):
@@ -51,6 +55,13 @@ class FakeLogging:
def __init__(self, litellm_call_id="k1"):
self.litellm_call_id = litellm_call_id
self.model_call_details = {}
+ self.dynamic_success_callbacks = []
+
+ def pre_call(self, *args, **kwargs):
+ return None
+
+ def post_call(self, *args, **kwargs):
+ return None
def _function_call_item(call_id="c1", name=LITELLM_CODE_EXECUTION_TOOL_NAME):
@@ -62,6 +73,17 @@ def _function_call_item(call_id="c1", name=LITELLM_CODE_EXECUTION_TOOL_NAME):
}
+def _chat_function_call_item(call_id="call_1", name=LITELLM_CODE_EXECUTION_TOOL_NAME):
+ return {
+ "id": call_id,
+ "type": "function",
+ "function": {
+ "name": name,
+ "arguments": '{"code":"print(40 + 2)"}',
+ },
+ }
+
+
class FakeResponse:
def __init__(self, output):
self.output = output
@@ -74,6 +96,18 @@ def _iter_messages(plan):
return patch.messages
+def test_interception_internal_key_prefix_sets_preserve_code_interpreter_state():
+ assert is_interception_internal_key("_code_interpreter_interception_active")
+ assert not is_interception_internal_key(
+ "_code_interpreter_interception_active",
+ prefixes=NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES,
+ )
+ assert is_interception_internal_key(
+ "_websearch_interception_converted_stream",
+ prefixes=NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES,
+ )
+
+
@pytest.mark.asyncio
async def test_build_plan_runs_code_and_feeds_output_back():
sandbox = FakeSandbox(stdout="42")
@@ -133,6 +167,30 @@ async def test_pre_call_converts_code_interpreter_tool():
assert LITELLM_CODE_EXECUTION_TOOL_NAME in names
+@pytest.mark.asyncio
+async def test_pre_call_converts_code_interpreter_tool_for_chat_completions():
+ logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
+ kwargs = {
+ "tools": [{"type": "code_interpreter", "container": {"type": "auto"}}],
+ "tool_choice": {"type": "code_interpreter"},
+ "custom_llm_provider": "openai",
+ }
+
+ result = await logger.async_pre_call_deployment_hook(kwargs, CallTypes.acompletion)
+
+ assert result is not None
+ tool = result["tools"][0]
+ assert tool["type"] == "function"
+ assert tool["function"]["name"] == LITELLM_CODE_EXECUTION_TOOL_NAME
+ assert tool["function"]["parameters"]["required"] == ["code"]
+ assert result["tool_choice"] == {
+ "type": "function",
+ "function": {"name": LITELLM_CODE_EXECUTION_TOOL_NAME},
+ }
+ assert result["litellm_metadata"][_ACTIVE_KEY] is True
+ assert result["litellm_metadata"][_SANDBOX_KEY] == result[_SANDBOX_KEY]
+
+
@pytest.mark.asyncio
@pytest.mark.parametrize(
"tool_choice",
@@ -184,9 +242,24 @@ async def test_pre_call_noop_on_non_responses():
"custom_llm_provider": "openai",
}
+ result = await logger.async_pre_call_deployment_hook(kwargs, CallTypes.aembedding)
+
+ assert result is None
+
+
+@pytest.mark.asyncio
+async def test_pre_call_noop_on_chat_completion_without_code_interpreter_tool():
+ logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
+ kwargs = {
+ "tools": [{"type": "web_search"}],
+ "custom_llm_provider": "openai",
+ }
+
result = await logger.async_pre_call_deployment_hook(kwargs, CallTypes.acompletion)
assert result is None
+ assert _ACTIVE_KEY not in kwargs
+ assert _SANDBOX_KEY not in kwargs
@pytest.mark.asyncio
@@ -524,14 +597,142 @@ async def test_gate_rechecks_provider_scope():
assert should_run is False
+@pytest.mark.asyncio
+async def test_chat_completion_gate_detects_code_execution_tool_call():
+ logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
+ response = {
+ "choices": [
+ {"message": {"tool_calls": [_chat_function_call_item(call_id="call_123")]}}
+ ]
+ }
+
+ should_run, payload = await logger.async_should_run_agentic_loop(
+ response=response,
+ model="gpt-5",
+ messages=[{"role": "user", "content": "x"}],
+ tools=[],
+ stream=False,
+ custom_llm_provider="openai",
+ kwargs={
+ _ACTIVE_KEY: True,
+ "_agentic_loop_api_surface": CHAT_COMPLETION_AGENTIC_SURFACE,
+ },
+ )
+
+ assert should_run is True
+ assert payload["tool_calls"][0]["id"] == "call_123"
+ assert payload["tool_calls"][0]["arguments"] == '{"code":"print(40 + 2)"}'
+
+
+@pytest.mark.asyncio
+async def test_chat_completion_gate_refuses_without_server_active_marker():
+ logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
+ response = {"choices": [{"message": {"tool_calls": [_chat_function_call_item()]}}]}
+
+ should_run, payload = await logger.async_should_run_agentic_loop(
+ response=response,
+ model="gpt-5",
+ messages=[{"role": "user", "content": "x"}],
+ tools=[],
+ stream=False,
+ custom_llm_provider="openai",
+ kwargs={"_agentic_loop_api_surface": CHAT_COMPLETION_AGENTIC_SURFACE},
+ )
+
+ assert should_run is False
+ assert payload == {}
+
+
+@pytest.mark.asyncio
+async def test_chat_completion_build_plan_runs_code_and_appends_tool_message():
+ sandbox = FakeSandbox(stdout="42")
+ logger = CodeInterpreterInterceptionLogger(sandbox_config=sandbox)
+ native_chat_tool = {"type": "code_interpreter", "container": {"type": "auto"}}
+
+ plan = await logger.async_build_agentic_loop_plan(
+ tools={
+ "tool_calls": [
+ {
+ "id": "call_1",
+ "name": LITELLM_CODE_EXECUTION_TOOL_NAME,
+ "arguments": '{"code":"print(40 + 2)"}',
+ }
+ ]
+ },
+ model="gpt-5",
+ messages=[{"role": "user", "content": "x"}],
+ response={
+ "choices": [{"message": {"tool_calls": [_chat_function_call_item()]}}]
+ },
+ anthropic_messages_provider_config=None,
+ anthropic_messages_optional_request_params={
+ "tools": [native_chat_tool],
+ "tool_choice": {"type": "code_interpreter", "container": {"type": "auto"}},
+ "temperature": 0,
+ },
+ logging_obj=FakeLogging(litellm_call_id="k1"),
+ stream=False,
+ kwargs={
+ "acompletion": True,
+ "litellm_call_id": "k1",
+ _ACTIVE_KEY: True,
+ _SANDBOX_KEY: "sbxkey1",
+ "_code_interpreter_interception_converted_stream": True,
+ "_agentic_loop_api_surface": CHAT_COMPLETION_AGENTIC_SURFACE,
+ },
+ )
+
+ assert sandbox.run_calls[0]["code"] == "print(40 + 2)"
+ patch = plan.request_patch
+ assert patch is not None
+ assert patch.tools == [
+ {
+ "type": "function",
+ "function": {
+ "name": LITELLM_CODE_EXECUTION_TOOL_NAME,
+ "description": "Execute python code in a sandbox and return stdout.",
+ "parameters": {
+ "type": "object",
+ "properties": {"code": {"type": "string"}},
+ "required": ["code"],
+ },
+ },
+ }
+ ]
+ assert patch.optional_params == {"temperature": 0}
+ assert patch.kwargs == {
+ "litellm_call_id": "k1",
+ _ACTIVE_KEY: True,
+ _SANDBOX_KEY: "sbxkey1",
+ "_code_interpreter_interception_converted_stream": True,
+ "_agentic_loop_api_surface": CHAT_COMPLETION_AGENTIC_SURFACE,
+ }
+ assert patch.messages is not None
+ assert patch.messages[-2]["role"] == "assistant"
+ assert patch.messages[-2]["tool_calls"][0]["id"] == "call_1"
+ assert patch.messages[-1] == {
+ "role": "tool",
+ "tool_call_id": "call_1",
+ "content": "42",
+ }
+ assert plan.metadata["code_interpreter_calls"][0]["code"] == "print(40 + 2)"
+
+
@pytest.mark.asyncio
async def test_pre_call_strips_client_forged_marker_on_initial_request():
- """A client cannot pre-set the active marker on the original request."""
+ """A client cannot pre-set the active marker on the original request: with no
+ native code_interpreter tool, any client-supplied interception markers in
+ litellm_metadata are scrubbed and the active flag in kwargs is cleared."""
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
kwargs = {
"tools": [{"type": "web_search"}],
"custom_llm_provider": "openai",
_ACTIVE_KEY: True,
+ "litellm_metadata": {
+ _ACTIVE_KEY: True,
+ _SANDBOX_KEY: "client-forged",
+ "safe_user_value": "kept",
+ },
}
await logger.async_pre_call_deployment_hook(kwargs, CallTypes.aresponses)
@@ -540,6 +741,42 @@ async def test_pre_call_strips_client_forged_marker_on_initial_request():
"no native code_interpreter tool was present, so a client-supplied "
"active marker must be cleared"
)
+ assert kwargs["litellm_metadata"] == {"safe_user_value": "kept"}
+
+
+@pytest.mark.asyncio
+async def test_pre_call_strips_forged_loop_controls_then_mints_own_markers():
+ """On an INITIAL request (no server-set _agentic_loop_depth) a client cannot
+ smuggle loop-control state: forged _agentic_loop_depth / max_agentic_loops and
+ interception markers in litellm_metadata are stripped before the interceptor
+ activates, so the only interception markers that survive are the ones the
+ server mints for the converted code_interpreter tool."""
+ logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
+ kwargs = {
+ "tools": [{"type": "code_interpreter", "container": {"type": "auto"}}],
+ "custom_llm_provider": "openai",
+ "litellm_metadata": {
+ _ACTIVE_KEY: True,
+ _SANDBOX_KEY: "client-forged",
+ "_agentic_loop_depth": 99,
+ "max_agentic_loops": 999,
+ "safe_user_value": "kept",
+ },
+ }
+
+ result = await logger.async_pre_call_deployment_hook(kwargs, CallTypes.acompletion)
+
+ assert result is not None
+ metadata = result["litellm_metadata"]
+ assert metadata["safe_user_value"] == "kept"
+ assert "_agentic_loop_depth" not in metadata, "forged loop depth must be stripped"
+ assert "max_agentic_loops" not in metadata, "forged loop cap must be stripped"
+ assert metadata[_ACTIVE_KEY] is True
+ assert metadata[_SANDBOX_KEY] == result[_SANDBOX_KEY]
+ assert metadata[_SANDBOX_KEY] != "client-forged", (
+ "the surviving sandbox key must be the server-minted one, not the forged "
+ "value the client supplied"
+ )
@pytest.mark.asyncio
diff --git a/tests/test_litellm/interactions/test_google_interactions_integration.py b/tests/test_litellm/interactions/test_google_interactions_integration.py
index cfff26d51ef..9c651cc94f5 100644
--- a/tests/test_litellm/interactions/test_google_interactions_integration.py
+++ b/tests/test_litellm/interactions/test_google_interactions_integration.py
@@ -299,7 +299,6 @@ class TestGoogleInteractionsResponseStructure:
assert hasattr(response, "outputs")
assert hasattr(response, "usage")
assert hasattr(response, "model") or hasattr(response, "agent")
- assert hasattr(response, "role")
assert hasattr(response, "created")
assert hasattr(response, "updated")
diff --git a/tests/test_litellm/interactions/test_openapi_compliance.py b/tests/test_litellm/interactions/test_openapi_compliance.py
index 44c0b6e5b02..209e99895db 100644
--- a/tests/test_litellm/interactions/test_openapi_compliance.py
+++ b/tests/test_litellm/interactions/test_openapi_compliance.py
@@ -162,7 +162,8 @@ class TestResponseCompliance:
# Keep this aligned with the live spec.
schema = spec_dict["components"]["schemas"]["Interaction"]
- # Output fields (readOnly).
+ # Output fields (readOnly). `role` was removed from the `Interaction`
+ # schema by Google; it now lives only on `Turn`.
output_fields = [
"id",
"status",
diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py
index 9b3152fae07..7f3d5a959a1 100644
--- a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py
+++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py
@@ -1499,19 +1499,20 @@ def _local_model_cost_map():
@pytest.mark.parametrize("data_residency", ["eu", "us"])
def test_data_residency_applies_uplift(data_residency, _local_model_cost_map):
- """gpt-5 should apply the regional processing uplift multiplier when
- data_residency is set."""
+ """gpt-5.4 should apply the regional processing uplift multiplier when
+ data_residency is set. gpt-5.4+ (released 2026-03-05) carry the 10% uplift;
+ gpt-5 and older models do not."""
from litellm.types.utils import Usage
usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500)
base = generic_cost_per_token(
- model="gpt-5",
+ model="gpt-5.4",
usage=usage,
custom_llm_provider="openai",
)
regional = generic_cost_per_token(
- model="gpt-5",
+ model="gpt-5.4",
usage=usage,
custom_llm_provider="openai",
data_residency=data_residency,
@@ -1526,6 +1527,23 @@ def test_data_residency_applies_uplift(data_residency, _local_model_cost_map):
assert regional[1] == pytest.approx(base[1] * 1.10, rel=1e-9)
+@pytest.mark.parametrize("model", ["gpt-5", "gpt-5-mini", "gpt-5-nano", "gpt-5-pro", "gpt-4o", "gpt-4.1"])
+def test_data_residency_no_uplift_for_pre_march_2026_models(model, _local_model_cost_map):
+ """Models released before 2026-03-05 must not have the regional uplift."""
+ from litellm.types.utils import Usage
+
+ usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500)
+
+ base = generic_cost_per_token(model=model, usage=usage, custom_llm_provider="openai")
+ regional = generic_cost_per_token(
+ model=model, usage=usage, custom_llm_provider="openai", data_residency="eu"
+ )
+
+ assert base == regional, (
+ f"{model} should not have a regional uplift, but cost changed with data_residency"
+ )
+
+
def test_data_residency_no_uplift_for_unmarked_model(_local_model_cost_map):
"""A model without a regional_processing_uplift_multiplier_* entry should
fall back to base pricing, not error."""
@@ -1555,12 +1573,12 @@ def test_data_residency_none_no_uplift(_local_model_cost_map):
usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500)
base = generic_cost_per_token(
- model="gpt-5",
+ model="gpt-5.4",
usage=usage,
custom_llm_provider="openai",
)
explicit_none = generic_cost_per_token(
- model="gpt-5",
+ model="gpt-5.4",
usage=usage,
custom_llm_provider="openai",
data_residency=None,
@@ -1576,13 +1594,13 @@ def test_data_residency_composes_with_service_tier(_local_model_cost_map):
usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500)
priority_base = generic_cost_per_token(
- model="gpt-5",
+ model="gpt-5.4",
usage=usage,
custom_llm_provider="openai",
service_tier="priority",
)
priority_eu = generic_cost_per_token(
- model="gpt-5",
+ model="gpt-5.4",
usage=usage,
custom_llm_provider="openai",
service_tier="priority",
diff --git a/tests/test_litellm/litellm_core_utils/test_chat_completion_agentic_loop.py b/tests/test_litellm/litellm_core_utils/test_chat_completion_agentic_loop.py
new file mode 100644
index 00000000000..f1196ab4692
--- /dev/null
+++ b/tests/test_litellm/litellm_core_utils/test_chat_completion_agentic_loop.py
@@ -0,0 +1,422 @@
+"""
+Tests for the provider-agnostic chat completion agentic loop dispatcher
+(`litellm/litellm_core_utils/chat_completion_agentic_loop.py`) and the
+code-interpreter interception integration that drives it.
+
+The load-bearing regression here protects a reviewer requirement: the internal
+agentic/interception control fields must NEVER reach the outbound provider HTTP
+request body. The relevant fields are:
+
+ _agentic_loop_depth
+ _agentic_loop_fingerprints
+ _agentic_loop_api_surface
+ max_agentic_loops
+ _code_interpreter_interception_active
+ _code_interpreter_interception_sandbox_key
+ _code_interpreter_interception_converted_stream
+
+A scrubber in gpt_transformation.py used to strip these. That scrubber was
+removed, so `test_internal_control_fields_never_leak_into_provider_body` proves
+they stay out of the body even without it.
+"""
+
+import os
+import sys
+from typing import Any, Dict, List, Optional, Tuple
+from unittest.mock import AsyncMock, MagicMock, patch
+
+import pytest
+
+sys.path.insert(0, os.path.abspath("../../../.."))
+
+import litellm
+from litellm.integrations.custom_logger import CustomLogger
+from litellm.integrations.code_interpreter_interception.handler import (
+ CodeInterpreterInterceptionLogger,
+)
+from litellm.litellm_core_utils.chat_completion_agentic_loop import (
+ maybe_run_chat_completion_agentic_loop,
+)
+from litellm.types.integrations.custom_logger import (
+ AgenticLoopPlan,
+ AgenticLoopRequestPatch,
+)
+from litellm.types.utils import (
+ Choices,
+ Function,
+ ChatCompletionMessageToolCall,
+ Message,
+ ModelResponse,
+)
+
+# The internal control fields that must never reach a provider request body.
+_INTERNAL_CONTROL_FIELDS = (
+ "_agentic_loop_depth",
+ "_agentic_loop_fingerprints",
+ "_agentic_loop_api_surface",
+ "max_agentic_loops",
+ "_code_interpreter_interception_active",
+ "_code_interpreter_interception_sandbox_key",
+ "_code_interpreter_interception_converted_stream",
+ "litellm_metadata",
+)
+
+
+@pytest.fixture
+def restore_callbacks():
+ """Save/restore litellm.callbacks so a registered fake logger never pollutes
+ other tests in the suite."""
+ saved = list(litellm.callbacks)
+ try:
+ yield
+ finally:
+ litellm.callbacks = saved
+
+
+class _SandboxResult:
+ def __init__(self, stdout: str) -> None:
+ self.stdout = stdout
+ self.error = None
+
+
+class FakeSandboxConfig:
+ """Injected sandbox so the interception loop runs no real network / E2B."""
+
+ def __init__(self) -> None:
+ self.created = 0
+ self.deleted = 0
+ self.run_codes: List[str] = []
+
+ async def acreate_sandbox(self) -> Any:
+ self.created += 1
+ return MagicMock(id="sandbox-123")
+
+ async def arun_code(self, container: Any, code: str) -> _SandboxResult:
+ self.run_codes.append(code)
+ return _SandboxResult(stdout="42\n")
+
+ async def adelete_sandbox(self, container: Any) -> None:
+ self.deleted += 1
+
+
+def _tool_call_model_response() -> ModelResponse:
+ return ModelResponse(
+ choices=[
+ Choices(
+ finish_reason="tool_calls",
+ message=Message(
+ role="assistant",
+ content=None,
+ tool_calls=[
+ ChatCompletionMessageToolCall(
+ id="call_abc",
+ type="function",
+ function=Function(
+ name="litellm_code_execution",
+ arguments='{"code": "print(6*7)"}',
+ ),
+ )
+ ],
+ ),
+ )
+ ]
+ )
+
+
+def _plain_model_response(content: str = "The answer is 42") -> ModelResponse:
+ return ModelResponse(
+ choices=[
+ Choices(
+ finish_reason="stop",
+ message=Message(role="assistant", content=content),
+ )
+ ]
+ )
+
+
+def _raw_response_for(model_response: ModelResponse) -> MagicMock:
+ """Wrap a ModelResponse as the OpenAI `with_raw_response.create` return value
+ (an object exposing `.headers` and `.parse()` -> something with model_dump)."""
+ parsed = MagicMock()
+ parsed.model_dump.return_value = model_response.model_dump()
+ raw = MagicMock()
+ raw.headers = {}
+ raw.parse.return_value = parsed
+ return raw
+
+
+# ---------------------------------------------------------------------------
+# A) PROVIDER-PAYLOAD REGRESSION
+# ---------------------------------------------------------------------------
+
+
+@pytest.mark.asyncio
+async def test_internal_control_fields_never_leak_into_provider_body(restore_callbacks):
+ """Drive a real acompletion with a native code_interpreter tool through the
+ interception logger + agentic loop, capturing every outbound OpenAI request
+ body. None of the internal control fields may appear at top-level or inside
+ extra_body on ANY of the captured calls."""
+ logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandboxConfig())
+ litellm.callbacks = [logger]
+
+ # First create -> model emits a code_execution tool call (triggers the loop).
+ # Second create -> model returns a plain answer (loop terminates).
+ create = AsyncMock(
+ side_effect=[
+ _raw_response_for(_tool_call_model_response()),
+ _raw_response_for(_plain_model_response()),
+ ]
+ )
+ mock_client = MagicMock()
+ mock_client.chat.completions.with_raw_response.create = create
+
+ response = await litellm.acompletion(
+ model="openai/gpt-4o-mini",
+ messages=[{"role": "user", "content": "what is 6*7?"}],
+ tools=[{"type": "code_interpreter"}],
+ tool_choice={"type": "code_interpreter"},
+ api_key="sk-test",
+ client=mock_client,
+ )
+
+ # The loop must have actually fired (sanity: two provider calls).
+ assert create.await_count == 2, (
+ "expected the agentic loop to issue a follow-up provider call; "
+ f"got {create.await_count} call(s)"
+ )
+
+ for idx, call in enumerate(create.await_args_list):
+ body = call.kwargs
+ extra_body = body.get("extra_body") or {}
+ for field in _INTERNAL_CONTROL_FIELDS:
+ assert field not in body, (
+ f"provider call #{idx}: internal field {field!r} leaked into "
+ f"top-level request body: {sorted(body.keys())}"
+ )
+ assert field not in extra_body, (
+ f"provider call #{idx}: internal field {field!r} leaked into "
+ f"extra_body: {sorted(extra_body.keys())}"
+ )
+ # The native code_interpreter tool must have been swapped for the
+ # function tool, never sent raw to OpenAI as a chat-completions request.
+ for tool in body.get("tools") or []:
+ assert tool.get("type") != "code_interpreter"
+
+ # The final response is the post-loop answer, not the tool-call turn.
+ assert response.choices[0].message.content == "The answer is 42"
+
+
+# ---------------------------------------------------------------------------
+# B) DISPATCHER UNIT TESTS
+# ---------------------------------------------------------------------------
+
+
+class _LoggingStub:
+ """Minimal logging_obj: dispatcher only reads dynamic_success_callbacks and
+ litellm_call_id off it."""
+
+ litellm_call_id = "call-test"
+ dynamic_success_callbacks: List[Any] = []
+
+
+class _GateOnlyLogger(CustomLogger):
+ """Overrides the gate to fire, but builds a plan from request_patch."""
+
+ def __init__(self, plan: AgenticLoopPlan, tool_calls: Dict[str, Any]) -> None:
+ super().__init__()
+ self._plan = plan
+ self._tool_calls = tool_calls
+ self.cleanup_calls = 0
+
+ async def async_should_run_agentic_loop(
+ self,
+ response: Any,
+ model: str,
+ messages: List[Dict[str, Any]],
+ tools: Optional[List[Dict[str, Any]]],
+ stream: bool,
+ custom_llm_provider: str,
+ kwargs: Dict[str, Any],
+ ) -> Tuple[bool, Dict[str, Any]]:
+ return True, self._tool_calls
+
+ async def async_build_agentic_loop_plan(
+ self,
+ tools: Dict[str, Any],
+ model: str,
+ messages: List[Dict[str, Any]],
+ response: Any,
+ anthropic_messages_provider_config: Any,
+ anthropic_messages_optional_request_params: Dict[str, Any],
+ logging_obj: Any,
+ stream: bool,
+ kwargs: Dict[str, Any],
+ ) -> AgenticLoopPlan:
+ return self._plan
+
+ async def async_agentic_loop_cleanup_hook(
+ self, plan: AgenticLoopPlan, kwargs: Dict[str, Any]
+ ) -> None:
+ self.cleanup_calls += 1
+
+
+def _patched_messages() -> List[Dict[str, Any]]:
+ return [
+ {"role": "user", "content": "what is 6*7?"},
+ {
+ "role": "assistant",
+ "tool_calls": [
+ {
+ "id": "call_abc",
+ "type": "function",
+ "function": {
+ "name": "litellm_code_execution",
+ "arguments": '{"code": "print(6*7)"}',
+ },
+ }
+ ],
+ },
+ {"role": "tool", "tool_call_id": "call_abc", "content": "42\n"},
+ ]
+
+
+@pytest.mark.asyncio
+async def test_dispatcher_returns_none_when_no_callback_gates(restore_callbacks):
+ """No callback overrides the gate -> dispatcher returns None so the caller
+ keeps the original response untouched."""
+ litellm.callbacks = []
+
+ result = await maybe_run_chat_completion_agentic_loop(
+ response=_plain_model_response(),
+ model="gpt-4o-mini",
+ messages=[{"role": "user", "content": "hi"}],
+ optional_params={},
+ kwargs={},
+ logging_obj=_LoggingStub(),
+ custom_llm_provider="openai",
+ stream=False,
+ )
+
+ assert result is None
+
+
+@pytest.mark.asyncio
+async def test_dispatcher_runs_followup_with_incremented_depth_and_patched_messages(
+ restore_callbacks,
+):
+ """A gating logger with a request_patch -> the dispatcher calls
+ litellm.acompletion exactly once with _agentic_loop_depth == 1 and the
+ patched messages. Loop-control state rides as litellm-level kwargs and is
+ mirrored into litellm_metadata; the provider-surface transient
+ _agentic_loop_api_surface is never forwarded. (Provider-body stripping of
+ these litellm-level kwargs is asserted separately in test A.)"""
+ followup = _plain_model_response("done")
+ plan = AgenticLoopPlan(
+ run_agentic_loop=True,
+ request_patch=AgenticLoopRequestPatch(messages=_patched_messages()),
+ )
+ logger = _GateOnlyLogger(plan=plan, tool_calls={"tool_calls": [{"id": "call_abc"}]})
+ litellm.callbacks = [logger]
+
+ acompletion_mock = AsyncMock(return_value=followup)
+ with patch.object(litellm, "acompletion", acompletion_mock):
+ result = await maybe_run_chat_completion_agentic_loop(
+ response=_tool_call_model_response(),
+ model="gpt-4o-mini",
+ messages=[{"role": "user", "content": "what is 6*7?"}],
+ optional_params={"temperature": 0.1},
+ kwargs={"_code_interpreter_interception_active": True},
+ logging_obj=_LoggingStub(),
+ custom_llm_provider="openai",
+ stream=False,
+ )
+
+ assert result is followup
+ acompletion_mock.assert_awaited_once()
+ call_kwargs = acompletion_mock.await_args.kwargs
+
+ assert call_kwargs["_agentic_loop_depth"] == 1
+ assert call_kwargs["messages"] == _patched_messages()
+ # Preserved non-internal optional param survives the rerun.
+ assert call_kwargs["temperature"] == 0.1
+ # Loop-control state is carried at the litellm level for the follow-up.
+ assert call_kwargs["max_agentic_loops"] >= 1
+ assert "_agentic_loop_fingerprints" in call_kwargs
+ # Interception markers are mirrored into litellm_metadata for the follow-up.
+ assert (
+ call_kwargs["litellm_metadata"]["_code_interpreter_interception_active"] is True
+ )
+ # The transient surface marker is NOT forwarded to the follow-up call.
+ assert "_agentic_loop_api_surface" not in call_kwargs
+ # Cleanup hook always runs.
+ assert logger.cleanup_calls == 1
+
+
+@pytest.mark.asyncio
+async def test_dispatcher_raises_when_depth_reaches_max_agentic_loops(
+ restore_callbacks,
+):
+ """depth >= max_agentic_loops -> ValueError mentioning max_agentic_loops,
+ before any follow-up call is attempted."""
+ logger = _GateOnlyLogger(
+ plan=AgenticLoopPlan(run_agentic_loop=True),
+ tool_calls={"tool_calls": [{"id": "call_abc"}]},
+ )
+ litellm.callbacks = [logger]
+
+ acompletion_mock = AsyncMock()
+ with patch.object(litellm, "acompletion", acompletion_mock):
+ with pytest.raises(ValueError, match="max_agentic_loops"):
+ await maybe_run_chat_completion_agentic_loop(
+ response=_tool_call_model_response(),
+ model="gpt-4o-mini",
+ messages=[{"role": "user", "content": "hi"}],
+ optional_params={},
+ kwargs={"_agentic_loop_depth": 3, "max_agentic_loops": 3},
+ logging_obj=_LoggingStub(),
+ custom_llm_provider="openai",
+ stream=False,
+ )
+
+ acompletion_mock.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+async def test_dispatcher_raises_on_repeated_tool_call_fingerprint(restore_callbacks):
+ """A tool_calls fingerprint already present in _agentic_loop_fingerprints ->
+ ValueError about the repeated fingerprint (cycle guard), with no follow-up
+ call."""
+ import json
+
+ # The dispatcher fingerprints the whole value the gate returns as its second
+ # tuple element, so the seeded fingerprint must mirror that dict exactly.
+ gate_tool_calls = {
+ "tool_calls": [{"id": "call_abc", "name": "litellm_code_execution"}]
+ }
+ fingerprint = json.dumps(gate_tool_calls, sort_keys=True, default=str)
+
+ logger = _GateOnlyLogger(
+ plan=AgenticLoopPlan(run_agentic_loop=True),
+ tool_calls=gate_tool_calls,
+ )
+ litellm.callbacks = [logger]
+
+ acompletion_mock = AsyncMock()
+ with patch.object(litellm, "acompletion", acompletion_mock):
+ with pytest.raises(ValueError, match="fingerprint"):
+ await maybe_run_chat_completion_agentic_loop(
+ response=_tool_call_model_response(),
+ model="gpt-4o-mini",
+ messages=[{"role": "user", "content": "hi"}],
+ optional_params={},
+ kwargs={
+ "_agentic_loop_depth": 0,
+ "max_agentic_loops": 3,
+ "_agentic_loop_fingerprints": [fingerprint],
+ },
+ logging_obj=_LoggingStub(),
+ custom_llm_provider="openai",
+ stream=False,
+ )
+
+ acompletion_mock.assert_not_awaited()
diff --git a/tests/test_litellm/litellm_core_utils/test_get_supported_openai_params.py b/tests/test_litellm/litellm_core_utils/test_get_supported_openai_params.py
index 3c280c6ba92..84900e3f2ed 100644
--- a/tests/test_litellm/litellm_core_utils/test_get_supported_openai_params.py
+++ b/tests/test_litellm/litellm_core_utils/test_get_supported_openai_params.py
@@ -132,3 +132,17 @@ def test_azure_base_model_detection_preserved():
assert params is not None
assert "reasoning_effort" in params
assert "tools" in params
+
+
+def test_sambanova_embeddings_request_returns_list_not_none():
+ """The sambanova embeddings branch resolved the config but dropped the result,
+ so embedding requests got ``None`` instead of the supported-params list while the
+ chat branch returned correctly. A list (the sambanova embeddings config exposes no
+ extra params, hence ``[]``) must reach the caller."""
+ embedding_params = get_supported_openai_params(
+ model="E5-Mistral-7B-Instruct",
+ custom_llm_provider="sambanova",
+ request_type="embeddings",
+ )
+
+ assert embedding_params == []
diff --git a/tests/test_litellm/litellm_core_utils/test_request_timeout_resolver.py b/tests/test_litellm/litellm_core_utils/test_request_timeout_resolver.py
new file mode 100644
index 00000000000..4e016622f1c
--- /dev/null
+++ b/tests/test_litellm/litellm_core_utils/test_request_timeout_resolver.py
@@ -0,0 +1,58 @@
+"""Unit tests for litellm.litellm_core_utils.request_timeout_resolver.
+
+The resolver decides whether ``litellm.request_timeout`` was *explicitly configured*
+(env REQUEST_TIMEOUT / litellm_settings, or a non-default runtime value) versus left
+at the package default. This is what lets request_timeout act as an independent
+per-attempt timeout instead of being indistinguishable from "nobody set it".
+"""
+
+import os
+import sys
+
+import pytest
+
+sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../..")))
+
+import litellm
+from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS
+from litellm.litellm_core_utils.request_timeout_resolver import (
+ get_configured_request_timeout,
+)
+
+
+@pytest.fixture
+def restore_request_timeout():
+ original_value = litellm.request_timeout
+ original_flag = litellm.request_timeout_explicitly_set
+ try:
+ yield
+ finally:
+ litellm.request_timeout = original_value
+ litellm.request_timeout_explicitly_set = original_flag
+
+
+def test_default_value_without_flag_is_unset(restore_request_timeout):
+ litellm.request_timeout = DEFAULT_REQUEST_TIMEOUT_SECONDS
+ litellm.request_timeout_explicitly_set = False
+ assert get_configured_request_timeout() is None
+
+
+def test_explicit_flag_returns_value(restore_request_timeout):
+ litellm.request_timeout = 300
+ litellm.request_timeout_explicitly_set = True
+ assert get_configured_request_timeout() == 300.0
+
+
+def test_explicit_flag_preserves_value_equal_to_default(restore_request_timeout):
+ # The case the bare ``!= default`` heuristic gets wrong: a user who explicitly
+ # configures the default value still means it explicitly.
+ litellm.request_timeout = DEFAULT_REQUEST_TIMEOUT_SECONDS
+ litellm.request_timeout_explicitly_set = True
+ assert get_configured_request_timeout() == float(DEFAULT_REQUEST_TIMEOUT_SECONDS)
+
+
+def test_non_default_runtime_value_treated_as_explicit(restore_request_timeout):
+ # SDK users assigning litellm.request_timeout directly (no flag) must keep working.
+ litellm.request_timeout = 300
+ litellm.request_timeout_explicitly_set = False
+ assert get_configured_request_timeout() == 300.0
diff --git a/tests/test_litellm/litellm_core_utils/test_sensitive_data_masker.py b/tests/test_litellm/litellm_core_utils/test_sensitive_data_masker.py
index 6808c4821c1..7239636fd48 100644
--- a/tests/test_litellm/litellm_core_utils/test_sensitive_data_masker.py
+++ b/tests/test_litellm/litellm_core_utils/test_sensitive_data_masker.py
@@ -126,6 +126,49 @@ def test_lists_with_sensitive_keys_are_masked():
assert masked["tags"] == ["prod", "test"]
+def test_short_secrets_are_fully_masked():
+ """
+ Regression test: secrets at or below the reveal threshold (visible_prefix +
+ visible_suffix, 8 by default) were returned verbatim instead of masked.
+ An exactly-8-char value hit masked_length == 0 and round-tripped unchanged;
+ anything shorter hit the early return. Both leaked short credentials (e.g. an
+ 8-char redis password) in plaintext through mask_dict.
+ """
+ masker = SensitiveDataMasker()
+
+ # Boundary: exactly 8 chars previously returned verbatim.
+ assert masker._mask_value("abcd1234") == "********"
+ # Below threshold previously hit the early return and leaked verbatim.
+ assert masker._mask_value("sk-12") == "*****"
+ # Values above the threshold must still partially reveal, not over-mask.
+ assert masker._mask_value("abcd12345") == "abcd*2345"
+
+ masked = masker.mask_dict({"redis_password": "pass1234", "api_key": "sk-7a"})
+ assert masked["redis_password"] == "********"
+ assert masked["api_key"] == "*****"
+
+
+def test_mask_short_values_false_keeps_short_values_readable():
+ """
+ mask_short_values=False opts out of full masking so short values are returned
+ as-is. This preserves the truncation use (e.g. CooldownCache shows the first 50
+ chars of an exception and only masks longer tails), while longer values are still
+ partially masked.
+ """
+ masker = SensitiveDataMasker(
+ visible_prefix=50, visible_suffix=0, mask_short_values=False
+ )
+
+ short = "Test exception for structure validation"
+ assert masker._mask_value(short) == short
+
+ long_value = "x" * 60
+ masked = masker._mask_value(long_value)
+ assert masked.startswith("x" * 50)
+ assert masked.endswith("*" * 10)
+ assert len(masked) == 60
+
+
def test_cost_per_token_fields_not_masked():
"""
Regression test: cost fields like input_cost_per_token contain "token" in their name
diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py
index 734e49161c4..81af0ad3e6f 100644
--- a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py
+++ b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py
@@ -878,6 +878,114 @@ def test_sync_streaming_bad_request_not_midstream(logging_obj: Logging):
assert "invalid maxOutputTokens" in str(excinfo.value)
+def _bedrock_error_event(exception_type: str):
+ """A mocked botocore event-stream error event: status_code is botocore's
+ hard-coded 400, with the real type in the :exception-type header."""
+ event = Mock()
+ event.to_response_dict = Mock(
+ return_value={
+ "status_code": 400,
+ "headers": {
+ ":exception-type": exception_type,
+ ":content-type": "application/json",
+ ":message-type": "exception",
+ },
+ "body": b'{"message":"Bedrock had an internal error."}',
+ }
+ )
+ return event
+
+
+@pytest.mark.asyncio
+async def test_bedrock_midstream_internal_server_error_wraps_for_fallback(
+ logging_obj: Logging,
+):
+ """End-to-end regression for https://github.com/BerriAI/litellm/issues/24608:
+ a Bedrock mid-stream internalServerException event (botocore stamps it 400)
+ must flow through the real decoder, gain its modeled 500 status, and wrap
+ into MidStreamFallbackError so the Router can run streaming fallback.
+
+ Calls the real AWSEventStreamDecoder, so reverting the decoder status fix
+ makes the decoder raise BedrockError(400) and the gate raises BadRequestError
+ directly -> this test fails without the fix."""
+ from litellm.exceptions import MidStreamFallbackError
+ from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder
+
+ decoder = AWSEventStreamDecoder(model="anthropic.claude-3-sonnet-20240229-v1:0")
+
+ async def _bedrock_stream():
+ decoder._parse_message_from_event(
+ _bedrock_error_event("internalServerException")
+ )
+ yield # unreachable; the line above raises
+
+ async def _make_call(**kwargs):
+ return _bedrock_stream()
+
+ response = CustomStreamWrapper(
+ completion_stream=None,
+ model="anthropic.claude-3-sonnet-20240229-v1:0",
+ logging_obj=logging_obj,
+ custom_llm_provider="bedrock",
+ make_call=_make_call,
+ )
+
+ with pytest.raises(MidStreamFallbackError):
+ await response.__anext__()
+
+
+@pytest.mark.asyncio
+async def test_bedrock_5xx_wraps_for_midstream_fallback(logging_obj: Logging):
+ """Gate contract: a Bedrock 5xx (here 503 serviceUnavailableException) wraps
+ into MidStreamFallbackError so the Router can run streaming fallback."""
+ from litellm.exceptions import MidStreamFallbackError
+ from litellm.llms.bedrock.chat.invoke_handler import BedrockError
+
+ async def _raise_503(**kwargs):
+ raise BedrockError(
+ status_code=503,
+ message="serviceUnavailableException Bedrock is unavailable.",
+ )
+
+ response = CustomStreamWrapper(
+ completion_stream=None,
+ model="anthropic.claude-3-sonnet-20240229-v1:0",
+ logging_obj=logging_obj,
+ custom_llm_provider="bedrock",
+ make_call=_raise_503,
+ )
+
+ with pytest.raises(MidStreamFallbackError):
+ await response.__anext__()
+
+
+@pytest.mark.asyncio
+async def test_bedrock_validation_error_raises_directly(logging_obj: Logging):
+ """Gate contract: a Bedrock validationException (400) is a client error and
+ must surface directly, never wrapped into MidStreamFallbackError."""
+ from litellm.exceptions import MidStreamFallbackError
+ from litellm.llms.bedrock.chat.invoke_handler import BedrockError
+
+ async def _raise_400(**kwargs):
+ raise BedrockError(
+ status_code=400,
+ message="validationException malformed input.",
+ )
+
+ response = CustomStreamWrapper(
+ completion_stream=None,
+ model="anthropic.claude-3-sonnet-20240229-v1:0",
+ logging_obj=logging_obj,
+ custom_llm_provider="bedrock",
+ make_call=_raise_400,
+ )
+
+ with pytest.raises(Exception) as excinfo:
+ await response.__anext__()
+ assert not isinstance(excinfo.value, MidStreamFallbackError)
+ assert getattr(excinfo.value, "status_code", None) == 400
+
+
@pytest.mark.asyncio
async def test_async_streaming_read_timeout_triggers_midstream_fallback(
logging_obj: Logging,
diff --git a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py
index 2876b56f516..b111b65e3af 100644
--- a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py
+++ b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py
@@ -1886,6 +1886,89 @@ def test_anthropic_model_supports_effort_param_rejects_non_supporting_models(mod
assert AnthropicConfig._model_supports_effort_param(model) is False
+@pytest.mark.parametrize(
+ "model",
+ [
+ "claude-opus-4-6",
+ "claude-opus-4-7",
+ "claude-opus-4-8",
+ "claude-opus-4-6-20260205",
+ "claude-opus-4-7-20260416",
+ ],
+)
+def test_anthropic_model_supports_speed_param_recognizes_supporting_models(model):
+ assert AnthropicConfig._model_supports_speed_param(model) is True
+
+
+@pytest.mark.parametrize(
+ "model",
+ [
+ "claude-sonnet-4-6",
+ "claude-fable-5",
+ "claude-3-haiku-20240307",
+ "vertex_ai/claude-opus-4-8",
+ "azure_ai/claude-opus-4-8",
+ "anthropic.claude-opus-4-8",
+ ],
+)
+def test_anthropic_model_supports_speed_param_rejects_non_supporting_models(model):
+ assert AnthropicConfig._model_supports_speed_param(model) is False
+
+
+@pytest.mark.parametrize("custom_llm_provider", ["vertex_ai", "azure_ai", "bedrock"])
+def test_anthropic_model_supports_speed_param_rejects_non_anthropic_providers(
+ custom_llm_provider,
+):
+ """Fast mode is direct-Anthropic-only. Vertex/Azure/Bedrock strip their prefix
+ before the shared transform runs, so the bare Opus id must still be rejected."""
+ assert (
+ AnthropicConfig._model_supports_speed_param(
+ "claude-opus-4-8", custom_llm_provider
+ )
+ is False
+ )
+ assert (
+ AnthropicConfig._model_supports_speed_param("claude-opus-4-8", "anthropic")
+ is True
+ )
+
+
+def test_vertex_anthropic_drops_speed_for_opus_with_drop_params(monkeypatch):
+ """Regression: vertex_ai Opus must drop ``speed`` even though the prefix-stripped
+ ``claude-opus-4-8`` maps to a fast-mode-capable direct-Anthropic entry."""
+ from litellm.llms.vertex_ai.vertex_ai_partner_models.anthropic.transformation import (
+ VertexAIAnthropicConfig,
+ )
+
+ monkeypatch.setattr(litellm, "drop_params", True)
+ result = VertexAIAnthropicConfig().transform_request(
+ model="claude-opus-4-8",
+ messages=[{"role": "user", "content": "Hello"}],
+ optional_params={"speed": "fast", "max_tokens": 1024},
+ litellm_params={},
+ headers={},
+ )
+
+ assert "speed" not in result
+
+
+def test_vertex_anthropic_raises_on_speed_without_drop_params(monkeypatch):
+ """Regression: vertex_ai Opus raises rather than forwarding an unsupported
+ ``speed`` when neither global nor per-request drop_params is set."""
+ from litellm.llms.vertex_ai.vertex_ai_partner_models.anthropic.transformation import (
+ VertexAIAnthropicConfig,
+ )
+
+ monkeypatch.setattr(litellm, "drop_params", False)
+ with pytest.raises(litellm.utils.UnsupportedParamsError, match="drop_params"):
+ VertexAIAnthropicConfig().map_openai_params(
+ non_default_params={"speed": "fast"},
+ optional_params={},
+ model="claude-opus-4-8",
+ drop_params=False,
+ )
+
+
def test_translate_system_message_skips_empty_string_content():
"""
Test that translate_system_message skips system messages with empty string content.
@@ -3766,6 +3849,61 @@ def test_fast_mode_parameter_mapping():
assert result["speed"] == "fast"
+def test_anthropic_drop_params_strips_speed_for_unsupported_models():
+ """``drop_params=True`` strips unsupported ``speed`` for non-Opus models."""
+ config = AnthropicConfig()
+ messages = [{"role": "user", "content": "Hello"}]
+
+ original = litellm.drop_params
+ litellm.drop_params = True
+ try:
+ result = config.transform_request(
+ model="claude-sonnet-4-6",
+ messages=messages,
+ optional_params={"speed": "fast", "max_tokens": 1024},
+ litellm_params={},
+ headers={},
+ )
+ finally:
+ litellm.drop_params = original
+
+ assert "speed" not in result
+
+
+def test_anthropic_drop_params_keeps_speed_for_supporting_models():
+ """``drop_params=True`` must not strip ``speed`` on Opus fast-mode models."""
+ config = AnthropicConfig()
+ messages = [{"role": "user", "content": "Hello"}]
+
+ original = litellm.drop_params
+ litellm.drop_params = True
+ try:
+ result = config.transform_request(
+ model="claude-opus-4-6",
+ messages=messages,
+ optional_params={"speed": "fast", "max_tokens": 1024},
+ litellm_params={},
+ headers={},
+ )
+ finally:
+ litellm.drop_params = original
+
+ assert result.get("speed") == "fast"
+
+
+def test_speed_raises_clean_error_without_drop_params(monkeypatch):
+ monkeypatch.setattr(litellm, "drop_params", False)
+ config = AnthropicConfig()
+
+ with pytest.raises(litellm.utils.UnsupportedParamsError, match="drop_params"):
+ config.map_openai_params(
+ non_default_params={"speed": "fast"},
+ optional_params={},
+ model="claude-sonnet-4-6",
+ drop_params=False,
+ )
+
+
def test_map_openai_params_max_tokens_normalized_to_int():
"""
Test that map_openai_params normalizes max_tokens to an integer (e.g. 0.7 -> 1).
diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_speed.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_speed.py
new file mode 100644
index 00000000000..6900f1062bf
--- /dev/null
+++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_speed.py
@@ -0,0 +1,117 @@
+import litellm
+import pytest
+from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
+ AnthropicMessagesConfig,
+)
+from litellm.llms.anthropic.experimental_pass_through.messages.utils import (
+ AnthropicMessagesRequestUtils,
+)
+
+
+def test_messages_drop_params_strips_speed_for_unsupported_models():
+ original = litellm.drop_params
+ litellm.drop_params = True
+ try:
+ optional_params = AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param(
+ params={
+ "max_tokens": 1024,
+ "speed": "fast",
+ "messages": [{"role": "user", "content": "Hello"}],
+ },
+ model="claude-sonnet-4-6",
+ drop_params=False,
+ )
+ config = AnthropicMessagesConfig()
+ headers, _ = config.validate_anthropic_messages_environment(
+ headers={},
+ model="claude-sonnet-4-6",
+ messages=[{"role": "user", "content": "Hello"}],
+ optional_params=dict(optional_params),
+ litellm_params={},
+ )
+ result = config.transform_anthropic_messages_request(
+ model="claude-sonnet-4-6",
+ messages=[{"role": "user", "content": "Hello"}],
+ anthropic_messages_optional_request_params=dict(optional_params),
+ litellm_params={},
+ headers=headers,
+ )
+ finally:
+ litellm.drop_params = original
+
+ assert "speed" not in optional_params
+ assert "speed" not in result
+ assert "fast-mode-2026-02-01" not in headers.get("anthropic-beta", "")
+
+
+def test_messages_drop_params_keeps_speed_for_supporting_models():
+ original = litellm.drop_params
+ litellm.drop_params = True
+ try:
+ optional_params = AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param(
+ params={"max_tokens": 1024, "speed": "fast"},
+ model="claude-opus-4-6",
+ drop_params=False,
+ )
+ config = AnthropicMessagesConfig()
+ headers, _ = config.validate_anthropic_messages_environment(
+ headers={},
+ model="claude-opus-4-6",
+ messages=[{"role": "user", "content": "Hello"}],
+ optional_params=dict(optional_params),
+ litellm_params={},
+ )
+ result = config.transform_anthropic_messages_request(
+ model="claude-opus-4-6",
+ messages=[{"role": "user", "content": "Hello"}],
+ anthropic_messages_optional_request_params=dict(optional_params),
+ litellm_params={},
+ headers=headers,
+ )
+ finally:
+ litellm.drop_params = original
+
+ assert optional_params.get("speed") == "fast"
+ assert result.get("speed") == "fast"
+ assert "fast-mode-2026-02-01" in headers.get("anthropic-beta", "")
+
+
+def test_messages_raises_when_speed_unsupported_and_drop_params_false(monkeypatch):
+ monkeypatch.setattr(litellm, "drop_params", False)
+
+ with pytest.raises(litellm.utils.UnsupportedParamsError, match="drop_params"):
+ AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param(
+ params={"max_tokens": 1024, "speed": "fast"},
+ model="claude-sonnet-4-6",
+ drop_params=False,
+ )
+
+
+def test_messages_drops_speed_for_vertex_opus_with_drop_params(monkeypatch):
+ """Regression: a vertex_ai Opus passthrough must drop ``speed`` even though the
+ prefix-stripped model id maps to a fast-mode-capable direct-Anthropic entry."""
+ monkeypatch.setattr(litellm, "drop_params", True)
+ optional_params = (
+ AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param(
+ params={"max_tokens": 1024, "speed": "fast"},
+ model="claude-opus-4-8",
+ drop_params=False,
+ custom_llm_provider="vertex_ai",
+ )
+ )
+
+ assert "speed" not in optional_params
+
+
+def test_messages_raises_for_vertex_opus_without_drop_params(monkeypatch):
+ """Regression: vertex_ai Opus passthrough raises rather than forwarding an
+ unsupported ``speed`` when drop_params is unset."""
+ monkeypatch.setattr(litellm, "drop_params", False)
+
+ with pytest.raises(litellm.utils.UnsupportedParamsError, match="drop_params"):
+ AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param(
+ params={"max_tokens": 1024, "speed": "fast"},
+ model="claude-opus-4-8",
+ drop_params=False,
+ custom_llm_provider="vertex_ai",
+ )
diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_request_optional_param_utils.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_request_optional_param_utils.py
index 3ce076640e8..f0252e13336 100644
--- a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_request_optional_param_utils.py
+++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_request_optional_param_utils.py
@@ -6,6 +6,7 @@ Regression tests for the /v1/messages request-parse fast paths:
while resolving the (static) type hints only once per process.
"""
+import litellm
from litellm.llms.anthropic.experimental_pass_through.messages.utils import (
AnthropicMessagesRequestUtils,
_anthropic_messages_optional_param_keys,
@@ -54,3 +55,36 @@ def test_empty_params():
)
== {}
)
+
+
+def test_drop_params_strips_speed_for_unsupported_model():
+ original = litellm.drop_params
+ litellm.drop_params = True
+ try:
+ result = (
+ AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param(
+ params={"speed": "fast", "temperature": 0.5},
+ model="claude-sonnet-4-6",
+ )
+ )
+ finally:
+ litellm.drop_params = original
+
+ assert result == {"temperature": 0.5}
+ assert "speed" not in result
+
+
+def test_drop_params_keeps_speed_for_supporting_model():
+ original = litellm.drop_params
+ litellm.drop_params = True
+ try:
+ result = (
+ AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param(
+ params={"speed": "fast"},
+ model="claude-opus-4-6",
+ )
+ )
+ finally:
+ litellm.drop_params = original
+
+ assert result == {"speed": "fast"}
diff --git a/tests/test_litellm/llms/apiserpent/test_apiserpent_search.py b/tests/test_litellm/llms/apiserpent/test_apiserpent_search.py
index 32838701949..bc26268ee92 100644
--- a/tests/test_litellm/llms/apiserpent/test_apiserpent_search.py
+++ b/tests/test_litellm/llms/apiserpent/test_apiserpent_search.py
@@ -66,9 +66,8 @@ class TestAPISerpentConfig:
assert headers["X-API-Key"] == "test-api-key"
assert headers["Content-Type"] == "application/json"
- @patch("litellm.llms.apiserpent.search.transformation.get_secret_str")
- def test_validate_environment_without_api_key(self, mock_get_secret):
- mock_get_secret.return_value = None
+ def test_validate_environment_without_api_key(self, monkeypatch):
+ monkeypatch.delenv("APISERPENT_API_KEY", raising=False)
with pytest.raises(ValueError, match="APISERPENT_API_KEY is not set"):
APISerpentSearchConfig().validate_environment({})
diff --git a/tests/test_litellm/llms/base_llm/search/test_base_search_transformation.py b/tests/test_litellm/llms/base_llm/search/test_base_search_transformation.py
new file mode 100644
index 00000000000..a1353d57038
--- /dev/null
+++ b/tests/test_litellm/llms/base_llm/search/test_base_search_transformation.py
@@ -0,0 +1,329 @@
+"""
+Regression tests for the host-aware server-credential fallback guard in
+``BaseSearchConfig``.
+
+A caller-supplied ``api_base`` is honored when building the request URL, so
+falling back to a server-configured secret while the caller controls the host
+would send the operator's credential to an attacker. The guard must refuse that
+combination for every provider that carries a server-managed secret, while
+leaving keyless providers and legitimate operator overrides untouched.
+"""
+
+from typing import Dict, Tuple, Type
+from unittest.mock import AsyncMock, patch
+
+import pytest
+
+import litellm
+from litellm.llms.apiserpent.search.transformation import APISerpentSearchConfig
+from litellm.llms.base_llm.search.transformation import (
+ BaseSearchConfig,
+ _is_trusted_search_api_base,
+)
+from litellm.llms.brave.search.transformation import BraveSearchConfig
+from litellm.llms.dataforseo.search.transformation import DataForSEOSearchConfig
+from litellm.llms.exa_ai.search.transformation import ExaAISearchConfig
+from litellm.llms.fastcrw.search.transformation import FastCRWSearchConfig
+from litellm.llms.firecrawl.search.transformation import FirecrawlSearchConfig
+from litellm.llms.google_pse.search.transformation import GooglePSESearchConfig
+from litellm.llms.linkup.search.transformation import LinkupSearchConfig
+from litellm.llms.parallel_ai.search.transformation import ParallelAISearchConfig
+from litellm.llms.perplexity.search.transformation import PerplexitySearchConfig
+from litellm.llms.searchapi.search.transformation import SearchAPIConfig
+from litellm.llms.searxng.search.transformation import SearXNGSearchConfig
+from litellm.llms.serper.search.transformation import SerperSearchConfig
+from litellm.llms.tavily.search.transformation import TavilySearchConfig
+from litellm.llms.tinyfish.search.transformation import TinyfishSearchConfig
+from litellm.llms.you_com.search.transformation import YouComSearchConfig
+
+ATTACKER_BASE = "https://attacker.example.com"
+
+# Every *_API_BASE override env var that could otherwise mark the attacker host
+# as trusted; cleared before each test so the suite is hermetic.
+_BASE_ENV_VARS = (
+ "SERPER_API_BASE",
+ "TAVILY_API_BASE",
+ "PERPLEXITY_API_BASE",
+ "APISERPENT_API_BASE",
+ "EXA_API_BASE",
+ "BRAVE_API_BASE",
+ "FIRECRAWL_API_BASE",
+ "LINKUP_API_BASE",
+ "SEARCHAPI_API_BASE",
+ "GOOGLE_PSE_API_BASE",
+ "PARALLEL_AI_API_BASE",
+ "YOUCOM_API_BASE",
+ "SEARXNG_API_BASE",
+ "DATAFORSEO_API_BASE",
+ "TINYFISH_API_BASE",
+ "CRW_API_BASE",
+)
+
+
+@pytest.fixture(autouse=True)
+def _clear_base_overrides(monkeypatch: pytest.MonkeyPatch) -> None:
+ for var in _BASE_ENV_VARS:
+ monkeypatch.delenv(var, raising=False)
+
+
+# (config, {server secret env vars}, caller_api_key honored as-is, extra env for full validate)
+ProviderSpec = Tuple[Type[BaseSearchConfig], Dict[str, str], str, Dict[str, str]]
+
+PROVIDERS: Tuple[ProviderSpec, ...] = (
+ (SerperSearchConfig, {"SERPER_API_KEY": "srv"}, "caller-key", {}),
+ (TavilySearchConfig, {"TAVILY_API_KEY": "srv"}, "caller-key", {}),
+ (PerplexitySearchConfig, {"PERPLEXITYAI_API_KEY": "srv"}, "caller-key", {}),
+ (APISerpentSearchConfig, {"APISERPENT_API_KEY": "srv"}, "caller-key", {}),
+ (ExaAISearchConfig, {"EXA_API_KEY": "srv"}, "caller-key", {}),
+ (BraveSearchConfig, {"BRAVE_API_KEY": "srv"}, "caller-key", {}),
+ (FirecrawlSearchConfig, {"FIRECRAWL_API_KEY": "srv"}, "caller-key", {}),
+ (LinkupSearchConfig, {"LINKUP_API_KEY": "srv"}, "caller-key", {}),
+ (SearchAPIConfig, {"SEARCHAPI_API_KEY": "srv"}, "caller-key", {}),
+ (
+ GooglePSESearchConfig,
+ {"GOOGLE_PSE_API_KEY": "srv"},
+ "caller-key",
+ {"GOOGLE_PSE_ENGINE_ID": "engine"},
+ ),
+ (ParallelAISearchConfig, {"PARALLEL_API_KEY": "srv"}, "caller-key", {}),
+ (YouComSearchConfig, {"YOUCOM_API_KEY": "srv"}, "caller-key", {}),
+ (SearXNGSearchConfig, {"SEARXNG_API_KEY": "srv"}, "caller-key", {}),
+ (
+ DataForSEOSearchConfig,
+ {"DATAFORSEO_LOGIN": "srv", "DATAFORSEO_PASSWORD": "pw"},
+ "login:password",
+ {},
+ ),
+ (TinyfishSearchConfig, {"TINYFISH_API_KEY": "srv"}, "caller-key", {}),
+ (FastCRWSearchConfig, {"CRW_API_KEY": "srv"}, "caller-key", {}),
+)
+
+_IDS = tuple(spec[0].__name__ for spec in PROVIDERS)
+
+
+@pytest.mark.parametrize(
+ "config_cls, server_env, caller_key, extra_env", PROVIDERS, ids=_IDS
+)
+def test_server_secret_refused_for_caller_api_base(
+ config_cls: Type[BaseSearchConfig],
+ server_env: Dict[str, str],
+ caller_key: str,
+ extra_env: Dict[str, str],
+ monkeypatch: pytest.MonkeyPatch,
+) -> None:
+ for key, value in {**server_env, **extra_env}.items():
+ monkeypatch.setenv(key, value)
+
+ with pytest.raises(ValueError, match="Refusing to send the server-configured"):
+ config_cls().validate_environment(headers={}, api_base=ATTACKER_BASE)
+
+
+@pytest.mark.parametrize(
+ "config_cls, server_env, caller_key, extra_env", PROVIDERS, ids=_IDS
+)
+def test_caller_supplied_key_is_honored_for_custom_api_base(
+ config_cls: Type[BaseSearchConfig],
+ server_env: Dict[str, str],
+ caller_key: str,
+ extra_env: Dict[str, str],
+ monkeypatch: pytest.MonkeyPatch,
+) -> None:
+ for key, value in {**server_env, **extra_env}.items():
+ monkeypatch.setenv(key, value)
+
+ # An explicit caller key is the caller's own credential, so pointing it at
+ # the caller's own host must be allowed.
+ config_cls().validate_environment(
+ headers={}, api_key=caller_key, api_base=ATTACKER_BASE
+ )
+
+
+@pytest.mark.parametrize(
+ "config_cls, server_env, caller_key, extra_env", PROVIDERS, ids=_IDS
+)
+def test_server_secret_used_without_caller_api_base(
+ config_cls: Type[BaseSearchConfig],
+ server_env: Dict[str, str],
+ caller_key: str,
+ extra_env: Dict[str, str],
+ monkeypatch: pytest.MonkeyPatch,
+) -> None:
+ for key, value in {**server_env, **extra_env}.items():
+ monkeypatch.setenv(key, value)
+
+ # No caller-supplied api_base -> the request targets the trusted default, so
+ # the server secret is still used and nothing is refused.
+ config_cls().validate_environment(headers={})
+
+
+def test_keyless_provider_allows_caller_api_base(
+ monkeypatch: pytest.MonkeyPatch,
+) -> None:
+ monkeypatch.delenv("SEARXNG_API_KEY", raising=False)
+
+ headers = SearXNGSearchConfig().validate_environment(
+ headers={}, api_base="https://my-searxng.internal"
+ )
+
+ assert "Authorization" not in headers
+
+
+def test_operator_env_base_override_is_trusted(
+ monkeypatch: pytest.MonkeyPatch,
+) -> None:
+ monkeypatch.setenv("SERPER_API_KEY", "srv")
+ monkeypatch.setenv("SERPER_API_BASE", "https://serper.internal.corp")
+
+ # Mirrors the second validate_environment call in the search handler, which
+ # receives the already-resolved operator base as api_base.
+ headers = SerperSearchConfig().validate_environment(
+ headers={}, api_base="https://serper.internal.corp/search"
+ )
+
+ assert headers["X-API-KEY"] == "srv"
+
+
+class TestResolveServerApiKey:
+ def test_caller_key_short_circuits(self) -> None:
+ result = BaseSearchConfig().resolve_server_api_key(
+ caller_api_key="mine",
+ caller_api_base=ATTACKER_BASE,
+ key_env_vars=("SERPER_API_KEY",),
+ base_env_var="SERPER_API_BASE",
+ default_api_base="https://google.serper.dev",
+ )
+ assert result == "mine"
+
+ def test_returns_none_when_no_server_secret(
+ self, monkeypatch: pytest.MonkeyPatch
+ ) -> None:
+ monkeypatch.delenv("SEARXNG_API_KEY", raising=False)
+ result = BaseSearchConfig().resolve_server_api_key(
+ caller_api_key=None,
+ caller_api_base=ATTACKER_BASE,
+ key_env_vars=("SEARXNG_API_KEY",),
+ base_env_var="SEARXNG_API_BASE",
+ default_api_base=None,
+ )
+ assert result is None
+
+ def test_first_set_env_var_wins(self, monkeypatch: pytest.MonkeyPatch) -> None:
+ monkeypatch.delenv("PARALLEL_AI_API_KEY", raising=False)
+ monkeypatch.setenv("PARALLEL_API_KEY", "second")
+ result = BaseSearchConfig().resolve_server_api_key(
+ caller_api_key=None,
+ caller_api_base=None,
+ key_env_vars=("PARALLEL_AI_API_KEY", "PARALLEL_API_KEY"),
+ base_env_var="PARALLEL_AI_API_BASE",
+ default_api_base="https://api.parallel.ai",
+ )
+ assert result == "second"
+
+
+class TestIsTrustedSearchApiBase:
+ def test_matches_default_host(self) -> None:
+ assert _is_trusted_search_api_base(
+ "https://google.serper.dev/search", "https://google.serper.dev", None
+ )
+
+ def test_foreign_host_untrusted(self) -> None:
+ assert not _is_trusted_search_api_base(
+ ATTACKER_BASE, "https://google.serper.dev", None
+ )
+
+ def test_env_override_host_trusted(self, monkeypatch: pytest.MonkeyPatch) -> None:
+ monkeypatch.setenv("SERPER_API_BASE", "https://serper.internal.corp")
+ assert _is_trusted_search_api_base(
+ "https://serper.internal.corp/search",
+ "https://google.serper.dev",
+ "SERPER_API_BASE",
+ )
+
+ def test_schemeless_candidate_untrusted(self) -> None:
+ # Without a scheme urlsplit puts the value in the path, leaving an empty
+ # netloc; an unparseable host must never be treated as trusted.
+ assert not _is_trusted_search_api_base(
+ "attacker.example.com", "https://google.serper.dev", None
+ )
+
+
+@pytest.mark.asyncio
+async def test_asearch_does_not_leak_server_key_to_caller_api_base(
+ monkeypatch: pytest.MonkeyPatch,
+) -> None:
+ """End-to-end regression on the reported vector: a search call with a foreign
+ api_base and no caller key must fail without any outbound request carrying the
+ server-configured key."""
+ monkeypatch.setenv("SERPER_API_KEY", "sk-server-secret")
+ monkeypatch.delenv("SERPER_API_BASE", raising=False)
+
+ with (
+ patch(
+ "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
+ new_callable=AsyncMock,
+ ) as mock_post,
+ patch(
+ "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get",
+ new_callable=AsyncMock,
+ ) as mock_get,
+ ):
+ with pytest.raises(Exception):
+ await litellm.asearch(
+ query="secrets",
+ search_provider="serper",
+ api_base=ATTACKER_BASE,
+ )
+
+ mock_post.assert_not_called()
+ mock_get.assert_not_called()
+
+
+@pytest.mark.parametrize(
+ "provider, key_env, server_key, extra_env",
+ [
+ ("searchapi", "SEARCHAPI_API_KEY", "sk-server-searchapi", {}),
+ (
+ "google_pse",
+ "GOOGLE_PSE_API_KEY",
+ "sk-server-google",
+ {"GOOGLE_PSE_ENGINE_ID": "engine-id"},
+ ),
+ ],
+)
+@pytest.mark.asyncio
+async def test_query_param_key_not_leaked_with_dummy_caller_key(
+ provider: str,
+ key_env: str,
+ server_key: str,
+ extra_env: Dict[str, str],
+ monkeypatch: pytest.MonkeyPatch,
+) -> None:
+ """Providers that send the key as a URL query param resolve it in
+ transform_search_request, not validate_environment. A caller who passes a
+ dummy api_key to clear the validate_environment short-circuit must not cause
+ the server key to be placed in the URL sent to their own api_base."""
+ monkeypatch.setenv(key_env, server_key)
+ for name, value in extra_env.items():
+ monkeypatch.setenv(name, value)
+
+ captured: Dict[str, str] = {}
+
+ async def fake_get(self, *args, **kwargs): # type: ignore[no-untyped-def]
+ captured["url"] = kwargs.get("url") or (args[0] if args else "")
+ raise RuntimeError("stop after capturing the outbound url")
+
+ with patch(
+ "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get",
+ fake_get,
+ ):
+ with pytest.raises(Exception):
+ await litellm.asearch(
+ query="secrets",
+ search_provider=provider,
+ api_key="sk-CALLER-DUMMY",
+ api_base=ATTACKER_BASE,
+ )
+
+ assert captured["url"], "expected an outbound request to be attempted"
+ assert server_key not in captured["url"]
+ assert "sk-CALLER-DUMMY" in captured["url"]
diff --git a/tests/test_litellm/llms/bedrock/test_mantle.py b/tests/test_litellm/llms/bedrock/test_mantle.py
index bbefdd621f0..f7f8f582abc 100644
--- a/tests/test_litellm/llms/bedrock/test_mantle.py
+++ b/tests/test_litellm/llms/bedrock/test_mantle.py
@@ -125,6 +125,74 @@ def test_mantle_messages_url_construction():
assert url == "https://bedrock-mantle.us-east-1.api.aws/anthropic/v1/messages"
+_VPC_ENDPOINT = "https://vpce-0a1b2c3d.bedrock-mantle.us-gov-west-1.vpce.amazonaws.com"
+
+
+def test_mantle_chat_url_honors_api_base_host():
+ config = AmazonMantleConfig()
+ url = config.get_complete_url(
+ api_base=_VPC_ENDPOINT,
+ api_key=None,
+ model="mantle/anthropic.claude-mythos-preview",
+ optional_params={"aws_region_name": "us-gov-west-1"},
+ litellm_params={},
+ )
+ assert url == f"{_VPC_ENDPOINT}/anthropic/v1/messages"
+
+
+def test_mantle_chat_url_honors_api_base_full_path_without_duplication():
+ config = AmazonMantleConfig()
+ full = f"{_VPC_ENDPOINT}/anthropic/v1/messages"
+ url = config.get_complete_url(
+ api_base=full,
+ api_key=None,
+ model="mantle/anthropic.claude-mythos-preview",
+ optional_params={"aws_region_name": "us-gov-west-1"},
+ litellm_params={},
+ )
+ assert url == full
+
+
+def test_mantle_messages_url_honors_api_base_host():
+ config = AmazonMantleMessagesConfig()
+ url = config.get_complete_url(
+ api_base=_VPC_ENDPOINT,
+ api_key=None,
+ model="mantle/anthropic.claude-mythos-preview",
+ optional_params={"aws_region_name": "us-gov-west-1"},
+ litellm_params={},
+ )
+ assert url == f"{_VPC_ENDPOINT}/anthropic/v1/messages"
+ assert "api.aws" not in url
+
+
+def test_mantle_messages_url_honors_api_base_with_trailing_slash():
+ config = AmazonMantleMessagesConfig()
+ url = config.get_complete_url(
+ api_base=f"{_VPC_ENDPOINT}/",
+ api_key=None,
+ model="mantle/anthropic.claude-mythos-preview",
+ optional_params={"aws_region_name": "us-gov-west-1"},
+ litellm_params={},
+ )
+ assert url == f"{_VPC_ENDPOINT}/anthropic/v1/messages"
+
+
+def test_mantle_messages_url_honors_aws_bedrock_runtime_endpoint():
+ config = AmazonMantleMessagesConfig()
+ url = config.get_complete_url(
+ api_base=None,
+ api_key=None,
+ model="mantle/anthropic.claude-mythos-preview",
+ optional_params={
+ "aws_region_name": "us-gov-west-1",
+ "aws_bedrock_runtime_endpoint": _VPC_ENDPOINT,
+ },
+ litellm_params={},
+ )
+ assert url == f"{_VPC_ENDPOINT}/anthropic/v1/messages"
+
+
def test_mantle_transform_request_strips_prefix_and_adds_model():
config = AmazonMantleConfig()
request = config.transform_request(
@@ -247,3 +315,35 @@ async def test_mantle_anthropic_messages_sends_workspace_header_and_clean_body()
assert requests[0]["path"] == "/anthropic/v1/messages"
assert requests[0]["headers"]["anthropic-workspace"] == "proj_abc123def456"
assert "aws_bedrock_project_id" not in requests[0]["body"]
+
+
+@pytest.mark.asyncio
+async def test_mantle_anthropic_messages_routes_to_vpc_api_base():
+ import litellm
+
+ urls = []
+
+ async def mock_post(self, url, data=None, headers=None, **kwargs):
+ urls.append(str(url))
+ return _anthropic_response(str(url))
+
+ try:
+ with patch(
+ "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
+ new=mock_post,
+ ):
+ await litellm.anthropic_messages(
+ model="bedrock/mantle/anthropic.claude-mythos-preview",
+ messages=[{"role": "user", "content": "hello"}],
+ max_tokens=10,
+ api_base=_VPC_ENDPOINT,
+ aws_access_key_id="fake-key",
+ aws_secret_access_key="fake-secret",
+ aws_region_name="us-gov-west-1",
+ )
+ finally:
+ await litellm.close_litellm_async_clients()
+
+ assert len(urls) == 1
+ assert urls[0] == f"{_VPC_ENDPOINT}/anthropic/v1/messages"
+ assert "api.aws" not in urls[0]
diff --git a/tests/test_litellm/llms/cloudflare/test_cloudflare_transformation.py b/tests/test_litellm/llms/cloudflare/test_cloudflare_transformation.py
index cecb6024de1..1a46015ceb9 100644
--- a/tests/test_litellm/llms/cloudflare/test_cloudflare_transformation.py
+++ b/tests/test_litellm/llms/cloudflare/test_cloudflare_transformation.py
@@ -3,25 +3,190 @@ import pytest
from litellm.llms.cloudflare.chat.transformation import CloudflareChatConfig
-def test_get_complete_url_encodes_model_path_segment():
+def test_supported_params_include_tools_and_tool_choice():
config = CloudflareChatConfig()
- assert (
- config.get_complete_url(
- api_base="https://api.cloudflare.com/client/v4/accounts/acct/ai/run/",
- api_key="cf-key",
- model="@cf/meta/llama?x=1#frag",
- optional_params={},
- litellm_params={},
- )
- == "https://api.cloudflare.com/client/v4/accounts/acct/ai/run/%40cf/meta/llama%3Fx%3D1%23frag"
+ params = config.get_supported_openai_params(model="@cf/meta/llama-2-7b-chat-int8")
+
+ assert "tools" in params
+ assert "tool_choice" in params
+ assert "stream" in params
+ assert "max_tokens" in params
+
+
+def test_get_complete_url_defaults_to_openai_compatible_endpoint(monkeypatch):
+ monkeypatch.setenv("CLOUDFLARE_ACCOUNT_ID", "acct")
+ config = CloudflareChatConfig()
+
+ url = config.get_complete_url(
+ api_base=None,
+ api_key="cf-key",
+ model="@cf/meta/llama-2-7b-chat-int8",
+ optional_params={},
+ litellm_params={},
)
- with pytest.raises(ValueError, match="dot path segment"):
+ assert (
+ url
+ == "https://api.cloudflare.com/client/v4/accounts/acct/ai/v1/chat/completions"
+ )
+ assert "/ai/run/" not in url
+
+
+def test_get_complete_url_appends_chat_completions_to_explicit_base():
+ config = CloudflareChatConfig()
+
+ url = config.get_complete_url(
+ api_base="https://api.cloudflare.com/client/v4/accounts/acct/ai/v1",
+ api_key="cf-key",
+ model="@cf/meta/llama-2-7b-chat-int8",
+ optional_params={},
+ litellm_params={},
+ )
+
+ assert (
+ url
+ == "https://api.cloudflare.com/client/v4/accounts/acct/ai/v1/chat/completions"
+ )
+ assert "/ai/run/" not in url
+
+
+def test_get_complete_url_is_idempotent_for_full_base():
+ config = CloudflareChatConfig()
+
+ url = config.get_complete_url(
+ api_base="https://api.cloudflare.com/client/v4/accounts/acct/ai/v1/chat/completions",
+ api_key="cf-key",
+ model="@cf/meta/llama-2-7b-chat-int8",
+ optional_params={},
+ litellm_params={},
+ )
+
+ assert (
+ url
+ == "https://api.cloudflare.com/client/v4/accounts/acct/ai/v1/chat/completions"
+ )
+
+
+def test_get_complete_url_falls_back_to_account_id_when_base_is_empty(monkeypatch):
+ monkeypatch.setenv("CLOUDFLARE_ACCOUNT_ID", "acct")
+ config = CloudflareChatConfig()
+
+ url = config.get_complete_url(
+ api_base="",
+ api_key="cf-key",
+ model="@cf/meta/llama-2-7b-chat-int8",
+ optional_params={},
+ litellm_params={},
+ )
+
+ assert (
+ url
+ == "https://api.cloudflare.com/client/v4/accounts/acct/ai/v1/chat/completions"
+ )
+
+
+def test_get_complete_url_raises_when_account_id_and_base_missing(monkeypatch):
+ monkeypatch.delenv("CLOUDFLARE_ACCOUNT_ID", raising=False)
+ config = CloudflareChatConfig()
+
+ with pytest.raises(ValueError, match="Missing CLOUDFLARE_ACCOUNT_ID"):
config.get_complete_url(
- api_base="https://api.cloudflare.com/client/v4/accounts/acct/ai/run/",
+ api_base=None,
api_key="cf-key",
- model="../../accounts/other",
+ model="@cf/meta/llama-2-7b-chat-int8",
optional_params={},
litellm_params={},
)
+
+
+def test_get_complete_url_raises_when_account_id_is_empty(monkeypatch):
+ monkeypatch.setenv("CLOUDFLARE_ACCOUNT_ID", " ")
+ config = CloudflareChatConfig()
+
+ with pytest.raises(ValueError, match="Missing CLOUDFLARE_ACCOUNT_ID"):
+ config.get_complete_url(
+ api_base=None,
+ api_key="cf-key",
+ model="@cf/meta/llama-2-7b-chat-int8",
+ optional_params={},
+ litellm_params={},
+ )
+
+
+def test_get_complete_url_migrates_legacy_ai_run_base():
+ config = CloudflareChatConfig()
+
+ url = config.get_complete_url(
+ api_base="https://api.cloudflare.com/client/v4/accounts/acct/ai/run/",
+ api_key="cf-key",
+ model="@cf/meta/llama-2-7b-chat-int8",
+ optional_params={},
+ litellm_params={},
+ )
+
+ assert (
+ url
+ == "https://api.cloudflare.com/client/v4/accounts/acct/ai/v1/chat/completions"
+ )
+ assert "/ai/run" not in url
+
+
+def test_transform_request_passes_tools_through_in_openai_format():
+ config = CloudflareChatConfig()
+ tools = [
+ {
+ "type": "function",
+ "function": {
+ "name": "get_weather",
+ "parameters": {
+ "type": "object",
+ "properties": {"city": {"type": "string"}},
+ },
+ },
+ }
+ ]
+ messages = [{"role": "user", "content": "weather in nyc?"}]
+
+ body = config.transform_request(
+ model="@cf/meta/llama-2-7b-chat-int8",
+ messages=messages,
+ optional_params={"tools": tools, "tool_choice": "auto"},
+ litellm_params={},
+ headers={},
+ )
+
+ assert body["messages"] == messages
+ assert body["model"] == "@cf/meta/llama-2-7b-chat-int8"
+ assert body["tools"] == tools
+ assert body["tool_choice"] == "auto"
+
+
+def test_validate_environment_requires_api_key():
+ config = CloudflareChatConfig()
+
+ with pytest.raises(ValueError, match="Missing Cloudflare API Key"):
+ config.validate_environment(
+ headers={},
+ model="@cf/meta/llama-2-7b-chat-int8",
+ messages=[],
+ optional_params={},
+ litellm_params={},
+ api_key=None,
+ )
+
+
+def test_validate_environment_sets_bearer_and_content_type():
+ config = CloudflareChatConfig()
+
+ headers = config.validate_environment(
+ headers={},
+ model="@cf/meta/llama-2-7b-chat-int8",
+ messages=[],
+ optional_params={},
+ litellm_params={},
+ api_key="cf-key",
+ )
+
+ assert headers["Authorization"] == "Bearer cf-key"
+ assert headers["Content-Type"] == "application/json"
diff --git a/tests/test_litellm/llms/custom_httpx/test_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_http_handler.py
index bf835b5d8f9..7bd1d7a6031 100644
--- a/tests/test_litellm/llms/custom_httpx/test_http_handler.py
+++ b/tests/test_litellm/llms/custom_httpx/test_http_handler.py
@@ -851,3 +851,56 @@ async def test_async_get_forwards_per_request_timeout():
}
finally:
await handler.close()
+
+
+class TestDefaultCachedClientTimeoutHonorsRequestTimeout:
+ """Cached default httpx clients must fall back to an explicit litellm.request_timeout.
+
+ Regression for LIT-2369: get_async_httpx_client / _get_httpx_client hardcoded a
+ 600s default and never consulted litellm.request_timeout, so provider calls with
+ no per-model timeout (e.g. Bedrock) hung for 600s.
+ """
+
+ @pytest.fixture
+ def restore_request_timeout(self):
+ original_value = litellm.request_timeout
+ original_flag = litellm.request_timeout_explicitly_set
+ try:
+ yield
+ finally:
+ litellm.request_timeout = original_value
+ litellm.request_timeout_explicitly_set = original_flag
+
+ def test_default_when_request_timeout_unset(self, restore_request_timeout):
+ from litellm.llms.custom_httpx.http_handler import (
+ _DEFAULT_TIMEOUT,
+ _default_cached_client_timeout,
+ )
+
+ litellm.request_timeout = litellm.constants.DEFAULT_REQUEST_TIMEOUT_SECONDS
+ litellm.request_timeout_explicitly_set = False
+ assert _default_cached_client_timeout() is _DEFAULT_TIMEOUT
+
+ def test_uses_explicit_request_timeout(self, restore_request_timeout):
+ from litellm.llms.custom_httpx.http_handler import (
+ _default_cached_client_timeout,
+ )
+
+ litellm.request_timeout = 300
+ litellm.request_timeout_explicitly_set = True
+ resolved = _default_cached_client_timeout()
+ assert resolved.read == 300.0
+ assert resolved.connect == 5.0
+
+ def test_cached_async_client_built_with_explicit_request_timeout(
+ self, restore_request_timeout
+ ):
+ from litellm.caching.llm_caching_handler import LLMClientCache
+ from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
+ from litellm.types.utils import LlmProviders
+
+ litellm.request_timeout = 300
+ litellm.request_timeout_explicitly_set = True
+ litellm.in_memory_llm_clients_cache = LLMClientCache()
+ client = get_async_httpx_client(llm_provider=LlmProviders.BEDROCK)
+ assert client.timeout.read == 300.0
diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py
index f7d445d0788..64ae30daa70 100644
--- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py
+++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py
@@ -10,13 +10,21 @@ sys.path.insert(
0, os.path.abspath("../../../..")
) # Adds the parent directory to the system path
import litellm
+from litellm.integrations.code_interpreter_interception.handler import (
+ CodeInterpreterInterceptionLogger,
+ LITELLM_CODE_EXECUTION_TOOL_NAME,
+)
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import (
BaseLLMHTTPHandler,
_google_genai_streaming_hidden_params,
)
+from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.router import GenericLiteLLMParams
+_ACTIVE_KEY = "_code_interpreter_interception_active"
+_SANDBOX_KEY = "_code_interpreter_interception_sandbox_key"
+
def test_prepare_fake_stream_request():
# Initialize the BaseLLMHTTPHandler
@@ -116,6 +124,117 @@ def test_response_api_handler_streams_when_provider_transform_adds_stream():
assert client.post.call_args.kwargs["json"]["stream"] is True
+def test_response_api_handler_runs_agentic_hooks_in_sync_path(monkeypatch):
+ handler = BaseLLMHTTPHandler()
+ config = Mock()
+ config.validate_environment.return_value = {}
+ config.get_complete_url.return_value = "https://chatgpt.example.com/responses"
+ config.transform_responses_api_request.return_value = {
+ "model": "gpt-5",
+ "input": "hi",
+ }
+ config.sign_request.return_value = ({}, None)
+ initial_response = Mock()
+ final_response = Mock()
+ config.transform_response_api_response.return_value = initial_response
+
+ client = HTTPHandler(client=httpx.Client())
+ client.post = Mock(
+ return_value=httpx.Response(
+ 200,
+ request=httpx.Request("POST", "https://chatgpt.example.com/responses"),
+ )
+ )
+ logging_obj = Mock()
+
+ monkeypatch.setattr(handler, "_has_agentic_completion_hook", Mock(return_value=True))
+ hook_mock = AsyncMock(return_value=final_response)
+ monkeypatch.setattr(handler, "_call_agentic_completion_hooks", hook_mock)
+
+ response = handler.response_api_handler(
+ model="gpt-5",
+ input="hi",
+ responses_api_provider_config=config,
+ response_api_optional_request_params={},
+ custom_llm_provider="openai",
+ litellm_params=GenericLiteLLMParams(),
+ logging_obj=logging_obj,
+ client=client,
+ )
+
+ assert response is final_response
+ hook_mock.assert_awaited_once()
+ assert hook_mock.call_args.kwargs["api_surface"] == "responses"
+ assert hook_mock.call_args.kwargs["messages"] == [
+ {"role": "user", "content": "hi"}
+ ]
+
+
+def test_response_api_handler_runs_responses_pre_call_hook_before_transform():
+ handler = BaseLLMHTTPHandler()
+ config = Mock()
+ config.validate_environment.return_value = {}
+ config.get_complete_url.return_value = "https://api.openai.com/v1/responses"
+ config.sign_request.return_value = ({}, None)
+ initial_response = ResponsesAPIResponse(
+ id="resp_1",
+ created_at=0,
+ output=[],
+ status="completed",
+ model="gpt-5",
+ )
+ config.transform_response_api_response.return_value = initial_response
+
+ def transform_responses_api_request(**kwargs):
+ return {
+ "model": kwargs["model"],
+ "input": kwargs["input"],
+ **kwargs["response_api_optional_request_params"],
+ }
+
+ config.transform_responses_api_request.side_effect = transform_responses_api_request
+ client = HTTPHandler(client=httpx.Client())
+ client.post = Mock(
+ return_value=httpx.Response(
+ 200,
+ request=httpx.Request("POST", "https://api.openai.com/v1/responses"),
+ )
+ )
+ logging_obj = Mock()
+ logging_obj.dynamic_success_callbacks = []
+
+ old_callbacks = list(litellm.callbacks)
+ litellm.callbacks = [CodeInterpreterInterceptionLogger()]
+ try:
+ response = handler.response_api_handler(
+ model="gpt-5",
+ input="use code",
+ responses_api_provider_config=config,
+ response_api_optional_request_params={
+ "tools": [{"type": "code_interpreter", "container": {"type": "auto"}}]
+ },
+ custom_llm_provider="openai",
+ litellm_params=GenericLiteLLMParams(api_key="sk-test"),
+ logging_obj=logging_obj,
+ client=client,
+ )
+ finally:
+ litellm.callbacks = old_callbacks
+
+ assert response is initial_response
+ transform_kwargs = config.transform_responses_api_request.call_args.kwargs
+ tools = transform_kwargs["response_api_optional_request_params"]["tools"]
+ assert not any(tool.get("type") == "code_interpreter" for tool in tools)
+ assert any(
+ tool.get("type") == "function"
+ and tool.get("name") == LITELLM_CODE_EXECUTION_TOOL_NAME
+ for tool in tools
+ )
+ hook_litellm_params = transform_kwargs["litellm_params"]
+ assert hook_litellm_params.get(_ACTIVE_KEY) is True
+ assert hook_litellm_params.get(_SANDBOX_KEY)
+
+
@pytest.mark.asyncio
async def test_async_response_api_handler_streams_when_provider_transform_adds_stream():
handler = BaseLLMHTTPHandler()
diff --git a/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py b/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py
index 3ee53bb46cd..7a3f372582f 100644
--- a/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py
+++ b/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py
@@ -719,3 +719,93 @@ class TestMistralFileHandling:
# Check that file_ids are modified to match Mistral's expected format
assert result[0]["content"][1]["file_id"] == "file-12345" # type: ignore
assert result[0]["content"][2]["file_id"] == "file-67890" # type: ignore
+
+
+class TestMistralStripsOutputOnlyFields:
+ """Mistral rejects unknown input fields with a 422 ``extra_forbidden``.
+
+ LiteLLM attaches ``reasoning_content`` / ``thinking_blocks`` to assistant
+ responses, so replaying an assistant turn verbatim must not forward them.
+ Regression for https://github.com/BerriAI/litellm/issues/30835.
+ """
+
+ def test_assistant_reasoning_content_is_dropped(self):
+ messages = cast(
+ List[AllMessageValues],
+ [
+ {"role": "user", "content": "Question?"},
+ {
+ "role": "assistant",
+ "content": "Follow-up",
+ "reasoning_content": "Some internal reasoning text.",
+ "thinking_blocks": [
+ {"type": "thinking", "thinking": "step", "signature": "mistral"}
+ ],
+ },
+ ],
+ )
+
+ result = cast(
+ List[AllMessageValues],
+ MistralConfig()._transform_messages(
+ messages=messages, model="mistral-medium-3-5"
+ ),
+ )
+
+ assistant_message = result[-1]
+ assert "reasoning_content" not in assistant_message
+ assert "thinking_blocks" not in assistant_message
+ assert assistant_message["content"] == "Follow-up"
+ assert assistant_message["role"] == "assistant"
+
+ def test_non_assistant_messages_are_untouched(self):
+ messages = cast(
+ List[AllMessageValues],
+ [{"role": "user", "content": "Question?", "reasoning_content": "noise"}],
+ )
+
+ result = cast(
+ List[AllMessageValues],
+ MistralConfig()._transform_messages(
+ messages=messages, model="mistral-medium-3-5"
+ ),
+ )
+
+ assert result[0].get("reasoning_content") == "noise"
+
+ def test_reasoning_content_dropped_when_image_present(self):
+ """The image branch returns early, so stripping must run before it."""
+ messages = cast(
+ List[AllMessageValues],
+ [
+ {
+ "role": "user",
+ "content": [
+ {"type": "text", "text": "Describe this"},
+ {
+ "type": "image_url",
+ "image_url": {"url": "https://example.com/cat.png"},
+ },
+ ],
+ },
+ {
+ "role": "assistant",
+ "content": "A cat.",
+ "reasoning_content": "leaked reasoning",
+ },
+ ],
+ )
+
+ with patch.object(
+ MistralConfig,
+ "_transform_messages_sync",
+ side_effect=lambda transformed, model: transformed,
+ ):
+ result = cast(
+ List[AllMessageValues],
+ MistralConfig()._transform_messages(
+ messages=messages, model="mistral-medium-3-5", is_async=False
+ ),
+ )
+
+ assert "reasoning_content" not in result[-1]
diff --git a/tests/test_litellm/llms/openai_like/test_json_providers.py b/tests/test_litellm/llms/openai_like/test_json_providers.py
index 39a4964f5f4..c8743e1809d 100644
--- a/tests/test_litellm/llms/openai_like/test_json_providers.py
+++ b/tests/test_litellm/llms/openai_like/test_json_providers.py
@@ -2,6 +2,7 @@
Tests for JSON-based provider configuration system.
"""
+import json
import os
import sys
from unittest.mock import MagicMock, patch
@@ -244,6 +245,99 @@ class TestPinstripes:
assert result["temperature"] == 0.7
+class TestDarkbloom:
+ def test_darkbloom_json_config_exists(self):
+ from litellm.llms.openai_like.json_loader import JSONProviderRegistry
+
+ darkbloom = JSONProviderRegistry.get("darkbloom")
+ assert darkbloom is not None
+ assert darkbloom.base_url == "https://api.darkbloom.dev/v1"
+ assert darkbloom.api_key_env == "DARKBLOOM_API_KEY"
+ assert darkbloom.api_base_env == "DARKBLOOM_API_BASE"
+ assert darkbloom.param_mappings.get("max_completion_tokens") == "max_tokens"
+
+ def test_darkbloom_provider_resolution(self):
+ from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
+
+ model, provider, api_key, api_base = get_llm_provider(
+ model="darkbloom/gemma-4-26b",
+ custom_llm_provider=None,
+ api_base=None,
+ api_key=None,
+ )
+
+ assert model == "gemma-4-26b"
+ assert provider == "darkbloom"
+ assert api_key is None
+ assert api_base == "https://api.darkbloom.dev/v1"
+
+ def test_darkbloom_dynamic_config(self):
+ from litellm.llms.openai_like.dynamic_config import create_config_class
+ from litellm.llms.openai_like.json_loader import JSONProviderRegistry
+
+ provider = JSONProviderRegistry.get("darkbloom")
+ config_class = create_config_class(provider)
+ config = config_class()
+
+ api_base, api_key = config._get_openai_compatible_provider_info(None, None)
+ assert api_base == "https://api.darkbloom.dev/v1"
+
+ api_base, api_key = config._get_openai_compatible_provider_info(
+ "https://custom.darkbloom.dev/v1", "test-key"
+ )
+ assert api_base == "https://custom.darkbloom.dev/v1"
+ assert api_key == "test-key"
+
+ def test_darkbloom_complete_url_appends_endpoint(self):
+ from litellm.llms.openai_like.dynamic_config import create_config_class
+ from litellm.llms.openai_like.json_loader import JSONProviderRegistry
+
+ provider = JSONProviderRegistry.get("darkbloom")
+ config_class = create_config_class(provider)
+ config = config_class()
+
+ url = config.get_complete_url(
+ api_base="https://api.darkbloom.dev/v1",
+ api_key="test-key",
+ model="darkbloom/gemma-4-26b",
+ optional_params={},
+ litellm_params={},
+ stream=True,
+ )
+
+ assert url == "https://api.darkbloom.dev/v1/chat/completions"
+
+ def test_darkbloom_provider_config_manager(self):
+ from litellm import LlmProviders
+ from litellm.utils import ProviderConfigManager
+
+ config = ProviderConfigManager.get_provider_chat_config(
+ model="gemma-4-26b", provider=LlmProviders.DARKBLOOM
+ )
+
+ assert config is not None
+ assert config.custom_llm_provider == "darkbloom"
+
+ def test_darkbloom_model_cost_map(self):
+ with open(
+ os.path.join(workspace_path, "model_prices_and_context_window.json")
+ ) as f:
+ model_cost = json.load(f)
+
+ expected_models = {
+ "darkbloom/gemma-4-26b": (3e-08, 1.65e-07),
+ "darkbloom/gpt-oss-20b": (1.45e-08, 7e-08),
+ }
+ for model, (input_cost, output_cost) in expected_models.items():
+ assert model in model_cost
+ assert model_cost[model]["litellm_provider"] == "darkbloom"
+ assert model_cost[model]["max_output_tokens"] == 32768
+ assert model_cost[model]["supports_function_calling"] is True
+ assert model_cost[model]["supports_tool_choice"] is True
+ assert model_cost[model]["input_cost_per_token"] == input_cost
+ assert model_cost[model]["output_cost_per_token"] == output_cost
+
+
class TestPublicAIIntegration:
"""Integration tests for PublicAI provider"""
diff --git a/tests/test_litellm/llms/parallel_ai/test_parallel_ai_search.py b/tests/test_litellm/llms/parallel_ai/test_parallel_ai_search.py
index b5c1a86205b..7be295826e3 100644
--- a/tests/test_litellm/llms/parallel_ai/test_parallel_ai_search.py
+++ b/tests/test_litellm/llms/parallel_ai/test_parallel_ai_search.py
@@ -293,7 +293,10 @@ class TestParallelAISearch:
],
)
@pytest.mark.asyncio
- async def test_custom_api_base_appends_v1_search(self, api_base):
+ async def test_custom_api_base_appends_v1_search(self, api_base, monkeypatch):
+ # Operator points at an internal base via the env override (a trusted
+ # host), so the server key is still used and the URL is normalized.
+ monkeypatch.setenv("PARALLEL_AI_API_BASE", api_base)
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
@@ -303,7 +306,6 @@ class TestParallelAISearch:
await litellm.asearch(
query="AI developments",
search_provider="parallel_ai",
- api_base=api_base,
)
call_args = mock_post.call_args
@@ -312,6 +314,23 @@ class TestParallelAISearch:
== "https://proxy.internal.example.com/v1/search"
)
+ @pytest.mark.asyncio
+ async def test_caller_api_base_without_key_is_refused(self, monkeypatch):
+ # A caller-supplied api_base (untrusted host) while relying on the
+ # server key must be refused without any outbound request.
+ monkeypatch.setenv("PARALLEL_API_KEY", "server-secret")
+ with patch(
+ "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
+ new_callable=AsyncMock,
+ ) as mock_post:
+ with pytest.raises(Exception, match="Refusing to send"):
+ await litellm.asearch(
+ query="AI developments",
+ search_provider="parallel_ai",
+ api_base="https://attacker.example.com",
+ )
+ mock_post.assert_not_called()
+
@pytest.mark.asyncio
async def test_missing_api_key_raises(self, monkeypatch):
monkeypatch.delenv("PARALLEL_API_KEY", raising=False)
diff --git a/tests/test_litellm/llms/perplexity/test_perplexity_cost_calculator.py b/tests/test_litellm/llms/perplexity/test_perplexity_cost_calculator.py
index e2d1ab72c5e..46c1e457d7c 100644
--- a/tests/test_litellm/llms/perplexity/test_perplexity_cost_calculator.py
+++ b/tests/test_litellm/llms/perplexity/test_perplexity_cost_calculator.py
@@ -9,7 +9,7 @@ import json
import math
import os
import sys
-from unittest.mock import Mock, patch
+from unittest.mock import patch
import pytest
@@ -120,10 +120,10 @@ class TestPerplexityCostCalculator:
# Expected costs:
# Input: 100 tokens * $2e-6 = $0.0002
# Output: 50 tokens * $8e-6 = $0.0004
- # Search: 3 queries * ($0.005 / 1000) = $0.000015
- # Total completion cost: $0.000415
+ # Search: 3 queries * $0.005 per request = $0.015
+ # Total completion cost: $0.0154
expected_prompt_cost = 100 * 2e-6
- expected_completion_cost = (50 * 8e-6) + (3 / 1000 * 0.005)
+ expected_completion_cost = (50 * 8e-6) + (3 * 0.005)
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6)
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-6)
@@ -195,10 +195,10 @@ class TestPerplexityCostCalculator:
# Total prompt cost = $0.00026
# Output (text): (50 - 15) tokens * $8e-6 = $0.00028
# Reasoning: 15 tokens * $3e-6 = $0.000045
- # Search: 2 queries * ($0.005 / 1000) = $0.00001
- # Total completion cost = $0.000335
+ # Search: 2 queries * $0.005 per request = $0.01
+ # Total completion cost = $0.010325
expected_prompt_cost = (100 * 2e-6) + (30 * 2e-6)
- expected_completion_cost = ((50 - 15) * 8e-6) + (15 * 3e-6) + (2 / 1000 * 0.005)
+ expected_completion_cost = ((50 - 15) * 8e-6) + (15 * 3e-6) + (2 * 0.005)
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6)
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-6)
@@ -311,7 +311,7 @@ class TestPerplexityCostCalculator:
# Calculate expected total cost (reasoning is a subset of completion_tokens)
expected_prompt_cost = (100 * 2e-6) + (15 * 2e-6) # Input + citation
expected_completion_cost = (
- ((50 - 10) * 8e-6) + (10 * 3e-6) + (1 / 1000 * 0.005)
+ ((50 - 10) * 8e-6) + (10 * 3e-6) + (1 * 0.005)
) # Output (text) + reasoning + search
expected_total = expected_prompt_cost + expected_completion_cost
@@ -361,7 +361,7 @@ class TestPerplexityCostCalculator:
expected_completion_cost = (
((50 - reasoning_tokens) * 8e-6)
+ (reasoning_tokens * 3e-6)
- + (search_queries / 1000 * 0.005)
+ + (search_queries * 0.005)
)
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6)
diff --git a/tests/test_litellm/llms/perplexity/test_perplexity_integration.py b/tests/test_litellm/llms/perplexity/test_perplexity_integration.py
index e59fbc9f272..8691e6a1ee5 100644
--- a/tests/test_litellm/llms/perplexity/test_perplexity_integration.py
+++ b/tests/test_litellm/llms/perplexity/test_perplexity_integration.py
@@ -9,7 +9,6 @@ import json
import math
import os
import sys
-from unittest.mock import Mock, patch
import pytest
@@ -106,8 +105,8 @@ class TestPerplexityIntegration:
expected_prompt_cost = (100 * 2e-6) + (citation_tokens * 2e-6)
expected_completion_cost = (
- ((50 - 10) * 8e-6) + (10 * 3e-6) + (2 / 1000 * 0.005)
- )
+ ((50 - 10) * 8e-6) + (10 * 3e-6) + (2 * 0.005)
+ ) # Output (text) + reasoning + search
expected_total = expected_prompt_cost + expected_completion_cost
assert math.isclose(total_cost, expected_total, rel_tol=1e-6)
@@ -152,8 +151,8 @@ class TestPerplexityIntegration:
expected_prompt_cost = (200 * 2e-6) + (40 * 2e-6)
expected_completion_cost = (
- ((100 - 25) * 8e-6) + (25 * 3e-6) + (3 / 1000 * 0.005)
- )
+ ((100 - 25) * 8e-6) + (25 * 3e-6) + (3 * 0.005)
+ ) # Output (text) + reasoning + search
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6)
assert math.isclose(completion_cost_val, expected_completion_cost, rel_tol=1e-6)
@@ -262,9 +261,9 @@ class TestPerplexityIntegration:
expected_prompt_cost = (50000 * 2e-6) + (5000 * 2e-6)
expected_completion_cost = (
- ((25000 - 10000) * 8e-6) + (10000 * 3e-6) + (100 / 1000 * 0.005)
- )
- expected_total = expected_prompt_cost + expected_completion_cost
+ ((25000 - 10000) * 8e-6) + (10000 * 3e-6) + (100 * 0.005)
+ ) # $0.65
+ expected_total = expected_prompt_cost + expected_completion_cost # $0.76
assert math.isclose(total_cost, expected_total, rel_tol=1e-6)
assert total_cost > 0.25
@@ -326,7 +325,7 @@ class TestPerplexityIntegration:
# Should calculate costs correctly
expected_prompt_cost = (100 * 2e-6) + (10 * 2e-6)
- expected_completion_cost = (50 * 8e-6) + (1 / 1000 * 0.005)
+ expected_completion_cost = (50 * 8e-6) + (1 * 0.005)
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6)
assert math.isclose(completion_cost_val, expected_completion_cost, rel_tol=1e-6)
diff --git a/tests/test_litellm/llms/tinyfish/test_tinyfish_search.py b/tests/test_litellm/llms/tinyfish/test_tinyfish_search.py
index 5496486765c..9870d30d488 100644
--- a/tests/test_litellm/llms/tinyfish/test_tinyfish_search.py
+++ b/tests/test_litellm/llms/tinyfish/test_tinyfish_search.py
@@ -63,23 +63,17 @@ class TestTinyfishSearchConfig:
assert headers["X-API-Key"] == "sk-tinyfish-test"
assert headers["Accept"] == "application/json"
- def test_validate_environment_from_env(self):
+ def test_validate_environment_from_env(self, monkeypatch):
+ monkeypatch.setenv("TINYFISH_API_KEY", "sk-from-env")
config = TinyfishSearchConfig()
- with patch(
- "litellm.llms.tinyfish.search.transformation.get_secret_str",
- return_value="sk-from-env",
- ):
- headers = config.validate_environment(headers={})
+ headers = config.validate_environment(headers={})
assert headers["X-API-Key"] == "sk-from-env"
- def test_validate_environment_missing_key(self):
+ def test_validate_environment_missing_key(self, monkeypatch):
+ monkeypatch.delenv("TINYFISH_API_KEY", raising=False)
config = TinyfishSearchConfig()
- with patch(
- "litellm.llms.tinyfish.search.transformation.get_secret_str",
- return_value=None,
- ):
- with pytest.raises(ValueError, match="TINYFISH_API_KEY"):
- config.validate_environment(headers={})
+ with pytest.raises(ValueError, match="TINYFISH_API_KEY"):
+ config.validate_environment(headers={})
def test_validate_environment_uses_api_base_kwarg(self):
config = TinyfishSearchConfig()
diff --git a/tests/test_litellm/llms/vertex_ai/realtime/test_vertex_ai_realtime_transformation.py b/tests/test_litellm/llms/vertex_ai/realtime/test_vertex_ai_realtime_transformation.py
index 1ebd704be34..1f171496cce 100644
--- a/tests/test_litellm/llms/vertex_ai/realtime/test_vertex_ai_realtime_transformation.py
+++ b/tests/test_litellm/llms/vertex_ai/realtime/test_vertex_ai_realtime_transformation.py
@@ -346,3 +346,74 @@ def test_vertex_does_not_warn_when_dropping_non_guardrail_session_update(caplog)
"Vertex AI Realtime" in record.message and "session.update" in record.message
for record in caplog.records
)
+
+
+async def test_async_realtime_does_not_forward_client_query_params_to_vertex_backend(
+ monkeypatch,
+):
+ """Regression: forwarding client ?model=/?intent= to the Vertex Live WSS URL causes 1007 errors.
+
+ Exercises ``async_realtime`` end-to-end so that re-adding ``_append_query_params``
+ (the reverted bug) would push ``model=``/``intent=`` onto the backend URL and fail here.
+ """
+ import websockets
+
+ from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
+
+ cfg = VertexAIRealtimeConfig(
+ access_token="tok", project="my-proj", location="us-central1"
+ )
+
+ captured = {}
+
+ def fake_connect(url, *args, **kwargs):
+ captured["url"] = url
+ raise RuntimeError("stop before establishing the backend connection")
+
+ monkeypatch.setattr(websockets, "connect", fake_connect)
+
+ await BaseLLMHTTPHandler().async_realtime(
+ model="gemini-live-2.5-flash-native-audio",
+ websocket=AsyncMock(),
+ logging_obj=MagicMock(),
+ provider_config=cfg,
+ headers={},
+ query_params={
+ "model": "gemini-live-2.5-flash-native-audio",
+ "intent": "chat",
+ },
+ )
+
+ assert "?" not in captured["url"]
+ assert "model=" not in captured["url"]
+ assert "intent=" not in captured["url"]
+
+
+def test_vertex_function_call_output_omits_id():
+ """Regression: Vertex Live rejects ``id`` on toolResponse.functionResponses (1007)."""
+ cfg = VertexAIRealtimeConfig(
+ access_token="tok", project="my-proj", location="us-central1"
+ )
+ cfg._tool_call_id_to_name["call_abc123"] = "terminate_call"
+
+ messages = cfg.transform_realtime_request(
+ json.dumps(
+ {
+ "type": "conversation.item.create",
+ "item": {
+ "type": "function_call_output",
+ "call_id": "call_abc123",
+ "output": '{"status": "ok"}',
+ },
+ }
+ ),
+ "gemini-live-2.5-flash-native-audio",
+ session_configuration_request="existing",
+ )
+
+ assert len(messages) == 1
+ payload = json.loads(messages[0])
+ function_response = payload["toolResponse"]["functionResponses"][0]
+ assert "id" not in function_response
+ assert function_response["name"] == "terminate_call"
+ assert function_response["response"] == {"status": "ok"}
diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py
new file mode 100644
index 00000000000..7e028064e4c
--- /dev/null
+++ b/tests/test_litellm/ocr/test_rust_bridge.py
@@ -0,0 +1,333 @@
+"""Tests for the optional Rust-backed OCR path (``litellm/ocr/rust_bridge.py``)."""
+
+import importlib
+import sys
+import types
+
+import httpx
+import pytest
+
+import litellm
+from litellm.llms.base_llm.ocr.transformation import OCRResponse
+
+# `litellm/__init__.py` does `from .ocr.main import *`, which binds the `ocr`
+# function onto `litellm.ocr` and shadows the submodule, so import the modules
+# explicitly via importlib rather than attribute traversal.
+ocr_main = importlib.import_module("litellm.ocr.main")
+rust_bridge = importlib.import_module("litellm.ocr.rust_bridge")
+
+MODEL = "mistral/mistral-ocr-latest"
+DOCUMENT = {"type": "document_url", "document_url": "https://example.com/doc.pdf"}
+
+FAKE_OCR_RESPONSE = {
+ "pages": [{"index": 0, "markdown": "hello world"}],
+ "model": "mistral-ocr-2505-completion",
+ "document_annotation": None,
+ "usage_info": {"pages_processed": 1},
+ "object": "ocr",
+}
+
+
+class RecordingBridge:
+ """A fake ``RustOcr`` callable that records the args it was handed."""
+
+ def __init__(self):
+ self.calls = []
+
+ def __call__(
+ self, model, document, api_key, api_base, optional_params, timeout_seconds
+ ):
+ self.calls.append(
+ {
+ "model": model,
+ "document": document,
+ "api_key": api_key,
+ "api_base": api_base,
+ "optional_params": optional_params,
+ "timeout_seconds": timeout_seconds,
+ }
+ )
+ return dict(FAKE_OCR_RESPONSE)
+
+
+class RecordingLogging:
+ """A spy standing in for ``LiteLLMLoggingObj`` to capture ``pre_call``."""
+
+ def __init__(self):
+ self.pre_call_kwargs = None
+
+ def pre_call(self, *, input, api_key, additional_args):
+ self.pre_call_kwargs = {
+ "input": input,
+ "api_key": api_key,
+ "additional_args": additional_args,
+ }
+
+
+class FakeOCRConfig:
+ """A stand-in ``BaseOCRConfig`` that echoes the request it would build."""
+
+ def validate_environment(
+ self, *, headers, model, api_key, api_base, litellm_params
+ ):
+ return {"authorization": f"Bearer {api_key}"}
+
+ def get_complete_url(self, *, api_base, model, optional_params, litellm_params):
+ return f"{api_base or 'https://api.mistral.ai/v1'}/ocr"
+
+
+@pytest.fixture(autouse=True)
+def _reset_rust_flag():
+ """Keep the global toggle isolated between tests."""
+ rust_bridge.use_litellm_rust(False, ocr=None)
+ yield
+ rust_bridge.use_litellm_rust(False, ocr=None)
+
+
+@pytest.fixture
+def fake_bridge():
+ """Enable the Rust path with an injected recording bridge (no native wheel)."""
+ bridge = RecordingBridge()
+ litellm.use_litellm_rust(True, ocr=bridge)
+ return bridge
+
+
+def test_use_litellm_rust_toggles_flag():
+ assert rust_bridge.rust_ocr_enabled() is False
+ litellm.use_litellm_rust()
+ assert rust_bridge.rust_ocr_enabled() is True
+ litellm.use_litellm_rust(False)
+ assert rust_bridge.rust_ocr_enabled() is False
+
+
+def test_load_rust_ocr_returns_injected_impl():
+ bridge = RecordingBridge()
+ litellm.use_litellm_rust(True, ocr=bridge)
+ assert rust_bridge.load_rust_ocr() is bridge
+
+
+def test_toggle_without_ocr_arg_preserves_injected_impl():
+ """Regression: routine enable/disable calls must not clobber a prior injection.
+
+ Earlier, ``use_litellm_rust()`` unconditionally assigned the keyword default
+ of ``None`` to ``_rust_ocr_impl``, silently dropping a custom bridge whenever
+ a caller toggled the flag without re-passing ``ocr=``.
+ """
+ bridge = RecordingBridge()
+ litellm.use_litellm_rust(True, ocr=bridge)
+
+ litellm.use_litellm_rust(False)
+ assert rust_bridge.load_rust_ocr() is bridge
+ litellm.use_litellm_rust(True)
+ assert rust_bridge.load_rust_ocr() is bridge
+
+
+def test_explicit_ocr_none_clears_injected_impl():
+ bridge = RecordingBridge()
+ litellm.use_litellm_rust(True, ocr=bridge)
+
+ litellm.use_litellm_rust(True, ocr=None)
+ assert rust_bridge.load_rust_ocr() is None
+
+
+def test_load_rust_ocr_none_when_extension_absent():
+ """With no injected impl and no compiled wheel, the loader returns None so the
+ caller degrades to the Python path instead of raising ImportError."""
+ litellm.use_litellm_rust(True) # no impl injected; extension isn't built in CI
+ assert rust_bridge.load_rust_ocr() is None
+
+
+def test_load_rust_ocr_uses_compiled_extension(monkeypatch):
+ """With no injected impl but a compiled ``litellm_python_bridge`` importable,
+ the loader returns the extension's ``ocr`` callable. The native wheel isn't
+ built in CI, so stand in a fake module via ``sys.modules``."""
+ fake_module = types.ModuleType("litellm_python_bridge")
+ fake_module.ocr = lambda **kwargs: dict(FAKE_OCR_RESPONSE) # type: ignore[attr-defined]
+ monkeypatch.setitem(sys.modules, "litellm_python_bridge", fake_module)
+
+ litellm.use_litellm_rust(True) # enabled, no impl injected -> import the extension
+ assert rust_bridge.load_rust_ocr() is fake_module.ocr
+
+
+def test_timeout_to_seconds_handles_float_timeout_and_none():
+ assert ocr_main._timeout_to_seconds(12.5) == 12.5
+ assert ocr_main._timeout_to_seconds(None) is None
+ assert ocr_main._timeout_to_seconds(httpx.Timeout(30.0, read=42.0)) == 42.0
+
+
+def test_run_rust_ocr_forwards_args_and_wraps_response():
+ bridge = RecordingBridge()
+ logging_obj = RecordingLogging()
+
+ response = ocr_main._run_rust_ocr(
+ rust_ocr=bridge,
+ logging_obj=logging_obj,
+ provider_config=FakeOCRConfig(),
+ resolve_api_key=lambda _name: None,
+ model="mistral-ocr-latest",
+ document=DOCUMENT,
+ api_key="sk-test",
+ api_base="https://proxy.internal",
+ optional_params={"include_image_base64": True},
+ litellm_params={},
+ timeout_seconds=12.5,
+ )
+
+ assert isinstance(response, OCRResponse)
+ assert response.pages[0].markdown == "hello world"
+ call = bridge.calls[0]
+ assert call == {
+ "model": "mistral-ocr-latest",
+ "document": DOCUMENT,
+ "api_key": "sk-test",
+ "api_base": "https://proxy.internal",
+ "optional_params": {"include_image_base64": True},
+ "timeout_seconds": 12.5,
+ }
+
+
+def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing():
+ """No explicit api_key: the resolver (get_secret_str in production) supplies it,
+ so secret-manager backends (AWS/Azure/GCP/Vault) work like the Python path."""
+ bridge = RecordingBridge()
+
+ ocr_main._run_rust_ocr(
+ rust_ocr=bridge,
+ logging_obj=RecordingLogging(),
+ provider_config=FakeOCRConfig(),
+ resolve_api_key=lambda name: (
+ "sk-from-vault" if name == "MISTRAL_API_KEY" else None
+ ),
+ model="mistral-ocr-latest",
+ document=DOCUMENT,
+ api_key=None,
+ api_base=None,
+ optional_params={},
+ litellm_params={},
+ timeout_seconds=None,
+ )
+
+ assert bridge.calls[0]["api_key"] == "sk-from-vault"
+
+
+def test_run_rust_ocr_prefers_explicit_key_over_resolver():
+ bridge = RecordingBridge()
+ resolver_calls = []
+
+ def _resolver(name):
+ resolver_calls.append(name)
+ return "sk-from-vault"
+
+ ocr_main._run_rust_ocr(
+ rust_ocr=bridge,
+ logging_obj=RecordingLogging(),
+ provider_config=FakeOCRConfig(),
+ resolve_api_key=_resolver,
+ model="mistral-ocr-latest",
+ document=DOCUMENT,
+ api_key="sk-explicit",
+ api_base=None,
+ optional_params={},
+ litellm_params={},
+ timeout_seconds=None,
+ )
+
+ assert bridge.calls[0]["api_key"] == "sk-explicit"
+ assert resolver_calls == [] # resolver never consulted when a key is supplied
+
+
+def test_run_rust_ocr_runs_pre_call_logging():
+ """The Rust shortcut must run pre_call so callbacks and spend tracking fire."""
+ logging_obj = RecordingLogging()
+
+ ocr_main._run_rust_ocr(
+ rust_ocr=RecordingBridge(),
+ logging_obj=logging_obj,
+ provider_config=FakeOCRConfig(),
+ resolve_api_key=lambda _name: None,
+ model="mistral-ocr-latest",
+ document=DOCUMENT,
+ api_key="sk-test",
+ api_base="https://api.mistral.ai/v1",
+ optional_params={"include_image_base64": True},
+ litellm_params={},
+ timeout_seconds=None,
+ )
+
+ assert logging_obj.pre_call_kwargs is not None
+ assert logging_obj.pre_call_kwargs["input"] == "OCR document processing"
+ additional_args = logging_obj.pre_call_kwargs["additional_args"]
+ complete_input = additional_args["complete_input_dict"]
+ assert complete_input["document"] == DOCUMENT
+ assert complete_input["include_image_base64"] is True
+ # The logged request mirrors what Rust sends: resolved URL + headers.
+ assert additional_args["api_base"] == "https://api.mistral.ai/v1/ocr"
+ assert additional_args["headers"] == {"authorization": "Bearer sk-test"}
+
+
+def test_ocr_routes_to_rust_when_enabled(fake_bridge):
+ response = litellm.ocr(
+ model=MODEL,
+ document=DOCUMENT,
+ api_key="sk-test",
+ include_image_base64=True,
+ )
+
+ assert isinstance(response, OCRResponse)
+ assert response.pages[0].markdown == "hello world"
+ assert len(fake_bridge.calls) == 1
+ call = fake_bridge.calls[0]
+ # Provider prefix is stripped before reaching the bridge.
+ assert call["model"] == "mistral-ocr-latest"
+ assert call["document"] == DOCUMENT
+ assert call["api_key"] == "sk-test"
+ # Raw OCR params ride along in optional_params; Rust filters to supported keys.
+ assert call["optional_params"].get("include_image_base64") is True
+
+
+def test_ocr_forwards_timeout_to_rust(fake_bridge):
+ """Caller-supplied timeout must flow into the Rust bridge so the fixed 600s
+ client ceiling doesn't silently override shorter deadlines."""
+ litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test", timeout=12.5)
+
+ assert fake_bridge.calls[0]["timeout_seconds"] == 12.5
+
+
+def test_ocr_passes_default_request_timeout_to_rust(fake_bridge):
+ """When no explicit timeout is given, the library default (request_timeout)
+ must still be forwarded so the Rust path matches the Python path's deadline."""
+ from litellm.constants import request_timeout
+
+ litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
+
+ assert fake_bridge.calls[0]["timeout_seconds"] == float(request_timeout)
+
+
+def test_ocr_does_not_route_to_rust_when_disabled():
+ """With the flag off, the bridge must not be consulted even if an impl exists."""
+ bridge = RecordingBridge()
+ litellm.use_litellm_rust(False, ocr=bridge)
+
+ assert rust_bridge.rust_ocr_enabled() is False
+ # The impl stays available for injection, but the disabled flag gates usage,
+ # so ocr() never reaches the Rust path (asserted via the enabled-path test).
+ assert bridge.calls == []
+
+
+def test_ocr_falls_back_to_python_when_bridge_unavailable(monkeypatch):
+ """Rust enabled but no bridge available (no injected impl, no compiled wheel):
+ ocr() must degrade to the Python HTTP handler instead of raising."""
+ litellm.use_litellm_rust(True) # enabled, but load_rust_ocr() returns None in CI
+
+ captured = {}
+
+ def fake_handler_ocr(**kwargs):
+ captured["called"] = True
+ return OCRResponse(pages=[], model="mistral-ocr-latest", object="ocr")
+
+ monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", fake_handler_ocr)
+
+ response = litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
+
+ assert captured.get("called") is True # Python path was used
+ assert isinstance(response, OCRResponse)
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py
new file mode 100644
index 00000000000..7885617aa46
--- /dev/null
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py
@@ -0,0 +1,57 @@
+"""Tests for the resolver dispatch skeleton.
+
+Every mode must reach its own arm and, until that arm is built, return a typed
+`not_implemented` CredError rather than silently producing no credential. Parametrizing over
+one config per mode also guards reachability: if a `case` were dropped, that mode would fall to
+the `assert_never` tail and raise here instead of returning the stub.
+"""
+
+import pytest
+from pydantic import SecretStr
+
+from litellm.proxy._experimental.mcp_server.outbound_credentials import (
+ ApiKeyConfig,
+ AuthorizationCodeConfig,
+ AuthSpecKind,
+ AwsSigV4Config,
+ ClientCredentialsConfig,
+ Error,
+ NoneConfig,
+ PassthroughConfig,
+ ServerSpec,
+ SharedKey,
+ Subject,
+ TokenExchangeConfig,
+ UpstreamCredentialProvider,
+)
+
+_ONE_CONFIG_PER_MODE = [
+ (AuthSpecKind.none, NoneConfig()),
+ (AuthSpecKind.api_key, ApiKeyConfig(key_source=SharedKey(value=SecretStr("k")))),
+ (AuthSpecKind.passthrough, PassthroughConfig()),
+ (AuthSpecKind.client_credentials, ClientCredentialsConfig()),
+ (AuthSpecKind.token_exchange, TokenExchangeConfig()),
+ (AuthSpecKind.authorization_code, AuthorizationCodeConfig()),
+ (AuthSpecKind.aws_sigv4, AwsSigV4Config(region="us-east-1")),
+]
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("kind, config", _ONE_CONFIG_PER_MODE)
+async def test_every_mode_reaches_its_arm_and_returns_not_implemented(kind, config):
+ spec = ServerSpec(
+ server_id="s", resource="https://upstream.example.com", config=config
+ )
+ subject = Subject(tenant_id="", subject_id="")
+
+ result = await UpstreamCredentialProvider().resolve_credentials(subject, spec)
+
+ assert isinstance(result, Error)
+ assert result.error.tag == "not_implemented"
+ assert kind.value in result.error.summary
+
+
+def test_all_seven_modes_are_covered():
+ # Guards that the parametrization (and therefore the dispatch) spans every AuthSpecKind, so a
+ # newly added mode without a test row is caught here rather than slipping through.
+ assert {kind for kind, _ in _ONE_CONFIG_PER_MODE} == set(AuthSpecKind)
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_debug.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_debug.py
index 468bd946ae9..d299239f68e 100644
--- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_debug.py
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_debug.py
@@ -41,9 +41,13 @@ class TestMask:
def test_empty_returns_none_label(self):
assert MCPDebug._mask("") == "(none)"
- def test_short_value_unchanged(self):
- # visible_prefix=6 + visible_suffix=4 = 10, so <= 10 chars unchanged
- assert MCPDebug._mask("sk-1234") == "sk-1234"
+ def test_short_value_masked(self):
+ # Short auth values must not be echoed verbatim in debug headers, even though
+ # visible_prefix + visible_suffix would otherwise reveal the whole value.
+ masked = MCPDebug._mask("sk-1234")
+ assert "sk-1234" not in masked
+ assert set(masked) == {"*"}
+ assert len(masked) == len("sk-1234")
def test_long_value_masked(self):
result = MCPDebug._mask("Bearer eyJhbGciOiJSUzI1NiIsInR5cCI6IkpXVCJ9")
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py
index 2558df8533b..cebc265a148 100644
--- a/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py
@@ -452,6 +452,430 @@ async def test_semantic_filter_hook_skips_no_tools():
print("✅ Hook correctly skips requests without tools")
+@pytest.mark.asyncio
+async def test_semantic_filter_hook_preserves_native_tools():
+ """
+ Regression test: mixed MCP + native tools.
+
+ Given: 5 MCP tools (registered in _tool_map) + 2 native OpenAI-format
+ function tools (not in _tool_map)
+ When: The hook filters tools
+ Then: The native tools must survive unconditionally, and only MCP
+ tools go through the semantic filter.
+ """
+ from litellm.proxy._experimental.mcp_server.semantic_tool_filter import (
+ SemanticMCPToolFilter,
+ )
+ from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook
+ from litellm.types.utils import Embedding, EmbeddingResponse
+
+ mock_router = Mock()
+
+ def mock_embedding_sync(*args, **kwargs):
+ return EmbeddingResponse(
+ data=[Embedding(embedding=[0.1] * 1536, index=0, object="embedding")],
+ model="text-embedding-3-small",
+ object="list",
+ usage={"prompt_tokens": 10, "total_tokens": 10},
+ )
+
+ async def mock_embedding_async(*args, **kwargs):
+ return mock_embedding_sync()
+
+ mock_router.embedding = mock_embedding_sync
+ mock_router.aembedding = mock_embedding_async
+
+ filter_instance = SemanticMCPToolFilter(
+ embedding_model="text-embedding-3-small",
+ litellm_router_instance=mock_router,
+ top_k=2,
+ similarity_threshold=0.3,
+ enabled=True,
+ )
+
+ # --- MCP tools (registered in the semantic router) ---
+ mcp_tools = [
+ MCPTool(
+ name=f"mcp_tool_{i}",
+ description=f"MCP tool {i}",
+ inputSchema={"type": "object"},
+ )
+ for i in range(5)
+ ]
+ filter_instance._build_router(mcp_tools)
+
+ # --- Native OpenAI-format function tools (NOT in _tool_map) ---
+ native_tools = [
+ {
+ "type": "function",
+ "function": {
+ "name": "get_current_weather",
+ "description": "Get the current weather",
+ "parameters": {"type": "object", "properties": {}},
+ },
+ },
+ {
+ "type": "function",
+ "function": {
+ "name": "search_web",
+ "description": "Search the web",
+ "parameters": {"type": "object", "properties": {}},
+ },
+ },
+ ]
+
+ # Combine: MCP tools + native tools
+ all_tools = list(mcp_tools) + native_tools
+
+ hook = SemanticToolFilterHook(filter_instance)
+
+ data = {
+ "model": "gpt-4",
+ "messages": [{"role": "user", "content": "What is the weather?"}],
+ "tools": all_tools,
+ "metadata": {},
+ }
+
+ result = await hook.async_pre_call_hook(
+ user_api_key_dict=Mock(),
+ cache=Mock(),
+ data=data,
+ call_type="completion",
+ )
+
+ assert result is not None, "Hook should return modified data"
+ filtered = result["tools"]
+
+ # Native tools must survive
+ native_in_result = [
+ t for t in filtered if isinstance(t, dict) and t.get("type") == "function"
+ ]
+ assert (
+ len(native_in_result) == 2
+ ), f"Both native tools must survive, got {len(native_in_result)}"
+
+ # MCP tools should be filtered (top_k=2)
+ mcp_in_result = [t for t in filtered if not isinstance(t, dict)]
+ assert (
+ len(mcp_in_result) <= 2
+ ), f"MCP tools should be filtered to top_k=2, got {len(mcp_in_result)}"
+
+ # Total should be native + filtered MCP
+ assert len(filtered) <= 4, f"Expected at most 4 tools, got {len(filtered)}"
+
+ # Filter stats should be emitted (MCP tools were present)
+ assert "litellm_semantic_filter_stats" in result["metadata"]
+
+ # Stats should report MCP-only counts, not inflated with native tools
+ stats = result["metadata"]["litellm_semantic_filter_stats"]
+ mcp_before, mcp_after = stats.split("->")
+ assert (
+ int(mcp_before) == 5
+ ), f"Stats 'from' should be MCP count (5), got {mcp_before}"
+ assert int(mcp_after) == len(
+ mcp_in_result
+ ), f"Stats 'to' should match filtered MCP count, got {mcp_after}"
+
+ print(
+ f"✅ Hook preserves native tools: {len(all_tools)} -> {len(filtered)} "
+ f"({len(native_in_result)} native + {len(mcp_in_result)} MCP), "
+ f"stats={stats}"
+ )
+
+
+@pytest.mark.asyncio
+async def test_semantic_filter_hook_all_native_tools():
+ """
+ Regression test: all-native request.
+
+ Given: Only native OpenAI-format function tools (none registered in
+ the MCP semantic router)
+ When: The hook processes the request
+ Then: All tools pass through, and NO spurious semantic filter response
+ headers are emitted (no litellm_semantic_filter_stats in metadata).
+ """
+ from litellm.proxy._experimental.mcp_server.semantic_tool_filter import (
+ SemanticMCPToolFilter,
+ )
+ from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook
+
+ mock_router = Mock()
+ filter_instance = SemanticMCPToolFilter(
+ embedding_model="text-embedding-3-small",
+ litellm_router_instance=mock_router,
+ top_k=3,
+ similarity_threshold=0.3,
+ enabled=True,
+ )
+
+ # Build router with some MCP tools (so tool_router is not None)
+ mcp_tools = [
+ MCPTool(
+ name="some_mcp_tool",
+ description="An MCP tool",
+ inputSchema={"type": "object"},
+ )
+ ]
+
+ from litellm.types.utils import Embedding, EmbeddingResponse
+
+ def mock_embedding_sync(*args, **kwargs):
+ return EmbeddingResponse(
+ data=[Embedding(embedding=[0.1] * 1536, index=0, object="embedding")],
+ model="text-embedding-3-small",
+ object="list",
+ usage={"prompt_tokens": 10, "total_tokens": 10},
+ )
+
+ async def mock_embedding_async(*args, **kwargs):
+ return mock_embedding_sync()
+
+ mock_router.embedding = mock_embedding_sync
+ mock_router.aembedding = mock_embedding_async
+
+ filter_instance._build_router(mcp_tools)
+
+ # --- Only native tools in the request ---
+ native_tools = [
+ {
+ "type": "function",
+ "function": {
+ "name": f"native_func_{i}",
+ "description": f"Native function {i}",
+ "parameters": {"type": "object", "properties": {}},
+ },
+ }
+ for i in range(3)
+ ]
+
+ hook = SemanticToolFilterHook(filter_instance)
+
+ data = {
+ "model": "gpt-4",
+ "messages": [{"role": "user", "content": "Hello"}],
+ "tools": native_tools,
+ "metadata": {},
+ }
+
+ result = await hook.async_pre_call_hook(
+ user_api_key_dict=Mock(),
+ cache=Mock(),
+ data=data,
+ call_type="completion",
+ )
+
+ assert result is not None, "Hook should return modified data"
+ filtered = result["tools"]
+
+ # All native tools must pass through
+ assert (
+ len(filtered) == 3
+ ), f"All 3 native tools must pass through, got {len(filtered)}"
+
+ # No spurious semantic filter stats (P2 fix)
+ assert (
+ "litellm_semantic_filter_stats" not in result["metadata"]
+ ), "Should NOT emit semantic filter stats for all-native-tool requests"
+
+ print(
+ f"✅ Hook passes through all {len(filtered)} native tools, "
+ f"no spurious filter headers emitted"
+ )
+
+
+@pytest.mark.asyncio
+async def test_semantic_filter_hook_responses_api_name_collision():
+ """
+ Regression test: Responses API native tool with MCP-matching name.
+
+ Given: A Responses-API native tool whose top-level ``name`` collides
+ with an MCP canonical name in ``_tool_map``
+ When: The hook classifies tools
+ Then: The native tool must NOT be sent to the semantic filter, even
+ though its name matches an MCP canonical.
+ """
+ from litellm.proxy._experimental.mcp_server.semantic_tool_filter import (
+ SemanticMCPToolFilter,
+ )
+ from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook
+ from litellm.types.utils import Embedding, EmbeddingResponse
+
+ mock_router = Mock()
+
+ def mock_embedding_sync(*args, **kwargs):
+ return EmbeddingResponse(
+ data=[Embedding(embedding=[0.1] * 1536, index=0, object="embedding")],
+ model="text-embedding-3-small",
+ object="list",
+ usage={"prompt_tokens": 10, "total_tokens": 10},
+ )
+
+ async def mock_embedding_async(*args, **kwargs):
+ return mock_embedding_sync()
+
+ mock_router.embedding = mock_embedding_sync
+ mock_router.aembedding = mock_embedding_async
+
+ filter_instance = SemanticMCPToolFilter(
+ embedding_model="text-embedding-3-small",
+ litellm_router_instance=mock_router,
+ top_k=2,
+ similarity_threshold=0.3,
+ enabled=True,
+ )
+
+ # Register an MCP tool with name "github-search"
+ mcp_tools = [
+ MCPTool(
+ name="github-search",
+ description="Search GitHub repos",
+ inputSchema={"type": "object"},
+ )
+ ]
+ filter_instance._build_router(mcp_tools)
+
+ # Responses API native tool with SAME name as MCP canonical
+ responses_api_tool = {
+ "type": "function",
+ "name": "github-search",
+ "description": "Caller-owned search tool",
+ "parameters": {"type": "object"},
+ }
+
+ hook = SemanticToolFilterHook(filter_instance)
+
+ # Verify classification: should be native, not MCP
+ assert not hook._is_mcp_tool(responses_api_tool), (
+ "Responses API tool with type=function + top-level name "
+ "should be classified as native, not MCP"
+ )
+
+ # Full hook test: all-native request should preserve tools
+ data = {
+ "model": "gpt-4",
+ "messages": [{"role": "user", "content": "Search GitHub"}],
+ "tools": [responses_api_tool],
+ "metadata": {},
+ }
+
+ result = await hook.async_pre_call_hook(
+ user_api_key_dict=Mock(),
+ cache=Mock(),
+ data=data,
+ call_type="completion",
+ )
+
+ # All tools are native → hook returns data with all tools preserved
+ filtered = (result or data)["tools"]
+ assert len(filtered) == 1, f"Native tool must survive, got {len(filtered)}"
+ assert filtered[0]["name"] == "github-search"
+
+ print("✅ Responses API tool with MCP-matching name correctly classified as native")
+
+
+@pytest.mark.asyncio
+async def test_semantic_filter_hook_preserves_tool_order():
+ """
+ Regression test: tool ordering preservation.
+
+ Given: An interleaved request [mcp_tool_A, native_tool, mcp_tool_B]
+ When: The hook filters tools (all MCP tools survive)
+ Then: The output order must match the original request order,
+ NOT native-first.
+ """
+ from litellm.proxy._experimental.mcp_server.semantic_tool_filter import (
+ SemanticMCPToolFilter,
+ )
+ from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook
+ from litellm.types.utils import Embedding, EmbeddingResponse
+
+ mock_router = Mock()
+
+ def mock_embedding_sync(*args, **kwargs):
+ return EmbeddingResponse(
+ data=[Embedding(embedding=[0.1] * 1536, index=0, object="embedding")],
+ model="text-embedding-3-small",
+ object="list",
+ usage={"prompt_tokens": 10, "total_tokens": 10},
+ )
+
+ async def mock_embedding_async(*args, **kwargs):
+ return mock_embedding_sync()
+
+ mock_router.embedding = mock_embedding_sync
+ mock_router.aembedding = mock_embedding_async
+
+ filter_instance = SemanticMCPToolFilter(
+ embedding_model="text-embedding-3-small",
+ litellm_router_instance=mock_router,
+ top_k=5,
+ similarity_threshold=0.3,
+ enabled=True,
+ )
+
+ # Register MCP tools
+ mcp_tool_a = MCPTool(
+ name="github-search",
+ description="Search GitHub",
+ inputSchema={"type": "object"},
+ )
+ mcp_tool_b = MCPTool(
+ name="github-issue",
+ description="Create GitHub issue",
+ inputSchema={"type": "object"},
+ )
+ filter_instance._build_router([mcp_tool_a, mcp_tool_b])
+
+ # Mock filter_tools to return both MCP tools (deterministic)
+ filter_instance.filter_tools = AsyncMock( # type: ignore[method-assign]
+ return_value=[mcp_tool_a, mcp_tool_b]
+ )
+
+ # Native tool (interleaved between MCP tools)
+ native_tool = {
+ "type": "function",
+ "function": {
+ "name": "weather_lookup",
+ "description": "Look up weather",
+ "parameters": {"type": "object", "properties": {}},
+ },
+ }
+
+ # Original order: [mcp_A, native, mcp_B]
+ original_tools = [mcp_tool_a, native_tool, mcp_tool_b]
+
+ hook = SemanticToolFilterHook(filter_instance)
+
+ data = {
+ "model": "gpt-4",
+ "messages": [{"role": "user", "content": "Search GitHub and check weather"}],
+ "tools": original_tools,
+ "metadata": {},
+ }
+
+ result = await hook.async_pre_call_hook(
+ user_api_key_dict=Mock(),
+ cache=Mock(),
+ data=data,
+ call_type="completion",
+ )
+
+ assert result is not None, "Hook should return modified data"
+ filtered = result["tools"]
+
+ # All tools should survive
+ assert len(filtered) == 3, f"Expected 3 tools, got {len(filtered)}"
+
+ # Order must be preserved: [mcp_A, native, mcp_B]
+ assert filtered[0] is mcp_tool_a, "First tool should be mcp_tool_a"
+ assert filtered[1] is native_tool, "Second tool should be native_tool"
+ assert filtered[2] is mcp_tool_b, "Third tool should be mcp_tool_b"
+
+ print(
+ "✅ Tool ordering preserved: [mcp_A, native, mcp_B] maintained after filtering"
+ )
+
+
class TestGetToolsByNames:
"""
Regression coverage for SemanticMCPToolFilter._get_tools_by_names
@@ -489,9 +913,7 @@ class TestGetToolsByNames:
{"name": "send_email", "description": "send mail"},
]
- matched = filter_instance._get_tools_by_names(
- ["send_email"], available_tools
- )
+ matched = filter_instance._get_tools_by_names(["send_email"], available_tools)
assert len(matched) == 1
assert matched[0]["name"] == "send_email"
@@ -503,9 +925,7 @@ class TestGetToolsByNames:
client_name = "litellm_" + canonical
available_tools = [{"name": client_name, "description": "scrape"}]
- matched = filter_instance._get_tools_by_names(
- [canonical], available_tools
- )
+ matched = filter_instance._get_tools_by_names([canonical], available_tools)
assert len(matched) == 1
# Must return the incoming tool unchanged so the client-facing
@@ -516,13 +936,9 @@ class TestGetToolsByNames:
"""Some clients use dash as alias separator; accept that too."""
filter_instance = self._make_filter()
canonical = "weather_svc-get_weather"
- available_tools = [
- {"name": "mcp-" + canonical, "description": "weather"}
- ]
+ available_tools = [{"name": "mcp-" + canonical, "description": "weather"}]
- matched = filter_instance._get_tools_by_names(
- [canonical], available_tools
- )
+ matched = filter_instance._get_tools_by_names([canonical], available_tools)
assert len(matched) == 1
assert matched[0]["name"] == "mcp-" + canonical
@@ -552,9 +968,7 @@ class TestGetToolsByNames:
{"name": "litellm_" + canonical, "description": "wrapped"},
]
- matched = filter_instance._get_tools_by_names(
- [canonical], available_tools
- )
+ matched = filter_instance._get_tools_by_names([canonical], available_tools)
assert len(matched) == 1
assert matched[0]["name"] == canonical
@@ -567,9 +981,7 @@ class TestGetToolsByNames:
separator-anchored suffixes of ``litellm_api-fs-read_file``.
"""
filter_instance = self._make_filter()
- available_tools = [
- {"name": "litellm_api-fs-read_file", "description": "read"}
- ]
+ available_tools = [{"name": "litellm_api-fs-read_file", "description": "read"}]
matched = filter_instance._get_tools_by_names(
["fs-read_file", "api-fs-read_file"], available_tools
@@ -590,9 +1002,7 @@ class TestGetToolsByNames:
{"name": "my_" + canonical, "description": "plain search"},
]
- matched = filter_instance._get_tools_by_names(
- [canonical], available_tools
- )
+ matched = filter_instance._get_tools_by_names([canonical], available_tools)
assert len(matched) == 1
assert matched[0]["name"] == "my_" + canonical
diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py
index 702bae77339..c8dc0ea5ed6 100644
--- a/tests/test_litellm/proxy/auth/test_auth_checks.py
+++ b/tests/test_litellm/proxy/auth/test_auth_checks.py
@@ -351,6 +351,43 @@ async def test_can_key_call_model_all_team_models_no_team_id_is_denied():
assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied
+@pytest.mark.asyncio
+async def test_can_team_access_model_all_team_models_expands_router_models():
+ from litellm import Router
+ from litellm.proxy._types import SpecialModelNames
+ from litellm.proxy.auth.auth_checks import can_team_access_model
+
+ team_object = LiteLLM_TeamTable(
+ team_id="team-123",
+ models=[SpecialModelNames.all_team_models.value],
+ )
+ router = Router(
+ model_list=[
+ {
+ "model_name": "allowed-model",
+ "litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test"},
+ }
+ ]
+ )
+
+ assert (
+ await can_team_access_model(
+ model="allowed-model",
+ team_object=team_object,
+ llm_router=router,
+ )
+ is True
+ )
+ with pytest.raises(ProxyException) as exc_info:
+ await can_team_access_model(
+ model="blocked-model",
+ team_object=team_object,
+ llm_router=router,
+ )
+
+ assert exc_info.value.type == ProxyErrorTypes.team_model_access_denied
+
+
@pytest.mark.asyncio
async def test_get_key_object_should_reconnect_once_on_db_connection_error():
mock_prisma_client = MagicMock()
diff --git a/tests/test_litellm/proxy/auth/test_model_checks.py b/tests/test_litellm/proxy/auth/test_model_checks.py
index 02b1f698132..261485e8965 100644
--- a/tests/test_litellm/proxy/auth/test_model_checks.py
+++ b/tests/test_litellm/proxy/auth/test_model_checks.py
@@ -543,3 +543,77 @@ async def test_get_available_models_for_user_expands_query_team_wildcard(
)
assert "openai/gpt-4o-mini" in result
+
+
+def test_get_key_models_all_team_models_recursive_team():
+ """GH#30619: when key and team both have all-team-models,
+ the sentinel should expand to proxy_model_list."""
+ from litellm.proxy.auth.model_checks import get_key_models
+ from litellm.proxy._types import SpecialModelNames
+
+ user_api_key_dict = type(
+ "obj", (object,),
+ {
+ "models": [SpecialModelNames.all_team_models.value],
+ "team_id": "team-1",
+ "team_models": [SpecialModelNames.all_team_models.value],
+ },
+ )()
+ proxy_model_list = ["model-a", "model-b"]
+ result = get_key_models(user_api_key_dict, proxy_model_list, {})
+ assert SpecialModelNames.all_team_models.value not in result
+ assert set(result) == {"model-a", "model-b"}
+
+
+def test_get_key_models_all_team_models_keeps_mixed_team_entries():
+ from litellm.proxy.auth.model_checks import get_key_models
+ from litellm.proxy._types import SpecialModelNames
+
+ user_api_key_dict = type(
+ "obj",
+ (object,),
+ {
+ "models": [SpecialModelNames.all_team_models.value],
+ "team_id": "team-1",
+ "team_models": [
+ SpecialModelNames.all_team_models.value,
+ "restricted-model",
+ ],
+ },
+ )()
+ result = get_key_models(user_api_key_dict, ["model-a", "model-b"], {})
+ assert SpecialModelNames.all_team_models.value not in result
+ assert set(result) == {"model-a", "model-b", "restricted-model"}
+
+
+def test_get_team_models_all_team_models_expands():
+ """GH#30619: all-team-models in team_models should expand."""
+ from litellm.proxy.auth.model_checks import get_team_models
+ from litellm.proxy._types import SpecialModelNames
+
+ result = get_team_models(
+ [SpecialModelNames.all_team_models.value],
+ ["model-a", "model-b"],
+ {},
+ )
+ assert SpecialModelNames.all_team_models.value not in result
+ assert set(result) == {"model-a", "model-b"}
+
+
+def test_get_team_models_all_team_models_expands_with_access_groups():
+ """GH#30619: all-team-models with include_model_access_groups
+ should include access group keys."""
+ from litellm.proxy.auth.model_checks import get_team_models
+ from litellm.proxy._types import SpecialModelNames
+
+ result = get_team_models(
+ [SpecialModelNames.all_team_models.value],
+ ["model-a", "model-b"],
+ {"group-1": ["g1-model"], "group-2": ["g2-model"]},
+ include_model_access_groups=True,
+ )
+ assert SpecialModelNames.all_team_models.value not in result
+ assert "model-a" in result
+ assert "model-b" in result
+ assert "group-1" in result
+ assert "group-2" in result
diff --git a/tests/test_litellm/proxy/db/test_db_url_settings.py b/tests/test_litellm/proxy/db/test_db_url_settings.py
index b2212068a5b..573bd5ae584 100644
--- a/tests/test_litellm/proxy/db/test_db_url_settings.py
+++ b/tests/test_litellm/proxy/db/test_db_url_settings.py
@@ -16,7 +16,11 @@ from unittest.mock import patch
import pytest
-from litellm.proxy.db.db_url_settings import DatabaseURLSettings
+from litellm.proxy.db.db_url_settings import (
+ DatabaseURLSettings,
+ unsupported_db_scheme,
+ unsupported_db_scheme_message,
+)
def _apply() -> bool:
@@ -27,6 +31,7 @@ def _apply() -> bool:
_MANAGED_DB_ENV_VARS = (
"IAM_TOKEN_DB_AUTH",
"DATABASE_URL",
+ "DIRECT_URL",
"DATABASE_URL_READ_REPLICA",
"DATABASE_HOST",
"DATABASE_PORT",
@@ -287,3 +292,87 @@ def test_password_reader_uses_own_credentials(monkeypatch):
os.environ["DATABASE_URL_READ_REPLICA"]
== "postgresql://litellm_ro:ro_pw@reader.example.com:5432/litellm_db"
)
+
+
+@pytest.mark.parametrize(
+ "url",
+ [
+ "postgresql://u:p@host:5432/db",
+ "postgres://u:p@host:5432/db",
+ "POSTGRESQL://u:p@host:5432/db",
+ "postgresql://host/db?schema=public",
+ ],
+)
+def test_unsupported_db_scheme_accepts_postgres(url):
+ assert unsupported_db_scheme(url) is None
+
+
+@pytest.mark.parametrize(
+ "url,scheme",
+ [
+ ("sqlite:///data/litellm.db", "sqlite"),
+ ("sqlite:///./local.db", "sqlite"),
+ ("mysql://u:p@host:3306/db", "mysql"),
+ ("mssql://host/db", "mssql"),
+ ],
+)
+def test_unsupported_db_scheme_rejects_non_postgres(url, scheme):
+ assert unsupported_db_scheme(url) == scheme
+
+
+def test_unsupported_db_scheme_does_not_echo_schemeless_credentials():
+ """A malformed schemeless DSN must not leak its embedded credentials
+ through the return value (which callers log)."""
+ leaky = "litellm:s3cr3t_password@db.internal:5432/litellm"
+
+ result = unsupported_db_scheme(leaky)
+
+ assert result is not None
+ assert "s3cr3t_password" not in result
+ assert "db.internal" not in result
+
+
+def test_apply_to_env_rejects_pinned_sqlite_writer(monkeypatch):
+ """Componentized entrypoints pin DATABASE_URL and call apply_to_env; a
+ sqlite writer must raise here rather than reach Prisma."""
+ monkeypatch.setenv("DATABASE_URL", "sqlite:///data/litellm.db")
+
+ with pytest.raises(RuntimeError, match="sqlite"):
+ _apply()
+
+ # The bad URL must not have been propagated as a usable connection string.
+ assert os.environ["DATABASE_URL"] == "sqlite:///data/litellm.db"
+
+
+def test_apply_to_env_rejects_pinned_sqlite_direct_url(monkeypatch):
+ """DIRECT_URL reaches Prisma the same way DATABASE_URL does; a non-postgres
+ direct URL must be rejected in apply_to_env, matching the CLI startup guard."""
+ monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@writer.example.com:5432/db")
+ monkeypatch.setenv("DIRECT_URL", "sqlite:///data/litellm.db")
+
+ with pytest.raises(RuntimeError, match="DIRECT_URL.*sqlite"):
+ _apply()
+
+
+def test_apply_to_env_rejects_pinned_non_postgres_reader(monkeypatch):
+ monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@writer.example.com:5432/db")
+ monkeypatch.setenv(
+ "DATABASE_URL_READ_REPLICA", "mysql://u:p@reader.example.com:3306/db"
+ )
+
+ with pytest.raises(RuntimeError, match="DATABASE_URL_READ_REPLICA.*mysql"):
+ _apply()
+
+
+def test_apply_to_env_accepts_pinned_postgres(monkeypatch):
+ monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@host:5432/db")
+
+ # Operator-pinned URL: nothing reassembled, no error.
+ assert _apply() is False
+
+
+def test_unsupported_db_scheme_message_names_var_and_scheme():
+ msg = unsupported_db_scheme_message("DIRECT_URL", "sqlite")
+ assert "DIRECT_URL" in msg
+ assert "sqlite" in msg
+ assert "postgresql://" in msg
diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py
index 0b5b5fb6ceb..f40904e234d 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py
@@ -1197,6 +1197,136 @@ class TestListMCPServers:
# Non-admin viewers get no env var config at all (not even names).
assert result.env_vars is None
+ @pytest.mark.asyncio
+ async def test_fetch_single_mcp_server_sanitizes_for_view_only_admin(self):
+ """PROXY_ADMIN_VIEW_ONLY must NOT see credential-bearing fields.
+
+ It previously passed the _user_has_admin_view gate (which also grants view-only
+ admins) and only had the explicit `credentials` field cleared, leaking secrets
+ embedded in url/static_headers/env_vars. Only a FULL PROXY_ADMIN may see those.
+ This test exercises the real role helpers (no patching of the gate)."""
+ mock_server = LiteLLM_MCPServerTable.model_construct(
+ server_id="leaky-server",
+ server_name="Leaky Server",
+ alias="Leaky Server",
+ transport=MCPTransport.http,
+ url="https://leaky.example.com/mcp?api_key=sk-embedded-in-url",
+ static_headers={"Authorization": "Bearer sk-secret-header"},
+ env={"UPSTREAM_TOKEN": "sk-secret-env"},
+ env_vars=[
+ {"name": "GLOBAL_KEY", "value": "super-secret", "scope": "global"},
+ ],
+ credentials={"auth_value": "sk-explicit-credential"},
+ )
+
+ mock_prisma_client = MagicMock()
+
+ mock_health_result = generate_mock_mcp_server_db_record(
+ server_id="leaky-server", alias="Leaky Server"
+ )
+ mock_health_result.status = "healthy"
+ mock_health_result.last_health_check = datetime.now()
+ mock_health_result.health_check_error = None
+
+ mock_manager = MagicMock()
+ mock_manager.add_server = AsyncMock()
+ mock_manager.health_check_server = AsyncMock(return_value=mock_health_result)
+
+ mock_user_auth = generate_mock_user_api_key_auth(
+ user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY
+ )
+
+ with (
+ patch(
+ "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
+ return_value=mock_prisma_client,
+ ),
+ patch(
+ "litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
+ AsyncMock(return_value=mock_server),
+ ),
+ patch(
+ "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
+ mock_manager,
+ ),
+ ):
+ from litellm.proxy.management_endpoints.mcp_management_endpoints import (
+ fetch_mcp_server,
+ )
+
+ result = await fetch_mcp_server(
+ request=_make_mock_request(),
+ server_id="leaky-server",
+ user_api_key_dict=mock_user_auth,
+ )
+
+ assert result.server_id == "leaky-server"
+ assert result.credentials is None
+ assert result.url is None
+ assert result.static_headers is None
+ assert result.env == {}
+ assert result.env_vars is None
+
+ @pytest.mark.asyncio
+ async def test_fetch_single_mcp_server_full_admin_still_sees_secrets(self):
+ """the fix must not over-redact for FULL PROXY_ADMIN,
+ who needs url/static_headers/env to populate the edit form."""
+ mock_server = LiteLLM_MCPServerTable.model_construct(
+ server_id="admin-server",
+ server_name="Admin Server",
+ alias="Admin Server",
+ transport=MCPTransport.http,
+ url="https://admin.example.com/mcp",
+ static_headers={"Authorization": "Bearer sk-secret-header"},
+ credentials={"auth_value": "sk-explicit-credential"},
+ )
+
+ mock_prisma_client = MagicMock()
+
+ mock_health_result = generate_mock_mcp_server_db_record(
+ server_id="admin-server", alias="Admin Server"
+ )
+ mock_health_result.status = "healthy"
+ mock_health_result.last_health_check = datetime.now()
+ mock_health_result.health_check_error = None
+
+ mock_manager = MagicMock()
+ mock_manager.add_server = AsyncMock()
+ mock_manager.health_check_server = AsyncMock(return_value=mock_health_result)
+
+ mock_user_auth = generate_mock_user_api_key_auth(
+ user_role=LitellmUserRoles.PROXY_ADMIN
+ )
+
+ with (
+ patch(
+ "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
+ return_value=mock_prisma_client,
+ ),
+ patch(
+ "litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
+ AsyncMock(return_value=mock_server),
+ ),
+ patch(
+ "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
+ mock_manager,
+ ),
+ ):
+ from litellm.proxy.management_endpoints.mcp_management_endpoints import (
+ fetch_mcp_server,
+ )
+
+ result = await fetch_mcp_server(
+ request=_make_mock_request(),
+ server_id="admin-server",
+ user_api_key_dict=mock_user_auth,
+ )
+
+ # credentials field is always redacted; the rest must survive for full admin.
+ assert result.credentials is None
+ assert result.url == "https://admin.example.com/mcp"
+ assert result.static_headers == {"Authorization": "Bearer sk-secret-header"}
+
class TestTeamScopedMCPServerAccess:
"""Tests for cross-team information disclosure and restricted key bypass fixes."""
@@ -3338,6 +3468,10 @@ async def test_store_mcp_oauth_user_credential_returns_status():
return_value=generate_mock_mcp_server_db_record(server_id=server_id)
),
),
+ patch(
+ "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view",
+ return_value=True,
+ ),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.store_user_oauth_credential",
new=AsyncMock(return_value=None),
@@ -3693,18 +3827,12 @@ def _server_with_env_vars(server_id: str = "srv-env"):
@pytest.mark.asyncio
-@pytest.mark.parametrize(
- "user_role, expected_global_value",
- [
- (LitellmUserRoles.PROXY_ADMIN, "super-secret"),
- (LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, ""),
- ],
-)
-async def test_fetch_single_mcp_server_redacts_global_env_for_view_only_admin(
- user_role, expected_global_value
-):
- """Read-only admins must not receive admin-supplied global env var secrets;
- full admins still see them so the edit form can pre-fill."""
+async def test_fetch_single_mcp_server_env_vars_full_admin_vs_view_only():
+ """full admins see admin-supplied global env var secrets so the edit
+ form can pre-fill; read-only admins now go through the non-admin sanitizer, which
+ drops env_vars entirely (the names alone, e.g. ADMIN_API_KEY, leak what secrets the
+ admin configured). Previously the view-only case merely blanked the global value
+ while keeping the names, which still leaked configuration metadata."""
server = _server_with_env_vars()
health_result = generate_mock_mcp_server_db_record(server_id=server.server_id)
@@ -3712,34 +3840,39 @@ async def test_fetch_single_mcp_server_redacts_global_env_for_view_only_admin(
health_result.last_health_check = datetime.now()
health_result.health_check_error = None
- with (
- patch(
- "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
- return_value=MagicMock(),
- ),
- patch(
- "litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
- AsyncMock(return_value=server),
- ),
- patch(
- "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager.add_server",
- AsyncMock(return_value=None),
- ),
- patch(
- "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager.health_check_server",
- AsyncMock(return_value=health_result),
- ),
- ):
- result = await mgmt_endpoints.fetch_mcp_server(
- request=_make_mock_request(),
- server_id=server.server_id,
- user_api_key_dict=generate_mock_user_api_key_auth(user_role=user_role),
- )
+ async def _fetch(user_role):
+ with (
+ patch(
+ "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
+ return_value=MagicMock(),
+ ),
+ patch(
+ "litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
+ AsyncMock(return_value=server),
+ ),
+ patch(
+ "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager.add_server",
+ AsyncMock(return_value=None),
+ ),
+ patch(
+ "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager.health_check_server",
+ AsyncMock(return_value=health_result),
+ ),
+ ):
+ return await mgmt_endpoints.fetch_mcp_server(
+ request=_make_mock_request(),
+ server_id=server.server_id,
+ user_api_key_dict=generate_mock_user_api_key_auth(user_role=user_role),
+ )
- by_name = {ev.name: ev for ev in result.env_vars}
- assert by_name["ADMIN_API_KEY"].value == expected_global_value
- # Per-user placeholders are always preserved.
+ full_admin = await _fetch(LitellmUserRoles.PROXY_ADMIN)
+ by_name = {ev.name: ev for ev in full_admin.env_vars}
+ assert by_name["ADMIN_API_KEY"].value == "super-secret"
assert by_name["USER_TOKEN"].value == "placeholder-hint"
+
+ view_only = await _fetch(LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY)
+ assert view_only.env_vars is None
+
# The source record must never be mutated.
assert {ev.name: ev.value for ev in server.env_vars}[
"ADMIN_API_KEY"
@@ -3747,18 +3880,67 @@ async def test_fetch_single_mcp_server_redacts_global_env_for_view_only_admin(
@pytest.mark.asyncio
-@pytest.mark.parametrize(
- "user_role, expected_global_value",
- [
- (LitellmUserRoles.PROXY_ADMIN, "super-secret"),
- (LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, ""),
- ],
-)
-async def test_fetch_all_mcp_servers_redacts_global_env_for_view_only_admin(
- user_role, expected_global_value
-):
+async def test_fetch_all_mcp_servers_env_vars_full_admin_vs_view_only():
+ """same posture as the single-server fetch. Full admins
+ keep the env var values; view-only admins get env_vars dropped via the non-admin
+ sanitizer rather than only having the global value blanked."""
server = _server_with_env_vars()
+ async def _fetch_all(user_role):
+ with (
+ patch(
+ "litellm.proxy.management_endpoints.mcp_management_endpoints._get_user_mcp_management_mode",
+ return_value="view_all",
+ ),
+ patch(
+ "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager.get_all_mcp_servers_unfiltered",
+ AsyncMock(return_value=[server]),
+ ),
+ patch(
+ "litellm.proxy.proxy_server.prisma_client",
+ None,
+ ),
+ ):
+ return await mgmt_endpoints.fetch_all_mcp_servers(
+ user_api_key_dict=generate_mock_user_api_key_auth(user_role=user_role),
+ )
+
+ full_admin = await _fetch_all(LitellmUserRoles.PROXY_ADMIN)
+ by_name = {ev.name: ev for ev in full_admin[0].env_vars}
+ assert by_name["ADMIN_API_KEY"].value == "super-secret"
+ assert by_name["USER_TOKEN"].value == "placeholder-hint"
+
+ view_only = await _fetch_all(LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY)
+ assert view_only[0].env_vars is None
+
+ assert {ev.name: ev.value for ev in server.env_vars}[
+ "ADMIN_API_KEY"
+ ] == "super-secret"
+
+
+def _leaky_list_server() -> "LiteLLM_MCPServerTable":
+ """A server whose url/static_headers/env carry embedded secrets, for the
+ list-endpoint sanitization tests. ``model_construct`` skips validation so
+ the raw values survive verbatim."""
+ return LiteLLM_MCPServerTable.model_construct(
+ server_id="leaky-list-server",
+ server_name="Leaky List Server",
+ alias="Leaky List Server",
+ transport=MCPTransport.http,
+ url="https://leaky.example.com/mcp?api_key=sk-embedded-in-url",
+ static_headers={"Authorization": "Bearer sk-secret-header"},
+ env={"UPSTREAM_TOKEN": "sk-secret-env"},
+ env_vars=[
+ {"name": "GLOBAL_KEY", "value": "super-secret", "scope": "global"},
+ ],
+ credentials={"auth_value": "sk-explicit-credential"},
+ )
+
+
+async def _fetch_all_via_view_all(user_role: LitellmUserRoles):
+ """Drive GET /v1/mcp/server in view_all mode for the given role using the
+ real role helpers (the full-admin gate is never patched)."""
+ server = _leaky_list_server()
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_user_mcp_management_mode",
@@ -3776,13 +3958,46 @@ async def test_fetch_all_mcp_servers_redacts_global_env_for_view_only_admin(
result = await mgmt_endpoints.fetch_all_mcp_servers(
user_api_key_dict=generate_mock_user_api_key_auth(user_role=user_role),
)
+ return server, result
- by_name = {ev.name: ev for ev in result[0].env_vars}
- assert by_name["ADMIN_API_KEY"].value == expected_global_value
- assert by_name["USER_TOKEN"].value == "placeholder-hint"
- assert {ev.name: ev.value for ev in server.env_vars}[
- "ADMIN_API_KEY"
- ] == "super-secret"
+
+@pytest.mark.asyncio
+async def test_list_mcp_servers_sanitized_for_view_only_admin():
+ """PROXY_ADMIN_VIEW_ONLY listing servers must go through the non-admin
+ sanitizer: url and static_headers cleared, env emptied, env_vars dropped.
+ A mutation swapping _user_is_full_admin() back to _user_has_admin_view()
+ (which also grants view-only admins) would return the raw url/headers and
+ fail this. The real role helpers are exercised; the gate is not patched."""
+ source, result = await _fetch_all_via_view_all(
+ LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY
+ )
+
+ assert len(result) == 1
+ sanitized = result[0]
+ assert sanitized.server_id == "leaky-list-server"
+ assert sanitized.url is None
+ assert sanitized.static_headers is None
+ assert sanitized.env == {}
+ assert sanitized.env_vars is None
+ assert sanitized.credentials is None
+
+ # The source record must never be mutated by sanitization.
+ assert source.url == "https://leaky.example.com/mcp?api_key=sk-embedded-in-url"
+ assert source.static_headers == {"Authorization": "Bearer sk-secret-header"}
+
+
+@pytest.mark.asyncio
+async def test_list_mcp_servers_full_admin_still_sees_secrets():
+ """The view-only redaction must not over-redact for a FULL PROXY_ADMIN,
+ who needs url/static_headers to populate the edit form. Only the explicit
+ credentials field is cleared for full admins on the list endpoint."""
+ _, result = await _fetch_all_via_view_all(LitellmUserRoles.PROXY_ADMIN)
+
+ assert len(result) == 1
+ raw = result[0]
+ assert raw.url == "https://leaky.example.com/mcp?api_key=sk-embedded-in-url"
+ assert raw.static_headers == {"Authorization": "Bearer sk-secret-header"}
+ assert raw.credentials is None
def _make_env_var_server(
@@ -4347,8 +4562,13 @@ class TestMCPUserEnvVarsAccessControl:
),
patch.object(
mgmt_endpoints,
- "get_all_mcp_servers_for_user",
- AsyncMock(return_value=[_make_env_var_server(server_id="other")]),
+ "build_effective_auth_contexts",
+ AsyncMock(return_value=[object()]),
+ ),
+ patch.object(
+ mgmt_endpoints.global_mcp_server_manager,
+ "get_allowed_mcp_servers",
+ AsyncMock(return_value=["other"]),
),
patch.object(mgmt_endpoints, "get_user_env_vars", get_user_env_vars),
):
@@ -4378,7 +4598,12 @@ class TestMCPUserEnvVarsAccessControl:
),
patch.object(
mgmt_endpoints,
- "get_all_mcp_servers_for_user",
+ "build_effective_auth_contexts",
+ AsyncMock(return_value=[object()]),
+ ),
+ patch.object(
+ mgmt_endpoints.global_mcp_server_manager,
+ "get_allowed_mcp_servers",
AsyncMock(return_value=[]),
),
patch.object(mgmt_endpoints, "merge_user_env_vars", merge_mock),
@@ -4412,7 +4637,12 @@ class TestMCPUserEnvVarsAccessControl:
),
patch.object(
mgmt_endpoints,
- "get_all_mcp_servers_for_user",
+ "build_effective_auth_contexts",
+ AsyncMock(return_value=[object()]),
+ ),
+ patch.object(
+ mgmt_endpoints.global_mcp_server_manager,
+ "get_allowed_mcp_servers",
AsyncMock(return_value=[]),
),
patch.object(mgmt_endpoints, "delete_user_env_vars", delete_mock),
@@ -4444,8 +4674,13 @@ class TestMCPUserEnvVarsAccessControl:
),
patch.object(
mgmt_endpoints,
- "get_all_mcp_servers_for_user",
- AsyncMock(return_value=[server]),
+ "build_effective_auth_contexts",
+ AsyncMock(return_value=[object()]),
+ ),
+ patch.object(
+ mgmt_endpoints.global_mcp_server_manager,
+ "get_allowed_mcp_servers",
+ AsyncMock(return_value=["srv-1"]),
),
patch.object(
mgmt_endpoints,
@@ -4465,13 +4700,13 @@ class TestMCPUserEnvVarsAccessControl:
@pytest.mark.asyncio
async def test_admin_bypasses_access_check(self):
- """Proxy admins must not be filtered by get_all_mcp_servers_for_user."""
+ """Proxy admins must not be filtered by the allowed-server check."""
server = _make_env_var_server(
server_id="srv-1",
env_vars=_ENV_VARS_MIXED,
static_headers=_STATIC_HEADERS_MIXED,
)
- access_list_mock = AsyncMock(return_value=[])
+ allowed_mock = AsyncMock(return_value=[])
with (
patch.object(
mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock()
@@ -4480,7 +4715,9 @@ class TestMCPUserEnvVarsAccessControl:
mgmt_endpoints, "get_mcp_server", AsyncMock(return_value=server)
),
patch.object(
- mgmt_endpoints, "get_all_mcp_servers_for_user", access_list_mock
+ mgmt_endpoints.global_mcp_server_manager,
+ "get_allowed_mcp_servers",
+ allowed_mock,
),
patch.object(
mgmt_endpoints, "get_user_env_vars", AsyncMock(return_value={})
@@ -4494,22 +4731,34 @@ class TestMCPUserEnvVarsAccessControl:
),
)
assert result.server_id == "srv-1"
- access_list_mock.assert_not_awaited()
+ allowed_mock.assert_not_awaited()
@pytest.mark.asyncio
async def test_non_admin_gets_403_not_404_for_inaccessible_server(self):
- """Authorization must run before the existence check so a non-admin
- cannot distinguish "server does not exist" (404) from "server exists but
- you lack access" (403) and enumerate server IDs."""
- get_mcp_server_mock = AsyncMock(return_value=None)
+ """A non-admin cannot distinguish "server does not exist" (404) from
+ "server exists but you lack access" (403): both collapse to 403 so server
+ ids stay non-enumerable, even when neither the DB nor the registry has the
+ server."""
with (
patch.object(
mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock()
),
- patch.object(mgmt_endpoints, "get_mcp_server", get_mcp_server_mock),
+ patch.object(
+ mgmt_endpoints, "get_mcp_server", AsyncMock(return_value=None)
+ ),
patch.object(
mgmt_endpoints,
- "get_all_mcp_servers_for_user",
+ "build_effective_auth_contexts",
+ AsyncMock(return_value=[object()]),
+ ),
+ patch.object(
+ mgmt_endpoints.global_mcp_server_manager,
+ "get_mcp_server_by_id",
+ MagicMock(return_value=None),
+ ),
+ patch.object(
+ mgmt_endpoints.global_mcp_server_manager,
+ "get_allowed_mcp_servers",
AsyncMock(return_value=[]),
),
):
@@ -4522,7 +4771,6 @@ class TestMCPUserEnvVarsAccessControl:
),
)
assert exc.value.status_code == 403
- get_mcp_server_mock.assert_not_awaited()
def test_oauth2_flow_accepted_on_create_request():
@@ -4572,3 +4820,217 @@ def test_oauth2_flow_defaults_to_none_when_omitted():
assert (
LiteLLM_MCPServerTable(server_id="srv-1", transport="http").oauth2_flow is None
)
+
+
+class TestPerUserCredentialConfigServerResolution:
+ """Per-user credential and env-var endpoints must resolve config-defined MCP
+ servers, which live only in the in-memory registry and never get a DB row, so
+ a user can store their BYOK key / OAuth token / env vars against them. The
+ same allowed-server authorization the MCP gateway enforces also gates these
+ writes for non-admins.
+ """
+
+ # 32-char sha256 stable id, the shape a config.yaml server gets.
+ CONFIG_SERVER_ID = "3a6a3f8633340371b49562c8c4682da9"
+
+ def _registry_only_manager(self, *, is_byok: bool = False):
+ """A manager mock where the server exists only in the registry (DB miss)."""
+ config_server = generate_mock_mcp_server_config_record(
+ server_id=self.CONFIG_SERVER_ID, name="Config Server"
+ )
+ record = generate_mock_mcp_server_db_record(
+ server_id=self.CONFIG_SERVER_ID
+ ).model_copy(update={"is_byok": is_byok})
+ manager = MagicMock()
+ 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.get_allowed_mcp_servers = AsyncMock(return_value=[])
+ return manager
+
+ @pytest.mark.asyncio
+ async def test_store_oauth_credential_resolves_config_server_for_admin(self):
+ """OBO token persists for a config-defined server (DB miss, registry hit).
+ Before the registry fallback this raised 404 "MCP Server not found"."""
+ manager = self._registry_only_manager()
+ store_mock = AsyncMock(return_value=None)
+ with (
+ patch.object(
+ mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock()
+ ),
+ patch.object(
+ mgmt_endpoints, "get_mcp_server", AsyncMock(return_value=None)
+ ),
+ patch.object(mgmt_endpoints, "global_mcp_server_manager", manager),
+ patch.object(mgmt_endpoints, "store_user_oauth_credential", store_mock),
+ patch.object(
+ mgmt_endpoints,
+ "get_user_oauth_credential",
+ AsyncMock(return_value={"expires_at": None}),
+ ),
+ ):
+ result = await mgmt_endpoints.store_mcp_oauth_user_credential(
+ server_id=self.CONFIG_SERVER_ID,
+ payload=mgmt_endpoints.MCPOAuthUserCredentialRequest(
+ access_token="tok", expires_in=3600
+ ),
+ user_api_key_dict=generate_mock_user_api_key_auth(user_id="admin"),
+ )
+ assert result.has_credential is True
+ store_mock.assert_awaited_once()
+ manager.get_mcp_server_by_id.assert_called_once_with(self.CONFIG_SERVER_ID)
+
+ @pytest.mark.asyncio
+ async def test_store_byok_credential_resolves_config_server_for_admin(self):
+ """BYOK key persists for a config-defined BYOK server (DB miss, registry
+ hit). Before the registry fallback this raised 404."""
+ manager = self._registry_only_manager(is_byok=True)
+ store_mock = AsyncMock(return_value=None)
+ with (
+ patch.object(
+ mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock()
+ ),
+ patch.object(
+ mgmt_endpoints, "get_mcp_server", AsyncMock(return_value=None)
+ ),
+ patch.object(mgmt_endpoints, "global_mcp_server_manager", manager),
+ patch.object(mgmt_endpoints, "store_user_credential", store_mock),
+ ):
+ result = await mgmt_endpoints.store_mcp_user_credential(
+ server_id=self.CONFIG_SERVER_ID,
+ payload=mgmt_endpoints.MCPUserCredentialRequest(credential="my-key"),
+ user_api_key_dict=generate_mock_user_api_key_auth(user_id="admin"),
+ )
+ assert result.has_credential is True
+ store_mock.assert_awaited_once()
+
+ @pytest.mark.asyncio
+ async def test_store_oauth_credential_forbidden_for_non_admin_without_access(self):
+ """A non-admin storing a credential for a server not in their allowed set
+ gets 403 and no row is written (the store endpoints had no authz before)."""
+ manager = self._registry_only_manager()
+ manager.get_allowed_mcp_servers = AsyncMock(return_value=[])
+ store_mock = AsyncMock(return_value=None)
+ with (
+ patch.object(
+ mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock()
+ ),
+ patch.object(
+ mgmt_endpoints, "get_mcp_server", AsyncMock(return_value=None)
+ ),
+ patch.object(mgmt_endpoints, "global_mcp_server_manager", manager),
+ patch.object(
+ mgmt_endpoints,
+ "build_effective_auth_contexts",
+ AsyncMock(return_value=[object()]),
+ ),
+ patch.object(mgmt_endpoints, "store_user_oauth_credential", store_mock),
+ ):
+ with pytest.raises(HTTPException) as exc:
+ await mgmt_endpoints.store_mcp_oauth_user_credential(
+ server_id=self.CONFIG_SERVER_ID,
+ payload=mgmt_endpoints.MCPOAuthUserCredentialRequest(
+ access_token="tok", expires_in=3600
+ ),
+ user_api_key_dict=generate_mock_user_api_key_auth(
+ user_id="alice", user_role=LitellmUserRoles.INTERNAL_USER
+ ),
+ )
+ assert exc.value.status_code == 403
+ store_mock.assert_not_awaited()
+
+ @pytest.mark.asyncio
+ async def test_store_oauth_credential_allowed_for_non_admin_with_access(self):
+ """A non-admin with the config server in their allowed set persists the
+ token; proves the non-admin authz uses the registry-aware allowed set."""
+ manager = self._registry_only_manager()
+ manager.get_allowed_mcp_servers = AsyncMock(
+ return_value=[self.CONFIG_SERVER_ID]
+ )
+ store_mock = AsyncMock(return_value=None)
+ with (
+ patch.object(
+ mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock()
+ ),
+ patch.object(
+ mgmt_endpoints, "get_mcp_server", AsyncMock(return_value=None)
+ ),
+ patch.object(mgmt_endpoints, "global_mcp_server_manager", manager),
+ patch.object(
+ mgmt_endpoints,
+ "build_effective_auth_contexts",
+ AsyncMock(return_value=[object()]),
+ ),
+ patch.object(mgmt_endpoints, "store_user_oauth_credential", store_mock),
+ patch.object(
+ mgmt_endpoints,
+ "get_user_oauth_credential",
+ AsyncMock(return_value={"expires_at": None}),
+ ),
+ ):
+ result = await mgmt_endpoints.store_mcp_oauth_user_credential(
+ server_id=self.CONFIG_SERVER_ID,
+ payload=mgmt_endpoints.MCPOAuthUserCredentialRequest(
+ access_token="tok", expires_in=3600
+ ),
+ user_api_key_dict=generate_mock_user_api_key_auth(
+ user_id="alice", user_role=LitellmUserRoles.INTERNAL_USER
+ ),
+ )
+ assert result.has_credential is True
+ store_mock.assert_awaited_once()
+
+ @pytest.mark.asyncio
+ async def test_store_env_vars_resolves_config_server_for_non_admin_with_access(
+ self,
+ ):
+ """Per-user env vars persist for a config server a non-admin may access.
+ The non-admin path previously used a DB-only access list that never
+ included config servers, so this 403'd before the fix."""
+ env_var_server = _make_env_var_server(
+ server_id=self.CONFIG_SERVER_ID,
+ env_vars=_ENV_VARS_MIXED,
+ static_headers=_STATIC_HEADERS_MIXED,
+ )
+ manager = MagicMock()
+ 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.get_allowed_mcp_servers = AsyncMock(
+ return_value=[self.CONFIG_SERVER_ID]
+ )
+ merge_mock = AsyncMock(return_value={"CORP_USERNAME": "alice"})
+ with (
+ patch.object(
+ mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock()
+ ),
+ patch.object(
+ mgmt_endpoints, "get_mcp_server", AsyncMock(return_value=None)
+ ),
+ patch.object(mgmt_endpoints, "global_mcp_server_manager", manager),
+ patch.object(
+ mgmt_endpoints,
+ "build_effective_auth_contexts",
+ AsyncMock(return_value=[object()]),
+ ),
+ patch.object(mgmt_endpoints, "merge_user_env_vars", merge_mock),
+ ):
+ result = await mgmt_endpoints.store_mcp_user_env_vars(
+ server_id=self.CONFIG_SERVER_ID,
+ payload=mgmt_endpoints.MCPUserEnvVarsRequest(
+ values={"CORP_USERNAME": "alice"}
+ ),
+ user_api_key_dict=generate_mock_user_api_key_auth(
+ user_id="alice", user_role=LitellmUserRoles.INTERNAL_USER
+ ),
+ )
+ merge_mock.assert_awaited_once()
+ _, _, _, updates, _ = merge_mock.await_args.args
+ assert updates == {"CORP_USERNAME": "alice"}
+ assert result.server_id == self.CONFIG_SERVER_ID
diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_config.py b/tests/test_litellm/proxy/proxy_server/test_routes_config.py
index e89ada5bdef..df14dc5b5dc 100644
--- a/tests/test_litellm/proxy/proxy_server/test_routes_config.py
+++ b/tests/test_litellm/proxy/proxy_server/test_routes_config.py
@@ -15,8 +15,6 @@ from __future__ import annotations
from unittest.mock import AsyncMock, MagicMock
-import pytest
-
from .conftest import VOLATILE_KEYS, normalize
@@ -248,6 +246,193 @@ def test_config_field_info_field_not_in_db(client, auth_as, mock_prisma, monkeyp
assert "not in DB" in response.json().get("detail", {}).get("error", "")
+def test_config_field_info_redacts_nested_secret_for_view_only_admin(
+ client, auth_as, mock_prisma, monkeypatch
+):
+ """A view-only admin reading a structured field must not receive nested
+ credentials. database_args carries aws_web_identity_token (a DynamoDB
+ role-assumption credential); it must come back redacted while non-secret
+ siblings like region_name stay visible."""
+ from litellm.proxy import proxy_server as ps
+ from litellm.proxy._types import LitellmUserRoles
+
+ table = _install_litellm_config(mock_prisma)
+ row = MagicMock()
+ row.param_value = {
+ "database_args": {
+ "region_name": "us-east-1",
+ "user_table_name": "LiteLLM_UserTable",
+ "aws_web_identity_token": "sk-super-secret-token",
+ }
+ }
+ table.find_first = AsyncMock(return_value=row)
+ monkeypatch.setattr(ps, "prisma_client", mock_prisma)
+
+ with auth_as(LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY):
+ response = client.get(
+ "/config/field/info", params={"field_name": "database_args"}
+ )
+ assert response.status_code == 200
+ value = response.json()["field_value"]
+ assert value["aws_web_identity_token"] == "REDACTED"
+ assert value["region_name"] == "us-east-1"
+ assert value["user_table_name"] == "LiteLLM_UserTable"
+
+
+def test_config_field_info_full_admin_sees_nested_secret(
+ client, auth_as, mock_prisma, monkeypatch
+):
+ """The redaction must not over-redact for a full PROXY_ADMIN, who needs
+ the real nested value to populate the edit form."""
+ from litellm.proxy import proxy_server as ps
+ from litellm.proxy._types import LitellmUserRoles
+
+ table = _install_litellm_config(mock_prisma)
+ row = MagicMock()
+ row.param_value = {
+ "database_args": {
+ "region_name": "us-east-1",
+ "aws_web_identity_token": "sk-super-secret-token",
+ }
+ }
+ table.find_first = AsyncMock(return_value=row)
+ monkeypatch.setattr(ps, "prisma_client", mock_prisma)
+
+ with auth_as(LitellmUserRoles.PROXY_ADMIN):
+ response = client.get(
+ "/config/field/info", params={"field_name": "database_args"}
+ )
+ assert response.status_code == 200
+ value = response.json()["field_value"]
+ assert value["aws_web_identity_token"] == "sk-super-secret-token"
+ assert value["region_name"] == "us-east-1"
+
+
+def test_config_field_info_redacts_top_level_scalar_for_view_only(
+ client, auth_as, mock_prisma, monkeypatch
+):
+ """The top-level scalar branch must also redact for a view-only admin.
+ database_url carries DB credentials and is not caught by the name masker,
+ so it is in the explicit secret set."""
+ from litellm.proxy import proxy_server as ps
+ from litellm.proxy._types import LitellmUserRoles
+
+ table = _install_litellm_config(mock_prisma)
+ row = MagicMock()
+ row.param_value = {"database_url": "postgresql://admin:p4ss@db:5432/litellm"}
+ table.find_first = AsyncMock(return_value=row)
+ monkeypatch.setattr(ps, "prisma_client", mock_prisma)
+
+ with auth_as(LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY):
+ response = client.get(
+ "/config/field/info", params={"field_name": "database_url"}
+ )
+ assert response.status_code == 200
+ assert response.json()["field_value"] == "REDACTED"
+
+
+def test_redact_general_setting_value_recurses_list_of_dicts():
+ """The list branch of the recursor redacts secret leaves inside each dict
+ while non-secret keys survive, and a full admin gets the value untouched."""
+ from litellm.proxy import proxy_server as ps
+
+ value = [
+ {"path": "/foo", "headers": {"Authorization": "Bearer sk-x"}},
+ {"path": "/bar", "client_secret": "sk-y"},
+ ]
+ redacted = ps._redact_general_setting_value(
+ "some_list_field", value, is_full_admin=False
+ )
+ assert redacted[0]["headers"]["Authorization"] == "REDACTED"
+ assert redacted[0]["path"] == "/foo"
+ assert redacted[1]["client_secret"] == "REDACTED"
+ assert redacted[1]["path"] == "/bar"
+ assert (
+ ps._redact_general_setting_value("some_list_field", value, is_full_admin=True)
+ == value
+ )
+
+
+def test_redact_secret_values_in_obj_fails_closed_at_max_depth():
+ """Past _REDACT_SECRET_MAX_DEPTH the whole subtree is replaced with
+ "REDACTED" rather than returned verbatim, so a secret buried below the cap
+ can never leak via depth-overrun. A future refactor that flips the cap
+ branch to fail-open would surface here."""
+ from litellm.proxy import proxy_server as ps
+
+ # leaf and wrap keys are both NON-secret so neither the key-name
+ # short-circuit nor the explicit-secret set catches the leak. The cap is
+ # the only thing standing between the secret and the response — flip the
+ # cap to fail-open and the secret comes back verbatim.
+ nested: object = {"notes": "sk-leak-bottom"}
+ for _ in range(ps._REDACT_SECRET_MAX_DEPTH + 2):
+ nested = {"wrap": nested}
+
+ out = ps._redact_general_setting_value(
+ "some_struct_field", nested, is_full_admin=False
+ )
+ # the secret must not survive anywhere in the returned tree
+ assert "sk-leak-bottom" not in repr(out)
+
+ # full admin is unaffected by the cap — the value comes back untouched
+ admin_out = ps._redact_general_setting_value(
+ "some_struct_field", nested, is_full_admin=True
+ )
+ assert admin_out is nested
+
+
+def test_config_list_redacts_pass_through_secret_for_view_only(
+ client, auth_as, mock_prisma, monkeypatch
+):
+ """/config/list must not leak pass_through_endpoints upstream credentials
+ to a view-only admin. pass_through_endpoints is a known secret-bearing
+ field, so a non-admin gets it redacted; a full admin still sees it."""
+ from litellm.proxy import proxy_server as ps
+ from litellm.proxy._types import LitellmUserRoles
+
+ table = _install_litellm_config(mock_prisma)
+ row = MagicMock()
+ row.param_value = {"max_parallel_requests": 3}
+ table.find_first = AsyncMock(return_value=row)
+ monkeypatch.setattr(ps, "prisma_client", mock_prisma)
+ monkeypatch.setattr(
+ ps,
+ "general_settings",
+ {
+ "pass_through_endpoints": [
+ {
+ "path": "/foo",
+ "target": "https://upstream.example.com",
+ "headers": {"Authorization": "Bearer sk-UPSTREAM-SECRET"},
+ }
+ ]
+ },
+ )
+
+ def _pass_through_value(body):
+ return next(
+ entry["field_value"]
+ for entry in body
+ if entry["field_name"] == "pass_through_endpoints"
+ )
+
+ with auth_as(LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY):
+ view_resp = client.get(
+ "/config/list", params={"config_type": "general_settings"}
+ )
+ assert view_resp.status_code == 200
+ assert "sk-UPSTREAM-SECRET" not in view_resp.text
+ assert _pass_through_value(view_resp.json()) == "REDACTED"
+
+ with auth_as(LitellmUserRoles.PROXY_ADMIN):
+ admin_resp = client.get(
+ "/config/list", params={"config_type": "general_settings"}
+ )
+ assert admin_resp.status_code == 200
+ admin_value = _pass_through_value(admin_resp.json())
+ assert admin_value[0]["headers"]["Authorization"] == "Bearer sk-UPSTREAM-SECRET"
+
+
# ---------------------------------------------------------------------------
# GET /config/list
# ---------------------------------------------------------------------------
diff --git a/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py b/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py
index 33de1ede917..699606b5277 100644
--- a/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py
+++ b/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py
@@ -16,8 +16,6 @@ Pins covered:
from __future__ import annotations
import json
-from typing import Any, AsyncIterator
-from unittest.mock import AsyncMock, MagicMock
import pytest
@@ -26,8 +24,11 @@ from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.proxy_server import (
_apply_streaming_chunk_hooks,
_fast_serialize_simple_model_response_stream,
+ _format_fallback_metadata_sse_event,
_format_streaming_sse_chunk,
_get_client_requested_model_for_streaming,
+ _get_streaming_fallback_metadata,
+ _is_positive_int_like,
_restamp_streaming_chunk_model,
_serialize_streaming_chunk,
async_assistants_data_generator,
@@ -71,6 +72,15 @@ async def _async_iter_raises(exc: Exception):
raise exc
+class _FakeStream:
+ def __init__(self, chunks, hidden_params=None):
+ self._chunks = chunks
+ self._hidden_params = hidden_params or {}
+
+ def __aiter__(self):
+ return _async_iter(self._chunks)
+
+
# ---------------------------------------------------------------------------
# data_generator
# ---------------------------------------------------------------------------
@@ -274,6 +284,34 @@ def test_restamp_streaming_chunk_model_overrides_model_on_dict():
assert logged is True
+def test_restamp_streaming_chunk_model_uses_fallback_model_from_metadata():
+ chunk = _simple_chunk(model="openai/internal-fallback")
+ new_chunk, logged = _restamp_streaming_chunk_model(
+ chunk=chunk,
+ requested_model_from_client="primary-model",
+ request_data={"litellm_call_id": "id-1"},
+ model_mismatch_logged=False,
+ fallback_was_attempted=True,
+ fallback_model_from_metadata="fallback-model",
+ )
+ assert new_chunk.model == "fallback-model"
+ assert logged is True
+
+
+def test_restamp_streaming_chunk_model_preserves_fallback_model_without_group():
+ chunk = _simple_chunk(model="openai/internal-fallback")
+ new_chunk, logged = _restamp_streaming_chunk_model(
+ chunk=chunk,
+ requested_model_from_client="primary-model",
+ request_data={},
+ model_mismatch_logged=False,
+ fallback_was_attempted=True,
+ fallback_model_from_metadata=None,
+ )
+ assert new_chunk.model == "openai/internal-fallback"
+ assert logged is False
+
+
def test_restamp_streaming_chunk_model_invalid_chunk_type_unchanged():
"""For a non-BaseModel, non-dict chunk the helper returns it as-is
along with the original ``model_mismatch_logged`` flag."""
@@ -288,6 +326,147 @@ def test_restamp_streaming_chunk_model_invalid_chunk_type_unchanged():
assert logged is False
+def test_is_positive_int_like_invalid_and_edge_values():
+ assert _is_positive_int_like(None) is False
+ assert _is_positive_int_like("not-a-number") is False
+ assert _is_positive_int_like(0) is False
+ assert _is_positive_int_like(-1) is False
+ assert _is_positive_int_like("1") is True
+ assert _is_positive_int_like(2) is True
+
+
+def test_get_streaming_fallback_metadata_reads_headers():
+ fallback_errors = [
+ {
+ "message": "litellm.RateLimitError: upstream limited request",
+ "type": "RateLimitError",
+ "param": None,
+ "code": "429",
+ }
+ ]
+ stream = _FakeStream(
+ [],
+ hidden_params={
+ "additional_headers": {
+ "x-litellm-attempted-fallbacks": "1",
+ "x-litellm-model-group": "fallback-model",
+ "x-litellm-fallback-errors": json.dumps(fallback_errors),
+ }
+ },
+ )
+ assert _get_streaming_fallback_metadata(stream) == (
+ True,
+ "fallback-model",
+ fallback_errors,
+ )
+
+
+def test_get_streaming_fallback_metadata_no_additional_headers():
+ stream = _FakeStream([], hidden_params={})
+ assert _get_streaming_fallback_metadata(stream) == (False, None, [])
+
+
+def test_get_streaming_fallback_metadata_zero_fallback_count():
+ stream = _FakeStream(
+ [],
+ hidden_params={
+ "additional_headers": {"x-litellm-attempted-fallbacks": 0}
+ },
+ )
+ assert _get_streaming_fallback_metadata(stream) == (False, None, [])
+
+
+def test_get_streaming_fallback_metadata_no_model_group_returns_none_model():
+ stream = _FakeStream(
+ [],
+ hidden_params={
+ "additional_headers": {
+ "x-litellm-attempted-fallbacks": 1,
+ }
+ },
+ )
+ was_attempted, fallback_model, errors = _get_streaming_fallback_metadata(stream)
+ assert was_attempted is True
+ assert fallback_model is None
+ assert errors == []
+
+
+def test_restamp_streaming_chunk_model_azure_router_preserves_model():
+ chunk = _simple_chunk(model="azure_ai/internal-deployment")
+ new_chunk, logged = _restamp_streaming_chunk_model(
+ chunk=chunk,
+ requested_model_from_client="azure_ai/model-router",
+ request_data={},
+ model_mismatch_logged=False,
+ )
+ assert new_chunk.model == "azure_ai/internal-deployment"
+ assert logged is False
+
+
+def test_restamp_streaming_chunk_model_fastest_response_preserves_model():
+ chunk = _simple_chunk(model="winning-model")
+ new_chunk, logged = _restamp_streaming_chunk_model(
+ chunk=chunk,
+ requested_model_from_client="gpt-4,claude-3",
+ request_data={"fastest_response": True},
+ model_mismatch_logged=False,
+ )
+ assert new_chunk.model == "winning-model"
+ assert logged is False
+
+
+def test_restamp_streaming_chunk_model_setattr_exception_logs_and_returns():
+ from pydantic import ConfigDict
+
+ class FrozenChunk(_simple_chunk().__class__):
+ model_config = ConfigDict(frozen=True)
+
+ chunk = FrozenChunk(
+ id="chatcmpl-test",
+ choices=[],
+ created=0,
+ model="openai/internal-x",
+ object="chat.completion.chunk",
+ )
+ new_chunk, logged = _restamp_streaming_chunk_model(
+ chunk=chunk,
+ requested_model_from_client="gpt-4",
+ request_data={"litellm_call_id": "test-id"},
+ model_mismatch_logged=False,
+ )
+ assert new_chunk.model == "openai/internal-x"
+ assert logged is True
+
+
+def test_format_fallback_metadata_sse_event():
+ fallback_errors = [
+ {
+ "message": "litellm.RateLimitError: upstream limited request",
+ "type": "RateLimitError",
+ "param": None,
+ "code": "429",
+ }
+ ]
+
+ event = _format_fallback_metadata_sse_event(
+ fallback_model="fallback-model",
+ fallback_errors=fallback_errors,
+ )
+
+ assert isinstance(event, str)
+ assert event.startswith("data: ")
+ payload = json.loads(event.removeprefix("data: ").removesuffix("\n\n"))
+ assert payload["choices"] == []
+ assert payload["litellm_fallback"] == {
+ "fallback_model": "fallback-model",
+ "errors": fallback_errors,
+ }
+ assert payload["id"] == "litellm-fallback-metadata"
+ assert payload["object"] == "chat.completion.chunk"
+ assert payload["model"] == "fallback-model"
+ assert isinstance(payload["created"], int)
+
+
# ---------------------------------------------------------------------------
# _fast_serialize_simple_model_response_stream
# ---------------------------------------------------------------------------
@@ -473,7 +652,7 @@ async def test_async_data_generator_yields_sse_chunks_and_done(monkeypatch):
# First chunk is bytes (fast path) wrapped via _format_streaming_sse_chunk.
first = out[0]
assert isinstance(first, bytes)
- payload = json.loads(first.removeprefix(b"data: ").rstrip(b"\n\n"))
+ payload = json.loads(first.removeprefix(b"data: ").removesuffix(b"\n\n"))
assert normalize(payload) == {
"id": "",
"object": "chat.completion.chunk",
@@ -488,6 +667,172 @@ async def test_async_data_generator_yields_sse_chunks_and_done(monkeypatch):
}
+@pytest.mark.asyncio
+async def test_async_data_generator_uses_response_fallback_metadata(monkeypatch):
+ _patch_logging_flags(monkeypatch)
+
+ response = _FakeStream(
+ [_simple_chunk(model="openai/internal-fallback", content="hello")],
+ hidden_params={
+ "additional_headers": {
+ "x-litellm-attempted-fallbacks": 1,
+ "x-litellm-model-group": "fallback-model",
+ }
+ },
+ )
+ out = []
+ async for line in async_data_generator(
+ response=response,
+ user_api_key_dict=_user_auth(),
+ request_data={"model": "primary-model", "include_fallback_errors": True},
+ ):
+ out.append(line)
+
+ first = out[0]
+ assert isinstance(first, bytes)
+ payload = json.loads(first.removeprefix(b"data: ").removesuffix(b"\n\n"))
+ assert payload["model"] == "fallback-model"
+
+
+@pytest.mark.asyncio
+async def test_async_data_generator_uses_chunk_fallback_metadata(monkeypatch):
+ _patch_logging_flags(monkeypatch)
+
+ chunk = _simple_chunk(model="openai/internal-fallback", content="hello")
+ chunk._hidden_params = {
+ "additional_headers": {
+ "x-litellm-attempted-fallbacks": 1,
+ "x-litellm-model-group": "fallback-model",
+ }
+ }
+ out = []
+ async for line in async_data_generator(
+ response=_async_iter([chunk]),
+ user_api_key_dict=_user_auth(),
+ request_data={"model": "primary-model"},
+ ):
+ out.append(line)
+
+ first = out[0]
+ assert isinstance(first, bytes)
+ payload = json.loads(first.removeprefix(b"data: ").removesuffix(b"\n\n"))
+ assert payload["model"] == "fallback-model"
+
+
+@pytest.mark.asyncio
+async def test_async_data_generator_switches_model_mid_stream_on_fallback(monkeypatch):
+ """Pre-fallback chunks keep the client-requested model; once a chunk carries
+ fallback metadata the model latches to the fallback group for the rest of the
+ stream. This pins the client-visible mid-stream model change."""
+ _patch_logging_flags(monkeypatch)
+
+ primary_chunk = _simple_chunk(model="openai/internal-primary", content="hi")
+ fallback_chunk = _simple_chunk(model="openai/internal-fallback", content="there")
+ fallback_chunk._hidden_params = {
+ "additional_headers": {
+ "x-litellm-attempted-fallbacks": 1,
+ "x-litellm-model-group": "fallback-model",
+ }
+ }
+ out = []
+ async for line in async_data_generator(
+ response=_async_iter([primary_chunk, fallback_chunk]),
+ user_api_key_dict=_user_auth(),
+ request_data={"model": "primary-model"},
+ ):
+ out.append(line)
+
+ first_payload = json.loads(out[0].removeprefix(b"data: ").removesuffix(b"\n\n"))
+ second_payload = json.loads(out[1].removeprefix(b"data: ").removesuffix(b"\n\n"))
+ assert first_payload["model"] == "primary-model"
+ assert second_payload["model"] == "fallback-model"
+
+
+@pytest.mark.asyncio
+async def test_async_data_generator_emits_fallback_error_metadata_event(monkeypatch):
+ _patch_logging_flags(monkeypatch)
+ monkeypatch.setitem(ps.general_settings, "expose_fallback_errors_to_caller", True)
+
+ fallback_errors = [
+ {
+ "message": "litellm.RateLimitError: upstream limited request",
+ "type": "RateLimitError",
+ "param": None,
+ "code": "429",
+ }
+ ]
+ response = _FakeStream(
+ [_simple_chunk(model="openai/internal-fallback", content="hello")],
+ hidden_params={
+ "additional_headers": {
+ "x-litellm-attempted-fallbacks": 1,
+ "x-litellm-model-group": "fallback-model",
+ "x-litellm-fallback-errors": json.dumps(fallback_errors),
+ }
+ },
+ )
+ out = []
+ async for line in async_data_generator(
+ response=response,
+ user_api_key_dict=_user_auth(),
+ request_data={"model": "primary-model", "include_fallback_errors": True},
+ ):
+ out.append(line)
+
+ assert isinstance(out[0], bytes)
+ chunk_payload = json.loads(out[0].removeprefix(b"data: ").removesuffix(b"\n\n"))
+ assert chunk_payload["model"] == "fallback-model"
+ assert isinstance(out[1], str)
+ assert out[1].startswith("data: ")
+ metadata_payload = json.loads(out[1].removeprefix("data: ").removesuffix("\n\n"))
+ assert metadata_payload["choices"] == []
+ assert metadata_payload["litellm_fallback"] == {
+ "fallback_model": "fallback-model",
+ "errors": fallback_errors,
+ }
+ assert metadata_payload["id"] == "litellm-fallback-metadata"
+ assert metadata_payload["object"] == "chat.completion.chunk"
+ assert metadata_payload["model"] == "fallback-model"
+ assert isinstance(metadata_payload["created"], int)
+
+
+@pytest.mark.asyncio
+async def test_async_data_generator_skips_fallback_error_event_without_opt_in(
+ monkeypatch,
+):
+ _patch_logging_flags(monkeypatch)
+
+ fallback_errors = [
+ {
+ "message": "litellm.RateLimitError: upstream limited request",
+ "type": "RateLimitError",
+ "param": None,
+ "code": "429",
+ }
+ ]
+ response = _FakeStream(
+ [_simple_chunk(model="openai/internal-fallback", content="hello")],
+ hidden_params={
+ "additional_headers": {
+ "x-litellm-attempted-fallbacks": 1,
+ "x-litellm-model-group": "fallback-model",
+ "x-litellm-fallback-errors": json.dumps(fallback_errors),
+ }
+ },
+ )
+ out = []
+ async for line in async_data_generator(
+ response=response,
+ user_api_key_dict=_user_auth(),
+ request_data={"model": "primary-model"},
+ ):
+ out.append(line)
+
+ assert isinstance(out[0], bytes)
+ payload = json.loads(out[0].removeprefix(b"data: ").removesuffix(b"\n\n"))
+ assert payload["model"] == "fallback-model"
+
+
@pytest.mark.asyncio
async def test_async_data_generator_mid_stream_exception_yields_error_payload(
monkeypatch,
diff --git a/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py b/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py
index 34405f20727..25e84fb59a7 100644
--- a/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py
+++ b/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py
@@ -14,7 +14,11 @@ from unittest.mock import AsyncMock, MagicMock
import pytest
import litellm.proxy.proxy_server as ps
-from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
+from litellm.proxy._types import (
+ LiteLLM_UserTable,
+ LitellmUserRoles,
+ UserAPIKeyAuth,
+)
from litellm.proxy.common_utils.model_listing_utils import TeamModelNameTranslator
from litellm.proxy.proxy_server import (
_get_proxy_model_info,
@@ -1360,3 +1364,83 @@ async def test_retrieve_model_by_inaccessible_public_name_404s(monkeypatch):
assert exc_info.value.status_code == 404
router.get_deployment_by_model_group_name.assert_not_called()
+
+
+def test_get_direct_access_models_expands_all_proxy_models_sentinel():
+ """A user provisioned with 'all-proxy-models' has direct access to every non-team
+ deployment. The sentinel must resolve via get_model_ids, not be looked up as a
+ literal model_name (which matches nothing). Regression for GH#22791."""
+ router = MagicMock()
+ router.get_model_ids.return_value = ["global-id-1", "global-id-2"]
+ router.get_model_list.return_value = []
+
+ user = LiteLLM_UserTable(
+ user_id="u",
+ models=[ps.SpecialModelNames.all_proxy_models.value],
+ teams=[],
+ )
+
+ result = ps.get_direct_access_models(user_db_object=user, llm_router=router)
+
+ assert result == ["global-id-1", "global-id-2"]
+ router.get_model_ids.assert_called_once_with(exclude_team_models=True)
+ router.get_model_list.assert_not_called()
+
+
+def test_get_direct_access_models_resolves_explicit_model_names():
+ """Without the sentinel, only the user's explicitly listed models resolve to ids;
+ the all-proxy-models shortcut must not fire."""
+ router = MagicMock()
+ router.get_model_list.return_value = [{"model_info": {"id": "gpt4o-id"}}]
+
+ user = LiteLLM_UserTable(user_id="u", models=["gpt-4o"], teams=[])
+
+ result = ps.get_direct_access_models(user_db_object=user, llm_router=router)
+
+ assert result == ["gpt4o-id"]
+ router.get_model_ids.assert_not_called()
+ router.get_model_list.assert_called_once_with(model_name="gpt-4o")
+
+
+@pytest.mark.asyncio
+async def test_populate_team_access_grants_all_proxy_models_user_direct_access(
+ monkeypatch,
+):
+ """An internal user provisioned with 'all-proxy-models' and no teams must see proxy
+ models on the Models+Endpoints page. Before the fix _populate_team_access_on_models
+ marked direct_access=False, so _filter_models_to_user_accessible dropped every model
+ and the page was empty. Regression for GH#22791."""
+ global_row = {
+ "model_name": "gpt-4o",
+ "litellm_params": {"model": "gpt-4o"},
+ "model_info": {"id": "global-id-1", "db_model": False},
+ }
+
+ router = MagicMock()
+ router.get_model_ids.return_value = ["global-id-1"]
+
+ user_row = LiteLLM_UserTable(
+ user_id="u",
+ user_role=LitellmUserRoles.INTERNAL_USER.value,
+ models=[ps.SpecialModelNames.all_proxy_models.value],
+ teams=[],
+ )
+ prisma_client = MagicMock()
+ prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row)
+
+ monkeypatch.setattr(ps, "get_all_team_models", AsyncMock(return_value={}))
+
+ caller = UserAPIKeyAuth(
+ user_id="u", user_role=LitellmUserRoles.INTERNAL_USER, team_models=[]
+ )
+
+ populated = await ps._populate_team_access_on_models(
+ user_api_key_dict=caller,
+ prisma_client=prisma_client,
+ llm_router=router,
+ all_models=[global_row],
+ )
+ visible = ps._filter_models_to_user_accessible(populated)
+
+ assert [m["model_info"]["id"] for m in visible] == ["global-id-1"]
+ assert visible[0]["model_info"]["direct_access"] is True
diff --git a/tests/test_litellm/proxy/test_proxy_cli.py b/tests/test_litellm/proxy/test_proxy_cli.py
index 56627c5be88..88dbec4020f 100644
--- a/tests/test_litellm/proxy/test_proxy_cli.py
+++ b/tests/test_litellm/proxy/test_proxy_cli.py
@@ -1708,6 +1708,57 @@ class TestRunServerDbSetup:
use_migrate=True, use_v2_resolver=False
)
+ @patch("subprocess.run")
+ @patch("atexit.register")
+ @patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database")
+ @patch("litellm.proxy.db.check_migration.check_prisma_schema_diff")
+ @patch("litellm.proxy.db.prisma_client.should_update_prisma_schema")
+ def test_startup_exits_on_non_postgres_database_url(
+ self,
+ mock_should_update_schema,
+ mock_check_schema_diff,
+ mock_setup_database,
+ mock_atexit_register,
+ mock_subprocess_run,
+ ):
+ """A sqlite DATABASE_URL must exit immediately, before any prisma call,
+ instead of stalling on a migration against the postgresql-only schema."""
+ from litellm.proxy.proxy_cli import run_server
+
+ mock_subprocess_run.return_value = MagicMock(returncode=0)
+ mock_should_update_schema.return_value = True
+
+ mock_proxy_module = MagicMock(
+ app=MagicMock(),
+ ProxyConfig=MagicMock(),
+ KeyManagementSettings=MagicMock(),
+ save_worker_config=MagicMock(),
+ )
+
+ clean_env = {
+ k: v
+ for k, v in os.environ.items()
+ if k not in ("DATABASE_URL", "DIRECT_URL")
+ }
+ clean_env["DATABASE_URL"] = "sqlite:///data/litellm.db"
+
+ with (
+ patch.dict(os.environ, clean_env, clear=True),
+ patch.dict(
+ "sys.modules",
+ {
+ "proxy_server": mock_proxy_module,
+ "litellm.proxy.proxy_server": mock_proxy_module,
+ },
+ ),
+ ):
+ with pytest.raises(SystemExit) as exc_info:
+ run_server.main(
+ ["--local", "--skip_server_startup"], standalone_mode=False
+ )
+ assert exc_info.value.code == 1
+ mock_setup_database.assert_not_called()
+
# --- Module-level helpers for worker startup hook tests ---
diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py
index c85c5ccb39f..a203fcc7ec0 100644
--- a/tests/test_litellm/proxy/test_proxy_server.py
+++ b/tests/test_litellm/proxy/test_proxy_server.py
@@ -8364,3 +8364,89 @@ def test_preserve_redacted_plugin_keys_sets_new_and_drops_orphan_placeholder():
[{"name": "p2", "url": "https://p2", "plugin_key": "***"}], existing
)
assert "plugin_key" not in new_plugin[0]
+
+
+def _config_field_info_client(monkeypatch, user_role):
+ import types
+ from unittest.mock import AsyncMock, MagicMock
+
+ from fastapi.testclient import TestClient
+
+ import litellm.proxy.proxy_server as ps
+ from litellm.proxy._types import UserAPIKeyAuth
+ from litellm.proxy.proxy_server import app
+
+ db_record = types.SimpleNamespace(
+ param_value={
+ "master_key": "sk-super-secret-master",
+ "database_url": "postgresql://user:p4ssw0rd@db:5432/litellm",
+ "pass_through_endpoints": [
+ {
+ "path": "/upstream",
+ "target": "https://upstream.example.com",
+ "headers": {"Authorization": "Bearer sk-upstream-secret"},
+ }
+ ],
+ "max_parallel_requests": 100,
+ }
+ )
+ mock_config_table = MagicMock()
+ mock_config_table.find_first = AsyncMock(return_value=db_record)
+ mock_prisma = MagicMock()
+ mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table)
+ monkeypatch.setattr(ps, "prisma_client", mock_prisma)
+ app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
+ user_id="u", user_role=user_role
+ )
+ return TestClient(app)
+
+
+def test_config_field_info_redacts_secrets_for_view_only_admin(monkeypatch):
+ """/config/field/info gates on _user_has_admin_view, which also grants
+ PROXY_ADMIN_VIEW_ONLY. A view-only admin reading master_key/database_url verbatim is
+ effectively a full admin. Secret-bearing fields must come back REDACTED for anyone who
+ is not a FULL PROXY_ADMIN, while non-secret fields stay readable."""
+ from litellm.proxy._types import LitellmUserRoles
+
+ client = _config_field_info_client(
+ monkeypatch, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY
+ )
+ try:
+ for secret_field in ("master_key", "database_url", "pass_through_endpoints"):
+ resp = client.get("/config/field/info", params={"field_name": secret_field})
+ assert resp.status_code == 200, resp.text
+ body = resp.json()
+ assert body["field_value"] == "REDACTED"
+ assert "secret" not in str(body["field_value"])
+ assert "p4ssw0rd" not in str(body["field_value"])
+
+ resp = client.get(
+ "/config/field/info", params={"field_name": "max_parallel_requests"}
+ )
+ assert resp.status_code == 200, resp.text
+ assert resp.json()["field_value"] == 100
+ finally:
+ app.dependency_overrides.clear()
+
+
+def test_config_field_info_returns_raw_secrets_for_full_admin(monkeypatch):
+ """the redaction must not over-apply. A FULL PROXY_ADMIN still
+ needs the real master_key value to populate the admin edit form."""
+ from litellm.proxy._types import LitellmUserRoles
+
+ client = _config_field_info_client(monkeypatch, LitellmUserRoles.PROXY_ADMIN)
+ try:
+ resp = client.get("/config/field/info", params={"field_name": "master_key"})
+ assert resp.status_code == 200, resp.text
+ assert resp.json()["field_value"] == "sk-super-secret-master"
+
+ resp = client.get(
+ "/config/field/info", params={"field_name": "pass_through_endpoints"}
+ )
+ assert resp.status_code == 200, resp.text
+ assert (
+ resp.json()["field_value"][0]["headers"]["Authorization"]
+ == "Bearer sk-upstream-secret"
+ )
+ finally:
+ app.dependency_overrides.clear()
diff --git a/tests/test_litellm/proxy/test_read_model_list.py b/tests/test_litellm/proxy/test_read_model_list.py
new file mode 100644
index 00000000000..7703d57e21d
--- /dev/null
+++ b/tests/test_litellm/proxy/test_read_model_list.py
@@ -0,0 +1,32 @@
+"""Tests for litellm.proxy.read_model_list (Rust AI gateway config bridge)."""
+
+from litellm.proxy.read_model_list import read_model_list
+
+
+def test_read_model_list_resolves_os_environ(monkeypatch, tmp_path):
+ """`os.environ/` markers in the model_list are resolved via ProxyConfig."""
+ monkeypatch.setenv("OPENAI_API_KEY", "sk-resolved-123")
+ config = tmp_path / "config.yaml"
+ config.write_text(
+ "model_list:\n"
+ " - model_name: gpt-realtime\n"
+ " litellm_params:\n"
+ " model: openai/gpt-realtime\n"
+ " api_key: os.environ/OPENAI_API_KEY\n"
+ )
+
+ model_list = read_model_list(str(config))
+
+ assert len(model_list) == 1
+ params = model_list[0]["litellm_params"]
+ assert model_list[0]["model_name"] == "gpt-realtime"
+ assert params["model"] == "openai/gpt-realtime"
+ assert params["api_key"] == "sk-resolved-123"
+
+
+def test_read_model_list_missing_key_returns_empty(tmp_path):
+ """A config without a model_list yields an empty list, not an error."""
+ config = tmp_path / "config.yaml"
+ config.write_text("general_settings: {}\n")
+
+ assert read_model_list(str(config)) == []
diff --git a/tests/test_litellm/router_utils/test_add_retry_fallback_headers.py b/tests/test_litellm/router_utils/test_add_retry_fallback_headers.py
new file mode 100644
index 00000000000..2aee0f0a4ef
--- /dev/null
+++ b/tests/test_litellm/router_utils/test_add_retry_fallback_headers.py
@@ -0,0 +1,142 @@
+import json
+
+from pydantic import BaseModel
+
+from litellm.router_utils.add_retry_fallback_headers import (
+ add_fallback_headers_to_response,
+ add_retry_headers_to_response,
+ get_fallback_errors_from_headers,
+ get_hidden_params_dict,
+)
+
+
+class StreamingWrapper:
+ def __init__(self):
+ self._hidden_params = {"additional_headers": {"x-existing": "keep"}}
+
+
+def test_add_fallback_headers_to_streaming_wrapper():
+ response = StreamingWrapper()
+
+ result = add_fallback_headers_to_response(
+ response=response,
+ attempted_fallbacks=1,
+ )
+
+ assert result is response
+ assert response._hidden_params["additional_headers"] == {
+ "x-existing": "keep",
+ "x-litellm-attempted-fallbacks": 1,
+ }
+
+
+def test_add_fallback_headers_serializes_fallback_errors():
+ response = StreamingWrapper()
+ fallback_errors = [
+ {
+ "message": "litellm.RateLimitError: upstream limited request",
+ "type": "RateLimitError",
+ "param": None,
+ "code": "429",
+ }
+ ]
+
+ result = add_fallback_headers_to_response(
+ response=response,
+ attempted_fallbacks=1,
+ fallback_errors=fallback_errors,
+ )
+
+ assert result is response
+ assert response._hidden_params["additional_headers"][
+ "x-litellm-attempted-fallbacks"
+ ] == 1
+ assert (
+ json.loads(
+ response._hidden_params["additional_headers"]["x-litellm-fallback-errors"]
+ )
+ == fallback_errors
+ )
+
+
+def test_add_retry_headers_to_streaming_wrapper():
+ response = StreamingWrapper()
+
+ result = add_retry_headers_to_response(
+ response=response,
+ attempted_retries=2,
+ max_retries=3,
+ )
+
+ assert result is response
+ assert response._hidden_params["additional_headers"] == {
+ "x-existing": "keep",
+ "x-litellm-attempted-retries": 2,
+ "x-litellm-max-retries": 3,
+ }
+
+
+def test_get_hidden_params_dict_with_pydantic_model_hidden_params():
+ class InnerHiddenParams(BaseModel):
+ additional_headers: dict = {}
+
+ class Response:
+ def __init__(self):
+ self._hidden_params = InnerHiddenParams(
+ additional_headers={"x-custom": "value"}
+ )
+
+ result = get_hidden_params_dict(Response())
+ assert result == {"additional_headers": {"x-custom": "value"}}
+
+
+def test_get_hidden_params_dict_with_no_hidden_params():
+ class PlainResponse:
+ pass
+
+ assert get_hidden_params_dict(PlainResponse()) == {}
+
+
+def test_add_fallback_headers_when_no_existing_additional_headers():
+ class NoHeadersWrapper:
+ def __init__(self):
+ self._hidden_params = {}
+
+ response = NoHeadersWrapper()
+ result = add_fallback_headers_to_response(response=response, attempted_fallbacks=2)
+
+ assert result is response
+ assert response._hidden_params["additional_headers"]["x-litellm-attempted-fallbacks"] == 2
+
+
+def test_add_fallback_headers_returns_none_when_response_is_none():
+ result = add_fallback_headers_to_response(response=None, attempted_fallbacks=1)
+ assert result is None
+
+
+def test_add_fallback_headers_returns_unchanged_when_response_has_no_hidden_params():
+ class PlainObject:
+ pass
+
+ obj = PlainObject()
+ result = add_fallback_headers_to_response(response=obj, attempted_fallbacks=1)
+ assert result is obj
+ assert not hasattr(obj, "_hidden_params")
+
+
+def test_get_fallback_errors_from_headers_existing_list_passthrough():
+ errors = [{"message": "err", "type": "T", "param": None, "code": "400"}]
+ result = get_fallback_errors_from_headers({"x-litellm-fallback-errors": errors})
+ assert result == errors
+
+
+def test_get_fallback_errors_from_headers_invalid_json_returns_empty():
+ result = get_fallback_errors_from_headers(
+ {"x-litellm-fallback-errors": "not-valid-json-{"}
+ )
+ assert result == []
+
+
+def test_get_fallback_errors_from_headers_missing_key_returns_empty():
+ result = get_fallback_errors_from_headers({})
+ assert result == []
diff --git a/tests/test_litellm/router_utils/test_fallback_event_handlers.py b/tests/test_litellm/router_utils/test_fallback_event_handlers.py
new file mode 100644
index 00000000000..ca647bdce55
--- /dev/null
+++ b/tests/test_litellm/router_utils/test_fallback_event_handlers.py
@@ -0,0 +1,139 @@
+import json
+
+import pytest
+
+from litellm.router_utils.fallback_event_handlers import run_async_fallback
+
+
+class StreamingWrapper:
+ def __init__(self):
+ self._hidden_params = {"additional_headers": {}}
+
+
+class FakeRouter:
+ def log_retry(self, kwargs, e):
+ return kwargs
+
+ async def async_function_with_fallbacks(self, *args, **kwargs):
+ return StreamingWrapper()
+
+
+class AlwaysFailRouter:
+ def log_retry(self, kwargs, e):
+ return kwargs
+
+ async def async_function_with_fallbacks(self, *args, **kwargs):
+ raise RuntimeError("fallback model also failed")
+
+
+@pytest.mark.asyncio
+async def test_run_async_fallback_adds_errors_when_opted_in():
+ response = await run_async_fallback(
+ litellm_router=FakeRouter(),
+ fallback_model_group=["fallback-model"],
+ original_model_group="primary-model",
+ original_exception=RuntimeError("upstream limited request"),
+ max_fallbacks=3,
+ fallback_depth=0,
+ include_fallback_errors=True,
+ )
+
+ additional_headers = response._hidden_params["additional_headers"]
+ assert additional_headers["x-litellm-attempted-fallbacks"] == 1
+ assert json.loads(additional_headers["x-litellm-fallback-errors"]) == [
+ {
+ "message": "upstream limited request",
+ "type": "RuntimeError",
+ "param": None,
+ "code": None,
+ }
+ ]
+
+
+@pytest.mark.asyncio
+async def test_run_async_fallback_omits_errors_without_opt_in():
+ response = await run_async_fallback(
+ litellm_router=FakeRouter(),
+ fallback_model_group=["fallback-model"],
+ original_model_group="primary-model",
+ original_exception=RuntimeError("upstream limited request"),
+ max_fallbacks=3,
+ fallback_depth=0,
+ )
+
+ additional_headers = response._hidden_params["additional_headers"]
+ assert additional_headers["x-litellm-attempted-fallbacks"] == 1
+ assert "x-litellm-fallback-errors" not in additional_headers
+
+
+@pytest.mark.asyncio
+async def test_run_async_fallback_raises_when_all_fallbacks_fail():
+ with pytest.raises(RuntimeError, match="fallback model also failed"):
+ await run_async_fallback(
+ litellm_router=AlwaysFailRouter(),
+ fallback_model_group=["fallback-model"],
+ original_model_group="primary-model",
+ original_exception=RuntimeError("original request failed"),
+ max_fallbacks=3,
+ fallback_depth=0,
+ include_fallback_errors=True,
+ )
+
+
+class RecordingRouter:
+ def __init__(self):
+ self.received_kwargs = None
+
+ def log_retry(self, kwargs, e):
+ return kwargs
+
+ async def async_function_with_fallbacks(self, *args, **kwargs):
+ self.received_kwargs = kwargs
+ return StreamingWrapper()
+
+
+@pytest.mark.asyncio
+async def test_run_async_fallback_forwards_include_fallback_errors_to_nested_call():
+ """A nested fallback (multi-hop) must keep collecting errors, so the opt-in
+ flag has to reach the nested async_function_with_fallbacks call."""
+ router = RecordingRouter()
+ await run_async_fallback(
+ litellm_router=router,
+ fallback_model_group=["fallback-model"],
+ original_model_group="primary-model",
+ original_exception=RuntimeError("upstream limited request"),
+ max_fallbacks=3,
+ fallback_depth=0,
+ include_fallback_errors=True,
+ )
+
+ assert router.received_kwargs.get("include_fallback_errors") is True
+
+
+@pytest.mark.asyncio
+async def test_run_async_fallback_does_not_forward_flag_without_opt_in():
+ router = RecordingRouter()
+ await run_async_fallback(
+ litellm_router=router,
+ fallback_model_group=["fallback-model"],
+ original_model_group="primary-model",
+ original_exception=RuntimeError("upstream limited request"),
+ max_fallbacks=3,
+ fallback_depth=0,
+ )
+
+ assert "include_fallback_errors" not in router.received_kwargs
+
+
+@pytest.mark.asyncio
+async def test_run_async_fallback_skips_original_model_group():
+ response = await run_async_fallback(
+ litellm_router=FakeRouter(),
+ fallback_model_group=["primary-model", "fallback-model"],
+ original_model_group="primary-model",
+ original_exception=RuntimeError("original failed"),
+ max_fallbacks=3,
+ fallback_depth=0,
+ )
+
+ assert response._hidden_params["additional_headers"]["x-litellm-attempted-fallbacks"] == 1
diff --git a/tests/test_litellm/sandbox/test_opensandbox_sandbox.py b/tests/test_litellm/sandbox/test_opensandbox_sandbox.py
new file mode 100644
index 00000000000..0d7bcbe1e53
--- /dev/null
+++ b/tests/test_litellm/sandbox/test_opensandbox_sandbox.py
@@ -0,0 +1,647 @@
+import json
+
+import httpx
+import pytest
+
+import litellm
+from litellm.llms.base_llm.sandbox.transformation import ContainerHandle
+from litellm.llms.opensandbox.sandbox.transformation import (
+ MAX_OUTPUT_BYTES,
+ OPEN_SANDBOX_DEFAULT_TEMPLATE,
+ OpenSandboxSandboxConfig,
+)
+from litellm.utils import ProviderConfigManager
+
+TEST_API_BASE = "https://sandbox.test/v1"
+
+
+def http_status_error(status_code, url="http://test"):
+ return httpx.HTTPStatusError(
+ f"status {status_code}",
+ request=httpx.Request("GET", url),
+ response=httpx.Response(status_code),
+ )
+
+
+def sse(data):
+ return f"data: {json.dumps(data)}"
+
+
+class FakeResponse:
+ def __init__(self, *, json_data=None, lines=None, status_code=200):
+ self._json = json_data
+ self._lines = lines or []
+ self.status_code = status_code
+
+ def json(self):
+ return self._json
+
+ def raise_for_status(self):
+ if self.status_code >= 400:
+ raise http_status_error(self.status_code)
+
+ async def aiter_lines(self):
+ for line in self._lines:
+ yield line
+
+
+class FakeHTTPClient:
+ def __init__(
+ self,
+ *,
+ create_json=None,
+ sandbox_states=None,
+ endpoint_json=None,
+ endpoint_responses=None,
+ execute_lines=None,
+ delete_status=204,
+ execute_raises=None,
+ ):
+ self.create_json = create_json or {
+ "id": "osb_123",
+ "status": {"state": "Running"},
+ "createdAt": "2026-01-01T00:00:00Z",
+ "entrypoint": ["/opt/code-interpreter/code-interpreter.sh"],
+ }
+ self.sandbox_states = list(
+ sandbox_states
+ or [
+ {
+ "id": "osb_123",
+ "status": {"state": "Running"},
+ "createdAt": "2026-01-01T00:00:00Z",
+ "entrypoint": ["/opt/code-interpreter/code-interpreter.sh"],
+ }
+ ]
+ )
+ self.endpoint_json = endpoint_json or {
+ "endpoint": "execd.local:44772",
+ "headers": {"X-EXECD-ACCESS-TOKEN": "execd-token"},
+ }
+ self.endpoint_responses = (
+ list(endpoint_responses) if endpoint_responses is not None else None
+ )
+ self.execute_lines = execute_lines or []
+ self.delete_status = delete_status
+ self.execute_raises = execute_raises
+ self.calls = []
+
+ async def post(self, url, headers=None, json=None, stream=False, **kwargs):
+ self.calls.append(("POST", url, headers, json, {"stream": stream}))
+ if url.endswith("/sandboxes"):
+ return FakeResponse(json_data=self.create_json)
+ if url.endswith("/code"):
+ if self.execute_raises is not None:
+ raise self.execute_raises
+ return FakeResponse(lines=self.execute_lines)
+ raise AssertionError(f"unexpected POST {url}")
+
+ async def get(self, url, headers=None, params=None, **kwargs):
+ self.calls.append(("GET", url, headers, None, params))
+ if "/endpoints/44772" in url:
+ if self.endpoint_responses is not None and self.endpoint_responses:
+ response = self.endpoint_responses.pop(0)
+ if isinstance(response, Exception):
+ raise response
+ if isinstance(response, FakeResponse):
+ return response
+ return FakeResponse(json_data=response)
+ return FakeResponse(json_data=self.endpoint_json)
+ if "/sandboxes/" in url:
+ state = self.sandbox_states.pop(0)
+ return FakeResponse(json_data=state)
+ raise AssertionError(f"unexpected GET {url}")
+
+ async def delete(self, url, headers=None, **kwargs):
+ self.calls.append(("DELETE", url, headers, None, None))
+ if not (200 <= self.delete_status < 300):
+ raise http_status_error(self.delete_status, url)
+ return FakeResponse(status_code=self.delete_status)
+
+
+def test_parse_sse_lines_maps_output_result_count_and_error():
+ lines = [
+ sse({"type": "stdout", "text": "hello\n"}),
+ sse({"type": "stderr", "text": "warn\n"}),
+ sse({"type": "result", "results": {"text/plain": "4"}}),
+ sse({"type": "execution_count", "execution_count": 7}),
+ sse(
+ {
+ "type": "error",
+ "error": {
+ "ename": "ValueError",
+ "evalue": "bad",
+ "traceback": ["Traceback"],
+ },
+ }
+ ),
+ ]
+
+ result = OpenSandboxSandboxConfig._parse_lines(lines)
+
+ assert result.stdout == "hello\n"
+ assert result.stderr == "warn\n"
+ assert result.results == [{"text/plain": "4"}]
+ assert result.execution_count == 7
+ assert result.error == {
+ "name": "ValueError",
+ "value": "bad",
+ "traceback": ["Traceback"],
+ }
+
+
+def test_parse_sse_lines_skips_non_json_and_control_lines():
+ lines = [
+ "event: message",
+ "not-json",
+ "",
+ sse({"type": "stdout", "text": "ok\n"}),
+ ]
+
+ result = OpenSandboxSandboxConfig._parse_lines(lines)
+
+ assert result.stdout == "ok\n"
+ assert result.error is None
+
+
+def test_parse_sse_lines_maps_fallback_shapes():
+ lines = [
+ "data:",
+ sse(["not-a-dict"]),
+ sse({"code": "BadRequest", "message": "nope"}),
+ sse({"type": "result", "text/plain": "4"}),
+ sse({"type": "error", "name": "RuntimeError", "text": "boom"}),
+ sse({"type": "execution_count", "execution_count": "8"}),
+ ]
+
+ result = OpenSandboxSandboxConfig._parse_lines(lines)
+
+ assert result.results == [{"text/plain": "4"}]
+ assert result.execution_count == 8
+ assert result.error == {
+ "name": "BadRequest",
+ "value": "nope",
+ "traceback": [],
+ }
+ fallback_error = OpenSandboxSandboxConfig._parse_lines(
+ [sse({"type": "error", "name": "RuntimeError", "text": "boom"})]
+ )
+ assert fallback_error.error == {
+ "name": "RuntimeError",
+ "value": "boom",
+ "traceback": [],
+ }
+ empty_string_error = OpenSandboxSandboxConfig._parse_lines(
+ [
+ sse(
+ {
+ "type": "error",
+ "error": {
+ "ename": "",
+ "name": "FallbackName",
+ "evalue": "",
+ "value": "fallback value",
+ "traceback": [],
+ },
+ }
+ )
+ ]
+ )
+ assert empty_string_error.error == {
+ "name": "",
+ "value": "",
+ "traceback": [],
+ }
+
+
+def test_static_helpers_cover_defaults_and_fallbacks(monkeypatch):
+ def fake_secret(key):
+ if key == "OPEN_SANDBOX_API_KEY":
+ return "env-key"
+ if key == "OPEN_SANDBOX_API_BASE":
+ return TEST_API_BASE
+ return None
+
+ monkeypatch.setattr(
+ "litellm.llms.opensandbox.sandbox.transformation.get_secret_str",
+ fake_secret,
+ )
+ config = OpenSandboxSandboxConfig()
+ handle = ContainerHandle(id="osb", provider="opensandbox", domain="http://x/v1")
+
+ assert config.validate_environment() == "env-key"
+ assert config.validate_environment(api_key="") == ""
+ assert config._api_key(api_key=None, handle=handle) == "env-key"
+
+ handle._hidden_params = {"api_key": "stored-key"}
+ assert config._api_key(api_key=None, handle=handle) == "stored-key"
+ assert config._http(None) is not None
+
+ body = config._create_body(
+ template=None,
+ timeout=None,
+ allow_internet_access=False,
+ metadata=None,
+ env_vars=None,
+ resource_limits=None,
+ resource_requests=None,
+ entrypoint=None,
+ network_policy={"egress": [{"domain": "example.com"}]},
+ secure_access=True,
+ )
+ assert body["networkPolicy"] == {"egress": [{"domain": "example.com"}]}
+ assert body["secureAccess"] is True
+
+ other_body = config._create_body(
+ template=None,
+ timeout=None,
+ allow_internet_access=False,
+ metadata=None,
+ env_vars=None,
+ resource_limits=None,
+ resource_requests=None,
+ entrypoint=None,
+ network_policy=None,
+ secure_access=False,
+ )
+ assert body["resourceLimits"] is not other_body["resourceLimits"]
+
+ assert config._sandbox_state(None) is None
+ assert config._sandbox_state({"status": "Running"}) is None
+ assert config._as_str_dict(None) == {}
+ assert config._endpoint_base_url("http://execd.local", "https://api/v1") == (
+ "http://execd.local"
+ )
+ assert config._api_base(None) == TEST_API_BASE
+ assert config._api_base("https://direct.test/v1/") == "https://direct.test/v1"
+ assert config._as_int("9") == 9
+ assert config._as_int("nope") is None
+ assert config._as_int(None) is None
+ assert isinstance(
+ ProviderConfigManager.get_provider_sandbox_config("opensandbox"),
+ OpenSandboxSandboxConfig,
+ )
+
+
+def test_api_base_requires_kwarg_or_env(monkeypatch):
+ monkeypatch.setattr(
+ "litellm.llms.opensandbox.sandbox.transformation.get_secret_str",
+ lambda key: None,
+ )
+
+ with pytest.raises(ValueError, match="api_base is required"):
+ OpenSandboxSandboxConfig._api_base(None)
+
+
+@pytest.mark.asyncio
+async def test_create_posts_default_body_and_omits_empty_api_key():
+ client = FakeHTTPClient()
+
+ handle = await OpenSandboxSandboxConfig().acreate_sandbox(
+ api_key="", api_base=TEST_API_BASE, client=client
+ )
+
+ method, url, headers, body, _ = client.calls[0]
+ assert method == "POST"
+ assert url == f"{TEST_API_BASE}/sandboxes"
+ assert "OPEN-SANDBOX-API-KEY" not in headers
+ assert body["image"] == {"uri": OPEN_SANDBOX_DEFAULT_TEMPLATE}
+ assert body["entrypoint"] == ["/opt/code-interpreter/code-interpreter.sh"]
+ assert body["timeout"] == 300
+ assert body["resourceLimits"] == {"cpu": "1", "memory": "2Gi"}
+ assert body["networkPolicy"] == {"defaultAction": "deny", "egress": []}
+ assert handle.id == "osb_123"
+ assert handle._hidden_params["execd_endpoint"] == "execd.local:44772"
+
+
+@pytest.mark.asyncio
+async def test_create_can_opt_into_internet_access():
+ client = FakeHTTPClient()
+
+ await OpenSandboxSandboxConfig().acreate_sandbox(
+ api_key="",
+ api_base=TEST_API_BASE,
+ allow_internet_access=True,
+ client=client,
+ )
+
+ _, _, _, body, _ = client.calls[0]
+ assert "networkPolicy" not in body
+
+
+@pytest.mark.asyncio
+async def test_create_custom_options_poll_and_endpoint_resolution():
+ client = FakeHTTPClient(
+ create_json={
+ "id": "osb_pending",
+ "status": {"state": "Pending"},
+ "createdAt": "2026-01-01T00:00:00Z",
+ "entrypoint": ["/bin/sh"],
+ },
+ sandbox_states=[
+ {
+ "id": "osb_pending",
+ "status": {"state": "Running"},
+ "createdAt": "2026-01-01T00:00:00Z",
+ "entrypoint": ["/bin/sh"],
+ }
+ ],
+ )
+
+ handle = await OpenSandboxSandboxConfig().acreate_sandbox(
+ template="custom/image:latest",
+ timeout=600,
+ allow_internet_access=False,
+ api_key="osb-key",
+ api_base="https://sandbox.example/v1",
+ metadata={"suite": "unit"},
+ env_vars={"PYTHONUNBUFFERED": "1"},
+ resource_limits={"cpu": "500m", "memory": "512Mi"},
+ resource_requests={"cpu": "250m", "memory": "256Mi"},
+ entrypoint=["/bin/sh", "-lc", "sleep 3600"],
+ use_server_proxy=True,
+ client=client,
+ )
+
+ _, create_url, create_headers, body, _ = client.calls[0]
+ _, poll_url, poll_headers, _, _ = client.calls[1]
+ _, endpoint_url, endpoint_headers, _, endpoint_params = client.calls[2]
+
+ assert create_url == "https://sandbox.example/v1/sandboxes"
+ assert create_headers["OPEN-SANDBOX-API-KEY"] == "osb-key"
+ assert body["image"] == {"uri": "custom/image:latest"}
+ assert body["entrypoint"] == ["/bin/sh", "-lc", "sleep 3600"]
+ assert body["metadata"] == {"suite": "unit"}
+ assert body["env"] == {"PYTHONUNBUFFERED": "1"}
+ assert body["resourceLimits"] == {"cpu": "500m", "memory": "512Mi"}
+ assert body["resourceRequests"] == {"cpu": "250m", "memory": "256Mi"}
+ assert body["networkPolicy"] == {"defaultAction": "deny", "egress": []}
+ assert poll_url == "https://sandbox.example/v1/sandboxes/osb_pending"
+ assert poll_headers["OPEN-SANDBOX-API-KEY"] == "osb-key"
+ assert endpoint_url.endswith("/sandboxes/osb_pending/endpoints/44772")
+ assert endpoint_headers["OPEN-SANDBOX-API-KEY"] == "osb-key"
+ assert endpoint_params == {"use_server_proxy": True}
+ assert handle.id == "osb_pending"
+
+
+@pytest.mark.asyncio
+async def test_create_waits_across_pending_state(monkeypatch):
+ client = FakeHTTPClient(
+ create_json={
+ "id": "osb_pending",
+ "status": {"state": "Pending"},
+ "createdAt": "2026-01-01T00:00:00Z",
+ },
+ sandbox_states=[
+ {"id": "osb_pending", "status": {"state": "Pending"}},
+ {"id": "osb_pending", "status": {"state": "Running"}},
+ ],
+ )
+ sleeps = []
+
+ async def fake_sleep(interval):
+ sleeps.append(interval)
+
+ monkeypatch.setattr(
+ "litellm.llms.opensandbox.sandbox.transformation.asyncio.sleep", fake_sleep
+ )
+
+ handle = await OpenSandboxSandboxConfig().acreate_sandbox(
+ api_key="",
+ api_base=TEST_API_BASE,
+ ready_timeout=1,
+ poll_interval=0.01,
+ client=client,
+ )
+
+ assert handle.id == "osb_pending"
+ assert sleeps == [0.01]
+
+
+@pytest.mark.asyncio
+async def test_create_raises_for_terminal_state():
+ client = FakeHTTPClient(
+ create_json={"id": "osb_failed", "status": {"state": "Pending"}},
+ sandbox_states=[
+ {"id": "osb_failed", "status": {"state": "Failed"}},
+ ],
+ )
+
+ with pytest.raises(ValueError, match="entered Failed"):
+ await OpenSandboxSandboxConfig().acreate_sandbox(
+ api_key="", api_base=TEST_API_BASE, client=client
+ )
+
+
+@pytest.mark.asyncio
+async def test_create_times_out_waiting_for_running():
+ client = FakeHTTPClient(
+ create_json={"id": "osb_slow", "status": {"state": "Pending"}},
+ sandbox_states=[
+ {"id": "osb_slow", "status": {"state": "Pending"}},
+ ],
+ )
+
+ with pytest.raises(TimeoutError, match="was not Running"):
+ await OpenSandboxSandboxConfig().acreate_sandbox(
+ api_key="",
+ api_base=TEST_API_BASE,
+ ready_timeout=0,
+ poll_interval=0,
+ client=client,
+ )
+
+
+@pytest.mark.asyncio
+async def test_create_waits_for_endpoint_resolution(monkeypatch):
+ client = FakeHTTPClient(
+ endpoint_responses=[
+ http_status_error(404, f"{TEST_API_BASE}/sandboxes/osb_123"),
+ {
+ "endpoint": "execd.local:44772",
+ "headers": {"X-EXECD-ACCESS-TOKEN": "execd-token"},
+ },
+ ],
+ )
+ sleeps = []
+
+ async def fake_sleep(interval):
+ sleeps.append(interval)
+
+ monkeypatch.setattr(
+ "litellm.llms.opensandbox.sandbox.transformation.asyncio.sleep", fake_sleep
+ )
+
+ handle = await OpenSandboxSandboxConfig().acreate_sandbox(
+ api_key="",
+ api_base=TEST_API_BASE,
+ ready_timeout=1,
+ poll_interval=0.01,
+ client=client,
+ )
+
+ endpoint_calls = [call for call in client.calls if "/endpoints/44772" in call[1]]
+ assert handle._hidden_params["execd_endpoint"] == "execd.local:44772"
+ assert len(endpoint_calls) == 2
+ assert sleeps == [0.01]
+
+
+@pytest.mark.asyncio
+async def test_create_raises_when_endpoint_is_missing():
+ client = FakeHTTPClient(endpoint_json={"headers": {"X": "y"}})
+
+ with pytest.raises(TimeoutError, match="execd endpoint.*not ready"):
+ await OpenSandboxSandboxConfig().acreate_sandbox(
+ api_key="", api_base=TEST_API_BASE, ready_timeout=0, client=client
+ )
+
+
+@pytest.mark.asyncio
+async def test_create_reraises_non_404_endpoint_error():
+ client = FakeHTTPClient(endpoint_responses=[http_status_error(500)])
+
+ with pytest.raises(httpx.HTTPStatusError):
+ await OpenSandboxSandboxConfig().acreate_sandbox(
+ api_key="", api_base=TEST_API_BASE, client=client
+ )
+
+
+@pytest.mark.asyncio
+async def test_run_code_resolves_bare_id_and_posts_sse_request():
+ client = FakeHTTPClient(
+ execute_lines=[
+ sse({"type": "stdout", "text": "42\n"}),
+ ]
+ )
+
+ result = await OpenSandboxSandboxConfig().arun_code(
+ container="osb_bare",
+ code="print(6*7)",
+ language="python",
+ api_key="",
+ api_base="http://sandbox.local/v1",
+ client=client,
+ )
+
+ endpoint_call = client.calls[0]
+ run_call = client.calls[1]
+ assert endpoint_call[0] == "GET"
+ assert (
+ endpoint_call[1] == "http://sandbox.local/v1/sandboxes/osb_bare/endpoints/44772"
+ )
+ assert run_call[0] == "POST"
+ assert run_call[1] == "http://execd.local:44772/code"
+ assert run_call[2]["X-EXECD-ACCESS-TOKEN"] == "execd-token"
+ assert run_call[3] == {
+ "code": "print(6*7)",
+ "context": {"language": "python"},
+ }
+ assert run_call[4] == {"stream": True}
+ assert result.stdout == "42\n"
+
+
+@pytest.mark.asyncio
+async def test_run_code_uses_https_for_scheme_less_endpoint_when_api_base_is_https():
+ client = FakeHTTPClient()
+ handle = ContainerHandle(
+ id="osb_https", provider="opensandbox", domain="https://sandbox.example/v1"
+ )
+ handle._hidden_params = {
+ "execd_endpoint": "execd.example/route/44772",
+ "execd_headers": {},
+ }
+
+ await OpenSandboxSandboxConfig().arun_code(
+ container=handle, code="print(1)", client=client
+ )
+
+ assert client.calls[0][1] == "https://execd.example/route/44772/code"
+
+
+@pytest.mark.asyncio
+async def test_run_code_aborts_on_output_over_cap():
+ client = FakeHTTPClient(execute_lines=["x" * (MAX_OUTPUT_BYTES + 1)])
+ handle = ContainerHandle(id="osb_big", provider="opensandbox", domain="http://x/v1")
+ handle._hidden_params = {"execd_endpoint": "execd.local:44772", "execd_headers": {}}
+
+ with pytest.raises(ValueError, match="exceeded"):
+ await OpenSandboxSandboxConfig().arun_code(
+ container=handle, code="print('x')", client=client
+ )
+
+
+@pytest.mark.asyncio
+async def test_delete_returns_false_on_404():
+ client = FakeHTTPClient(delete_status=404)
+
+ ok = await OpenSandboxSandboxConfig().adelete_sandbox(
+ container="osb_gone",
+ api_key="",
+ api_base="http://sandbox.local/v1",
+ client=client,
+ )
+
+ assert ok is False
+
+
+@pytest.mark.asyncio
+async def test_delete_reraises_non_404_http_error():
+ client = FakeHTTPClient(delete_status=500)
+
+ with pytest.raises(httpx.HTTPStatusError):
+ await OpenSandboxSandboxConfig().adelete_sandbox(
+ container="osb_err",
+ api_key="",
+ api_base="http://sandbox.local/v1",
+ client=client,
+ )
+
+
+@pytest.mark.asyncio
+async def test_public_lifecycle_create_run_delete():
+ client = FakeHTTPClient(
+ execute_lines=[
+ sse({"type": "stdout", "text": "42\n"}),
+ ]
+ )
+
+ container = await litellm.acreate_sandbox(
+ provider="opensandbox", api_key="", api_base=TEST_API_BASE, client=client
+ )
+ result = await litellm.arun_code(
+ provider="opensandbox",
+ container=container,
+ code="print(6*7)",
+ api_key="",
+ client=client,
+ )
+ ok = await litellm.adelete_sandbox(
+ provider="opensandbox",
+ container=container,
+ api_key="",
+ client=client,
+ )
+
+ assert container.id == "osb_123"
+ assert result.stdout == "42\n"
+ assert ok is True
+
+
+@pytest.mark.asyncio
+async def test_code_interpreter_tool_deletes_even_when_run_raises():
+ client = FakeHTTPClient(execute_raises=RuntimeError("boom"))
+
+ with pytest.raises(RuntimeError, match="boom"):
+ await litellm.acode_interpreter_tool(
+ provider="opensandbox",
+ code="1/0",
+ api_key="",
+ api_base=TEST_API_BASE,
+ client=client,
+ )
+
+ assert [call[0] for call in client.calls] == ["POST", "GET", "POST", "DELETE"]
+ assert client.calls[0][1].endswith("/sandboxes")
+ assert client.calls[1][1].endswith("/endpoints/44772")
+ assert client.calls[2][1].endswith("/code")
+ assert client.calls[3][1].endswith("/sandboxes/osb_123")
diff --git a/tests/test_litellm/test_cloudflare_workers_ai_model_metadata.py b/tests/test_litellm/test_cloudflare_workers_ai_model_metadata.py
new file mode 100644
index 00000000000..9ca4515239a
--- /dev/null
+++ b/tests/test_litellm/test_cloudflare_workers_ai_model_metadata.py
@@ -0,0 +1,89 @@
+"""
+Regression tests for the Cloudflare Workers AI text-generation catalog in the
+model-cost map.
+
+The Cloudflare list was badly stale (only 4 ancient entries). These tests pin
+the newly added current Workers AI models (sourced from Cloudflare's live
+``/ai/models/search?task=Text Generation`` catalog) and guard against the root
+``model_prices_and_context_window.json`` and the bundled
+``litellm/model_prices_and_context_window_backup.json`` drifting out of sync for
+the ``cloudflare/`` namespace.
+"""
+
+import json
+import os
+
+import pytest
+
+import litellm
+
+ROOT_MAP = os.path.join(
+ os.path.dirname(os.path.dirname(litellm.__file__)),
+ "model_prices_and_context_window.json",
+)
+BACKUP_MAP = os.path.join(
+ os.path.dirname(litellm.__file__),
+ "model_prices_and_context_window_backup.json",
+)
+
+
+@pytest.fixture(autouse=True)
+def _use_local_model_cost_map(monkeypatch):
+ original_model_cost = litellm.model_cost
+ monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
+ litellm.model_cost = litellm.get_model_cost_map(url="")
+ try:
+ yield
+ finally:
+ litellm.model_cost = original_model_cost
+
+
+def _load(path: str) -> dict:
+ with open(path, encoding="utf-8") as f:
+ return json.load(f)
+
+
+def _cloudflare_keys(data: dict) -> set:
+ return {k for k in data if k.startswith("cloudflare/")}
+
+
+def test_glm_5_2_entry_is_present_and_well_formed():
+ entry = litellm.model_cost["cloudflare/@cf/zai-org/glm-5.2"]
+ assert entry["litellm_provider"] == "cloudflare"
+ assert entry["mode"] == "chat"
+ assert entry["supports_function_calling"] is True
+ assert entry["input_cost_per_token"] > 0
+ assert entry["output_cost_per_token"] > 0
+
+
+def test_vision_model_is_flagged_supports_vision():
+ entry = litellm.model_cost["cloudflare/@cf/meta/llama-3.2-11b-vision-instruct"]
+ assert entry["litellm_provider"] == "cloudflare"
+ assert entry.get("supports_vision") is True
+
+
+def test_additional_current_models_are_present():
+ for key in (
+ "cloudflare/@cf/openai/gpt-oss-120b",
+ "cloudflare/@cf/meta/llama-3.3-70b-instruct-fp8-fast",
+ ):
+ entry = litellm.model_cost[key]
+ assert entry["litellm_provider"] == "cloudflare"
+ assert entry["mode"] == "chat"
+ assert entry["supports_function_calling"] is True
+ assert entry["input_cost_per_token"] > 0
+ assert entry["output_cost_per_token"] > 0
+
+
+def test_root_and_backup_have_identical_cloudflare_keys():
+ if not os.path.exists(ROOT_MAP):
+ pytest.skip("root cost map only ships in source checkouts")
+ assert _cloudflare_keys(_load(ROOT_MAP)) == _cloudflare_keys(_load(BACKUP_MAP))
+
+
+def test_root_and_backup_cloudflare_entries_are_byte_for_byte_equal():
+ if not os.path.exists(ROOT_MAP):
+ pytest.skip("root cost map only ships in source checkouts")
+ root = {k: v for k, v in _load(ROOT_MAP).items() if k.startswith("cloudflare/")}
+ backup = {k: v for k, v in _load(BACKUP_MAP).items() if k.startswith("cloudflare/")}
+ assert root == backup
diff --git a/tests/test_litellm/test_completion_timeout_resolution.py b/tests/test_litellm/test_completion_timeout_resolution.py
index a76cc6f7de8..7eb79e90e60 100644
--- a/tests/test_litellm/test_completion_timeout_resolution.py
+++ b/tests/test_litellm/test_completion_timeout_resolution.py
@@ -63,8 +63,9 @@ def test_global_timeout_from_litellm_settings():
)
-def test_global_timeout_package_default_coerced_to_600_for_completion():
- """Package default 6000s → 600s for completion-only path."""
+def test_explicit_global_timeout_6000_is_preserved():
+ """The caller passes the explicitly-configured value (or None); an explicit
+ 6000 must be honored, not silently coerced to 600."""
assert (
CompletionTimeout.resolve(
None,
@@ -73,7 +74,7 @@ def test_global_timeout_package_default_coerced_to_600_for_completion():
global_timeout=6000.0,
supports_httpx_timeout=supports_httpx_timeout,
)
- == 600.0
+ == 6000.0
)
diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py
index c2aa4a095d4..3e2150848a7 100644
--- a/tests/test_litellm/test_router.py
+++ b/tests/test_litellm/test_router.py
@@ -4901,3 +4901,107 @@ def test_is_deployment_blocked_static_helper_reflects_blocked_flag():
)
is True
)
+
+
+class TestRouterRequestTimeoutPropagation:
+ """litellm_settings.request_timeout must act as an independent per-attempt timeout.
+
+ Regression for LIT-2369: request_timeout was shadowed by router_settings.timeout,
+ so Bedrock (and other provider) calls fell back to the hardcoded 600s httpx
+ default instead of the configured value.
+ """
+
+ def _make_router(self, timeout=None, stream_timeout=None):
+ return litellm.Router(
+ model_list=[
+ {
+ "model_name": "test-model",
+ "litellm_params": {
+ "model": "openai/gpt-4",
+ "api_key": "sk-test",
+ },
+ }
+ ],
+ timeout=timeout,
+ stream_timeout=stream_timeout,
+ )
+
+ @pytest.fixture
+ def explicit_request_timeout(self):
+ original_value = litellm.request_timeout
+ original_flag = litellm.request_timeout_explicitly_set
+ litellm.request_timeout = 300
+ litellm.request_timeout_explicitly_set = True
+ try:
+ yield 300
+ finally:
+ litellm.request_timeout = original_value
+ litellm.request_timeout_explicitly_set = original_flag
+
+ def test_request_timeout_stored_independently_when_both_set(
+ self, explicit_request_timeout
+ ):
+ router = self._make_router(timeout=330)
+ assert router.timeout == 330
+ assert router.request_timeout == 300
+
+ def test_request_timeout_none_when_not_explicitly_configured(self):
+ original_value = litellm.request_timeout
+ original_flag = litellm.request_timeout_explicitly_set
+ litellm.request_timeout = litellm.constants.DEFAULT_REQUEST_TIMEOUT_SECONDS
+ litellm.request_timeout_explicitly_set = False
+ try:
+ router = self._make_router(timeout=330)
+ assert router.timeout == 330
+ assert router.request_timeout is None
+ finally:
+ litellm.request_timeout = original_value
+ litellm.request_timeout_explicitly_set = original_flag
+
+ def test_non_stream_prefers_request_timeout_over_router_timeout(
+ self, explicit_request_timeout
+ ):
+ router = self._make_router(timeout=330)
+ assert router._get_non_stream_timeout(kwargs={}, data={}) == 300
+
+ def test_stream_prefers_request_timeout_over_router_timeout(
+ self, explicit_request_timeout
+ ):
+ router = self._make_router(timeout=330)
+ # stream=True resolves through _get_stream_timeout; request_timeout must win.
+ assert router._get_timeout(kwargs={"stream": True}, data={}) == 300
+
+ def test_explicit_stream_timeout_still_wins_over_request_timeout(
+ self, explicit_request_timeout
+ ):
+ router = self._make_router(timeout=330, stream_timeout=45)
+ assert router._get_stream_timeout(kwargs={}, data={}) == 45
+
+ def test_non_stream_falls_through_to_router_timeout_without_request_timeout(self):
+ original_value = litellm.request_timeout
+ original_flag = litellm.request_timeout_explicitly_set
+ litellm.request_timeout = litellm.constants.DEFAULT_REQUEST_TIMEOUT_SECONDS
+ litellm.request_timeout_explicitly_set = False
+ try:
+ router = self._make_router(timeout=330)
+ assert router._get_non_stream_timeout(kwargs={}, data={}) == 330
+ finally:
+ litellm.request_timeout = original_value
+ litellm.request_timeout_explicitly_set = original_flag
+
+ def test_per_deployment_timeout_overrides_request_timeout(
+ self, explicit_request_timeout
+ ):
+ router = self._make_router(timeout=330)
+ assert router._get_non_stream_timeout(kwargs={}, data={"timeout": 120}) == 120
+
+ def test_per_request_timeout_overrides_request_timeout(
+ self, explicit_request_timeout
+ ):
+ router = self._make_router(timeout=330)
+ assert (
+ router._get_non_stream_timeout(
+ kwargs={"timeout": 60}, data={"timeout": 120}
+ )
+ == 60
+ )
diff --git a/tests/test_litellm/test_router_per_deployment_num_retries.py b/tests/test_litellm/test_router_per_deployment_num_retries.py
index 154ba579e4e..af2372616a6 100644
--- a/tests/test_litellm/test_router_per_deployment_num_retries.py
+++ b/tests/test_litellm/test_router_per_deployment_num_retries.py
@@ -4,8 +4,9 @@ GitHub Issue: #18968 - Per-deployment max_retries/num_retries in litellm_params
"""
import pytest
-from unittest.mock import MagicMock, patch
+from unittest.mock import patch
+import litellm
from litellm import Router
@@ -188,3 +189,133 @@ class TestPerDeploymentNumRetries:
# Verify num_retries was converted from string to int
assert exc.num_retries == 6
+
+
+class TestNumRetriesNoneGuard:
+ """
+ Regression tests for the num_retries=None TypeError in async_function_with_retries.
+
+ When num_retries reaches async_function_with_retries as None - e.g. a caller passes
+ num_retries=None explicitly (dict.get() does not fall back on an existing None value),
+ an auto_router/complexity_router path does not propagate it, or
+ Router.update_settings(num_retries=None) is used - AND the underlying call fails with a
+ retryable error, the comparison `if num_retries > 0:` raised:
+
+ TypeError: '>' not supported between instances of 'NoneType' and 'int'
+
+ This masked the real upstream error (rate limit / connection / 5xx) behind a TypeError.
+ Related issues: #23316, #25889, #23699, #28126.
+ """
+
+ @staticmethod
+ def _mock_router(num_retries=2):
+ return Router(
+ model_list=[
+ {
+ "model_name": "mock-model",
+ "litellm_params": {
+ "model": "gpt-4o-mini",
+ "mock_response": "ok",
+ },
+ }
+ ],
+ num_retries=num_retries,
+ )
+
+ def test_update_kwargs_normalises_explicit_none_to_router_default(self):
+ """
+ _update_kwargs_before_fallbacks must normalise an explicit num_retries=None to
+ the router default (not leave it as None), while preserving an explicit 0.
+ """
+ router = self._mock_router(num_retries=4)
+
+ # explicit None -> router default
+ kwargs = {"num_retries": None}
+ router._update_kwargs_before_fallbacks(model="mock-model", kwargs=kwargs)
+ assert kwargs["num_retries"] == 4
+
+ # explicit 0 is preserved (retries stay disabled)
+ kwargs = {"num_retries": 0}
+ router._update_kwargs_before_fallbacks(model="mock-model", kwargs=kwargs)
+ assert kwargs["num_retries"] == 0
+
+ # absent -> router default (unchanged behaviour)
+ kwargs = {}
+ router._update_kwargs_before_fallbacks(model="mock-model", kwargs=kwargs)
+ assert kwargs["num_retries"] == 4
+
+ # explicit None with router default also None -> 0 (mirrors the downstream guard)
+ router.num_retries = None # simulate update_settings(num_retries=None) (#28126)
+ kwargs = {"num_retries": None}
+ router._update_kwargs_before_fallbacks(model="mock-model", kwargs=kwargs)
+ assert kwargs["num_retries"] == 0
+
+ @pytest.mark.asyncio
+ async def test_acompletion_num_retries_none_does_not_raise_typeerror(self):
+ """
+ Per-request num_retries=None + a retryable error must NOT raise TypeError.
+ The router falls back to its configured num_retries and retries the (transient)
+ error, so the request succeeds.
+ """
+ router = self._mock_router(num_retries=2)
+ with patch("asyncio.sleep", return_value=None):
+ response = await router.acompletion(
+ model="mock-model",
+ messages=[{"role": "user", "content": "hi"}],
+ num_retries=None, # the trigger
+ mock_testing_rate_limit_error=True, # retryable error path
+ )
+ assert response.choices[0].message.content == "ok"
+
+ @pytest.mark.asyncio
+ async def test_async_function_with_retries_none_falls_back_to_zero(self):
+ """
+ When both the per-request value AND the router-level setting are None
+ (e.g. after Router.update_settings(num_retries=None), #28126), num_retries must
+ fall back to 0 and the real retryable error must surface - not a TypeError.
+ """
+ router = self._mock_router(num_retries=0)
+ router.num_retries = None # simulate update_settings(num_retries=None)
+
+ async def failing_fn(*args, **kwargs):
+ raise litellm.RateLimitError(
+ message="boom", model="mock-model", llm_provider="openai"
+ )
+
+ with patch("asyncio.sleep", return_value=None):
+ with pytest.raises(litellm.RateLimitError):
+ await router.async_function_with_retries(
+ original_function=failing_fn,
+ model="mock-model",
+ messages=[{"role": "user", "content": "hi"}],
+ num_retries=None,
+ )
+
+ @pytest.mark.asyncio
+ async def test_async_function_with_retries_none_falls_back_to_router_default(self):
+ """
+ A None per-request num_retries falls back to the router-level setting, so retries
+ still happen (original_function is invoked more than once) before the real error
+ is raised - proving None did not silently disable retries or crash.
+ """
+ router = self._mock_router(num_retries=3)
+ calls = {"n": 0}
+
+ async def failing_fn(*args, **kwargs):
+ calls["n"] += 1
+ raise litellm.InternalServerError(
+ message="boom", model="mock-model", llm_provider="openai"
+ )
+
+ with patch("asyncio.sleep", return_value=None):
+ with pytest.raises(litellm.InternalServerError):
+ await router.async_function_with_retries(
+ original_function=failing_fn,
+ model="mock-model",
+ messages=[{"role": "user", "content": "hi"}],
+ metadata={}, # populated by acompletion in the real path; log_retry needs it
+ num_retries=None,
+ )
+
+ # 1 initial attempt + at least 1 retry -> proves None fell back to a positive int
+ assert calls["n"] >= 2
diff --git a/tests/test_litellm/test_router_streaming_fallback_metadata.py b/tests/test_litellm/test_router_streaming_fallback_metadata.py
new file mode 100644
index 00000000000..6ed70dc7cfe
--- /dev/null
+++ b/tests/test_litellm/test_router_streaming_fallback_metadata.py
@@ -0,0 +1,187 @@
+import json
+from unittest.mock import MagicMock
+
+import pytest
+
+import litellm
+from litellm.proxy.proxy_server import _should_include_fallback_errors
+from litellm.router import Router
+from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict
+
+
+def test_apply_fallback_hidden_params_copies_from_fallback_response():
+ fallback_errors = [
+ {
+ "message": "litellm.RateLimitError: upstream limited request",
+ "type": "RateLimitError",
+ "param": None,
+ "code": "429",
+ }
+ ]
+ chunk = litellm.ModelResponseStream(
+ id="test",
+ model="openai/internal-fallback",
+ choices=[],
+ )
+ chunk._hidden_params = {
+ "additional_headers": {"x-existing-chunk-header": "keep"},
+ "model_id": "chunk-model-id",
+ }
+ fallback_response = MagicMock()
+ fallback_response._hidden_params = {
+ "additional_headers": {
+ "x-litellm-attempted-fallbacks": 1,
+ "x-litellm-model-group": "fallback-model",
+ "x-litellm-fallback-errors": json.dumps(fallback_errors),
+ },
+ "api_base": "https://fallback.example",
+ }
+
+ Router._apply_fallback_hidden_params_to_item(
+ fallback_item=chunk,
+ prepared_fallback_hidden_params=Router._prepare_fallback_hidden_params(
+ fallback_response
+ ),
+ )
+
+ assert chunk._hidden_params["api_base"] == "https://fallback.example"
+ assert chunk._hidden_params["model_id"] == "chunk-model-id"
+ assert chunk._hidden_params["additional_headers"] == {
+ "x-existing-chunk-header": "keep",
+ "x-litellm-attempted-fallbacks": 1,
+ "x-litellm-model-group": "fallback-model",
+ "x-litellm-fallback-errors": json.dumps(fallback_errors),
+ }
+
+
+def _two_group_fallback_router() -> Router:
+ return litellm.Router(
+ model_list=[
+ {
+ "model_name": "primary-model",
+ "litellm_params": {"model": "openai/gpt-fake", "api_key": "sk-fake"},
+ },
+ {
+ "model_name": "fallback-model",
+ "litellm_params": {"model": "openai/gpt-fake-2", "api_key": "sk-fake"},
+ },
+ ],
+ fallbacks=[{"primary-model": ["fallback-model"]}],
+ )
+
+
+def _additional_headers(response: object) -> dict:
+ return get_hidden_params_dict(response).get("additional_headers", {})
+
+
+@pytest.mark.asyncio
+async def test_include_fallback_errors_propagates_through_router():
+ router = _two_group_fallback_router()
+
+ response = await router.acompletion(
+ model="primary-model",
+ messages=[{"role": "user", "content": "Hello"}],
+ mock_testing_fallbacks=True,
+ mock_response="fallback success",
+ include_fallback_errors=True,
+ )
+
+ headers = _additional_headers(response)
+ assert headers["x-litellm-attempted-fallbacks"] == 1
+ errors = json.loads(headers["x-litellm-fallback-errors"])
+ assert isinstance(errors, list) and len(errors) >= 1
+ assert set(errors[0].keys()) == {"message", "type", "param", "code"}
+
+
+@pytest.mark.asyncio
+async def test_router_omits_fallback_errors_without_opt_in():
+ router = _two_group_fallback_router()
+
+ response = await router.acompletion(
+ model="primary-model",
+ messages=[{"role": "user", "content": "Hello"}],
+ mock_testing_fallbacks=True,
+ mock_response="fallback success",
+ )
+
+ headers = _additional_headers(response)
+ assert headers["x-litellm-attempted-fallbacks"] == 1
+ assert "x-litellm-fallback-errors" not in headers
+
+
+def test_prepare_fallback_hidden_params_no_additional_headers():
+ class FakeResponse:
+ _hidden_params = {"api_base": "http://example.com"}
+
+ hidden_params, headers = Router._prepare_fallback_hidden_params(FakeResponse())
+ assert hidden_params == {"api_base": "http://example.com"}
+ assert headers == {}
+
+
+def test_apply_fallback_hidden_params_to_item_none_item():
+ Router._apply_fallback_hidden_params_to_item(
+ None, ({"api_base": "http://fallback.example"}, {"x-custom": "value"})
+ )
+
+
+def test_apply_fallback_hidden_params_to_item_no_existing_additional_headers():
+ class FakeChunk:
+ _hidden_params = {"model_id": "test-id"}
+
+ chunk = FakeChunk()
+ Router._apply_fallback_hidden_params_to_item(
+ chunk,
+ (
+ {"api_base": "http://fallback.example"},
+ {"x-litellm-attempted-fallbacks": 1},
+ ),
+ )
+
+ assert chunk._hidden_params["api_base"] == "http://fallback.example"
+ assert chunk._hidden_params["model_id"] == "test-id"
+ assert chunk._hidden_params["additional_headers"] == {
+ "x-litellm-attempted-fallbacks": 1
+ }
+
+
+@pytest.mark.asyncio
+async def test_set_response_headers_adds_model_group_to_streaming_wrapper():
+ class StreamingWrapper:
+ def __init__(self):
+ self._hidden_params = {"additional_headers": {"x-existing": "keep"}}
+
+ router = litellm.Router(model_list=[])
+ response = StreamingWrapper()
+
+ result = await router.set_response_headers(
+ response=response,
+ model_group="fallback-model",
+ )
+
+ assert result is response
+ assert response._hidden_params["additional_headers"] == {
+ "x-existing": "keep",
+ "x-litellm-model-group": "fallback-model",
+ }
+
+
+def test_should_include_fallback_errors_gated_by_operator_setting():
+ request_data: dict = {"include_fallback_errors": True}
+
+ import litellm.proxy.proxy_server as ps
+
+ original = ps.general_settings.copy() if isinstance(ps.general_settings, dict) else {}
+ try:
+ ps.general_settings = {}
+ assert _should_include_fallback_errors(request_data) is False
+
+ ps.general_settings = {"expose_fallback_errors_to_caller": False}
+ assert _should_include_fallback_errors(request_data) is False
+
+ ps.general_settings = {"expose_fallback_errors_to_caller": True}
+ assert _should_include_fallback_errors(request_data) is True
+
+ ps.general_settings = {"expose_fallback_errors_to_caller": True}
+ assert _should_include_fallback_errors({}) is False
+ finally:
+ ps.general_settings = original
diff --git a/tests/test_litellm/test_type_check_gate.py b/tests/test_litellm/test_type_check_gate.py
index 18374c5db4b..e99ad0a4f41 100644
--- a/tests/test_litellm/test_type_check_gate.py
+++ b/tests/test_litellm/test_type_check_gate.py
@@ -56,29 +56,58 @@ def test_paths_outside_repo_are_skipped():
def test_at_or_under_ceiling_passes():
budget = {"no-any-return": {"baseline": 5, "slack": 0}}
- assert gate.evaluate({"no-any-return": 5}, budget) == []
+ assert gate.evaluate({"no-any-return": 5}, {}, budget) == []
def test_one_more_error_than_ceiling_fails():
budget = {"no-any-return": {"baseline": 5, "slack": 0}}
- assert gate.evaluate({"no-any-return": 6}, budget) == [
- gate.Breach("no-any-return", 6, 5)
+ assert gate.evaluate({"no-any-return": 6}, {}, budget) == [
+ gate.Breach("no-any-return", 6, 5, 6)
]
def test_slack_absorbs_small_increase_then_fails_past_it():
budget = {"arg-type": {"baseline": 5, "slack": 5}}
- assert gate.evaluate({"arg-type": 10}, budget) == []
- assert gate.evaluate({"arg-type": 11}, budget) == [gate.Breach("arg-type", 11, 10)]
+ assert gate.evaluate({"arg-type": 10}, {}, budget) == []
+ assert gate.evaluate({"arg-type": 11}, {}, budget) == [
+ gate.Breach("arg-type", 11, 10, 11)
+ ]
def test_unbudgeted_new_code_uses_default_slack():
- assert gate.evaluate({"brand-new": gate.DEFAULT_SLACK}, {}) == []
- assert gate.evaluate({"brand-new": gate.DEFAULT_SLACK + 1}, {}) == [
- gate.Breach("brand-new", gate.DEFAULT_SLACK + 1, gate.DEFAULT_SLACK)
+ assert gate.evaluate({"brand-new": gate.DEFAULT_SLACK}, {}, {}) == []
+ assert gate.evaluate({"brand-new": gate.DEFAULT_SLACK + 1}, {}, {}) == [
+ gate.Breach(
+ "brand-new",
+ gate.DEFAULT_SLACK + 1,
+ gate.DEFAULT_SLACK,
+ gate.DEFAULT_SLACK + 1,
+ )
]
+def test_drift_already_over_cap_in_base_is_not_blamed_on_a_flat_change():
+ # The bystander case: a rule sits over its ceiling because two earlier PRs
+ # summed past it. A PR that branches off that base and adds nothing must pass
+ # -- total > cap but total == base, so the `> base` guard spares it.
+ budget = {"arg-type": {"baseline": 5, "slack": 5}}
+ assert gate.evaluate({"arg-type": 12}, {"arg-type": 12}, budget) == []
+
+
+def test_change_that_grows_an_over_cap_rule_is_blamed_for_only_what_it_added():
+ # Over cap AND above base: blamed, and `added` is the delta vs base, not the
+ # whole overage, so the message points at this change's contribution.
+ budget = {"arg-type": {"baseline": 5, "slack": 5}}
+ assert gate.evaluate({"arg-type": 14}, {"arg-type": 12}, budget) == [
+ gate.Breach("arg-type", 14, 10, 2)
+ ]
+
+
+def test_reducing_an_over_cap_rule_below_base_passes():
+ budget = {"arg-type": {"baseline": 5, "slack": 5}}
+ assert gate.evaluate({"arg-type": 11}, {"arg-type": 12}, budget) == []
+
+
def test_no_output_against_a_nonempty_budget_is_a_vacuous_run():
# A crashed type checker emits nothing; the gate must not certify it as clean.
budget = {"no-untyped-def": {"baseline": 4888, "slack": 10}}
diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py
index d94a86d8e55..cf7afc1af68 100644
--- a/tests/test_litellm/test_utils.py
+++ b/tests/test_litellm/test_utils.py
@@ -865,6 +865,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"supports_service_tier": {"type": "boolean"},
"supports_preset": {"type": "boolean"},
"supports_output_config": {"type": "boolean"},
+ "supports_speed": {"type": "boolean"},
"bedrock_output_config_effort_ceiling": {
"type": "string",
"enum": ["low", "medium", "high", "max", "xhigh"],
diff --git a/tests/test_litellm/types/test_completion.py b/tests/test_litellm/types/test_completion.py
index f24b00df3fc..cd51913c5dd 100644
--- a/tests/test_litellm/types/test_completion.py
+++ b/tests/test_litellm/types/test_completion.py
@@ -8,9 +8,16 @@ Usage:
pytest tests/test_litellm/types/test_completion.py -v
"""
+import dataclasses
from typing import List
-from litellm.types.completion import CompletionRequest, ChatCompletionMessageParam
+import pytest
+
+from litellm.types.completion import (
+ ChatCompletionMessageParam,
+ CompletionRequest,
+ _CompletionDispatchContext,
+)
def test_completion_request_messages_type_validation():
@@ -146,3 +153,55 @@ def test_completion_request_with_all_params():
assert request.presence_penalty == 0.0
assert request.stream is False
assert request.n == 1
+
+
+def _build_dispatch_context() -> _CompletionDispatchContext:
+ return _CompletionDispatchContext(
+ _azure_detection_model="gpt-4o",
+ acompletion=False,
+ api_base=None,
+ api_key=None,
+ api_version=None,
+ client=None,
+ custom_llm_provider="openai",
+ custom_prompt_dict={},
+ extra_headers=None,
+ headers={},
+ hf_model_name=None,
+ kwargs={},
+ litellm_params={},
+ logger_fn=None,
+ logging=None, # type: ignore[arg-type]
+ max_retries=None,
+ max_tokens=None,
+ messages=[],
+ metadata=None,
+ model="gpt-4o",
+ model_response=None, # type: ignore[arg-type]
+ optional_params={},
+ organization=None,
+ provider_config=None,
+ shared_session=None,
+ stream=None,
+ temperature=None,
+ text_completion=False,
+ timeout=None,
+ top_p=None,
+ )
+
+
+def test_dispatch_context_is_frozen():
+ """A helper must not be able to re-route the call by rebinding a dispatch
+ input mid-flight; this pins the frozen invariant the dispatch shape relies on."""
+ ctx = _build_dispatch_context()
+ with pytest.raises(dataclasses.FrozenInstanceError):
+ ctx.model = "claude-haiku-4-5" # type: ignore[misc]
+ with pytest.raises(dataclasses.FrozenInstanceError):
+ ctx.custom_llm_provider = "anthropic" # type: ignore[misc]
+
+
+def test_dispatch_context_uses_slots():
+ """slots=True keeps the per-call context lightweight (no per-instance __dict__)."""
+ ctx = _build_dispatch_context()
+ assert not hasattr(ctx, "__dict__")
+ assert hasattr(type(ctx), "__slots__")
diff --git a/ui/litellm-dashboard/package-lock.json b/ui/litellm-dashboard/package-lock.json
index 3beae6526e2..2afc145d15b 100644
--- a/ui/litellm-dashboard/package-lock.json
+++ b/ui/litellm-dashboard/package-lock.json
@@ -6882,16 +6882,16 @@
}
},
"node_modules/form-data": {
- "version": "4.0.5",
- "resolved": "https://registry.npmjs.org/form-data/-/form-data-4.0.5.tgz",
- "integrity": "sha512-8RipRLol37bNs2bhoV67fiTEvdTrbMUYcFTiy3+wuuOnUog2QBHCZWXDRijWQfAkhBj2Uf5UnVaiWwA5vdd82w==",
+ "version": "4.0.6",
+ "resolved": "https://registry.npmjs.org/form-data/-/form-data-4.0.6.tgz",
+ "integrity": "sha512-vKatAh4SlVfgbv+YtmhiRjhEMJsYpsG1Y2rMQtR+SVSbytsSD1YGzDIcrAJmdFec88u/+VoGmxnl+80gL1tRCQ==",
"license": "MIT",
"dependencies": {
"asynckit": "^0.4.0",
"combined-stream": "^1.0.8",
"es-set-tostringtag": "^2.1.0",
- "hasown": "^2.0.2",
- "mime-types": "^2.1.12"
+ "hasown": "^2.0.4",
+ "mime-types": "^2.1.35"
},
"engines": {
"node": ">= 6"
@@ -7248,9 +7248,9 @@
}
},
"node_modules/hasown": {
- "version": "2.0.3",
- "resolved": "https://registry.npmjs.org/hasown/-/hasown-2.0.3.tgz",
- "integrity": "sha512-ej4AhfhfL2Q2zpMmLo7U1Uv9+PyhIZpgQLGT1F9miIGmiCJIoCgSmczFdrc97mWT4kVY72KA+WnnhJ5pghSvSg==",
+ "version": "2.0.4",
+ "resolved": "https://registry.npmjs.org/hasown/-/hasown-2.0.4.tgz",
+ "integrity": "sha512-T2UbfbBEF32wiepXIsMlTW9+dDYC6wMh/t/vYA4tuOMKqWz/n3vr1NFSxQiyP+zk2mXsoMA/i/7qV6LKut1t1A==",
"license": "MIT",
"dependencies": {
"function-bind": "^1.1.2"
@@ -8149,10 +8149,20 @@
"license": "MIT"
},
"node_modules/js-yaml": {
- "version": "4.1.1",
- "resolved": "https://registry.npmjs.org/js-yaml/-/js-yaml-4.1.1.tgz",
- "integrity": "sha512-qQKT4zQxXl8lLwBtHMWwaTcGfFOZviOJet3Oy/xmGk2gZH677CJM9EvtfdSkgWcATZhj/55JZ0rmy3myCT5lsA==",
+ "version": "4.2.0",
+ "resolved": "https://registry.npmjs.org/js-yaml/-/js-yaml-4.2.0.tgz",
+ "integrity": "sha512-ePWsvanv0DWuDRsW8dnt+R4jQ31SCRCQ7hhNcPXZPsoBZiemuZNYGf7adZdqX2D86j6rvKp3RpCxVTSb8WQlOw==",
"dev": true,
+ "funding": [
+ {
+ "type": "github",
+ "url": "https://github.com/sponsors/puzrin"
+ },
+ {
+ "type": "github",
+ "url": "https://github.com/sponsors/nodeca"
+ }
+ ],
"license": "MIT",
"dependencies": {
"argparse": "^2.0.1"
@@ -13726,9 +13736,9 @@
}
},
"node_modules/ws": {
- "version": "8.20.1",
- "resolved": "https://registry.npmjs.org/ws/-/ws-8.20.1.tgz",
- "integrity": "sha512-It4dO0K5v//JtTXuPkfEOaI3uUN87iYPnqo/ZzqCoG3g8uhA66QUMs/SrM0YK7/NAu+r4LMh/9dq2A7k+rHs+w==",
+ "version": "8.21.0",
+ "resolved": "https://registry.npmjs.org/ws/-/ws-8.21.0.tgz",
+ "integrity": "sha512-Vsp28b7DRcimFQvrqu2Wek3z1iYxDCWqHYB8Qsnk/S4RfaCQzPGPyBNuVjJV3cd6UiKtUtp6sNM77gWvzcCH+g==",
"devOptional": true,
"license": "MIT",
"engines": {
diff --git a/ui/litellm-dashboard/package.json b/ui/litellm-dashboard/package.json
index c0899be8639..23da0bcd636 100644
--- a/ui/litellm-dashboard/package.json
+++ b/ui/litellm-dashboard/package.json
@@ -83,11 +83,11 @@
},
"overrides": {
"prismjs": "1.30.0",
- "js-yaml": "4.1.1",
+ "js-yaml": "4.2.0",
"glob": "13.0.0",
"minimatch": "10.2.4",
"lodash": "4.18.1",
- "ws": "8.20.1",
+ "ws": "8.21.0",
"braces": "3.0.3",
"axios": "1.13.6",
"postcss": "8.5.13",
diff --git a/ui/litellm-dashboard/src/components/OldTeams.test.tsx b/ui/litellm-dashboard/src/components/OldTeams.test.tsx
index d777ba1b0dc..0b6e5786aaf 100644
--- a/ui/litellm-dashboard/src/components/OldTeams.test.tsx
+++ b/ui/litellm-dashboard/src/components/OldTeams.test.tsx
@@ -1097,3 +1097,52 @@ describe("OldTeams - delete team warning copy", () => {
);
});
});
+
+describe("OldTeams - LIT-2530 organization stays optional for proxy admin with a single org", () => {
+ beforeEach(() => {
+ vi.clearAllMocks();
+ mockTeamInfoView.mockClear();
+ vi.mocked(fetchAvailableModelsForTeamOrKey).mockResolvedValue(["gpt-4"]);
+ vi.mocked(fetchMCPAccessGroups).mockResolvedValue([]);
+ vi.mocked(getGuardrailsList).mockResolvedValue({ guardrails: [] });
+ vi.mocked(teamListCall).mockResolvedValue({ teams: [], total: 0, page: 1, page_size: 100, total_pages: 1 });
+ vi.mocked(teamCreateCall).mockResolvedValue({
+ team_id: "new-team-1",
+ team_alias: "No Org Team",
+ models: ["gpt-4"],
+ organization_id: null,
+ keys: [],
+ members_with_roles: [],
+ spend: 0,
+ });
+ mockUseOrganizations.mockReturnValue({
+ data: [{ organization_id: "org-1", organization_alias: "Org 1", models: [], members: [] }],
+ });
+ });
+
+ it("creates a team with no organization when exactly one organization exists", async () => {
+ renderWithQueryClient();
+
+ const createButton = screen.getAllByRole("button", { name: /create team/i })[0];
+ act(() => {
+ fireEvent.click(createButton);
+ });
+
+ await waitFor(() => {
+ expect(screen.getByLabelText(/team name/i)).toBeInTheDocument();
+ });
+
+ fireEvent.change(screen.getByLabelText(/team name/i), { target: { value: "No Org Team" } });
+ fireEvent.change(screen.getByTestId("create-team-models-select"), { target: { value: "gpt-4" } });
+
+ const submitButtons = screen.getAllByRole("button", { name: /create team/i });
+ fireEvent.click(submitButtons[submitButtons.length - 1]);
+
+ await waitFor(() => {
+ expect(teamCreateCall).toHaveBeenCalledWith(
+ "test-token",
+ expect.objectContaining({ team_alias: "No Org Team", organization_id: null }),
+ );
+ });
+ });
+});
diff --git a/ui/litellm-dashboard/src/components/OldTeams.tsx b/ui/litellm-dashboard/src/components/OldTeams.tsx
index adfec4bdf6a..be9015e3730 100644
--- a/ui/litellm-dashboard/src/components/OldTeams.tsx
+++ b/ui/litellm-dashboard/src/components/OldTeams.tsx
@@ -262,14 +262,15 @@ const Teams: React.FC = ({ accessToken, userID, userRole, premiumUser
useEffect(() => {
if (isTeamModalVisible) {
const adminOrgs = getAdminOrganizations(userRole, userID, organizations);
+ const isOrgAdmin = userRole !== "Admin";
- // If there's exactly one organization the user is admin for, preselect it
- if (adminOrgs.length === 1) {
+ // Org admins must scope a team to an org, so with exactly one we preselect it.
+ // Proxy admins can create org-less teams, so the field stays optional regardless of org count.
+ if (isOrgAdmin && adminOrgs.length === 1) {
const org = adminOrgs[0];
form.setFieldValue("organization_id", org.organization_id);
setCurrentOrgForCreateTeam(org);
} else {
- // Reset the organization selection for multiple orgs
form.setFieldValue("organization_id", currentOrg?.organization_id || null);
setCurrentOrgForCreateTeam(currentOrg);
}
@@ -1132,7 +1133,7 @@ const Teams: React.FC = ({ accessToken, userID, userRole, premiumUser
: []
}
help={
- isSingleOrg
+ isOrgAdmin && isSingleOrg
? "You can only create teams within this organization"
: isOrgAdmin
? "required"
@@ -1142,7 +1143,7 @@ const Teams: React.FC = ({ accessToken, userID, userRole, premiumUser