Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_fix_budget_window_delete

This commit is contained in:
ryan-crabbe-berri 2026-06-23 13:22:51 -07:00
commit 62eb37d83d
74 changed files with 6330 additions and 434 deletions

BIN
.github/deploy-on-aws.png vendored Normal file

Binary file not shown.

After

Width:  |  Height:  |  Size: 4 KiB

BIN
.github/deploy-on-gcp.png vendored Normal file

Binary file not shown.

After

Width:  |  Height:  |  Size: 4.8 KiB

View file

@ -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: |

65
.github/workflows/test-rust.yml vendored Normal file
View file

@ -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

View file

@ -32,6 +32,7 @@ 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

View file

@ -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

142
README.md
View file

@ -6,10 +6,10 @@
</p>
<p align="center">Open Source AI Gateway for 100+ LLMs. Self-hosted. Enterprise-ready. Call any LLM in OpenAI format.</p>
<p align="center">
<a href="https://render.com/deploy?repo=https://github.com/BerriAI/litellm" target="_blank" rel="nofollow"><img src="https://render.com/images/deploy-to-render-button.svg" alt="Deploy to Render"></a>
<a href="https://railway.com/deploy/RhvhdC?referralCode=7mRv9K&utm_medium=integration&utm_source=template&utm_campaign=generic">
<img src="https://railway.com/button.svg" alt="Deploy on Railway">
</a>
<a href="https://render.com/deploy?repo=https://github.com/BerriAI/litellm" target="_blank" rel="nofollow"><img src="https://render.com/images/deploy-to-render-button.svg" alt="Deploy to Render" height="40"></a>
<a href="https://railway.com/deploy/RhvhdC?referralCode=7mRv9K&utm_medium=integration&utm_source=template&utm_campaign=generic"><img src="https://railway.com/button.svg" alt="Deploy on Railway" height="40"></a>
<a href="https://console.aws.amazon.com/cloudshell/home" target="_blank" rel="nofollow"><img src="./.github/deploy-on-aws.png" alt="Deploy on AWS" height="40"></a>
<a href="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" target="_blank" rel="nofollow"><img src="./.github/deploy-on-gcp.png" alt="Deploy on GCP" height="40"></a>
</p>
</p>
<h4 align="center"><a href="https://docs.litellm.ai/docs/simple_proxy" target="_blank">LiteLLM Proxy Server (AI Gateway)</a> | <a href="https://docs.litellm.ai/docs/enterprise#hosted-litellm-proxy" target="_blank"> Hosted Proxy</a> | <a href="https://litellm.ai/enterprise"target="_blank">Enterprise Tier</a> | <a href="https://www.litellm.ai/ai-gateway" target="_blank">Website</a></h4>
@ -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
[![Launch in AWS CloudShell](https://img.shields.io/badge/Launch-AWS_CloudShell-FF9900?logo=amazon-aws&logoColor=white)](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
[![Open in Cloud Shell](https://gstatic.com/cloudssh/images/open-btn.png)](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

View file

@ -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,

1
litellm-rust/.gitignore vendored Normal file
View file

@ -0,0 +1 @@
/target/

88
litellm-rust/CLAUDE.md Normal file
View file

@ -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/<route>/` owns the route contract, shared types, and provider
template traits. For OCR, this means `core/src/ocr`.
- `providers/src/<provider>/<route>/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.

1498
litellm-rust/Cargo.lock generated Normal file

File diff suppressed because it is too large Load diff

21
litellm-rust/Cargo.toml Normal file
View file

@ -0,0 +1,21 @@
[workspace]
members = [
"crates/core",
"crates/providers",
"crates/python-bridge",
]
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" }
pyo3 = "0.23.5"
reqwest = { version = "0.12", default-features = false, features = ["blocking", "json", "rustls-tls"] }
serde = { version = "1.0", features = ["derive"] }
serde_json = "1.0"
thiserror = "2.0"

34
litellm-rust/README.md Normal file
View file

@ -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/<provider>/<route>/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
```

View file

@ -0,0 +1,38 @@
# 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.
## 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.

View file

@ -0,0 +1,11 @@
[package]
name = "litellm-core"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
[dependencies]
serde.workspace = true
serde_json.workspace = true
thiserror.workspace = true

View file

@ -0,0 +1,33 @@
use thiserror::Error;
pub type CoreResult<T> = Result<T, CoreError>;
#[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),
}
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",
}
}

View file

@ -0,0 +1,4 @@
pub mod error;
pub mod ocr;
pub use error::{CoreError, CoreResult};

View file

@ -0,0 +1,2 @@
pub mod transformation;
pub mod types;

View file

@ -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<String, Value>) -> Map<String, Value> {
let mut mapped_params = Map::new();
for (param, value) in non_default_params {
if self.supported_ocr_params().contains(&param.as_str()) {
mapped_params.insert(param.clone(), value.clone());
}
}
mapped_params
}
fn transform_ocr_request(
&self,
model: &str,
document: Value,
optional_params: Map<String, Value>,
) -> CoreResult<OcrRequestData>;
fn transform_ocr_response(
&self,
model: &str,
response_json: Value,
) -> CoreResult<OcrResponseData>;
}

View file

@ -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<Value>,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct OcrResponseData {
pub pages: Vec<Value>,
pub model: String,
pub document_annotation: Option<Value>,
pub usage_info: Option<Value>,
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,
})
}
}

View file

@ -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/<provider>/<route>/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.

View file

@ -0,0 +1,14 @@
[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
[dev-dependencies]
serde_json.workspace = true

View file

@ -0,0 +1,2 @@
pub mod mistral;
pub mod ocr;

View file

@ -0,0 +1 @@
pub mod ocr;

View file

@ -0,0 +1 @@
pub mod transformation;

View file

@ -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<String>,
) -> CoreResult<String> {
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<String, Value>,
) -> CoreResult<OcrRequestData> {
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<OcrResponseData> {
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<String, Value>) -> Map<String, Value> {
MISTRAL_OCR_CONFIG.map_ocr_params(non_default_params)
}
pub fn transform_ocr_request(
model: &str,
document: Value,
optional_params: Map<String, Value>,
) -> CoreResult<OcrRequestData> {
MISTRAL_OCR_CONFIG.transform_ocr_request(model, document, optional_params)
}
pub fn transform_ocr_response(model: &str, response_json: Value) -> CoreResult<OcrResponseData> {
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()));
}
}

View file

@ -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<reqwest::blocking::Client> = 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<String, Value>,
timeout: Option<Duration>,
) -> CoreResult<Value> {
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()));
}
}

View file

@ -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.

View file

@ -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

View file

@ -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<T, F>(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)
}

View file

@ -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<Value> {
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<Py<PyAny>> {
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<PyAny>,
api_key: Option<String>,
api_base: Option<String>,
optional_params: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<Py<PyAny>> {
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<Py<PyAny>> {
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(())
}

View file

@ -1405,6 +1405,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 *

View file

@ -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 (

View file

@ -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

View file

@ -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."

View file

@ -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,

View file

@ -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(

View file

@ -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}")

View file

@ -1879,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)

View file

@ -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."

View file

@ -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."

View file

@ -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."

View file

@ -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."

View file

@ -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:

View file

@ -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."

View file

@ -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(

View file

@ -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."

View file

@ -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."

View file

@ -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"

View file

@ -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."

View file

@ -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."

View file

@ -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."

View file

@ -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

View file

@ -81,6 +81,9 @@ 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,
@ -654,6 +657,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
@ -4372,13 +4408,7 @@ def _complete_cloudflare(ctx: _CompletionDispatchContext) -> _CompletionDispatch
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/"
)
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(

View file

@ -10684,6 +10684,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",

View file

@ -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,

View file

@ -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)

View file

@ -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):
"""

View file

@ -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",

View file

@ -10684,6 +10684,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",

View file

@ -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 = "<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__":

View file

@ -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()

View file

@ -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 = {}

View file

@ -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"

View file

@ -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

View file

@ -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()

View file

@ -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({})

View file

@ -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"]

View file

@ -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"

View file

@ -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)

View file

@ -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()

View file

@ -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)

View file

@ -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

View file

@ -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}}