diff --git a/.github/deploy-on-aws.png b/.github/deploy-on-aws.png new file mode 100644 index 00000000000..06d41f2a5e0 Binary files /dev/null and b/.github/deploy-on-aws.png differ diff --git a/.github/deploy-on-gcp.png b/.github/deploy-on-gcp.png new file mode 100644 index 00000000000..e831a8c2e4e Binary files /dev/null and b/.github/deploy-on-gcp.png differ diff --git a/.github/workflows/test-linting.yml b/.github/workflows/test-linting.yml index de7e1b68346..950d6ca31a6 100644 --- a/.github/workflows/test-linting.yml +++ b/.github/workflows/test-linting.yml @@ -14,7 +14,7 @@ permissions: jobs: lint: runs-on: ubuntu-latest - timeout-minutes: 10 + timeout-minutes: 15 steps: - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 @@ -87,9 +87,11 @@ jobs: run: | uv run --no-sync python -c "import openai; print(f'OpenAI version: {openai.__version__}')" - - name: Run basedpyright type checking + - name: Check basedpyright budget (delta vs base) + env: + BASE_SHA: ${{ github.event.pull_request.base.sha }} run: | - (uv run --no-sync basedpyright --outputjson || true) | uv run --no-sync python scripts/type_check_gate.py + (uv run --no-sync basedpyright --outputjson || true) | uv run --no-sync python scripts/type_check_gate.py --base "$BASE_SHA" - name: Check for circular imports run: | diff --git a/.github/workflows/test-rust.yml b/.github/workflows/test-rust.yml new file mode 100644 index 00000000000..3d0a159cdc7 --- /dev/null +++ b/.github/workflows/test-rust.yml @@ -0,0 +1,65 @@ +name: LiteLLM Rust + +on: + push: + paths: + - "litellm-rust/**" + - ".github/workflows/test-rust.yml" + pull_request: + branches: + - main + - litellm_internal_staging + - litellm_oss_branch + - "litellm_**" + paths: + - "litellm-rust/**" + - ".github/workflows/test-rust.yml" + +permissions: + contents: read + +concurrency: + group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }} + cancel-in-progress: true + +jobs: + rust-checks: + name: rustfmt, clippy, test + runs-on: ubuntu-latest + timeout-minutes: 10 + defaults: + run: + working-directory: litellm-rust + env: + CARGO_TERM_COLOR: always + + steps: + - name: Checkout repository + uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 + with: + persist-credentials: false + + - name: Set up Rust + run: | + rustup toolchain install stable --profile minimal --component clippy,rustfmt + rustup default stable + + - name: Cache Cargo registry and target + uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0 + with: + path: | + ~/.cargo/registry + ~/.cargo/git + litellm-rust/target + key: ${{ runner.os }}-cargo-${{ hashFiles('litellm-rust/Cargo.lock') }} + restore-keys: | + ${{ runner.os }}-cargo- + + - name: Check Rust formatting + run: cargo fmt --check + + - name: Run Clippy + run: cargo clippy --workspace --all-targets --locked -- -D warnings + + - name: Run Rust tests + run: cargo test --workspace --locked diff --git a/.github/workflows/test-unit-misc.yml b/.github/workflows/test-unit-misc.yml index c29c2d632f2..2226d519331 100644 --- a/.github/workflows/test-unit-misc.yml +++ b/.github/workflows/test-unit-misc.yml @@ -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 diff --git a/Makefile b/Makefile index 27150aec938..076eac0f4a7 100644 --- a/Makefile +++ b/Makefile @@ -125,7 +125,8 @@ lint-ruff-FULL-dev: install-dev else echo "No changed .py files to check."; fi lint-basedpyright: install-dev - ($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py + git fetch origin litellm_internal_staging + ($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py --base origin/litellm_internal_staging lint-basedpyright-budget-update: install-dev ($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py --update diff --git a/README.md b/README.md index b26ad39eada..3d0f7282d7c 100644 --- a/README.md +++ b/README.md @@ -6,10 +6,10 @@

Open Source AI Gateway for 100+ LLMs. Self-hosted. Enterprise-ready. Call any LLM in OpenAI format.

- Deploy to Render - - Deploy on Railway - + Deploy to Render + Deploy on Railway + Deploy on AWS + Deploy on GCP

LiteLLM Proxy Server (AI Gateway) | Hosted Proxy | Enterprise Tier | Website

@@ -406,6 +406,140 @@ You can use LiteLLM through either the Proxy Server or Python SDK. Both give you Support for more providers. Missing a provider or LLM Platform, raise a [feature request](https://github.com/BerriAI/litellm/issues/new?assignees=&labels=enhancement&projects=&template=feature_request.yml&title=%5BFeature%5D%3A+). +### Deploy on AWS or GCP with Terraform + +Run the LiteLLM proxy as a production-ready componentized stack (gateway, backend, UI on separate services; managed Postgres + Redis + object store) using the published Terraform modules. Both modules are on the [public Terraform Registry](https://registry.terraform.io/namespaces/BerriAI) — no auth needed. + +#### AWS — ECS Fargate + Aurora + ElastiCache + ALB + +[![Launch in AWS CloudShell](https://img.shields.io/badge/Launch-AWS_CloudShell-FF9900?logo=amazon-aws&logoColor=white)](https://console.aws.amazon.com/cloudshell/home) — opens an in-browser shell, already authenticated to your AWS account. Once inside, run: + +```bash +git clone https://github.com/BerriAI/litellm.git +cd litellm/terraform/litellm/aws/examples/default +cp terraform.tfvars.example terraform.tfvars # edit region/tenant/env +terraform init && terraform apply +``` + +[Module page →](https://registry.terraform.io/modules/BerriAI/litellm/aws/latest) + +Or call the module from your own root config: + +```hcl +# main.tf +terraform { + required_version = ">= 1.6.0" + required_providers { + aws = { source = "hashicorp/aws", version = "~> 5.60" } + } +} + +provider "aws" { + region = "us-west-2" +} + +module "litellm" { + source = "BerriAI/litellm/aws" + version = "~> 1.89" + + region = "us-west-2" + azs = ["us-west-2a", "us-west-2b"] + tenant = "acme" + env = "prod" + + # Production: provide an ACM cert. Without one, set allow_plaintext_alb = true + # (dev/trial only). + # acm_certificate_arn = "arn:aws:acm:us-west-2:111122223333:certificate/..." + allow_plaintext_alb = true +} + +output "litellm_url" { + value = module.litellm.alb_dns_name +} +``` + +```bash +terraform init +terraform apply +``` + +Provider API keys live in AWS Secrets Manager; reference ARNs via `gateway_extra_secrets`. Full input list and architecture diagram on the [registry page](https://registry.terraform.io/modules/BerriAI/litellm/aws/latest?tab=inputs). + +#### GCP — Cloud Run + Cloud SQL + Memorystore + HTTPS LB + +[![Open in Cloud Shell](https://gstatic.com/cloudssh/images/open-btn.png)](https://ssh.cloud.google.com/cloudshell/editor?cloudshell_git_repo=https%3A%2F%2Fgithub.com%2FBerriAI%2Flitellm&cloudshell_workspace=terraform%2Flitellm%2Fgcp%2Fexamples%2Fdefault&cloudshell_tutorial=TUTORIAL.md&cloudshell_image=gcr.io/ds-artifacts-cloudshell/deploystack_custom_image&shellonly=true) + +Real 1-click. Opens Cloud Shell, clones this repo, and walks you through `terraform apply` via a built-in [DeployStack tutorial](./terraform/litellm/gcp/examples/default/TUTORIAL.md) — pick the project, the tutorial sets up the Artifact Registry remote repo, writes `terraform.tfvars` from your answers, and runs apply. + +[Module page →](https://registry.terraform.io/modules/BerriAI/litellm/google/latest) + +To call the module from your own config instead, Cloud Run can't pull from `ghcr.io` directly, so first set up a one-time Artifact Registry remote repo backed by GHCR: + +```bash +gcloud artifacts repositories create litellm \ + --location=us-central1 \ + --repository-format=docker \ + --mode=remote-repository \ + --remote-docker-repo=https://ghcr.io \ + --project=my-gcp-project +``` + +Then: + +```hcl +# main.tf +terraform { + required_version = ">= 1.6.0" + required_providers { + google = { source = "hashicorp/google", version = "~> 6.10" } + google-beta = { source = "hashicorp/google-beta", version = "~> 6.10" } + } +} + +provider "google" { project = "my-gcp-project"; region = "us-central1" } +provider "google-beta" { project = "my-gcp-project"; region = "us-central1" } + +module "litellm" { + source = "BerriAI/litellm/google" + version = "~> 1.89" + + project_id = "my-gcp-project" + region = "us-central1" + tenant = "acme" + env = "prod" + + # Replace my-gcp-project with your GCP project ID (same value as project_id above). + image_registry = "us-central1-docker.pkg.dev/my-gcp-project/litellm/berriai" + + # Production: provide DNS already pointing at the LB IP for Google-managed certs. + # Without one, set allow_plaintext_lb = true (dev/trial only). + # lb_domains = ["proxy.example.com"] + allow_plaintext_lb = true +} + +output "litellm_url" { + value = module.litellm.load_balancer_url +} +``` + +```bash +terraform init +terraform apply +``` + +Provider API keys live in Secret Manager; reference resource IDs (e.g. `projects/my-gcp-project/secrets/openai-api-key`) via `gateway_extra_secrets`. Full input list and architecture diagram on the [registry page](https://registry.terraform.io/modules/BerriAI/litellm/google/latest?tab=inputs). + +#### Both stacks include + +- The full componentized split (gateway / backend / UI as independent services) +- Managed Postgres (writer + reader) and Redis +- Versioned object store for proxy state + file uploads +- An auto-generated `LITELLM_MASTER_KEY` in your cloud's secret manager +- A one-off migration job that runs `prisma migrate deploy` before the proxy starts +- The same `proxy_config` surface as the [Helm chart](./helm/litellm/) — pass YAML as a typed map + +The Terraform modules live at [`terraform/litellm/aws/`](./terraform/litellm/aws/) and [`terraform/litellm/gcp/`](./terraform/litellm/gcp/) in this repo; the registry entries are read-only mirrors updated on each release. + ### Run in Developer Mode #### Services 1. Setup .env file in root diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 7ba7656e407..f5b0a9aaf81 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -121,7 +121,7 @@ }, "reportReturnType": { "baseline": 126, - "slack": 13 + "slack": 100 }, "reportTypedDictNotRequiredAccess": { "baseline": 20, @@ -157,7 +157,7 @@ }, "reportUnnecessaryComparison": { "baseline": 683, - "slack": 10 + "slack": 100 }, "reportUnnecessaryContains": { "baseline": 4, diff --git a/litellm-rust/.gitignore b/litellm-rust/.gitignore new file mode 100644 index 00000000000..b83d22266ac --- /dev/null +++ b/litellm-rust/.gitignore @@ -0,0 +1 @@ +/target/ diff --git a/litellm-rust/CLAUDE.md b/litellm-rust/CLAUDE.md new file mode 100644 index 00000000000..1d2987e0a1a --- /dev/null +++ b/litellm-rust/CLAUDE.md @@ -0,0 +1,88 @@ +# CLAUDE.md + +This file defines the rules for Rust work in LiteLLM. + +## Core Boundary + +The `core` and `providers` crates describe work; hosts execute work. + +Route-level Rust structure mirrors LiteLLM's Python responsibilities: +- `core/src//` owns the route contract, shared types, and provider + template traits. For OCR, this means `core/src/ocr`. +- `providers/src///transformation.rs` owns the + provider-specific transform. For Mistral OCR, this means + `providers/src/mistral/ocr/transformation.rs`. +- Future network execution belongs in a host/transport layer such as + `llm_http_handler`, not inside `core` or `providers`. + +Allowed in `core` and `providers`: +- Pure request transforms +- Pure response transforms +- Pure stream chunk normalization +- Shared data types and validation errors +- Deterministic token/cost helper logic + +Not allowed in `core` or `providers`: +- Network calls +- Environment variable or secret reads +- Filesystem access +- Database or cache access +- Provider SDK signing or auth flows +- Logging callbacks, spend writes, or custom callbacks +- Global mutable runtime state + +Python owns rollout state and fallback while Rust is being introduced. Rust +paths must be off by default until parity tests prove equivalence with Python. + +## Production Bar + +Rust code in this workspace is held to a strict parity and robustness bar from +the first PR: + +- Correctness parity is proven with tests. Do not rely on README claims or + manual inspection for a port that mirrors Python behavior. +- Every provider transform must have unit tests for supported-parameter + filtering, request body shape, response normalization, missing/null fields, + and bad-input errors. +- When Rust is exposed through Python, add Python tests that prove disabled, + enabled, and unavailable-bridge fallback behavior. +- Avoid panics on user/provider input. Return typed errors and let the host map + them to Python exceptions or HTTP responses. +- OCR handles documents that often contain personal data. Do not log document + contents, base64 payloads, provider response bodies, or secrets. +- Error messages must be useful but data-minimized. Truncate or sanitize any + upstream body before it crosses a host boundary. +- Treat empty or whitespace-only credentials, URLs, and config values as absent + at the host/config resolution layer. +- Preserve Python output shape intentionally. If a field is always serialized as + `null` for Python parity, leave a short comment explaining that parity choice. + +## Host I/O Rules + +These rules apply when adding future crates or modules that execute network I/O, +such as `ai-gateway`, router hosts, or standalone servers: + +- Set connect and full-request timeouts. No unbounded waits. +- Reuse HTTP clients; do not construct clients per request. +- Prefer rustls TLS for portable Python wheels and Linux images unless there is + a documented reason not to. +- Add request IDs and structured tracing at the host layer, without logging OCR + document contents or secrets. +- Do not echo raw upstream response bodies to callers. Sanitize and bound them. +- Avoid `expect`/`unwrap` in server startup and request paths unless the panic is + impossible by construction and documented. + +## Checks + +Run these before pushing Rust changes. The same checks run in GitHub Actions +for changes under `litellm-rust/`. + +```bash +cd litellm-rust +cargo fmt --check +cargo clippy --workspace --all-targets -- -D warnings +cargo test --workspace +``` + +When a Rust path is exposed through Python, add Python parity tests that compare +the existing Python output with the Rust-backed output. diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock new file mode 100644 index 00000000000..0fb2fcb2921 --- /dev/null +++ b/litellm-rust/Cargo.lock @@ -0,0 +1,1498 @@ +# This file is automatically @generated by Cargo. +# It is not intended for manual editing. +version = 4 + +[[package]] +name = "atomic-waker" +version = "1.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1505bd5d3d116872e7271a6d4e16d81d0c8570876c8de68093a09ac269d8aac0" + +[[package]] +name = "autocfg" +version = "1.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f2032f911046de80f0a198e0901378627c33f59ea0ac00e363d481118bd70a53" + +[[package]] +name = "base64" +version = "0.22.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "72b3254f16251a8381aa12e40e3c4d2f0199f8c6508fbecb9d91f575e0fbb8c6" + +[[package]] +name = "bitflags" +version = "2.13.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b4388bee8683e3d04af747c73422af53102d2bd24d9eadb6cbc100baef4b43f8" + +[[package]] +name = "bumpalo" +version = "3.20.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "72f5acc6cb2ba439de613abc23857ec3d78374d8ed5ac84e9d11336e87da8649" + +[[package]] +name = "bytes" +version = "1.12.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8ae3f5d315924270530207e2a68396c3cc547f6dca3fbdca317cfb1a51edb593" + +[[package]] +name = "cc" +version = "1.2.65" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e228eec9be7c17ccb640b59b36a5cd805ea2a564a4c5e162c2f659fea30d3b96" +dependencies = [ + "find-msvc-tools", + "shlex", +] + +[[package]] +name = "cfg-if" +version = "1.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9330f8b2ff13f34540b44e946ef35111825727b38d33286ef986142615121801" + +[[package]] +name = "cfg_aliases" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724" + +[[package]] +name = "displaydoc" +version = "0.2.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ac70aa55017e108007fbaf5aa0f54b021c98f92ff8af59d42eda9da96e3dd4f" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "find-msvc-tools" +version = "0.1.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5baebc0774151f905a1a2cc41989300b1e6fbb29aff0ceffa1064fdd3088d582" + +[[package]] +name = "form_urlencoded" +version = "1.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cb4cb245038516f5f85277875cdaa4f7d2c9a0fa0468de06ed190163b1581fcf" +dependencies = [ + "percent-encoding", +] + +[[package]] +name = "futures-channel" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "07bbe89c50d7a535e539b8c17bc0b49bdb77747034daa8087407d655f3f7cc1d" +dependencies = [ + "futures-core", + "futures-sink", +] + +[[package]] +name = "futures-core" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7e3450815272ef58cec6d564423f6e755e25379b217b0bc688e295ba24df6b1d" + +[[package]] +name = "futures-io" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cecba35d7ad927e23624b22ad55235f2239cfa44fd10428eecbeba6d6a717718" + +[[package]] +name = "futures-sink" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c39754e157331b013978ec91992bde1ac089843443c49cbc7f46150b0fad0893" + +[[package]] +name = "futures-task" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "037711b3d59c33004d3856fbdc83b99d4ff37a24768fa1be9ce3538a1cde4393" + +[[package]] +name = "futures-util" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "389ca41296e6190b48053de0321d02a77f32f8a5d2461dd38762c0593805c6d6" +dependencies = [ + "futures-core", + "futures-io", + "futures-sink", + "futures-task", + "memchr", + "pin-project-lite", + "slab", +] + +[[package]] +name = "getrandom" +version = "0.2.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ff2abc00be7fca6ebc474524697ae276ad847ad0a6b3faa4bcb027e9a4614ad0" +dependencies = [ + "cfg-if", + "js-sys", + "libc", + "wasi", + "wasm-bindgen", +] + +[[package]] +name = "getrandom" +version = "0.3.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "899def5c37c4fd7b2664648c28120ecec138e4d395b459e5ca34f9cce2dd77fd" +dependencies = [ + "cfg-if", + "js-sys", + "libc", + "r-efi", + "wasip2", + "wasm-bindgen", +] + +[[package]] +name = "heck" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" + +[[package]] +name = "http" +version = "1.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6970f50e31d6fc17d3fa27329444bfa74e196cf62e95052a3f6fee181dba6425" +dependencies = [ + "bytes", + "itoa", +] + +[[package]] +name = "http-body" +version = "1.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1efedce1fb8e6913f23e0c92de8e62cd5b772a67e7b3946df930a62566c93184" +dependencies = [ + "bytes", + "http", +] + +[[package]] +name = "http-body-util" +version = "0.1.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b021d93e26becf5dc7e1b75b1bed1fd93124b374ceb73f43d4d4eafec896a64a" +dependencies = [ + "bytes", + "futures-core", + "http", + "http-body", + "pin-project-lite", +] + +[[package]] +name = "httparse" +version = "1.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6dbf3de79e51f3d586ab4cb9d5c3e2c14aa28ed23d180cf89b4df0454a69cc87" + +[[package]] +name = "hyper" +version = "1.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "55281c53a1894c864990125767da440a4e630446785086f52523b20033b74498" +dependencies = [ + "atomic-waker", + "bytes", + "futures-channel", + "futures-core", + "http", + "http-body", + "httparse", + "itoa", + "pin-project-lite", + "smallvec", + "tokio", + "want", +] + +[[package]] +name = "hyper-rustls" +version = "0.27.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "33ca68d021ef39cf6463ab54c1d0f5daf03377b70561305bb89a8f83aab66e0f" +dependencies = [ + "http", + "hyper", + "hyper-util", + "rustls", + "tokio", + "tokio-rustls", + "tower-service", + "webpki-roots", +] + +[[package]] +name = "hyper-util" +version = "0.1.20" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "96547c2556ec9d12fb1578c4eaf448b04993e7fb79cbaad930a656880a6bdfa0" +dependencies = [ + "base64", + "bytes", + "futures-channel", + "futures-util", + "http", + "http-body", + "hyper", + "ipnet", + "libc", + "percent-encoding", + "pin-project-lite", + "socket2", + "tokio", + "tower-service", + "tracing", +] + +[[package]] +name = "icu_collections" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2984d1cd16c883d7935b9e07e44071dca8d917fd52ecc02c04d5fa0b5a3f191c" +dependencies = [ + "displaydoc", + "potential_utf", + "utf8_iter", + "yoke", + "zerofrom", + "zerovec", +] + +[[package]] +name = "icu_locale_core" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "92219b62b3e2b4d88ac5119f8904c10f8f61bf7e95b640d25ba3075e6cac2c29" +dependencies = [ + "displaydoc", + "litemap", + "tinystr", + "writeable", + "zerovec", +] + +[[package]] +name = "icu_normalizer" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c56e5ee99d6e3d33bd91c5d85458b6005a22140021cc324cea84dd0e72cff3b4" +dependencies = [ + "icu_collections", + "icu_normalizer_data", + "icu_properties", + "icu_provider", + "smallvec", + "zerovec", +] + +[[package]] +name = "icu_normalizer_data" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "da3be0ae77ea334f4da67c12f149704f19f81d1adf7c51cf482943e84a2bad38" + +[[package]] +name = "icu_properties" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bee3b67d0ea5c2cca5003417989af8996f8604e34fb9ddf96208a033901e70de" +dependencies = [ + "icu_collections", + "icu_locale_core", + "icu_properties_data", + "icu_provider", + "zerotrie", + "zerovec", +] + +[[package]] +name = "icu_properties_data" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8e2bbb201e0c04f7b4b3e14382af113e17ba4f63e2c9d2ee626b720cbce54a14" + +[[package]] +name = "icu_provider" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "139c4cf31c8b5f33d7e199446eff9c1e02decfc2f0eec2c8d71f65befa45b421" +dependencies = [ + "displaydoc", + "icu_locale_core", + "writeable", + "yoke", + "zerofrom", + "zerotrie", + "zerovec", +] + +[[package]] +name = "idna" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3b0875f23caa03898994f6ddc501886a45c7d3d62d04d2d90788d47be1b1e4de" +dependencies = [ + "idna_adapter", + "smallvec", + "utf8_iter", +] + +[[package]] +name = "idna_adapter" +version = "1.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cb68373c0d6620ef8105e855e7745e18b0d00d3bdb07fb532e434244cdb9a714" +dependencies = [ + "icu_normalizer", + "icu_properties", +] + +[[package]] +name = "indoc" +version = "2.0.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "79cf5c93f93228cf8efb3ba362535fb11199ac548a09ce117c9b1adc3030d706" +dependencies = [ + "rustversion", +] + +[[package]] +name = "ipnet" +version = "2.12.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d98f6fed1fde3f8c21bc40a1abb88dd75e67924f9cffc3ef95607bad8017f8e2" + +[[package]] +name = "itoa" +version = "1.0.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8f42a60cbdf9a97f5d2305f08a87dc4e09308d1276d28c869c684d7777685682" + +[[package]] +name = "js-sys" +version = "0.3.102" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "03d04c30968dffe80775bd4d7fb676131cd04a1fb46d2686dbffbaec2d9dfd31" +dependencies = [ + "cfg-if", + "futures-util", + "wasm-bindgen", +] + +[[package]] +name = "libc" +version = "0.2.186" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "68ab91017fe16c622486840e4c83c9a37afeff978bd239b5293d61ece587de66" + +[[package]] +name = "litellm-core" +version = "0.1.0" +dependencies = [ + "serde", + "serde_json", + "thiserror", +] + +[[package]] +name = "litellm-providers" +version = "0.1.0" +dependencies = [ + "litellm-core", + "reqwest", + "serde_json", +] + +[[package]] +name = "litellm-python-bridge" +version = "0.1.0" +dependencies = [ + "litellm-core", + "litellm-providers", + "pyo3", + "serde_json", +] + +[[package]] +name = "litemap" +version = "0.8.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "92daf443525c4cce67b150400bc2316076100ce0b3686209eb8cf3c31612e6f0" + +[[package]] +name = "log" +version = "0.4.33" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0ceec5bc11778974d1bcb055b18002eba7f4b3518b6a0081b3af5f21666da9ad" + +[[package]] +name = "lru-slab" +version = "0.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "112b39cec0b298b6c1999fee3e31427f74f676e4cb9879ed1a121b43661a4154" + +[[package]] +name = "memchr" +version = "2.8.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "88904434abc2901f197fe8cc55f0445e7ded921dba5911dad2e2b39b48e663c4" + +[[package]] +name = "memoffset" +version = "0.9.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "488016bfae457b036d996092f6cb448677611ce4449e970ceaf42695203f218a" +dependencies = [ + "autocfg", +] + +[[package]] +name = "mio" +version = "1.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "02bd0af71c67b473010cbbc60715ee815645a4dc942899111f494b4b737d6fda" +dependencies = [ + "libc", + "wasi", + "windows-sys 0.61.2", +] + +[[package]] +name = "once_cell" +version = "1.21.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50" + +[[package]] +name = "percent-encoding" +version = "2.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9b4f627cb1b25917193a259e49bdad08f671f8d9708acfd5fe0a8c1455d87220" + +[[package]] +name = "pin-project-lite" +version = "0.2.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a89322df9ebe1c1578d689c92318e070967d1042b512afbe49518723f4e6d5cd" + +[[package]] +name = "portable-atomic" +version = "1.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c33a9471896f1c69cecef8d20cbe2f7accd12527ce60845ff44c153bb2a21b49" + +[[package]] +name = "potential_utf" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0103b1cef7ec0cf76490e969665504990193874ea05c85ff9bab8b911d0a0564" +dependencies = [ + "zerovec", +] + +[[package]] +name = "ppv-lite86" +version = "0.2.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "85eae3c4ed2f50dcfe72643da4befc30deadb458a9b590d720cde2f2b1e97da9" +dependencies = [ + "zerocopy", +] + +[[package]] +name = "proc-macro2" +version = "1.0.106" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8fd00f0bb2e90d81d1044c2b32617f68fcb9fa3bb7640c23e9c748e53fb30934" +dependencies = [ + "unicode-ident", +] + +[[package]] +name = "pyo3" +version = "0.23.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7778bffd85cf38175ac1f545509665d0b9b92a198ca7941f131f85f7a4f9a872" +dependencies = [ + "cfg-if", + "indoc", + "libc", + "memoffset", + "once_cell", + "portable-atomic", + "pyo3-build-config", + "pyo3-ffi", + "pyo3-macros", + "unindent", +] + +[[package]] +name = "pyo3-build-config" +version = "0.23.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "94f6cbe86ef3bf18998d9df6e0f3fc1050a8c5efa409bf712e661a4366e010fb" +dependencies = [ + "once_cell", + "target-lexicon", +] + +[[package]] +name = "pyo3-ffi" +version = "0.23.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e9f1b4c431c0bb1c8fb0a338709859eed0d030ff6daa34368d3b152a63dfdd8d" +dependencies = [ + "libc", + "pyo3-build-config", +] + +[[package]] +name = "pyo3-macros" +version = "0.23.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fbc2201328f63c4710f68abdf653c89d8dbc2858b88c5d88b0ff38a75288a9da" +dependencies = [ + "proc-macro2", + "pyo3-macros-backend", + "quote", + "syn", +] + +[[package]] +name = "pyo3-macros-backend" +version = "0.23.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fca6726ad0f3da9c9de093d6f116a93c1a38e417ed73bf138472cf4064f72028" +dependencies = [ + "heck", + "proc-macro2", + "pyo3-build-config", + "quote", + "syn", +] + +[[package]] +name = "quinn" +version = "0.11.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0c1a41e437b6bbd489372cd4971de128e85c855f56c57f283d20ff016cf7c0a8" +dependencies = [ + "bytes", + "cfg_aliases", + "pin-project-lite", + "quinn-proto", + "quinn-udp", + "rustc-hash", + "rustls", + "socket2", + "thiserror", + "tokio", + "tracing", + "web-time", +] + +[[package]] +name = "quinn-proto" +version = "0.11.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4fcb935c5bec503c2f0e306bdd3e58bb9029dcb14fa8d9ac76e3a5256ac0763e" +dependencies = [ + "bytes", + "getrandom 0.3.4", + "lru-slab", + "rand", + "ring", + "rustc-hash", + "rustls", + "rustls-pki-types", + "slab", + "thiserror", + "tinyvec", + "tracing", + "web-time", +] + +[[package]] +name = "quinn-udp" +version = "0.5.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "addec6a0dcad8a8d96a771f815f0eaf55f9d1805756410b39f5fa81332574cbd" +dependencies = [ + "cfg_aliases", + "libc", + "once_cell", + "socket2", + "tracing", + "windows-sys 0.60.2", +] + +[[package]] +name = "quote" +version = "1.0.46" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dfbc457d0c7a0759a614551b11a6409e5951f6c7537be1f1b7682b9ae9230368" +dependencies = [ + "proc-macro2", +] + +[[package]] +name = "r-efi" +version = "5.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "69cdb34c158ceb288df11e18b4bd39de994f6657d83847bdffdbd7f346754b0f" + +[[package]] +name = "rand" +version = "0.9.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "44c5af06bb1b7d3216d91932aed5265164bf384dc89cd6ba05cf59a35f5f76ea" +dependencies = [ + "rand_chacha", + "rand_core", +] + +[[package]] +name = "rand_chacha" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d3022b5f1df60f26e1ffddd6c66e8aa15de382ae63b3a0c1bfc0e4d3e3f325cb" +dependencies = [ + "ppv-lite86", + "rand_core", +] + +[[package]] +name = "rand_core" +version = "0.9.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "76afc826de14238e6e8c374ddcc1fa19e374fd8dd986b0d2af0d02377261d83c" +dependencies = [ + "getrandom 0.3.4", +] + +[[package]] +name = "reqwest" +version = "0.12.28" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "eddd3ca559203180a307f12d114c268abf583f59b03cb906fd0b3ff8646c1147" +dependencies = [ + "base64", + "bytes", + "futures-channel", + "futures-core", + "futures-util", + "http", + "http-body", + "http-body-util", + "hyper", + "hyper-rustls", + "hyper-util", + "js-sys", + "log", + "percent-encoding", + "pin-project-lite", + "quinn", + "rustls", + "rustls-pki-types", + "serde", + "serde_json", + "serde_urlencoded", + "sync_wrapper", + "tokio", + "tokio-rustls", + "tower", + "tower-http", + "tower-service", + "url", + "wasm-bindgen", + "wasm-bindgen-futures", + "web-sys", + "webpki-roots", +] + +[[package]] +name = "ring" +version = "0.17.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a4689e6c2294d81e88dc6261c768b63bc4fcdb852be6d1352498b114f61383b7" +dependencies = [ + "cc", + "cfg-if", + "getrandom 0.2.17", + "libc", + "untrusted", + "windows-sys 0.52.0", +] + +[[package]] +name = "rustc-hash" +version = "2.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "94300abf3f1ae2e2b8ffb7b58043de3d399c73fa6f4b73826402a5c457614dbe" + +[[package]] +name = "rustls" +version = "0.23.41" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6b92b125634d9b795e7beca796cc790df15a7fb38323bf3196fda83292d06b1f" +dependencies = [ + "once_cell", + "ring", + "rustls-pki-types", + "rustls-webpki", + "subtle", + "zeroize", +] + +[[package]] +name = "rustls-pki-types" +version = "1.14.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "30a7197ae7eb376e574fe940d068c30fe0462554a3ddbe4eca7838e049c937a9" +dependencies = [ + "web-time", + "zeroize", +] + +[[package]] +name = "rustls-webpki" +version = "0.103.13" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "61c429a8649f110dddef65e2a5ad240f747e85f7758a6bccc7e5777bd33f756e" +dependencies = [ + "ring", + "rustls-pki-types", + "untrusted", +] + +[[package]] +name = "rustversion" +version = "1.0.22" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b39cdef0fa800fc44525c84ccb54a029961a8215f9619753635a9c0d2538d46d" + +[[package]] +name = "ryu" +version = "1.0.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9774ba4a74de5f7b1c1451ed6cd5285a32eddb5cccb8cc655a4e50009e06477f" + +[[package]] +name = "serde" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9a8e94ea7f378bd32cbbd37198a4a91436180c5bb472411e48b5ec2e2124ae9e" +dependencies = [ + "serde_core", + "serde_derive", +] + +[[package]] +name = "serde_core" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "41d385c7d4ca58e59fc732af25c3983b67ac852c1a25000afe1175de458b67ad" +dependencies = [ + "serde_derive", +] + +[[package]] +name = "serde_derive" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d540f220d3187173da220f885ab66608367b6574e925011a9353e4badda91d79" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "serde_json" +version = "1.0.150" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e8014e44b4736ed0538adeecded0fce2a272f22dc9578a7eb6b2d9993c74cfb9" +dependencies = [ + "itoa", + "memchr", + "serde", + "serde_core", + "zmij", +] + +[[package]] +name = "serde_urlencoded" +version = "0.7.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d3491c14715ca2294c4d6a88f15e84739788c1d030eed8c110436aafdaa2f3fd" +dependencies = [ + "form_urlencoded", + "itoa", + "ryu", + "serde", +] + +[[package]] +name = "shlex" +version = "2.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f8fadd59c855ef2080decdef8ff161eb6661b86933c9d82e5ba29dc602a55aba" + +[[package]] +name = "slab" +version = "0.4.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0c790de23124f9ab44544d7ac05d60440adc586479ce501c1d6d7da3cd8c9cf5" + +[[package]] +name = "smallvec" +version = "1.15.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8ed6a63f02c8539c91a8685a86f4099661ba3da017932f6ebbea6de3f0fa7c90" + +[[package]] +name = "socket2" +version = "0.6.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "52d1cfed4120b4d927bf7c0f86d2087a4a7d6027c906d9f9d525a80573b9be51" +dependencies = [ + "libc", + "windows-sys 0.61.2", +] + +[[package]] +name = "stable_deref_trait" +version = "1.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6ce2be8dc25455e1f91df71bfa12ad37d7af1092ae736f3a6cd0e37bc7810596" + +[[package]] +name = "subtle" +version = "2.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "13c2bddecc57b384dee18652358fb23172facb8a2c51ccc10d74c157bdea3292" + +[[package]] +name = "syn" +version = "2.0.118" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1b9ae57f904213ebb649ce6895b8a66c66f0203b9319718f69a5612a065b1422" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + +[[package]] +name = "sync_wrapper" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0bf256ce5efdfa370213c1dabab5935a12e49f2c58d15e9eac2870d3b4f27263" +dependencies = [ + "futures-core", +] + +[[package]] +name = "synstructure" +version = "0.13.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "728a70f3dbaf5bab7f0c4b1ac8d7ae5ea60a4b5549c8a5914361c99147a709d2" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "target-lexicon" +version = "0.12.16" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "61c41af27dd6d1e27b1b16b489db798443478cef1f06a660c96db617ba5de3b1" + +[[package]] +name = "thiserror" +version = "2.0.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4288b5bcbc7920c07a1149a35cf9590a2aa808e0bc1eafaade0b80947865fbc4" +dependencies = [ + "thiserror-impl", +] + +[[package]] +name = "thiserror-impl" +version = "2.0.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ebc4ee7f67670e9b64d05fa4253e753e016c6c95ff35b89b7941d6b856dec1d5" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "tinystr" +version = "0.8.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c8323304221c2a851516f22236c5722a72eaa19749016521d6dff0824447d96d" +dependencies = [ + "displaydoc", + "zerovec", +] + +[[package]] +name = "tinyvec" +version = "1.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3e61e67053d25a4e82c844e8424039d9745781b3fc4f32b8d55ed50f5f667ef3" +dependencies = [ + "tinyvec_macros", +] + +[[package]] +name = "tinyvec_macros" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" + +[[package]] +name = "tokio" +version = "1.52.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8fc7f01b389ac15039e4dc9531aa973a135d7a4135281b12d7c1bc79fd57fffe" +dependencies = [ + "bytes", + "libc", + "mio", + "pin-project-lite", + "socket2", + "windows-sys 0.61.2", +] + +[[package]] +name = "tokio-rustls" +version = "0.26.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1729aa945f29d91ba541258c8df89027d5792d85a8841fb65e8bf0f4ede4ef61" +dependencies = [ + "rustls", + "tokio", +] + +[[package]] +name = "tower" +version = "0.5.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ebe5ef63511595f1344e2d5cfa636d973292adc0eec1f0ad45fae9f0851ab1d4" +dependencies = [ + "futures-core", + "futures-util", + "pin-project-lite", + "sync_wrapper", + "tokio", + "tower-layer", + "tower-service", +] + +[[package]] +name = "tower-http" +version = "0.6.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4cfcf7e2740e6fc6d4d688b4ef00650406bb94adf4731e43c096c3a19fe40840" +dependencies = [ + "bitflags", + "bytes", + "futures-util", + "http", + "http-body", + "pin-project-lite", + "tower", + "tower-layer", + "tower-service", + "url", +] + +[[package]] +name = "tower-layer" +version = "0.3.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "121c2a6cda46980bb0fcd1647ffaf6cd3fc79a013de288782836f6df9c48780e" + +[[package]] +name = "tower-service" +version = "0.3.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8df9b6e13f2d32c91b9bd719c00d1958837bc7dec474d94952798cc8e69eeec3" + +[[package]] +name = "tracing" +version = "0.1.44" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "63e71662fa4b2a2c3a26f570f037eb95bb1f85397f3cd8076caed2f026a6d100" +dependencies = [ + "pin-project-lite", + "tracing-core", +] + +[[package]] +name = "tracing-core" +version = "0.1.36" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "db97caf9d906fbde555dd62fa95ddba9eecfd14cb388e4f491a66d74cd5fb79a" +dependencies = [ + "once_cell", +] + +[[package]] +name = "try-lock" +version = "0.2.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e421abadd41a4225275504ea4d6566923418b7f05506fbc9c0fe86ba7396114b" + +[[package]] +name = "unicode-ident" +version = "1.0.24" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" + +[[package]] +name = "unindent" +version = "0.2.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7264e107f553ccae879d21fbea1d6724ac785e8c3bfc762137959b5802826ef3" + +[[package]] +name = "untrusted" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8ecb6da28b8a351d773b68d5825ac39017e680750f980f3a1a85cd8dd28a47c1" + +[[package]] +name = "url" +version = "2.5.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ff67a8a4397373c3ef660812acab3268222035010ab8680ec4215f38ba3d0eed" +dependencies = [ + "form_urlencoded", + "idna", + "percent-encoding", + "serde", +] + +[[package]] +name = "utf8_iter" +version = "1.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6c140620e7ffbb22c2dee59cafe6084a59b5ffc27a8859a5f0d494b5d52b6be" + +[[package]] +name = "want" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bfa7760aed19e106de2c7c0b581b509f2f25d3dacaf737cb82ac61bc6d760b0e" +dependencies = [ + "try-lock", +] + +[[package]] +name = "wasi" +version = "0.11.1+wasi-snapshot-preview1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ccf3ec651a847eb01de73ccad15eb7d99f80485de043efb2f370cd654f4ea44b" + +[[package]] +name = "wasip2" +version = "1.0.4+wasi-0.2.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b67efb37e106e55ce722a510d6b5f9c17f083e5fc79afc2badeb12cc313d9487" +dependencies = [ + "wit-bindgen", +] + +[[package]] +name = "wasm-bindgen" +version = "0.2.125" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8ddb3f79143bced6de84270411622a2699cee572fc0875aeaf1e7867cf9fca1a" +dependencies = [ + "cfg-if", + "once_cell", + "rustversion", + "wasm-bindgen-macro", + "wasm-bindgen-shared", +] + +[[package]] +name = "wasm-bindgen-futures" +version = "0.4.75" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "503b14d284f2c8dac03b819967e155ea753f573586193b2b2c95990cb5d69280" +dependencies = [ + "js-sys", + "wasm-bindgen", +] + +[[package]] +name = "wasm-bindgen-macro" +version = "0.2.125" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4e21a184b13fb19e157296e2c46056aec9092264fab83e4ba59e68c61b323c3d" +dependencies = [ + "quote", + "wasm-bindgen-macro-support", +] + +[[package]] +name = "wasm-bindgen-macro-support" +version = "0.2.125" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fecefd9c35bd935a20fc3fc344b5f29138961e4f47fb03297d88f2587afb5ebd" +dependencies = [ + "bumpalo", + "proc-macro2", + "quote", + "syn", + "wasm-bindgen-shared", +] + +[[package]] +name = "wasm-bindgen-shared" +version = "0.2.125" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "23939e44bb9a5d7576fa2b563dc2e136628f1224e88a8deed09e04858b77871f" +dependencies = [ + "unicode-ident", +] + +[[package]] +name = "web-sys" +version = "0.3.102" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a6430a72df5eb332242960fe84b3002a241163998241eb596d4f739b9757061d" +dependencies = [ + "js-sys", + "wasm-bindgen", +] + +[[package]] +name = "web-time" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5a6580f308b1fad9207618087a65c04e7a10bc77e02c8e84e9b00dd4b12fa0bb" +dependencies = [ + "js-sys", + "wasm-bindgen", +] + +[[package]] +name = "webpki-roots" +version = "1.0.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bf85cb06032201fa7c6f829d7db5a7e5aa45bcc0655327713065f6f0576731bf" +dependencies = [ + "rustls-pki-types", +] + +[[package]] +name = "windows-link" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5" + +[[package]] +name = "windows-sys" +version = "0.52.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "282be5f36a8ce781fad8c8ae18fa3f9beff57ec1b52cb3de0789201425d9a33d" +dependencies = [ + "windows-targets 0.52.6", +] + +[[package]] +name = "windows-sys" +version = "0.60.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f2f500e4d28234f72040990ec9d39e3a6b950f9f22d3dba18416c35882612bcb" +dependencies = [ + "windows-targets 0.53.5", +] + +[[package]] +name = "windows-sys" +version = "0.61.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ae137229bcbd6cdf0f7b80a31df61766145077ddf49416a728b02cb3921ff3fc" +dependencies = [ + "windows-link", +] + +[[package]] +name = "windows-targets" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9b724f72796e036ab90c1021d4780d4d3d648aca59e491e6b98e725b84e99973" +dependencies = [ + "windows_aarch64_gnullvm 0.52.6", + "windows_aarch64_msvc 0.52.6", + "windows_i686_gnu 0.52.6", + "windows_i686_gnullvm 0.52.6", + "windows_i686_msvc 0.52.6", + "windows_x86_64_gnu 0.52.6", + "windows_x86_64_gnullvm 0.52.6", + "windows_x86_64_msvc 0.52.6", +] + +[[package]] +name = "windows-targets" +version = "0.53.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4945f9f551b88e0d65f3db0bc25c33b8acea4d9e41163edf90dcd0b19f9069f3" +dependencies = [ + "windows-link", + "windows_aarch64_gnullvm 0.53.1", + "windows_aarch64_msvc 0.53.1", + "windows_i686_gnu 0.53.1", + "windows_i686_gnullvm 0.53.1", + "windows_i686_msvc 0.53.1", + "windows_x86_64_gnu 0.53.1", + "windows_x86_64_gnullvm 0.53.1", + "windows_x86_64_msvc 0.53.1", +] + +[[package]] +name = "windows_aarch64_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32a4622180e7a0ec044bb555404c800bc9fd9ec262ec147edd5989ccd0c02cd3" + +[[package]] +name = "windows_aarch64_gnullvm" +version = "0.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a9d8416fa8b42f5c947f8482c43e7d89e73a173cead56d044f6a56104a6d1b53" + +[[package]] +name = "windows_aarch64_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09ec2a7bb152e2252b53fa7803150007879548bc709c039df7627cabbd05d469" + +[[package]] +name = "windows_aarch64_msvc" +version = "0.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b9d782e804c2f632e395708e99a94275910eb9100b2114651e04744e9b125006" + +[[package]] +name = "windows_i686_gnu" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8e9b5ad5ab802e97eb8e295ac6720e509ee4c243f69d781394014ebfe8bbfa0b" + +[[package]] +name = "windows_i686_gnu" +version = "0.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "960e6da069d81e09becb0ca57a65220ddff016ff2d6af6a223cf372a506593a3" + +[[package]] +name = "windows_i686_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0eee52d38c090b3caa76c563b86c3a4bd71ef1a819287c19d586d7334ae8ed66" + +[[package]] +name = "windows_i686_gnullvm" +version = "0.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fa7359d10048f68ab8b09fa71c3daccfb0e9b559aed648a8f95469c27057180c" + +[[package]] +name = "windows_i686_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "240948bc05c5e7c6dabba28bf89d89ffce3e303022809e73deaefe4f6ec56c66" + +[[package]] +name = "windows_i686_msvc" +version = "0.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1e7ac75179f18232fe9c285163565a57ef8d3c89254a30685b57d83a38d326c2" + +[[package]] +name = "windows_x86_64_gnu" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "147a5c80aabfbf0c7d901cb5895d1de30ef2907eb21fbbab29ca94c5b08b1a78" + +[[package]] +name = "windows_x86_64_gnu" +version = "0.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9c3842cdd74a865a8066ab39c8a7a473c0778a3f29370b5fd6b4b9aa7df4a499" + +[[package]] +name = "windows_x86_64_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "24d5b23dc417412679681396f2b49f3de8c1473deb516bd34410872eff51ed0d" + +[[package]] +name = "windows_x86_64_gnullvm" +version = "0.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0ffa179e2d07eee8ad8f57493436566c7cc30ac536a3379fdf008f47f6bb7ae1" + +[[package]] +name = "windows_x86_64_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec" + +[[package]] +name = "windows_x86_64_msvc" +version = "0.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d6bbff5f0aada427a1e5a6da5f1f98158182f26556f345ac9e04d36d0ebed650" + +[[package]] +name = "wit-bindgen" +version = "0.57.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ebf944e87a7c253233ad6766e082e3cd714b5d03812acc24c318f549614536e" + +[[package]] +name = "writeable" +version = "0.6.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ffae5123b2d3fc086436f8834ae3ab053a283cfac8fe0a0b8eaae044768a4c4" + +[[package]] +name = "yoke" +version = "0.8.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "709fe23a0424b6a435d82152b1bd3fdfb0833487d5fa90d05d42762a9891fef5" +dependencies = [ + "stable_deref_trait", + "yoke-derive", + "zerofrom", +] + +[[package]] +name = "yoke-derive" +version = "0.8.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "de844c262c8848816172cef550288e7dc6c7b7814b4ee56b3e1553f275f1858e" +dependencies = [ + "proc-macro2", + "quote", + "syn", + "synstructure", +] + +[[package]] +name = "zerocopy" +version = "0.8.52" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ce1022995ff5ff5d841ad7d994facc23098cd40152f2c1d11cd607c6f530653f" +dependencies = [ + "zerocopy-derive", +] + +[[package]] +name = "zerocopy-derive" +version = "0.8.52" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ae7f38b72ec2a254e2b87ef277cf2cd4fb97cbebf944faa6f33354da0867930" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "zerofrom" +version = "0.1.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0ec05a11813ea801ff6d75110ad09cd0824ddba17dfe17128ea0d5f68e6c5272" +dependencies = [ + "zerofrom-derive", +] + +[[package]] +name = "zerofrom-derive" +version = "0.1.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "11532158c46691caf0f2593ea8358fed6bbf68a0315e80aae9bd41fbade684a1" +dependencies = [ + "proc-macro2", + "quote", + "syn", + "synstructure", +] + +[[package]] +name = "zeroize" +version = "1.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e13c156562582aa81c60cb29407084cdb54c4164760106ab78e6c5b0858cf64e" + +[[package]] +name = "zerotrie" +version = "0.2.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0f9152d31db0792fa83f70fb2f83148effb5c1f5b8c7686c3459e361d9bc20bf" +dependencies = [ + "displaydoc", + "yoke", + "zerofrom", +] + +[[package]] +name = "zerovec" +version = "0.11.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "90f911cbc359ab6af17377d242225f4d75119aec87ea711a880987b18cd7b239" +dependencies = [ + "yoke", + "zerofrom", + "zerovec-derive", +] + +[[package]] +name = "zerovec-derive" +version = "0.11.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "625dc425cab0dca6dc3c3319506e6593dcb08a9f387ea3b284dbd52a92c40555" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "zmij" +version = "1.0.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b8848ee67ecc8aedbaf3e4122217aff892639231befc6a1b58d29fff4c2cabaa" diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml new file mode 100644 index 00000000000..fdc5f5efde1 --- /dev/null +++ b/litellm-rust/Cargo.toml @@ -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" diff --git a/litellm-rust/README.md b/litellm-rust/README.md new file mode 100644 index 00000000000..15ad1855420 --- /dev/null +++ b/litellm-rust/README.md @@ -0,0 +1,34 @@ +# LiteLLM Rust + +This workspace contains the staged Rust implementation for LiteLLM. + +Rust starts as a pure transform core used by the existing Python host. Python +continues to own auth, configuration, network I/O, retries, routing, logging, +callbacks, spend tracking, and customer plugins until each Rust path has parity +coverage and production evidence. + +## Layout + +```text +crates/ + core/ Route contracts, shared pure types, errors, and templates. + src/ocr/ + providers/ Provider-specific pure transforms. + src/mistral/ocr/transformation.rs + python-bridge/ PyO3 bridge for Python LiteLLM. +``` + +The folder shape should follow the Python provider tree: +`providers/src///transformation.rs`. The bridge should expose +one function per top-level route, starting with `ocr(payload)`. + +## Checks + +Run these before pushing Rust changes. GitHub Actions runs the same checks for +changes under `litellm-rust/`. + +```bash +cargo fmt --check +cargo clippy --workspace --all-targets -- -D warnings +cargo test --workspace +``` diff --git a/litellm-rust/crates/core/CLAUDE.md b/litellm-rust/crates/core/CLAUDE.md new file mode 100644 index 00000000000..51e0d215e6a --- /dev/null +++ b/litellm-rust/crates/core/CLAUDE.md @@ -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. diff --git a/litellm-rust/crates/core/Cargo.toml b/litellm-rust/crates/core/Cargo.toml new file mode 100644 index 00000000000..e54002fe5e8 --- /dev/null +++ b/litellm-rust/crates/core/Cargo.toml @@ -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 diff --git a/litellm-rust/crates/core/src/error.rs b/litellm-rust/crates/core/src/error.rs new file mode 100644 index 00000000000..645e261f76d --- /dev/null +++ b/litellm-rust/crates/core/src/error.rs @@ -0,0 +1,33 @@ +use thiserror::Error; + +pub type CoreResult = Result; + +#[derive(Debug, Error, PartialEq, Eq)] +pub enum CoreError { + #[error("expected {expected}, got {actual}")] + InvalidType { + expected: &'static str, + actual: &'static str, + }, + #[error("missing required field: {0}")] + MissingField(&'static str), + #[error("invalid response: {0}")] + InvalidResponse(String), + #[error("{0}")] + Auth(String), + #[error("OCR request failed with status {status}: {body}")] + Http { status: u16, body: String }, + #[error("OCR network error: {0}")] + Network(String), +} + +pub fn json_type_name(value: &serde_json::Value) -> &'static str { + match value { + serde_json::Value::Null => "null", + serde_json::Value::Bool(_) => "bool", + serde_json::Value::Number(_) => "number", + serde_json::Value::String(_) => "string", + serde_json::Value::Array(_) => "array", + serde_json::Value::Object(_) => "object", + } +} diff --git a/litellm-rust/crates/core/src/lib.rs b/litellm-rust/crates/core/src/lib.rs new file mode 100644 index 00000000000..a88204867c4 --- /dev/null +++ b/litellm-rust/crates/core/src/lib.rs @@ -0,0 +1,4 @@ +pub mod error; +pub mod ocr; + +pub use error::{CoreError, CoreResult}; diff --git a/litellm-rust/crates/core/src/ocr/mod.rs b/litellm-rust/crates/core/src/ocr/mod.rs new file mode 100644 index 00000000000..ec2fbb969a6 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/mod.rs @@ -0,0 +1,2 @@ +pub mod transformation; +pub mod types; diff --git a/litellm-rust/crates/core/src/ocr/transformation.rs b/litellm-rust/crates/core/src/ocr/transformation.rs new file mode 100644 index 00000000000..7353d9d22c4 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/transformation.rs @@ -0,0 +1,32 @@ +use serde_json::{Map, Value}; + +use crate::CoreResult; + +use super::types::{OcrRequestData, OcrResponseData}; + +pub trait OcrProviderConfig { + fn supported_ocr_params(&self) -> &'static [&'static str]; + + fn map_ocr_params(&self, non_default_params: &Map) -> Map { + let mut mapped_params = Map::new(); + for (param, value) in non_default_params { + if self.supported_ocr_params().contains(¶m.as_str()) { + mapped_params.insert(param.clone(), value.clone()); + } + } + mapped_params + } + + fn transform_ocr_request( + &self, + model: &str, + document: Value, + optional_params: Map, + ) -> CoreResult; + + fn transform_ocr_response( + &self, + model: &str, + response_json: Value, + ) -> CoreResult; +} diff --git a/litellm-rust/crates/core/src/ocr/types.rs b/litellm-rust/crates/core/src/ocr/types.rs new file mode 100644 index 00000000000..1a72b8f1d66 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/types.rs @@ -0,0 +1,29 @@ +use serde::{Deserialize, Serialize}; +use serde_json::Value; + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct OcrRequestData { + pub data: Value, + pub files: Option, +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct OcrResponseData { + pub pages: Vec, + pub model: String, + pub document_annotation: Option, + pub usage_info: Option, + pub object: String, +} + +impl OcrResponseData { + pub fn into_json(self) -> Value { + serde_json::json!({ + "pages": self.pages, + "model": self.model, + "document_annotation": self.document_annotation, + "usage_info": self.usage_info, + "object": self.object, + }) + } +} diff --git a/litellm-rust/crates/providers/CLAUDE.md b/litellm-rust/crates/providers/CLAUDE.md new file mode 100644 index 00000000000..0f7fdcda2aa --- /dev/null +++ b/litellm-rust/crates/providers/CLAUDE.md @@ -0,0 +1,53 @@ +# CLAUDE.md + +Rules for `litellm-rust/crates/providers`. + +## Responsibility + +`providers` owns provider-specific pure transforms. It mirrors the existing +Python provider modules closely enough that parity review is mechanical. + +Provider files should map to the Python provider tree: + +```text +providers/src///transformation.rs +``` + +For example, Mistral OCR lives at +`providers/src/mistral/ocr/transformation.rs`, matching +`litellm/llms/mistral/ocr/transformation.py`. + +Allowed: +- Provider request transforms. +- Provider response normalization. +- Supported-parameter filtering. +- Provider-specific validation that does not require I/O or secrets. + +Not allowed: +- HTTP clients or provider SDK calls. +- Environment variable reads. +- API key resolution or auth header construction. +- Logging, callbacks, spend tracking, retries, routing, cooldowns, or fallbacks. +- Panics on bad user/provider input. + +## Required Tests + +Every provider transform must include focused unit tests for: +- Supported params matching the Python provider config. +- Unknown params being dropped or transformed the same way as Python. +- Request body shape matching Python output. +- Response normalization with complete, missing, null, and extra fields. +- Bad input returning typed errors. + +For OCR specifically, assume documents can contain personal data. Tests should +prove transforms do not copy document contents into error messages. + +## Implementation Rules + +- Prefer static supported-parameter lists over allocating strings on every call. +- Keep transforms deterministic and allocation-conscious, but choose clarity over + premature micro-optimization for tiny parameter lists. +- Use typed errors from `core`; avoid stringly-typed error plumbing. +- Add comments only when they explain Python-parity decisions or provider quirks. +- Put route-level provider dispatch in a route file such as `providers/src/ocr.rs`. + Do not move provider-specific transform logic into the Python bridge. diff --git a/litellm-rust/crates/providers/Cargo.toml b/litellm-rust/crates/providers/Cargo.toml new file mode 100644 index 00000000000..b8ed2066079 --- /dev/null +++ b/litellm-rust/crates/providers/Cargo.toml @@ -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 diff --git a/litellm-rust/crates/providers/src/lib.rs b/litellm-rust/crates/providers/src/lib.rs new file mode 100644 index 00000000000..1a0ca5b7e42 --- /dev/null +++ b/litellm-rust/crates/providers/src/lib.rs @@ -0,0 +1,2 @@ +pub mod mistral; +pub mod ocr; diff --git a/litellm-rust/crates/providers/src/mistral/mod.rs b/litellm-rust/crates/providers/src/mistral/mod.rs new file mode 100644 index 00000000000..3621ff6a2fd --- /dev/null +++ b/litellm-rust/crates/providers/src/mistral/mod.rs @@ -0,0 +1 @@ +pub mod ocr; diff --git a/litellm-rust/crates/providers/src/mistral/ocr/mod.rs b/litellm-rust/crates/providers/src/mistral/ocr/mod.rs new file mode 100644 index 00000000000..f239b6921fa --- /dev/null +++ b/litellm-rust/crates/providers/src/mistral/ocr/mod.rs @@ -0,0 +1 @@ +pub mod transformation; diff --git a/litellm-rust/crates/providers/src/mistral/ocr/transformation.rs b/litellm-rust/crates/providers/src/mistral/ocr/transformation.rs new file mode 100644 index 00000000000..fd691177783 --- /dev/null +++ b/litellm-rust/crates/providers/src/mistral/ocr/transformation.rs @@ -0,0 +1,292 @@ +use litellm_core::error::{json_type_name, CoreError, CoreResult}; +use litellm_core::ocr::transformation::OcrProviderConfig; +use litellm_core::ocr::types::{OcrRequestData, OcrResponseData}; +use serde_json::{Map, Value}; + +const SUPPORTED_OCR_PARAMS: &[&str] = &[ + "pages", + "include_image_base64", + "image_limit", + "image_min_size", + "bbox_annotation_format", + "document_annotation_format", + "document_annotation_prompt", + "extract_header", + "extract_footer", + "table_format", + "confidence_scores_granularity", + "id", +]; + +/// Default Mistral API base, used when the caller does not override `api_base`. +pub const MISTRAL_DEFAULT_API_BASE: &str = "https://api.mistral.ai/v1"; + +/// Environment variable holding the Mistral API key. +pub const MISTRAL_API_KEY_ENV: &str = "MISTRAL_API_KEY"; + +/// Error message raised when no Mistral API key can be resolved. +pub const MISSING_KEY_MESSAGE: &str = "Missing Mistral API Key - A call is being made to Mistral but no key is set either in the environment variables or via params"; + +/// Build the complete OCR endpoint URL, de-duplicating a trailing `/v1`. +/// +/// Blank/whitespace `api_base` is treated as absent (guard at resolution time). +pub fn complete_url(api_base: Option<&str>) -> String { + let base = api_base + .map(str::trim) + .filter(|base| !base.is_empty()) + .unwrap_or(MISTRAL_DEFAULT_API_BASE) + .trim_end_matches('/'); + + if base.ends_with("/v1") { + format!("{base}/ocr") + } else { + format!("{base}/v1/ocr") + } +} + +/// Resolve the Mistral API key from the explicit param or the environment. +/// +/// Blank/whitespace values are treated as absent. Returns `CoreError::Auth` +/// when no usable key is available. +/// +/// Note: the env fallback only reads the process environment. Secret-manager +/// backends (AWS/Azure/GCP/Vault) are resolved on the Python side and passed in +/// via `api_key`; this fallback is a last resort for direct/standalone use. +pub fn resolve_api_key( + api_key: Option<&str>, + env_lookup: &dyn Fn(&str) -> Option, +) -> CoreResult { + api_key + .map(str::trim) + .filter(|key| !key.is_empty()) + .map(str::to_string) + .or_else(|| env_lookup(MISTRAL_API_KEY_ENV).filter(|key| !key.trim().is_empty())) + .ok_or_else(|| CoreError::Auth(MISSING_KEY_MESSAGE.to_string())) +} + +pub struct MistralOcrConfig; + +pub const MISTRAL_OCR_CONFIG: MistralOcrConfig = MistralOcrConfig; + +impl OcrProviderConfig for MistralOcrConfig { + fn supported_ocr_params(&self) -> &'static [&'static str] { + SUPPORTED_OCR_PARAMS + } + + fn transform_ocr_request( + &self, + model: &str, + document: Value, + optional_params: Map, + ) -> CoreResult { + if !document.is_object() { + return Err(CoreError::InvalidType { + expected: "object", + actual: json_type_name(&document), + }); + } + + let mut data = Map::new(); + data.insert("model".to_string(), Value::String(model.to_string())); + data.insert("document".to_string(), document); + for (param, value) in optional_params { + data.insert(param, value); + } + + Ok(OcrRequestData { + data: Value::Object(data), + files: None, + }) + } + + fn transform_ocr_response( + &self, + model: &str, + response_json: Value, + ) -> CoreResult { + let response_object = response_json + .as_object() + .ok_or_else(|| CoreError::InvalidType { + expected: "object", + actual: json_type_name(&response_json), + })?; + + let pages = response_object + .get("pages") + .and_then(Value::as_array) + .cloned() + .unwrap_or_default(); + let model = response_object + .get("model") + .and_then(Value::as_str) + .unwrap_or(model) + .to_string(); + let document_annotation = response_object.get("document_annotation").cloned(); + let usage_info = response_object.get("usage_info").cloned(); + + Ok(OcrResponseData { + pages, + model, + document_annotation, + usage_info, + object: "ocr".to_string(), + }) + } +} + +pub fn supported_ocr_params() -> &'static [&'static str] { + MISTRAL_OCR_CONFIG.supported_ocr_params() +} + +pub fn map_ocr_params(non_default_params: &Map) -> Map { + MISTRAL_OCR_CONFIG.map_ocr_params(non_default_params) +} + +pub fn transform_ocr_request( + model: &str, + document: Value, + optional_params: Map, +) -> CoreResult { + MISTRAL_OCR_CONFIG.transform_ocr_request(model, document, optional_params) +} + +pub fn transform_ocr_response(model: &str, response_json: Value) -> CoreResult { + MISTRAL_OCR_CONFIG.transform_ocr_response(model, response_json) +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + #[test] + fn supported_params_match_python_mistral_ocr_config() { + assert_eq!( + supported_ocr_params(), + &[ + "pages", + "include_image_base64", + "image_limit", + "image_min_size", + "bbox_annotation_format", + "document_annotation_format", + "document_annotation_prompt", + "extract_header", + "extract_footer", + "table_format", + "confidence_scores_granularity", + "id", + ] + ); + } + + #[test] + fn map_ocr_params_drops_unknown_params() { + let params = json!({ + "extract_header": true, + "unsupported_param": "value", + "pages": [0, 1] + }); + let mapped = map_ocr_params(params.as_object().unwrap()); + + assert_eq!(mapped.get("extract_header"), Some(&json!(true))); + assert_eq!(mapped.get("pages"), Some(&json!([0, 1]))); + assert!(!mapped.contains_key("unsupported_param")); + } + + #[test] + fn transform_ocr_request_builds_mistral_body() { + let document = json!({ + "type": "document_url", + "document_url": "https://example.com/doc.pdf" + }); + let optional_params = json!({ + "include_image_base64": true, + "table_format": "html" + }) + .as_object() + .unwrap() + .clone(); + + let result = transform_ocr_request("mistral-ocr-latest", document.clone(), optional_params) + .expect("request should transform"); + + assert_eq!( + result.data, + json!({ + "model": "mistral-ocr-latest", + "document": document, + "include_image_base64": true, + "table_format": "html" + }) + ); + assert_eq!(result.files, None); + } + + #[test] + fn transform_ocr_request_rejects_non_object_document() { + let err = transform_ocr_request("mistral-ocr-latest", json!("bad"), Map::new()) + .expect_err("string document should be rejected"); + + assert_eq!( + err, + CoreError::InvalidType { + expected: "object", + actual: "string", + } + ); + } + + #[test] + fn transform_ocr_response_normalizes_mistral_json() { + let response = json!({ + "pages": [{"index": 0, "markdown": "hello"}], + "model": "mistral-ocr-2505-completion", + "document_annotation": null, + "usage_info": {"pages_processed": 1} + }); + + let result = transform_ocr_response("mistral-ocr-latest", response) + .expect("response should transform"); + + assert_eq!(result.pages, vec![json!({"index": 0, "markdown": "hello"})]); + assert_eq!(result.model, "mistral-ocr-2505-completion"); + assert_eq!(result.document_annotation, Some(Value::Null)); + assert_eq!(result.usage_info, Some(json!({"pages_processed": 1}))); + assert_eq!(result.object, "ocr"); + } + + #[test] + fn complete_url_defaults_and_dedupes_v1() { + assert_eq!(complete_url(None), "https://api.mistral.ai/v1/ocr"); + assert_eq!(complete_url(Some(" ")), "https://api.mistral.ai/v1/ocr"); + assert_eq!( + complete_url(Some("https://proxy.internal")), + "https://proxy.internal/v1/ocr" + ); + assert_eq!( + complete_url(Some("https://proxy.internal/v1/")), + "https://proxy.internal/v1/ocr" + ); + } + + #[test] + fn resolve_api_key_prefers_param_then_env() { + let no_env = |_: &str| None; + assert_eq!( + resolve_api_key(Some("sk-param"), &no_env).unwrap(), + "sk-param" + ); + + let with_env = |key: &str| (key == MISTRAL_API_KEY_ENV).then(|| "sk-env".to_string()); + assert_eq!(resolve_api_key(None, &with_env).unwrap(), "sk-env"); + // Blank param falls through to the environment. + assert_eq!(resolve_api_key(Some(" "), &with_env).unwrap(), "sk-env"); + } + + #[test] + fn resolve_api_key_errors_when_absent() { + let err = resolve_api_key(None, &|_| None).expect_err("missing key should error"); + assert_eq!(err, CoreError::Auth(MISSING_KEY_MESSAGE.to_string())); + } +} diff --git a/litellm-rust/crates/providers/src/ocr.rs b/litellm-rust/crates/providers/src/ocr.rs new file mode 100644 index 00000000000..dcd56a5f0b4 --- /dev/null +++ b/litellm-rust/crates/providers/src/ocr.rs @@ -0,0 +1,127 @@ +//! End-to-end OCR orchestration. +//! +//! Owns the whole Mistral OCR call so the Python side stays a thin bridge: +//! resolve the API key, build the URL + body via the pure transforms, POST it, +//! and normalize the response. The HTTP client is built once and reused. + +use std::sync::OnceLock; +use std::time::Duration; + +use litellm_core::error::CoreError; +use litellm_core::ocr::transformation::OcrProviderConfig; +use litellm_core::CoreResult; +use serde_json::{Map, Value}; + +use crate::mistral::ocr::transformation as mistral; +use crate::mistral::ocr::transformation::MISTRAL_OCR_CONFIG; + +/// OCR over large documents can take a while; bound it generously rather than +/// hanging forever on an unresponsive upstream. The client-level limit is the +/// outer ceiling; callers can tighten it per request via ``run_ocr``'s ``timeout``. +const OCR_TIMEOUT_SECS: u64 = 600; + +/// Maximum upstream body characters retained in error messages. OCR responses +/// can echo document contents and prompts; keep enough for debugging without +/// forwarding sensitive payloads across the host boundary. +const ERROR_BODY_MAX_CHARS: usize = 256; + +/// Process-wide blocking HTTP client (connection pool + TLS reused across calls). +fn http_client() -> &'static reqwest::blocking::Client { + static CLIENT: OnceLock = OnceLock::new(); + CLIENT.get_or_init(|| { + reqwest::blocking::Client::builder() + .timeout(Duration::from_secs(OCR_TIMEOUT_SECS)) + .build() + .expect("failed to build reqwest client") + }) +} + +fn truncate_error_body(body: &str) -> String { + if body.chars().count() <= ERROR_BODY_MAX_CHARS { + return body.to_string(); + } + let truncated: String = body.chars().take(ERROR_BODY_MAX_CHARS).collect(); + format!("{truncated}... (truncated)") +} + +/// Perform a Mistral OCR call end to end and return the normalized response as +/// JSON (the shape the Python `OCRResponse` model expects). +/// +/// Blocking: intended to be called with the GIL released from the Python bridge. +pub fn run_ocr( + model: &str, + document: Value, + api_key: Option<&str>, + api_base: Option<&str>, + optional_params: Map, + timeout: Option, +) -> CoreResult { + let config = &MISTRAL_OCR_CONFIG; + + let api_key = mistral::resolve_api_key(api_key, &|key| std::env::var(key).ok())?; + let url = mistral::complete_url(api_base); + let filtered_params = config.map_ocr_params(&optional_params); + let body = config + .transform_ocr_request(model, document, filtered_params)? + .data; + + let mut request = http_client().post(&url).bearer_auth(&api_key).json(&body); + if let Some(duration) = timeout { + request = request.timeout(duration); + } + + let response = request + .send() + .map_err(|err| CoreError::Network(err.to_string()))?; + + let status = response.status(); + let text = response + .text() + .map_err(|err| CoreError::Network(err.to_string()))?; + + if !status.is_success() { + return Err(CoreError::Http { + status: status.as_u16(), + body: truncate_error_body(&text), + }); + } + + let response_json: Value = serde_json::from_str(&text) + .map_err(|err| CoreError::InvalidResponse(format!("invalid OCR response JSON: {err}")))?; + + Ok(config + .transform_ocr_response(model, response_json)? + .into_json()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn truncate_error_body_passes_short_strings_through() { + let body = "Unauthorized"; + assert_eq!(truncate_error_body(body), "Unauthorized"); + } + + #[test] + fn truncate_error_body_caps_long_payloads() { + let body = "x".repeat(ERROR_BODY_MAX_CHARS + 50); + let truncated = truncate_error_body(&body); + + assert!(truncated.ends_with("... (truncated)")); + let prefix_chars = truncated + .strip_suffix("... (truncated)") + .expect("truncated marker present") + .chars() + .count(); + assert_eq!(prefix_chars, ERROR_BODY_MAX_CHARS); + } + + #[test] + fn truncate_error_body_does_not_split_multibyte_chars() { + let body = "é".repeat(ERROR_BODY_MAX_CHARS + 10); + let truncated = truncate_error_body(&body); + assert!(truncated.is_char_boundary(truncated.len())); + } +} diff --git a/litellm-rust/crates/python-bridge/CLAUDE.md b/litellm-rust/crates/python-bridge/CLAUDE.md new file mode 100644 index 00000000000..efa1a554c9c --- /dev/null +++ b/litellm-rust/crates/python-bridge/CLAUDE.md @@ -0,0 +1,36 @@ +# CLAUDE.md + +Rules for `litellm-rust/crates/python-bridge`. + +## Responsibility + +`python-bridge` is the PyO3 boundary between Python LiteLLM and Rust transforms. +Keep this crate thin. It adapts Python objects to Rust payloads and returns +Python-compatible dictionaries. + +## Bridge Shape + +- Prefer one stable method per top-level LiteLLM route, for example + `ocr(payload)`. +- Do not add one exported PyO3 function per provider helper unless there is a + measured reason. +- Provider dispatch belongs in Rust route modules such as + `litellm_providers::ocr`, not in this PyO3 crate. +- Python owns rollout state and fallback. Rust should return errors; Python + decides whether to raise or fall back. + +## Data Handling + +- OCR payloads can contain personal data and large base64 images. Do not log + payloads or provider responses. +- Avoid copying large payloads more than needed. The current JSON round-trip is + acceptable for the first scaffold, but future performance work should evaluate + direct PyO3 conversion before expanding Rust coverage to image-heavy paths. +- Do not expose raw Rust errors that include document contents or upstream + bodies. + +## Tests + +- `cargo test --workspace` must compile this crate. +- Python tests must cover bridge disabled, bridge enabled, and module-missing + fallback behavior for every exposed route. diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml new file mode 100644 index 00000000000..80b6478daac --- /dev/null +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -0,0 +1,16 @@ +[package] +name = "litellm-python-bridge" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[lib] +name = "litellm_python_bridge" +crate-type = ["cdylib"] + +[dependencies] +litellm-core.workspace = true +litellm-providers.workspace = true +pyo3 = { workspace = true, features = ["extension-module"] } +serde_json.workspace = true diff --git a/litellm-rust/crates/python-bridge/src/gil.rs b/litellm-rust/crates/python-bridge/src/gil.rs new file mode 100644 index 00000000000..dc1b591735c --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/gil.rs @@ -0,0 +1,32 @@ +//! GIL accounting. +//! +//! A single chokepoint for releasing the GIL around blocking work. Every +//! blocking call in the bridge goes through [`release_gil`] instead of calling +//! `Python::allow_threads` directly, so the release count stays accurate and we +//! have one place to extend later (timing histograms, per-call labels, etc.). + +use std::sync::atomic::{AtomicU64, Ordering}; + +use pyo3::prelude::*; + +/// Number of times the bridge has released the GIL since process start. +static GIL_RELEASES: AtomicU64 = AtomicU64::new(0); + +/// Release the GIL around `f`, recording the release. +/// +/// `f` must not touch any Python state — that is what makes releasing the GIL +/// safe. Returning the value back to Python re-acquires the GIL at the call +/// site, after `f` has finished. +pub fn release_gil(py: Python<'_>, f: F) -> T +where + F: FnOnce() -> T + Send, + T: Send, +{ + GIL_RELEASES.fetch_add(1, Ordering::Relaxed); + py.allow_threads(f) +} + +/// Total GIL releases performed by the bridge so far. +pub fn release_count() -> u64 { + GIL_RELEASES.load(Ordering::Relaxed) +} diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs new file mode 100644 index 00000000000..15e93f7b00c --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -0,0 +1,100 @@ +use std::time::Duration; + +use litellm_core::error::CoreError; +use litellm_providers::ocr::run_ocr; +use pyo3::exceptions::{PyRuntimeError, PyValueError}; +use pyo3::prelude::*; +use pyo3::types::{PyAny, PyDict}; +use serde_json::{Map, Value}; + +mod gil; + +fn py_to_json(py: Python<'_>, value: &Bound<'_, PyAny>) -> PyResult { + let json = py.import("json")?; + let encoded: String = json.call_method1("dumps", (value,))?.extract()?; + serde_json::from_str(&encoded).map_err(|err| PyValueError::new_err(err.to_string())) +} + +fn json_to_py(py: Python<'_>, value: Value) -> PyResult> { + let json = py.import("json")?; + let encoded = + serde_json::to_string(&value).map_err(|err| PyValueError::new_err(err.to_string()))?; + Ok(json.call_method1("loads", (encoded,))?.unbind()) +} + +/// Map a core error to the closest Python exception. Caller-input problems +/// (auth, bad types, missing fields) -> `ValueError`; everything else +/// (network, upstream status, parse failures) -> `RuntimeError`. +fn core_error_to_pyerr(err: CoreError) -> PyErr { + match err { + CoreError::Auth(message) => PyValueError::new_err(message), + CoreError::InvalidType { .. } | CoreError::MissingField(_) => { + PyValueError::new_err(err.to_string()) + } + other => PyRuntimeError::new_err(other.to_string()), + } +} + +/// Perform a Mistral OCR call end to end and return the response as a dict. +#[pyfunction] +#[pyo3(signature = (model, document, api_key=None, api_base=None, optional_params=None, timeout_seconds=None))] +fn ocr( + py: Python<'_>, + model: String, + document: Py, + api_key: Option, + api_base: Option, + optional_params: Option>, + timeout_seconds: Option, +) -> PyResult> { + let document = py_to_json(py, document.bind(py))?; + + let optional_params = match optional_params { + Some(params) => match py_to_json(py, params.bind(py))? { + Value::Object(map) => map, + _ => return Err(PyValueError::new_err("optional_params must be a dict")), + }, + None => Map::new(), + }; + + let timeout = timeout_seconds.and_then(|secs| { + if secs.is_finite() && secs > 0.0 { + Some(Duration::from_secs_f64(secs)) + } else { + None + } + }); + + // Release the GIL during the blocking HTTP call (counted for observability). + let result = gil::release_gil(py, || { + run_ocr( + &model, + document, + api_key.as_deref(), + api_base.as_deref(), + optional_params, + timeout, + ) + }); + + match result { + Ok(value) => json_to_py(py, value), + Err(err) => Err(core_error_to_pyerr(err)), + } +} + +/// Bridge GIL accounting, e.g. `{"releases": 12}`. Lets the Python side observe +/// how often the bridge has dropped the GIL for blocking work. +#[pyfunction] +fn gil_stats(py: Python<'_>) -> PyResult> { + let stats = PyDict::new(py); + stats.set_item("releases", gil::release_count())?; + Ok(stats.into_any().unbind()) +} + +#[pymodule] +fn litellm_python_bridge(module: &Bound<'_, PyModule>) -> PyResult<()> { + module.add_function(wrap_pyfunction!(ocr, module)?)?; + module.add_function(wrap_pyfunction!(gil_stats, module)?)?; + Ok(()) +} diff --git a/litellm/__init__.py b/litellm/__init__.py index b1ad63d72b0..c8bea6953e6 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -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 * diff --git a/litellm/integrations/code_interpreter_interception/handler.py b/litellm/integrations/code_interpreter_interception/handler.py index da8149eab9b..362581937d7 100644 --- a/litellm/integrations/code_interpreter_interception/handler.py +++ b/litellm/integrations/code_interpreter_interception/handler.py @@ -9,9 +9,11 @@ captured stdout back through the typed agentic loop plan. import json import time import uuid -from typing import Any, cast +from typing import Any, Literal, TypedDict, cast import litellm +from pydantic import ValidationError + from litellm._logging import verbose_logger from litellm.integrations.custom_logger import CustomLogger from litellm.types.integrations.code_interpreter_interception import ( @@ -20,15 +22,93 @@ from litellm.types.integrations.code_interpreter_interception import ( from litellm.types.integrations.custom_logger import ( AgenticLoopPlan, AgenticLoopRequestPatch, + CHAT_COMPLETION_AGENTIC_SURFACE, + NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES, + is_interception_internal_key, +) +from litellm.types.llms.openai import ( + ChatCompletionAssistantMessage, + ChatCompletionAssistantToolCall, + ChatCompletionToolMessage, +) +from litellm.types.utils import ( + CallTypes, + ChatCompletionMessageToolCall, + ModelResponse, ) -from litellm.types.utils import CallTypes LITELLM_CODE_EXECUTION_TOOL_NAME = "litellm_code_execution" _INTERCEPTION_ACTIVE_KEY = "_code_interpreter_interception_active" _SANDBOX_KEY = "_code_interpreter_interception_sandbox_key" +_CONVERTED_STREAM_KEY = "_code_interpreter_interception_converted_stream" +_LITELLM_METADATA_KEY = "litellm_metadata" _CACHE_TTL_SECONDS = 15 * 60 +class CodeExecutionToolCall(TypedDict, total=False): + id: str | None + call_id: str | None + type: Literal["function"] + name: str + arguments: str + + +class CodeInterpreterLogOutput(TypedDict): + type: Literal["logs"] + logs: str + + +class CodeInterpreterCall(TypedDict): + id: str + type: Literal["code_interpreter_call"] + status: Literal["completed"] + code: str + container_id: str | None + outputs: list[CodeInterpreterLogOutput] + + +class CodeExecutionFunctionParameters(TypedDict): + type: Literal["object"] + properties: dict[str, dict[str, str]] + required: list[str] + + +class ResponsesFunctionTool(TypedDict): + type: Literal["function"] + name: str + description: str + parameters: CodeExecutionFunctionParameters + + +class ChatCompletionFunctionDefinition(TypedDict): + name: str + description: str + parameters: CodeExecutionFunctionParameters + + +class ChatCompletionFunctionTool(TypedDict): + type: Literal["function"] + function: ChatCompletionFunctionDefinition + + +CodeExecutionFunctionTool = ResponsesFunctionTool | ChatCompletionFunctionTool + + +class ResponsesFunctionToolChoice(TypedDict): + type: Literal["function"] + name: str + + +class ChatCompletionFunctionToolChoice(TypedDict): + type: Literal["function"] + function: dict[str, str] + + +CodeExecutionFunctionToolChoice = ( + ResponsesFunctionToolChoice | ChatCompletionFunctionToolChoice +) + + def _resolve_sandbox_tool(sandbox_tool_name: str | None) -> dict[str, Any] | None: try: from litellm.sandbox.sandbox_tools import resolve_sandbox_tool @@ -97,9 +177,15 @@ class CodeInterpreterInterceptionLogger(CustomLogger): if not kwargs.get("_agentic_loop_depth"): kwargs.pop(_INTERCEPTION_ACTIVE_KEY, None) kwargs.pop(_SANDBOX_KEY, None) + self._strip_interception_metadata(kwargs) if not self.enabled: return None - if call_type not in (CallTypes.responses, CallTypes.aresponses): + if call_type not in ( + CallTypes.responses, + CallTypes.aresponses, + CallTypes.completion, + CallTypes.acompletion, + ): return None if ( self.enabled_providers is not None @@ -120,18 +206,10 @@ class CodeInterpreterInterceptionLogger(CustomLogger): kwargs[_SANDBOX_KEY] = uuid.uuid4().hex if kwargs.get("stream"): kwargs["stream"] = False - kwargs["_code_interpreter_interception_converted_stream"] = True + kwargs[_CONVERTED_STREAM_KEY] = True + self._write_interception_metadata(kwargs) - function_tool = { - "type": "function", - "name": LITELLM_CODE_EXECUTION_TOOL_NAME, - "description": "Execute python code in a sandbox and return stdout.", - "parameters": { - "type": "object", - "properties": {"code": {"type": "string"}}, - "required": ["code"], - }, - } + function_tool = self._get_function_tool(call_type=call_type) kwargs["tools"] = [ ( function_tool @@ -141,19 +219,90 @@ class CodeInterpreterInterceptionLogger(CustomLogger): for tool in tools ] if self._tool_choice_targets_code_interpreter(kwargs.get("tool_choice")): - kwargs["tool_choice"] = { - "type": "function", - "name": LITELLM_CODE_EXECUTION_TOOL_NAME, - } + kwargs["tool_choice"] = self._get_function_tool_choice(call_type=call_type) return kwargs + @staticmethod + def _strip_interception_metadata(kwargs: dict[str, Any]) -> None: + metadata = kwargs.get(_LITELLM_METADATA_KEY) + if not isinstance(metadata, dict): + return + filtered_metadata = { + key: value + for key, value in metadata.items() + if not is_interception_internal_key(key) + and not key.startswith("_agentic_loop") + and key != "max_agentic_loops" + } + if filtered_metadata: + kwargs[_LITELLM_METADATA_KEY] = filtered_metadata + else: + kwargs.pop(_LITELLM_METADATA_KEY, None) + + @staticmethod + def _write_interception_metadata(kwargs: dict[str, Any]) -> None: + metadata = kwargs.get(_LITELLM_METADATA_KEY) + metadata = dict(metadata) if isinstance(metadata, dict) else {} + for key in (_INTERCEPTION_ACTIVE_KEY, _SANDBOX_KEY, _CONVERTED_STREAM_KEY): + if key in kwargs: + metadata[key] = kwargs[key] + kwargs[_LITELLM_METADATA_KEY] = metadata + + @staticmethod + def _get_function_parameters() -> CodeExecutionFunctionParameters: + return { + "type": "object", + "properties": {"code": {"type": "string"}}, + "required": ["code"], + } + + def _get_function_tool( + self, call_type: CallTypes | None + ) -> CodeExecutionFunctionTool: + description = "Execute python code in a sandbox and return stdout." + if call_type in (CallTypes.completion, CallTypes.acompletion): + return { + "type": "function", + "function": { + "name": LITELLM_CODE_EXECUTION_TOOL_NAME, + "description": description, + "parameters": self._get_function_parameters(), + }, + } + return { + "type": "function", + "name": LITELLM_CODE_EXECUTION_TOOL_NAME, + "description": description, + "parameters": self._get_function_parameters(), + } + + @staticmethod + def _get_function_tool_choice( + call_type: CallTypes | None, + ) -> CodeExecutionFunctionToolChoice: + if call_type in (CallTypes.completion, CallTypes.acompletion): + return { + "type": "function", + "function": {"name": LITELLM_CODE_EXECUTION_TOOL_NAME}, + } + return { + "type": "function", + "name": LITELLM_CODE_EXECUTION_TOOL_NAME, + } + @staticmethod def _tool_choice_targets_code_interpreter(tool_choice: Any) -> bool: if not isinstance(tool_choice, dict): return False + function = tool_choice.get("function") return ( tool_choice.get("type") == "code_interpreter" or tool_choice.get("name") == "code_interpreter" + or tool_choice.get("name") == LITELLM_CODE_EXECUTION_TOOL_NAME + or ( + isinstance(function, dict) + and function.get("name") == LITELLM_CODE_EXECUTION_TOOL_NAME + ) ) def _resolve_provider(self, kwargs: dict[str, Any]) -> str | None: @@ -188,7 +337,12 @@ class CodeInterpreterInterceptionLogger(CustomLogger): ): return False, {} - tool_calls = self._extract_code_execution_tool_calls(response=response) + tool_calls = ( + self._extract_chat_completion_code_execution_tool_calls(response=response) + if kwargs.get("_agentic_loop_api_surface") + == CHAT_COMPLETION_AGENTIC_SURFACE + else self._extract_code_execution_tool_calls(response=response) + ) if not tool_calls: return False, {} @@ -206,15 +360,24 @@ class CodeInterpreterInterceptionLogger(CustomLogger): stream: bool, kwargs: dict, ) -> AgenticLoopPlan: + if kwargs.get("_agentic_loop_api_surface") == CHAT_COMPLETION_AGENTIC_SURFACE: + return await self._build_chat_completion_agentic_loop_plan( + tools=tools, + model=model, + messages=messages, + optional_params=anthropic_messages_optional_request_params, + kwargs=kwargs, + ) + await self._prune_expired_cache() - tool_calls = cast(list[dict[str, Any]], tools.get("tool_calls", [])) + tool_calls = cast(list[CodeExecutionToolCall], tools.get("tool_calls", [])) sandbox_key = kwargs.get(_SANDBOX_KEY) container, params = await self._get_or_create_container(cache_key=sandbox_key) try: - container_id = getattr(container, "id", None) + container_id = cast(str | None, getattr(container, "id", None)) input_list = self._normalize_messages(messages) - code_interpreter_calls = [] + code_interpreter_calls: list[CodeInterpreterCall] = [] for tool_call in tool_calls: arguments = tool_call.get("arguments", "") code = self._parse_code(arguments) @@ -256,9 +419,12 @@ class CodeInterpreterInterceptionLogger(CustomLogger): request_patch = AgenticLoopRequestPatch( model=model, messages=input_list, - tools=optional_params.get("tools"), - optional_params={k: v for k, v in optional_params.items() if k != "tools"}, - kwargs={k: v for k, v in kwargs.items() if k != "litellm_logging_obj"}, + tools=self._get_followup_tools( + tools=optional_params.get("tools"), + call_type=CallTypes.responses, + ), + optional_params=self._get_followup_optional_params(optional_params), + kwargs=self._filter_agentic_loop_kwargs(kwargs), ) return AgenticLoopPlan( @@ -271,12 +437,134 @@ class CodeInterpreterInterceptionLogger(CustomLogger): }, ) + async def _build_chat_completion_agentic_loop_plan( + self, + tools: dict[str, object], + model: str, + messages: list[dict], + optional_params: dict[str, object], + kwargs: dict[str, object], + ) -> AgenticLoopPlan: + await self._prune_expired_cache() + tool_calls = cast(list[CodeExecutionToolCall], tools.get("tool_calls", [])) + sandbox_key = cast(str | None, kwargs.get(_SANDBOX_KEY)) + container, params = await self._get_or_create_container(cache_key=sandbox_key) + + try: + container_id = cast(str | None, getattr(container, "id", None)) + tool_results = [ + await self._build_chat_completion_tool_result( + container=container, + params=params, + tool_call=tool_call, + container_id=container_id, + ) + for tool_call in tool_calls + ] + except Exception: + await self._delete_container_for_cache_key(sandbox_key) + raise + tool_messages = [result[0] for result in tool_results] + code_interpreter_calls = [result[1] for result in tool_results] + + request_patch = AgenticLoopRequestPatch( + model=model, + messages=list(messages) + + [self._build_chat_completion_assistant_message(tool_calls)] + + tool_messages, + tools=self._get_followup_tools( + tools=optional_params.get("tools"), + call_type=CallTypes.completion, + ), + optional_params=self._get_followup_optional_params(optional_params), + kwargs=self._filter_agentic_loop_kwargs(kwargs), + ) + + return AgenticLoopPlan( + run_agentic_loop=True, + request_patch=request_patch, + metadata={ + "tool_type": "code_interpreter", + "sandbox_key": sandbox_key or "", + "code_interpreter_calls": code_interpreter_calls, + "response_format": "openai", + }, + ) + + async def _build_chat_completion_tool_result( + self, + container: object, + params: dict[str, Any] | None, + tool_call: CodeExecutionToolCall, + container_id: str | None, + ) -> tuple[ChatCompletionToolMessage, CodeInterpreterCall]: + arguments = tool_call.get("arguments", "") + code = self._parse_code(arguments) + stdout = await self._run_tool_call( + container=container, params=params, arguments=arguments + ) + tool_call_id = ( + tool_call.get("id") or tool_call.get("call_id") or uuid.uuid4().hex + ) + return ( + { + "role": "tool", + "tool_call_id": tool_call_id, + "content": stdout, + }, + { + "id": f"ci_{uuid.uuid4().hex}", + "type": "code_interpreter_call", + "status": "completed", + "code": code, + "container_id": container_id, + "outputs": [{"type": "logs", "logs": stdout}] if stdout else [], + }, + ) + async def async_agentic_loop_cleanup_hook( self, plan: AgenticLoopPlan, kwargs: dict ) -> None: metadata = plan.metadata or {} if plan else {} await self._delete_container_for_cache_key(metadata.get("sandbox_key")) + @staticmethod + def _filter_agentic_loop_kwargs(kwargs: dict[str, object]) -> dict[str, object]: + return { + k: v + for k, v in kwargs.items() + if k not in {"litellm_logging_obj", "acompletion"} + and not is_interception_internal_key( + k, prefixes=NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES + ) + } + + def _get_followup_tools( + self, tools: object, call_type: CallTypes | None + ) -> list[dict[str, Any]] | None: + if not isinstance(tools, list): + return None + return [ + ( + self._get_function_tool(call_type=call_type) + if isinstance(tool, dict) and tool.get("type") == "code_interpreter" + else tool + ) + for tool in tools + ] + + def _get_followup_optional_params( + self, optional_params: dict[str, object] + ) -> dict[str, object]: + drop_tool_choice = self._tool_choice_targets_code_interpreter( + optional_params.get("tool_choice") + ) + return { + k: v + for k, v in optional_params.items() + if k != "tools" and not (k == "tool_choice" and drop_tool_choice) + } + async def async_post_agentic_loop_response_hook( self, response: Any, plan: AgenticLoopPlan, kwargs: dict ) -> Any: @@ -420,7 +708,9 @@ class CodeInterpreterInterceptionLogger(CustomLogger): return list(messages) return [] - def _extract_code_execution_tool_calls(self, response: Any) -> list[dict[str, Any]]: + def _extract_code_execution_tool_calls( + self, response: object + ) -> list[CodeExecutionToolCall]: if isinstance(response, dict): output = response.get("output", []) else: @@ -446,6 +736,82 @@ class CodeInterpreterInterceptionLogger(CustomLogger): if self._is_code_execution_call(item) ] + def _extract_chat_completion_code_execution_tool_calls( + self, response: ModelResponse | dict[str, Any] + ) -> list[CodeExecutionToolCall]: + model_response = self._to_model_response(response) + if model_response is None: + return [] + choices = model_response.choices or [] + if not choices: + return [] + message = choices[0].message + tool_calls = message.tool_calls or [] + + return [ + normalized + for tool_call in tool_calls + if (normalized := self._normalize_chat_completion_tool_call(tool_call)) + is not None + ] + + @staticmethod + def _normalize_chat_completion_tool_call( + tool_call: ChatCompletionMessageToolCall, + ) -> CodeExecutionToolCall | None: + if ( + tool_call.type != "function" + or tool_call.function.name != LITELLM_CODE_EXECUTION_TOOL_NAME + ): + return None + + arguments = tool_call.function.arguments + if isinstance(arguments, dict): + arguments = json.dumps(arguments) + elif not isinstance(arguments, str): + arguments = "" if arguments is None else str(arguments) + + return { + "id": tool_call.id, + "call_id": tool_call.id, + "type": "function", + "name": LITELLM_CODE_EXECUTION_TOOL_NAME, + "arguments": arguments, + } + + @staticmethod + def _build_chat_completion_assistant_message( + tool_calls: list[CodeExecutionToolCall], + ) -> ChatCompletionAssistantMessage: + return { + "role": "assistant", + "tool_calls": [ + cast( + ChatCompletionAssistantToolCall, + { + "id": tool_call.get("id"), + "type": "function", + "function": { + "name": LITELLM_CODE_EXECUTION_TOOL_NAME, + "arguments": tool_call.get("arguments", ""), + }, + }, + ) + for tool_call in tool_calls + ], + } + + @staticmethod + def _to_model_response( + response: ModelResponse | dict[str, Any], + ) -> ModelResponse | None: + if isinstance(response, ModelResponse): + return response + try: + return ModelResponse(**response) + except (TypeError, ValidationError): + return None + def _is_code_execution_call(self, item: Any) -> bool: if isinstance(item, dict): return ( diff --git a/litellm/litellm_core_utils/chat_completion_agentic_loop.py b/litellm/litellm_core_utils/chat_completion_agentic_loop.py new file mode 100644 index 00000000000..938e892bd50 --- /dev/null +++ b/litellm/litellm_core_utils/chat_completion_agentic_loop.py @@ -0,0 +1,332 @@ +# this is a patch to allow for agentic loops covering llm_http_handler.py and openai sdk based calling flows for the .completion() api + +import json +from typing import cast + +from litellm._logging import verbose_logger +from litellm.integrations.custom_logger import CustomLogger +from litellm.types.integrations.custom_logger import ( + CHAT_COMPLETION_AGENTIC_SURFACE, + NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES, + AgenticLoopPlan, + AgenticLoopRequestPatch, + is_interception_internal_key, +) +from litellm.types.utils import ModelResponse +from litellm.utils import CustomStreamWrapper + +_FOLLOWUP_INTERNAL_PARAMS = frozenset( + ( + "acompletion", + "litellm_logging_obj", + "custom_llm_provider", + "model_alias_map", + "stream_response", + "custom_prompt_dict", + "_agentic_loop_api_surface", + ) +) + + +def _gate_overridden(callback: CustomLogger) -> bool: + base = CustomLogger.async_should_run_agentic_loop + func = type(callback).async_should_run_agentic_loop + return getattr(func, "__func__", func) is not getattr(base, "__func__", base) + + +def _build_plan_overridden(callback: CustomLogger) -> bool: + base = CustomLogger.async_build_agentic_loop_plan + func = type(callback).async_build_agentic_loop_plan + return getattr(func, "__func__", func) is not getattr(base, "__func__", base) + + +def _post_hook_overridden(callback: CustomLogger) -> bool: + base = CustomLogger.async_post_agentic_loop_response_hook + func = type(callback).async_post_agentic_loop_response_hook + return getattr(func, "__func__", func) is not getattr(base, "__func__", base) + + +def _coerce_int(value: object, default: int) -> int: + return int(value) if isinstance(value, (int, str)) else default + + +def _agentic_loop_settings(kwargs: dict[str, object]) -> tuple[int, int, list[str]]: + depth = _coerce_int(kwargs.get("_agentic_loop_depth"), 0) + max_loops = max(_coerce_int(kwargs.get("max_agentic_loops"), 3), 1) + raw_fingerprints = kwargs.get("_agentic_loop_fingerprints") + fingerprints = ( + [str(fp) for fp in raw_fingerprints] + if isinstance(raw_fingerprints, list) + else [] + ) + return depth, max_loops, fingerprints + + +def _fingerprint_tools(tool_calls: object) -> str: + try: + return json.dumps(tool_calls, sort_keys=True, default=str) + except Exception: + return str(tool_calls) + + +def _check_agentic_loop_safety( + tool_calls: object, + fingerprints: list[str], + depth: int, + max_loops: int, + model: str, +) -> str: + fingerprint = _fingerprint_tools(tool_calls) + if fingerprint in fingerprints: + raise ValueError( + "Agentic loop detected repeated tool-call fingerprint; aborting rerun" + ) + if depth >= max_loops: + raise ValueError(f"Exceeded max_agentic_loops={max_loops} for model={model}") + return fingerprint + + +def _wrap_response_as_fake_stream(response: object) -> object: + if getattr(response, "object", None) == "chat.completion.chunk": + return response + if not hasattr(response, "choices"): + return response + from litellm.llms.base_llm.base_model_iterator import ( + convert_model_response_to_streaming, + ) + + return convert_model_response_to_streaming(cast(ModelResponse, response)) + + +def _add_agentic_loop_metadata(kwargs_for_followup: dict[str, object]) -> None: + metadata = kwargs_for_followup.get("litellm_metadata") + metadata = dict(metadata) if isinstance(metadata, dict) else {} + for key, value in kwargs_for_followup.items(): + if ( + key.startswith("_agentic_loop") + or key == "max_agentic_loops" + or is_interception_internal_key(key) + ): + metadata[key] = value + kwargs_for_followup["litellm_metadata"] = metadata + + +def _filter_followup_kwargs(source: dict[str, object]) -> dict[str, object]: + return { + k: v + for k, v in source.items() + if not is_interception_internal_key( + k, prefixes=NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES + ) + and k not in _FOLLOWUP_INTERNAL_PARAMS + } + + +async def _execute_chat_completion_agentic_plan( + *, + plan: AgenticLoopPlan, + callback: CustomLogger, + model: str, + optional_params: dict[str, object], + kwargs: dict[str, object], + logging_obj: object, + custom_llm_provider: str, + depth: int, + max_loops: int, + fingerprints: list[str], + fingerprint: str, +) -> object: + import litellm + + patch = plan.request_patch or AgenticLoopRequestPatch() + if patch.messages is None: + raise ValueError("Agentic loop plan missing patched messages") + + full_model_name = patch.model or model + if "/" not in full_model_name: + full_model_name = f"{custom_llm_provider}/{full_model_name}" + + optional_params_for_followup = {**optional_params, **patch.optional_params} + if patch.tools is not None: + optional_params_for_followup["tools"] = patch.tools + if "tool_choice" not in patch.optional_params: + optional_params_for_followup.pop("tool_choice", None) + + kwargs_for_followup = _filter_followup_kwargs(kwargs) + kwargs_for_followup.update( + { + k: v + for k, v in _filter_followup_kwargs(patch.kwargs).items() + if k not in optional_params_for_followup + } + ) + kwargs_for_followup["_agentic_loop_depth"] = depth + 1 + kwargs_for_followup["max_agentic_loops"] = max_loops + kwargs_for_followup["_agentic_loop_fingerprints"] = fingerprints + [fingerprint] + _add_agentic_loop_metadata(kwargs_for_followup) + + try: + response_followup = await litellm.acompletion( + model=full_model_name, + messages=patch.messages, + **optional_params_for_followup, + **kwargs_for_followup, + ) + if _post_hook_overridden(callback): + try: + response_followup = ( + await callback.async_post_agentic_loop_response_hook( + response=response_followup, plan=plan, kwargs=kwargs + ) + ) + except Exception as e: + _call_id = getattr(logging_obj, "litellm_call_id", "unknown") + verbose_logger.exception( + "LiteLLM.AgenticHookError: Exception in " + "async_post_agentic_loop_response_hook [call_id=%s model=%s]: %s", + _call_id, + model, + str(e), + ) + if kwargs.get("_code_interpreter_interception_converted_stream") and not depth: + return _wrap_response_as_fake_stream(response_followup) + return response_followup + finally: + try: + await callback.async_agentic_loop_cleanup_hook(plan=plan, kwargs=kwargs) + except Exception as e: + _call_id = getattr(logging_obj, "litellm_call_id", "unknown") + verbose_logger.exception( + "LiteLLM.AgenticHookError: Exception in " + "async_agentic_loop_cleanup_hook [call_id=%s model=%s]: %s", + _call_id, + model, + str(e), + ) + + +async def maybe_run_chat_completion_agentic_loop( + *, + response: ModelResponse, + model: str, + messages: list, + optional_params: dict, + kwargs: dict, + logging_obj: object, + custom_llm_provider: str, + stream: bool, +) -> ModelResponse | CustomStreamWrapper | None: + import litellm + + callbacks = litellm.callbacks + ( + getattr(logging_obj, "dynamic_success_callbacks", None) or [] + ) + depth, max_loops, fingerprints = _agentic_loop_settings(kwargs) + tools = optional_params.get("tools", []) + + for callback in callbacks: + if not isinstance(callback, CustomLogger): + continue + if not _gate_overridden(callback): + continue + + gate_kwargs = { + **kwargs, + "_agentic_loop_api_surface": CHAT_COMPLETION_AGENTIC_SURFACE, + "custom_llm_provider": custom_llm_provider, + } + try: + should_run, tool_calls = await callback.async_should_run_agentic_loop( + response=response, + model=model, + messages=messages, + tools=tools, + stream=stream, + custom_llm_provider=custom_llm_provider, + kwargs=gate_kwargs, + ) + except Exception as e: + verbose_logger.exception( + "LiteLLM.AgenticHookError: Exception in chat completion agentic gate: %s", + str(e), + ) + continue + + if not should_run: + continue + + fingerprint = _check_agentic_loop_safety( + tool_calls=tool_calls, + fingerprints=fingerprints, + depth=depth, + max_loops=max_loops, + model=model, + ) + + try: + plan_kwargs = { + **kwargs, + "_agentic_loop_api_surface": CHAT_COMPLETION_AGENTIC_SURFACE, + "custom_llm_provider": custom_llm_provider, + } + if not _build_plan_overridden(callback): + return await callback.async_run_agentic_loop( + tools=tool_calls, + model=model, + messages=messages, + response=response, + anthropic_messages_provider_config=None, + anthropic_messages_optional_request_params=optional_params, + logging_obj=logging_obj, + stream=stream, + kwargs=plan_kwargs, + ) + + plan = await callback.async_build_agentic_loop_plan( + tools=tool_calls, + model=model, + messages=messages, + response=response, + anthropic_messages_provider_config=None, + anthropic_messages_optional_request_params=optional_params, + logging_obj=logging_obj, + stream=stream, + kwargs=plan_kwargs, + ) + + if plan.response_override is not None: + return plan.response_override + if plan.terminate: + return response + if not plan.run_agentic_loop: + continue + + return await _execute_chat_completion_agentic_plan( + plan=plan, + callback=callback, + model=model, + optional_params=optional_params, + kwargs=kwargs, + logging_obj=logging_obj, + custom_llm_provider=custom_llm_provider, + depth=depth, + max_loops=max_loops, + fingerprints=fingerprints, + fingerprint=fingerprint, + ) + except Exception as e: + verbose_logger.exception( + "LiteLLM.AgenticHookError: Exception in chat completion agentic hooks: %s", + str(e), + ) + + if ( + kwargs.get("_code_interpreter_interception_converted_stream") + and not depth + and hasattr(response, "choices") + ): + return cast( + "ModelResponse | CustomStreamWrapper", + _wrap_response_as_fake_stream(response), + ) + return None diff --git a/litellm/llms/apiserpent/search/transformation.py b/litellm/llms/apiserpent/search/transformation.py index 1eb7d34c875..bc11875ba12 100644 --- a/litellm/llms/apiserpent/search/transformation.py +++ b/litellm/llms/apiserpent/search/transformation.py @@ -53,7 +53,13 @@ class APISerpentSearchConfig(BaseSearchConfig): api_base: Optional[str] = None, **kwargs, ) -> Dict: - api_key = api_key or get_secret_str("APISERPENT_API_KEY") + api_key = self.resolve_server_api_key( + caller_api_key=api_key, + caller_api_base=api_base, + key_env_vars=("APISERPENT_API_KEY",), + base_env_var="APISERPENT_API_BASE", + default_api_base=APISERPENT_BASE, + ) if not api_key: raise ValueError( "APISERPENT_API_KEY is not set. Set `APISERPENT_API_KEY` environment variable." diff --git a/litellm/llms/base_llm/search/transformation.py b/litellm/llms/base_llm/search/transformation.py index 4dfe86685fb..1581d8bb064 100644 --- a/litellm/llms/base_llm/search/transformation.py +++ b/litellm/llms/base_llm/search/transformation.py @@ -3,11 +3,13 @@ Base Search transformation configuration. """ from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union +from urllib.parse import urlsplit import httpx from pydantic import PrivateAttr from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.secret_managers.main import get_secret_str from litellm.types.llms.base import LiteLLMPydanticObjectBase if TYPE_CHECKING: @@ -16,6 +18,29 @@ else: LiteLLMLoggingObj = Any +def _search_host(url: str) -> str: + return urlsplit(url).netloc.lower() + + +def _is_trusted_search_api_base( + caller_api_base: str, + default_api_base: str | None, + base_env_var: str | None, +) -> bool: + candidate = _search_host(caller_api_base) + if not candidate: + return False + trusted = { + _search_host(base) + for base in ( + default_api_base, + get_secret_str(base_env_var) if base_env_var else None, + ) + if base + } + return candidate in trusted + + class SearchResult(LiteLLMPydanticObjectBase): """Single search result.""" @@ -86,6 +111,60 @@ class BaseSearchConfig: "max_tokens_per_page", } + def _assert_trusted_api_base_for_server_credential( + self, + caller_api_base: str | None, + default_api_base: str | None, + base_env_var: str | None, + credential_name: str, + ) -> None: + """ + Block sending a server-managed credential to a caller-chosen host. + + A caller-supplied api_base is honored when constructing the request URL, so + falling back to a server-configured secret while the caller controls the host + leaks that secret. The provider default and the operator's own api_base + override are the only trusted destinations for a server-managed credential. + """ + if not caller_api_base: + return + if _is_trusted_search_api_base(caller_api_base, default_api_base, base_env_var): + return + raise ValueError( + f"Refusing to send the server-configured {credential_name} to the " + f"caller-supplied api_base '{caller_api_base}'. Pass an explicit api_key " + f"when overriding api_base for this search provider." + ) + + def resolve_server_api_key( + self, + *, + caller_api_key: str | None, + caller_api_base: str | None, + key_env_vars: tuple[str, ...], + base_env_var: str | None, + default_api_base: str | None, + ) -> str | None: + """ + Resolve a single-secret search API key, falling back to a server-managed + secret only when the request targets a trusted host. + + Returns the caller's key when provided, otherwise the first set + server-managed secret (or None when none is set, for keyless providers). + """ + if caller_api_key: + return caller_api_key + server_key = next( + (key for key in (get_secret_str(var) for var in key_env_vars) if key), + None, + ) + if server_key is None: + return None + self._assert_trusted_api_base_for_server_credential( + caller_api_base, default_api_base, base_env_var, key_env_vars[0] + ) + return server_key + def validate_environment( self, headers: Dict, diff --git a/litellm/llms/brave/search/transformation.py b/litellm/llms/brave/search/transformation.py index 9dfcd6bc75a..8ffe7dcb126 100644 --- a/litellm/llms/brave/search/transformation.py +++ b/litellm/llms/brave/search/transformation.py @@ -115,7 +115,13 @@ class BraveSearchConfig(BaseSearchConfig): """ Validate environment and return headers. """ - api_key = api_key or get_secret_str("BRAVE_API_KEY") + api_key = self.resolve_server_api_key( + caller_api_key=api_key, + caller_api_base=api_base, + key_env_vars=("BRAVE_API_KEY",), + base_env_var="BRAVE_API_BASE", + default_api_base=self.BRAVE_API_BASE, + ) if not api_key: raise ValueError( diff --git a/litellm/llms/cloudflare/chat/transformation.py b/litellm/llms/cloudflare/chat/transformation.py index 66e253f304d..68f08741cc5 100644 --- a/litellm/llms/cloudflare/chat/transformation.py +++ b/litellm/llms/cloudflare/chat/transformation.py @@ -1,26 +1,15 @@ -import json -import time -from typing import AsyncIterator, Iterator, List, Optional, Union +from typing import List, Optional, Union import httpx -import litellm -from litellm.litellm_core_utils.url_utils import encode_url_path_segments -from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator -from litellm.llms.base_llm.chat.transformation import ( - BaseConfig, - BaseLLMException, - LiteLLMLoggingObj, +from litellm._logging import verbose_logger +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig +from litellm.secret_managers.main import ( + get_secret_str, + normalize_nonempty_secret_str, ) -from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import AllMessageValues -from litellm.types.utils import ( - ChatCompletionToolCallChunk, - ChatCompletionUsageBlock, - GenericStreamingChunk, - ModelResponse, - Usage, -) class CloudflareError(BaseLLMException): @@ -34,26 +23,46 @@ class CloudflareError(BaseLLMException): message=message, request=self.request, response=self.response, - ) # Call the base class constructor with the parameters it needs + ) -class CloudflareChatConfig(BaseConfig): - max_tokens: Optional[int] = None - stream: Optional[bool] = None - - def __init__( +class CloudflareChatConfig(OpenAIGPTConfig): + def get_complete_url( self, - max_tokens: Optional[int] = None, + api_base: Optional[str], + api_key: Optional[str], + model: str, + optional_params: dict, + litellm_params: dict, stream: Optional[bool] = None, - ) -> None: - locals_ = locals().copy() - for key, value in locals_.items(): - if key != "self" and value is not None: - setattr(self.__class__, key, value) + ) -> str: + return super().get_complete_url( + api_base=self._resolve_api_base(api_base), + api_key=api_key, + model=model, + optional_params=optional_params, + litellm_params=litellm_params, + stream=stream, + ) - @classmethod - def get_config(cls): - return super().get_config() + @staticmethod + def _resolve_api_base(api_base: Optional[str]) -> str: + if not api_base: + account_id = normalize_nonempty_secret_str( + get_secret_str("CLOUDFLARE_ACCOUNT_ID") + ) + if account_id is None: + raise ValueError( + "Missing CLOUDFLARE_ACCOUNT_ID - set CLOUDFLARE_ACCOUNT_ID in the environment or pass api_base explicitly" + ) + return f"https://api.cloudflare.com/client/v4/accounts/{account_id}/ai/v1" + trimmed = api_base.rstrip("/") + if trimmed.endswith("/ai/run"): + verbose_logger.warning( + "Cloudflare api_base ending in '/ai/run' is the legacy Workers AI path and no longer serves OpenAI-compatible requests; rewriting to the '/ai/v1' endpoint" + ) + return f"{trimmed[: -len('/ai/run')]}/ai/v1" + return api_base def validate_environment( self, @@ -67,107 +76,18 @@ class CloudflareChatConfig(BaseConfig): ) -> dict: if api_key is None: raise ValueError( - "Missing CloudflareError API Key - A call is being made to cloudflare but no key is set either in the environment variables or via params" + "Missing Cloudflare API Key - A call is being made to cloudflare but no key is set either in the environment variables or via params" ) - headers = { - "accept": "application/json", - "content-type": "apbplication/json", - "Authorization": "Bearer " + api_key, - } - return headers - - def get_complete_url( - self, - api_base: Optional[str], - api_key: Optional[str], - model: str, - optional_params: dict, - litellm_params: dict, - stream: Optional[bool] = None, - ) -> str: - if api_base is None: - account_id = get_secret_str("CLOUDFLARE_ACCOUNT_ID") - api_base = ( - f"https://api.cloudflare.com/client/v4/accounts/{account_id}/ai/run/" - ) - encoded_model = encode_url_path_segments(model, field_name="model") - return api_base + encoded_model - - def get_supported_openai_params(self, model: str) -> List[str]: - return [ - "stream", - "max_tokens", - ] - - def map_openai_params( - self, - non_default_params: dict, - optional_params: dict, - model: str, - drop_params: bool, - ) -> dict: - supported_openai_params = self.get_supported_openai_params(model=model) - for param, value in non_default_params.items(): - if param == "max_completion_tokens": - optional_params["max_tokens"] = value - elif param in supported_openai_params: - optional_params[param] = value - return optional_params - - def transform_request( - self, - model: str, - messages: List[AllMessageValues], - optional_params: dict, - litellm_params: dict, - headers: dict, - ) -> dict: - config = litellm.CloudflareChatConfig.get_config() - for k, v in config.items(): - if k not in optional_params: - optional_params[k] = v - - data = { - "messages": messages, - **optional_params, - } - return data - - def transform_response( - self, - model: str, - raw_response: httpx.Response, - model_response: ModelResponse, - logging_obj: LiteLLMLoggingObj, - request_data: dict, - messages: List[AllMessageValues], - optional_params: dict, - litellm_params: dict, - encoding: str, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, - ) -> ModelResponse: - completion_response = raw_response.json() - - # Support both "response" and "response_text" keys (newer models like Nemotron use "response_text") - result = completion_response["result"] - model_response.choices[0].message.content = result.get("response") if result.get("response") is not None else result.get("response_text", "") # type: ignore - - prompt_tokens = litellm.utils.get_token_count(messages=messages, model=model) - completion_tokens = len( - encoding.encode(model_response["choices"][0]["message"].get("content", "")) + return super().validate_environment( + headers=headers, + model=model, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + api_key=api_key, + api_base=api_base, ) - model_response.created = int(time.time()) - model_response.model = "cloudflare/" + model - usage = Usage( - prompt_tokens=prompt_tokens, - completion_tokens=completion_tokens, - total_tokens=prompt_tokens + completion_tokens, - ) - setattr(model_response, "usage", usage) - return model_response - def get_error_class( self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] ) -> BaseLLMException: @@ -175,48 +95,3 @@ class CloudflareChatConfig(BaseConfig): status_code=status_code, message=error_message, ) - - def get_model_response_iterator( - self, - streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse], - sync_stream: bool, - json_mode: Optional[bool] = False, - ): - return CloudflareChatResponseIterator( - streaming_response=streaming_response, - sync_stream=sync_stream, - json_mode=json_mode, - ) - - -class CloudflareChatResponseIterator(BaseModelResponseIterator): - def chunk_parser(self, chunk: dict) -> GenericStreamingChunk: - try: - text = "" - tool_use: Optional[ChatCompletionToolCallChunk] = None - is_finished = False - finish_reason = "" - usage: Optional[ChatCompletionUsageBlock] = None - provider_specific_fields = None - - index = int(chunk.get("index", 0)) - - if "response" in chunk and chunk["response"] is not None: - text = chunk["response"] - elif "response_text" in chunk and chunk["response_text"] is not None: - text = chunk["response_text"] - - returned_chunk = GenericStreamingChunk( - text=text, - tool_use=tool_use, - is_finished=is_finished, - finish_reason=finish_reason, - usage=usage, - index=index, - provider_specific_fields=provider_specific_fields, - ) - - return returned_chunk - - except json.JSONDecodeError: - raise ValueError(f"Failed to decode JSON from chunk: {chunk}") diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 138f2410c89..948c90f9f99 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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) diff --git a/litellm/llms/dataforseo/search/transformation.py b/litellm/llms/dataforseo/search/transformation.py index 27c10d740b5..701db586b72 100644 --- a/litellm/llms/dataforseo/search/transformation.py +++ b/litellm/llms/dataforseo/search/transformation.py @@ -61,9 +61,18 @@ class DataForSEOSearchConfig(BaseSearchConfig): password = get_secret_str("DATAFORSEO_PASSWORD") # If api_key is provided in "login:password" format, use it + caller_supplied_credentials = bool(api_key and ":" in api_key) if api_key and ":" in api_key: login, password = api_key.split(":", 1) + if not caller_supplied_credentials and login and password: + self._assert_trusted_api_base_for_server_credential( + api_base, + self.DATAFORSEO_API_BASE, + "DATAFORSEO_API_BASE", + "DATAFORSEO_LOGIN", + ) + if not login: raise ValueError( "DATAFORSEO_LOGIN is not set. Set `DATAFORSEO_LOGIN` environment variable or pass credentials in api_key parameter." diff --git a/litellm/llms/exa_ai/search/transformation.py b/litellm/llms/exa_ai/search/transformation.py index 7a34ededa6b..5cfd14aeaa9 100644 --- a/litellm/llms/exa_ai/search/transformation.py +++ b/litellm/llms/exa_ai/search/transformation.py @@ -65,7 +65,13 @@ class ExaAISearchConfig(BaseSearchConfig): """ Validate environment and return headers. """ - api_key = api_key or get_secret_str("EXA_API_KEY") + api_key = self.resolve_server_api_key( + caller_api_key=api_key, + caller_api_base=api_base, + key_env_vars=("EXA_API_KEY",), + base_env_var="EXA_API_BASE", + default_api_base=self.EXA_AI_API_BASE, + ) if not api_key: raise ValueError( "EXA_API_KEY is not set. Set `EXA_API_KEY` environment variable." diff --git a/litellm/llms/fastcrw/search/transformation.py b/litellm/llms/fastcrw/search/transformation.py index ce702266e7b..b571a659cac 100644 --- a/litellm/llms/fastcrw/search/transformation.py +++ b/litellm/llms/fastcrw/search/transformation.py @@ -57,7 +57,13 @@ class FastCRWSearchConfig(BaseSearchConfig): """ Validate environment and return headers. """ - api_key = api_key or get_secret_str("CRW_API_KEY") + api_key = self.resolve_server_api_key( + caller_api_key=api_key, + caller_api_base=api_base, + key_env_vars=("CRW_API_KEY",), + base_env_var="CRW_API_BASE", + default_api_base=self.FASTCRW_API_BASE, + ) if not api_key: raise ValueError( "CRW_API_KEY is not set. Set `CRW_API_KEY` environment variable." diff --git a/litellm/llms/firecrawl/search/transformation.py b/litellm/llms/firecrawl/search/transformation.py index 18cf1d28c4d..7e01ba58706 100644 --- a/litellm/llms/firecrawl/search/transformation.py +++ b/litellm/llms/firecrawl/search/transformation.py @@ -61,7 +61,13 @@ class FirecrawlSearchConfig(BaseSearchConfig): """ Validate environment and return headers. """ - api_key = api_key or get_secret_str("FIRECRAWL_API_KEY") + api_key = self.resolve_server_api_key( + caller_api_key=api_key, + caller_api_base=api_base, + key_env_vars=("FIRECRAWL_API_KEY",), + base_env_var="FIRECRAWL_API_BASE", + default_api_base=self.FIRECRAWL_API_BASE, + ) if not api_key: raise ValueError( "FIRECRAWL_API_KEY is not set. Set `FIRECRAWL_API_KEY` environment variable." diff --git a/litellm/llms/google_pse/search/transformation.py b/litellm/llms/google_pse/search/transformation.py index a8aa109cbf0..5cd3f2085a8 100644 --- a/litellm/llms/google_pse/search/transformation.py +++ b/litellm/llms/google_pse/search/transformation.py @@ -85,7 +85,13 @@ class GooglePSESearchConfig(BaseSearchConfig): Google PSE uses API key as a query parameter, not in headers. This method is called but headers are not used for authentication. """ - api_key = api_key or get_secret_str("GOOGLE_PSE_API_KEY") + api_key = self.resolve_server_api_key( + caller_api_key=api_key, + caller_api_base=api_base, + key_env_vars=("GOOGLE_PSE_API_KEY",), + base_env_var="GOOGLE_PSE_API_BASE", + default_api_base=self.GOOGLE_PSE_API_BASE, + ) if not api_key: raise ValueError( "GOOGLE_PSE_API_KEY is not set. Set `GOOGLE_PSE_API_KEY` environment variable." @@ -137,6 +143,7 @@ class GooglePSESearchConfig(BaseSearchConfig): query: Union[str, List[str]], optional_params: dict, api_key: Optional[str] = None, + api_base: str | None = None, search_engine_id: Optional[str] = None, **kwargs, ) -> Dict: @@ -165,8 +172,16 @@ class GooglePSESearchConfig(BaseSearchConfig): # Google PSE only supports single string queries query = " ".join(query) - # Get API credentials - api_key = api_key or get_secret_str("GOOGLE_PSE_API_KEY") + # Get API credentials. The key is sent as a query param to api_base, so + # resolve it host-aware to avoid leaking a server-managed key to a + # caller-supplied host. + api_key = self.resolve_server_api_key( + caller_api_key=api_key, + caller_api_base=api_base, + key_env_vars=("GOOGLE_PSE_API_KEY",), + base_env_var="GOOGLE_PSE_API_BASE", + default_api_base=self.GOOGLE_PSE_API_BASE, + ) search_engine_id = search_engine_id or get_secret_str("GOOGLE_PSE_ENGINE_ID") if not api_key: diff --git a/litellm/llms/linkup/search/transformation.py b/litellm/llms/linkup/search/transformation.py index 2b17d5642ac..d27ae038f9e 100644 --- a/litellm/llms/linkup/search/transformation.py +++ b/litellm/llms/linkup/search/transformation.py @@ -61,7 +61,13 @@ class LinkupSearchConfig(BaseSearchConfig): """ Validate environment and return headers. """ - api_key = api_key or get_secret_str("LINKUP_API_KEY") + api_key = self.resolve_server_api_key( + caller_api_key=api_key, + caller_api_base=api_base, + key_env_vars=("LINKUP_API_KEY",), + base_env_var="LINKUP_API_BASE", + default_api_base=self.LINKUP_API_BASE, + ) if not api_key: raise ValueError( "LINKUP_API_KEY is not set. Set `LINKUP_API_KEY` environment variable." diff --git a/litellm/llms/parallel_ai/search/transformation.py b/litellm/llms/parallel_ai/search/transformation.py index 85602bf1d86..35a0d84df40 100644 --- a/litellm/llms/parallel_ai/search/transformation.py +++ b/litellm/llms/parallel_ai/search/transformation.py @@ -67,10 +67,12 @@ class ParallelAISearchConfig(BaseSearchConfig): api_base: Optional[str] = None, **kwargs, ) -> Dict: - api_key = ( - api_key - or get_secret_str("PARALLEL_AI_API_KEY") - or get_secret_str("PARALLEL_API_KEY") + api_key = self.resolve_server_api_key( + caller_api_key=api_key, + caller_api_base=api_base, + key_env_vars=("PARALLEL_AI_API_KEY", "PARALLEL_API_KEY"), + base_env_var="PARALLEL_AI_API_BASE", + default_api_base=self.PARALLEL_AI_API_BASE, ) if not api_key: raise ValueError( diff --git a/litellm/llms/perplexity/search/transformation.py b/litellm/llms/perplexity/search/transformation.py index ea96f87957c..55de52c5384 100644 --- a/litellm/llms/perplexity/search/transformation.py +++ b/litellm/llms/perplexity/search/transformation.py @@ -50,7 +50,13 @@ class PerplexitySearchConfig(BaseSearchConfig): """ Validate environment and return headers. """ - api_key = api_key or get_secret_str("PERPLEXITYAI_API_KEY") + api_key = self.resolve_server_api_key( + caller_api_key=api_key, + caller_api_base=api_base, + key_env_vars=("PERPLEXITYAI_API_KEY",), + base_env_var="PERPLEXITY_API_BASE", + default_api_base=self.PERPLEXITY_API_BASE, + ) if not api_key: raise ValueError( "PERPLEXITYAI_API_KEY is not set. Set `PERPLEXITYAI_API_KEY` environment variable." diff --git a/litellm/llms/searchapi/search/transformation.py b/litellm/llms/searchapi/search/transformation.py index c04e1377f9c..ae8413684cc 100644 --- a/litellm/llms/searchapi/search/transformation.py +++ b/litellm/llms/searchapi/search/transformation.py @@ -74,7 +74,13 @@ class SearchAPIConfig(BaseSearchConfig): """ Validate environment and return headers. """ - api_key = api_key or get_secret_str("SEARCHAPI_API_KEY") + api_key = self.resolve_server_api_key( + caller_api_key=api_key, + caller_api_base=api_base, + key_env_vars=("SEARCHAPI_API_KEY",), + base_env_var="SEARCHAPI_API_BASE", + default_api_base=self.SEARCHAPI_API_BASE, + ) if not api_key: raise ValueError( @@ -114,6 +120,7 @@ class SearchAPIConfig(BaseSearchConfig): query: Union[str, List[str]], optional_params: dict, api_key: Optional[str] = None, + api_base: str | None = None, search_engine_id: Optional[str] = None, **kwargs, ) -> Dict: @@ -137,8 +144,16 @@ class SearchAPIConfig(BaseSearchConfig): if isinstance(query, list): query = " ".join(query) - # Get API key from parameter or environment - api_key = api_key or get_secret_str("SEARCHAPI_API_KEY") + # Get API key from parameter or environment. The key is sent as a query + # param to api_base, so resolve it host-aware to avoid leaking a + # server-managed key to a caller-supplied host. + api_key = self.resolve_server_api_key( + caller_api_key=api_key, + caller_api_base=api_base, + key_env_vars=("SEARCHAPI_API_KEY",), + base_env_var="SEARCHAPI_API_BASE", + default_api_base=self.SEARCHAPI_API_BASE, + ) if not api_key: raise ValueError( "SEARCHAPI_API_KEY is not set. Set `SEARCHAPI_API_KEY` environment variable." diff --git a/litellm/llms/searxng/search/transformation.py b/litellm/llms/searxng/search/transformation.py index ee6f3895721..ff68be5709e 100644 --- a/litellm/llms/searxng/search/transformation.py +++ b/litellm/llms/searxng/search/transformation.py @@ -61,7 +61,13 @@ class SearXNGSearchConfig(BaseSearchConfig): Some instances may require authentication via headers. """ # SearXNG typically doesn't require API keys, but support optional auth - api_key = api_key or get_secret_str("SEARXNG_API_KEY") + api_key = self.resolve_server_api_key( + caller_api_key=api_key, + caller_api_base=api_base, + key_env_vars=("SEARXNG_API_KEY",), + base_env_var="SEARXNG_API_BASE", + default_api_base=None, + ) if api_key: headers["Authorization"] = f"Bearer {api_key}" headers["Content-Type"] = "application/json" diff --git a/litellm/llms/serper/search/transformation.py b/litellm/llms/serper/search/transformation.py index 0daccbe652b..dd43f2d2dc9 100644 --- a/litellm/llms/serper/search/transformation.py +++ b/litellm/llms/serper/search/transformation.py @@ -55,7 +55,13 @@ class SerperSearchConfig(BaseSearchConfig): """ Validate environment and return headers. """ - api_key = api_key or get_secret_str("SERPER_API_KEY") + api_key = self.resolve_server_api_key( + caller_api_key=api_key, + caller_api_base=api_base, + key_env_vars=("SERPER_API_KEY",), + base_env_var="SERPER_API_BASE", + default_api_base=self.SERPER_API_BASE, + ) if not api_key: raise ValueError( "SERPER_API_KEY is not set. Set `SERPER_API_KEY` environment variable." diff --git a/litellm/llms/tavily/search/transformation.py b/litellm/llms/tavily/search/transformation.py index ec96db96f36..647cfb5fa84 100644 --- a/litellm/llms/tavily/search/transformation.py +++ b/litellm/llms/tavily/search/transformation.py @@ -64,7 +64,13 @@ class TavilySearchConfig(BaseSearchConfig): """ Validate environment and return headers. """ - api_key = api_key or get_secret_str("TAVILY_API_KEY") + api_key = self.resolve_server_api_key( + caller_api_key=api_key, + caller_api_base=api_base, + key_env_vars=("TAVILY_API_KEY",), + base_env_var="TAVILY_API_BASE", + default_api_base=self.TAVILY_API_BASE, + ) if not api_key: raise ValueError( "TAVILY_API_KEY is not set. Set `TAVILY_API_KEY` environment variable." diff --git a/litellm/llms/tinyfish/search/transformation.py b/litellm/llms/tinyfish/search/transformation.py index c4949380e3a..b92f7ca1aff 100644 --- a/litellm/llms/tinyfish/search/transformation.py +++ b/litellm/llms/tinyfish/search/transformation.py @@ -67,7 +67,13 @@ class TinyfishSearchConfig(BaseSearchConfig): api_base: str | None = None, **kwargs: object, ) -> dict[str, str]: - resolved_key = api_key or get_secret_str("TINYFISH_API_KEY") + resolved_key = self.resolve_server_api_key( + caller_api_key=api_key, + caller_api_base=api_base, + key_env_vars=("TINYFISH_API_KEY",), + base_env_var="TINYFISH_API_BASE", + default_api_base=self.TINYFISH_API_BASE, + ) if not resolved_key: raise ValueError( "TINYFISH_API_KEY is not set. Set `TINYFISH_API_KEY` environment variable." diff --git a/litellm/llms/you_com/search/transformation.py b/litellm/llms/you_com/search/transformation.py index 3c94b991735..0c7916e4c05 100644 --- a/litellm/llms/you_com/search/transformation.py +++ b/litellm/llms/you_com/search/transformation.py @@ -64,7 +64,13 @@ class YouComSearchConfig(BaseSearchConfig): endpoint with the `X-API-Key` header. Otherwise fall through to the keyless free tier; no auth header is required. """ - api_key = api_key or get_secret_str("YOUCOM_API_KEY") + api_key = self.resolve_server_api_key( + caller_api_key=api_key, + caller_api_base=api_base, + key_env_vars=("YOUCOM_API_KEY",), + base_env_var="YOUCOM_API_BASE", + default_api_base=self.YOU_COM_API_BASE, + ) headers["Content-Type"] = "application/json" # Pin Accept-Encoding to identity: the keyless `api.you.com/v1/agents/search` # endpoint advertises gzip content-encoding but returns body bytes the diff --git a/litellm/main.py b/litellm/main.py index 4fade2ac4b0..a80109f4bff 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -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( diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 4f022e1f882..1dc984d5fda 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -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", diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index b27082c361a..3a9ef8db804 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -10,7 +10,7 @@ import os import re from functools import partial from io import IOBase -from typing import Any, Coroutine, Dict, Optional, Union +from typing import Any, Callable, Coroutine, Dict, Optional, Union, cast import httpx @@ -20,6 +20,7 @@ from litellm.constants import request_timeout from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig, OCRResponse from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler +from litellm.ocr.rust_bridge import RustOcr, load_rust_ocr, rust_ocr_enabled from litellm.types.router import GenericLiteLLMParams from litellm.utils import ProviderConfigManager, client @@ -28,6 +29,82 @@ base_llm_http_handler = BaseLLMHTTPHandler() ################################################# +def _timeout_to_seconds( + timeout: Optional[Union[float, httpx.Timeout]], +) -> Optional[float]: + """Convert the Python OCR timeout to a single seconds value for the Rust bridge. + + The Rust HTTP client takes one duration; ``httpx.Timeout`` carries separate + connect/read/write/pool values, so pick the read deadline as the closest + analog to a total-request timeout. + """ + if timeout is None: + return None + if isinstance(timeout, httpx.Timeout): + return timeout.read + return float(timeout) + + +def _run_rust_ocr( + rust_ocr: RustOcr, + logging_obj: LiteLLMLoggingObj, + provider_config: BaseOCRConfig, + resolve_api_key: Callable[[str], Optional[str]], + model: str, + document: dict[str, object], + api_key: Optional[str], + api_base: Optional[str], + optional_params: dict[str, object], + litellm_params: dict[str, object], + timeout_seconds: Optional[float], +) -> OCRResponse: + """Run the Mistral OCR call through the Rust bridge and wrap the result. + + Resolves the key the same way the Python path does so secret-manager backends + (AWS/Azure/GCP/Vault) work; the Rust bridge's own fallback only reads the + process environment. The request that Rust actually sends (resolved URL and + headers) is mirrored into pre_call so logs match the wire. Dependencies are + injected so this stays unit-testable without patching module globals. + """ + resolved_api_key = api_key or resolve_api_key("MISTRAL_API_KEY") + resolved_headers = provider_config.validate_environment( + headers={}, + model=model, + api_key=resolved_api_key, + api_base=api_base, + litellm_params=litellm_params, + ) + resolved_complete_url = provider_config.get_complete_url( + api_base=api_base, + model=model, + optional_params=optional_params, + litellm_params=litellm_params, + ) + logging_obj.pre_call( + input="OCR document processing", + api_key=resolved_api_key, + additional_args={ + "complete_input_dict": { + "model": model, + "document": document, + **optional_params, + }, + "api_base": resolved_complete_url, + "headers": resolved_headers, + }, + ) + return OCRResponse.model_validate( + rust_ocr( + model=model, + document=document, + api_key=resolved_api_key, + api_base=api_base, + optional_params=optional_params, + timeout_seconds=timeout_seconds, + ) + ) + + @client async def aocr( model: str, @@ -220,7 +297,7 @@ def ocr( """ local_vars = locals() try: - litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore + litellm_logging_obj = cast(LiteLLMLoggingObj, kwargs.pop("litellm_logging_obj")) litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("aocr", False) is True @@ -261,7 +338,6 @@ def ocr( if dynamic_api_base: api_base = dynamic_api_base - # Get provider config ocr_provider_config: Optional[BaseOCRConfig] = ( ProviderConfigManager.get_provider_ocr_config( model=model, @@ -278,17 +354,14 @@ def ocr( f"OCR call - model: {model}, provider: {custom_llm_provider}" ) - # Get litellm params using GenericLiteLLMParams (same as responses API) litellm_params = GenericLiteLLMParams(**kwargs) - # Extract OCR-specific parameters from kwargs supported_params = ocr_provider_config.get_supported_ocr_params(model=model) non_default_params = {} for param in supported_params: if param in kwargs: non_default_params[param] = kwargs.pop(param) - # Map parameters to provider-specific format optional_params = ocr_provider_config.map_ocr_params( non_default_params=non_default_params, optional_params={}, @@ -297,7 +370,8 @@ def ocr( verbose_logger.debug(f"OCR optional_params after mapping: {optional_params}") - # Pre Call logging + effective_timeout = timeout or request_timeout + litellm_logging_obj.update_from_kwargs( kwargs=kwargs, model=model, @@ -309,12 +383,35 @@ def ocr( custom_llm_provider=custom_llm_provider, ) - # Call the handler - pass document dict directly + # Optional Rust path: hand the whole Mistral OCR call to the Rust bridge. + if custom_llm_provider == "mistral" and rust_ocr_enabled(): + rust_ocr = load_rust_ocr() + if rust_ocr is None: + verbose_logger.debug( + "Rust OCR bridge unavailable; falling back to Python path" + ) + else: + from litellm.secret_managers.main import get_secret_str + + return _run_rust_ocr( + rust_ocr=rust_ocr, + logging_obj=litellm_logging_obj, + provider_config=ocr_provider_config, + resolve_api_key=get_secret_str, + model=model, + document=document, + api_key=api_key, + api_base=api_base, + optional_params=optional_params, + litellm_params=dict(litellm_params), + timeout_seconds=_timeout_to_seconds(effective_timeout), + ) + response = base_llm_http_handler.ocr( model=model, - document=document, # Pass the entire document dict + document=document, optional_params=optional_params, - timeout=timeout or request_timeout, + timeout=effective_timeout, logging_obj=litellm_logging_obj, api_key=api_key, api_base=api_base, diff --git a/litellm/ocr/rust_bridge.py b/litellm/ocr/rust_bridge.py new file mode 100644 index 00000000000..61f9e8ca69a --- /dev/null +++ b/litellm/ocr/rust_bridge.py @@ -0,0 +1,74 @@ +""" +Optional Rust-backed OCR path. + +Enable with ``litellm.use_litellm_rust()``; the sync ``litellm.ocr()`` entrypoint +then routes supported Mistral calls through the compiled ``litellm_python_bridge`` +extension, which performs the whole OCR call (URL, headers, HTTP, parse) in Rust. + +No module-level ``litellm`` imports keep this a leaf so ``litellm/ocr/main.py`` +can import it statically without forming an import cycle. +""" + +from __future__ import annotations + +from typing import Final, Protocol, cast + + +class RustOcr(Protocol): + """Signature of the compiled ``litellm_python_bridge.ocr`` entrypoint.""" + + def __call__( + self, + model: str, + document: dict[str, object], + api_key: str | None, + api_base: str | None, + optional_params: dict[str, object], + timeout_seconds: float | None, + ) -> dict[str, object]: ... + + +class _Unset: + """Sentinel type so ``ocr=None`` can clear a prior injection while omission preserves it.""" + + +_UNSET: Final[_Unset] = _Unset() + +_rust_ocr_enabled = False +_rust_ocr_impl: RustOcr | None = None + + +def use_litellm_rust( + enabled: bool = True, *, ocr: RustOcr | None | _Unset = _UNSET +) -> None: + """Route supported OCR calls through the Rust ``litellm_python_bridge`` extension. + + ``ocr`` injects the bridge callable; when omitted the compiled extension is + loaded on demand and any previously injected bridge is preserved. Pass + ``ocr=None`` explicitly to clear a prior injection. + """ + global _rust_ocr_enabled, _rust_ocr_impl + _rust_ocr_enabled = enabled + if not isinstance(ocr, _Unset): + _rust_ocr_impl = ocr + + +def rust_ocr_enabled() -> bool: + """Whether the Rust OCR path has been turned on via ``use_litellm_rust()``.""" + return _rust_ocr_enabled + + +def load_rust_ocr() -> RustOcr | None: + """Return the Rust OCR callable, or ``None`` when no bridge is available. + + Prefers an injected implementation, otherwise loads the compiled + ``litellm_python_bridge`` extension; a missing extension yields ``None`` so + the caller can fall back to the Python path instead of hard-failing. + """ + if _rust_ocr_impl is not None: + return _rust_ocr_impl + try: + import litellm_python_bridge + except ImportError: + return None + return cast(RustOcr, litellm_python_bridge.ocr) diff --git a/litellm/types/integrations/custom_logger.py b/litellm/types/integrations/custom_logger.py index b5726a11ca0..26a0be36ef4 100644 --- a/litellm/types/integrations/custom_logger.py +++ b/litellm/types/integrations/custom_logger.py @@ -2,6 +2,25 @@ from typing import Any, Dict, List, Optional from pydantic import BaseModel, Field +CHAT_COMPLETION_AGENTIC_SURFACE = "chat_completions" +CODE_INTERPRETER_INTERCEPTION_PREFIX = "_code_interpreter_interception" +NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES = frozenset( + ("_websearch_interception", "_compression_interception") +) +INTERCEPTION_INTERNAL_PREFIXES = frozenset( + ( + *NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES, + CODE_INTERPRETER_INTERCEPTION_PREFIX, + ) +) + + +def is_interception_internal_key( + key: str, + prefixes: frozenset[str] = INTERCEPTION_INTERNAL_PREFIXES, +) -> bool: + return any(key.startswith(prefix) for prefix in prefixes) + class StandardCustomLoggerInitParams(BaseModel): """ diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 00b095f33ca..24d6e84fba7 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3163,8 +3163,26 @@ class CustomPricingLiteLLMParams(BaseModel): regional_processing_uplift_multiplier_us: Optional[float] = None +# Server-controlled fields that bound or drive an interceptor's agentic loop +# (depth, cycle fingerprints, ceiling, code-interpreter sandbox state). Listed +# in all_litellm_params so they are treated as LiteLLM-level and excluded from +# get_non_default_completion_params; otherwise the OpenAI param builder sweeps +# any unrecognized top-level key into extra_body and leaks them to the provider. +# This is what lets the loop carry state across rerun calls without a provider +# scrubber. +agentic_loop_internal_litellm_params = [ + "_agentic_loop_depth", + "_agentic_loop_fingerprints", + "_agentic_loop_api_surface", + "max_agentic_loops", + "_code_interpreter_interception_active", + "_code_interpreter_interception_sandbox_key", + "_code_interpreter_interception_converted_stream", +] + all_litellm_params = ( - [ + agentic_loop_internal_litellm_params + + [ "metadata", "litellm_metadata", "litellm_trace_id", diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 56baa5c573f..f02efbf6595 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -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", diff --git a/scripts/type_check_gate.py b/scripts/type_check_gate.py index 0f9a44703f9..2ef332d91ea 100644 --- a/scripts/type_check_gate.py +++ b/scripts/type_check_gate.py @@ -1,21 +1,22 @@ #!/usr/bin/env python3 -"""Per-rule count gate for basedpyright. +"""Delta-vs-base per-rule gate for basedpyright. basedpyright's ``--outputjson`` is reduced to a count of errors per *rule* (``reportAny``, ``reportArgumentType``, ...) and checked against a committed budget of the form ``{rule: {baseline, slack}}``, the same shape as -``ruff-strict-budget.json``. A rule fails when its codebase-wide total exceeds -``baseline + slack``. Counts ignore file, line, and column, so a violation -moving anywhere in the tree is invisible; only the per-rule total moves the -needle. +``ruff-strict-budget.json``. A rule fails only when its codebase-wide total is +both over its ceiling (``baseline + slack``) *and* higher than the count on the +base it merges into, so a change is blamed for the errors it adds, never for +drift that already sits in the base. That ``> base`` guard is what stops an +unrelated PR from inheriting a red once two PRs each land near the ceiling and +their sum crosses it: the bystander's count equals its base, so it is spared, +while any PR that actually grows the rule past the cap still fails. -Unlike ``ruff_strict_gate.py`` this does *not* re-run the tool on the merge base -to compute a delta: a second basedpyright pass is minutes and gigabytes, whereas -ruff is milliseconds. The committed budget is the baseline instead -- exactly -how the previous per-file gate worked -- so keep it fresh with ``--update`` -(ratchet), which re-captures every rule's count from the current tree while -preserving each rule's slack. Tool output is read from stdin, so the caller -decides how to invoke basedpyright (and from which cwd). +Head counts are read from stdin (the caller runs basedpyright once and pipes +``--outputjson`` in); the base count is a second basedpyright pass over a +detached worktree at the merge-base, run under the same environment so import +resolution matches. ``--update`` re-captures the absolute per-rule baselines for +the ratchet, preserving each rule's slack. ``--outputjson`` is used rather than text diagnostics because the latter wrap across lines, leaving the ``(reportRule)`` on a continuation line away from the @@ -24,13 +25,21 @@ carries an unambiguous ``rule`` field. """ import argparse +import contextlib import json +import shutil +import subprocess import sys +import tempfile from collections import Counter +from collections.abc import Iterator, Mapping from pathlib import Path -from typing import Mapping, NamedTuple +from typing import NamedTuple REPO_ROOT = Path(__file__).resolve().parent.parent +BUDGET_PATH = REPO_ROOT / "basedpyright-code-budget.json" +PYRIGHT_CONFIG = REPO_ROOT / "pyrightconfig.json" +DEFAULT_BASE = "origin/litellm_internal_staging" # Bucket for a basedpyright diagnostic with no `rule`. Counted so it's gated. UNCODED = "" @@ -45,6 +54,7 @@ class Breach(NamedTuple): code: str total: int cap: int + added: int def _seed_slack(baseline: int) -> int: @@ -54,18 +64,19 @@ def _seed_slack(baseline: int) -> int: return 10 if baseline >= 50 else 3 -def _to_repo_relative(raw: str) -> str | None: +def _to_relative(raw: str, root: Path) -> str | None: path = Path(raw) - absolute = path if path.is_absolute() else Path.cwd() / path + absolute = path if path.is_absolute() else root / path try: - return absolute.resolve().relative_to(REPO_ROOT).as_posix() + return absolute.resolve().relative_to(root).as_posix() except ValueError: return None -def count_basedpyright(payload: str) -> dict[str, int]: - """Count in-repo basedpyright errors per rule from `--outputjson`. Warnings - and information are ignored; only `severity == "error"` is gated.""" +def count_basedpyright(payload: str, root: Path = REPO_ROOT) -> dict[str, int]: + """Count in-tree basedpyright errors per rule from `--outputjson`. Warnings + and information are ignored; only `severity == "error"` is gated. Files + outside `root` (the venv's site-packages, say) are dropped.""" try: data = json.loads(payload or "{}") except json.JSONDecodeError as exc: @@ -79,21 +90,62 @@ def count_basedpyright(payload: str) -> dict[str, int]: for diag in data.get("generalDiagnostics", []): if diag.get("severity") != "error": continue - if _to_repo_relative(diag.get("file", "")) is None: + if _to_relative(diag.get("file", ""), root) is None: continue counts[diag.get("rule") or UNCODED] += 1 return dict(counts) +def _run(cmd: list[str], cwd: Path = REPO_ROOT) -> str: + proc = subprocess.run(cmd, cwd=cwd, capture_output=True, text=True) + if proc.returncode not in (0, 1): + sys.stderr.write(proc.stderr) + raise SystemExit(f"{cmd[0]} exited {proc.returncode}") + return proc.stdout + + +@contextlib.contextmanager +def _temp_worktree(ref: str) -> Iterator[Path]: + parent = Path(tempfile.mkdtemp(prefix="bpr_base_")) + worktree = parent / "wt" + try: + _run(["git", "worktree", "add", "--detach", str(worktree), ref]) + yield worktree + finally: + subprocess.run( + ["git", "worktree", "remove", "--force", str(worktree)], + cwd=REPO_ROOT, + capture_output=True, + text=True, + ) + shutil.rmtree(parent, ignore_errors=True) + + +def base_counts(ref: str) -> dict[str, int]: + """basedpyright error counts per rule for the merge-base tree. The head + config is copied in so the base is judged by today's rules, and the run uses + the head environment's basedpyright (on PATH) so imports resolve the same.""" + exe = shutil.which("basedpyright") or "basedpyright" + with _temp_worktree(ref) as worktree: + shutil.copy(PYRIGHT_CONFIG, worktree / "pyrightconfig.json") + proc = subprocess.run( + [exe, "--outputjson"], cwd=worktree, capture_output=True, text=True + ) + return count_basedpyright(proc.stdout, root=worktree) + + def evaluate( - counts: Mapping[str, int], budget: Mapping[str, Mapping[str, int]] + head: Mapping[str, int], + base: Mapping[str, int], + budget: Mapping[str, Mapping[str, int]], ) -> list[Breach]: breaches = [] - for code, total in counts.items(): + for code, total in head.items(): spec = budget.get(code) cap = spec["baseline"] + spec["slack"] if spec else DEFAULT_SLACK - if total > cap: - breaches.append(Breach(code, total, cap)) + prior = base.get(code, 0) + if total > cap and total > prior: + breaches.append(Breach(code, total, cap, total - prior)) return sorted(breaches) @@ -107,9 +159,6 @@ def is_vacuous_run( return not counts and any(spec["baseline"] for spec in budget.values()) -BUDGET_PATH = REPO_ROOT / "basedpyright-code-budget.json" - - def cmd_update(counts: Mapping[str, int]) -> None: existing = json.loads(BUDGET_PATH.read_text()) if BUDGET_PATH.exists() else {} budget = { @@ -127,9 +176,10 @@ def cmd_update(counts: Mapping[str, int]) -> None: ) -def cmd_check(counts: Mapping[str, int]) -> None: +def cmd_check(base_ref: str) -> None: budget = json.loads(BUDGET_PATH.read_text()) - if is_vacuous_run(counts, budget): + head = count_basedpyright(sys.stdin.read()) + if is_vacuous_run(head, budget): expected = sum(spec["baseline"] for spec in budget.values()) print( f"FAIL: basedpyright produced no errors, but {BUDGET_PATH.name} expects " @@ -137,27 +187,44 @@ def cmd_check(counts: Mapping[str, int]) -> None: f"nothing; refusing to certify a vacuous run." ) raise SystemExit(1) - breaches = evaluate(counts, budget) + base_point = _run(["git", "merge-base", base_ref, "HEAD"]).strip() or base_ref + base = base_counts(base_point) + if is_vacuous_run(base, budget): + print( + f"FAIL: basedpyright produced no errors for the base tree at " + f"{base_point[:12]}, so every rule would look freshly added. The base " + f"pass almost certainly crashed; refusing to blame this change for it." + ) + raise SystemExit(1) + breaches = evaluate(head, base, budget) if not breaches: print( - f"OK: every rule is within its basedpyright ceiling ({sum(counts.values())} errors total)" + f"OK: every rule is within its basedpyright ceiling or no higher than base ({sum(head.values())} errors total)" ) return print("FAIL: basedpyright errors exceed the per-rule ceiling:") for breach in breaches: - print(f" {breach.code}: {breach.total} errors over cap {breach.cap}") + print( + f" {breach.code}: total {breach.total} over cap {breach.cap} (this change added {breach.added})" + ) print( - "Resolve the new errors, or run 'make lint-basedpyright-budget-update' if the ceiling should move." + "Reduce the new errors or remove an equal number elsewhere; the ceiling is " + "baseline + slack in basedpyright-code-budget.json." ) + summary = "; ".join(f"{b.code} {b.total}/{b.cap} (+{b.added})" for b in breaches) + print(f"BREACHED RULES: {summary}") raise SystemExit(1) def main() -> None: parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--base", default=DEFAULT_BASE) parser.add_argument("--update", action="store_true") args = parser.parse_args() - counts = count_basedpyright(sys.stdin.read()) - cmd_update(counts) if args.update else cmd_check(counts) + if args.update: + cmd_update(count_basedpyright(sys.stdin.read())) + else: + cmd_check(args.base) if __name__ == "__main__": diff --git a/tests/llm_translation/test_cloudflare.py b/tests/llm_translation/test_cloudflare.py index 5a6a0008398..54c5d9e4e07 100644 --- a/tests/llm_translation/test_cloudflare.py +++ b/tests/llm_translation/test_cloudflare.py @@ -9,9 +9,7 @@ import pytest from litellm import acompletion, completion from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler -FAKE_API_BASE = ( - "https://fake-cloudflare.example.com/client/v4/accounts/fake-acct/ai/run/" -) +FAKE_API_BASE = "https://fake-cloudflare.example.com/client/v4/accounts/fake-acct/ai/v1" FAKE_API_KEY = "fake-cf-api-key" @@ -26,28 +24,78 @@ def _make_mock_response(json_data: Dict[str, Any]) -> MagicMock: def _chat_response() -> Dict[str, Any]: return { - "result": { - "response": "I am a large language model created to assist you.", - }, - "success": True, - "errors": [], - "messages": [], + "id": "chatcmpl-cf", + "object": "chat.completion", + "created": 1234567890, + "model": "@cf/meta/llama-2-7b-chat-int8", + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": "I am a large language model created to assist you.", + }, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 8, "completion_tokens": 11, "total_tokens": 19}, + } + + +def _tool_call_response() -> Dict[str, Any]: + return { + "id": "chatcmpl-cf-tools", + "object": "chat.completion", + "created": 1234567890, + "model": "@cf/meta/llama-2-7b-chat-int8", + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": { + "name": "get_weather", + "arguments": '{"city": "New York"}', + }, + } + ], + }, + "finish_reason": "tool_calls", + } + ], + "usage": {"prompt_tokens": 20, "completion_tokens": 9, "total_tokens": 29}, } def _streaming_chunks() -> list[str]: + base = { + "id": "chatcmpl-cf", + "object": "chat.completion.chunk", + "created": 1234567890, + "model": "@cf/meta/llama-2-7b-chat-int8", + } return [ - json.dumps({"response": "I am"}), - json.dumps({"response": " a language"}), - json.dumps({"response": " model."}), - ] - - -def _streaming_chunks_response_text() -> list[str]: - return [ - json.dumps({"response_text": "I am"}), - json.dumps({"response_text": " a language"}), - json.dumps({"response_text": " model."}), + json.dumps({**base, "choices": [{"index": 0, "delta": {"content": "I am"}}]}), + json.dumps( + {**base, "choices": [{"index": 0, "delta": {"content": " a language"}}]} + ), + json.dumps( + { + **base, + "choices": [ + { + "index": 0, + "delta": {"content": " model."}, + "finish_reason": "stop", + } + ], + } + ), ] @@ -85,6 +133,48 @@ def test_completion_cloudflare(sync_mode): assert response.choices[0].message.content is not None assert "language model" in response.choices[0].message.content.lower() + called_url = mock_post.call_args.kwargs.get("url") or mock_post.call_args.args[0] + assert called_url.endswith("/ai/v1/chat/completions") + assert "/ai/run/" not in called_url + + +def test_completion_cloudflare_tool_calls_sent_to_openai_endpoint(): + messages = [{"role": "user", "content": "weather in New York?"}] + tools = [ + { + "type": "function", + "function": { + "name": "get_weather", + "parameters": { + "type": "object", + "properties": {"city": {"type": "string"}}, + "required": ["city"], + }, + }, + } + ] + mock_resp = _make_mock_response(_tool_call_response()) + + with patch.object(HTTPHandler, "post", return_value=mock_resp) as mock_post: + response = completion( + model="cloudflare/@cf/meta/llama-2-7b-chat-int8", + messages=messages, + tools=tools, + tool_choice="auto", + api_base=FAKE_API_BASE, + api_key=FAKE_API_KEY, + ) + mock_post.assert_called_once() + + sent_body = json.loads(mock_post.call_args.kwargs["data"]) + assert sent_body["tools"] == tools + assert sent_body["tool_choice"] == "auto" + + assert response.choices[0].finish_reason == "tool_calls" + tool_calls = response.choices[0].message.tool_calls + assert tool_calls is not None and len(tool_calls) == 1 + assert tool_calls[0].function.name == "get_weather" + @pytest.mark.parametrize("sync_mode", [True, False]) def test_completion_cloudflare_stream(sync_mode): @@ -153,76 +243,3 @@ def test_completion_cloudflare_stream(sync_mode): if c.choices[0].delta.content ) assert "language" in content.lower() - - -@pytest.mark.parametrize("sync_mode", [True, False]) -def test_completion_cloudflare_stream_response_text(sync_mode): - """Newer Cloudflare Workers AI models (e.g. Nemotron) emit `response_text` - instead of `response` in streamed chunks. The iterator must surface that - text so streaming output is not silently empty. - """ - messages = [{"role": "user", "content": "what llm are you"}] - raw_chunks = _streaming_chunks_response_text() - - if sync_mode: - - def _iter_lines(): - for chunk in raw_chunks: - yield f"data: {chunk}" - yield "data: [DONE]" - - mock_resp = MagicMock() - mock_resp.iter_lines.return_value = _iter_lines() - mock_resp.status_code = 200 - mock_resp.headers = {"content-type": "text/event-stream"} - - with patch.object(HTTPHandler, "post", return_value=mock_resp) as mock_post: - response = completion( - model="cloudflare/@cf/nvidia/nemotron-mini-4b-instruct", - messages=messages, - max_tokens=15, - stream=True, - api_base=FAKE_API_BASE, - api_key=FAKE_API_KEY, - ) - chunks_received = list(response) - mock_post.assert_called_once() - else: - - async def _aiter_lines(): - for chunk in raw_chunks: - yield f"data: {chunk}" - yield "data: [DONE]" - - mock_resp = MagicMock() - mock_resp.aiter_lines.return_value = _aiter_lines() - mock_resp.status_code = 200 - mock_resp.headers = {"content-type": "text/event-stream"} - - async def _run(): - with patch.object( - AsyncHTTPHandler, "post", new_callable=AsyncMock, return_value=mock_resp - ) as mock_post: - resp = await acompletion( - model="cloudflare/@cf/nvidia/nemotron-mini-4b-instruct", - messages=messages, - max_tokens=15, - stream=True, - api_base=FAKE_API_BASE, - api_key=FAKE_API_KEY, - ) - received = [] - async for chunk in resp: - received.append(chunk) - mock_post.assert_called_once() - return received - - chunks_received = asyncio.run(_run()) - - assert len(chunks_received) > 0 - content = "".join( - c.choices[0].delta.content - for c in chunks_received - if c.choices[0].delta.content - ) - assert "language" in content.lower() diff --git a/tests/search_tests/test_searchapi_search.py b/tests/search_tests/test_searchapi_search.py index 5ef9d922b89..d16868502a4 100644 --- a/tests/search_tests/test_searchapi_search.py +++ b/tests/search_tests/test_searchapi_search.py @@ -46,10 +46,9 @@ class TestSearchAPIConfig: assert result["Content-Type"] == "application/json" - @patch("litellm.llms.searchapi.search.transformation.get_secret_str") - def test_validate_environment_without_api_key(self, mock_get_secret): + def test_validate_environment_without_api_key(self, monkeypatch): """Test environment validation without API key raises error.""" - mock_get_secret.return_value = None + monkeypatch.delenv("SEARCHAPI_API_KEY", raising=False) config = SearchAPIConfig() headers = {} diff --git a/tests/search_tests/test_searxng_search.py b/tests/search_tests/test_searxng_search.py index 45b0f3214d9..c12d44183b0 100644 --- a/tests/search_tests/test_searxng_search.py +++ b/tests/search_tests/test_searxng_search.py @@ -318,13 +318,11 @@ class TestSearXNGSearchHeaders: assert headers["Content-Type"] == "application/json" assert headers["Authorization"] == "Bearer test-key-123" - def test_headers_with_env_api_key(self): + def test_headers_with_env_api_key(self, monkeypatch): """Test that headers use SEARXNG_API_KEY from env.""" - with patch( - "litellm.llms.searxng.search.transformation.get_secret_str", - return_value="env-key-456", - ): - headers = self.config.validate_environment(headers={}) + monkeypatch.setenv("SEARXNG_API_KEY", "env-key-456") + + headers = self.config.validate_environment(headers={}) assert headers["Authorization"] == "Bearer env-key-456" diff --git a/tests/test_litellm/integrations/code_interpreter_interception/test_handler.py b/tests/test_litellm/integrations/code_interpreter_interception/test_handler.py index f33814b86df..7ff58ba6324 100644 --- a/tests/test_litellm/integrations/code_interpreter_interception/test_handler.py +++ b/tests/test_litellm/integrations/code_interpreter_interception/test_handler.py @@ -1,8 +1,8 @@ """ Unit tests for CodeInterpreterInterceptionLogger. -All sandbox dependencies are injected (dependency injection, no monkeypatch): -a FakeSandbox stands in for the real e2b config and records how it is called. +All sandbox dependencies are injected: a FakeSandbox stands in for the real e2b +config and records how it is called. """ import time @@ -12,13 +12,17 @@ import pytest from litellm.integrations.code_interpreter_interception.handler import ( CodeInterpreterInterceptionLogger, LITELLM_CODE_EXECUTION_TOOL_NAME, + _INTERCEPTION_ACTIVE_KEY as _ACTIVE_KEY, + _SANDBOX_KEY, +) +from litellm.types.integrations.custom_logger import ( + CHAT_COMPLETION_AGENTIC_SURFACE, + NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES, + is_interception_internal_key, ) from litellm.llms.base_llm.sandbox.transformation import CodeExecutionResult from litellm.types.utils import CallTypes -_ACTIVE_KEY = "_code_interpreter_interception_active" -_SANDBOX_KEY = "_code_interpreter_interception_sandbox_key" - class FakeHandle: def __init__(self, sandbox_id="sbx_fake"): @@ -51,6 +55,13 @@ class FakeLogging: def __init__(self, litellm_call_id="k1"): self.litellm_call_id = litellm_call_id self.model_call_details = {} + self.dynamic_success_callbacks = [] + + def pre_call(self, *args, **kwargs): + return None + + def post_call(self, *args, **kwargs): + return None def _function_call_item(call_id="c1", name=LITELLM_CODE_EXECUTION_TOOL_NAME): @@ -62,6 +73,17 @@ def _function_call_item(call_id="c1", name=LITELLM_CODE_EXECUTION_TOOL_NAME): } +def _chat_function_call_item(call_id="call_1", name=LITELLM_CODE_EXECUTION_TOOL_NAME): + return { + "id": call_id, + "type": "function", + "function": { + "name": name, + "arguments": '{"code":"print(40 + 2)"}', + }, + } + + class FakeResponse: def __init__(self, output): self.output = output @@ -74,6 +96,18 @@ def _iter_messages(plan): return patch.messages +def test_interception_internal_key_prefix_sets_preserve_code_interpreter_state(): + assert is_interception_internal_key("_code_interpreter_interception_active") + assert not is_interception_internal_key( + "_code_interpreter_interception_active", + prefixes=NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES, + ) + assert is_interception_internal_key( + "_websearch_interception_converted_stream", + prefixes=NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES, + ) + + @pytest.mark.asyncio async def test_build_plan_runs_code_and_feeds_output_back(): sandbox = FakeSandbox(stdout="42") @@ -133,6 +167,30 @@ async def test_pre_call_converts_code_interpreter_tool(): assert LITELLM_CODE_EXECUTION_TOOL_NAME in names +@pytest.mark.asyncio +async def test_pre_call_converts_code_interpreter_tool_for_chat_completions(): + logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox()) + kwargs = { + "tools": [{"type": "code_interpreter", "container": {"type": "auto"}}], + "tool_choice": {"type": "code_interpreter"}, + "custom_llm_provider": "openai", + } + + result = await logger.async_pre_call_deployment_hook(kwargs, CallTypes.acompletion) + + assert result is not None + tool = result["tools"][0] + assert tool["type"] == "function" + assert tool["function"]["name"] == LITELLM_CODE_EXECUTION_TOOL_NAME + assert tool["function"]["parameters"]["required"] == ["code"] + assert result["tool_choice"] == { + "type": "function", + "function": {"name": LITELLM_CODE_EXECUTION_TOOL_NAME}, + } + assert result["litellm_metadata"][_ACTIVE_KEY] is True + assert result["litellm_metadata"][_SANDBOX_KEY] == result[_SANDBOX_KEY] + + @pytest.mark.asyncio @pytest.mark.parametrize( "tool_choice", @@ -184,9 +242,24 @@ async def test_pre_call_noop_on_non_responses(): "custom_llm_provider": "openai", } + result = await logger.async_pre_call_deployment_hook(kwargs, CallTypes.aembedding) + + assert result is None + + +@pytest.mark.asyncio +async def test_pre_call_noop_on_chat_completion_without_code_interpreter_tool(): + logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox()) + kwargs = { + "tools": [{"type": "web_search"}], + "custom_llm_provider": "openai", + } + result = await logger.async_pre_call_deployment_hook(kwargs, CallTypes.acompletion) assert result is None + assert _ACTIVE_KEY not in kwargs + assert _SANDBOX_KEY not in kwargs @pytest.mark.asyncio @@ -524,14 +597,142 @@ async def test_gate_rechecks_provider_scope(): assert should_run is False +@pytest.mark.asyncio +async def test_chat_completion_gate_detects_code_execution_tool_call(): + logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox()) + response = { + "choices": [ + {"message": {"tool_calls": [_chat_function_call_item(call_id="call_123")]}} + ] + } + + should_run, payload = await logger.async_should_run_agentic_loop( + response=response, + model="gpt-5", + messages=[{"role": "user", "content": "x"}], + tools=[], + stream=False, + custom_llm_provider="openai", + kwargs={ + _ACTIVE_KEY: True, + "_agentic_loop_api_surface": CHAT_COMPLETION_AGENTIC_SURFACE, + }, + ) + + assert should_run is True + assert payload["tool_calls"][0]["id"] == "call_123" + assert payload["tool_calls"][0]["arguments"] == '{"code":"print(40 + 2)"}' + + +@pytest.mark.asyncio +async def test_chat_completion_gate_refuses_without_server_active_marker(): + logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox()) + response = {"choices": [{"message": {"tool_calls": [_chat_function_call_item()]}}]} + + should_run, payload = await logger.async_should_run_agentic_loop( + response=response, + model="gpt-5", + messages=[{"role": "user", "content": "x"}], + tools=[], + stream=False, + custom_llm_provider="openai", + kwargs={"_agentic_loop_api_surface": CHAT_COMPLETION_AGENTIC_SURFACE}, + ) + + assert should_run is False + assert payload == {} + + +@pytest.mark.asyncio +async def test_chat_completion_build_plan_runs_code_and_appends_tool_message(): + sandbox = FakeSandbox(stdout="42") + logger = CodeInterpreterInterceptionLogger(sandbox_config=sandbox) + native_chat_tool = {"type": "code_interpreter", "container": {"type": "auto"}} + + plan = await logger.async_build_agentic_loop_plan( + tools={ + "tool_calls": [ + { + "id": "call_1", + "name": LITELLM_CODE_EXECUTION_TOOL_NAME, + "arguments": '{"code":"print(40 + 2)"}', + } + ] + }, + model="gpt-5", + messages=[{"role": "user", "content": "x"}], + response={ + "choices": [{"message": {"tool_calls": [_chat_function_call_item()]}}] + }, + anthropic_messages_provider_config=None, + anthropic_messages_optional_request_params={ + "tools": [native_chat_tool], + "tool_choice": {"type": "code_interpreter", "container": {"type": "auto"}}, + "temperature": 0, + }, + logging_obj=FakeLogging(litellm_call_id="k1"), + stream=False, + kwargs={ + "acompletion": True, + "litellm_call_id": "k1", + _ACTIVE_KEY: True, + _SANDBOX_KEY: "sbxkey1", + "_code_interpreter_interception_converted_stream": True, + "_agentic_loop_api_surface": CHAT_COMPLETION_AGENTIC_SURFACE, + }, + ) + + assert sandbox.run_calls[0]["code"] == "print(40 + 2)" + patch = plan.request_patch + assert patch is not None + assert patch.tools == [ + { + "type": "function", + "function": { + "name": LITELLM_CODE_EXECUTION_TOOL_NAME, + "description": "Execute python code in a sandbox and return stdout.", + "parameters": { + "type": "object", + "properties": {"code": {"type": "string"}}, + "required": ["code"], + }, + }, + } + ] + assert patch.optional_params == {"temperature": 0} + assert patch.kwargs == { + "litellm_call_id": "k1", + _ACTIVE_KEY: True, + _SANDBOX_KEY: "sbxkey1", + "_code_interpreter_interception_converted_stream": True, + "_agentic_loop_api_surface": CHAT_COMPLETION_AGENTIC_SURFACE, + } + assert patch.messages is not None + assert patch.messages[-2]["role"] == "assistant" + assert patch.messages[-2]["tool_calls"][0]["id"] == "call_1" + assert patch.messages[-1] == { + "role": "tool", + "tool_call_id": "call_1", + "content": "42", + } + assert plan.metadata["code_interpreter_calls"][0]["code"] == "print(40 + 2)" + + @pytest.mark.asyncio async def test_pre_call_strips_client_forged_marker_on_initial_request(): - """A client cannot pre-set the active marker on the original request.""" + """A client cannot pre-set the active marker on the original request: with no + native code_interpreter tool, any client-supplied interception markers in + litellm_metadata are scrubbed and the active flag in kwargs is cleared.""" logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox()) kwargs = { "tools": [{"type": "web_search"}], "custom_llm_provider": "openai", _ACTIVE_KEY: True, + "litellm_metadata": { + _ACTIVE_KEY: True, + _SANDBOX_KEY: "client-forged", + "safe_user_value": "kept", + }, } await logger.async_pre_call_deployment_hook(kwargs, CallTypes.aresponses) @@ -540,6 +741,42 @@ async def test_pre_call_strips_client_forged_marker_on_initial_request(): "no native code_interpreter tool was present, so a client-supplied " "active marker must be cleared" ) + assert kwargs["litellm_metadata"] == {"safe_user_value": "kept"} + + +@pytest.mark.asyncio +async def test_pre_call_strips_forged_loop_controls_then_mints_own_markers(): + """On an INITIAL request (no server-set _agentic_loop_depth) a client cannot + smuggle loop-control state: forged _agentic_loop_depth / max_agentic_loops and + interception markers in litellm_metadata are stripped before the interceptor + activates, so the only interception markers that survive are the ones the + server mints for the converted code_interpreter tool.""" + logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox()) + kwargs = { + "tools": [{"type": "code_interpreter", "container": {"type": "auto"}}], + "custom_llm_provider": "openai", + "litellm_metadata": { + _ACTIVE_KEY: True, + _SANDBOX_KEY: "client-forged", + "_agentic_loop_depth": 99, + "max_agentic_loops": 999, + "safe_user_value": "kept", + }, + } + + result = await logger.async_pre_call_deployment_hook(kwargs, CallTypes.acompletion) + + assert result is not None + metadata = result["litellm_metadata"] + assert metadata["safe_user_value"] == "kept" + assert "_agentic_loop_depth" not in metadata, "forged loop depth must be stripped" + assert "max_agentic_loops" not in metadata, "forged loop cap must be stripped" + assert metadata[_ACTIVE_KEY] is True + assert metadata[_SANDBOX_KEY] == result[_SANDBOX_KEY] + assert metadata[_SANDBOX_KEY] != "client-forged", ( + "the surviving sandbox key must be the server-minted one, not the forged " + "value the client supplied" + ) @pytest.mark.asyncio diff --git a/tests/test_litellm/litellm_core_utils/test_chat_completion_agentic_loop.py b/tests/test_litellm/litellm_core_utils/test_chat_completion_agentic_loop.py new file mode 100644 index 00000000000..f1196ab4692 --- /dev/null +++ b/tests/test_litellm/litellm_core_utils/test_chat_completion_agentic_loop.py @@ -0,0 +1,422 @@ +""" +Tests for the provider-agnostic chat completion agentic loop dispatcher +(`litellm/litellm_core_utils/chat_completion_agentic_loop.py`) and the +code-interpreter interception integration that drives it. + +The load-bearing regression here protects a reviewer requirement: the internal +agentic/interception control fields must NEVER reach the outbound provider HTTP +request body. The relevant fields are: + + _agentic_loop_depth + _agentic_loop_fingerprints + _agentic_loop_api_surface + max_agentic_loops + _code_interpreter_interception_active + _code_interpreter_interception_sandbox_key + _code_interpreter_interception_converted_stream + +A scrubber in gpt_transformation.py used to strip these. That scrubber was +removed, so `test_internal_control_fields_never_leak_into_provider_body` proves +they stay out of the body even without it. +""" + +import os +import sys +from typing import Any, Dict, List, Optional, Tuple +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../../..")) + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.integrations.code_interpreter_interception.handler import ( + CodeInterpreterInterceptionLogger, +) +from litellm.litellm_core_utils.chat_completion_agentic_loop import ( + maybe_run_chat_completion_agentic_loop, +) +from litellm.types.integrations.custom_logger import ( + AgenticLoopPlan, + AgenticLoopRequestPatch, +) +from litellm.types.utils import ( + Choices, + Function, + ChatCompletionMessageToolCall, + Message, + ModelResponse, +) + +# The internal control fields that must never reach a provider request body. +_INTERNAL_CONTROL_FIELDS = ( + "_agentic_loop_depth", + "_agentic_loop_fingerprints", + "_agentic_loop_api_surface", + "max_agentic_loops", + "_code_interpreter_interception_active", + "_code_interpreter_interception_sandbox_key", + "_code_interpreter_interception_converted_stream", + "litellm_metadata", +) + + +@pytest.fixture +def restore_callbacks(): + """Save/restore litellm.callbacks so a registered fake logger never pollutes + other tests in the suite.""" + saved = list(litellm.callbacks) + try: + yield + finally: + litellm.callbacks = saved + + +class _SandboxResult: + def __init__(self, stdout: str) -> None: + self.stdout = stdout + self.error = None + + +class FakeSandboxConfig: + """Injected sandbox so the interception loop runs no real network / E2B.""" + + def __init__(self) -> None: + self.created = 0 + self.deleted = 0 + self.run_codes: List[str] = [] + + async def acreate_sandbox(self) -> Any: + self.created += 1 + return MagicMock(id="sandbox-123") + + async def arun_code(self, container: Any, code: str) -> _SandboxResult: + self.run_codes.append(code) + return _SandboxResult(stdout="42\n") + + async def adelete_sandbox(self, container: Any) -> None: + self.deleted += 1 + + +def _tool_call_model_response() -> ModelResponse: + return ModelResponse( + choices=[ + Choices( + finish_reason="tool_calls", + message=Message( + role="assistant", + content=None, + tool_calls=[ + ChatCompletionMessageToolCall( + id="call_abc", + type="function", + function=Function( + name="litellm_code_execution", + arguments='{"code": "print(6*7)"}', + ), + ) + ], + ), + ) + ] + ) + + +def _plain_model_response(content: str = "The answer is 42") -> ModelResponse: + return ModelResponse( + choices=[ + Choices( + finish_reason="stop", + message=Message(role="assistant", content=content), + ) + ] + ) + + +def _raw_response_for(model_response: ModelResponse) -> MagicMock: + """Wrap a ModelResponse as the OpenAI `with_raw_response.create` return value + (an object exposing `.headers` and `.parse()` -> something with model_dump).""" + parsed = MagicMock() + parsed.model_dump.return_value = model_response.model_dump() + raw = MagicMock() + raw.headers = {} + raw.parse.return_value = parsed + return raw + + +# --------------------------------------------------------------------------- +# A) PROVIDER-PAYLOAD REGRESSION +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_internal_control_fields_never_leak_into_provider_body(restore_callbacks): + """Drive a real acompletion with a native code_interpreter tool through the + interception logger + agentic loop, capturing every outbound OpenAI request + body. None of the internal control fields may appear at top-level or inside + extra_body on ANY of the captured calls.""" + logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandboxConfig()) + litellm.callbacks = [logger] + + # First create -> model emits a code_execution tool call (triggers the loop). + # Second create -> model returns a plain answer (loop terminates). + create = AsyncMock( + side_effect=[ + _raw_response_for(_tool_call_model_response()), + _raw_response_for(_plain_model_response()), + ] + ) + mock_client = MagicMock() + mock_client.chat.completions.with_raw_response.create = create + + response = await litellm.acompletion( + model="openai/gpt-4o-mini", + messages=[{"role": "user", "content": "what is 6*7?"}], + tools=[{"type": "code_interpreter"}], + tool_choice={"type": "code_interpreter"}, + api_key="sk-test", + client=mock_client, + ) + + # The loop must have actually fired (sanity: two provider calls). + assert create.await_count == 2, ( + "expected the agentic loop to issue a follow-up provider call; " + f"got {create.await_count} call(s)" + ) + + for idx, call in enumerate(create.await_args_list): + body = call.kwargs + extra_body = body.get("extra_body") or {} + for field in _INTERNAL_CONTROL_FIELDS: + assert field not in body, ( + f"provider call #{idx}: internal field {field!r} leaked into " + f"top-level request body: {sorted(body.keys())}" + ) + assert field not in extra_body, ( + f"provider call #{idx}: internal field {field!r} leaked into " + f"extra_body: {sorted(extra_body.keys())}" + ) + # The native code_interpreter tool must have been swapped for the + # function tool, never sent raw to OpenAI as a chat-completions request. + for tool in body.get("tools") or []: + assert tool.get("type") != "code_interpreter" + + # The final response is the post-loop answer, not the tool-call turn. + assert response.choices[0].message.content == "The answer is 42" + + +# --------------------------------------------------------------------------- +# B) DISPATCHER UNIT TESTS +# --------------------------------------------------------------------------- + + +class _LoggingStub: + """Minimal logging_obj: dispatcher only reads dynamic_success_callbacks and + litellm_call_id off it.""" + + litellm_call_id = "call-test" + dynamic_success_callbacks: List[Any] = [] + + +class _GateOnlyLogger(CustomLogger): + """Overrides the gate to fire, but builds a plan from request_patch.""" + + def __init__(self, plan: AgenticLoopPlan, tool_calls: Dict[str, Any]) -> None: + super().__init__() + self._plan = plan + self._tool_calls = tool_calls + self.cleanup_calls = 0 + + async def async_should_run_agentic_loop( + self, + response: Any, + model: str, + messages: List[Dict[str, Any]], + tools: Optional[List[Dict[str, Any]]], + stream: bool, + custom_llm_provider: str, + kwargs: Dict[str, Any], + ) -> Tuple[bool, Dict[str, Any]]: + return True, self._tool_calls + + async def async_build_agentic_loop_plan( + self, + tools: Dict[str, Any], + model: str, + messages: List[Dict[str, Any]], + response: Any, + anthropic_messages_provider_config: Any, + anthropic_messages_optional_request_params: Dict[str, Any], + logging_obj: Any, + stream: bool, + kwargs: Dict[str, Any], + ) -> AgenticLoopPlan: + return self._plan + + async def async_agentic_loop_cleanup_hook( + self, plan: AgenticLoopPlan, kwargs: Dict[str, Any] + ) -> None: + self.cleanup_calls += 1 + + +def _patched_messages() -> List[Dict[str, Any]]: + return [ + {"role": "user", "content": "what is 6*7?"}, + { + "role": "assistant", + "tool_calls": [ + { + "id": "call_abc", + "type": "function", + "function": { + "name": "litellm_code_execution", + "arguments": '{"code": "print(6*7)"}', + }, + } + ], + }, + {"role": "tool", "tool_call_id": "call_abc", "content": "42\n"}, + ] + + +@pytest.mark.asyncio +async def test_dispatcher_returns_none_when_no_callback_gates(restore_callbacks): + """No callback overrides the gate -> dispatcher returns None so the caller + keeps the original response untouched.""" + litellm.callbacks = [] + + result = await maybe_run_chat_completion_agentic_loop( + response=_plain_model_response(), + model="gpt-4o-mini", + messages=[{"role": "user", "content": "hi"}], + optional_params={}, + kwargs={}, + logging_obj=_LoggingStub(), + custom_llm_provider="openai", + stream=False, + ) + + assert result is None + + +@pytest.mark.asyncio +async def test_dispatcher_runs_followup_with_incremented_depth_and_patched_messages( + restore_callbacks, +): + """A gating logger with a request_patch -> the dispatcher calls + litellm.acompletion exactly once with _agentic_loop_depth == 1 and the + patched messages. Loop-control state rides as litellm-level kwargs and is + mirrored into litellm_metadata; the provider-surface transient + _agentic_loop_api_surface is never forwarded. (Provider-body stripping of + these litellm-level kwargs is asserted separately in test A.)""" + followup = _plain_model_response("done") + plan = AgenticLoopPlan( + run_agentic_loop=True, + request_patch=AgenticLoopRequestPatch(messages=_patched_messages()), + ) + logger = _GateOnlyLogger(plan=plan, tool_calls={"tool_calls": [{"id": "call_abc"}]}) + litellm.callbacks = [logger] + + acompletion_mock = AsyncMock(return_value=followup) + with patch.object(litellm, "acompletion", acompletion_mock): + result = await maybe_run_chat_completion_agentic_loop( + response=_tool_call_model_response(), + model="gpt-4o-mini", + messages=[{"role": "user", "content": "what is 6*7?"}], + optional_params={"temperature": 0.1}, + kwargs={"_code_interpreter_interception_active": True}, + logging_obj=_LoggingStub(), + custom_llm_provider="openai", + stream=False, + ) + + assert result is followup + acompletion_mock.assert_awaited_once() + call_kwargs = acompletion_mock.await_args.kwargs + + assert call_kwargs["_agentic_loop_depth"] == 1 + assert call_kwargs["messages"] == _patched_messages() + # Preserved non-internal optional param survives the rerun. + assert call_kwargs["temperature"] == 0.1 + # Loop-control state is carried at the litellm level for the follow-up. + assert call_kwargs["max_agentic_loops"] >= 1 + assert "_agentic_loop_fingerprints" in call_kwargs + # Interception markers are mirrored into litellm_metadata for the follow-up. + assert ( + call_kwargs["litellm_metadata"]["_code_interpreter_interception_active"] is True + ) + # The transient surface marker is NOT forwarded to the follow-up call. + assert "_agentic_loop_api_surface" not in call_kwargs + # Cleanup hook always runs. + assert logger.cleanup_calls == 1 + + +@pytest.mark.asyncio +async def test_dispatcher_raises_when_depth_reaches_max_agentic_loops( + restore_callbacks, +): + """depth >= max_agentic_loops -> ValueError mentioning max_agentic_loops, + before any follow-up call is attempted.""" + logger = _GateOnlyLogger( + plan=AgenticLoopPlan(run_agentic_loop=True), + tool_calls={"tool_calls": [{"id": "call_abc"}]}, + ) + litellm.callbacks = [logger] + + acompletion_mock = AsyncMock() + with patch.object(litellm, "acompletion", acompletion_mock): + with pytest.raises(ValueError, match="max_agentic_loops"): + await maybe_run_chat_completion_agentic_loop( + response=_tool_call_model_response(), + model="gpt-4o-mini", + messages=[{"role": "user", "content": "hi"}], + optional_params={}, + kwargs={"_agentic_loop_depth": 3, "max_agentic_loops": 3}, + logging_obj=_LoggingStub(), + custom_llm_provider="openai", + stream=False, + ) + + acompletion_mock.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_dispatcher_raises_on_repeated_tool_call_fingerprint(restore_callbacks): + """A tool_calls fingerprint already present in _agentic_loop_fingerprints -> + ValueError about the repeated fingerprint (cycle guard), with no follow-up + call.""" + import json + + # The dispatcher fingerprints the whole value the gate returns as its second + # tuple element, so the seeded fingerprint must mirror that dict exactly. + gate_tool_calls = { + "tool_calls": [{"id": "call_abc", "name": "litellm_code_execution"}] + } + fingerprint = json.dumps(gate_tool_calls, sort_keys=True, default=str) + + logger = _GateOnlyLogger( + plan=AgenticLoopPlan(run_agentic_loop=True), + tool_calls=gate_tool_calls, + ) + litellm.callbacks = [logger] + + acompletion_mock = AsyncMock() + with patch.object(litellm, "acompletion", acompletion_mock): + with pytest.raises(ValueError, match="fingerprint"): + await maybe_run_chat_completion_agentic_loop( + response=_tool_call_model_response(), + model="gpt-4o-mini", + messages=[{"role": "user", "content": "hi"}], + optional_params={}, + kwargs={ + "_agentic_loop_depth": 0, + "max_agentic_loops": 3, + "_agentic_loop_fingerprints": [fingerprint], + }, + logging_obj=_LoggingStub(), + custom_llm_provider="openai", + stream=False, + ) + + acompletion_mock.assert_not_awaited() diff --git a/tests/test_litellm/llms/apiserpent/test_apiserpent_search.py b/tests/test_litellm/llms/apiserpent/test_apiserpent_search.py index 32838701949..bc26268ee92 100644 --- a/tests/test_litellm/llms/apiserpent/test_apiserpent_search.py +++ b/tests/test_litellm/llms/apiserpent/test_apiserpent_search.py @@ -66,9 +66,8 @@ class TestAPISerpentConfig: assert headers["X-API-Key"] == "test-api-key" assert headers["Content-Type"] == "application/json" - @patch("litellm.llms.apiserpent.search.transformation.get_secret_str") - def test_validate_environment_without_api_key(self, mock_get_secret): - mock_get_secret.return_value = None + def test_validate_environment_without_api_key(self, monkeypatch): + monkeypatch.delenv("APISERPENT_API_KEY", raising=False) with pytest.raises(ValueError, match="APISERPENT_API_KEY is not set"): APISerpentSearchConfig().validate_environment({}) diff --git a/tests/test_litellm/llms/base_llm/search/test_base_search_transformation.py b/tests/test_litellm/llms/base_llm/search/test_base_search_transformation.py new file mode 100644 index 00000000000..a1353d57038 --- /dev/null +++ b/tests/test_litellm/llms/base_llm/search/test_base_search_transformation.py @@ -0,0 +1,329 @@ +""" +Regression tests for the host-aware server-credential fallback guard in +``BaseSearchConfig``. + +A caller-supplied ``api_base`` is honored when building the request URL, so +falling back to a server-configured secret while the caller controls the host +would send the operator's credential to an attacker. The guard must refuse that +combination for every provider that carries a server-managed secret, while +leaving keyless providers and legitimate operator overrides untouched. +""" + +from typing import Dict, Tuple, Type +from unittest.mock import AsyncMock, patch + +import pytest + +import litellm +from litellm.llms.apiserpent.search.transformation import APISerpentSearchConfig +from litellm.llms.base_llm.search.transformation import ( + BaseSearchConfig, + _is_trusted_search_api_base, +) +from litellm.llms.brave.search.transformation import BraveSearchConfig +from litellm.llms.dataforseo.search.transformation import DataForSEOSearchConfig +from litellm.llms.exa_ai.search.transformation import ExaAISearchConfig +from litellm.llms.fastcrw.search.transformation import FastCRWSearchConfig +from litellm.llms.firecrawl.search.transformation import FirecrawlSearchConfig +from litellm.llms.google_pse.search.transformation import GooglePSESearchConfig +from litellm.llms.linkup.search.transformation import LinkupSearchConfig +from litellm.llms.parallel_ai.search.transformation import ParallelAISearchConfig +from litellm.llms.perplexity.search.transformation import PerplexitySearchConfig +from litellm.llms.searchapi.search.transformation import SearchAPIConfig +from litellm.llms.searxng.search.transformation import SearXNGSearchConfig +from litellm.llms.serper.search.transformation import SerperSearchConfig +from litellm.llms.tavily.search.transformation import TavilySearchConfig +from litellm.llms.tinyfish.search.transformation import TinyfishSearchConfig +from litellm.llms.you_com.search.transformation import YouComSearchConfig + +ATTACKER_BASE = "https://attacker.example.com" + +# Every *_API_BASE override env var that could otherwise mark the attacker host +# as trusted; cleared before each test so the suite is hermetic. +_BASE_ENV_VARS = ( + "SERPER_API_BASE", + "TAVILY_API_BASE", + "PERPLEXITY_API_BASE", + "APISERPENT_API_BASE", + "EXA_API_BASE", + "BRAVE_API_BASE", + "FIRECRAWL_API_BASE", + "LINKUP_API_BASE", + "SEARCHAPI_API_BASE", + "GOOGLE_PSE_API_BASE", + "PARALLEL_AI_API_BASE", + "YOUCOM_API_BASE", + "SEARXNG_API_BASE", + "DATAFORSEO_API_BASE", + "TINYFISH_API_BASE", + "CRW_API_BASE", +) + + +@pytest.fixture(autouse=True) +def _clear_base_overrides(monkeypatch: pytest.MonkeyPatch) -> None: + for var in _BASE_ENV_VARS: + monkeypatch.delenv(var, raising=False) + + +# (config, {server secret env vars}, caller_api_key honored as-is, extra env for full validate) +ProviderSpec = Tuple[Type[BaseSearchConfig], Dict[str, str], str, Dict[str, str]] + +PROVIDERS: Tuple[ProviderSpec, ...] = ( + (SerperSearchConfig, {"SERPER_API_KEY": "srv"}, "caller-key", {}), + (TavilySearchConfig, {"TAVILY_API_KEY": "srv"}, "caller-key", {}), + (PerplexitySearchConfig, {"PERPLEXITYAI_API_KEY": "srv"}, "caller-key", {}), + (APISerpentSearchConfig, {"APISERPENT_API_KEY": "srv"}, "caller-key", {}), + (ExaAISearchConfig, {"EXA_API_KEY": "srv"}, "caller-key", {}), + (BraveSearchConfig, {"BRAVE_API_KEY": "srv"}, "caller-key", {}), + (FirecrawlSearchConfig, {"FIRECRAWL_API_KEY": "srv"}, "caller-key", {}), + (LinkupSearchConfig, {"LINKUP_API_KEY": "srv"}, "caller-key", {}), + (SearchAPIConfig, {"SEARCHAPI_API_KEY": "srv"}, "caller-key", {}), + ( + GooglePSESearchConfig, + {"GOOGLE_PSE_API_KEY": "srv"}, + "caller-key", + {"GOOGLE_PSE_ENGINE_ID": "engine"}, + ), + (ParallelAISearchConfig, {"PARALLEL_API_KEY": "srv"}, "caller-key", {}), + (YouComSearchConfig, {"YOUCOM_API_KEY": "srv"}, "caller-key", {}), + (SearXNGSearchConfig, {"SEARXNG_API_KEY": "srv"}, "caller-key", {}), + ( + DataForSEOSearchConfig, + {"DATAFORSEO_LOGIN": "srv", "DATAFORSEO_PASSWORD": "pw"}, + "login:password", + {}, + ), + (TinyfishSearchConfig, {"TINYFISH_API_KEY": "srv"}, "caller-key", {}), + (FastCRWSearchConfig, {"CRW_API_KEY": "srv"}, "caller-key", {}), +) + +_IDS = tuple(spec[0].__name__ for spec in PROVIDERS) + + +@pytest.mark.parametrize( + "config_cls, server_env, caller_key, extra_env", PROVIDERS, ids=_IDS +) +def test_server_secret_refused_for_caller_api_base( + config_cls: Type[BaseSearchConfig], + server_env: Dict[str, str], + caller_key: str, + extra_env: Dict[str, str], + monkeypatch: pytest.MonkeyPatch, +) -> None: + for key, value in {**server_env, **extra_env}.items(): + monkeypatch.setenv(key, value) + + with pytest.raises(ValueError, match="Refusing to send the server-configured"): + config_cls().validate_environment(headers={}, api_base=ATTACKER_BASE) + + +@pytest.mark.parametrize( + "config_cls, server_env, caller_key, extra_env", PROVIDERS, ids=_IDS +) +def test_caller_supplied_key_is_honored_for_custom_api_base( + config_cls: Type[BaseSearchConfig], + server_env: Dict[str, str], + caller_key: str, + extra_env: Dict[str, str], + monkeypatch: pytest.MonkeyPatch, +) -> None: + for key, value in {**server_env, **extra_env}.items(): + monkeypatch.setenv(key, value) + + # An explicit caller key is the caller's own credential, so pointing it at + # the caller's own host must be allowed. + config_cls().validate_environment( + headers={}, api_key=caller_key, api_base=ATTACKER_BASE + ) + + +@pytest.mark.parametrize( + "config_cls, server_env, caller_key, extra_env", PROVIDERS, ids=_IDS +) +def test_server_secret_used_without_caller_api_base( + config_cls: Type[BaseSearchConfig], + server_env: Dict[str, str], + caller_key: str, + extra_env: Dict[str, str], + monkeypatch: pytest.MonkeyPatch, +) -> None: + for key, value in {**server_env, **extra_env}.items(): + monkeypatch.setenv(key, value) + + # No caller-supplied api_base -> the request targets the trusted default, so + # the server secret is still used and nothing is refused. + config_cls().validate_environment(headers={}) + + +def test_keyless_provider_allows_caller_api_base( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv("SEARXNG_API_KEY", raising=False) + + headers = SearXNGSearchConfig().validate_environment( + headers={}, api_base="https://my-searxng.internal" + ) + + assert "Authorization" not in headers + + +def test_operator_env_base_override_is_trusted( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("SERPER_API_KEY", "srv") + monkeypatch.setenv("SERPER_API_BASE", "https://serper.internal.corp") + + # Mirrors the second validate_environment call in the search handler, which + # receives the already-resolved operator base as api_base. + headers = SerperSearchConfig().validate_environment( + headers={}, api_base="https://serper.internal.corp/search" + ) + + assert headers["X-API-KEY"] == "srv" + + +class TestResolveServerApiKey: + def test_caller_key_short_circuits(self) -> None: + result = BaseSearchConfig().resolve_server_api_key( + caller_api_key="mine", + caller_api_base=ATTACKER_BASE, + key_env_vars=("SERPER_API_KEY",), + base_env_var="SERPER_API_BASE", + default_api_base="https://google.serper.dev", + ) + assert result == "mine" + + def test_returns_none_when_no_server_secret( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + monkeypatch.delenv("SEARXNG_API_KEY", raising=False) + result = BaseSearchConfig().resolve_server_api_key( + caller_api_key=None, + caller_api_base=ATTACKER_BASE, + key_env_vars=("SEARXNG_API_KEY",), + base_env_var="SEARXNG_API_BASE", + default_api_base=None, + ) + assert result is None + + def test_first_set_env_var_wins(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("PARALLEL_AI_API_KEY", raising=False) + monkeypatch.setenv("PARALLEL_API_KEY", "second") + result = BaseSearchConfig().resolve_server_api_key( + caller_api_key=None, + caller_api_base=None, + key_env_vars=("PARALLEL_AI_API_KEY", "PARALLEL_API_KEY"), + base_env_var="PARALLEL_AI_API_BASE", + default_api_base="https://api.parallel.ai", + ) + assert result == "second" + + +class TestIsTrustedSearchApiBase: + def test_matches_default_host(self) -> None: + assert _is_trusted_search_api_base( + "https://google.serper.dev/search", "https://google.serper.dev", None + ) + + def test_foreign_host_untrusted(self) -> None: + assert not _is_trusted_search_api_base( + ATTACKER_BASE, "https://google.serper.dev", None + ) + + def test_env_override_host_trusted(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("SERPER_API_BASE", "https://serper.internal.corp") + assert _is_trusted_search_api_base( + "https://serper.internal.corp/search", + "https://google.serper.dev", + "SERPER_API_BASE", + ) + + def test_schemeless_candidate_untrusted(self) -> None: + # Without a scheme urlsplit puts the value in the path, leaving an empty + # netloc; an unparseable host must never be treated as trusted. + assert not _is_trusted_search_api_base( + "attacker.example.com", "https://google.serper.dev", None + ) + + +@pytest.mark.asyncio +async def test_asearch_does_not_leak_server_key_to_caller_api_base( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """End-to-end regression on the reported vector: a search call with a foreign + api_base and no caller key must fail without any outbound request carrying the + server-configured key.""" + monkeypatch.setenv("SERPER_API_KEY", "sk-server-secret") + monkeypatch.delenv("SERPER_API_BASE", raising=False) + + with ( + patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post, + patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get", + new_callable=AsyncMock, + ) as mock_get, + ): + with pytest.raises(Exception): + await litellm.asearch( + query="secrets", + search_provider="serper", + api_base=ATTACKER_BASE, + ) + + mock_post.assert_not_called() + mock_get.assert_not_called() + + +@pytest.mark.parametrize( + "provider, key_env, server_key, extra_env", + [ + ("searchapi", "SEARCHAPI_API_KEY", "sk-server-searchapi", {}), + ( + "google_pse", + "GOOGLE_PSE_API_KEY", + "sk-server-google", + {"GOOGLE_PSE_ENGINE_ID": "engine-id"}, + ), + ], +) +@pytest.mark.asyncio +async def test_query_param_key_not_leaked_with_dummy_caller_key( + provider: str, + key_env: str, + server_key: str, + extra_env: Dict[str, str], + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Providers that send the key as a URL query param resolve it in + transform_search_request, not validate_environment. A caller who passes a + dummy api_key to clear the validate_environment short-circuit must not cause + the server key to be placed in the URL sent to their own api_base.""" + monkeypatch.setenv(key_env, server_key) + for name, value in extra_env.items(): + monkeypatch.setenv(name, value) + + captured: Dict[str, str] = {} + + async def fake_get(self, *args, **kwargs): # type: ignore[no-untyped-def] + captured["url"] = kwargs.get("url") or (args[0] if args else "") + raise RuntimeError("stop after capturing the outbound url") + + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get", + fake_get, + ): + with pytest.raises(Exception): + await litellm.asearch( + query="secrets", + search_provider=provider, + api_key="sk-CALLER-DUMMY", + api_base=ATTACKER_BASE, + ) + + assert captured["url"], "expected an outbound request to be attempted" + assert server_key not in captured["url"] + assert "sk-CALLER-DUMMY" in captured["url"] diff --git a/tests/test_litellm/llms/cloudflare/test_cloudflare_transformation.py b/tests/test_litellm/llms/cloudflare/test_cloudflare_transformation.py index cecb6024de1..1a46015ceb9 100644 --- a/tests/test_litellm/llms/cloudflare/test_cloudflare_transformation.py +++ b/tests/test_litellm/llms/cloudflare/test_cloudflare_transformation.py @@ -3,25 +3,190 @@ import pytest from litellm.llms.cloudflare.chat.transformation import CloudflareChatConfig -def test_get_complete_url_encodes_model_path_segment(): +def test_supported_params_include_tools_and_tool_choice(): config = CloudflareChatConfig() - assert ( - config.get_complete_url( - api_base="https://api.cloudflare.com/client/v4/accounts/acct/ai/run/", - api_key="cf-key", - model="@cf/meta/llama?x=1#frag", - optional_params={}, - litellm_params={}, - ) - == "https://api.cloudflare.com/client/v4/accounts/acct/ai/run/%40cf/meta/llama%3Fx%3D1%23frag" + params = config.get_supported_openai_params(model="@cf/meta/llama-2-7b-chat-int8") + + assert "tools" in params + assert "tool_choice" in params + assert "stream" in params + assert "max_tokens" in params + + +def test_get_complete_url_defaults_to_openai_compatible_endpoint(monkeypatch): + monkeypatch.setenv("CLOUDFLARE_ACCOUNT_ID", "acct") + config = CloudflareChatConfig() + + url = config.get_complete_url( + api_base=None, + api_key="cf-key", + model="@cf/meta/llama-2-7b-chat-int8", + optional_params={}, + litellm_params={}, ) - with pytest.raises(ValueError, match="dot path segment"): + assert ( + url + == "https://api.cloudflare.com/client/v4/accounts/acct/ai/v1/chat/completions" + ) + assert "/ai/run/" not in url + + +def test_get_complete_url_appends_chat_completions_to_explicit_base(): + config = CloudflareChatConfig() + + url = config.get_complete_url( + api_base="https://api.cloudflare.com/client/v4/accounts/acct/ai/v1", + api_key="cf-key", + model="@cf/meta/llama-2-7b-chat-int8", + optional_params={}, + litellm_params={}, + ) + + assert ( + url + == "https://api.cloudflare.com/client/v4/accounts/acct/ai/v1/chat/completions" + ) + assert "/ai/run/" not in url + + +def test_get_complete_url_is_idempotent_for_full_base(): + config = CloudflareChatConfig() + + url = config.get_complete_url( + api_base="https://api.cloudflare.com/client/v4/accounts/acct/ai/v1/chat/completions", + api_key="cf-key", + model="@cf/meta/llama-2-7b-chat-int8", + optional_params={}, + litellm_params={}, + ) + + assert ( + url + == "https://api.cloudflare.com/client/v4/accounts/acct/ai/v1/chat/completions" + ) + + +def test_get_complete_url_falls_back_to_account_id_when_base_is_empty(monkeypatch): + monkeypatch.setenv("CLOUDFLARE_ACCOUNT_ID", "acct") + config = CloudflareChatConfig() + + url = config.get_complete_url( + api_base="", + api_key="cf-key", + model="@cf/meta/llama-2-7b-chat-int8", + optional_params={}, + litellm_params={}, + ) + + assert ( + url + == "https://api.cloudflare.com/client/v4/accounts/acct/ai/v1/chat/completions" + ) + + +def test_get_complete_url_raises_when_account_id_and_base_missing(monkeypatch): + monkeypatch.delenv("CLOUDFLARE_ACCOUNT_ID", raising=False) + config = CloudflareChatConfig() + + with pytest.raises(ValueError, match="Missing CLOUDFLARE_ACCOUNT_ID"): config.get_complete_url( - api_base="https://api.cloudflare.com/client/v4/accounts/acct/ai/run/", + api_base=None, api_key="cf-key", - model="../../accounts/other", + model="@cf/meta/llama-2-7b-chat-int8", optional_params={}, litellm_params={}, ) + + +def test_get_complete_url_raises_when_account_id_is_empty(monkeypatch): + monkeypatch.setenv("CLOUDFLARE_ACCOUNT_ID", " ") + config = CloudflareChatConfig() + + with pytest.raises(ValueError, match="Missing CLOUDFLARE_ACCOUNT_ID"): + config.get_complete_url( + api_base=None, + api_key="cf-key", + model="@cf/meta/llama-2-7b-chat-int8", + optional_params={}, + litellm_params={}, + ) + + +def test_get_complete_url_migrates_legacy_ai_run_base(): + config = CloudflareChatConfig() + + url = config.get_complete_url( + api_base="https://api.cloudflare.com/client/v4/accounts/acct/ai/run/", + api_key="cf-key", + model="@cf/meta/llama-2-7b-chat-int8", + optional_params={}, + litellm_params={}, + ) + + assert ( + url + == "https://api.cloudflare.com/client/v4/accounts/acct/ai/v1/chat/completions" + ) + assert "/ai/run" not in url + + +def test_transform_request_passes_tools_through_in_openai_format(): + config = CloudflareChatConfig() + tools = [ + { + "type": "function", + "function": { + "name": "get_weather", + "parameters": { + "type": "object", + "properties": {"city": {"type": "string"}}, + }, + }, + } + ] + messages = [{"role": "user", "content": "weather in nyc?"}] + + body = config.transform_request( + model="@cf/meta/llama-2-7b-chat-int8", + messages=messages, + optional_params={"tools": tools, "tool_choice": "auto"}, + litellm_params={}, + headers={}, + ) + + assert body["messages"] == messages + assert body["model"] == "@cf/meta/llama-2-7b-chat-int8" + assert body["tools"] == tools + assert body["tool_choice"] == "auto" + + +def test_validate_environment_requires_api_key(): + config = CloudflareChatConfig() + + with pytest.raises(ValueError, match="Missing Cloudflare API Key"): + config.validate_environment( + headers={}, + model="@cf/meta/llama-2-7b-chat-int8", + messages=[], + optional_params={}, + litellm_params={}, + api_key=None, + ) + + +def test_validate_environment_sets_bearer_and_content_type(): + config = CloudflareChatConfig() + + headers = config.validate_environment( + headers={}, + model="@cf/meta/llama-2-7b-chat-int8", + messages=[], + optional_params={}, + litellm_params={}, + api_key="cf-key", + ) + + assert headers["Authorization"] == "Bearer cf-key" + assert headers["Content-Type"] == "application/json" diff --git a/tests/test_litellm/llms/parallel_ai/test_parallel_ai_search.py b/tests/test_litellm/llms/parallel_ai/test_parallel_ai_search.py index b5c1a86205b..7be295826e3 100644 --- a/tests/test_litellm/llms/parallel_ai/test_parallel_ai_search.py +++ b/tests/test_litellm/llms/parallel_ai/test_parallel_ai_search.py @@ -293,7 +293,10 @@ class TestParallelAISearch: ], ) @pytest.mark.asyncio - async def test_custom_api_base_appends_v1_search(self, api_base): + async def test_custom_api_base_appends_v1_search(self, api_base, monkeypatch): + # Operator points at an internal base via the env override (a trusted + # host), so the server key is still used and the URL is normalized. + monkeypatch.setenv("PARALLEL_AI_API_BASE", api_base) with patch( "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", new_callable=AsyncMock, @@ -303,7 +306,6 @@ class TestParallelAISearch: await litellm.asearch( query="AI developments", search_provider="parallel_ai", - api_base=api_base, ) call_args = mock_post.call_args @@ -312,6 +314,23 @@ class TestParallelAISearch: == "https://proxy.internal.example.com/v1/search" ) + @pytest.mark.asyncio + async def test_caller_api_base_without_key_is_refused(self, monkeypatch): + # A caller-supplied api_base (untrusted host) while relying on the + # server key must be refused without any outbound request. + monkeypatch.setenv("PARALLEL_API_KEY", "server-secret") + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + with pytest.raises(Exception, match="Refusing to send"): + await litellm.asearch( + query="AI developments", + search_provider="parallel_ai", + api_base="https://attacker.example.com", + ) + mock_post.assert_not_called() + @pytest.mark.asyncio async def test_missing_api_key_raises(self, monkeypatch): monkeypatch.delenv("PARALLEL_API_KEY", raising=False) diff --git a/tests/test_litellm/llms/tinyfish/test_tinyfish_search.py b/tests/test_litellm/llms/tinyfish/test_tinyfish_search.py index 5496486765c..9870d30d488 100644 --- a/tests/test_litellm/llms/tinyfish/test_tinyfish_search.py +++ b/tests/test_litellm/llms/tinyfish/test_tinyfish_search.py @@ -63,23 +63,17 @@ class TestTinyfishSearchConfig: assert headers["X-API-Key"] == "sk-tinyfish-test" assert headers["Accept"] == "application/json" - def test_validate_environment_from_env(self): + def test_validate_environment_from_env(self, monkeypatch): + monkeypatch.setenv("TINYFISH_API_KEY", "sk-from-env") config = TinyfishSearchConfig() - with patch( - "litellm.llms.tinyfish.search.transformation.get_secret_str", - return_value="sk-from-env", - ): - headers = config.validate_environment(headers={}) + headers = config.validate_environment(headers={}) assert headers["X-API-Key"] == "sk-from-env" - def test_validate_environment_missing_key(self): + def test_validate_environment_missing_key(self, monkeypatch): + monkeypatch.delenv("TINYFISH_API_KEY", raising=False) config = TinyfishSearchConfig() - with patch( - "litellm.llms.tinyfish.search.transformation.get_secret_str", - return_value=None, - ): - with pytest.raises(ValueError, match="TINYFISH_API_KEY"): - config.validate_environment(headers={}) + with pytest.raises(ValueError, match="TINYFISH_API_KEY"): + config.validate_environment(headers={}) def test_validate_environment_uses_api_base_kwarg(self): config = TinyfishSearchConfig() diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py new file mode 100644 index 00000000000..7e028064e4c --- /dev/null +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -0,0 +1,333 @@ +"""Tests for the optional Rust-backed OCR path (``litellm/ocr/rust_bridge.py``).""" + +import importlib +import sys +import types + +import httpx +import pytest + +import litellm +from litellm.llms.base_llm.ocr.transformation import OCRResponse + +# `litellm/__init__.py` does `from .ocr.main import *`, which binds the `ocr` +# function onto `litellm.ocr` and shadows the submodule, so import the modules +# explicitly via importlib rather than attribute traversal. +ocr_main = importlib.import_module("litellm.ocr.main") +rust_bridge = importlib.import_module("litellm.ocr.rust_bridge") + +MODEL = "mistral/mistral-ocr-latest" +DOCUMENT = {"type": "document_url", "document_url": "https://example.com/doc.pdf"} + +FAKE_OCR_RESPONSE = { + "pages": [{"index": 0, "markdown": "hello world"}], + "model": "mistral-ocr-2505-completion", + "document_annotation": None, + "usage_info": {"pages_processed": 1}, + "object": "ocr", +} + + +class RecordingBridge: + """A fake ``RustOcr`` callable that records the args it was handed.""" + + def __init__(self): + self.calls = [] + + def __call__( + self, model, document, api_key, api_base, optional_params, timeout_seconds + ): + self.calls.append( + { + "model": model, + "document": document, + "api_key": api_key, + "api_base": api_base, + "optional_params": optional_params, + "timeout_seconds": timeout_seconds, + } + ) + return dict(FAKE_OCR_RESPONSE) + + +class RecordingLogging: + """A spy standing in for ``LiteLLMLoggingObj`` to capture ``pre_call``.""" + + def __init__(self): + self.pre_call_kwargs = None + + def pre_call(self, *, input, api_key, additional_args): + self.pre_call_kwargs = { + "input": input, + "api_key": api_key, + "additional_args": additional_args, + } + + +class FakeOCRConfig: + """A stand-in ``BaseOCRConfig`` that echoes the request it would build.""" + + def validate_environment( + self, *, headers, model, api_key, api_base, litellm_params + ): + return {"authorization": f"Bearer {api_key}"} + + def get_complete_url(self, *, api_base, model, optional_params, litellm_params): + return f"{api_base or 'https://api.mistral.ai/v1'}/ocr" + + +@pytest.fixture(autouse=True) +def _reset_rust_flag(): + """Keep the global toggle isolated between tests.""" + rust_bridge.use_litellm_rust(False, ocr=None) + yield + rust_bridge.use_litellm_rust(False, ocr=None) + + +@pytest.fixture +def fake_bridge(): + """Enable the Rust path with an injected recording bridge (no native wheel).""" + bridge = RecordingBridge() + litellm.use_litellm_rust(True, ocr=bridge) + return bridge + + +def test_use_litellm_rust_toggles_flag(): + assert rust_bridge.rust_ocr_enabled() is False + litellm.use_litellm_rust() + assert rust_bridge.rust_ocr_enabled() is True + litellm.use_litellm_rust(False) + assert rust_bridge.rust_ocr_enabled() is False + + +def test_load_rust_ocr_returns_injected_impl(): + bridge = RecordingBridge() + litellm.use_litellm_rust(True, ocr=bridge) + assert rust_bridge.load_rust_ocr() is bridge + + +def test_toggle_without_ocr_arg_preserves_injected_impl(): + """Regression: routine enable/disable calls must not clobber a prior injection. + + Earlier, ``use_litellm_rust()`` unconditionally assigned the keyword default + of ``None`` to ``_rust_ocr_impl``, silently dropping a custom bridge whenever + a caller toggled the flag without re-passing ``ocr=``. + """ + bridge = RecordingBridge() + litellm.use_litellm_rust(True, ocr=bridge) + + litellm.use_litellm_rust(False) + assert rust_bridge.load_rust_ocr() is bridge + litellm.use_litellm_rust(True) + assert rust_bridge.load_rust_ocr() is bridge + + +def test_explicit_ocr_none_clears_injected_impl(): + bridge = RecordingBridge() + litellm.use_litellm_rust(True, ocr=bridge) + + litellm.use_litellm_rust(True, ocr=None) + assert rust_bridge.load_rust_ocr() is None + + +def test_load_rust_ocr_none_when_extension_absent(): + """With no injected impl and no compiled wheel, the loader returns None so the + caller degrades to the Python path instead of raising ImportError.""" + litellm.use_litellm_rust(True) # no impl injected; extension isn't built in CI + assert rust_bridge.load_rust_ocr() is None + + +def test_load_rust_ocr_uses_compiled_extension(monkeypatch): + """With no injected impl but a compiled ``litellm_python_bridge`` importable, + the loader returns the extension's ``ocr`` callable. The native wheel isn't + built in CI, so stand in a fake module via ``sys.modules``.""" + fake_module = types.ModuleType("litellm_python_bridge") + fake_module.ocr = lambda **kwargs: dict(FAKE_OCR_RESPONSE) # type: ignore[attr-defined] + monkeypatch.setitem(sys.modules, "litellm_python_bridge", fake_module) + + litellm.use_litellm_rust(True) # enabled, no impl injected -> import the extension + assert rust_bridge.load_rust_ocr() is fake_module.ocr + + +def test_timeout_to_seconds_handles_float_timeout_and_none(): + assert ocr_main._timeout_to_seconds(12.5) == 12.5 + assert ocr_main._timeout_to_seconds(None) is None + assert ocr_main._timeout_to_seconds(httpx.Timeout(30.0, read=42.0)) == 42.0 + + +def test_run_rust_ocr_forwards_args_and_wraps_response(): + bridge = RecordingBridge() + logging_obj = RecordingLogging() + + response = ocr_main._run_rust_ocr( + rust_ocr=bridge, + logging_obj=logging_obj, + provider_config=FakeOCRConfig(), + resolve_api_key=lambda _name: None, + model="mistral-ocr-latest", + document=DOCUMENT, + api_key="sk-test", + api_base="https://proxy.internal", + optional_params={"include_image_base64": True}, + litellm_params={}, + timeout_seconds=12.5, + ) + + assert isinstance(response, OCRResponse) + assert response.pages[0].markdown == "hello world" + call = bridge.calls[0] + assert call == { + "model": "mistral-ocr-latest", + "document": DOCUMENT, + "api_key": "sk-test", + "api_base": "https://proxy.internal", + "optional_params": {"include_image_base64": True}, + "timeout_seconds": 12.5, + } + + +def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing(): + """No explicit api_key: the resolver (get_secret_str in production) supplies it, + so secret-manager backends (AWS/Azure/GCP/Vault) work like the Python path.""" + bridge = RecordingBridge() + + ocr_main._run_rust_ocr( + rust_ocr=bridge, + logging_obj=RecordingLogging(), + provider_config=FakeOCRConfig(), + resolve_api_key=lambda name: ( + "sk-from-vault" if name == "MISTRAL_API_KEY" else None + ), + model="mistral-ocr-latest", + document=DOCUMENT, + api_key=None, + api_base=None, + optional_params={}, + litellm_params={}, + timeout_seconds=None, + ) + + assert bridge.calls[0]["api_key"] == "sk-from-vault" + + +def test_run_rust_ocr_prefers_explicit_key_over_resolver(): + bridge = RecordingBridge() + resolver_calls = [] + + def _resolver(name): + resolver_calls.append(name) + return "sk-from-vault" + + ocr_main._run_rust_ocr( + rust_ocr=bridge, + logging_obj=RecordingLogging(), + provider_config=FakeOCRConfig(), + resolve_api_key=_resolver, + model="mistral-ocr-latest", + document=DOCUMENT, + api_key="sk-explicit", + api_base=None, + optional_params={}, + litellm_params={}, + timeout_seconds=None, + ) + + assert bridge.calls[0]["api_key"] == "sk-explicit" + assert resolver_calls == [] # resolver never consulted when a key is supplied + + +def test_run_rust_ocr_runs_pre_call_logging(): + """The Rust shortcut must run pre_call so callbacks and spend tracking fire.""" + logging_obj = RecordingLogging() + + ocr_main._run_rust_ocr( + rust_ocr=RecordingBridge(), + logging_obj=logging_obj, + provider_config=FakeOCRConfig(), + resolve_api_key=lambda _name: None, + model="mistral-ocr-latest", + document=DOCUMENT, + api_key="sk-test", + api_base="https://api.mistral.ai/v1", + optional_params={"include_image_base64": True}, + litellm_params={}, + timeout_seconds=None, + ) + + assert logging_obj.pre_call_kwargs is not None + assert logging_obj.pre_call_kwargs["input"] == "OCR document processing" + additional_args = logging_obj.pre_call_kwargs["additional_args"] + complete_input = additional_args["complete_input_dict"] + assert complete_input["document"] == DOCUMENT + assert complete_input["include_image_base64"] is True + # The logged request mirrors what Rust sends: resolved URL + headers. + assert additional_args["api_base"] == "https://api.mistral.ai/v1/ocr" + assert additional_args["headers"] == {"authorization": "Bearer sk-test"} + + +def test_ocr_routes_to_rust_when_enabled(fake_bridge): + response = litellm.ocr( + model=MODEL, + document=DOCUMENT, + api_key="sk-test", + include_image_base64=True, + ) + + assert isinstance(response, OCRResponse) + assert response.pages[0].markdown == "hello world" + assert len(fake_bridge.calls) == 1 + call = fake_bridge.calls[0] + # Provider prefix is stripped before reaching the bridge. + assert call["model"] == "mistral-ocr-latest" + assert call["document"] == DOCUMENT + assert call["api_key"] == "sk-test" + # Raw OCR params ride along in optional_params; Rust filters to supported keys. + assert call["optional_params"].get("include_image_base64") is True + + +def test_ocr_forwards_timeout_to_rust(fake_bridge): + """Caller-supplied timeout must flow into the Rust bridge so the fixed 600s + client ceiling doesn't silently override shorter deadlines.""" + litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test", timeout=12.5) + + assert fake_bridge.calls[0]["timeout_seconds"] == 12.5 + + +def test_ocr_passes_default_request_timeout_to_rust(fake_bridge): + """When no explicit timeout is given, the library default (request_timeout) + must still be forwarded so the Rust path matches the Python path's deadline.""" + from litellm.constants import request_timeout + + litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test") + + assert fake_bridge.calls[0]["timeout_seconds"] == float(request_timeout) + + +def test_ocr_does_not_route_to_rust_when_disabled(): + """With the flag off, the bridge must not be consulted even if an impl exists.""" + bridge = RecordingBridge() + litellm.use_litellm_rust(False, ocr=bridge) + + assert rust_bridge.rust_ocr_enabled() is False + # The impl stays available for injection, but the disabled flag gates usage, + # so ocr() never reaches the Rust path (asserted via the enabled-path test). + assert bridge.calls == [] + + +def test_ocr_falls_back_to_python_when_bridge_unavailable(monkeypatch): + """Rust enabled but no bridge available (no injected impl, no compiled wheel): + ocr() must degrade to the Python HTTP handler instead of raising.""" + litellm.use_litellm_rust(True) # enabled, but load_rust_ocr() returns None in CI + + captured = {} + + def fake_handler_ocr(**kwargs): + captured["called"] = True + return OCRResponse(pages=[], model="mistral-ocr-latest", object="ocr") + + monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", fake_handler_ocr) + + response = litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test") + + assert captured.get("called") is True # Python path was used + assert isinstance(response, OCRResponse) diff --git a/tests/test_litellm/test_cloudflare_workers_ai_model_metadata.py b/tests/test_litellm/test_cloudflare_workers_ai_model_metadata.py new file mode 100644 index 00000000000..9ca4515239a --- /dev/null +++ b/tests/test_litellm/test_cloudflare_workers_ai_model_metadata.py @@ -0,0 +1,89 @@ +""" +Regression tests for the Cloudflare Workers AI text-generation catalog in the +model-cost map. + +The Cloudflare list was badly stale (only 4 ancient entries). These tests pin +the newly added current Workers AI models (sourced from Cloudflare's live +``/ai/models/search?task=Text Generation`` catalog) and guard against the root +``model_prices_and_context_window.json`` and the bundled +``litellm/model_prices_and_context_window_backup.json`` drifting out of sync for +the ``cloudflare/`` namespace. +""" + +import json +import os + +import pytest + +import litellm + +ROOT_MAP = os.path.join( + os.path.dirname(os.path.dirname(litellm.__file__)), + "model_prices_and_context_window.json", +) +BACKUP_MAP = os.path.join( + os.path.dirname(litellm.__file__), + "model_prices_and_context_window_backup.json", +) + + +@pytest.fixture(autouse=True) +def _use_local_model_cost_map(monkeypatch): + original_model_cost = litellm.model_cost + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + litellm.model_cost = litellm.get_model_cost_map(url="") + try: + yield + finally: + litellm.model_cost = original_model_cost + + +def _load(path: str) -> dict: + with open(path, encoding="utf-8") as f: + return json.load(f) + + +def _cloudflare_keys(data: dict) -> set: + return {k for k in data if k.startswith("cloudflare/")} + + +def test_glm_5_2_entry_is_present_and_well_formed(): + entry = litellm.model_cost["cloudflare/@cf/zai-org/glm-5.2"] + assert entry["litellm_provider"] == "cloudflare" + assert entry["mode"] == "chat" + assert entry["supports_function_calling"] is True + assert entry["input_cost_per_token"] > 0 + assert entry["output_cost_per_token"] > 0 + + +def test_vision_model_is_flagged_supports_vision(): + entry = litellm.model_cost["cloudflare/@cf/meta/llama-3.2-11b-vision-instruct"] + assert entry["litellm_provider"] == "cloudflare" + assert entry.get("supports_vision") is True + + +def test_additional_current_models_are_present(): + for key in ( + "cloudflare/@cf/openai/gpt-oss-120b", + "cloudflare/@cf/meta/llama-3.3-70b-instruct-fp8-fast", + ): + entry = litellm.model_cost[key] + assert entry["litellm_provider"] == "cloudflare" + assert entry["mode"] == "chat" + assert entry["supports_function_calling"] is True + assert entry["input_cost_per_token"] > 0 + assert entry["output_cost_per_token"] > 0 + + +def test_root_and_backup_have_identical_cloudflare_keys(): + if not os.path.exists(ROOT_MAP): + pytest.skip("root cost map only ships in source checkouts") + assert _cloudflare_keys(_load(ROOT_MAP)) == _cloudflare_keys(_load(BACKUP_MAP)) + + +def test_root_and_backup_cloudflare_entries_are_byte_for_byte_equal(): + if not os.path.exists(ROOT_MAP): + pytest.skip("root cost map only ships in source checkouts") + root = {k: v for k, v in _load(ROOT_MAP).items() if k.startswith("cloudflare/")} + backup = {k: v for k, v in _load(BACKUP_MAP).items() if k.startswith("cloudflare/")} + assert root == backup diff --git a/tests/test_litellm/test_type_check_gate.py b/tests/test_litellm/test_type_check_gate.py index 18374c5db4b..e99ad0a4f41 100644 --- a/tests/test_litellm/test_type_check_gate.py +++ b/tests/test_litellm/test_type_check_gate.py @@ -56,29 +56,58 @@ def test_paths_outside_repo_are_skipped(): def test_at_or_under_ceiling_passes(): budget = {"no-any-return": {"baseline": 5, "slack": 0}} - assert gate.evaluate({"no-any-return": 5}, budget) == [] + assert gate.evaluate({"no-any-return": 5}, {}, budget) == [] def test_one_more_error_than_ceiling_fails(): budget = {"no-any-return": {"baseline": 5, "slack": 0}} - assert gate.evaluate({"no-any-return": 6}, budget) == [ - gate.Breach("no-any-return", 6, 5) + assert gate.evaluate({"no-any-return": 6}, {}, budget) == [ + gate.Breach("no-any-return", 6, 5, 6) ] def test_slack_absorbs_small_increase_then_fails_past_it(): budget = {"arg-type": {"baseline": 5, "slack": 5}} - assert gate.evaluate({"arg-type": 10}, budget) == [] - assert gate.evaluate({"arg-type": 11}, budget) == [gate.Breach("arg-type", 11, 10)] + assert gate.evaluate({"arg-type": 10}, {}, budget) == [] + assert gate.evaluate({"arg-type": 11}, {}, budget) == [ + gate.Breach("arg-type", 11, 10, 11) + ] def test_unbudgeted_new_code_uses_default_slack(): - assert gate.evaluate({"brand-new": gate.DEFAULT_SLACK}, {}) == [] - assert gate.evaluate({"brand-new": gate.DEFAULT_SLACK + 1}, {}) == [ - gate.Breach("brand-new", gate.DEFAULT_SLACK + 1, gate.DEFAULT_SLACK) + assert gate.evaluate({"brand-new": gate.DEFAULT_SLACK}, {}, {}) == [] + assert gate.evaluate({"brand-new": gate.DEFAULT_SLACK + 1}, {}, {}) == [ + gate.Breach( + "brand-new", + gate.DEFAULT_SLACK + 1, + gate.DEFAULT_SLACK, + gate.DEFAULT_SLACK + 1, + ) ] +def test_drift_already_over_cap_in_base_is_not_blamed_on_a_flat_change(): + # The bystander case: a rule sits over its ceiling because two earlier PRs + # summed past it. A PR that branches off that base and adds nothing must pass + # -- total > cap but total == base, so the `> base` guard spares it. + budget = {"arg-type": {"baseline": 5, "slack": 5}} + assert gate.evaluate({"arg-type": 12}, {"arg-type": 12}, budget) == [] + + +def test_change_that_grows_an_over_cap_rule_is_blamed_for_only_what_it_added(): + # Over cap AND above base: blamed, and `added` is the delta vs base, not the + # whole overage, so the message points at this change's contribution. + budget = {"arg-type": {"baseline": 5, "slack": 5}} + assert gate.evaluate({"arg-type": 14}, {"arg-type": 12}, budget) == [ + gate.Breach("arg-type", 14, 10, 2) + ] + + +def test_reducing_an_over_cap_rule_below_base_passes(): + budget = {"arg-type": {"baseline": 5, "slack": 5}} + assert gate.evaluate({"arg-type": 11}, {"arg-type": 12}, budget) == [] + + def test_no_output_against_a_nonempty_budget_is_a_vacuous_run(): # A crashed type checker emits nothing; the gate must not certify it as clean. budget = {"no-untyped-def": {"baseline": 4888, "slack": 10}}