This commit is contained in:
Yujong Lee 2026-09-12 08:07:31 -07:00
parent 43a19d81ab
commit c6fc5e185f
29 changed files with 619 additions and 816 deletions

View file

@ -1,29 +0,0 @@
# Adding a provider / route to litellm-rust
Everything for a route lives in `crates/core/src/<route>/`; `crates/core/src/messages` is the reference. A host (the axum gateway, the Python bridge) only calls the route's entrypoint.
1. **Entrypoint** — `mod.rs`: `pub async fn <route>(request) -> CoreResult<Response>`, the Rust equivalent of `litellm.<route>()`, plus a `<route>_stream` variant when the route streams. It is the only thing a host touches.
2. **Transform contract** — `transformation.rs`: a `…ProviderConfig` trait (URL build + request/response transforms) with types in `types.rs`.
3. **Provider config** — `crates/core/src/providers/<provider>/<route>/transformation.rs`: implement that trait as a `const <PROVIDER>_<ROUTE>_CONFIG`, mirroring the Python provider tree. Add parity unit tests.
4. **Prepare + handler** — `prepare.rs` resolves provider/model, credentials, auth headers, and URL, then transforms the request; `handler.rs` performs the provider call through the shared client in `client.rs` and transforms the response.
## Coding standards
Before writing new logic, look for an existing base to extend. When a change is
“the same behavior for one more provider/endpoint/integration”, the codebase
almost always already has a shared abstraction for it (for example, provider
`BaseConfig` transformation classes in `litellm/llms/base_llm/`, shared
helpers in `litellm_core_utils/`, typed request/response models, or factory
functions). Find it first with a search, then add the new variant by inheriting
from or composing that base, overriding only what genuinely differs (model
name, parameter mapping, or auth).
Never copy an existing implementation and edit it in place, and never hand-roll
a parallel version of logic a base already provides. If you catch yourself
writing a second copy of a pattern that exists twice already, stop and extract a
base instead: put the shared shape in one place and make both call sites thin
variants of it. The test for a good abstraction is that adding the next provider
is a few declarative lines, not a new file of duplicated flow. Only diverge from
the base when behavior is genuinely different, and say so explicitly in the PR.
**Calling:** hosts invoke the core entrypoint — the Python bridge and the `ai-gateway` route service both call `litellm_core::messages::messages`. Never add a provider handler to `ai-gateway`. Register new modules in `lib.rs` / `mod.rs`, then run the commands under "Checks" in [CLAUDE.md](CLAUDE.md).

View file

@ -1,45 +0,0 @@
# AGENTS.md
litellm-rust has six crates. A crate is a layer or shared foundation, not a route. Routes (ocr, realtime, chat) and providers (mistral, openai) are modules inside the layers.
## Crates
| Crate | Role |
|-------|------|
| litellm-core | The LiteLLM SDK in Rust. One public entrypoint per top-level call (`messages::messages()`), owning types, transforms, provider resolution, auth, and the provider HTTP call. Call it, get a typed response. |
| litellm-token-counter | Standalone input token counting shared by host integrations without pulling in the full SDK. |
| litellm-config | Config-loading boundary. Returns resolved core deployment data and optionally delegates loading to Python. |
| litellm-ai-gateway | The axum server (behind the `server` feature) plus the WebSocket hosts. Translates HTTP/WS to core entrypoints; owns no provider logic and no handlers. |
| litellm-python-interop | Domain-neutral PyO3 foundation for GIL handling and typed Python/Serde conversion. |
| litellm-python-bridge | PyO3 cdylib exposing LiteLLM Rust APIs to the Python SDK. Owns API registration, domain wiring, and Python exception mapping. |
Dependency direction is acyclic: `litellm-config` depends on `litellm-core`, the gateway depends on both, and `litellm-python-bridge` depends on the domain layers, `litellm-token-counter`, and `litellm-python-interop`. The token counter and interop foundations depend on no LiteLLM domain crate.
## Where a route lives
A top-level LiteLLM call is a module under `crates/core/src/<route>/`, shaped like `messages`:
```
core/src/messages/
mod.rs # pub async fn messages(..) -> CoreResult<..> (+ messages_stream for SSE)
types.rs # request/response types, MessagesRequest
transformation.rs # the provider template trait
prepare.rs # provider resolution, auth headers, URL
handler.rs # the provider call
client.rs # the shared reqwest client
```
Handlers never live in `ai-gateway`. `ocr`, `audio_transcription`, and `realtime` are still hosted there from before this rule; they move to `core` as they are touched.
Adding a crate: default to a module. A new crate requires a real trigger: separate artifact (binary/cdylib), proc-macro, shared foundation, or publishable standalone. A new provider or route is none of these.
Adding a crate fails crates/core/tests/workspace_crate_allowlist.rs until you update its allowlist and this file — intentional.
## Style
All Rust in `litellm-rust/` follows the official Rust Style Guide:
https://doc.rust-lang.org/style-guide/
`rustfmt` implements its formatting by default, so run `cargo fmt` before committing; CI gates every PR on `cargo fmt --check`. Do not hand-format against rustfmt or add a `rustfmt.toml` that diverges from the default style.
Beyond formatting, follow the guide's naming and idiom conventions rustfmt cannot auto-apply: `snake_case` items/functions/modules, `UpperCamelCase` types/traits/variants, `SCREAMING_SNAKE_CASE` constants/statics (acronyms as one word, e.g. `HttpClient`), and the import grouping and item ordering it prescribes. See CLAUDE.md for the detailed version.

View file

@ -1,189 +0,0 @@
# CLAUDE.md
This file defines the rules for Rust work in LiteLLM.
## Provider Coding Standards
Before writing new logic, look for an existing base to extend. When a change is
“the same behavior for one more provider/endpoint/integration”, the codebase
almost always already has a shared abstraction for it (for example, provider
`BaseConfig` transformation classes in `litellm/llms/base_llm/`, shared
helpers in `litellm_core_utils/`, typed request/response models, or factory
functions). Find it first with a search, then add the new variant by inheriting
from or composing that base, overriding only what genuinely differs (model
name, parameter mapping, or auth).
Never copy an existing implementation and edit it in place, and never hand-roll
a parallel version of logic a base already provides. If you catch yourself
writing a second copy of a pattern that exists twice already, stop and extract a
base instead: put the shared shape in one place and make both call sites thin
variants of it. The test for a good abstraction is that adding the next provider
is a few declarative lines, not a new file of duplicated flow. Only diverge from
the base when behavior is genuinely different, and say so explicitly in the PR.
## Crates (see AGENTS.md)
`litellm-core` **is** the LiteLLM SDK in Rust: it makes the LLM call.
`litellm-config` is the config-loading boundary and returns resolved core types.
`litellm-ai-gateway` is an HTTP/WebSocket server in front of it, and
`litellm-python-bridge` exposes it to the Python SDK. `litellm-python-interop`
holds domain-neutral PyO3 primitives shared by Python-facing Rust code. A crate
is a layer or shared foundation, not a route; add modules, not crates.
## Core Boundary
`litellm-core` owns the whole call. The Rust equivalent of `litellm.messages()`
is `litellm_core::messages::messages(request).await`: you call it, it does the
provider call, and you get a typed non-streaming response back.
Route-level Rust structure mirrors LiteLLM's Python responsibilities:
- `core/src/<route>/` owns the route end to end: the public entrypoint fn named
after the route in `mod.rs`, the request/response types (`types.rs`), the
provider template trait (`transformation.rs`), the provider/auth/URL
resolution (`prepare.rs`), the HTTP client (`client.rs`), and the handler that
performs the call (`handler.rs`). `core/src/messages` is the reference.
- `core/src/providers/<provider>/<route>/transformation.rs` owns the
provider-specific transform. For Anthropic Messages, this means
`core/src/providers/anthropic/messages/transformation.rs`.
- Handlers live in `core`, never in a host. `ai-gateway` must not contain a
route handler that talks to a provider; its axum route reads the HTTP request,
picks a deployment, and calls the `core` entrypoint. `python-bridge` marshals
Python objects and calls the same entrypoint.
Streaming keeps the same shape: the route entrypoint has a `<route>_stream`
variant in `core` that returns the upstream response so a host can splice it to
its own caller; the host still owns no provider logic.
Call-hook and lifecycle instrumentation, including phase timing, usage
accumulation, and callback payload construction, always lives in `core`.
Hosts feed observed events into core and dispatch the completed payloads through
their I/O logger; hosts must not own callback orchestration.
Allowed in `core`:
- The public entrypoint for a top-level LiteLLM call
- Request/response transforms and stream chunk normalization
- Provider resolution, auth header construction, and URL building
- The provider HTTP call itself, through a shared reused client with connect and
request timeouts
- Shared data types and validation errors
- Deterministic token/cost helper logic
Not allowed in `core`:
- Serving HTTP: axum routes, extractors, and transport concerns stay in the host
- Filesystem access
- Database access
- Config file reading and rollout state
- Logging callbacks, spend writes, or custom callbacks
- Global mutable runtime state
Env reads in `core` are limited to credential fallback inside a route's
`prepare.rs` (the `env_lookup` closure), mirroring what the Python SDK does when
no key is passed. Everything else config-shaped is resolved by the host and
passed in.
Routes still hosted in `ai-gateway` (`ocr`, `audio_transcription`, `realtime`)
predate this rule and are being moved into `core` route modules; do not add new
ones there, and prefer moving one when you touch it.
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.
A new provider/route may instead be implemented rust-only with no Python
reference; then the Python interface is a thin dispatch that calls Rust with no
fallback, and you state the rust-only choice explicitly in the PR. Either way
the Python side stays minimal (it only marshals inputs and calls the Rust
interface), never add a per-route feature flag, and never push provider
dispatch into `litellm/main.py`; put it in a thin dispatch class under
`litellm/llms/<provider>/<route>/`.
## 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.
## Network I/O Rules
These rules apply to every module that executes network I/O, whether it is a
`core` route handler or a host such as `ai-gateway`:
- 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.
## Rust Style Guide
All Rust in `litellm-rust/` follows the official Rust Style Guide:
https://doc.rust-lang.org/style-guide/
`rustfmt` implements the guide's formatting rules by default, so the mechanical
side is enforced for you: run `cargo fmt` before committing and CI gates every
PR on `cargo fmt --check` (see Checks). Do not hand-format against rustfmt or add
a `rustfmt.toml` that diverges from the default style; the default style *is* the
guide.
The guide also covers conventions rustfmt cannot auto-apply; follow these too:
- Naming: `snake_case` for items, functions, and modules; `UpperCamelCase` for
types, traits, and enum variants; `SCREAMING_SNAKE_CASE` for constants and
statics; acronyms count as one word (`HttpClient`, not `HTTPClient`).
- Ordering and grouping the guide prescribes: imports grouped std / external /
crate-local, derives before other attributes, and consistent item order.
- Idioms the guide recommends over the formatter fighting you (e.g. prefer
restructuring an over-long expression rather than forcing an awkward wrap).
## Constants
Magic numbers and fixed strings go in a crate-level `constants.rs`, never
hardcoded inline — the Rust mirror of Python's `litellm/constants.py`.
- Each crate that needs them has `src/constants.rs` (declared `mod constants;`);
import from it (`use crate::constants::...`). Don't scatter `const` values at
the top of feature modules.
- An env-overridable tunable still lives in `constants.rs` as its `DEFAULT_*`
value; the env read (with fallback to that default) happens at the host/config
resolution layer, not in `core`/`providers`.
- Exception: a value that is purely local to one function and has no meaning
elsewhere may stay inline, but prefer `constants.rs` when in doubt.
## 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 clippy -p litellm-core --all-targets --features bedrock-auth -- -D warnings
# the ai-gateway binary + server code is behind the `server` feature
cargo clippy -p litellm-ai-gateway --all-targets --all-features -- -D warnings
cargo test --workspace
cargo test -p litellm-core --features bedrock-auth
# the `auth`, `routes`, `state` and `realtime` tests only exist under `server`
cargo test -p litellm-ai-gateway --features server
```
When a Rust path is exposed through Python, add Python parity tests that compare
the existing Python output with the Rust-backed output.

View file

@ -1950,6 +1950,7 @@ dependencies = [
"base64 0.22.1",
"data-url",
"gcp_auth",
"mime_guess",
"moka",
"rand 0.8.7",
"reqwest 0.12.28",

View file

@ -1,56 +0,0 @@
# LiteLLM Rust
This workspace contains the staged Rust implementation for LiteLLM.
`litellm-core` is the LiteLLM SDK in Rust: one entrypoint per top-level call
that makes the LLM call and hands back a typed response, the same shape as
`litellm.messages()` in Python.
```rust
let response = litellm_core::messages::messages(MessagesRequest {
model: "claude-sonnet-4-5",
body,
api_key: Some(key),
..
})
.await?;
```
Python continues to own configuration, retries, routing policy, logging,
callbacks, spend tracking, and customer plugins until each Rust path has parity
coverage and production evidence.
## Crates
| Crate | Role |
|-------|------|
| litellm-core | The SDK. Per-route entrypoints (`messages::messages()`), types, provider transforms (modules under `providers/`), provider resolution, auth, the provider HTTP call, and the router. |
| litellm-config | Config-loading boundary. Returns resolved deployments and optionally delegates loading to Python. |
| litellm-ai-gateway | The axum server (behind the `server` feature) and WebSocket hosts. Translates HTTP/WS to core entrypoints; no provider handlers. |
| litellm-python-interop | Domain-neutral PyO3 foundation for GIL handling and typed Python/Serde conversion. |
| litellm-python-bridge | PyO3 cdylib exposing LiteLLM Rust APIs to the Python SDK. Owns API registration, domain wiring, and Python exception mapping. |
Dependency direction is acyclic: config depends on core, the gateway depends on config and core, and the Python bridge depends on the domain layers and Python interop.
## Layout
```text
crates/
core/ The SDK: route modules + provider transforms.
src/messages/ mod.rs (entrypoint), types, transformation, prepare, handler, client
src/providers/anthropic/messages/transformation.rs
config/ Config loading and resolved deployments.
ai-gateway/ Axum server + WebSocket hosts; calls core entrypoints.
python-interop/ Domain-neutral PyO3 conversion and GIL primitives.
python-bridge/ PyO3 API adapter for Python LiteLLM.
```
The folder shape follows the Python provider tree:
`core/src/providers/<provider>/<route>/transformation.rs`. The bridge exposes one
function per top-level route, mirroring the core entrypoints.
## Checks
Run the commands under "Checks" in [CLAUDE.md](CLAUDE.md) before pushing Rust
changes. That list is the single source of truth and matches what GitHub Actions
runs for changes under `litellm-rust/`.

View file

@ -1,53 +0,0 @@
# Provider coding standards (litellm-rust)
Rules for adding or changing an LLM provider/route in `litellm-rust`. `messages` (`core/src/messages`, `ANTHROPIC_MESSAGES_CONFIG`) is the reference: a route is a `core` module with a public entrypoint that makes the call and returns a typed response.
## Provider resolution
1. Always resolve the provider/model first with `get_custom_llm_provider` (`core/src/routing_utils/provider.rs`). Nothing downstream may branch on a raw model string.
2. Model/provider is resolved once, in `prepare.rs`, and passed down as typed fields. Don't re-resolve or re-parse it in transforms or handlers.
## Transforms and the base config
3. Every route defines a base config trait with `transform_request` + `transform_response` (+ `complete_url`, `supported_params`), living in `core/src/<route>/transformation.rs` (e.g. `AnthropicMessagesProviderConfig`, mirroring `OcrProviderConfig`).
4. Each provider implements that trait as a `const <PROVIDER>_<ROUTE>_CONFIG` in `core/src/providers/<provider>/<route>/transformation.rs`, mirroring the Python provider tree.
5. Individual configs implement only the request/response transforms. Shared behavior (param filtering, defaults) stays as trait default methods so future providers inherit existing logic instead of reimplementing it.
6. Prefer composition: a provider that extends another reuses the base trait's defaults or wraps another config; don't copy transform bodies between providers.
## Boundaries
7. Layers never cross: `core` = the call itself (entrypoint, types, transforms, provider resolution, auth headers, provider HTTP, lifecycle hooks); `ai-gateway` = serving HTTP/WS (routing, extractors, auth of *our* callers, streaming to the client); `python-bridge` = thin PyO3 adapter. Hosts call the core entrypoint; they never build a provider request.
8. Generic/route files contain zero provider-specific branches. A provider is one module under `core/src/providers/<provider>/<route>/`; a route is a module, never a new crate.
9. Route entry point stays thin: `core::<route>::<route>()` -> `prepare_*` -> handler (or `CallLifecycle::run_request`, which owns the pre_call -> during_call -> provider call -> success/failure order and phase timing). Axum handlers validate and delegate to a service that calls the entrypoint; no business logic in them.
10. Constants (URLs, env-var names, API versions, error messages) live in a crate `constants.rs`, never inline. Config-shaped env reads happen at the host/config layer with the `DEFAULT_*` fallback defined in `constants.rs`; the only env read in `core` is the credential fallback in a route's `prepare.rs`.
## Types and errors
11. Typed contracts only: no bare `serde_json::Value` / `String` / `Vec<String>` as a transform input or output. Parse wire bytes into typed structs/enums at the host edge; a `type` discriminator is a typed field, not a raw string.
12. Model failures as values: return typed `CoreError`, don't panic. No `unwrap`/`expect`/`panic!` on user or provider input.
13. No mutation: build values in one shot (comprehensions/iterators, `collect`), prefer immutable bindings and owned typed structs over seeding-and-mutating.
14. Early returns over deep nesting; small focused files over god modules.
15. Preserve Python output shape intentionally. If a field is always serialized as `null` for parity, keep it and pin it with a test.
## Safety and data minimization
16. Never log request/response bodies, base64 payloads, document contents, or secrets. Truncate and bound any upstream body before it crosses a host boundary.
17. Treat empty/whitespace credentials, URLs, and config values as absent at the host resolution layer.
18. Network I/O sets connect + request timeouts (no unbounded waits), reuses a shared HTTP client, and prefers rustls TLS.
## Tests and rollout
19. Every provider transform ships tests for: supported-param filtering, request body shape, response normalization, missing/null fields, bad input, and `*_match_python` fixture parity.
20. Lifecycle/hook tests cover hook order, success + failure callback payloads, pre-call guardrail blocking before any provider I/O, during-call body mutation, and provider-error mapping.
21. When a route has a Python reference implementation, the Rust path stays off by default and behind Python parity tests (disabled / enabled-equals-Python / bridge-unavailable fallback) until parity is proven. A new provider/route may instead be implemented rust-only with no Python reference; then the Python interface is a thin dispatch to Rust with no fallback, and tests cover the rust-backed path plus the unavailable-bridge error. State the rust-only choice explicitly in the PR.
## Python bridge (SDK side)
22. A Python -> Rust bridge keeps the Python side minimal: the Python interface only marshals inputs and calls the Rust interface, with no transform, handler, or business logic. Aim for well under 100 lines of interface code per route; if the Python grows past that, the logic belongs in Rust.
23. Do not bloat `litellm/main.py`. A route's provider dispatch lives in a thin dispatch class under `litellm/llms/<provider>/<route>/` that calls the Rust bridge; `main.py` only instantiates it and calls its sync/async method.
24. Do not add new feature flags unless explicitly requested. Reuse the existing LiteLLM Rust rollout mechanism (`litellm.rust`); never introduce a per-route env flag such as `LITELLM_USE_RUST_<ROUTE>`.
## Checks before push
25. Run, and keep green, the commands under "Checks" in `litellm-rust/CLAUDE.md`.
That list is the single source of truth and matches what GitHub Actions runs.

View file

@ -1,54 +0,0 @@
# ai-gateway — folder architecture
The Axum server that fronts the Rust gateway. It owns transport + config + auth
only; deployment selection lives in `core::router`, and the LLM call itself
(transforms, auth headers, provider HTTP) lives behind a `core` route entrypoint
such as `litellm_core::messages::messages`. No provider handler lives here.
```
src/
main.rs # entrypoint: build AppState (router + master key), bind, serve
state.rs # AppState — shared Arc<Router> + master_key
auth/ # authentication as an axum extractor — added to handler args
mod.rs # RequireMasterKey: FromRequestParts, single master key (LITELLM_MASTER_KEY)
routes/ # one module per route, all matching the same template
AGENTS.md # ← the route template (read this before adding a route)
mod.rs # app(): merges every module's router()
health.rs # simple route (one file): router() + liveness/readiness
realtime/ # route with logic → axum surface + a no-axum service:
mod.rs # router() + handler + WS<->events adapter (the axum surface)
service.rs # business logic (select deployment, call provider) — no axum, testable
```
## Rules
- **Routes follow one template.** Each route module exposes
`pub fn router() -> Router<AppState>`; `routes/mod.rs` only merges them. Simple
routes are one file; non-trivial routes are a folder (`handler`/`service`/
`transport`). See `routes/AGENTS.md`.
- **Auth is an extractor.** Add `crate::auth::RequireMasterKey` to a handler's
args; it runs during extraction. Never re-implement the check per route.
- **Handlers are thin.** A handler validates and delegates to its `service`. No
business logic, no provider calls, no transforms in handlers.
- **Services call `core`, they don't reimplement it.** A `service` picks the
deployment and calls the `core` route entrypoint. Provider resolution, auth
headers, URL building, and the HTTP call are `core`'s job; a service that
builds a provider request itself is a bug (`routes/messages/service.rs` is
the reference).
- **State is shared and cheap to clone.** Long-lived handles live behind `Arc` in
`state.rs`; read env/config only in `main.rs` when building state.
## Auth (interim)
A single **master key** (`LITELLM_MASTER_KEY`), enforced by the
`auth::RequireMasterKey` extractor: any caller presenting it as
`Authorization: Bearer <key>` may invoke the gateway. Fails closed (500) when
unset; constant-time compare. The server binds `127.0.0.1` by default (`HOST` to
override). Full per-key auth + budgets/rate-limits are delegated to the Python
proxy in a later phase. Health routes don't add the extractor (unauthenticated).
## Python interop
Python-backed loading lives in `litellm-config` and is **load-time only**. The
gateway's `python-config` feature forwards to that crate. The realtime data path
never takes the GIL.

View file

@ -1,14 +0,0 @@
# ai-gateway architecture
The Rust ai-gateway does LLM inference (realtime WebSocket). Spend tracking is an
API callback: it POSTs each finished session to the LiteLLM proxy, which records
spend and runs the usual callbacks.
```mermaid
flowchart LR
C[client] <--> G[Rust ai-gateway<br/>LLM inference]
G <--> O[OpenAI realtime]
G -. spend tracking callback .-> P[litellm proxy]
F[litellm-config<br/>load-time only] --> G
F -. Python backend .-> P
```

View file

@ -6,10 +6,6 @@ license.workspace = true
repository.workspace = true
autotests = false
[[test]]
name = "workspace_crate_allowlist"
path = "tests/workspace_crate_allowlist.rs"
[dependencies]
base64.workspace = true
azure_core.workspace = true
@ -17,6 +13,7 @@ azure_identity.workspace = true
data-url = "0.3.2"
gcp_auth.workspace = true
moka.workspace = true
mime_guess = "2.0.5"
rand.workspace = true
reqwest.workspace = true
serde.workspace = true

View file

@ -48,7 +48,7 @@ pub(crate) const MEDIA_CONNECT_TIMEOUT_SECS: u64 = 10;
pub(crate) const OCR_HTTP_TIMEOUT_SECS: u64 = 600;
pub(crate) const OCR_CONNECT_TIMEOUT_SECS: u64 = 10;
pub(crate) const OCR_INLINE_MAX_BYTES: usize = 50 * 1024 * 1024;
pub const OCR_INLINE_MAX_BYTES: usize = 50 * 1024 * 1024;
pub(crate) const OCR_DOWNLOAD_MAX_BYTES: u64 = 50 * 1024 * 1024;
pub(crate) const OCR_MAX_FETCH_REDIRECTS: usize = 10;
pub(crate) const OCR_POLL_TIMEOUT_SECS: u64 = 120;

View file

@ -2,11 +2,11 @@ use base64::{Engine, engine::general_purpose::STANDARD};
use data_url::mime::Mime;
use data_url::{DataUrl, DataUrlError, forgiving_base64::DecodeError};
use reqwest::Url;
use serde_json::{Map, Value};
use serde_json::Map;
use super::error::{OcrError, OcrRequestError, OcrResponseError};
use super::types::{OcrConnection, OcrDocument};
use crate::constants::OCR_MAX_FETCH_REDIRECTS;
use crate::constants::{OCR_INLINE_MAX_BYTES, OCR_MAX_FETCH_REDIRECTS};
use crate::error::{MediaError, TransportError};
use crate::media::{DownloadPolicy, MediaFetcher};
@ -14,32 +14,34 @@ pub fn encode_file_document(
bytes: &[u8],
file_name: Option<&str>,
mime_type: Option<&str>,
) -> Result<Value, OcrRequestError> {
) -> Result<OcrDocument, OcrRequestError> {
if bytes.is_empty() {
return Err(OcrRequestError::RequestField {
path: "document.file".into(),
});
return Err(OcrRequestError::EmptyFile);
}
let mime_type = mime_type.map(str::trim);
if mime_type.is_some_and(|value| !valid_mime_type(value)) {
return Err(OcrRequestError::RequestField {
path: "document.mime_type".into(),
});
if bytes.len() > OCR_INLINE_MAX_BYTES {
return Err(OcrRequestError::InlineDocumentTooLarge);
}
if let Some(value) = mime_type
&& !valid_mime_type(value)
{
return Err(OcrRequestError::InvalidMimeType(value.into()));
}
let mime_type = mime_type
.map(str::to_string)
.or_else(|| file_name.and_then(mime_type_for_name).map(str::to_string))
.or_else(|| file_name.map(|name| mime_type_for_name(name).to_string()))
.unwrap_or_else(|| "application/octet-stream".into());
let source = format!("data:{mime_type};base64,{}", STANDARD.encode(bytes));
let (kind, field) = if mime_type.starts_with("image/") {
("image_url", "image_url")
Ok(if mime_type.starts_with("image/") {
OcrDocument::ImageUrl {
image_url: source,
extra_fields: Map::new(),
}
} else {
("document_url", "document_url")
};
Ok(Value::Object(Map::from_iter([
("type".into(), Value::String(kind.into())),
(field.into(), Value::String(source)),
])))
OcrDocument::DocumentUrl {
document_url: source,
extra_fields: Map::new(),
}
})
}
fn valid_mime_type(value: &str) -> bool {
@ -48,22 +50,39 @@ fn valid_mime_type(value: &str) -> bool {
};
!kind.is_empty()
&& !subtype.is_empty()
&& value.bytes().all(|byte| {
byte.is_ascii_alphanumeric() || matches!(byte, b'/' | b'.' | b'+' | b'-' | b'_')
&& kind.chars().chain(subtype.chars()).all(|character| {
character.is_alphanumeric() || matches!(character, '.' | '+' | '-' | '_')
})
}
fn mime_type_for_name(name: &str) -> Option<&'static str> {
let extension = name.rsplit_once('.')?.1;
pub fn mime_type_for_name(name: &str) -> &'static str {
let extension = std::path::Path::new(name)
.extension()
.and_then(|value| value.to_str())
.unwrap_or_default();
match extension.to_ascii_lowercase().as_str() {
"pdf" => Some("application/pdf"),
"png" => Some("image/png"),
"jpg" | "jpeg" => Some("image/jpeg"),
"gif" => Some("image/gif"),
"webp" => Some("image/webp"),
"tiff" | "tif" => Some("image/tiff"),
"bmp" => Some("image/bmp"),
_ => None,
"pdf" => "application/pdf",
"png" => "image/png",
"jpg" | "jpeg" => "image/jpeg",
"gif" => "image/gif",
"webp" => "image/webp",
"tiff" | "tif" => "image/tiff",
"bmp" => "image/bmp",
_ => mime_guess::from_path(name)
.first_raw()
.unwrap_or("application/octet-stream"),
}
}
pub fn upload_mime_type<'a>(file_name: Option<&str>, content_type: Option<&'a str>) -> &'a str {
match content_type
.and_then(|value| value.split(';').next())
.map(str::trim)
{
Some(value) if !value.is_empty() && value != "application/octet-stream" => value,
_ => file_name
.map(mime_type_for_name)
.unwrap_or("application/octet-stream"),
}
}
@ -176,24 +195,43 @@ mod tests {
fn file_bytes_are_encoded_with_core_owned_mime_policy() {
assert_eq!(
encode_file_document(b"abc", Some("scan.png"), None).unwrap(),
serde_json::json!({
"type": "image_url",
"image_url": "data:image/png;base64,YWJj"
})
OcrDocument::ImageUrl {
image_url: "data:image/png;base64,YWJj".into(),
extra_fields: Map::new(),
}
);
assert_eq!(
encode_file_document(b"abc", None, Some("application/pdf")).unwrap(),
serde_json::json!({
"type": "document_url",
"document_url": "data:application/pdf;base64,YWJj"
})
document("data:application/pdf;base64,YWJj")
);
}
#[test]
fn file_encoding_enforces_decoded_size_limit() {
let bytes = vec![b'a'; OCR_INLINE_MAX_BYTES + 1];
assert_eq!(
encode_file_document(&bytes, None, None),
Err(OcrRequestError::InlineDocumentTooLarge)
);
let document = encode_file_document(&bytes[..OCR_INLINE_MAX_BYTES], None, None).unwrap();
let inline = InlineDocument::parse(document.source()).unwrap().unwrap();
assert_eq!(
inline.decode(OCR_INLINE_MAX_BYTES).unwrap(),
bytes[..OCR_INLINE_MAX_BYTES]
);
}
#[test]
fn file_encoding_rejects_empty_bytes_and_invalid_explicit_mime() {
assert!(encode_file_document(b"", None, None).is_err());
assert!(encode_file_document(b"abc", None, Some("text/plain;bad")).is_err());
for mime in [
"text/plain;bad",
"text/plain/extra",
" text/plain",
"text/plain\n",
] {
assert!(encode_file_document(b"abc", None, Some(mime)).is_err());
}
}
#[test]

View file

@ -4,6 +4,10 @@ use crate::error::TransportError;
#[derive(Debug, Clone, PartialEq, Eq, Error)]
pub enum OcrRequestError {
#[error("File is empty or could not be read")]
EmptyFile,
#[error("Invalid MIME type: {0}")]
InvalidMimeType(String),
#[error(
"Cohere Parse only accepts `image_url` documents; document_url and PDF inputs are not supported"
)]

View file

@ -12,11 +12,12 @@ pub mod types;
pub mod wire;
pub use client::{OcrClient, ocr};
pub use document::encode_file_document;
pub use document::{encode_file_document, mime_type_for_name, upload_mime_type};
pub use lifecycle::{
NativeOutcome, NativeResult, NoopOcrHost, OcrAdmission, OcrCall, OcrCallStep, OcrDecline,
OcrHookHost, OcrHost, OcrHostOperation, OcrHostResult,
};
pub use prepare::{credential_default_fields, credential_index};
pub use types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument};
#[cfg(test)]

View file

@ -6,6 +6,21 @@ use super::error::{OcrError, OcrRequestError};
use super::hooks::OcrDuringCallRequest;
use super::types::{LiteLLMOcrRequest, OcrDocument};
pub fn credential_index(requested: &str, names: &[String]) -> Option<usize> {
names.iter().position(|name| name == requested)
}
pub fn credential_default_fields<'a>(
supplied: &[String],
credential_fields: &'a [String],
) -> Vec<&'a str> {
credential_fields
.iter()
.filter(|name| !supplied.contains(name))
.map(String::as_str)
.collect()
}
#[derive(Debug, Deserialize)]
pub(crate) struct ParsedProviderParams<T> {
#[serde(flatten)]

View file

@ -1,115 +0,0 @@
//! Enforcement: the litellm-rust workspace has exactly six crates.
//!
//! `core` (the Rust SDK), `token-counter` (standalone input token counting),
//! `config` (the config-loading boundary),
//! `ai-gateway` (the HTTP/WebSocket host),
//! `python-interop` (domain-neutral PyO3 primitives), and `python-bridge` (the
//! PyO3 cdylib). Adding or removing a crate must be a
//! deliberate act: this test fails until the allowlist here is updated, forcing
//! whoever changes the crate set to justify the new crate per the rule that a
//! crate is a layer needing independent compilation / its own deps / a separate
//! artifact — and to keep `litellm-rust/AGENTS.md` in sync.
//!
//! Std-only (no toml crate): we scan the workspace manifest's `members = [...]`
//! block and the `crates/` directory directly.
use std::collections::BTreeSet;
use std::fs;
use std::path::{Path, PathBuf};
/// The one true crate set. Update BOTH this and `litellm-rust/AGENTS.md` when the
/// workspace legitimately gains or loses a crate.
const EXPECTED_MEMBERS: &[&str] = &[
"crates/core",
"crates/token-counter",
"crates/config",
"crates/ai-gateway",
"crates/python-interop",
"crates/python-bridge",
];
/// The crate subdirectory names that must exist under `crates/`.
const EXPECTED_CRATE_DIRS: &[&str] = &[
"core",
"token-counter",
"config",
"ai-gateway",
"python-interop",
"python-bridge",
];
const MISMATCH: &str = "litellm-rust crate set changed — update this allowlist AND litellm-rust/AGENTS.md, and justify the crate per the rule (crate = layer needing independent compilation / its own deps / a separate artifact).";
/// Absolute path to the workspace root (`litellm-rust/`).
fn workspace_root() -> PathBuf {
// CARGO_MANIFEST_DIR is `.../litellm-rust/crates/core`; the workspace root is
// two levels up.
Path::new(concat!(env!("CARGO_MANIFEST_DIR"), "/../.."))
.canonicalize()
.expect("workspace root should resolve")
}
/// Parse the `members = [ ... ]` array out of the workspace `[workspace]` table.
///
/// Minimal hand-rolled scan: find `members`, then collect every double-quoted
/// string up to the closing `]`. Good enough for our fixed manifest shape and
/// keeps this test dependency-free.
fn parse_members(manifest: &str) -> BTreeSet<String> {
let after_members = manifest
.split_once("members")
.map(|(_, rest)| rest)
.expect("workspace manifest should declare members");
let open = after_members.find('[').expect("members should be an array");
let close = after_members[open..]
.find(']')
.map(|offset| open + offset)
.expect("members array should be closed");
let body = &after_members[open + 1..close];
let mut members = BTreeSet::new();
let mut rest = body;
while let Some(start) = rest.find('"') {
let after_quote = &rest[start + 1..];
let end = after_quote
.find('"')
.expect("opening quote should be matched");
members.insert(after_quote[..end].to_string());
rest = &after_quote[end + 1..];
}
members
}
/// The crate subdirectory names under `crates/`.
///
/// A directory counts as a crate only when it holds a `Cargo.toml`; non-crate
/// directories (e.g. docs like `CODING_STANDARDS/`) are ignored so they can live
/// under `crates/` without tripping the crate-proliferation guard.
fn crate_dirs(root: &Path) -> BTreeSet<String> {
fs::read_dir(root.join("crates"))
.expect("crates/ directory should exist")
.filter_map(Result::ok)
.filter(|entry| entry.file_type().map(|ty| ty.is_dir()).unwrap_or(false))
.filter(|entry| entry.path().join("Cargo.toml").is_file())
.map(|entry| entry.file_name().to_string_lossy().into_owned())
.collect()
}
#[test]
fn workspace_members_match_allowlist() {
let root = workspace_root();
let manifest = fs::read_to_string(root.join("Cargo.toml"))
.expect("workspace Cargo.toml should be readable");
let actual = parse_members(&manifest);
let expected: BTreeSet<String> = EXPECTED_MEMBERS.iter().map(|s| s.to_string()).collect();
assert_eq!(actual, expected, "{MISMATCH}");
}
#[test]
fn crates_directory_matches_allowlist() {
let root = workspace_root();
let actual = crate_dirs(&root);
let expected: BTreeSet<String> = EXPECTED_CRATE_DIRS.iter().map(|s| s.to_string()).collect();
assert_eq!(actual, expected, "{MISMATCH}");
}

View file

@ -14,6 +14,8 @@ use tokio::sync::Mutex;
use crate::errors::ocr_error_to_pyerr;
use crate::execution::{run_async_value, run_sync_value};
mod preparation;
pub(crate) trait PythonRoute: Send + Sync {
fn state(&self) -> &PythonCallState;
fn state_mut(&mut self) -> &mut PythonCallState;
@ -532,12 +534,7 @@ impl PythonCallState {
}
pub fn prepare(&mut self, py: Python<'_>) -> PyResult<()> {
self.kwargs = py
.import("litellm.rust_bridge.lifecycle")?
.getattr("prepare")?
.call1((&self.kwargs, self.logger(py)?))?
.cast_into::<PyDict>()?
.unbind();
self.kwargs = preparation::prepare(py, self.kwargs.bind(py), &self.logger(py)?)?.unbind();
Ok(())
}

View file

@ -0,0 +1,59 @@
use litellm_core::ocr::{credential_default_fields, credential_index};
use pyo3::prelude::*;
use pyo3::types::{PyDict, PyList};
pub(super) fn prepare<'py>(
py: Python<'py>,
kwargs: &Bound<'py, PyDict>,
logger: &Bound<'py, PyAny>,
) -> PyResult<Bound<'py, PyDict>> {
let arguments = kwargs.copy()?;
arguments.set_item("litellm_logging_obj", logger)?;
let litellm = py.import("litellm")?;
inherit_credentials(py, &litellm, &arguments)?;
py.import("litellm.rust_bridge.lifecycle")?
.getattr("check_limits")?
.call1((&arguments,))?;
Ok(arguments)
}
fn inherit_credentials(
py: Python<'_>,
litellm: &Bound<'_, PyModule>,
arguments: &Bound<'_, PyDict>,
) -> PyResult<()> {
let Some(requested) = arguments
.get_item("litellm_credential_name")?
.filter(|value| !value.is_none())
else {
return Ok(());
};
if !requested.is_truthy()? {
return Ok(());
}
let requested: String = requested.extract()?;
let credentials = litellm.getattr("credential_list")?.cast_into::<PyList>()?;
let names = credentials
.iter()
.map(|credential| credential.getattr("credential_name")?.extract::<String>())
.collect::<PyResult<Vec<_>>>()?;
let Some(index) = credential_index(&requested, &names) else {
py.import("litellm._logging")?.getattr("verbose_logger")?.call_method1(
"warning",
("litellm_credential_name=%s matched none of the %d loaded credentials; the request runs without it", requested, names.len()),
)?;
return Ok(());
};
let values = credentials
.get_item(index)?
.getattr("credential_values")?
.cast_into::<PyDict>()?;
let supplied: Vec<String> = arguments.keys().extract()?;
let fields: Vec<String> = values.keys().extract()?;
for name in credential_default_fields(&supplied, &fields) {
if let Some(value) = values.get_item(name)? {
arguments.set_item(name, value)?;
}
}
Ok(())
}

View file

@ -10,10 +10,12 @@ mod audio_transcription;
mod chat_completions;
mod messages;
mod ocr;
mod ocr_document;
mod ocr_lifecycle;
pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
ocr::register(module)?;
ocr_document::register(module)?;
ocr_lifecycle::register(module)?;
audio_transcription::register(module)?;
messages::register(module)?;

View file

@ -0,0 +1,148 @@
use std::io::Read;
use std::path::PathBuf;
use pyo3::exceptions::{PyFileNotFoundError, PyTypeError, PyValueError};
use pyo3::prelude::*;
use pyo3::pybacked::PyBackedBytes;
use pyo3::types::{PyBytes, PyDict, PyString};
use litellm_core::constants::OCR_INLINE_MAX_BYTES;
use litellm_core::ocr::{OcrDocument, encode_file_document, mime_type_for_name, upload_mime_type};
use litellm_python_interop::to_py_preserving_errors;
enum FileBytes {
Python(PyBackedBytes),
Native(Vec<u8>),
}
impl AsRef<[u8]> for FileBytes {
fn as_ref(&self) -> &[u8] {
match self {
Self::Python(bytes) => bytes,
Self::Native(bytes) => bytes,
}
}
}
fn read_file_input(
py: Python<'_>,
file: &Bound<'_, PyAny>,
) -> PyResult<(FileBytes, Option<String>)> {
if file.is_instance_of::<PyString>() {
return Err(PyValueError::new_err(
"OCR file input does not accept bare str values. Pass bytes, a pathlib.Path, or a file-like object.",
));
}
if file.is_instance(&py.import("os")?.getattr("PathLike")?)? {
let path: PathBuf = file.extract()?;
let name = path
.file_name()
.map(|value| value.to_string_lossy().into_owned());
let bytes = py
.detach(|| {
let mut bytes = Vec::new();
std::fs::File::open(&path)?
.take(OCR_INLINE_MAX_BYTES as u64 + 1)
.read_to_end(&mut bytes)?;
Ok::<_, std::io::Error>(bytes)
})
.map_err(|error| {
if error.kind() == std::io::ErrorKind::NotFound {
PyFileNotFoundError::new_err(format!("File not found: {}", path.display()))
} else {
error.into()
}
})?;
return Ok((FileBytes::Native(bytes), name));
}
if file.is_instance_of::<PyBytes>() {
return Ok((FileBytes::Python(file.extract()?), None));
}
let reader = file
.getattr_opt("read")?
.filter(|value| value.is_callable());
let Some(reader) = reader else {
return Err(PyValueError::new_err(format!(
"Unsupported file input type: {}. Expected pathlib.Path, bytes, or a file-like object.",
file.get_type(),
)));
};
let name = file
.getattr_opt("name")?
.filter(|value| !value.is_none())
.map(|value| value.extract::<String>())
.transpose()?;
let value = reader.call0()?;
let bytes = if value.is_instance_of::<PyString>() {
FileBytes::Native(value.extract::<String>()?.into_bytes())
} else if value.is_instance_of::<PyBytes>() {
FileBytes::Python(value.extract()?)
} else {
return Err(PyTypeError::new_err(format!(
"OCR file read must return bytes or str, got {}",
value.get_type(),
)));
};
Ok((bytes, name))
}
pub(super) fn file_document(py: Python<'_>, document: &Bound<'_, PyAny>) -> PyResult<OcrDocument> {
let file = document.get_item("file").map_err(|error| {
if error.is_instance_of::<pyo3::exceptions::PyKeyError>(py) {
PyValueError::new_err("document with type='file' must include a 'file' field containing a pathlib.Path, file-like object, or bytes")
} else {
error
}
})?;
if file.is_none() {
return Err(PyValueError::new_err(
"document with type='file' must include a 'file' field containing a pathlib.Path, file-like object, or bytes",
));
}
let (bytes, name) = read_file_input(py, &file)?;
let mime = document
.cast::<PyDict>()?
.get_item("mime_type")?
.map(|value| value.extract::<String>())
.transpose()?;
py.detach(|| encode_file_document(bytes.as_ref(), name.as_deref(), mime.as_deref()))
.map_err(|error| PyValueError::new_err(error.to_string()))
}
#[pyfunction]
fn _ocr_file_document(py: Python<'_>, document: Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
to_py_preserving_errors(py, &file_document(py, &document)?)
}
#[pyfunction]
fn _ocr_mime_type(file_name: &str) -> String {
mime_type_for_name(file_name).into()
}
#[pyfunction]
#[pyo3(signature = (file_content, file_name=None, content_type=None))]
fn _ocr_upload_document(
py: Python<'_>,
file_content: &Bound<'_, PyBytes>,
file_name: Option<&str>,
content_type: Option<&str>,
) -> PyResult<Py<PyAny>> {
let bytes: PyBackedBytes = file_content.extract()?;
let document = py
.detach(|| {
encode_file_document(
&bytes,
None,
Some(upload_mime_type(file_name, content_type)),
)
})
.map_err(|error| PyValueError::new_err(error.to_string()))?;
to_py_preserving_errors(py, &document)
}
pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
module.add("_OCR_MAX_FILE_BYTES", OCR_INLINE_MAX_BYTES)?;
module.add_function(wrap_pyfunction!(_ocr_upload_document, module)?)?;
module.add_function(wrap_pyfunction!(_ocr_file_document, module)?)?;
module.add_function(wrap_pyfunction!(_ocr_mime_type, module)?)
}

View file

@ -2,7 +2,7 @@ use serde_json::{Map, Value};
use std::sync::Arc;
use pyo3::prelude::*;
use pyo3::types::{PyBytes, PyDict, PyTuple};
use pyo3::types::{PyDict, PyTuple};
use litellm_core::auth::{ResolvedCredential, SecretValue};
use litellm_core::ocr::hooks::{OcrDuringCallRequest, OcrPostCallRequest, OcrPreCallRequest};
@ -363,22 +363,8 @@ fn extract_document(py: Python<'_>, document: &Bound<'_, PyAny>) -> PyResult<Val
if document.get_item("type")?.extract::<String>()? != "file" {
return from_py(document);
}
let file = document.get_item("file")?;
let (bytes, name): (Py<PyBytes>, Option<String>) = py
.import("litellm.rust_bridge.ocr_lifecycle")?
.getattr("read_file_input")?
.call1((file,))?
.extract()?;
let mime_type = document
.get_item("mime_type")
.ok()
.and_then(|value| value.extract::<String>().ok());
litellm_core::ocr::encode_file_document(
bytes.bind(py).as_bytes(),
name.as_deref(),
mime_type.as_deref(),
)
.map_err(|error| ocr_error_to_pyerr(error.into()))
serde_json::to_value(super::ocr_document::file_document(py, document)?)
.map_err(|error| pyo3::exceptions::PyValueError::new_err(error.to_string()))
}
fn retained_document(

View file

@ -1,36 +1,10 @@
import base64
import mimetypes
import os
import re
from io import IOBase
from typing import Final, Literal, Protocol
from collections.abc import Mapping
from os import PathLike
from typing import Final, Literal, Protocol, cast # noqa: TID251 # native callables are validated when loaded
from typing_extensions import NotRequired, ReadOnly, TypedDict
from litellm._logging import verbose_logger
_MIME_PATTERN: Final = re.compile(r"^[\w.+-]+/[\w.+-]+$")
_MIME_TYPE_MAP: Final = {
".pdf": "application/pdf",
".png": "image/png",
".jpg": "image/jpeg",
".jpeg": "image/jpeg",
".gif": "image/gif",
".webp": "image/webp",
".tiff": "image/tiff",
".tif": "image/tiff",
".bmp": "image/bmp",
}
def get_mime_type(file_path: str) -> str:
ext: Final = os.path.splitext(file_path)[1].lower()
mime: Final = _MIME_TYPE_MAP.get(ext)
if mime:
return mime
guessed, _ = mimetypes.guess_type(file_path)
return guessed or "application/octet-stream"
from litellm.rust_bridge.bindings import NativeBinding
class FileReader(Protocol):
@ -39,76 +13,61 @@ class FileReader(Protocol):
class FileDocument(TypedDict):
type: ReadOnly[Literal["file"]]
file: ReadOnly[bytes | os.PathLike[str] | FileReader]
file: ReadOnly[bytes | PathLike[str] | FileReader]
mime_type: ReadOnly[NotRequired[str]]
class NativeFileDocument(Protocol):
def __call__(self, document: Mapping[str, object]) -> dict[str, str]: ...
class NativeUploadDocument(Protocol):
def __call__(self, file_content: bytes, file_name: str | None, content_type: str | None) -> dict[str, str]: ...
class NativeMimeType(Protocol):
def __call__(self, file_name: str) -> str: ...
_FILE_DOCUMENT: Final = NativeBinding(
"_ocr_file_document", validate=lambda value: cast(NativeFileDocument, value) if callable(value) else None
)
_UPLOAD_DOCUMENT: Final = NativeBinding(
"_ocr_upload_document", validate=lambda value: cast(NativeUploadDocument, value) if callable(value) else None
)
_MAX_FILE_BYTES: Final = NativeBinding(
"_OCR_MAX_FILE_BYTES", validate=lambda value: value if isinstance(value, int) and value > 0 else None
)
_MIME_TYPE: Final = NativeBinding(
"_ocr_mime_type", validate=lambda value: cast(NativeMimeType, value) if callable(value) else None
)
def get_mime_type(file_path: str) -> str:
native: Final = _MIME_TYPE.load()
if native is None:
raise RuntimeError("Rust OCR document preparation is unavailable")
return native(file_path)
def get_max_file_bytes() -> int:
limit: Final = _MAX_FILE_BYTES.load()
if limit is None:
raise RuntimeError("Rust OCR document preparation is unavailable")
return limit
def convert_file_document_to_url_document(document: FileDocument) -> dict[str, str]:
file_input: Final = document.get("file")
if file_input is None:
raise ValueError(
"document with type='file' must include a 'file' field containing "
"a pathlib.Path, file-like object, or bytes"
)
native: Final = _FILE_DOCUMENT.load()
if native is None:
raise RuntimeError("Rust OCR document preparation is unavailable")
return native(document)
file_bytes: bytes
mime_type: str = "application/octet-stream"
file_name: str | None = None
if isinstance(file_input, str):
raise ValueError(
"OCR file input does not accept bare str values. Pass bytes, "
"a pathlib.Path, or a file-like object. To OCR a local file "
"from a path, call open(path, 'rb') yourself."
)
if isinstance(file_input, os.PathLike):
file_path: Final = str(file_input)
if not os.path.isfile(file_path):
raise FileNotFoundError(f"File not found: {file_path}")
mime_type = get_mime_type(file_path)
file_name = os.path.basename(file_path)
with open(file_path, "rb") as file:
file_bytes = file.read()
elif isinstance(file_input, bytes):
file_bytes = file_input
elif isinstance(file_input, IOBase) or hasattr(file_input, "read"):
if hasattr(file_input, "name"):
file_name = getattr(file_input, "name", None)
if file_name:
mime_type = get_mime_type(file_name)
file_bytes = file_input.read()
if isinstance(file_bytes, str):
file_bytes = file_bytes.encode("utf-8")
else:
raise ValueError(
f"Unsupported file input type: {type(file_input)}. Expected pathlib.Path, bytes, or a file-like object."
)
if not file_bytes:
raise ValueError("File is empty or could not be read")
if "mime_type" in document:
mime_type = document["mime_type"]
if not _MIME_PATTERN.match(mime_type):
raise ValueError(f"Invalid MIME type: {mime_type}")
base64_data: Final = base64.b64encode(file_bytes).decode("utf-8")
data_uri: Final = f"data:{mime_type};base64,{base64_data}"
if mime_type.startswith("image/"):
verbose_logger.debug(
"OCR file input: Converted file to image_url data URI (mime=%s, size=%s bytes, name=%s)",
mime_type,
len(file_bytes),
file_name,
)
return {"type": "image_url", "image_url": data_uri}
verbose_logger.debug(
"OCR file input: Converted file to document_url data URI (mime=%s, size=%s bytes, name=%s)",
mime_type,
len(file_bytes),
file_name,
)
return {"type": "document_url", "document_url": data_uri}
def convert_upload_to_url_document(
file_content: bytes, filename: str | None, content_type: str | None
) -> dict[str, str]:
native: Final = _UPLOAD_DOCUMENT.load()
if native is None:
raise RuntimeError("Rust OCR document preparation is unavailable")
return native(file_content, filename, content_type)

View file

@ -15,7 +15,7 @@ from litellm.llms.base_llm.ocr.transformation import (
OCRResponse,
parse_ocr_request_format,
)
from litellm.ocr.main import convert_file_document_to_url_document, get_mime_type
from litellm.ocr.input import convert_upload_to_url_document, get_max_file_bytes
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
@ -28,24 +28,7 @@ def _build_document_from_upload(
filename: str | None,
content_type: str | None,
) -> dict[str, str]:
"""
Convert uploaded file bytes into a Mistral-format document dict with base64 data URI.
Delegates to convert_file_document_to_url_document after resolving MIME type
from the upload's content_type header or filename.
"""
mime_type = content_type.split(";")[0].strip() if content_type else None
if not mime_type or mime_type == "application/octet-stream":
if filename:
mime_type = get_mime_type(filename)
return convert_file_document_to_url_document(
{
"type": "file",
"file": file_content,
"mime_type": mime_type or "application/octet-stream",
}
)
return convert_upload_to_url_document(file_content, filename, content_type)
def _with_request_format(data: Mapping[str, Any], request: Request) -> Mapping[str, Any]:
@ -120,7 +103,7 @@ async def _parse_multipart_form(request: Request) -> dict[str, Any]:
# Seek to start in case the file was already partially read by middleware
await uploaded_file.seek(0)
file_content: Final = await uploaded_file.read()
file_content: Final = await uploaded_file.read(get_max_file_bytes() + 1)
if not file_content:
raise ValueError("Uploaded file is empty")

View file

@ -52,10 +52,6 @@ async def drive(execution: Execution) -> object:
execution.close()
class CredentialLoader(Protocol):
def __call__(self, kwargs: dict[str, object]) -> None: ...
class MetadataUpdater(Protocol):
def __call__(
self,
@ -97,22 +93,13 @@ def setup(
return CallSetup(logger, prepared)
def prepare(kwargs: Mapping[str, object], logger: Logging) -> dict[str, object]:
def check_limits(kwargs: Mapping[str, object]) -> None:
import litellm
from litellm import utils
arguments: Final = { # mutable-ok: credential loader updates an owned kwargs dict
**kwargs,
"litellm_logging_obj": logger,
}
load_credentials: Final = cast( # cast-ok: legacy credential loader mutates a concrete kwargs dict
CredentialLoader, utils.load_credentials_from_list
)
load_credentials(arguments)
current_cost: Final = litellm._current_cost # pyright: ignore[reportPrivateUsage] # shared SDK budget counter has no public accessor
if litellm.max_budget and current_cost > litellm.max_budget:
raise litellm.BudgetExceededError(current_cost=current_cost, max_budget=litellm.max_budget)
metadata: Final = arguments.get("metadata")
metadata: Final = kwargs.get("metadata")
if isinstance(metadata, Mapping):
typed_metadata: Final = cast( # cast-ok: runtime Mapping check establishes read-only metadata
Mapping[str, object], metadata
@ -125,7 +112,6 @@ def prepare(kwargs: Mapping[str, object], logger: Logging) -> dict[str, object]:
>= litellm.num_retries_per_request
):
raise RuntimeError("Max retries per request hit!")
return arguments
def finalize(

View file

@ -1,8 +1,6 @@
from __future__ import annotations
from collections.abc import Awaitable, Mapping, Sequence
from os import PathLike
from pathlib import Path
from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables
import litellm
@ -43,7 +41,7 @@ NATIVE_OCR_LIFECYCLE: Final = NativeBinding("_ocr_lifecycle", validate=_binding)
def select(request: LiteLLMOcrRequest) -> NativeOcrLifecycle | None:
if litellm.cache is not None or request.kwargs.get("caching") or request.kwargs.get("aocr"):
if request.kwargs.get("aocr"):
return None
return NATIVE_OCR_LIFECYCLE.load()
@ -52,35 +50,6 @@ def arguments(request: LiteLLMOcrRequest) -> Mapping[str, object]:
return request.kwargs
class FileReader(Protocol):
def __call__(self) -> object: ...
def read_file_input(file_input: object) -> tuple[bytes, str | None]:
if isinstance(file_input, str):
raise ValueError(
"OCR file input does not accept bare str values. Pass bytes, a pathlib.Path, or a file-like object."
)
if isinstance(file_input, PathLike):
path: Final = Path(
cast(PathLike[str], file_input)
) # cast-ok: Path validates the path protocol at its consumption point
return path.read_bytes(), path.name
if isinstance(file_input, bytes):
return file_input, None
reader: Final[object] = getattr(file_input, "read", None)
if callable(reader):
data: Final = cast(FileReader, reader)() # cast-ok: read is callable and its return is validated below
encoded: Final = data.encode("utf-8") if isinstance(data, str) else data
if not isinstance(encoded, bytes):
raise TypeError(f"OCR file read must return bytes or str, got {type(encoded)}")
name: Final = getattr(file_input, "name", None)
return encoded, name if isinstance(name, str) else None
raise ValueError(
f"Unsupported file input type: {type(file_input)}. Expected pathlib.Path, bytes, or a file-like object."
)
def call_azure_ad_token_provider(provider: object) -> str:
if not callable(provider):
raise TypeError("Azure AD token provider must be callable")

View file

@ -14,6 +14,7 @@ import os
import tempfile
from io import BytesIO
from pathlib import Path
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import orjson
@ -480,3 +481,37 @@ class TestProxySecurityGuard:
"data:application/pdf;base64,"
)
assert result["model"] == "mistral/mistral-ocr-latest"
@pytest.mark.asyncio
async def test_proxy_upload_stops_reading_at_size_limit() -> None:
from starlette.datastructures import UploadFile
from litellm.ocr.input import get_max_file_bytes
from litellm.proxy.ocr_endpoints.endpoints import _parse_multipart_form
limit: Final = get_max_file_bytes()
with tempfile.TemporaryFile() as stream:
stream.truncate(limit * 2)
upload: Final = UploadFile(file=stream, filename="large.pdf")
request: Final = MagicMock(form=AsyncMock(return_value=FormData({"file": upload})))
with pytest.raises(ValueError, match="exceeds the size limit"):
await _parse_multipart_form(request)
assert stream.tell() == limit + 1
@pytest.mark.asyncio
async def test_proxy_upload_filename_is_only_metadata(tmp_path: Path) -> None:
from starlette.datastructures import UploadFile
from litellm.proxy.ocr_endpoints.endpoints import _parse_multipart_form
secret: Final = tmp_path / "secret.pdf"
secret.write_bytes(b"server secret")
upload: Final = UploadFile(file=BytesIO(b"uploaded bytes"), filename=str(secret))
request: Final = MagicMock(form=AsyncMock(return_value=FormData({"file": upload})))
result: Final = await _parse_multipart_form(request)
assert result["document"] == {
"type": "document_url",
"document_url": "data:application/pdf;base64,dXBsb2FkZWQgYnl0ZXM=",
}

View file

@ -1,25 +0,0 @@
# Rust OCR bridge tests
This suite covers OCR requests through LiteLLM's compiled Rust extension. OCR behavior tests live under `ocr/`; reusable OCR request, callback, and recording-server fixtures live under `support/`
A test name identifies the OCR entrypoint or callback under test and its expected observable result. Parameter IDs state the execution mode or credential case. Keep multiple assertions together only when they prove one request, mutation, failure, or callback lifecycle behavior. Record callback observations and assert them after the callback returns because production logging can swallow callback exceptions
`ocr/test_requests.py` covers provider payloads, file preparation, endpoint and credential resolution, normalized responses, errors, timeouts, and Azure token-provider behavior. `ocr/test_callbacks.py` covers callback inputs, mutations, ordering, context, failure handling, concurrency, and cleanup. `ocr/test_guardrails.py` covers post-call blocking and response replacement. `ocr/test_lifecycle.py` checks final object identity, finalization failures, caller-task context, cancellation, nested requests, executor scheduling, deferred release, Reducto upload/parse and Azure submission/poll boundaries. These tests use the public OCR APIs. `ocr/test_dispatch.py` covers enabled native dispatch and rejection when native execution is disabled. `test_ocr.py` exercises the compiled Rust transport directly
Run `make test-rust-extension` as the acceptance command. It builds a fresh wheel, installs that wheel into a temporary environment, requires `LITELLM_RUST=1`, and runs this suite with isolated Python imports
Collection fails when `LITELLM_RUST=1` is set but the compiled `_native` module cannot be imported. The autouse fixture isolates callbacks, both logging executor references, and configuration state. All tests are strict
The native lifecycle supports Mistral, Azure, Vertex and Reducto workflows. Caller-supplied synchronous Azure token providers execute inline when requested by core. Disabled or unavailable native execution and unsupported caching requests raise an error rather than invoking a legacy Python OCR provider
Rust lifecycle ordering lives in `core/src/call_lifecycle/host.rs`, with provider work owned by `core/src/ocr/lifecycle.rs`. The native execution handle and Python reference ownership live in `python-bridge/src/lifecycle.rs`. One ordinary Python coroutine in `litellm/rust_bridge/lifecycle.py` awaits Rust-selected operations in the caller task through `start`, `resume_value`, `resume_error` and idempotent `close`. Tagged Await/Complete steps preserve awaitable final values. The hand-written Rust coroutine protocol has been removed
The extension explicitly requires the GIL and detaches Rust-only synchronous waits. Native results stay in Rust, while retained Python roots and exceptions participate in GC. Request projection and file reads happen after lifecycle setup and applicable deployment hooks. Cancellation during failure logging propagates, while deployment-failure observers preserve the original provider error. Native cancellation waits for the owned provider task through core; synchronous close and GC signal cancellation without claiming to await termination
The native-backed driver probe is `litellm-rust/crates/python-bridge/tests/lifecycle.py`, invoked by Rust unit tests. It covers custom awaitables, task/thread/loop identity, context writes, exception identity, repeated cancellation, re-entry and cycles. Native typing, serialization benchmark additions and token-counter changes are separate follow-ups
## Local validation
The final lifecycle validation run passed `cargo fmt --check`, workspace Clippy with warnings denied, core Clippy with `bedrock-auth`, gateway Clippy with all features, workspace tests, core tests with `bedrock-auth`, and gateway tests with `server`. The installed-wheel acceptance command passed 109 tests on GIL-enabled CPython 3.12.13 with the ABI3 extension. Focused Ruff and basedpyright checks also passed
The installed-wheel run reported two existing Pydantic warnings that ReadOnly TypedDict fields are not runtime mutation guards. Credential-dependent live Bedrock and OpenAI realtime Rust tests remained explicitly ignored. These local results cover controlled provider dependencies and do not establish live-provider or free-threaded Python acceptance

View file

@ -48,3 +48,24 @@ def test_public_ocr_uses_native_route_independently_of_flag(ocr_server: Recordin
assert response.pages[0].markdown == "native OCR response"
assert len(ocr_server.requests) == 1
assert not ocr_server.requests[0].headers.get("user-agent", "").startswith("python-httpx")
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
@pytest.mark.parametrize("caching", [None, False, True])
async def test_ocr_does_not_depend_on_chat_cache(
ocr_server: RecordingServer, monkeypatch: pytest.MonkeyPatch, asynchronous: bool, caching: bool | None
) -> None:
from litellm.caching.caching import Cache
monkeypatch.setattr(litellm, "cache", Cache(type="local", supported_call_types=["completion", "acompletion"]))
arguments: Final = {
"model": OCR_MODEL,
"document": OCR_DOCUMENT,
"api_key": "test-key",
"api_base": ocr_server.base_url,
"caching": caching,
}
response: Final = await litellm.aocr(**arguments) if asynchronous else litellm.ocr(**arguments)
assert response.pages[0].markdown == "native OCR response"
assert len(ocr_server.requests) == 1

View file

@ -772,3 +772,28 @@ async def test_vertex_deepseek_public_lifecycle_normalizes_before_success(ocr_se
ocr_server.requests[0].path
== "/v1/projects/project-1/locations/europe-west4/endpoints/openapi/chat/completions"
)
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
@pytest.mark.parametrize("limit", ["budget", "retries"])
async def test_shared_call_limits_still_reject_before_reading_ocr_file(
ocr_server: RecordingServer, monkeypatch: pytest.MonkeyPatch, asynchronous: bool, limit: str
) -> None:
ocr_server.expected_requests = 0
reads: Final = []
class File:
def read(self):
reads.append("read")
return b"abc"
monkeypatch.setattr(litellm, "max_budget", 1 if limit == "budget" else None)
monkeypatch.setattr(litellm, "_current_cost", 2)
monkeypatch.setattr(litellm, "num_retries_per_request", 1 if limit == "retries" else None)
expected: Final = litellm.BudgetExceededError if limit == "budget" else RuntimeError
arguments: Final = {"document": {"type": "file", "file": File()}, "metadata": {"previous_models": ["earlier"]}}
with pytest.raises(expected, match=r"Budget has been exceeded|Max retries per request hit"):
await call_aocr(ocr_server, **arguments) if asynchronous else call_ocr(ocr_server, **arguments)
assert reads == []
assert ocr_server.requests == []

View file

@ -6,13 +6,13 @@ import pytest
import litellm
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from tests.test_litellm_rust.support.callback_recorder import RecordingLogger
from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec
from tests.test_litellm_rust.support.requests import (
OCR_DOCUMENT,
OCR_RESPONSE,
call_native_aocr,
call_native_ocr,
)
from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec
pytestmark = pytest.mark.requires_rust_extension
@ -456,3 +456,160 @@ async def test_native_azure_ocr_rejects_coroutine_returned_by_sync_token_provide
coroutine.close()
assert calls == []
assert ocr_server.requests == []
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
@pytest.mark.parametrize("explicit_key", [False, True])
async def test_native_ocr_inherits_named_credentials_without_overwriting_arguments(
ocr_server: RecordingServer, monkeypatch: pytest.MonkeyPatch, asynchronous: bool, explicit_key: bool
) -> None:
from litellm.models.credentials import CredentialItem
pages: Final = [0]
opaque: Final = object()
monkeypatch.setattr(
litellm,
"credential_list",
[
CredentialItem(credential_name="other", credential_info={}, credential_values={"api_key": "wrong-key"}),
CredentialItem(
credential_name="ocr-test",
credential_info={},
credential_values={
"api_key": "credential-key",
"api_base": ocr_server.base_url,
"pages": pages,
"opaque": opaque,
},
),
CredentialItem(credential_name="ocr-test", credential_info={}, credential_values={"api_key": "later-key"}),
],
)
class Observer(RecordingLogger):
def log_pre_api_call(self, model, messages, kwargs):
super().log_pre_api_call(model, messages, kwargs)
pages.append(2)
arguments: Final = {
"model": "mistral/mistral-ocr-latest",
"document": OCR_DOCUMENT,
"litellm_credential_name": "ocr-test",
"callbacks": [Observer()],
**({"api_key": "explicit-key"} if explicit_key else {}),
}
response: Final = await litellm.aocr(**arguments) if asynchronous else litellm.ocr(**arguments)
assert response.pages[0].markdown == "native OCR response"
assert (
ocr_server.requests[0].headers["authorization"]
== f"Bearer {'explicit-key' if explicit_key else 'credential-key'}"
)
assert ocr_server.requests[0].body["pages"] == [0, 2]
@pytest.mark.parametrize("source", ["sdk", "proxy"])
@pytest.mark.parametrize(
"filename,mime", [("scan.PNG", "image/png"), ("document.pdf", "application/pdf"), ("note.txt", "text/plain")]
)
def test_ocr_file_helpers_use_native_document_preparation(source: str, filename: str, mime: str) -> None:
from io import BytesIO
from litellm.ocr.input import convert_file_document_to_url_document, get_mime_type
from litellm.proxy.ocr_endpoints.endpoints import _build_document_from_upload
file: Final = BytesIO(b"abc")
file.name = filename
document: Final = (
convert_file_document_to_url_document({"type": "file", "file": file})
if source == "sdk"
else _build_document_from_upload(b"abc", filename, "application/octet-stream; charset=utf-8")
)
field: Final = "image_url" if mime.startswith("image/") else "document_url"
assert get_mime_type(filename) == mime
assert document == {"type": field, field: f"data:{mime};base64,YWJj"}
@pytest.mark.parametrize("attribute", ["read", "name"])
def test_native_file_preparation_preserves_property_errors(attribute: str) -> None:
from litellm.ocr.input import convert_file_document_to_url_document
failure: Final = LookupError("file property failed")
class File:
def __getattribute__(self, name: str):
if name == attribute:
raise failure
return super().__getattribute__(name)
def read(self):
return b"abc"
with pytest.raises(LookupError) as caught:
convert_file_document_to_url_document({"type": "file", "file": File()})
assert caught.value is failure
@pytest.mark.parametrize("kind", ["bytes", "path", "reader"])
def test_native_file_preparation_rejects_oversized_input(kind: str, tmp_path: Path) -> None:
from litellm.ocr.input import FileDocument, convert_file_document_to_url_document, get_max_file_bytes
limit: Final = get_max_file_bytes()
path: Final = tmp_path / "large.pdf"
with path.open("wb") as stream:
stream.truncate(limit + 1)
class Reader:
def read(self) -> bytes:
return b"a" * (limit + 1)
document: Final[FileDocument] = {
"type": "file",
"file": path if kind == "path" else Reader() if kind == "reader" else b"a" * (limit + 1),
}
with pytest.raises(ValueError, match="exceeds the size limit"):
convert_file_document_to_url_document(document)
@pytest.mark.parametrize("kind", ["str", "path", "reader"])
def test_native_upload_binding_rejects_filesystem_inputs(kind: str, tmp_path: Path) -> None:
from io import BytesIO
from typing import cast # noqa: TID251 # deliberately invalid inputs exercise the native runtime boundary
from litellm.ocr.input import convert_upload_to_url_document
path: Final = tmp_path / "secret.pdf"
path.write_bytes(b"server secret")
source: Final = str(path) if kind == "str" else path if kind == "path" else BytesIO(b"abc")
with pytest.raises(TypeError):
convert_upload_to_url_document(cast(bytes, source), "document.pdf", None)
@pytest.mark.parametrize("extra_bytes", [0, 1])
def test_native_upload_enforces_file_size_limit(extra_bytes: int) -> None:
import base64
from litellm.ocr.input import convert_upload_to_url_document, get_max_file_bytes
content: Final = b"a" * (get_max_file_bytes() + extra_bytes)
if extra_bytes:
with pytest.raises(ValueError, match="exceeds the size limit"):
convert_upload_to_url_document(content, "scan.pdf", None)
return
document: Final = convert_upload_to_url_document(content, "scan.pdf", None)
assert document["type"] == "document_url"
assert base64.b64decode(document["document_url"].split(",", 1)[1]) == content
def test_native_file_preparation_preserves_reader_exception() -> None:
from litellm.ocr.input import convert_file_document_to_url_document
failure: Final = RuntimeError("reader failed")
class Reader:
def read(self) -> bytes:
raise failure
with pytest.raises(RuntimeError) as caught:
convert_file_document_to_url_document({"type": "file", "file": Reader()})
assert caught.value is failure