mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_fix_budget_window_delete
This commit is contained in:
commit
62eb37d83d
74 changed files with 6330 additions and 434 deletions
BIN
.github/deploy-on-aws.png
vendored
Normal file
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
BIN
.github/deploy-on-gcp.png
vendored
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 4.8 KiB |
8
.github/workflows/test-linting.yml
vendored
8
.github/workflows/test-linting.yml
vendored
|
|
@ -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
65
.github/workflows/test-rust.yml
vendored
Normal 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
|
||||
1
.github/workflows/test-unit-misc.yml
vendored
1
.github/workflows/test-unit-misc.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
3
Makefile
3
Makefile
|
|
@ -125,7 +125,8 @@ lint-ruff-FULL-dev: install-dev
|
|||
else echo "No changed .py files to check."; fi
|
||||
|
||||
lint-basedpyright: install-dev
|
||||
($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py
|
||||
git fetch origin litellm_internal_staging
|
||||
($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py --base origin/litellm_internal_staging
|
||||
|
||||
lint-basedpyright-budget-update: install-dev
|
||||
($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py --update
|
||||
|
|
|
|||
142
README.md
142
README.md
|
|
@ -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
|
||||
|
||||
[](https://console.aws.amazon.com/cloudshell/home) — opens an in-browser shell, already authenticated to your AWS account. Once inside, run:
|
||||
|
||||
```bash
|
||||
git clone https://github.com/BerriAI/litellm.git
|
||||
cd litellm/terraform/litellm/aws/examples/default
|
||||
cp terraform.tfvars.example terraform.tfvars # edit region/tenant/env
|
||||
terraform init && terraform apply
|
||||
```
|
||||
|
||||
[Module page →](https://registry.terraform.io/modules/BerriAI/litellm/aws/latest)
|
||||
|
||||
Or call the module from your own root config:
|
||||
|
||||
```hcl
|
||||
# main.tf
|
||||
terraform {
|
||||
required_version = ">= 1.6.0"
|
||||
required_providers {
|
||||
aws = { source = "hashicorp/aws", version = "~> 5.60" }
|
||||
}
|
||||
}
|
||||
|
||||
provider "aws" {
|
||||
region = "us-west-2"
|
||||
}
|
||||
|
||||
module "litellm" {
|
||||
source = "BerriAI/litellm/aws"
|
||||
version = "~> 1.89"
|
||||
|
||||
region = "us-west-2"
|
||||
azs = ["us-west-2a", "us-west-2b"]
|
||||
tenant = "acme"
|
||||
env = "prod"
|
||||
|
||||
# Production: provide an ACM cert. Without one, set allow_plaintext_alb = true
|
||||
# (dev/trial only).
|
||||
# acm_certificate_arn = "arn:aws:acm:us-west-2:111122223333:certificate/..."
|
||||
allow_plaintext_alb = true
|
||||
}
|
||||
|
||||
output "litellm_url" {
|
||||
value = module.litellm.alb_dns_name
|
||||
}
|
||||
```
|
||||
|
||||
```bash
|
||||
terraform init
|
||||
terraform apply
|
||||
```
|
||||
|
||||
Provider API keys live in AWS Secrets Manager; reference ARNs via `gateway_extra_secrets`. Full input list and architecture diagram on the [registry page](https://registry.terraform.io/modules/BerriAI/litellm/aws/latest?tab=inputs).
|
||||
|
||||
#### GCP — Cloud Run + Cloud SQL + Memorystore + HTTPS LB
|
||||
|
||||
[](https://ssh.cloud.google.com/cloudshell/editor?cloudshell_git_repo=https%3A%2F%2Fgithub.com%2FBerriAI%2Flitellm&cloudshell_workspace=terraform%2Flitellm%2Fgcp%2Fexamples%2Fdefault&cloudshell_tutorial=TUTORIAL.md&cloudshell_image=gcr.io/ds-artifacts-cloudshell/deploystack_custom_image&shellonly=true)
|
||||
|
||||
Real 1-click. Opens Cloud Shell, clones this repo, and walks you through `terraform apply` via a built-in [DeployStack tutorial](./terraform/litellm/gcp/examples/default/TUTORIAL.md) — pick the project, the tutorial sets up the Artifact Registry remote repo, writes `terraform.tfvars` from your answers, and runs apply.
|
||||
|
||||
[Module page →](https://registry.terraform.io/modules/BerriAI/litellm/google/latest)
|
||||
|
||||
To call the module from your own config instead, Cloud Run can't pull from `ghcr.io` directly, so first set up a one-time Artifact Registry remote repo backed by GHCR:
|
||||
|
||||
```bash
|
||||
gcloud artifacts repositories create litellm \
|
||||
--location=us-central1 \
|
||||
--repository-format=docker \
|
||||
--mode=remote-repository \
|
||||
--remote-docker-repo=https://ghcr.io \
|
||||
--project=my-gcp-project
|
||||
```
|
||||
|
||||
Then:
|
||||
|
||||
```hcl
|
||||
# main.tf
|
||||
terraform {
|
||||
required_version = ">= 1.6.0"
|
||||
required_providers {
|
||||
google = { source = "hashicorp/google", version = "~> 6.10" }
|
||||
google-beta = { source = "hashicorp/google-beta", version = "~> 6.10" }
|
||||
}
|
||||
}
|
||||
|
||||
provider "google" { project = "my-gcp-project"; region = "us-central1" }
|
||||
provider "google-beta" { project = "my-gcp-project"; region = "us-central1" }
|
||||
|
||||
module "litellm" {
|
||||
source = "BerriAI/litellm/google"
|
||||
version = "~> 1.89"
|
||||
|
||||
project_id = "my-gcp-project"
|
||||
region = "us-central1"
|
||||
tenant = "acme"
|
||||
env = "prod"
|
||||
|
||||
# Replace my-gcp-project with your GCP project ID (same value as project_id above).
|
||||
image_registry = "us-central1-docker.pkg.dev/my-gcp-project/litellm/berriai"
|
||||
|
||||
# Production: provide DNS already pointing at the LB IP for Google-managed certs.
|
||||
# Without one, set allow_plaintext_lb = true (dev/trial only).
|
||||
# lb_domains = ["proxy.example.com"]
|
||||
allow_plaintext_lb = true
|
||||
}
|
||||
|
||||
output "litellm_url" {
|
||||
value = module.litellm.load_balancer_url
|
||||
}
|
||||
```
|
||||
|
||||
```bash
|
||||
terraform init
|
||||
terraform apply
|
||||
```
|
||||
|
||||
Provider API keys live in Secret Manager; reference resource IDs (e.g. `projects/my-gcp-project/secrets/openai-api-key`) via `gateway_extra_secrets`. Full input list and architecture diagram on the [registry page](https://registry.terraform.io/modules/BerriAI/litellm/google/latest?tab=inputs).
|
||||
|
||||
#### Both stacks include
|
||||
|
||||
- The full componentized split (gateway / backend / UI as independent services)
|
||||
- Managed Postgres (writer + reader) and Redis
|
||||
- Versioned object store for proxy state + file uploads
|
||||
- An auto-generated `LITELLM_MASTER_KEY` in your cloud's secret manager
|
||||
- A one-off migration job that runs `prisma migrate deploy` before the proxy starts
|
||||
- The same `proxy_config` surface as the [Helm chart](./helm/litellm/) — pass YAML as a typed map
|
||||
|
||||
The Terraform modules live at [`terraform/litellm/aws/`](./terraform/litellm/aws/) and [`terraform/litellm/gcp/`](./terraform/litellm/gcp/) in this repo; the registry entries are read-only mirrors updated on each release.
|
||||
|
||||
### Run in Developer Mode
|
||||
#### Services
|
||||
1. Setup .env file in root
|
||||
|
|
|
|||
|
|
@ -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
1
litellm-rust/.gitignore
vendored
Normal file
|
|
@ -0,0 +1 @@
|
|||
/target/
|
||||
88
litellm-rust/CLAUDE.md
Normal file
88
litellm-rust/CLAUDE.md
Normal 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
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
21
litellm-rust/Cargo.toml
Normal 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
34
litellm-rust/README.md
Normal 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
|
||||
```
|
||||
38
litellm-rust/crates/core/CLAUDE.md
Normal file
38
litellm-rust/crates/core/CLAUDE.md
Normal 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.
|
||||
11
litellm-rust/crates/core/Cargo.toml
Normal file
11
litellm-rust/crates/core/Cargo.toml
Normal 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
|
||||
33
litellm-rust/crates/core/src/error.rs
Normal file
33
litellm-rust/crates/core/src/error.rs
Normal 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",
|
||||
}
|
||||
}
|
||||
4
litellm-rust/crates/core/src/lib.rs
Normal file
4
litellm-rust/crates/core/src/lib.rs
Normal file
|
|
@ -0,0 +1,4 @@
|
|||
pub mod error;
|
||||
pub mod ocr;
|
||||
|
||||
pub use error::{CoreError, CoreResult};
|
||||
2
litellm-rust/crates/core/src/ocr/mod.rs
Normal file
2
litellm-rust/crates/core/src/ocr/mod.rs
Normal file
|
|
@ -0,0 +1,2 @@
|
|||
pub mod transformation;
|
||||
pub mod types;
|
||||
32
litellm-rust/crates/core/src/ocr/transformation.rs
Normal file
32
litellm-rust/crates/core/src/ocr/transformation.rs
Normal 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(¶m.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>;
|
||||
}
|
||||
29
litellm-rust/crates/core/src/ocr/types.rs
Normal file
29
litellm-rust/crates/core/src/ocr/types.rs
Normal 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,
|
||||
})
|
||||
}
|
||||
}
|
||||
53
litellm-rust/crates/providers/CLAUDE.md
Normal file
53
litellm-rust/crates/providers/CLAUDE.md
Normal 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.
|
||||
14
litellm-rust/crates/providers/Cargo.toml
Normal file
14
litellm-rust/crates/providers/Cargo.toml
Normal 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
|
||||
2
litellm-rust/crates/providers/src/lib.rs
Normal file
2
litellm-rust/crates/providers/src/lib.rs
Normal file
|
|
@ -0,0 +1,2 @@
|
|||
pub mod mistral;
|
||||
pub mod ocr;
|
||||
1
litellm-rust/crates/providers/src/mistral/mod.rs
Normal file
1
litellm-rust/crates/providers/src/mistral/mod.rs
Normal file
|
|
@ -0,0 +1 @@
|
|||
pub mod ocr;
|
||||
1
litellm-rust/crates/providers/src/mistral/ocr/mod.rs
Normal file
1
litellm-rust/crates/providers/src/mistral/ocr/mod.rs
Normal file
|
|
@ -0,0 +1 @@
|
|||
pub mod transformation;
|
||||
292
litellm-rust/crates/providers/src/mistral/ocr/transformation.rs
Normal file
292
litellm-rust/crates/providers/src/mistral/ocr/transformation.rs
Normal 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()));
|
||||
}
|
||||
}
|
||||
127
litellm-rust/crates/providers/src/ocr.rs
Normal file
127
litellm-rust/crates/providers/src/ocr.rs
Normal 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()));
|
||||
}
|
||||
}
|
||||
36
litellm-rust/crates/python-bridge/CLAUDE.md
Normal file
36
litellm-rust/crates/python-bridge/CLAUDE.md
Normal 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.
|
||||
16
litellm-rust/crates/python-bridge/Cargo.toml
Normal file
16
litellm-rust/crates/python-bridge/Cargo.toml
Normal 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
|
||||
32
litellm-rust/crates/python-bridge/src/gil.rs
Normal file
32
litellm-rust/crates/python-bridge/src/gil.rs
Normal 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)
|
||||
}
|
||||
100
litellm-rust/crates/python-bridge/src/lib.rs
Normal file
100
litellm-rust/crates/python-bridge/src/lib.rs
Normal 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(())
|
||||
}
|
||||
|
|
@ -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 *
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
332
litellm/litellm_core_utils/chat_completion_agentic_loop.py
Normal file
332
litellm/litellm_core_utils/chat_completion_agentic_loop.py
Normal 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
|
||||
|
|
@ -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."
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
74
litellm/ocr/rust_bridge.py
Normal file
74
litellm/ocr/rust_bridge.py
Normal 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)
|
||||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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__":
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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 = {}
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
@ -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({})
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
333
tests/test_litellm/ocr/test_rust_bridge.py
Normal file
333
tests/test_litellm/ocr/test_rust_bridge.py
Normal 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)
|
||||
|
|
@ -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
|
||||
|
|
@ -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}}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue