mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
wip
This commit is contained in:
parent
43a19d81ab
commit
c6fc5e185f
29 changed files with 619 additions and 816 deletions
|
|
@ -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).
|
||||
|
|
@ -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.
|
||||
|
|
@ -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.
|
||||
1
litellm-rust/Cargo.lock
generated
1
litellm-rust/Cargo.lock
generated
|
|
@ -1950,6 +1950,7 @@ dependencies = [
|
|||
"base64 0.22.1",
|
||||
"data-url",
|
||||
"gcp_auth",
|
||||
"mime_guess",
|
||||
"moka",
|
||||
"rand 0.8.7",
|
||||
"reqwest 0.12.28",
|
||||
|
|
|
|||
|
|
@ -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/`.
|
||||
|
|
@ -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.
|
||||
|
|
@ -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.
|
||||
|
|
@ -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
|
||||
```
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)]
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
|
|
|
|||
|
|
@ -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}");
|
||||
}
|
||||
|
|
@ -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(())
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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(())
|
||||
}
|
||||
|
|
@ -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)?;
|
||||
|
|
|
|||
148
litellm-rust/crates/python-bridge/src/routes/ocr_document.rs
Normal file
148
litellm-rust/crates/python-bridge/src/routes/ocr_document.rs
Normal 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)?)
|
||||
}
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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=",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue