refactor(ocr): move public OCR into native lifecycle

This commit is contained in:
Yujong Lee 2026-09-11 23:22:14 -07:00
parent 4d5c83df93
commit da1882d18f
61 changed files with 5270 additions and 2852 deletions

View file

@ -1,167 +0,0 @@
# Call lifecycle
`litellm_core::call_lifecycle` is the shared execution wrapper for LiteLLM call
types migrated to Rust. It owns lifecycle ordering, phase timing, and trace
observer calls. It must not know about OCR, chat, messages, responses,
completions, provider auth, request transforms, or response normalization.
Call-type modules own their domain behavior. For example, OCR owns document
payloads, OCR provider transforms, safe document fetch, guardrail payload shape,
callback payload shape, and provider HTTP execution.
## Runtime order
Every wrapped call runs in this order:
1. `async_pre_call_hook`
2. `async_during_call_hook`
3. provider call
4. `async_log_success_event` or `async_log_failure_event`
`async_pre_call_hook` receives the initial LiteLLM request shape. It is where
pre-call custom guardrails run.
`async_during_call_hook` converts the initial request into the provider-ready
request. It is where provider config selection, parameter mapping, auth/header
resolution, request transforms, and during-call guardrails belong.
The provider call receives only the provider-ready request. It should execute
I/O and call the provider response transform.
Success and failure callbacks receive `CallLifecycleTiming`. Callback failures
must not replace the original provider or guardrail result.
## Trace contract
The lifecycle runner records:
- full call start and end time
- `pre_call` phase timing
- `during_call` phase timing
- `provider_call` phase timing
- `success_callback` phase timing
- `failure_callback` phase timing
`CallLifecycleObserver` receives phase start and end events. The default
observer is a no-op. Future OTEL support should implement this observer instead
of editing OCR, chat, messages, responses, completions, or provider modules.
## Required shape
Each migrated call type should use this folder shape:
```text
litellm-rust/crates/ai-gateway/src/<call_type>/
mod.rs # thin public entrypoint
types.rs # public request, prepared request, provider request, response types
prepare.rs # model/provider/callback/guardrail setup
hooks.rs # CallLifecycleHooks implementation
handler.rs # provider I/O and response normalization
tests.rs # call-type lifecycle and handler tests
```
Provider transforms can live in `litellm-rust/crates/core/src/providers/...`.
Shared call-type helpers can live beside the call type, but generic lifecycle
code stays in this folder.
## Core API
The prepared request implements `CallLifecycleRequest`:
```rust
impl CallLifecycleRequest for PreparedMessagesRequest {
fn lifecycle_context(&self) -> CallLifecycleContext {
CallLifecycleContext::new(
"messages",
self.model.clone(),
self.custom_llm_provider.clone(),
self.litellm_call_id.clone(),
)
}
}
```
The call-type hooks implement `CallLifecycleHooks`:
```rust
impl CallLifecycleHooks<
PreparedMessagesRequest,
ProviderMessagesRequest,
MessagesResponse,
> for MessagesLifecycleHooks {
fn async_pre_call_hook(...) {
// run pre-call custom guardrails against the LiteLLM request shape
}
fn async_during_call_hook(...) {
// map params, validate env, transform request, run during-call guardrails
}
fn async_log_success_event(...) {
// call async_log_success_event on configured custom loggers
}
fn async_log_failure_event(...) {
// call async_log_failure_event without swallowing the original error
}
}
```
The public entrypoint stays thin:
```rust
pub async fn messages(request: MessagesRequest<'_>) -> CoreResult<MessagesResponse> {
let PreparedMessagesCall { request, hooks } = prepare_messages_call(request)?;
CallLifecycle::default()
.run_request(request, &hooks, execute_messages_provider_call)
.await
}
```
Use `run_request` for new call types. Keep `run` available only for specialized
tests or existing code that already has a `CallLifecycleContext`.
## Adding a new call type
1. Add `<call_type>/types.rs`
Define the public request accepted by the bridge, the prepared request used by
the lifecycle runner, and the provider request consumed by the handler.
2. Implement `CallLifecycleRequest`
Return `call_type`, `model`, `custom_llm_provider`, and `litellm_call_id`.
Do not put provider-specific logic here.
3. Add `<call_type>/prepare.rs`
Resolve model/provider once, generate or preserve `litellm_call_id`, construct
callback and guardrail runners, and return `Prepared<CallType>Call`.
4. Add `<call_type>/hooks.rs`
Implement `CallLifecycleHooks`. Put pre-call guardrail payload construction,
provider config selection, param mapping, request transform, during-call
guardrail payload construction, and callback payload construction here.
5. Add `<call_type>/handler.rs`
Execute the provider request and normalize the provider response. Do not repeat
provider-specific transforms here; call the provider config.
6. Add tests
Cover hook order, success callback payload, failure callback payload, pre-call
guardrail blocking before provider I/O, during-call body mutation, and provider
error mapping.
## Review checklist
- Core lifecycle has no call-type or provider-specific branches
- Public call-type entrypoint only prepares and calls `run_request`
- Provider behavior lives behind provider config/transformation code
- Hook method names map to the Python custom logger and guardrail concepts
- Phase timing is recorded once in lifecycle, not separately per call type
- Callback failures never hide the original provider or guardrail error
- Tests prove the provider socket is not touched when pre-call guardrails block

View file

@ -0,0 +1,94 @@
pub enum HostStep<V, S> {
Ready(V),
Suspend(S),
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum HostPhase {
Setup,
DeploymentPreCall,
Prepare,
Execute,
ConstructResponse,
DeploymentPostCall,
Finalize,
Success,
MapFailure,
DeploymentFailure,
Failure,
AsyncFailure,
Complete,
}
#[derive(Clone, Debug)]
pub enum HostFailure {
Error(crate::Error),
Cancelled(crate::Error),
}
pub struct HostLifecycle {
phase: HostPhase,
asynchronous: bool,
}
impl HostLifecycle {
pub fn new(asynchronous: bool) -> Self {
Self {
phase: HostPhase::Setup,
asynchronous,
}
}
pub fn phase(&self) -> HostPhase {
self.phase
}
pub fn accept(&mut self, result: Result<(), HostFailure>) -> Option<crate::Error> {
if let Err(failure) = result {
if self.phase == HostPhase::DeploymentFailure {
self.phase = HostPhase::Failure;
return None;
}
let error = match failure {
HostFailure::Cancelled(error) => {
self.phase = HostPhase::Complete;
return Some(error);
}
HostFailure::Error(error) => error,
};
match self.phase {
HostPhase::Failure | HostPhase::AsyncFailure => {
self.advance();
return None;
}
HostPhase::Success => self.phase = HostPhase::Complete,
HostPhase::Execute | HostPhase::ConstructResponse => {
self.phase = HostPhase::MapFailure;
}
_ => self.phase = HostPhase::Failure,
}
return Some(error);
}
self.advance();
None
}
fn advance(&mut self) {
self.phase = match self.phase {
HostPhase::Setup if self.asynchronous => HostPhase::DeploymentPreCall,
HostPhase::Setup | HostPhase::DeploymentPreCall => HostPhase::Prepare,
HostPhase::Prepare => HostPhase::Execute,
HostPhase::Execute => HostPhase::ConstructResponse,
HostPhase::ConstructResponse if self.asynchronous => HostPhase::DeploymentPostCall,
HostPhase::ConstructResponse | HostPhase::DeploymentPostCall => HostPhase::Finalize,
HostPhase::Finalize => HostPhase::Success,
HostPhase::MapFailure if self.asynchronous => HostPhase::DeploymentFailure,
HostPhase::MapFailure | HostPhase::DeploymentFailure => HostPhase::Failure,
HostPhase::Failure if self.asynchronous => HostPhase::AsyncFailure,
HostPhase::Failure
| HostPhase::AsyncFailure
| HostPhase::Success
| HostPhase::Complete => HostPhase::Complete,
};
}
}

View file

@ -3,6 +3,10 @@ use std::time::{Instant, SystemTime, UNIX_EPOCH};
use crate::Error;
pub mod host;
#[cfg(test)]
#[path = "../../tests/host_lifecycle.rs"]
mod host_tests;
pub mod types;
pub use types::{

View file

@ -1,6 +1,6 @@
use thiserror::Error as ThisError;
#[derive(Debug, ThisError, PartialEq, Eq)]
#[derive(Clone, Debug, ThisError, PartialEq, Eq)]
pub enum Error {
#[error("expected {expected}, got {actual}")]
InvalidType {

View file

@ -10,7 +10,6 @@ use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
use crate::ocr::prepare::{credential_env, transform_request_body};
use crate::ocr::registry::OcrProvider;
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrResponseFormat};
use crate::ocr::wire::DecodedOcrResponse;
use crate::providers::azure_ai::auth::AzureAuthInputs;
use crate::url_utils::ApiUrl;
@ -32,18 +31,19 @@ impl OcrAdapter for AzureDocumentIntelligenceAdapter {
client: &OcrClient,
) -> Result<reqwest::Request, OcrError> {
let params = map_ocr_params(request)?;
let config = AzureAuthInputs::from_sourced_optional_params(
let mut config = AzureAuthInputs::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)
.map_err(Error::from)?;
config.azure_ad_token_provider = request.azure_ad_token_provider.clone();
let headers = validate_environment(&request.connection, &config, &credential_env).await?;
let endpoint = nonblank(request.connection.api_base.clone())
.or_else(|| nonblank(credential_env(AZURE_DI_ENDPOINT_ENV)))
.ok_or_else(|| Error::Auth("Missing Azure Document Intelligence API Base - Set AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT or pass api_base".into()))?;
let url = get_complete_url(&endpoint, &request.model, &params)?;
let body = document_intelligence::transform_ocr_request(request.document.clone())?;
transform_request_body(client, request, &url, &headers, body, |_| Ok(())).await
transform_request_body(client, request, &url, &headers, false, body, |_| Ok(())).await
}
fn transform_ocr_response(
@ -61,7 +61,7 @@ impl OcrAdapter for AzureDocumentIntelligenceAdapter {
url: &str,
headers: &[(String, String)],
request: &LiteLLMOcrRequest,
) -> Result<DecodedOcrResponse<Self::ProviderResponse>, OcrError> {
) -> Result<Vec<u8>, OcrError> {
polling::read_operation_response(
client.polling_http(),
response,
@ -69,8 +69,10 @@ impl OcrAdapter for AzureDocumentIntelligenceAdapter {
headers,
&request.connection,
request.response_format()? == OcrResponseFormat::Native,
&request.hooks,
)
.await
.map(|decoded| decoded.text.into_bytes())
}
}

View file

@ -1,3 +1,4 @@
use std::sync::Arc;
use std::time::Duration;
use reqwest::Url;
@ -9,6 +10,7 @@ use crate::ocr::codecs::document_intelligence::{
AzureDocumentIntelligenceOperation, OperationStatus,
};
use crate::ocr::error::{OcrError, OcrPollingError, OcrResponseError};
use crate::ocr::hooks::OcrHooks;
use crate::ocr::types::OcrConnection;
use crate::ocr::wire::DecodedOcrResponse;
@ -19,23 +21,29 @@ pub(super) async fn read_operation_response(
headers: &[(String, String)],
connection: &OcrConnection,
native: bool,
hooks: &Arc<dyn OcrHooks>,
) -> Result<DecodedOcrResponse<AzureDocumentIntelligenceOperation>, OcrError> {
if response.status() != reqwest::StatusCode::ACCEPTED {
return read_json_response(response, native).await;
let bytes = crate::ocr::client::read_response_bytes(response).await?;
crate::ocr::handler::post_call(hooks, &bytes).await?;
return Ok(crate::ocr::wire::decode_response(&bytes, native)?);
}
let location = response
.headers()
.get("operation-location")
.and_then(|value| value.to_str().ok())
.ok_or(OcrPollingError::PollLocation)?;
.ok_or(OcrPollingError::PollLocation)?
.to_string();
let original = Url::parse(original_url).map_err(|_| OcrPollingError::PollOrigin)?;
let operation = Url::parse(location).map_err(|_| OcrPollingError::PollOrigin)?;
let operation = Url::parse(&location).map_err(|_| OcrPollingError::PollOrigin)?;
if original.origin() != operation.origin()
|| !operation.username().is_empty()
|| operation.password().is_some()
{
return Err(OcrPollingError::PollOrigin.into());
}
let bytes = crate::ocr::client::read_response_bytes(response).await?;
crate::ocr::handler::post_call(hooks, &bytes).await?;
poll_operation(http_client, operation, headers, connection, native).await
}

View file

@ -33,13 +33,16 @@ impl OcrAdapter for AzureMistralAdapter {
known: params,
extra_params: _extra_params,
} = _prepare_ocr_request::<MistralOcrParams>(request)?;
let config = AzureAuthInputs::from_sourced_optional_params(
let mut config = AzureAuthInputs::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)
.map_err(Error::from)?;
let headers = validate_environment(&request.connection, &config, &credential_env).await?;
config.azure_ad_token_provider = request.azure_ad_token_provider.clone();
let url = get_complete_url(request.connection.api_base.as_deref(), &credential_env)?;
let headers = validate_environment(&request.connection, &config, &credential_env).await?;
let retains_document = !request.document.source().starts_with("http://")
&& !request.document.source().starts_with("https://");
let document = inline_remote_document(
client.document_fetcher(),
request.document.clone(),
@ -47,9 +50,15 @@ impl OcrAdapter for AzureMistralAdapter {
)
.await?;
let body = mistral::transform_ocr_request(&request.model, document, &params)?;
transform_request_body(client, request, &url, &headers, body, |body| {
validate_inline_document(&body.document)
})
transform_request_body(
client,
request,
&url,
&headers,
retains_document,
body,
|body| validate_inline_document(&body.document),
)
.await
}
@ -89,6 +98,9 @@ async fn validate_environment(
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Vec<(String, String)>, OcrError> {
if crate::http_utils::has_header(&connection.extra_headers, "authorization") {
if config.azure_ad_token_provider.is_some() {
super::resolve_entra(config, env_lookup).await?;
}
super::validate_destination(connection, connection.extra_headers_source)?;
return Ok(connection.extra_headers.clone());
}

View file

@ -33,7 +33,7 @@ impl OcrAdapter for MistralAdapter {
let url = get_complete_url(request.connection.api_base.as_deref())?;
let body =
mistral::transform_ocr_request(&request.model, request.document.clone(), &params)?;
transform_request_body(client, request, &url, &headers, body, |_| Ok(())).await
transform_request_body(client, request, &url, &headers, true, body, |_| Ok(())).await
}
fn transform_ocr_response(

View file

@ -5,8 +5,7 @@ use serde::de::DeserializeOwned;
use super::OcrClient;
use super::error::{OcrError, OcrResponseError};
use super::registry::OcrProvider;
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrResponseFormat};
use super::wire::DecodedOcrResponse;
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
mod azure;
mod mistral;
@ -55,12 +54,12 @@ pub(crate) trait OcrAdapter: Send + Sync + Sized + 'static {
_url: &str,
_headers: &[(String, String)],
request: &LiteLLMOcrRequest,
) -> impl Future<Output = Result<DecodedOcrResponse<Self::ProviderResponse>, OcrError>> + Send
{
let retain_native = request
.response_format()
.map(|format| format == OcrResponseFormat::Native);
async move { super::client::read_json_response(response, retain_native?).await }
) -> impl Future<Output = Result<Vec<u8>, OcrError>> + Send {
async move {
let bytes = super::client::read_response_bytes(response).await?;
super::handler::post_call(&request.hooks, &bytes).await?;
Ok(bytes)
}
}
}

View file

@ -27,7 +27,7 @@ impl OcrAdapter for ReductoLegacyAdapter {
} = _prepare_ocr_request::<ReductoLegacyParams>(request)?;
let headers = super::validate_environment(&request.connection, &credential_env)?;
let url = super::get_complete_url(request.connection.api_base.as_deref(), "parse")?;
let document = guardrail_document(request, &url).await?;
let (document, headers) = guardrail_document(request, &url, &headers).await?;
let document =
super::prepare_document(client, document, &request.connection, &headers).await?;
let body = reducto::transform_legacy_ocr_request(&request.model, document, &params)?;

View file

@ -27,7 +27,7 @@ impl OcrAdapter for ReductoV3Adapter {
} = _prepare_ocr_request::<ReductoV3Params>(request)?;
let headers = super::validate_environment(&request.connection, &credential_env)?;
let url = super::get_complete_url(request.connection.api_base.as_deref(), "parse")?;
let document = guardrail_document(request, &url).await?;
let (document, headers) = guardrail_document(request, &url, &headers).await?;
let document =
super::prepare_document(client, document, &request.connection, &headers).await?;
let body = reducto::transform_v3_ocr_request(&request.model, document, &params)?;

View file

@ -57,9 +57,15 @@ impl OcrAdapter for VertexDeepSeekAdapter {
let document = request.document.clone();
let body =
deepseek::transform_ocr_request(&provider_model(&request.model), document, &params)?;
transform_request_body(client, request, &url, &authentication.headers, body, |_| {
Ok(())
})
transform_request_body(
client,
request,
&url,
&authentication.headers,
false,
body,
|_| Ok(()),
)
.await
}

View file

@ -54,6 +54,8 @@ impl OcrAdapter for VertexMistralAdapter {
&location,
&request.model,
)?;
let retains_document = !request.document.source().starts_with("http://")
&& !request.document.source().starts_with("https://");
let document = inline_remote_document(
client.document_fetcher(),
request.document.clone(),
@ -66,6 +68,7 @@ impl OcrAdapter for VertexMistralAdapter {
request,
&url,
&authentication.headers,
retains_document,
body,
|body| validate_inline_document(&body.document),
)

View file

@ -4,7 +4,6 @@ use std::time::Duration;
use serde::de::DeserializeOwned;
use super::error::OcrError;
use super::handler::perform_ocr_request;
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
use super::wire::{DecodedOcrResponse, decode_response};
use crate::Error;
@ -32,6 +31,10 @@ impl OcrClient {
})
}
pub fn shared() -> Result<Self, Error> {
shared_client()
}
#[tracing::instrument(
name = "ocr",
target = "litellm::function_trace",
@ -39,7 +42,34 @@ impl OcrClient {
skip_all
)]
pub async fn perform(&self, request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, Error> {
perform_ocr_request(self, request).await
use super::{
NativeOutcome, OcrAdmission, OcrCall, OcrCallStep, OcrHookHost, OcrHost,
OcrHostOperation, OcrHostResult,
};
let host = OcrHookHost::new(request.hooks.clone());
let mut request = Some(request);
let NativeOutcome::Completed(mut call) = OcrCall::admit(self.clone(), OcrAdmission::all())
else {
return Err(Error::InvalidRequest(
"native OCR host admission declined".into(),
));
};
let mut result = None;
loop {
match call.resume(result.take()).await? {
OcrCallStep::Host(OcrHostOperation::ProjectRequest) => {
result = Some(OcrHostResult::Request(Ok((
Box::new(request.take().ok_or_else(|| {
Error::InvalidRequest("OCR request was already projected".into())
})?),
false,
))))
}
OcrCallStep::Host(operation) => result = Some(host.invoke(operation).await),
OcrCallStep::Complete(response) => return Ok(response),
}
}
}
pub(crate) fn provider_http(&self) -> &reqwest::Client {
@ -77,7 +107,7 @@ fn no_redirect_http() -> Result<reqwest::Client, TransportError> {
.map_err(TransportError::from)
}
pub async fn ocr(request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, Error> {
pub(crate) fn shared_client() -> Result<OcrClient, Error> {
static CLIENT: OnceLock<Result<OcrClient, TransportError>> = OnceLock::new();
let client = CLIENT
.get_or_init(|| {
@ -88,18 +118,24 @@ pub async fn ocr(request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, Error
.and_then(OcrClient::new)
})
.clone()?;
client.perform(request).await
Ok(client)
}
pub async fn ocr(request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, Error> {
shared_client()?.perform(request).await
}
pub async fn read_json_response<T: DeserializeOwned>(
response: reqwest::Response,
native: bool,
) -> Result<DecodedOcrResponse<T>, OcrError> {
let bytes = read_response_bytes(response).await?;
Ok(decode_response(&bytes, native)?)
}
pub(crate) async fn read_response_bytes(response: reqwest::Response) -> Result<Vec<u8>, OcrError> {
let status = response.status();
let bytes = response
.bytes()
.await
.map_err(crate::error::TransportError::from)?;
let bytes = response.bytes().await.map_err(transport_error)?;
if !status.is_success() {
return Err(crate::error::TransportError::Http {
status: status.as_u16(),
@ -107,5 +143,15 @@ pub async fn read_json_response<T: DeserializeOwned>(
}
.into());
}
Ok(decode_response(&bytes, native)?)
Ok(bytes.to_vec())
}
pub(crate) fn transport_error(error: reqwest::Error) -> Error {
if error.is_timeout() {
return Error::Http {
status: 408,
body: "OCR request timed out".into(),
};
}
crate::error::TransportError::from(error).into()
}

View file

@ -2,6 +2,7 @@ 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 super::error::{OcrError, OcrRequestError, OcrResponseError};
use super::types::{OcrConnection, OcrDocument};
@ -9,6 +10,63 @@ use crate::constants::OCR_MAX_FETCH_REDIRECTS;
use crate::error::{MediaError, TransportError};
use crate::media::{DownloadPolicy, MediaFetcher};
pub fn encode_file_document(
bytes: &[u8],
file_name: Option<&str>,
mime_type: Option<&str>,
) -> Result<Value, OcrRequestError> {
if bytes.is_empty() {
return Err(OcrRequestError::RequestField {
path: "document.file".into(),
});
}
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(),
});
}
let mime_type = mime_type
.map(str::to_string)
.or_else(|| file_name.and_then(mime_type_for_name).map(str::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")
} else {
("document_url", "document_url")
};
Ok(Value::Object(Map::from_iter([
("type".into(), Value::String(kind.into())),
(field.into(), Value::String(source)),
])))
}
fn valid_mime_type(value: &str) -> bool {
let Some((kind, subtype)) = value.split_once('/') else {
return false;
};
!kind.is_empty()
&& !subtype.is_empty()
&& value.bytes().all(|byte| {
byte.is_ascii_alphanumeric() || matches!(byte, b'/' | b'.' | b'+' | b'-' | b'_')
})
}
fn mime_type_for_name(name: &str) -> Option<&'static str> {
let extension = name.rsplit_once('.')?.1;
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,
}
}
pub(crate) struct InlineDocument<'a>(DataUrl<'a>);
impl<'a> InlineDocument<'a> {
@ -114,6 +172,30 @@ mod tests {
}
}
#[test]
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"
})
);
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"
})
);
}
#[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());
}
#[test]
fn decodes_data_urls_and_limits_decoded_size() {
for (source, expected) in [

View file

@ -1,15 +1,17 @@
use super::OcrClient;
use super::adapters::OcrAdapter;
use super::hooks::OcrLifecycleHooks;
use super::hooks::{OcrHooks, OcrLifecycleHooks, OcrPostCallRequest};
use super::registry::OcrAdapterKind;
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
use crate::Error;
use crate::call_lifecycle::{CallLifecycle, CallLifecycleContext};
use std::sync::Arc;
pub(crate) async fn perform_ocr_request(
client: &OcrClient,
request: LiteLLMOcrRequest,
) -> Result<LiteLLMOcrResponse, Error> {
request.response_format()?;
let context = CallLifecycleContext::new(
"ocr",
request.model.clone(),
@ -23,27 +25,71 @@ pub(crate) async fn perform_ocr_request(
hooks: request.hooks.clone(),
provider_name: context.custom_llm_provider.clone(),
};
CallLifecycle::default().run(context, request, &hooks, |request| async move {
macro_rules! execute_selected_adapter {
CallLifecycle::default()
.run(context, request, &hooks, |request| async move {
PreparedOcrCall::prepare(client.clone(), request)
.await?
.execute()
.await?
.normalize()
})
.await
}
pub(crate) struct PreparedOcrCall {
client: OcrClient,
request: LiteLLMOcrRequest,
http: reqwest::Request,
}
impl PreparedOcrCall {
pub(crate) async fn prepare(
client: OcrClient,
request: LiteLLMOcrRequest,
) -> Result<Self, Error> {
macro_rules! prepare_adapter {
($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => {
match request.adapter {
$( OcrAdapterKind::$variant => execute_ocr_provider_call(client, &$instance, request).await, )+
$( OcrAdapterKind::$variant => $instance.prepare_request(&request, &client).await?, )+
}
};
}
super::adapters::for_each_ocr_adapter!(execute_selected_adapter)
}).await
let http = super::adapters::for_each_ocr_adapter!(prepare_adapter);
Ok(Self {
client,
request,
http,
})
}
pub(crate) async fn execute(self) -> Result<OcrProviderResponse, Error> {
let url = self.http.url().to_string();
let headers = request_headers(&self.http)?;
let response = crate::http_utils::http_request(reqwest::RequestBuilder::from_parts(
self.client.provider_http().clone(),
self.http,
))
.await
.map_err(super::client::transport_error)?;
macro_rules! read_adapter {
($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => {
match self.request.adapter {
$( OcrAdapterKind::$variant => {
let bytes = $instance.read_response(&self.client, response, &url, &headers, &self.request).await?;
Ok(OcrProviderResponse {
request: self.request,
bytes,
})
}, )+
}
};
}
super::adapters::for_each_ocr_adapter!(read_adapter)
}
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
async fn execute_ocr_provider_call<A: OcrAdapter>(
client: &OcrClient,
adapter: &A,
request: LiteLLMOcrRequest,
) -> Result<LiteLLMOcrResponse, Error> {
let provider_request = adapter.prepare_request(&request, client).await?;
let url = provider_request.url().to_string();
let headers = provider_request
fn request_headers(request: &reqwest::Request) -> Result<Vec<(String, String)>, Error> {
request
.headers()
.iter()
.map(|(name, value)| {
@ -53,20 +99,39 @@ async fn execute_ocr_provider_call<A: OcrAdapter>(
.map_err(|_| super::error::OcrRequestError::RequestField {
path: "headers".into(),
})
.map_err(Error::from)
})
.collect::<Result<Vec<_>, _>>()?;
let response = crate::http_utils::http_request(reqwest::RequestBuilder::from_parts(
client.provider_http().clone(),
provider_request,
))
.await
.map_err(crate::error::TransportError::from)?;
let decoded = adapter
.read_response(client, response, &url, &headers, &request)
.await?;
let response = adapter.transform_ocr_response(&request, decoded.data)?;
Ok(LiteLLMOcrResponse {
provider_native_response: decoded.native,
..response
})
.collect()
}
macro_rules! provider_data {
($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => {
impl OcrProviderResponse {
pub(crate) fn normalize(self) -> Result<LiteLLMOcrResponse, Error> {
let native = self.request.response_format()? == super::types::OcrResponseFormat::Native;
match self.request.adapter {
$( OcrAdapterKind::$variant => {
let decoded = super::wire::decode_response::<<$adapter as OcrAdapter>::ProviderResponse>(&self.bytes, native)?;
let response = $instance.transform_ocr_response(&self.request, decoded.data)?;
Ok(LiteLLMOcrResponse { provider_native_response: decoded.native, ..response })
}, )+
}
}
}
};
}
pub(crate) struct OcrProviderResponse {
request: LiteLLMOcrRequest,
bytes: Vec<u8>,
}
pub(crate) async fn post_call(hooks: &Arc<dyn OcrHooks>, bytes: &[u8]) -> Result<(), Error> {
let original_response = serde_json::Value::String(String::from_utf8_lossy(bytes).into_owned());
hooks
.post_call(OcrPostCallRequest { original_response })
.await?;
Ok(())
}
super::adapters::for_each_ocr_adapter!(provider_data);

View file

@ -24,7 +24,15 @@ pub struct OcrDuringCallRequest {
pub model: String,
pub custom_llm_provider: String,
pub url: String,
pub headers: Vec<(String, String)>,
pub body: Value,
#[serde(skip)]
pub retained_fields: Vec<String>,
}
#[derive(Clone, Debug, Serialize)]
pub struct OcrPostCallRequest {
pub original_response: Value,
}
pub trait OcrHooks: Send + Sync {
@ -40,6 +48,9 @@ pub trait OcrHooks: Send + Sync {
) -> OcrHookFuture<'_, OcrDuringCallRequest> {
Box::pin(async move { Ok(request) })
}
fn post_call(&self, request: OcrPostCallRequest) -> OcrHookFuture<'_, OcrPostCallRequest> {
Box::pin(async move { Ok(request) })
}
fn success<'a>(
&'a self,
_context: &'a CallLifecycleContext,

View file

@ -0,0 +1,612 @@
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use tokio::sync::{mpsc, oneshot};
use super::handler::perform_ocr_request;
use super::hooks::{
OcrDuringCallRequest, OcrHookFuture, OcrHooks, OcrLogFuture, OcrPostCallRequest,
OcrPreCallRequest,
};
use super::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrClient};
use crate::AuthError;
use crate::Error;
use crate::auth::{ResolvedCredential, TokenFuture, TokenProvider, TokenProviderHandle};
use crate::call_lifecycle::host::{HostFailure, HostLifecycle, HostPhase};
use crate::call_lifecycle::{CallLifecycleContext, CallLifecycleTiming};
pub type NativeResult<T> = Result<NativeOutcome<T>, Error>;
#[derive(Debug, PartialEq, Eq)]
pub enum NativeOutcome<T> {
Completed(T),
Declined(OcrDecline),
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum OcrDecline {
ProviderWorkflow,
HostOperations,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct OcrAdmission {
pub provider_workflow: bool,
pub host_operations: bool,
pub azure_ad_token_provider: bool,
pub asynchronous: bool,
}
impl OcrAdmission {
pub const fn all() -> Self {
Self {
provider_workflow: true,
host_operations: true,
azure_ad_token_provider: false,
asynchronous: false,
}
}
}
#[derive(Clone, Debug)]
pub enum OcrHostOperation {
ProjectRequest,
Lifecycle(HostPhase),
ConstructResponse(Arc<LiteLLMOcrResponse>),
MapFailure(Error),
Success {
context: CallLifecycleContext,
response: Arc<LiteLLMOcrResponse>,
timing: CallLifecycleTiming,
},
Failure {
context: CallLifecycleContext,
error: Error,
timing: CallLifecycleTiming,
},
AcquireAzureAdToken,
PreCall(OcrPreCallRequest),
DuringCall(OcrDuringCallRequest),
PostCall(OcrPostCallRequest),
}
pub enum OcrHostResult {
Request(Result<(Box<LiteLLMOcrRequest>, bool), Error>),
Lifecycle(Result<(), HostFailure>),
AzureAdToken(Result<ResolvedCredential, AuthError>),
PreCall(Result<OcrPreCallRequest, Error>),
DuringCall(Result<OcrDuringCallRequest, Error>),
PostCall(Result<OcrPostCallRequest, Error>),
}
pub enum OcrCallStep {
Host(OcrHostOperation),
Complete(LiteLLMOcrResponse),
}
pub struct OcrCall {
lifecycle: HostLifecycle,
execution: OcrExecution,
response: Option<Arc<LiteLLMOcrResponse>>,
error: Option<Error>,
pending: bool,
completed: bool,
projecting: bool,
}
impl OcrCall {
pub fn admit(client: OcrClient, admission: OcrAdmission) -> NativeOutcome<Self> {
if !admission.provider_workflow {
return NativeOutcome::Declined(OcrDecline::ProviderWorkflow);
}
if !admission.host_operations {
return NativeOutcome::Declined(OcrDecline::HostOperations);
}
NativeOutcome::Completed(Self {
lifecycle: HostLifecycle::new(admission.asynchronous),
execution: OcrExecution::new(client, admission.azure_ad_token_provider),
response: None,
error: None,
pending: false,
completed: false,
projecting: false,
})
}
pub async fn resume(&mut self, result: Option<OcrHostResult>) -> Result<OcrCallStep, Error> {
if self.completed {
return Err(Error::InvalidRequest(
"OCR call cannot be resumed after completion".into(),
));
}
if self.pending != result.is_some() {
return Err(Error::InvalidRequest(
"OCR host operation result does not match pending state".into(),
));
}
match &result {
Some(OcrHostResult::Lifecycle(Ok(())))
if self.lifecycle.phase() == HostPhase::Execute =>
{
return Err(Error::InvalidRequest(
"OCR provider operation requires a typed result".into(),
));
}
Some(result)
if !matches!(result, OcrHostResult::Lifecycle(_))
&& self.lifecycle.phase() != HostPhase::Execute =>
{
return Err(Error::InvalidRequest(
"unexpected OCR provider operation result".into(),
));
}
_ => {}
}
self.pending = false;
let provider_result = match result {
Some(OcrHostResult::Request(result)) if self.projecting => {
self.projecting = false;
match result {
Ok((request, azure_ad_token_provider)) => {
self.execution.request = Some(*request);
self.execution.azure_ad_token_provider = azure_ad_token_provider;
}
Err(error) => self.accept(Err(HostFailure::Error(error))),
}
None
}
Some(OcrHostResult::Request(_)) => {
return Err(Error::InvalidRequest(
"unexpected OCR request projection".into(),
));
}
Some(OcrHostResult::Lifecycle(result)) => {
self.accept(result);
None
}
result => result,
};
if self.lifecycle.phase() == HostPhase::Execute {
if self.execution.request.is_none()
&& self.execution.execution.is_none()
&& !self.execution.completed
{
self.projecting = true;
return Ok(self.host_step(OcrHostOperation::ProjectRequest));
}
match self.execution.resume(provider_result).await {
Ok(OcrCallStep::Host(operation)) => return Ok(self.host_step(operation)),
Ok(OcrCallStep::Complete(response)) => {
self.response = Some(Arc::new(response));
self.accept(Ok(()));
}
Err(error) => self.accept(Err(HostFailure::Error(error))),
}
}
if self.error.is_some() {
self.execution.stop().await;
}
let operation = match self.lifecycle.phase() {
HostPhase::Complete => {
self.completed = true;
return match self.error.take() {
Some(error) => Err(error),
None => self
.response
.take()
.map(Arc::unwrap_or_clone)
.map(OcrCallStep::Complete)
.ok_or_else(|| {
Error::InvalidRequest("OCR completed without a response".into())
}),
};
}
HostPhase::ConstructResponse => OcrHostOperation::ConstructResponse(
self.response
.as_ref()
.ok_or_else(|| Error::InvalidRequest("missing OCR response".into()))?
.clone(),
),
HostPhase::MapFailure => OcrHostOperation::MapFailure(
self.error
.as_ref()
.ok_or_else(|| Error::InvalidRequest("missing OCR failure".into()))?
.clone(),
),
HostPhase::Success | HostPhase::Failure => {
let snapshot = self
.execution
.terminal
.lock()
.unwrap_or_else(|error| error.into_inner())
.clone();
match (self.lifecycle.phase(), snapshot) {
(HostPhase::Success, Some((context, timing))) => OcrHostOperation::Success {
context,
response: self
.response
.as_ref()
.ok_or_else(|| Error::InvalidRequest("missing OCR response".into()))?
.clone(),
timing,
},
(HostPhase::Failure, Some((context, timing))) => OcrHostOperation::Failure {
context,
error: self
.error
.as_ref()
.ok_or_else(|| Error::InvalidRequest("missing OCR failure".into()))?
.clone(),
timing,
},
(phase, _) => OcrHostOperation::Lifecycle(phase),
}
}
phase => OcrHostOperation::Lifecycle(phase),
};
Ok(self.host_step(operation))
}
fn accept(&mut self, result: Result<(), HostFailure>) {
let cancelled = matches!(&result, Err(HostFailure::Cancelled(_)));
if let Some(error) = self.lifecycle.accept(result) {
if cancelled {
self.error = Some(error);
} else {
self.error.get_or_insert(error);
}
self.execution.cancel();
}
}
pub async fn interrupt(&mut self, failure: HostFailure) -> Result<OcrCallStep, Error> {
if self.completed {
return Err(Error::InvalidRequest(
"OCR call cannot be interrupted after completion".into(),
));
}
self.pending = false;
self.accept(Err(failure));
self.resume(None).await
}
fn host_step(&mut self, operation: OcrHostOperation) -> OcrCallStep {
self.pending = true;
OcrCallStep::Host(operation)
}
}
struct PendingOperation {
operation: OcrHostOperation,
result: oneshot::Sender<OcrHostResult>,
}
struct OcrExecution {
client: Option<OcrClient>,
request: Option<LiteLLMOcrRequest>,
operations_tx: mpsc::UnboundedSender<PendingOperation>,
operations_rx: mpsc::UnboundedReceiver<PendingOperation>,
pending_result: Option<oneshot::Sender<OcrHostResult>>,
execution: Option<tokio::task::JoinHandle<Result<LiteLLMOcrResponse, Error>>>,
completed: bool,
azure_ad_token_provider: bool,
terminal: Arc<std::sync::Mutex<Option<(CallLifecycleContext, CallLifecycleTiming)>>>,
}
impl OcrExecution {
fn new(client: OcrClient, azure_ad_token_provider: bool) -> Self {
let (operations_tx, operations_rx) = mpsc::unbounded_channel();
Self {
client: Some(client),
request: None,
operations_tx,
operations_rx,
pending_result: None,
execution: None,
completed: false,
azure_ad_token_provider,
terminal: Arc::default(),
}
}
pub async fn resume(&mut self, result: Option<OcrHostResult>) -> Result<OcrCallStep, Error> {
if self.completed {
return Err(Error::InvalidRequest(
"OCR call cannot be resumed after completion".into(),
));
}
match (self.pending_result.take(), result) {
(Some(sender), Some(result)) => sender
.send(result)
.map_err(|_| Error::InvalidRequest("OCR host operation was abandoned".into()))?,
(None, None) if self.execution.is_none() => self.start(),
(Some(sender), None) => {
self.pending_result = Some(sender);
return Err(Error::InvalidRequest(
"OCR host operation result is required".into(),
));
}
(None, Some(_)) => {
return Err(Error::InvalidRequest(
"unexpected OCR host operation result".into(),
));
}
(None, None) => {}
}
let execution = self.execution.as_mut().ok_or_else(|| {
Error::InvalidRequest("OCR call cannot be resumed after completion".into())
})?;
tokio::select! {
operation = self.operations_rx.recv() => {
let operation = operation.ok_or_else(|| Error::InvalidRequest("OCR operation channel closed".into()))?;
self.pending_result = Some(operation.result);
Ok(OcrCallStep::Host(operation.operation))
}
result = execution => {
self.execution = None;
self.completed = true;
result
.map_err(|error| Error::Network(format!("OCR execution task failed: {error}")))?
.map(OcrCallStep::Complete)
}
}
}
fn start(&mut self) {
let client = self.client.take().expect("admitted OCR call has a client");
let mut request = self
.request
.take()
.expect("admitted OCR call has a request");
let has_guardrails = request.hooks.has_guardrails();
if self.azure_ad_token_provider {
request.azure_ad_token_provider = Some(TokenProviderHandle::new(Arc::new(
OcrAzureAdTokenProvider {
operations: self.operations_tx.clone(),
},
)));
}
request.hooks = Arc::new(ProtocolHooks {
operations: self.operations_tx.clone(),
has_guardrails,
terminal: self.terminal.clone(),
});
self.execution = Some(tokio::spawn(async move {
perform_ocr_request(&client, request).await
}));
}
fn cancel(&mut self) {
self.pending_result = None;
if let Some(execution) = &self.execution {
execution.abort();
}
}
async fn stop(&mut self) {
self.cancel();
if let Some(execution) = self.execution.as_mut() {
let _ = execution.await;
}
self.execution = None;
}
}
impl Drop for OcrExecution {
fn drop(&mut self) {
if let Some(execution) = &self.execution {
execution.abort();
}
}
}
struct ProtocolHooks {
operations: mpsc::UnboundedSender<PendingOperation>,
has_guardrails: bool,
terminal: Arc<std::sync::Mutex<Option<(CallLifecycleContext, CallLifecycleTiming)>>>,
}
#[derive(Debug)]
struct OcrAzureAdTokenProvider {
operations: mpsc::UnboundedSender<PendingOperation>,
}
impl TokenProvider for OcrAzureAdTokenProvider {
fn acquire(&self) -> TokenFuture<'_> {
Box::pin(async move {
let (result, receiver) = oneshot::channel();
self.operations
.send(PendingOperation {
operation: OcrHostOperation::AcquireAzureAdToken,
result,
})
.map_err(|_| {
AuthError::AzureTokenAcquisition("OCR host driver was abandoned".into())
})?;
match receiver.await.map_err(|_| {
AuthError::AzureTokenAcquisition(
"OCR token provider operation was abandoned".into(),
)
})? {
OcrHostResult::AzureAdToken(result) => result,
_ => Err(AuthError::AzureTokenAcquisition(
"invalid OCR token provider host result".into(),
)),
}
})
}
}
impl ProtocolHooks {
async fn invoke(&self, operation: OcrHostOperation) -> Result<OcrHostResult, Error> {
let (result, receiver) = oneshot::channel();
self.operations
.send(PendingOperation { operation, result })
.map_err(|_| Error::InvalidRequest("OCR host driver was abandoned".into()))?;
receiver
.await
.map_err(|_| Error::InvalidRequest("OCR host operation was abandoned".into()))
}
}
impl OcrHooks for ProtocolHooks {
fn has_guardrails(&self) -> bool {
self.has_guardrails
}
fn pre_call(&self, request: OcrPreCallRequest) -> OcrHookFuture<'_, OcrPreCallRequest> {
Box::pin(async move {
match self.invoke(OcrHostOperation::PreCall(request)).await? {
OcrHostResult::PreCall(result) => result,
_ => Err(Error::InvalidRequest(
"invalid OCR pre-call host result".into(),
)),
}
})
}
fn during_call(
&self,
request: OcrDuringCallRequest,
) -> OcrHookFuture<'_, OcrDuringCallRequest> {
Box::pin(async move {
match self.invoke(OcrHostOperation::DuringCall(request)).await? {
OcrHostResult::DuringCall(result) => result,
_ => Err(Error::InvalidRequest(
"invalid OCR during-call host result".into(),
)),
}
})
}
fn post_call(&self, request: OcrPostCallRequest) -> OcrHookFuture<'_, OcrPostCallRequest> {
Box::pin(async move {
match self.invoke(OcrHostOperation::PostCall(request)).await? {
OcrHostResult::PostCall(result) => result,
_ => Err(Error::InvalidRequest(
"invalid OCR post-call host result".into(),
)),
}
})
}
fn success<'a>(
&'a self,
context: &'a CallLifecycleContext,
_response: &'a LiteLLMOcrResponse,
timing: &'a CallLifecycleTiming,
) -> OcrLogFuture<'a> {
Box::pin(async move {
*self
.terminal
.lock()
.unwrap_or_else(|error| error.into_inner()) =
Some((context.clone(), timing.clone()));
})
}
fn failure<'a>(
&'a self,
context: &'a CallLifecycleContext,
_error: &'a Error,
timing: &'a CallLifecycleTiming,
) -> OcrLogFuture<'a> {
Box::pin(async move {
*self
.terminal
.lock()
.unwrap_or_else(|error| error.into_inner()) =
Some((context.clone(), timing.clone()));
})
}
}
pub type OcrHostFuture<'a> = Pin<Box<dyn Future<Output = OcrHostResult> + Send + 'a>>;
pub trait OcrHost: Send + Sync {
fn invoke(&self, operation: OcrHostOperation) -> OcrHostFuture<'_>;
}
pub struct NoopOcrHost;
impl OcrHost for NoopOcrHost {
fn invoke(&self, operation: OcrHostOperation) -> OcrHostFuture<'_> {
Box::pin(async move {
match operation {
OcrHostOperation::ProjectRequest => OcrHostResult::Request(Err(
Error::InvalidRequest("OCR host has no request projection".into()),
)),
OcrHostOperation::Lifecycle(_)
| OcrHostOperation::ConstructResponse(_)
| OcrHostOperation::MapFailure(_)
| OcrHostOperation::Success { .. }
| OcrHostOperation::Failure { .. } => OcrHostResult::Lifecycle(Ok(())),
OcrHostOperation::AcquireAzureAdToken => {
OcrHostResult::AzureAdToken(Err(AuthError::AzureTokenAcquisition(
"OCR host has no Azure AD token provider".into(),
)))
}
OcrHostOperation::PreCall(request) => OcrHostResult::PreCall(Ok(request)),
OcrHostOperation::DuringCall(request) => OcrHostResult::DuringCall(Ok(request)),
OcrHostOperation::PostCall(request) => OcrHostResult::PostCall(Ok(request)),
}
})
}
}
pub struct OcrHookHost {
hooks: Arc<dyn OcrHooks>,
}
impl OcrHookHost {
pub fn new(hooks: Arc<dyn OcrHooks>) -> Self {
Self { hooks }
}
}
impl OcrHost for OcrHookHost {
fn invoke(&self, operation: OcrHostOperation) -> OcrHostFuture<'_> {
Box::pin(async move {
match operation {
OcrHostOperation::ProjectRequest => OcrHostResult::Request(Err(
Error::InvalidRequest("OCR hook host has no request projection".into()),
)),
OcrHostOperation::Success {
context,
response,
timing,
} => {
self.hooks.success(&context, &response, &timing).await;
OcrHostResult::Lifecycle(Ok(()))
}
OcrHostOperation::Failure {
context,
error,
timing,
} => {
self.hooks.failure(&context, &error, &timing).await;
OcrHostResult::Lifecycle(Ok(()))
}
OcrHostOperation::Lifecycle(_)
| OcrHostOperation::ConstructResponse(_)
| OcrHostOperation::MapFailure(_) => OcrHostResult::Lifecycle(Ok(())),
OcrHostOperation::AcquireAzureAdToken => {
OcrHostResult::AzureAdToken(Err(AuthError::AzureTokenAcquisition(
"OCR hook host has no Azure AD token provider".into(),
)))
}
OcrHostOperation::PreCall(request) => {
OcrHostResult::PreCall(self.hooks.pre_call(request).await)
}
OcrHostOperation::DuringCall(request) => {
OcrHostResult::DuringCall(self.hooks.during_call(request).await)
}
OcrHostOperation::PostCall(request) => {
OcrHostResult::PostCall(self.hooks.post_call(request).await)
}
}
})
}
}

View file

@ -5,12 +5,18 @@ mod document;
pub mod error;
mod handler;
pub mod hooks;
mod lifecycle;
mod prepare;
mod registry;
pub mod types;
pub mod wire;
pub use client::{OcrClient, ocr};
pub use document::encode_file_document;
pub use lifecycle::{
NativeOutcome, NativeResult, NoopOcrHost, OcrAdmission, OcrCall, OcrCallStep, OcrDecline,
OcrHookHost, OcrHost, OcrHostOperation, OcrHostResult,
};
pub use types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument};
#[cfg(test)]

View file

@ -62,34 +62,48 @@ pub(crate) async fn transform_request_body<B>(
request: &LiteLLMOcrRequest,
url: &str,
headers: &[(String, String)],
retains_document: bool,
body: B,
validate: impl FnOnce(&B) -> Result<(), OcrRequestError>,
) -> Result<reqwest::Request, OcrError>
where
B: Serialize + DeserializeOwned,
{
let body = if request.hooks.has_guardrails() {
let (body, headers) = if request.hooks.has_guardrails() {
let body = serde_json::to_value(body).map_err(|_| OcrRequestError::RequestField {
path: "body".into(),
})?;
let retained_fields = request
.optional_params
.keys()
.filter(|name| body.get(*name).is_some())
.cloned()
.chain(retains_document.then(|| "document".to_string()))
.collect();
let changed = request
.hooks
.during_call(OcrDuringCallRequest {
model: request.model.clone(),
custom_llm_provider: request.adapter.provider().as_str().into(),
url: url.into(),
body: serde_json::to_value(body).map_err(|_| OcrRequestError::RequestField {
path: "body".into(),
})?,
headers: headers.to_vec(),
body,
retained_fields,
})
.await?;
let body = OcrWireBody::<B>::decode(changed.body)?;
validate(&body.body)?;
body
(body, changed.headers)
} else {
OcrWireBody {
body,
extra: Map::new(),
}
(
OcrWireBody {
body,
extra: Map::new(),
},
headers.to_vec(),
)
};
build_http_request(client, request, url, headers, &body)
build_http_request(client, request, url, &headers, &body)
}
pub(crate) fn build_http_request<B: Serialize>(
@ -113,9 +127,10 @@ pub(crate) fn build_http_request<B: Serialize>(
pub(crate) async fn guardrail_document(
request: &LiteLLMOcrRequest,
url: &str,
) -> Result<OcrDocument, OcrError> {
headers: &[(String, String)],
) -> Result<(OcrDocument, Vec<(String, String)>), OcrError> {
if !request.hooks.has_guardrails() {
return Ok(request.document.clone());
return Ok((request.document.clone(), headers.to_vec()));
}
let changed = request
.hooks
@ -123,14 +138,17 @@ pub(crate) async fn guardrail_document(
model: request.model.clone(),
custom_llm_provider: request.adapter.provider().as_str().into(),
url: url.into(),
headers: headers.to_vec(),
body: serde_json::to_value(&request.document).map_err(|_| {
OcrRequestError::RequestField {
path: "document".into(),
}
})?,
retained_fields: Vec::new(),
})
.await?;
super::wire::decode_request_value(changed.body, "guardrail.document").map_err(OcrError::from)
let document = super::wire::decode_request_value(changed.body, "guardrail.document")?;
Ok((document, changed.headers))
}
#[derive(Serialize)]

View file

@ -8,7 +8,7 @@ use serde_json::{Map, Value};
use super::hooks::{NoopOcrHooks, OcrHooks};
use super::registry::{OcrAdapterKind, resolve_wire_adapter};
use crate::Error;
use crate::auth::InputSource;
use crate::auth::{InputSource, TokenProviderHandle};
use crate::constants::OCR_HTTP_TIMEOUT_SECS;
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
@ -95,6 +95,7 @@ pub struct LiteLLMOcrRequest {
pub litellm_call_id: Option<String>,
pub optional_params: Map<String, Value>,
pub input_sources: BTreeMap<String, InputSource>,
pub azure_ad_token_provider: Option<TokenProviderHandle>,
pub(crate) adapter: OcrAdapterKind,
}
@ -115,6 +116,7 @@ impl LiteLLMOcrRequest {
litellm_call_id: None,
optional_params,
input_sources: BTreeMap::new(),
azure_ad_token_provider: None,
adapter: adapter_kind,
})
}
@ -132,6 +134,10 @@ impl LiteLLMOcrRequest {
.map(|format| format.unwrap_or_default())
}
pub fn provider_name(&self) -> &'static str {
self.adapter.provider().as_str()
}
pub fn with_host_hooks(
self,
hooks: Arc<dyn OcrHooks>,

View file

@ -13,10 +13,52 @@ use serde::{
};
use serde_json::{Map, Value};
const COMMON_OPTION_FIELDS: &[&str] = &["req_format", "extra_body"];
const MISTRAL_OPTION_FIELDS: &[&str] = &[
"pages",
"include_image_base64",
"image_limit",
"image_min_size",
"bbox_annotation_format",
"document_annotation_format",
"document_annotation_prompt",
"extract_header",
"extract_footer",
"table_format",
"confidence_scores_granularity",
"include_blocks",
"id",
];
const DEEPSEEK_OPTION_FIELDS: &[&str] =
&["stream", "temperature", "max_tokens", "top_p", "n", "stop"];
const DOCUMENT_INTELLIGENCE_OPTION_FIELDS: &[&str] = &["pages", "features"];
const REDUCTO_V3_OPTION_FIELDS: &[&str] = &["formatting", "retrieval", "settings"];
const REDUCTO_LEGACY_OPTION_FIELDS: &[&str] = &["enhance"];
const AZURE_AUTH_OPTION_FIELDS: &[&str] = &[
"azure_ad_token",
"tenant_id",
"client_id",
"client_secret",
"azure_scope",
"azure_authority_host",
"azure_credential",
"azure_federated_token_file",
"enable_azure_ad_token_refresh",
];
const VERTEX_AUTH_OPTION_FIELDS: &[&str] = &[
"vertex_credentials",
"vertex_ai_credentials",
"vertex_project",
"vertex_ai_project",
"vertex_location",
"vertex_ai_location",
];
#[derive(Debug)]
pub struct DecodedOcrResponse<T> {
pub data: T,
pub native: Option<Value>,
pub text: String,
}
#[derive(Deserialize)]
@ -39,6 +81,37 @@ pub fn is_supported_request(model: &str, custom_llm_provider: Option<&str>) -> b
super::registry::resolve_wire_adapter(model, custom_llm_provider).is_ok()
}
pub fn consumed_optional_param_names(
model: &str,
custom_llm_provider: Option<&str>,
) -> Result<Vec<&'static str>, Error> {
use super::registry::OcrAdapterKind;
let (_, adapter) = super::registry::resolve_wire_adapter(model, custom_llm_provider)?;
let provider_fields: &[&str] = match adapter {
OcrAdapterKind::Mistral | OcrAdapterKind::AzureMistral | OcrAdapterKind::VertexMistral => {
MISTRAL_OPTION_FIELDS
}
OcrAdapterKind::AzureDocumentIntelligence => DOCUMENT_INTELLIGENCE_OPTION_FIELDS,
OcrAdapterKind::ReductoV3 => REDUCTO_V3_OPTION_FIELDS,
OcrAdapterKind::ReductoLegacy => REDUCTO_LEGACY_OPTION_FIELDS,
OcrAdapterKind::VertexDeepSeek => DEEPSEEK_OPTION_FIELDS,
};
let auth_fields: &[&str] = match adapter {
OcrAdapterKind::AzureMistral | OcrAdapterKind::AzureDocumentIntelligence => {
AZURE_AUTH_OPTION_FIELDS
}
OcrAdapterKind::VertexMistral | OcrAdapterKind::VertexDeepSeek => VERTEX_AUTH_OPTION_FIELDS,
_ => &[],
};
Ok(COMMON_OPTION_FIELDS
.iter()
.chain(provider_fields)
.chain(auth_fields)
.copied()
.collect())
}
pub fn decode_request(wire: OcrWireRequest) -> Result<LiteLLMOcrRequest, Error> {
let api_key_source = source_for(&wire.input_sources, "api_key");
let api_base_source = source_for(&wire.input_sources, "api_base");
@ -134,7 +207,11 @@ pub fn decode_response<T: DeserializeOwned>(
} else {
None
};
Ok(DecodedOcrResponse { data, native })
Ok(DecodedOcrResponse {
data,
native,
text: String::from_utf8_lossy(bytes).into_owned(),
})
}
pub fn decode_pre_call_result(
@ -169,3 +246,22 @@ pub fn decode_during_call_result(
..original
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn option_projection_is_provider_specific_and_excludes_opaque_fields() {
let mistral = consumed_optional_param_names("mistral/model", None).unwrap();
assert!(mistral.contains(&"pages"));
assert!(mistral.contains(&"req_format"));
assert!(!mistral.contains(&"vertex_project"));
assert!(!mistral.contains(&"opaque_extension"));
let vertex = consumed_optional_param_names("vertex_ai/deepseek-ocr", None).unwrap();
assert!(vertex.contains(&"temperature"));
assert!(vertex.contains(&"vertex_credentials"));
assert!(!vertex.contains(&"pages"));
}
}

View file

@ -81,7 +81,7 @@ impl AzureAuthService {
AzureCredentialPlan::Caller(caller) => {
let credential = caller.acquire().await?;
if credential.secret().expose().is_empty() {
return Err(AuthError::EmptyAzureToken);
return Ok(None);
}
Ok(Some(Sourced::new(credential, InputSource::Deployment)))
}

View file

@ -1,4 +1,5 @@
use serde_json::{Value, json};
use std::sync::{Arc, Mutex};
use super::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
use super::wire::{OcrWireRequest, decode_request};
@ -169,6 +170,46 @@ async fn accepted_response_polls_to_success_with_only_credentials() {
}
}
struct SubmissionBoundary {
request_count: Arc<Mutex<Vec<String>>>,
}
impl super::hooks::OcrHooks for SubmissionBoundary {
fn post_call(
&self,
request: super::hooks::OcrPostCallRequest,
) -> super::hooks::OcrHookFuture<'_, super::hooks::OcrPostCallRequest> {
Box::pin(async move {
assert_eq!(self.request_count.lock().unwrap().len(), 1);
assert_eq!(request.original_response, json!(r#"{"submitted":true}"#));
Ok(request)
})
}
}
#[tokio::test]
async fn accepted_response_runs_post_call_before_polling() {
let (base, seen, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({"submitted": true}),
},
MockResponse::json(json!({"status":"succeeded"})),
])
.await;
let request = super::LiteLLMOcrRequest {
hooks: Arc::new(SubmissionBoundary {
request_count: seen.clone(),
}),
..wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({}))
};
perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(seen.lock().unwrap().len(), 2);
}
#[tokio::test]
async fn polling_forwards_bearer_credentials() {
let (base, seen, server) = mock_server(vec![

View file

@ -0,0 +1,116 @@
use crate::Error;
use crate::call_lifecycle::host::{HostFailure, HostLifecycle, HostPhase};
fn run(fail_at: Option<HostPhase>, asynchronous: bool) -> (Vec<HostPhase>, Vec<Error>) {
let mut lifecycle = HostLifecycle::new(asynchronous);
let mut events = Vec::new();
let mut failures = Vec::new();
while lifecycle.phase() != HostPhase::Complete {
let phase = lifecycle.phase();
events.push(phase);
let result = if Some(phase) == fail_at {
Err(HostFailure::Error(Error::InvalidRequest(
"selected failure".into(),
)))
} else {
Ok(())
};
if let Some(error) = lifecycle.accept(result) {
failures.push(error);
}
}
(events, failures)
}
#[test]
fn public_outcome_is_finalized_before_a_single_terminal_dispatch() {
for asynchronous in [false, true] {
let (events, failures) = run(None, asynchronous);
assert!(failures.is_empty());
assert_eq!(
&events[events.len() - 2..],
&[HostPhase::Finalize, HostPhase::Success]
);
assert_eq!(
events
.iter()
.filter(|phase| **phase == HostPhase::Execute)
.count(),
1
);
assert_eq!(
events.contains(&HostPhase::DeploymentPostCall),
asynchronous
);
}
}
#[test]
fn only_provider_and_response_construction_failures_use_provider_mapping() {
for phase in [
HostPhase::Setup,
HostPhase::DeploymentPreCall,
HostPhase::Prepare,
HostPhase::Execute,
HostPhase::ConstructResponse,
HostPhase::DeploymentPostCall,
HostPhase::Finalize,
] {
let (events, failures) = run(Some(phase), true);
assert_eq!(failures.len(), 1);
assert!(!events.contains(&HostPhase::Success));
let mapped = matches!(phase, HostPhase::Execute | HostPhase::ConstructResponse);
assert_eq!(events.contains(&HostPhase::MapFailure), mapped);
assert_eq!(events.contains(&HostPhase::DeploymentFailure), mapped);
assert_eq!(
&events[events.len() - 2..],
&[HostPhase::Failure, HostPhase::AsyncFailure]
);
assert!(
events
.iter()
.filter(|phase| **phase == HostPhase::Execute)
.count()
<= 1
);
}
}
#[test]
fn failure_handler_errors_do_not_replace_selected_failure_or_suppress_async_dispatch() {
let mut lifecycle = HostLifecycle::new(true);
while lifecycle.phase() != HostPhase::Execute {
lifecycle.accept(Ok(()));
}
let selected = Error::InvalidRequest("provider".into());
assert_eq!(
lifecycle.accept(Err(HostFailure::Error(selected.clone()))),
Some(selected)
);
lifecycle.accept(Ok(()));
for phase in [
HostPhase::DeploymentFailure,
HostPhase::Failure,
HostPhase::AsyncFailure,
] {
assert_eq!(lifecycle.phase(), phase);
assert_eq!(
lifecycle.accept(Err(HostFailure::Error(Error::InvalidRequest(
"callback".into()
)))),
None
);
}
assert_eq!(lifecycle.phase(), HostPhase::Complete);
}
#[test]
fn cancellation_skips_terminal_dispatch() {
let mut lifecycle = HostLifecycle::new(true);
let error = Error::InvalidRequest("cancelled".into());
assert_eq!(
lifecycle.accept(Err(HostFailure::Cancelled(error.clone()))),
Some(error)
);
assert_eq!(lifecycle.phase(), HostPhase::Complete);
}

View file

@ -3,9 +3,16 @@ use std::sync::{Arc, Mutex};
use serde_json::{Value, json};
use super::OcrClient;
use super::hooks::{OcrHookFuture, OcrHooks, OcrLogFuture, OcrPreCallRequest};
use super::hooks::{
OcrDuringCallRequest, OcrHookFuture, OcrHooks, OcrLogFuture, OcrPostCallRequest,
OcrPreCallRequest,
};
use super::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
use super::wire::{OcrWireRequest, decode_request};
use super::{
NativeOutcome, NoopOcrHost, OcrAdmission, OcrCall, OcrCallStep, OcrDecline, OcrHost,
OcrHostOperation, OcrHostResult,
};
use crate::call_lifecycle::{CallLifecycleContext, CallLifecycleTiming};
#[test]
@ -148,6 +155,13 @@ impl OcrHooks for RecordingHooks {
})
}
fn post_call(&self, request: OcrPostCallRequest) -> OcrHookFuture<'_, OcrPostCallRequest> {
Box::pin(async move {
self.events.lock().unwrap().push("post");
Ok(request)
})
}
fn success<'a>(
&'a self,
_context: &'a CallLifecycleContext,
@ -171,6 +185,38 @@ impl OcrHooks for RecordingHooks {
}
}
struct HeaderEditHooks;
impl OcrHooks for HeaderEditHooks {
fn has_guardrails(&self) -> bool {
true
}
fn during_call(
&self,
mut request: OcrDuringCallRequest,
) -> OcrHookFuture<'_, OcrDuringCallRequest> {
request
.headers
.push(("x-core-callback".into(), "edited".into()));
Box::pin(async move { Ok(request) })
}
}
#[tokio::test]
async fn lifecycle_sends_headers_returned_by_the_typed_during_call_operation() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let request = super::LiteLLMOcrRequest {
hooks: Arc::new(HeaderEditHooks),
..wire_request("mistral/model", &base, json!({}))
};
perform_ocr(request).await.unwrap();
server.await.unwrap();
assert!(seen.lock().unwrap()[0].contains("x-core-callback: edited"));
}
#[tokio::test]
async fn lifecycle_orders_hooks_and_emits_one_success() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
@ -185,7 +231,10 @@ async fn lifecycle_orders_hooks_and_emits_one_success() {
};
perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(*events.lock().unwrap(), ["pre", "during", "success"]);
assert_eq!(
*events.lock().unwrap(),
["pre", "during", "post", "success"]
);
assert_eq!(seen.lock().unwrap().len(), 1);
}
@ -227,3 +276,351 @@ async fn upstream_failure_emits_one_terminal_failure() {
assert_eq!(*events.lock().unwrap(), ["pre", "during", "failure"]);
assert_eq!(seen.lock().unwrap().len(), 1);
}
struct AdmissionSpy {
effects: Arc<Mutex<usize>>,
}
impl OcrHooks for AdmissionSpy {
fn has_guardrails(&self) -> bool {
*self.effects.lock().unwrap() += 1;
true
}
fn pre_call(&self, request: OcrPreCallRequest) -> OcrHookFuture<'_, OcrPreCallRequest> {
*self.effects.lock().unwrap() += 1;
Box::pin(async move { Ok(request) })
}
}
#[test]
fn admission_declines_without_invoking_hooks_or_transport() {
for (admission, expected) in [
(
OcrAdmission {
provider_workflow: false,
host_operations: true,
azure_ad_token_provider: false,
asynchronous: false,
},
OcrDecline::ProviderWorkflow,
),
(
OcrAdmission {
provider_workflow: true,
host_operations: false,
azure_ad_token_provider: false,
asynchronous: false,
},
OcrDecline::HostOperations,
),
] {
let outcome = OcrCall::admit(super::test_support::ocr_client(), admission);
assert!(matches!(outcome, NativeOutcome::Declined(reason) if reason == expected));
}
}
#[tokio::test]
async fn fallible_host_phases_do_not_replay_or_reach_transport() {
for failure_phase in ["pre", "during"] {
let request = super::LiteLLMOcrRequest {
hooks: Arc::new(AdmissionSpy {
effects: Arc::new(Mutex::new(0)),
}),
..wire_request("mistral/model", "http://127.0.0.1:1", json!({}))
};
let NativeOutcome::Completed(mut call) =
OcrCall::admit(super::test_support::ocr_client(), OcrAdmission::all())
else {
panic!("supported call declined")
};
let mut request = Some(request);
let mut result = None;
let mut phases = Vec::new();
let error = loop {
match call.resume(result.take()).await {
Ok(OcrCallStep::Host(operation)) => match operation {
OcrHostOperation::Lifecycle(_)
| OcrHostOperation::ConstructResponse(_)
| OcrHostOperation::MapFailure(_)
| OcrHostOperation::Success { .. }
| OcrHostOperation::Failure { .. } => {
result = Some(OcrHostResult::Lifecycle(Ok(())))
}
OcrHostOperation::ProjectRequest => {
result = Some(OcrHostResult::Request(Ok((
Box::new(request.take().unwrap()),
false,
))))
}
OcrHostOperation::AcquireAzureAdToken => {
panic!("test request has no token provider")
}
OcrHostOperation::PreCall(request) => {
phases.push("pre");
result = Some(OcrHostResult::PreCall(if failure_phase == "pre" {
Err(crate::Error::InvalidRequest("pre failed".into()))
} else {
Ok(request)
}));
}
OcrHostOperation::DuringCall(request) => {
phases.push("during");
result = Some(OcrHostResult::DuringCall(if failure_phase == "during" {
Err(crate::Error::InvalidRequest("during failed".into()))
} else {
Ok(request)
}));
}
OcrHostOperation::PostCall(_) => panic!("transport should not be reached"),
},
Err(error) => break error,
Ok(OcrCallStep::Complete(_)) => panic!("failed call completed"),
}
};
assert!(matches!(error, crate::Error::InvalidRequest(_)));
assert_eq!(
phases
.iter()
.filter(|phase| **phase == failure_phase)
.count(),
1
);
}
}
#[tokio::test]
async fn invalid_provider_response_runs_post_call_before_normalization_failure() {
let (base, seen, server) =
mock_server(vec![MockResponse::json(json!({"pages":"invalid"}))]).await;
let mut request = Some(wire_request("mistral/model", &base, json!({})));
let NativeOutcome::Completed(mut call) =
OcrCall::admit(super::test_support::ocr_client(), OcrAdmission::all())
else {
panic!("supported call declined")
};
let host = NoopOcrHost;
let mut result = None;
let mut post_calls = Vec::new();
let error = loop {
match call.resume(result.take()).await {
Ok(OcrCallStep::Host(OcrHostOperation::ProjectRequest)) => {
result = Some(OcrHostResult::Request(Ok((
Box::new(request.take().unwrap()),
false,
))));
}
Ok(OcrCallStep::Host(operation)) => {
if let OcrHostOperation::PostCall(request) = &operation {
post_calls.push(request.original_response.clone());
}
result = Some(host.invoke(operation).await);
}
Err(error) => break error,
Ok(OcrCallStep::Complete(_)) => panic!("invalid provider response completed"),
}
};
server.await.unwrap();
assert!(matches!(error, crate::Error::InvalidResponse(_)));
assert_eq!(seen.lock().unwrap().len(), 1);
assert_eq!(post_calls, [json!(r#"{"pages":"invalid"}"#)]);
}
#[tokio::test]
async fn direct_native_host_drives_the_same_state_machine() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"pages":[{"index":0,"markdown":"native"}]
}))])
.await;
let request = super::LiteLLMOcrRequest {
hooks: Arc::new(AdmissionSpy {
effects: Arc::new(Mutex::new(0)),
}),
..wire_request("mistral/model", &base, json!({}))
};
let NativeOutcome::Completed(mut call) = OcrCall::admit(
super::test_support::ocr_client(),
OcrAdmission {
asynchronous: true,
..OcrAdmission::all()
},
) else {
panic!("supported call declined")
};
let mut request = Some(request);
let host = NoopOcrHost;
let mut result = None;
let mut operations = Vec::new();
let response = loop {
match call.resume(result.take()).await.unwrap() {
OcrCallStep::Host(operation) => {
operations.push(match &operation {
OcrHostOperation::ProjectRequest => "ProjectRequest".into(),
OcrHostOperation::Lifecycle(phase) => format!("{phase:?}"),
OcrHostOperation::PreCall(_) => "PreCall".into(),
OcrHostOperation::DuringCall(_) => "DuringCall".into(),
OcrHostOperation::PostCall(_) => "PostCall".into(),
OcrHostOperation::ConstructResponse(_) => "ConstructResponse".into(),
OcrHostOperation::Success { response, .. } => {
assert_eq!(response.pages[0]["markdown"], "native");
"Success".into()
}
_ => panic!("unexpected OCR operation"),
});
result = Some(match operation {
OcrHostOperation::ProjectRequest => {
OcrHostResult::Request(Ok((Box::new(request.take().unwrap()), false)))
}
operation => host.invoke(operation).await,
});
}
OcrCallStep::Complete(response) => break response,
}
};
server.await.unwrap();
assert_eq!(response.pages[0]["markdown"], "native");
assert_eq!(seen.lock().unwrap().len(), 1);
assert_eq!(
operations,
[
"Setup",
"DeploymentPreCall",
"Prepare",
"ProjectRequest",
"PreCall",
"DuringCall",
"PostCall",
"ConstructResponse",
"DeploymentPostCall",
"Finalize",
"Success",
]
);
assert!(matches!(
call.resume(None).await,
Err(crate::Error::InvalidRequest(_))
));
}
#[tokio::test]
async fn public_finalization_failure_never_dispatches_success_or_replays_provider() {
use crate::call_lifecycle::host::{HostFailure, HostPhase};
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let mut request = Some(wire_request("mistral/model", &base, json!({})));
let NativeOutcome::Completed(mut call) = OcrCall::admit(
super::test_support::ocr_client(),
OcrAdmission {
asynchronous: true,
..OcrAdmission::all()
},
) else {
panic!("supported call declined")
};
let selected = crate::Error::InvalidRequest("public metadata failed".into());
let host = NoopOcrHost;
let mut result = None;
let mut failures = Vec::new();
let error = loop {
match call.resume(result.take()).await {
Ok(OcrCallStep::Host(operation)) => {
result = Some(match operation {
OcrHostOperation::Lifecycle(HostPhase::Finalize) => {
OcrHostResult::Lifecycle(Err(HostFailure::Error(selected.clone())))
}
OcrHostOperation::Failure { error, .. } => {
assert_eq!(error, selected);
failures.push("sync");
OcrHostResult::Lifecycle(Err(HostFailure::Error(
crate::Error::InvalidRequest("failure callback failed".into()),
)))
}
OcrHostOperation::Lifecycle(HostPhase::AsyncFailure) => {
failures.push("async");
OcrHostResult::Lifecycle(Ok(()))
}
OcrHostOperation::Success { .. }
| OcrHostOperation::MapFailure(_)
| OcrHostOperation::Lifecycle(HostPhase::DeploymentFailure) => {
panic!("finalization failure used provider/success dispatch")
}
OcrHostOperation::ProjectRequest => {
OcrHostResult::Request(Ok((Box::new(request.take().unwrap()), false)))
}
operation => host.invoke(operation).await,
});
}
Ok(OcrCallStep::Complete(_)) => panic!("failed call completed successfully"),
Err(error) => break error,
}
};
server.await.unwrap();
assert_eq!(error, selected);
assert_eq!(failures, ["sync", "async"]);
assert_eq!(seen.lock().unwrap().len(), 1);
}
#[tokio::test]
async fn cancellation_at_provider_hook_prevents_execution_and_further_resumption() {
use crate::call_lifecycle::host::HostFailure;
let request = super::LiteLLMOcrRequest {
hooks: Arc::new(AdmissionSpy {
effects: Arc::new(Mutex::new(0)),
}),
..wire_request("mistral/model", "http://127.0.0.1:1", json!({}))
};
let NativeOutcome::Completed(mut call) =
OcrCall::admit(super::test_support::ocr_client(), OcrAdmission::all())
else {
panic!("supported call declined")
};
let mut request = Some(request);
let host = NoopOcrHost;
let mut result = None;
loop {
match call.resume(result.take()).await.unwrap() {
OcrCallStep::Host(OcrHostOperation::PreCall(_)) => break,
OcrCallStep::Host(OcrHostOperation::ProjectRequest) => {
result = Some(OcrHostResult::Request(Ok((
Box::new(request.take().unwrap()),
false,
))))
}
OcrCallStep::Host(operation) => result = Some(host.invoke(operation).await),
OcrCallStep::Complete(_) => panic!("provider executed before pre-call result"),
}
}
let selected = crate::Error::InvalidRequest("cancelled".into());
assert!(matches!(
call.interrupt(HostFailure::Cancelled(selected.clone())).await,
Err(error) if error == selected
));
assert!(
call.resume(Some(OcrHostResult::Lifecycle(Ok(()))))
.await
.is_err()
);
}
#[tokio::test]
async fn missing_host_result_preserves_pending_operation() {
use crate::call_lifecycle::host::HostPhase;
let NativeOutcome::Completed(mut call) =
OcrCall::admit(super::test_support::ocr_client(), OcrAdmission::all())
else {
panic!("supported call declined")
};
assert!(matches!(
call.resume(None).await.unwrap(),
OcrCallStep::Host(OcrHostOperation::Lifecycle(HostPhase::Setup))
));
assert!(call.resume(None).await.is_err());
assert!(matches!(
call.resume(Some(OcrHostResult::Lifecycle(Ok(()))))
.await
.unwrap(),
OcrCallStep::Host(OcrHostOperation::Lifecycle(HostPhase::Prepare))
));
}

View file

@ -3,7 +3,7 @@ use std::sync::Arc;
use rstest::rstest;
use serde_json::{Value, json};
use super::hooks::{OcrDuringCallRequest, OcrHookFuture, OcrHooks};
use super::hooks::{OcrDuringCallRequest, OcrHookFuture, OcrHooks, OcrPostCallRequest};
use super::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
fn request_body(request: &str) -> Value {
@ -100,6 +100,42 @@ async fn data_uri_upload_preserves_multipart_headers(#[case] model: &str) {
assert!(requests[1].starts_with("POST /parse "));
}
struct ParseBoundary {
request_count: Arc<std::sync::Mutex<Vec<String>>>,
}
impl OcrHooks for ParseBoundary {
fn post_call(&self, request: OcrPostCallRequest) -> OcrHookFuture<'_, OcrPostCallRequest> {
Box::pin(async move {
assert_eq!(self.request_count.lock().unwrap().len(), 2);
assert_eq!(
request.original_response,
json!(r#"{"result":{"chunks":[]}}"#)
);
Ok(request)
})
}
}
#[tokio::test]
async fn post_call_stays_after_reducto_upload_and_parse() {
let (base, seen, server) = mock_server(vec![
MockResponse::json(json!({"file_id":"reducto://uploaded.pdf"})),
MockResponse::json(json!({"result":{"chunks":[]}})),
])
.await;
let request = super::LiteLLMOcrRequest {
hooks: Arc::new(ParseBoundary {
request_count: seen.clone(),
}),
..wire_request("reducto/parse-v3", &base, json!({}))
};
perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(seen.lock().unwrap().len(), 2);
}
#[rstest]
#[case(json!({"file_id":""}))]
#[case(json!({}))]

View file

@ -1,3 +1,42 @@
litellm-python-bridge is the PyO3 cdylib that exposes LiteLLM Rust APIs to the Python SDK. Keep API registration, domain dependency wiring, request assembly, and Python exception mapping here. Put domain-neutral Python/Serde conversion and GIL primitives in litellm-python-interop.
Keep it thin: no business logic, no transforms, no I/O orchestration — just marshal in/out and call the core entrypoint.
- Target invariants, not completion claims; these supersede older conflicting bridge guidance
- Keep this crate the product-specific PyO3 consumer of `litellm-python-interop`
- Own registration, input projection, retained Python state, callback invocation, public response/error construction and host scheduling
- Keep value-oriented execution, sync waiting, nested-runtime checks, signal polling and panic containment in `execution.rs`; native async work uses `pyo3-async-runtimes`, Serde output uses `Pythonized<T>`
- Core owns typed native state, admission, lifecycle sequencing, provider preparation/I/O, normalization and terminal-outcome/dispatch decisions
- Python, Rust SDK and gateway use one lifecycle-bearing core route entrypoint; provider helpers stay private, never bridge-accessible transport drivers
- Built-in provider/config/secret/auth/document preparation stays in Rust; caller-authored callbacks and focused Python-file reads run only at core-selected points
- Target GIL-enabled CPython explicitly with `#[pymodule(gil_used = true)]`; detach Rust-only work
- Free-threading requires separate runtime/concurrency validation; omitting the attribute does not opt out on PyO3 0.28+
- Preserve public argument binding and Python object provenance
- Retain complete boundary arguments, opaque unknown values, aliases, omitted/default distinctions and deliberate copies; preserve the established deployment-hook kwargs view
- Retain independently captured body/header roots; in-place mutation and logging-envelope field replacement have different effects
- Project only consumed fields at reference read points; no eager whole-graph serialization or equality-based alias reconstruction
- Preserve provider-specific upload/submission/poll observation and encoding boundaries; signed/build-captured bytes must not be silently reserialized
- Only core's typed, effect-free admission may return `Declined`; conversion errors and all post-admission failures are terminal
- Admission cannot invoke hooks, acquire credentials, consume files/iterators, prepare requests or perform I/O
- Disabled/unavailable native execution or an admission decline may select legacy once; callback exceptions never authorize fallback or replay
- Use one ordinary inline `async def` driver in `litellm/rust_bridge/lifecycle.py`, with the native handle in `src/lifecycle.rs`
- Contract: `start`, `resume_value`, `resume_error`, idempotent `close`; explicitly tagged `Await`/`Complete` preserve awaitable final values
- Validate Created/Running/Suspended/Closed protocol states; core alone chooses lifecycle phases and result/error policy
- Defer effectful setup/context reads/timestamps until start; unstarted-handle destruction releases inputs independently of Python `finally`
- Catch only the selected await's errors; start/resume errors propagate, `GeneratorExit` closes without further awaits
- Inline hooks preserve caller task/thread/loop and context writes; `into_future` creates a separate task and cannot satisfy this contract
- Delivery follows the binding, not callable type; keep direct, awaited, worker, background and deferred behavior distinct
- Finalize fallible public response/error construction, replacements and metadata under core control before terminal dispatch
- Success/failure handler entry receives the exact selected public response/exception; logging projections/redaction/snapshots retain their own copy contracts
- Ordinary failure-callback errors cannot suppress later eligible sync/async callbacks or replace the mapped provider error; control-flow exceptions have phase-specific policy
- Dispatch errors never replay provider work/accepted dispatch or trigger the opposite outcome; proxy acceptance/rejection releases core-owned deferred success at most once
- Make ownership safe across suspension, re-entry, cancellation and GC
- Keep native provider state typed in core; do not shuttle it through opaque Python transport/response classes
- Prefer one retained `Py<PyBaseException>` via `PyErr::into_value(py)`; reconstruct transient `PyErr`s, preserving identity, traceback, cause and context
- Traverse every owned Python edge, including duplicate references; traversal cannot call Python
- Take state out and mark Running under a short borrow, release borrows/locks before Python invocation, publish terminal state before finalizer-capable drops
- Close/GC/deferred release are idempotent and re-entry-safe, including during Rust unwinding; release only owned references, never clear caller containers or mask the selected error
- Cancellation signaling is not termination; retain captures until work actually finishes and use a Rust-selected awaited acknowledgement where required, never synchronous close/GC
- Verify behavior through a fresh, provenance-checked installed extension and positive native execution evidence before replacing the custom coroutine
- Cover admitted provider workflows, binding/read-point/identity behavior, failure continuation, finalization, no replay, deferred gates, re-entry, GC and cancellation termination
- Measure real conversion/copy costs before optimizing; preserve input contracts and capture lifetimes with `PyBackedBytes`, and lookup timing when interning names
- Ship accurate `_native.pyi` declarations and typing markers; distinguish Future-returning bindings from coroutine-returning bindings
- References: [ownership](https://pyo3.rs/v0.29.2/types.html), [GC](https://pyo3.rs/v0.29.2/class/protocols.html#garbage-collector-integration), [exception transfer](https://docs.rs/pyo3/0.29.2/pyo3/struct.PyErr.html#method.into_value), [re-entry](https://pyo3.rs/v0.29.2/class/call.html)
- [GIL policy](https://pyo3.rs/v0.29.2/free-threading.html), [experimental async limits](https://pyo3.rs/v0.29.2/async-await.html), [task conversion](https://docs.rs/pyo3-async-runtimes/0.29.0/pyo3_async_runtimes/fn.into_future_with_locals.html), [native cancellation/delivery](https://docs.rs/pyo3-async-runtimes/0.29.0/pyo3_async_runtimes/tokio/fn.future_into_py.html)
- [performance](https://pyo3.rs/v0.29.2/performance.html), [PyBackedBytes](https://docs.rs/pyo3/0.29.2/pyo3/pybacked/struct.PyBackedBytes.html), [typing](https://pyo3.rs/v0.29.2/python-typing-hints.html)

View file

@ -69,11 +69,24 @@ pub(crate) fn ocr_error_to_pyerr(err: Error) -> PyErr {
Error::MissingField("document_url" | "image_url") => {
PyValueError::new_err("Document URL is required")
}
Error::Http { status, body } => RustUpstreamError::new_err((status, body)),
Error::Http { status, body } => ocr_upstream_error(status, body),
Error::Network(message) if message.contains("timed out") => {
ocr_upstream_error(408, message)
}
other => core_error_to_pyerr(other),
}
}
fn ocr_upstream_error(status: u16, message: String) -> PyErr {
let error = RustUpstreamError::new_err((status, message.clone()));
Python::attach(|py| {
let value = error.value(py);
value.setattr("status_code", status).ok();
value.setattr("message", message).ok();
});
error
}
#[cfg(test)]
mod ocr_error_tests {
use super::*;

View file

@ -28,6 +28,27 @@ where
)
}
pub(crate) fn run_sync_value<T, F>(py: Python<'_>, future: F) -> PyResult<T>
where
T: Send + 'static,
F: Future<Output = PyResult<T>> + Send + 'static,
{
run_sync_value_on(py, pyo3_async_runtimes::tokio::get_runtime(), future)
}
fn run_sync_value_on<T, F>(py: Python<'_>, runtime: &Runtime, future: F) -> PyResult<T>
where
T: Send + 'static,
F: Future<Output = PyResult<T>> + Send + 'static,
{
if Handle::try_current().is_ok() {
return Err(PyRuntimeError::new_err(
"synchronous native routes cannot run from a Tokio context; use the async route",
));
}
release_gil(py, move || runtime.block_on(wait_for_sync_result(future)))?
}
fn run_sync_on<T, E, F>(
py: Python<'_>,
runtime: &Runtime,
@ -39,14 +60,9 @@ where
E: Send + 'static,
F: Future<Output = Result<T, E>> + Send + 'static,
{
if Handle::try_current().is_ok() {
return Err(PyRuntimeError::new_err(
"synchronous native routes cannot run from a Tokio context; use the async route",
));
}
let result = release_gil(py, move || runtime.block_on(wait_for_sync_result(future)))?;
let result = map_core_result(result, map_error)?;
let result = run_sync_value_on(py, runtime, async move {
map_core_result(future.await, map_error)
})?;
Pythonized(result).into_pyobject(py).map(Bound::unbind)
}
@ -60,13 +76,20 @@ where
E: Send + 'static,
F: Future<Output = Result<T, E>> + Send + 'static,
{
pyo3_async_runtimes::tokio::future_into_py(py, async move {
let result = catch_future_panic(future).await?;
let result = map_core_result(result, map_error)?;
run_async_value(py, async move {
let result = map_core_result(future.await, map_error)?;
Ok(Pythonized(result))
})
}
pub(crate) fn run_async_value<T, F>(py: Python<'_>, future: F) -> PyResult<Bound<'_, PyAny>>
where
T: for<'py> IntoPyObject<'py> + Send + 'static,
F: Future<Output = PyResult<T>> + Send + 'static,
{
pyo3_async_runtimes::tokio::future_into_py(py, async move { catch_future_panic(future).await? })
}
fn map_core_result<T, E>(result: Result<T, E>, map_error: fn(E) -> PyErr) -> PyResult<T> {
match result {
Ok(value) => Ok(value),

View file

@ -4,6 +4,7 @@ mod errors;
mod execution;
#[cfg(feature = "trace-parity")]
mod function_trace;
mod lifecycle;
mod marshal;
mod routes;
mod token_counter;
@ -64,7 +65,7 @@ impl ResponsesWebSocketConnection {
}
}
#[pymodule(gil_used = false)]
#[pymodule(gil_used = true)]
mod _native {
use pyo3::prelude::*;

File diff suppressed because it is too large Load diff

View file

@ -10,12 +10,15 @@ mod audio_transcription;
mod chat_completions;
mod messages;
mod ocr;
mod ocr_lifecycle;
pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
ocr::register(module)?;
ocr_lifecycle::register(module)?;
audio_transcription::register(module)?;
messages::register(module)?;
chat_completions::register(module)?;
#[cfg(feature = "trace-parity")]
{
let trace = PyModule::new(module.py(), "_trace")?;

View file

@ -0,0 +1,478 @@
use serde_json::{Map, Value};
use std::sync::Arc;
use pyo3::prelude::*;
use pyo3::types::{PyBytes, PyDict, PyTuple};
use litellm_core::auth::{ResolvedCredential, SecretValue};
use litellm_core::ocr::hooks::{OcrDuringCallRequest, OcrPostCallRequest, OcrPreCallRequest};
use litellm_core::ocr::wire::{OcrWireRequest, consumed_optional_param_names, decode_request};
use litellm_core::ocr::{
NativeOutcome, OcrAdmission, OcrCall, OcrClient, OcrHostOperation, OcrHostResult,
};
use litellm_python_interop::{from_py, to_py};
use crate::errors::{RustBridgeDeclined, ocr_error_to_pyerr};
use crate::lifecycle::{PythonCallState, PythonRoute, missing_state, now, run_call};
struct PythonOcrHost {
state: PythonCallState,
request: Option<Py<PyAny>>,
pre_call: Option<OcrPreCallRequest>,
document: Option<Py<PyAny>>,
api_key: Option<Py<PyAny>>,
azure_ad_token_provider: Option<Py<PyAny>>,
provider: String,
retained_fields: Option<Py<PyDict>>,
body: Option<Py<PyDict>>,
headers: Option<Py<PyDict>>,
}
struct AdmittedOcrCall {
request: litellm_core::ocr::LiteLLMOcrRequest,
document: Py<PyAny>,
api_key: Py<PyAny>,
azure_ad_token_provider: Option<Py<PyAny>>,
provider: String,
}
impl PythonOcrHost {
fn pre_call(
&mut self,
py: Python<'_>,
request: OcrPreCallRequest,
) -> PyResult<OcrPreCallRequest> {
let kwargs = self.state.kwargs.bind(py);
let retained_fields = PyDict::new(py);
for name in request
.optional_params
.as_object()
.ok_or_else(missing_state)?
.keys()
{
if let Some(value) = kwargs.get_item(name)? {
retained_fields.set_item(name, value)?;
}
}
retained_fields.set_item(
"document",
self.document.as_ref().ok_or_else(missing_state)?,
)?;
self.retained_fields = Some(retained_fields.unbind());
self.pre_call = Some(request.clone());
Ok(request)
}
fn acquire_azure_ad_token(&self, py: Python<'_>) -> PyResult<ResolvedCredential> {
let provider = self
.azure_ad_token_provider
.as_ref()
.ok_or_else(missing_state)?;
let token: String = py
.import("litellm.rust_bridge.ocr_lifecycle")?
.getattr("call_azure_ad_token_provider")?
.call1((provider,))?
.extract()?;
Ok(ResolvedCredential::AccessToken {
token: SecretValue::new(token),
expires_on: None,
})
}
fn python_pre_call(
&mut self,
py: Python<'_>,
request: OcrDuringCallRequest,
) -> PyResult<OcrDuringCallRequest> {
let pre_call = self.pre_call.as_ref().ok_or_else(missing_state)?;
let body = to_py(py, &request.body)?
.into_bound(py)
.cast_into::<PyDict>()?;
if let Some(retained) = &self.retained_fields {
for name in &request.retained_fields {
if let Some(value) = retained.bind(py).get_item(name)? {
body.set_item(name, value)?;
}
}
}
let headers = PyDict::new(py);
for (name, value) in &request.headers {
headers.set_item(name, value)?;
}
self.body = Some(body.clone().unbind());
self.headers = Some(headers.clone().unbind());
let logger = self.state.logger(py)?;
let redact = py
.import("litellm.rust_bridge.ocr")?
.getattr("redact_logging_params")?;
let update = PyDict::new(py);
update.set_item("kwargs", redact.call1((&self.state.kwargs,))?)?;
update.set_item("model", &pre_call.model)?;
update.set_item(
"optional_params",
redact.call1((to_py(py, &pre_call.optional_params)?,))?,
)?;
let params = PyDict::new(py);
params.set_item(
"litellm_call_id",
self.state.kwargs.bind(py).get_item("litellm_call_id")?,
)?;
params.set_item("api_base", &request.url)?;
update.set_item("litellm_params", params)?;
update.set_item("custom_llm_provider", &pre_call.custom_llm_provider)?;
logger.call_method("update_from_kwargs", (), Some(&update))?;
let additional = PyDict::new(py);
additional.set_item("complete_input_dict", &body)?;
additional.set_item("headers", &headers)?;
additional.set_item("api_base", &request.url)?;
let kwargs = PyDict::new(py);
kwargs.set_item("input", "OCR document processing")?;
kwargs.set_item("api_key", &self.api_key)?;
kwargs.set_item("additional_args", additional)?;
logger.call_method("pre_call", (), Some(&kwargs))?;
let headers = headers
.iter()
.map(|(name, value)| Ok((name.extract::<String>()?, value.extract::<String>()?)))
.collect::<PyResult<Vec<_>>>()?;
Ok(OcrDuringCallRequest {
body: from_py(&body)?,
headers,
..request
})
}
fn python_post_call(
&mut self,
py: Python<'_>,
request: OcrPostCallRequest,
) -> PyResult<OcrPostCallRequest> {
let logger = self.state.logger(py)?;
let kwargs = PyDict::new(py);
kwargs.set_item("original_response", to_py(py, &request.original_response)?)?;
let additional = PyDict::new(py);
additional.set_item("complete_input_dict", &self.body)?;
additional.set_item("headers", &self.headers)?;
kwargs.set_item("additional_args", additional)?;
logger.call_method("post_call", (), Some(&kwargs))?;
Ok(request)
}
}
impl PythonRoute for PythonOcrHost {
fn state(&self) -> &PythonCallState {
&self.state
}
fn state_mut(&mut self) -> &mut PythonCallState {
&mut self.state
}
fn invoke(&mut self, py: Python<'_>, operation: OcrHostOperation) -> PyResult<OcrHostResult> {
Ok(match operation {
OcrHostOperation::ProjectRequest => {
let projected = project_request(
py,
self.request.as_ref().ok_or_else(missing_state)?.bind(py),
self.state.kwargs.bind(py),
)?;
self.document = Some(projected.document);
self.api_key = Some(projected.api_key);
self.azure_ad_token_provider = projected.azure_ad_token_provider;
self.provider = projected.provider;
OcrHostResult::Request(Ok((
Box::new(projected.request),
self.azure_ad_token_provider.is_some(),
)))
}
OcrHostOperation::AcquireAzureAdToken => {
OcrHostResult::AzureAdToken(Ok(self.acquire_azure_ad_token(py)?))
}
OcrHostOperation::PreCall(request) => {
OcrHostResult::PreCall(Ok(self.pre_call(py, request)?))
}
OcrHostOperation::DuringCall(request) => {
OcrHostResult::DuringCall(Ok(self.python_pre_call(py, request)?))
}
OcrHostOperation::PostCall(request) => {
OcrHostResult::PostCall(Ok(self.python_post_call(py, request)?))
}
OcrHostOperation::ConstructResponse(response) => {
self.state.end = Some(now(py)?);
self.state.response = Some(
py.import("litellm.rust_bridge.ocr")?
.getattr("_response")?
.call1((to_py(py, response.as_ref())?,))?
.unbind(),
);
OcrHostResult::Lifecycle(Ok(()))
}
OcrHostOperation::MapFailure(error) => {
if self.state.error.is_none() {
self.state.retain_error(py, ocr_error_to_pyerr(error));
}
if self.state.end.is_none() {
self.state.end = Some(now(py)?);
}
let error = self.state.error.as_ref().ok_or_else(missing_state)?;
let request = self.request.as_ref().ok_or_else(missing_state)?.bind(py);
let mapped = py
.import("litellm.rust_bridge.ocr_lifecycle")?
.getattr("map_failure")?
.call1((error, request, &self.provider))?;
self.state.retain_error(py, PyErr::from_value(mapped));
OcrHostResult::Lifecycle(Ok(()))
}
OcrHostOperation::Lifecycle(_)
| OcrHostOperation::Success { .. }
| OcrHostOperation::Failure { .. } => return Err(missing_state()),
})
}
fn cleanup(&mut self) {
self.request = None;
self.pre_call = None;
self.document = None;
self.api_key = None;
self.azure_ad_token_provider = None;
self.retained_fields = None;
self.body = None;
self.headers = None;
}
fn traverse(&self, visit: &pyo3::gc::PyVisit<'_>) -> Result<(), pyo3::gc::PyTraverseError> {
visit.call(&self.request)?;
visit.call(&self.document)?;
visit.call(&self.api_key)?;
visit.call(&self.azure_ad_token_provider)?;
visit.call(&self.retained_fields)?;
visit.call(&self.body)?;
visit.call(&self.headers)
}
}
fn project_request(
py: Python<'_>,
request: &Bound<'_, PyAny>,
kwargs: &Bound<'_, PyDict>,
) -> PyResult<AdmittedOcrCall> {
let argument = |name: &str| {
kwargs
.get_item(name)?
.map(Ok)
.unwrap_or_else(|| request.getattr(name))
};
let model: String = argument("model")?.extract()?;
let custom_llm_provider: Option<String> = argument("custom_llm_provider")?.extract()?;
let document = argument("document")?;
let wire_document = extract_document(py, &document)?;
let retained_document = retained_document(py, &document, &wire_document)?;
let api_key = argument("api_key")?;
let request_kwargs = kwargs;
let consumed = consumed_optional_param_names(&model, custom_llm_provider.as_deref())
.map_err(ocr_error_to_pyerr)?;
let optional_params = extract_optional_params(request_kwargs, &consumed)?;
let input_sources = extract_input_sources(request_kwargs, &consumed)?;
let azure_ad_token_provider = request_kwargs
.get_item("azure_ad_token_provider")?
.filter(|provider| provider.is_callable() && provider.is_truthy().unwrap_or(false))
.map(Bound::unbind);
let wire = OcrWireRequest {
model,
document: wire_document,
api_key: api_key.extract()?,
api_base: argument("api_base")?.extract()?,
custom_llm_provider,
extra_headers: argument("extra_headers")?
.extract::<Option<Py<PyAny>>>()?
.map(|value| from_py(value.bind(py)))
.transpose()?,
optional_params,
input_sources,
timeout_seconds: argument("timeout")?
.extract::<Option<Py<PyAny>>>()?
.map(|value| {
py.import("litellm.rust_bridge.timeouts")?
.getattr("timeout_to_seconds")?
.call1((value,))?
.extract()
})
.transpose()?
.flatten(),
};
let request = decode_request(wire).map_err(ocr_error_to_pyerr)?;
let provider = request.provider_name().to_string();
let request = request.with_host_hooks(Arc::new(BridgeOcrHooks), None);
Ok(AdmittedOcrCall {
request,
document: retained_document,
api_key: api_key.unbind(),
azure_ad_token_provider,
provider,
})
}
fn extract_optional_params(
kwargs: &Bound<'_, PyDict>,
consumed: &[&str],
) -> PyResult<Map<String, Value>> {
let mut optional_params = Map::new();
for name in consumed {
if let Some(value) = kwargs.get_item(name)? {
optional_params.insert((*name).to_string(), from_py(&value)?);
}
}
Ok(optional_params)
}
fn extract_input_sources(
kwargs: &Bound<'_, PyDict>,
consumed: &[&str],
) -> PyResult<std::collections::BTreeMap<String, litellm_core::auth::InputSource>> {
let Some(proxy_request) = kwargs.get_item("proxy_server_request")? else {
return Ok(Default::default());
};
let proxy_request = proxy_request.cast_into::<PyDict>()?;
let body_fields = proxy_request
.get_item("body_fields")?
.or(proxy_request.get_item("body")?);
let credential_fields = proxy_request.get_item("credential_fields")?;
let mut sources = std::collections::BTreeMap::new();
for name in consumed
.iter()
.copied()
.chain(["api_key", "api_base", "extra_headers"])
{
let present = body_fields
.as_ref()
.is_some_and(|fields| fields.contains(name).unwrap_or(false))
|| credential_fields
.as_ref()
.is_some_and(|fields| fields.contains(name).unwrap_or(false));
if present {
sources.insert(name.to_string(), litellm_core::auth::InputSource::Request);
}
}
Ok(sources)
}
fn extract_document(py: Python<'_>, document: &Bound<'_, PyAny>) -> PyResult<Value> {
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()))
}
fn retained_document(
py: Python<'_>,
document: &Bound<'_, PyAny>,
wire_document: &Value,
) -> PyResult<Py<PyAny>> {
if document.get_item("type")?.extract::<String>()? == "file" {
to_py(py, wire_document)
} else {
Ok(document.clone().unbind())
}
}
fn admitted_call(outcome: NativeOutcome<OcrCall>) -> PyResult<OcrCall> {
match outcome {
NativeOutcome::Completed(call) => Ok(call),
NativeOutcome::Declined(reason) => Err(RustBridgeDeclined::new_err(format!(
"native OCR admission declined: {reason:?}"
))),
}
}
struct BridgeOcrHooks;
impl litellm_core::ocr::hooks::OcrHooks for BridgeOcrHooks {
fn has_guardrails(&self) -> bool {
true
}
}
#[pyfunction]
fn _ocr_lifecycle(
py: Python<'_>,
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
asynchronous: bool,
) -> PyResult<Py<PyAny>> {
let client = OcrClient::shared().map_err(ocr_error_to_pyerr)?;
let call = admitted_call(OcrCall::admit(
client,
OcrAdmission {
asynchronous,
..OcrAdmission::all()
},
))?;
let host = PythonOcrHost {
state: PythonCallState::new(
py,
args.unbind(),
kwargs.copy()?.unbind(),
asynchronous,
if asynchronous { "aocr" } else { "ocr" },
)?,
request: Some(request.unbind()),
pre_call: None,
document: None,
api_key: None,
azure_ad_token_provider: None,
provider: String::new(),
retained_fields: None,
body: None,
headers: None,
};
run_call(py, call, host)
}
pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
module.add_function(wrap_pyfunction!(_ocr_lifecycle, module)?)
}
#[cfg(test)]
mod tests {
use litellm_core::Error;
use litellm_core::ocr::OcrDecline;
use pyo3::exceptions::PyValueError;
use super::*;
#[test]
fn typed_initial_decline_uses_bridge_decline_contract() {
Python::initialize();
Python::attach(|py| {
let Err(error) = admitted_call(NativeOutcome::Declined(OcrDecline::HostOperations))
else {
panic!("unsupported host operations should decline admission");
};
assert!(error.is_instance_of::<RustBridgeDeclined>(py));
});
}
#[test]
fn post_admission_error_does_not_use_bridge_decline_contract() {
Python::initialize();
Python::attach(|py| {
let error = ocr_error_to_pyerr(Error::InvalidRequest("callback result".into()));
assert!(error.is_instance_of::<PyValueError>(py));
assert!(!error.is_instance_of::<RustBridgeDeclined>(py));
});
}
}

View file

@ -0,0 +1,186 @@
import asyncio
import gc
import threading
import weakref
from contextvars import ContextVar
async def exercise():
caller = asyncio.current_task()
thread = threading.get_ident()
loop = asyncio.get_running_loop()
marker = ContextVar("driver", default="before")
entered = asyncio.Event()
released = asyncio.Event()
result = object()
class CustomAwaitable:
def __await__(self):
return operation().__await__()
async def operation():
assert asyncio.current_task() is caller
assert threading.get_ident() == thread
assert asyncio.get_running_loop() is loop
marker.set("inside")
entered.set()
await released.wait()
assert asyncio.current_task() is caller
assert marker.get() == "inside"
return result
async def release():
await entered.wait()
released.set()
releaser = asyncio.create_task(release())
execution = await_execution(CustomAwaitable())
try:
execution.resume_value(None)
except RuntimeError:
pass
else:
raise AssertionError("resumed an unstarted execution")
wrapped = drive(execution)
try:
wrapped.send(1)
except TypeError:
pass
else:
raise AssertionError("accepted initial value")
assert await wrapped is result
assert marker.get() == "inside"
await releaser
execution.close()
execution.close()
try:
await wrapped
except RuntimeError:
pass
else:
raise AssertionError("accepted coroutine reuse")
final_awaitable = CustomAwaitable()
assert await drive(calling_execution(lambda: final_awaitable)) is final_awaitable
cause = KeyError("cause")
failure = ValueError("original")
async def failing():
await asyncio.sleep(0)
raise failure from cause
try:
await drive(await_execution(failing()))
except ValueError as error:
assert error is failure
assert error.__cause__ is cause
names = []
traceback = error.__traceback__
while traceback:
names.append(traceback.tb_frame.f_code.co_name)
traceback = traceback.tb_next
assert "failing" in names
else:
raise AssertionError("lost original exception")
for suppress in (False, True):
pending = asyncio.Event()
cleanup_entered = asyncio.Event()
cleanup_release = asyncio.Event()
cleaned = []
async def cancel_operation():
try:
pending.set()
await asyncio.Event().wait()
except asyncio.CancelledError:
if suppress:
return result
raise
finally:
cleanup_entered.set()
try:
await cleanup_release.wait()
except asyncio.CancelledError:
await cleanup_release.wait()
cleaned.append(asyncio.current_task())
task = asyncio.create_task(drive(await_execution(cancel_operation())))
await pending.wait()
task.cancel()
await cleanup_entered.wait()
assert not task.done()
task.cancel()
await asyncio.sleep(0)
cleanup_release.set()
if suppress:
assert await task is result
else:
try:
await task
except asyncio.CancelledError:
pass
else:
raise AssertionError("lost cancellation")
assert cleaned == [task]
observed = []
def reenter():
try:
active.start()
except RuntimeError as error:
observed.append(str(error))
return result
active = calling_execution(reenter)
assert await drive(active) is result
assert observed == ["execution is already running"]
class Finalizer:
def __call__(self):
return result
def __del__(self):
self.owner.close()
observed.append("released")
def cycle(started):
callback = Finalizer()
execution = calling_execution(callback)
callback.owner = execution
if started:
assert execution.start().value is result
return weakref.ref(callback)
for started in (False, True):
reference = cycle(started)
gc.collect()
assert reference() is None
assert observed[-2:] == ["released", "released"]
class Awaitable:
def __await__(self):
try:
yield self
finally:
observed.append("unwound")
def abandoned(started):
awaitable = Awaitable()
coroutine = drive(await_execution(awaitable))
awaitable.owner = coroutine
if started:
assert coroutine.send(None) is awaitable
coroutine.close()
return weakref.ref(awaitable)
for started in (False, True):
reference = abandoned(started)
gc.collect()
assert reference() is None
assert observed[-1] == "unwound"
asyncio.run(asyncio.wait_for(exercise(), 10))

View file

@ -1 +1,16 @@
litellm-python-interop is the domain-neutral PyO3 foundation. Keep generic Python/Serde conversion and interpreter primitives here. Do not add LiteLLM domain crates, route types, API registration, or cdylib build features.
- Target invariants; implementation and runtime validation may lag these rules
- Keep this crate a small, domain-neutral foundation: Python/Serde conversion and interpreter-boundary utilities
- No LiteLLM domain dependencies, route types, callback policy, public API registration or cdylib build features
- Generic code alone does not justify extraction: runtime integration stays in `python-bridge/src/execution.rs`, host adaptation in its `lifecycle.rs`
- Use standard PyO3 ownership and conversion APIs
- Prefer `Bound<'py, T>` for attached operations/results, `Py<T>` for retention; binding/unbinding does not copy payloads
- Use `pythonize` for selected Serde data, never a JSON-text round trip; share conversion with `Pythonized<T>`
- Preserve `PythonizeError`'s standard conversion into `PyErr`; do not stringify original Python exceptions into new `ValueError`s
- Keep serializer-panic containment in `Pythonized<T>`: async output conversion can run in an unjoined blocking task and otherwise strand delivery
- Use `Python::detach` for Rust-only work; Python operations require attachment
- Keep diagnostic counters in the consumer; wrapper invocations do not measure every interpreter release
- Release exclusive class borrows/locks before Python calls or decrements that can invoke finalizers; expose retained Python edges to GC without calling Python during traversal
- Keep coroutine driving in the shared Python driver and native adapter
- Driver: `litellm/rust_bridge/lifecycle.py`; handle: `python-bridge/src/lifecycle.rs`; native-backed behavior tests: `python-bridge/tests/lifecycle.py`
- References: [ownership](https://pyo3.rs/v0.29.2/types.html), [conversions](https://pyo3.rs/v0.29.2/conversions/traits.html), [pythonize errors](https://docs.rs/pythonize/0.29.0/src/pythonize/error.rs.html)
- [GC](https://pyo3.rs/v0.29.2/class/protocols.html#garbage-collector-integration), [re-entry](https://pyo3.rs/v0.29.2/class/call.html), [parallelism](https://pyo3.rs/v0.29.2/parallelism.html), [async delivery source](https://docs.rs/pyo3-async-runtimes/0.29.0/src/pyo3_async_runtimes/generic.rs.html)

View file

@ -1,7 +1,6 @@
use std::any::Any;
use std::panic::{AssertUnwindSafe, catch_unwind};
use pyo3::exceptions::PyValueError;
use pyo3::panic::PanicException;
use pyo3::prelude::*;
use serde::Serialize;
@ -11,16 +10,21 @@ pub fn from_py<T>(value: &Bound<'_, PyAny>) -> PyResult<T>
where
T: DeserializeOwned,
{
pythonize::depythonize(value).map_err(|error| PyValueError::new_err(error.to_string()))
pythonize::depythonize(value).map_err(PyErr::from)
}
pub fn to_py<T>(py: Python<'_>, value: &T) -> PyResult<Py<PyAny>>
where
T: Serialize + ?Sized,
{
pythonize::pythonize(py, value)
.map(Bound::unbind)
.map_err(|error| PyValueError::new_err(error.to_string()))
pythonize_bound(py, value).map(Bound::unbind)
}
fn pythonize_bound<'py, T>(py: Python<'py>, value: &T) -> PyResult<Bound<'py, PyAny>>
where
T: Serialize + ?Sized,
{
pythonize::pythonize(py, value).map_err(PyErr::from)
}
pub struct Pythonized<T>(pub T);
@ -34,9 +38,7 @@ where
type Error = PyErr;
fn into_pyobject(self, py: Python<'py>) -> PyResult<Self::Output> {
catch_unwind(AssertUnwindSafe(|| pythonize::pythonize(py, &self.0)))
.map_err(panic_to_pyerr)?
.map_err(|error| PyValueError::new_err(error.to_string()))
catch_unwind(AssertUnwindSafe(|| pythonize_bound(py, &self.0))).map_err(panic_to_pyerr)?
}
}
@ -89,4 +91,41 @@ mod tests {
assert_eq!(error.to_string(), "PanicException: serializer panicked");
});
}
#[test]
fn depythonize_preserves_python_exception_identity_and_traceback() {
Python::initialize();
Python::attach(|py| {
let locals = pyo3::types::PyDict::new(py);
py.run(
pyo3::ffi::c_str!(
r#"
failure = LookupError('conversion failed')
cause = ValueError('cause')
class Broken:
def __index__(self):
raise failure from cause
value = Broken()
"#
),
Some(&locals),
Some(&locals),
)
.unwrap();
let error = from_py::<i64>(&locals.get_item("value").unwrap().unwrap()).unwrap_err();
assert!(
error
.value(py)
.is(locals.get_item("failure").unwrap().unwrap())
);
assert!(
error
.cause(py)
.unwrap()
.value(py)
.is(locals.get_item("cause").unwrap().unwrap())
);
assert!(error.traceback(py).is_some());
});
}
}

102
litellm/ocr/input.py Normal file
View file

@ -0,0 +1,102 @@
import base64
import mimetypes
import os
import re
from io import IOBase
from typing import Any, Final
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"
def convert_file_document_to_url_document(document: dict[str, Any]) -> 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"
)
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}

View file

@ -1,461 +1,18 @@
"""
Main OCR function for LiteLLM.
"""
import asyncio
import base64
import mimetypes
import os
import re
from collections.abc import Callable, Coroutine, Mapping
from dataclasses import dataclass
from io import IOBase
from types import MappingProxyType
from typing import Any, Final, cast
from collections.abc import Awaitable, Mapping
from typing import Final, cast # noqa: TID251 # native binding selects a sync result or an async awaitable
import httpx
import litellm
from litellm._logging import verbose_logger
from litellm.constants import request_timeout
from litellm.litellm_core_utils.call_completion import CallCompletion
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.azure_ai.ocr.common_utils import (
is_azure_cohere_parse_model,
is_azure_document_intelligence_model,
)
from litellm.llms.base_llm.ocr.transformation import (
OCR_REQUEST_FORMAT_PARAM,
BaseOCRConfig,
OCRResponse,
parse_ocr_request_format,
)
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.rust_bridge import ocr as rust_ocr_bridge
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.ocr.input import convert_file_document_to_url_document, get_mime_type
from litellm.rust_bridge.bindings import native_exception_types
from litellm.rust_bridge.configuration import rust_enabled
from litellm.types.router import GenericLiteLLMParams
from litellm.utils import ProviderConfigManager, client
from litellm.rust_bridge.ocr import LiteLLMOcrRequest
from litellm.rust_bridge.ocr_lifecycle import select
####### ENVIRONMENT VARIABLES ###################
base_llm_http_handler = BaseLLMHTTPHandler()
#################################################
__all__ = ("aocr", "convert_file_document_to_url_document", "get_mime_type", "ocr")
@dataclass
class _PreparedOCRRequest:
model: str
document: dict[str, Any]
api_key: str | None
api_base: str | None
custom_llm_provider: str
extra_headers: dict[str, object] | None
provider_config: BaseOCRConfig
optional_params: dict[str, object]
litellm_params: dict[str, object]
effective_timeout: float | httpx.Timeout
litellm_logging_obj: LiteLLMLoggingObj
caller_supplied_api_key: bool = True
caller_supplied_api_base: bool = True
_RUST_OCR_PROVIDERS: Final = frozenset({"mistral", "azure_ai", "vertex_ai"})
_RUST_OCR_CONFIG_FIELDS: Final = frozenset(
{
"azure_ad_token",
"tenant_id",
"client_id",
"client_secret",
"azure_scope",
"azure_authority_host",
"azure_credential",
"azure_federated_token_file",
"vertex_credentials",
"vertex_ai_credentials",
"vertex_project",
"vertex_ai_project",
"vertex_location",
"vertex_ai_location",
}
)
_RUST_OCR_SECRET_FIELDS: Final = frozenset(
{"azure_ad_token", "client_secret", "azure_federated_token_file", "vertex_credentials", "vertex_ai_credentials"}
)
def _prepare_ocr_request(
model: str,
document: Mapping[str, object],
api_key: str | None,
api_base: str | None,
timeout: float | httpx.Timeout | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
kwargs: dict[str, object],
) -> _PreparedOCRRequest:
litellm_logging_obj: Final = cast(LiteLLMLoggingObj, kwargs.pop("litellm_logging_obj"))
litellm_call_id: Final = cast(str | None, kwargs.get("litellm_call_id", None))
if not isinstance(document, dict):
raise ValueError(f"document must be a dict with 'type' and URL/file field, got {type(document)}")
doc_type = document.get("type")
if doc_type == "file":
document = convert_file_document_to_url_document(document)
doc_type = document.get("type")
if doc_type not in ["document_url", "image_url"]:
raise ValueError(f"Invalid document type: {doc_type}. Must be 'document_url', 'image_url', or 'file'")
caller_supplied_api_key: Final = api_key is not None
caller_supplied_api_base: Final = api_base is not None
(
model,
custom_llm_provider,
dynamic_api_key,
dynamic_api_base,
) = litellm.get_llm_provider(
model=model,
custom_llm_provider=custom_llm_provider,
api_base=api_base,
api_key=api_key,
)
suppress_dynamic_api_base: Final = (
not caller_supplied_api_base
and custom_llm_provider == "azure_ai"
and is_azure_document_intelligence_model(model)
)
if dynamic_api_key:
api_key = dynamic_api_key
if dynamic_api_base and not suppress_dynamic_api_base:
api_base = dynamic_api_base
ocr_provider_config: Final = ProviderConfigManager.get_provider_ocr_config(
model=model,
provider=litellm.LlmProviders(custom_llm_provider),
)
if ocr_provider_config is None:
raise ValueError(f"OCR is not supported for provider: {custom_llm_provider}")
verbose_logger.debug("OCR call - model: %s, provider: %s", model, custom_llm_provider)
litellm_params: Final = GenericLiteLLMParams.model_validate(kwargs)
supported_params: Final = ocr_provider_config.get_supported_ocr_params(model=model)
requested_format: Final = kwargs.get(OCR_REQUEST_FORMAT_PARAM)
if requested_format is not None:
try:
parsed_format: Final = parse_ocr_request_format(requested_format)
except ValueError as e:
raise litellm.exceptions.UnsupportedParamsError(
message=f"{e}", model=model, llm_provider=custom_llm_provider
) from e
if OCR_REQUEST_FORMAT_PARAM not in supported_params and parsed_format == "native":
raise litellm.exceptions.UnsupportedParamsError(
message=(
f"`{OCR_REQUEST_FORMAT_PARAM}='native'` is not supported for provider: {custom_llm_provider}, "
f"model: {model}"
),
model=model,
llm_provider=custom_llm_provider,
)
non_default_params: Final = {}
for param in supported_params:
if param in kwargs:
non_default_params[param] = kwargs.pop(param)
optional_params: Final = ocr_provider_config.map_ocr_params(
non_default_params=non_default_params,
optional_params={},
model=model,
)
verbose_logger.debug("OCR optional_params after mapping: %s", optional_params)
effective_timeout: Final = timeout or request_timeout
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model,
optional_params=optional_params,
litellm_params={
"litellm_call_id": litellm_call_id,
"api_base": api_base,
},
custom_llm_provider=custom_llm_provider,
)
return _PreparedOCRRequest(
model=model,
document=document,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
provider_config=ocr_provider_config,
optional_params=cast(dict[str, object], optional_params),
litellm_params=dict(litellm_params),
effective_timeout=effective_timeout,
litellm_logging_obj=litellm_logging_obj,
caller_supplied_api_key=caller_supplied_api_key,
caller_supplied_api_base=caller_supplied_api_base,
)
def _rust_ocr_provider(request: rust_ocr_bridge.LiteLLMOcrRequest) -> str | None:
if request.custom_llm_provider is not None:
return request.custom_llm_provider
prefix: Final = request.model.partition("/")[0]
if prefix in _RUST_OCR_PROVIDERS:
return prefix
if request.model.startswith("mistral-ocr"):
return "mistral"
return None
def _rust_ocr_supported(request: rust_ocr_bridge.LiteLLMOcrRequest) -> bool:
provider: Final = _rust_ocr_provider(request)
if provider not in _RUST_OCR_PROVIDERS or request.kwargs.get(OCR_REQUEST_FORMAT_PARAM) == "native":
return False
if provider == "azure_ai":
return (
not is_azure_cohere_parse_model(request.model)
and not callable(request.kwargs.get("azure_ad_token_provider"))
and request.kwargs.get("azure_username") is None
and request.kwargs.get("azure_password") is None
)
return True
def _rust_bridge_optional_params(
request: rust_ocr_bridge.LiteLLMOcrRequest,
resolve_secret: Callable[[str], str | None],
) -> Mapping[str, object]:
optional_params: Final = MappingProxyType(
{
name: value
for name, value in request.kwargs.items()
if (name not in GenericLiteLLMParams.model_fields or name in _RUST_OCR_CONFIG_FIELDS)
and name not in {"litellm_logging_obj", "aocr", "litellm_call_id", "proxy_server_request"}
}
)
provider: Final = _rust_ocr_provider(request)
if provider == "azure_ai" and litellm.enable_azure_ad_token_refresh is True:
return MappingProxyType({**optional_params, "enable_azure_ad_token_refresh": True})
if provider != "vertex_ai":
return optional_params
project: Final = (
request.kwargs.get("vertex_project")
or request.kwargs.get("vertex_ai_project")
or litellm.vertex_project
or resolve_secret("VERTEXAI_PROJECT")
)
location: Final = (
request.kwargs.get("vertex_location")
or request.kwargs.get("vertex_ai_location")
or litellm.vertex_location
or resolve_secret("VERTEXAI_LOCATION")
or resolve_secret("VERTEX_LOCATION")
)
credentials: Final = (
request.kwargs.get("vertex_credentials")
or request.kwargs.get("vertex_ai_credentials")
or resolve_secret("VERTEXAI_CREDENTIALS")
)
vertex_params: Final = MappingProxyType(
{
name: value
for name, value in (
("vertex_project", project),
("vertex_location", location),
("vertex_credentials", credentials),
)
if value is not None
}
)
return MappingProxyType({**optional_params, **vertex_params})
def _rust_bridge_input_sources(
request: rust_ocr_bridge.LiteLLMOcrRequest,
optional_params: Mapping[str, object],
) -> Mapping[str, str]:
proxy_request: Final = request.kwargs.get("proxy_server_request")
if not isinstance(proxy_request, Mapping):
return MappingProxyType({})
proxy_request_mapping: Final = cast( # cast-ok: runtime Mapping check loses generic key and value types
Mapping[object, object], proxy_request
)
body_value: Final = proxy_request_mapping.get("body")
if not isinstance(body_value, Mapping):
return MappingProxyType({})
body: Final = cast( # cast-ok: runtime Mapping check loses generic key and value types
Mapping[object, object], body_value
)
credential_fields_value: Final = proxy_request_mapping.get("credential_fields", ())
credential_fields: Final = (
frozenset(name for name in credential_fields_value if isinstance(name, str))
if isinstance(credential_fields_value, (list, tuple, set, frozenset))
else frozenset()
)
names: Final = frozenset(optional_params) | frozenset({"api_key", "api_base", "extra_headers"})
request_sources: Final = MappingProxyType(
{name: "request" for name in names if name in body or name in credential_fields}
)
if litellm.enable_azure_ad_token_refresh is True and "enable_azure_ad_token_refresh" in optional_params:
return MappingProxyType({**request_sources, "enable_azure_ad_token_refresh": "deployment"})
return request_sources
def _marshal_rust_ocr_request(
request: rust_ocr_bridge.LiteLLMOcrRequest,
resolve_secret: Callable[[str], str | None],
) -> rust_ocr_bridge.LiteLLMOcrRequest:
if not isinstance(request.document, dict):
raise TypeError(f"document must be a dict with 'type' and URL/file field, got {type(request.document)}")
document: Final = (
convert_file_document_to_url_document(request.document)
if request.document.get("type") == "file"
else request.document
)
provider: Final = _rust_ocr_provider(request)
api_key: Final = request.api_key or resolve_secret("MISTRAL_API_KEY") if provider == "mistral" else request.api_key
optional_params: Final = _rust_bridge_optional_params(request, resolve_secret)
input_sources: Final = _rust_bridge_input_sources(request, optional_params)
logged_optional_params: Final = MappingProxyType(
{name: "****" if name in _RUST_OCR_SECRET_FIELDS else value for name, value in optional_params.items()}
)
logged_kwargs: Final = MappingProxyType(
{
name: "****" if name in _RUST_OCR_SECRET_FIELDS else value
for name, value in request.kwargs.items()
if name != "proxy_server_request"
}
)
logging_obj: Final = cast( # cast-ok: bridge kwargs carry the prepared logging object
LiteLLMLoggingObj, request.kwargs["litellm_logging_obj"]
)
logging_obj.update_from_kwargs(
kwargs=dict(logged_kwargs), # mutable-ok: logging API requires an owned dict
model=request.model,
optional_params=dict(logged_optional_params), # mutable-ok: logging API requires an owned dict
litellm_params={
"litellm_call_id": request.kwargs.get("litellm_call_id"),
"api_base": request.api_base,
}, # mutable-ok: legacy logging requires a concrete params dict
custom_llm_provider=provider,
)
logging_obj.pre_call(
input="OCR document processing",
api_key=api_key,
additional_args={ # mutable-ok: pre_call mutates the additional_args dict
"complete_input_dict": {
"model": request.model,
"document": document,
**logged_optional_params,
}, # mutable-ok: callbacks consume a JSON-serializable request dict
"api_base": request.api_base or "",
"headers": request.extra_headers or {}, # mutable-ok: logging callbacks consume a concrete headers dict
},
)
return rust_ocr_bridge.LiteLLMOcrRequest(
model=request.model,
document=document,
api_key=api_key,
api_base=request.api_base,
timeout=request.timeout if request.timeout is not None else request_timeout,
custom_llm_provider=request.custom_llm_provider,
extra_headers=request.extra_headers,
kwargs=optional_params,
input_sources=input_sources,
)
def _map_rust_ocr_error(
error: Exception,
request: rust_ocr_bridge.LiteLLMOcrRequest,
exception_types: tuple[type[BaseException], type[BaseException]] | None,
) -> Exception:
if exception_types is None or not isinstance(error, exception_types[1]):
return error
provider: Final = _rust_ocr_provider(request)
if provider is None:
return error
provider_config: Final = ProviderConfigManager.get_provider_ocr_config(
model=request.model.removeprefix(f"{provider}/"), provider=litellm.LlmProviders(provider)
)
if provider_config is None:
return error
error_args: Final = cast( # cast-ok: Python exceptions expose positional args as a tuple
tuple[object, ...], error.args
)
status: Final = error_args[0] if error_args and isinstance(error_args[0], int) else 500
message: Final = str(error_args[1]) if len(error_args) > 1 else str(error)
error_factory: Final = cast( # cast-ok: provider configs expose heterogeneous exception factories
Callable[..., Exception], provider_config.get_error_class
)
return error_factory(
error_message=message, status_code=status or 500, headers={}
) # mutable-ok: provider error factories require a concrete headers dict
def _run_rust_ocr(
request: rust_ocr_bridge.LiteLLMOcrRequest,
resolve_api_key: Callable[[str], str | None],
) -> OCRResponse | None:
if rust_ocr_bridge.load_rust_ocr() is None:
return None
marshalled: Final = _marshal_rust_ocr_request(request, resolve_api_key)
input_sources: Final = marshalled.input_sources
try:
response: Final = rust_ocr_bridge.ocr(
model=marshalled.model,
document=dict(marshalled.document), # mutable-ok: PyO3 OCR binding requires a concrete dict
api_key=marshalled.api_key,
api_base=marshalled.api_base,
custom_llm_provider=marshalled.custom_llm_provider,
extra_headers=marshalled.extra_headers,
optional_params=dict(marshalled.kwargs), # mutable-ok: PyO3 OCR binding requires a concrete dict
input_sources=input_sources,
timeout=marshalled.timeout,
)
except Exception as error:
raise _map_rust_ocr_error(error, request, native_exception_types()) from error
return OCRResponse.model_validate(response) if response is not None else None
async def _run_rust_aocr(
request: rust_ocr_bridge.LiteLLMOcrRequest,
resolve_api_key: Callable[[str], str | None],
) -> OCRResponse | None:
if rust_ocr_bridge.load_rust_aocr() is None:
return None
marshalled: Final = _marshal_rust_ocr_request(request, resolve_api_key)
input_sources: Final = marshalled.input_sources
try:
response: Final = await rust_ocr_bridge.aocr(
model=marshalled.model,
document=dict(marshalled.document), # mutable-ok: PyO3 OCR binding requires a concrete dict
api_key=marshalled.api_key,
api_base=marshalled.api_base,
custom_llm_provider=marshalled.custom_llm_provider,
extra_headers=marshalled.extra_headers,
optional_params=dict(marshalled.kwargs), # mutable-ok: PyO3 OCR binding requires a concrete dict
input_sources=input_sources,
timeout=marshalled.timeout,
)
except Exception as error:
raise _map_rust_ocr_error(error, request, native_exception_types()) from error
return OCRResponse.model_validate(response) if response is not None else None
@client
async def aocr(
def _bind_request(
model: str,
document: Mapping[str, object],
api_key: str | None = None,
@ -463,79 +20,9 @@ async def aocr(
timeout: float | httpx.Timeout | None = None,
custom_llm_provider: str | None = None,
extra_headers: dict[str, object] | None = None,
*,
_litellm_call_completion: CallCompletion | None = None,
**kwargs: object,
) -> OCRResponse:
"""
Async OCR function.
Args:
model: Model name (e.g., "mistral/mistral-ocr-latest")
document: Document to process in Mistral format:
{"type": "document_url", "document_url": "https://..."} for PDFs/docs,
{"type": "image_url", "image_url": "https://..."} for images, or
{"type": "file", "file": <path/bytes/file-obj>} for local files
api_key: Optional API key
api_base: Optional API base URL
timeout: Optional timeout
custom_llm_provider: Optional custom LLM provider
extra_headers: Optional extra headers
**kwargs: Additional parameters (e.g., include_image_base64, pages, image_limit)
Returns:
OCRResponse in Mistral OCR format with pages, model, usage_info, etc.
Example:
```python
import litellm
# OCR with PDF
response = await litellm.aocr(
model="mistral/mistral-ocr-latest",
document={
"type": "document_url",
"document_url": "https://arxiv.org/pdf/2201.04234"
},
include_image_base64=True
)
# OCR with image
response = await litellm.aocr(
model="mistral/mistral-ocr-latest",
document={
"type": "image_url",
"image_url": "https://example.com/image.png"
}
)
# OCR with base64 encoded PDF
response = await litellm.aocr(
model="mistral/mistral-ocr-latest",
document={
"type": "document_url",
"document_url": f"data:application/pdf;base64,{base64_pdf}"
}
)
# OCR with local file
response = await litellm.aocr(
model="mistral/mistral-ocr-latest",
document={"type": "file", "file": "/path/to/document.pdf"}
)
```
"""
completion_kwargs: Final[dict[str, object]] = {
"model": model,
"document": document,
"api_key": api_key,
"api_base": api_base,
"timeout": timeout,
"custom_llm_provider": custom_llm_provider,
"extra_headers": extra_headers,
"kwargs": kwargs,
}
request: Final = rust_ocr_bridge.LiteLLMOcrRequest(
**kwargs: object, # kwargs-ok: public OCR accepts provider-specific options
) -> LiteLLMOcrRequest:
return LiteLLMOcrRequest(
model=model,
document=document,
api_key=api_key,
@ -544,345 +31,40 @@ async def aocr(
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
kwargs=kwargs,
call_completion=_litellm_call_completion,
)
def _public_request(name: str, args: tuple[object, ...], kwargs: dict[str, object]) -> LiteLLMOcrRequest:
try:
if rust_enabled() and _rust_ocr_supported(request):
from litellm.secret_managers.main import get_secret_str
rust_response: Final = await _run_rust_aocr(
request=request,
resolve_api_key=get_secret_str,
)
if rust_response is None:
verbose_logger.debug("Async Rust OCR bridge unavailable; falling back to Python path")
else:
return rust_response
prepared: Final = _prepare_ocr_request(
model=model,
document=document,
api_key=api_key,
api_base=api_base,
timeout=timeout,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
kwargs=kwargs,
)
model = prepared.model
custom_llm_provider = prepared.custom_llm_provider
completion_kwargs.update({"model": model, "custom_llm_provider": custom_llm_provider})
response = base_llm_http_handler.ocr(
model=prepared.model,
document=prepared.document,
optional_params=prepared.optional_params,
timeout=prepared.effective_timeout,
logging_obj=prepared.litellm_logging_obj,
api_key=prepared.api_key,
api_base=prepared.api_base,
custom_llm_provider=prepared.custom_llm_provider,
aocr=True,
headers=prepared.extra_headers,
provider_config=prepared.provider_config,
litellm_params=prepared.litellm_params,
)
if asyncio.iscoroutine(response):
response = await response
if response is None:
raise ValueError(f"Got an unexpected None response from the OCR API: {response}")
return response
except Exception as e:
error_provider: Final = custom_llm_provider or _rust_ocr_provider(request)
error_model: Final = model.removeprefix(f"{error_provider}/") if error_provider else model
raise litellm.exception_type(
model=error_model,
custom_llm_provider=error_provider,
original_exception=e,
completion_kwargs=completion_kwargs,
extra_kwargs=kwargs,
)
return _bind_request(*args, **kwargs) # pyright: ignore[reportArgumentType] # Python binds the public arguments before native validation
except TypeError as error:
raise TypeError(str(error).replace("_bind_request()", f"{name}()")) from None
#################################################
# Public utilities — used by the SDK and the proxy
#################################################
_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:
"""
Determine MIME type from file path extension.
Falls back to mimetypes.guess_type, then to 'application/octet-stream'.
"""
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"
def convert_file_document_to_url_document(document: dict[str, Any]) -> dict[str, str]:
"""
Convert a file-type document dict to a document_url-type document dict
with an inline base64 data URI.
Accepts document dicts like:
{"type": "file", "file": Path("/path/to/doc.pdf")} # pathlib.Path
{"type": "file", "file": <binary file-like object>} # file-like object (BinaryIO)
{"type": "file", "file": b"raw bytes"} # raw bytes
Bare ``str`` paths are not accepted — pass a ``pathlib.Path`` or
``open(path, "rb")`` instead. See the str check below for the rationale.
Returns:
{"type": "document_url", "document_url": "data:<mime>;base64,<data>"}
or {"type": "image_url", "image_url": "data:<mime>;base64,<data>"}
"""
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"
)
file_bytes: bytes
mime_type: str = "application/octet-stream"
file_name: str | None = None
if isinstance(file_input, str):
# Bare strings are rejected here. The OCR ``document`` accepts a
# ``{"type": "file", "file": <value>}`` shape, and when this helper
# runs in a proxy request handler ``<value>`` is attacker-controlled.
# Opening it as a path is an arbitrary local file read on the proxy
# host, which is then base64-encoded and forwarded to the OCR
# provider — an exfiltration primitive.
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):
# os.PathLike (pathlib.Path and custom __fspath__ classes) is a
# Python-level type that HTTP form values can't fabricate.
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 f:
file_bytes = f.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}
@client
def ocr(
model: str,
document: Mapping[str, object],
api_key: str | None = None,
api_base: str | None = None,
timeout: float | httpx.Timeout | None = None,
custom_llm_provider: str | None = None,
extra_headers: dict[str, object] | None = None,
*,
_litellm_call_completion: CallCompletion | None = None,
**kwargs: object,
) -> OCRResponse | Coroutine[object, object, OCRResponse]:
"""
Synchronous OCR function.
Args:
model: Model name (e.g., "mistral/mistral-ocr-latest")
document: Document to process in Mistral format:
{"type": "document_url", "document_url": "https://..."} for PDFs/docs,
{"type": "image_url", "image_url": "https://..."} for images, or
{"type": "file", "file": <path/bytes/file-obj>} for local files
api_key: Optional API key
api_base: Optional API base URL
timeout: Optional timeout
custom_llm_provider: Optional custom LLM provider
extra_headers: Optional extra headers
**kwargs: Additional parameters (e.g., include_image_base64, pages, image_limit)
Returns:
OCRResponse in Mistral OCR format with pages, model, usage_info, etc.
Example:
```python
import litellm
# OCR with PDF
response = litellm.ocr(
model="mistral/mistral-ocr-latest",
document={
"type": "document_url",
"document_url": "https://arxiv.org/pdf/2201.04234"
},
include_image_base64=True
)
# OCR with image
response = litellm.ocr(
model="mistral/mistral-ocr-latest",
document={
"type": "image_url",
"image_url": "https://example.com/image.png"
}
)
# OCR with base64 encoded PDF
response = litellm.ocr(
model="mistral/mistral-ocr-latest",
document={
"type": "document_url",
"document_url": f"data:application/pdf;base64,{base64_pdf}"
}
)
# OCR with local file
response = litellm.ocr(
model="mistral/mistral-ocr-latest",
document={"type": "file", "file": "/path/to/document.pdf"}
)
# Access pages
for page in response.pages:
print(f"Page {page.index}: {page.markdown}")
```
"""
completion_kwargs: Final[dict[str, object]] = {
"model": model,
"document": document,
"api_key": api_key,
"api_base": api_base,
"timeout": timeout,
"custom_llm_provider": custom_llm_provider,
"extra_headers": extra_headers,
"kwargs": kwargs,
}
request: Final = rust_ocr_bridge.LiteLLMOcrRequest(
model=model,
document=document,
api_key=api_key,
api_base=api_base,
timeout=timeout,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
kwargs=kwargs,
call_completion=_litellm_call_completion,
)
def ocr(*args: object, **kwargs: object) -> OCRResponse:
request: Final = _public_request("ocr", args, kwargs)
native: Final = select(request)
if native is None:
raise RuntimeError("Rust OCR is unavailable or does not support this request")
try:
_is_async: Final = kwargs.pop("aocr", False) is True
completion_kwargs["aocr"] = _is_async
if rust_enabled() and _rust_ocr_supported(request):
from litellm.secret_managers.main import get_secret_str
return cast(OCRResponse, native(request, args, kwargs, False)) # cast-ok: False selects the synchronous result
except _decline_types() as error:
raise RuntimeError(f"Rust OCR declined the request: {error}") from error
rust_response: Final = _run_rust_ocr(
request=request,
resolve_api_key=get_secret_str,
)
if rust_response is None:
verbose_logger.debug("Rust OCR bridge unavailable; falling back to Python path")
else:
return rust_response
prepared: Final = _prepare_ocr_request(
model=model,
document=document,
api_key=api_key,
api_base=api_base,
kwargs=kwargs,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout=timeout,
)
model = prepared.model
custom_llm_provider = prepared.custom_llm_provider
completion_kwargs.update({"model": model, "custom_llm_provider": custom_llm_provider})
async def aocr(*args: object, **kwargs: object) -> OCRResponse:
request: Final = _public_request("aocr", args, kwargs)
native: Final = select(request)
if native is None:
raise RuntimeError("Rust OCR is unavailable or does not support this request")
try:
return await cast(
Awaitable[OCRResponse], native(request, args, kwargs, True)
) # cast-ok: True selects the asynchronous result
except _decline_types() as error:
raise RuntimeError(f"Rust OCR declined the request: {error}") from error
response: Final = base_llm_http_handler.ocr(
model=prepared.model,
document=prepared.document,
optional_params=prepared.optional_params,
timeout=prepared.effective_timeout,
logging_obj=prepared.litellm_logging_obj,
api_key=prepared.api_key,
api_base=prepared.api_base,
custom_llm_provider=prepared.custom_llm_provider,
aocr=_is_async,
headers=prepared.extra_headers,
provider_config=prepared.provider_config,
litellm_params=prepared.litellm_params,
)
return response
except Exception as e:
error_provider: Final = custom_llm_provider or _rust_ocr_provider(request)
error_model: Final = model.removeprefix(f"{error_provider}/") if error_provider else model
raise litellm.exception_type(
model=error_model,
custom_llm_provider=error_provider,
original_exception=e,
completion_kwargs=completion_kwargs,
extra_kwargs=kwargs,
)
def _decline_types() -> tuple[type[BaseException], ...]:
exception_types: Final = native_exception_types()
return (exception_types[0],) if exception_types is not None else ()

View file

@ -3181,6 +3181,10 @@ class ProxyBaseLLMRequestProcessing:
Extracted as a static method so tests can exercise the production
gating logic directly rather than reimplementing the finally block.
"""
pending: Final = getattr(logging_obj, "_native_pending_logging", None)
if pending is not None:
logging_obj._native_pending_logging = None # rebind-ok: consume the native release signal once
pending.release(not exception_raised)
_enqueue_fn: Final = getattr(logging_obj, "_enqueue_deferred_logging", None)
if _enqueue_fn is None:
return

View file

@ -315,7 +315,6 @@ class UnifiedLLMGuardrails(CustomLogger):
if call_type is None:
call_type = _infer_call_type(call_type=None, completion_response=response)
# Fallback: resolve call_type from logging_obj for pass-through endpoints
if call_type is None:
litellm_logging_obj: Final = data.get("litellm_logging_obj")
logging_call_type: Final = (
@ -324,6 +323,8 @@ class UnifiedLLMGuardrails(CustomLogger):
if logging_call_type in (
CallTypes.pass_through.value,
CallTypes.allm_passthrough_route.value,
CallTypes.ocr.value,
CallTypes.aocr.value,
):
call_type = logging_call_type

View file

@ -0,0 +1,144 @@
from __future__ import annotations
import datetime
import uuid
from collections.abc import Awaitable, Mapping
from dataclasses import dataclass
from typing import (
TYPE_CHECKING,
Final,
Protocol,
cast, # noqa: TID251 # bounded compatibility calls into legacy Python integrations
)
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging
@dataclass(frozen=True, slots=True)
class Await:
awaitable: Awaitable[object]
@dataclass(frozen=True, slots=True)
class Complete:
value: object
class Execution(Protocol):
def start(self) -> Await | Complete: ...
def resume_value(self, value: object) -> Await | Complete: ...
def resume_error(self, error: BaseException) -> Await | Complete: ...
def close(self) -> None: ...
async def drive(execution: Execution) -> object:
try:
step = execution.start() # rebind-ok: the execution protocol advances after each selected await
while isinstance(step, Await):
try:
value = await step.awaitable # rebind-ok: each selected await produces the next protocol input
except GeneratorExit:
raise
except BaseException as error:
step = execution.resume_error(error) # rebind-ok: advance the execution protocol
else:
step = execution.resume_value(value) # rebind-ok: advance the execution protocol
return step.value
finally:
execution.close()
class CredentialLoader(Protocol):
def __call__(self, kwargs: dict[str, object]) -> None: ...
class MetadataUpdater(Protocol):
def __call__(
self,
result: object,
logging_obj: Logging,
model: str | None,
kwargs: dict[str, object],
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
@dataclass(frozen=True, slots=True)
class CallSetup:
logger: Logging
kwargs: dict[str, object]
def setup(
call_type: str,
args: tuple[object, ...],
kwargs: Mapping[str, object],
start_time: datetime.datetime,
asynchronous: bool,
) -> CallSetup:
from litellm import utils
from litellm.litellm_core_utils.litellm_logging import Logging
arguments: Final = { # mutable-ok: function_setup consumes an owned kwargs dict
"litellm_call_id": str(uuid.uuid4()),
**kwargs,
}
supplied: Final = arguments.get("litellm_logging_obj")
if isinstance(supplied, Logging):
return CallSetup(supplied, arguments)
logger, prepared = utils.function_setup(
call_type, utils.Rules(), start_time, *args, is_async_call=asynchronous, **arguments
)
return CallSetup(logger, prepared)
def prepare(kwargs: Mapping[str, object], logger: Logging) -> dict[str, object]:
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")
if isinstance(metadata, Mapping):
typed_metadata: Final = cast( # cast-ok: runtime Mapping check establishes read-only metadata
Mapping[str, object], metadata
)
previous: Final = typed_metadata.get("previous_models")
if (
isinstance(previous, list)
and litellm.num_retries_per_request is not None
and len(cast(list[object], previous)) # cast-ok: runtime list check establishes the retry history
>= litellm.num_retries_per_request
):
raise RuntimeError("Max retries per request hit!")
return arguments
def finalize(
response: object,
logger: Logging,
kwargs: dict[str, object],
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None:
from litellm.litellm_core_utils.llm_response_utils import response_metadata
model: Final = kwargs.get("model")
update: Final = cast( # cast-ok: legacy metadata function accepts concrete kwargs
MetadataUpdater, response_metadata.update_response_metadata
)
update(response, logger, model if isinstance(model, str) else None, kwargs, start_time, end_time)

View file

@ -2,47 +2,30 @@
from __future__ import annotations
from collections.abc import Awaitable, Callable, Mapping, Sequence
from collections.abc import Awaitable, Mapping
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final, Protocol, cast # noqa: TID251 # native extension exposes dynamically typed callables
import httpx
import litellm
from litellm.constants import request_timeout
from litellm.litellm_core_utils.call_completion import CallCompletion
from litellm.llms.azure_ai.ocr.common_utils import is_azure_cohere_parse_model
from litellm.llms.base_llm.ocr.transformation import PROVIDER_NATIVE_RESPONSE_KEY, OCRResponse
from litellm.rust_bridge.bindings import NativeBinding, native_exception_types
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.timeouts import timeout_to_seconds as _timeout_to_seconds
from litellm.types.router import GenericLiteLLMParams
from litellm.utils import ProviderConfigManager
_RUST_OCR_PROVIDERS: Final = frozenset({"mistral", "azure_ai", "vertex_ai"})
_RUST_OCR_CONFIG_FIELDS: Final = frozenset(
{
"azure_ad_token",
"tenant_id",
"client_id",
"client_secret",
"azure_scope",
"azure_authority_host",
"azure_credential",
"azure_federated_token_file",
"vertex_credentials",
"vertex_ai_credentials",
"vertex_project",
"vertex_ai_project",
"vertex_location",
"vertex_ai_location",
}
)
_RUST_OCR_SECRET_FIELDS: Final = frozenset(
{"azure_ad_token", "client_secret", "azure_federated_token_file", "vertex_credentials", "vertex_ai_credentials"}
)
def redact_logging_params(params: Mapping[str, object]) -> dict[str, object]:
return { # mutable-ok: Logging.update_from_kwargs requires concrete params
name: "****" if name in _RUST_OCR_SECRET_FIELDS else value
for name, value in params.items()
if name != "proxy_server_request"
}
@dataclass(frozen=True, slots=True)
class LiteLLMOcrRequest:
model: str
@ -53,7 +36,6 @@ class LiteLLMOcrRequest:
custom_llm_provider: str | None
extra_headers: dict[str, object] | None
kwargs: Mapping[str, object]
call_completion: CallCompletion | None = None
input_sources: Mapping[str, str] | None = None
@ -89,26 +71,6 @@ class RustAocr(Protocol):
raise NotImplementedError
class _OCRLogging(Protocol):
def update_from_kwargs(
self,
*,
kwargs: dict[str, object],
model: str,
optional_params: dict[str, object],
litellm_params: dict[str, object],
custom_llm_provider: str | None,
) -> None: ...
def pre_call(
self,
*,
input: str,
api_key: str | None,
additional_args: dict[str, object],
) -> None: ...
def _as_ocr(value: object) -> RustOcr | None:
return cast(RustOcr, value) if callable(value) else None
@ -129,205 +91,6 @@ def load_rust_aocr() -> RustAocr | None:
return _AOCR.load()
def provider(request: LiteLLMOcrRequest) -> str | None:
if request.custom_llm_provider is not None:
return request.custom_llm_provider
prefix: Final = request.model.partition("/")[0]
if prefix in _RUST_OCR_PROVIDERS:
return prefix
if request.model.startswith("mistral-ocr"):
return "mistral"
return None
def supported(request: LiteLLMOcrRequest) -> bool:
request_provider: Final = provider(request)
if request_provider not in _RUST_OCR_PROVIDERS:
return False
if request_provider == "azure_ai":
return (
not is_azure_cohere_parse_model(request.model)
and not callable(request.kwargs.get("azure_ad_token_provider"))
and request.kwargs.get("azure_username") is None
and request.kwargs.get("azure_password") is None
)
return True
def _optional_params(request: LiteLLMOcrRequest, resolve_secret: Callable[[str], str | None]) -> Mapping[str, object]:
optional_params: Final = MappingProxyType(
{
name: value
for name, value in request.kwargs.items()
if (name not in GenericLiteLLMParams.model_fields or name in _RUST_OCR_CONFIG_FIELDS)
and name not in ("litellm_logging_obj", "aocr", "litellm_call_id", "proxy_server_request")
}
)
request_provider: Final = provider(request)
if request_provider == "azure_ai" and litellm.enable_azure_ad_token_refresh is True:
return MappingProxyType({**optional_params, "enable_azure_ad_token_refresh": True})
if request_provider != "vertex_ai":
return optional_params
project: Final = (
request.kwargs.get("vertex_project")
or request.kwargs.get("vertex_ai_project")
or litellm.vertex_project
or resolve_secret("VERTEXAI_PROJECT")
)
location: Final = (
request.kwargs.get("vertex_location")
or request.kwargs.get("vertex_ai_location")
or litellm.vertex_location
or resolve_secret("VERTEXAI_LOCATION")
or resolve_secret("VERTEX_LOCATION")
)
credentials: Final = (
request.kwargs.get("vertex_credentials")
or request.kwargs.get("vertex_ai_credentials")
or resolve_secret("VERTEXAI_CREDENTIALS")
)
vertex_params: Final = MappingProxyType(
{
name: value
for name, value in (
("vertex_project", project),
("vertex_location", location),
("vertex_credentials", credentials),
)
if value is not None
}
)
return MappingProxyType({**optional_params, **vertex_params})
def _input_sources(request: LiteLLMOcrRequest, optional_params: Mapping[str, object]) -> Mapping[str, str]:
proxy_request_value: Final = request.kwargs.get("proxy_server_request")
if not isinstance(proxy_request_value, Mapping):
return MappingProxyType({})
proxy_request: Final = cast( # cast-ok: runtime Mapping check narrows metadata with unknown key and value types
Mapping[object, object], proxy_request_value
)
credential_fields_value: Final = proxy_request.get("credential_fields", ())
credential_fields: Final = (
frozenset(name for name in credential_fields_value if isinstance(name, str))
if isinstance(credential_fields_value, (list, tuple, set, frozenset))
else frozenset()
)
request_fields_value: Final = proxy_request.get("body_fields")
request_fields: Sequence[object]
if isinstance(request_fields_value, Sequence) and not isinstance(request_fields_value, (str, bytes)):
request_fields = cast( # cast-ok: runtime Sequence check excludes scalar strings and bytes
Sequence[object], request_fields_value
)
else:
body_value: Final = proxy_request.get("body")
request_fields = (
tuple(cast(Mapping[object, object], body_value)) # cast-ok: runtime Mapping check establishes iterable keys
if isinstance(body_value, Mapping)
else ()
)
names: Final = frozenset(optional_params) | frozenset({"api_key", "api_base", "extra_headers"})
request_sources: Final = MappingProxyType(
{name: "request" for name in names if name in request_fields or name in credential_fields}
)
if litellm.enable_azure_ad_token_refresh is True and "enable_azure_ad_token_refresh" in optional_params:
return MappingProxyType({**request_sources, "enable_azure_ad_token_refresh": "deployment"})
return request_sources
def _marshal(
request: LiteLLMOcrRequest,
resolve_secret: Callable[[str], str | None],
convert_file_document: Callable[[dict[str, object]], dict[str, str]],
) -> LiteLLMOcrRequest:
if not isinstance(request.document, dict):
raise TypeError(f"document must be a dict with 'type' and URL/file field, got {type(request.document)}")
document: Final = (
convert_file_document(request.document) if request.document.get("type") == "file" else request.document
)
request_provider: Final = provider(request)
api_key: Final = (
request.api_key or resolve_secret("MISTRAL_API_KEY") if request_provider == "mistral" else request.api_key
)
optional_params: Final = _optional_params(request, resolve_secret)
input_sources: Final = _input_sources(request, optional_params)
logged_optional_params: Final = MappingProxyType(
{name: "****" if name in _RUST_OCR_SECRET_FIELDS else value for name, value in optional_params.items()}
)
logged_kwargs: Final = MappingProxyType(
{
name: "****" if name in _RUST_OCR_SECRET_FIELDS else value
for name, value in request.kwargs.items()
if name != "proxy_server_request"
}
)
logging_obj: Final = cast( # cast-ok: client decorator injects the logging object through untyped kwargs
_OCRLogging, request.kwargs["litellm_logging_obj"]
)
logging_obj.update_from_kwargs(
kwargs=dict(logged_kwargs), # mutable-ok: legacy logging mutates its kwargs copy
model=request.model,
optional_params=dict(logged_optional_params), # mutable-ok: legacy logging requires concrete dict params
litellm_params={ # mutable-ok: legacy logging requires a concrete params dict
"litellm_call_id": request.kwargs.get("litellm_call_id"),
"api_base": request.api_base,
},
custom_llm_provider=request_provider,
)
logging_obj.pre_call(
input="OCR document processing",
api_key=api_key,
additional_args={ # mutable-ok: pre_call mutates the additional_args dict
"complete_input_dict": { # mutable-ok: callbacks consume a JSON-serializable request dict
"model": request.model,
"document": document,
**logged_optional_params,
},
"api_base": request.api_base or "",
"headers": request.extra_headers or {}, # mutable-ok: logging callbacks consume a concrete headers dict
},
)
return LiteLLMOcrRequest(
model=request.model,
document=document,
api_key=api_key,
api_base=request.api_base,
timeout=request.timeout if request.timeout is not None else request_timeout,
custom_llm_provider=request.custom_llm_provider,
extra_headers=request.extra_headers,
kwargs=optional_params,
call_completion=request.call_completion,
input_sources=input_sources,
)
def _map_error(error: Exception, request: LiteLLMOcrRequest) -> Exception:
exception_types: Final = native_exception_types()
if exception_types is None or not isinstance(error, exception_types[1]):
return error
request_provider: Final = provider(request)
if request_provider is None:
return error
provider_config: Final = ProviderConfigManager.get_provider_ocr_config(
model=request.model.removeprefix(f"{request_provider}/"), provider=litellm.LlmProviders(request_provider)
)
if provider_config is None:
return error
error_args: Final = cast( # cast-ok: BaseException.args exposes Any while native errors carry scalar args
tuple[object, ...], error.args
)
status: Final = error_args[0] if error_args and isinstance(error_args[0], int) else 500
message: Final = str(error_args[1]) if len(error_args) > 1 else str(error)
error_factory: Final = cast( # cast-ok: legacy provider error factories have untyped callable parameters
Callable[..., Exception], provider_config.get_error_class
)
return error_factory(
error_message=message,
status_code=status or 500,
headers={}, # mutable-ok: provider error factories require a concrete headers dict
)
def _response(response: Mapping[str, object]) -> OCRResponse:
provider_native_response: Final = response.get(PROVIDER_NATIVE_RESPONSE_KEY)
normalized: Final = OCRResponse.model_validate(
@ -338,56 +101,6 @@ def _response(response: Mapping[str, object]) -> OCRResponse:
return normalized
def run(
request: LiteLLMOcrRequest,
resolve_secret: Callable[[str], str | None],
convert_file_document: Callable[[dict[str, object]], dict[str, str]],
) -> OCRResponse | None:
if load_rust_ocr() is None:
return None
marshalled: Final = _marshal(request, resolve_secret, convert_file_document)
try:
response: Final = ocr(
model=marshalled.model,
document=dict(marshalled.document), # mutable-ok: PyO3 OCR binding requires a concrete dict
api_key=marshalled.api_key,
api_base=marshalled.api_base,
custom_llm_provider=marshalled.custom_llm_provider,
extra_headers=marshalled.extra_headers,
optional_params=dict(marshalled.kwargs), # mutable-ok: PyO3 OCR binding requires a concrete dict
input_sources=marshalled.input_sources,
timeout=marshalled.timeout,
)
except Exception as error:
raise _map_error(error, request) from error
return _response(response) if response is not None else None
async def arun(
request: LiteLLMOcrRequest,
resolve_secret: Callable[[str], str | None],
convert_file_document: Callable[[dict[str, object]], dict[str, str]],
) -> OCRResponse | None:
if load_rust_aocr() is None:
return None
marshalled: Final = _marshal(request, resolve_secret, convert_file_document)
try:
response: Final = await aocr(
model=marshalled.model,
document=dict(marshalled.document), # mutable-ok: PyO3 OCR binding requires a concrete dict
api_key=marshalled.api_key,
api_base=marshalled.api_base,
custom_llm_provider=marshalled.custom_llm_provider,
extra_headers=marshalled.extra_headers,
optional_params=dict(marshalled.kwargs), # mutable-ok: PyO3 OCR binding requires a concrete dict
input_sources=marshalled.input_sources,
timeout=marshalled.timeout,
)
except Exception as error:
raise _map_error(error, request) from error
return _response(response) if response is not None else None
def ocr(
*,
model: str,

View file

@ -0,0 +1,121 @@
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
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.configuration import rust_enabled
from litellm.rust_bridge.ocr import LiteLLMOcrRequest
class NativeOcrLifecycle(Protocol):
def __call__(
self,
request: LiteLLMOcrRequest,
args: Sequence[object],
kwargs: Mapping[str, object],
asynchronous: bool,
) -> OCRResponse | Awaitable[OCRResponse]: ...
class ExceptionMapper(Protocol):
def __call__(
self,
*,
model: str,
custom_llm_provider: str | None,
original_exception: Exception,
completion_kwargs: dict[str, object],
extra_kwargs: dict[str, object],
) -> Exception: ...
def _binding(value: object) -> NativeOcrLifecycle | None:
if not callable(value):
return None
return cast("NativeOcrLifecycle", value) # cast-ok: callable validated at the native binding boundary
NATIVE_OCR_LIFECYCLE: Final = NativeBinding("_ocr_lifecycle", validate=_binding)
def select(request: LiteLLMOcrRequest) -> NativeOcrLifecycle | None:
if not rust_enabled():
return None
if litellm.cache is not None or request.kwargs.get("caching") or request.kwargs.get("aocr"):
return None
return NATIVE_OCR_LIFECYCLE.load()
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")
try:
token: Final = provider()
if not isinstance(token, str):
raise TypeError(f"Azure AD token must be a string, got {type(token)}")
return token
except TypeError:
raise
except Exception as error:
raise RuntimeError(f"Failed to get Azure AD token: {error}") from error
def map_failure(error: Exception, request: LiteLLMOcrRequest, request_provider: str) -> Exception:
if isinstance(error, ValueError) and "Invalid `req_format`" in str(error):
return litellm.BadRequestError(
message=str(error),
model=request.model,
llm_provider=request_provider,
)
mapper: Final = cast( # cast-ok: bounded adapter for the legacy public exception mapper
ExceptionMapper, litellm.exception_type
)
try:
return mapper(
model=request.model.removeprefix(f"{request_provider}/"),
custom_llm_provider=request_provider,
original_exception=error,
completion_kwargs=dict(arguments(request)), # mutable-ok: exception mapper requires owned kwargs
extra_kwargs=dict(request.kwargs), # mutable-ok: exception mapper requires owned kwargs
)
except Exception as public_error:
public_error.__context__ = error
return public_error

View file

@ -4,7 +4,6 @@ import datetime
import weakref
from collections.abc import Callable, Coroutine
from concurrent.futures import Future, ThreadPoolExecutor
from importlib import import_module
from threading import get_ident
from typing import Final
from unittest.mock import AsyncMock, MagicMock
@ -15,7 +14,6 @@ import litellm
from litellm.litellm_core_utils import thread_pool_executor
from litellm.litellm_core_utils.call_completion import CallCompletion, PythonCompletion
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.rust_bridge import ocr as rust_ocr_bridge
from litellm.utils import client
@ -313,7 +311,7 @@ async def test_async_ocr_wrapper_reports_metadata_failure_without_success(
@pytest.mark.parametrize("asynchronous", [False, True])
@pytest.mark.asyncio
async def test_ocr_completion_stays_separate_from_marshaled_provider_options(
async def test_wrapper_completion_stays_separate_from_provider_options(
monkeypatch: pytest.MonkeyPatch, asynchronous: bool
) -> None:
native_completion: Final = RecordingCompletion()
@ -321,45 +319,22 @@ async def test_ocr_completion_stays_separate_from_marshaled_provider_options(
metadata: Final = {"request": "shared"}
pages: Final = [0, 2]
def run(
request: rust_ocr_bridge.LiteLLMOcrRequest,
resolve_secret: Callable[[str], str | None],
convert_file_document: Callable[[dict[str, object]], dict[str, str]],
) -> OCRResponse:
assert request.kwargs["metadata"] is metadata
assert "_litellm_call_completion" not in request.kwargs
marshalled: Final = rust_ocr_bridge._marshal(request, resolve_secret, convert_file_document)
assert "_litellm_call_completion" not in marshalled.kwargs
assert marshalled.kwargs["pages"] is pages
assert marshalled.call_completion is request.call_completion
assert marshalled.call_completion is not None
assert marshalled.call_completion.attach(native_completion)
def ocr(*, _litellm_call_completion: CallCompletion, **kwargs: object) -> OCRResponse:
assert kwargs["metadata"] is metadata
assert kwargs["pages"] is pages
assert "_litellm_call_completion" not in kwargs
assert _litellm_call_completion.attach(native_completion)
return response
async def arun(
request: rust_ocr_bridge.LiteLLMOcrRequest,
resolve_secret: Callable[[str], str | None],
convert_file_document: Callable[[dict[str, object]], dict[str, str]],
) -> OCRResponse:
return run(request, resolve_secret, convert_file_document)
async def aocr(*, _litellm_call_completion: CallCompletion, **kwargs: object) -> OCRResponse:
return ocr(_litellm_call_completion=_litellm_call_completion, **kwargs)
def run_bridge(
request: rust_ocr_bridge.LiteLLMOcrRequest,
resolve_api_key: Callable[[str], str | None],
) -> OCRResponse:
return run(request, resolve_api_key, lambda document: {})
async def arun_bridge(
request: rust_ocr_bridge.LiteLLMOcrRequest,
resolve_api_key: Callable[[str], str | None],
) -> OCRResponse:
return await arun(request, resolve_api_key, lambda document: {})
ocr_main: Final = import_module("litellm.ocr.main")
monkeypatch.setattr(ocr_main, "rust_enabled", lambda: True)
monkeypatch.setattr(ocr_main, "_rust_ocr_supported", lambda request: True)
monkeypatch.setattr(ocr_main, "_run_rust_ocr", run_bridge)
monkeypatch.setattr(ocr_main, "_run_rust_aocr", arun_bridge)
monkeypatch.setattr(
"litellm.utils.function_setup",
MagicMock(return_value=(MagicMock(), {"metadata": metadata, "pages": pages})),
)
monkeypatch.setattr("litellm.utils.load_credentials_from_list", MagicMock())
wrapped: Final = client(aocr if asynchronous else ocr)
arguments: Final = {
"model": "mistral/mistral-ocr-latest",
"document": {"type": "document_url", "document_url": "https://example.com/doc.pdf"},
@ -368,7 +343,7 @@ async def test_ocr_completion_stays_separate_from_marshaled_provider_options(
"pages": pages,
}
result: Final = await litellm.aocr(**arguments) if asynchronous else litellm.ocr(**arguments)
result: Final = await wrapped(**arguments) if asynchronous else wrapped(**arguments)
assert result is response
assert native_completion.successes == [response]

View file

@ -1,73 +0,0 @@
"""
Regression tests for Azure Document Intelligence api_base ownership in OCR.
`azure_ai` exposes two OCR services on one provider; the `doc-intelligence`
sub-route must defer environment resolution to Rust, not accept the generic
`AZURE_AI_API_BASE` fallback that `get_llm_provider` injects. An explicitly
supplied api_base is still always honoured.
"""
from litellm.llms.azure_ai.ocr.common_utils import (
is_azure_document_intelligence_model,
)
from litellm.ocr.main import _prepare_ocr_request
_DOC = {"type": "document_url", "document_url": "https://example.com/doc.pdf"}
_AZURE_AI_API_BASE = "https://generic-azure-ai.example.com"
class _FakeLogging:
def update_from_kwargs(self, **kwargs: object) -> None:
return None
def _prepare(model: str, api_base: str | None):
return _prepare_ocr_request(
model=model,
document=dict(_DOC),
api_key="test-key",
api_base=api_base,
timeout=None,
custom_llm_provider=None,
extra_headers=None,
kwargs={"litellm_logging_obj": _FakeLogging()},
)
class TestIsAzureDocumentIntelligenceModel:
def test_matches_doc_intelligence_route(self):
assert is_azure_document_intelligence_model("doc-intelligence/prebuilt-layout")
def test_matches_documentintelligence_and_is_case_insensitive(self):
assert is_azure_document_intelligence_model("azure_ai/DocumentIntelligence/x")
def test_does_not_match_mistral_route(self):
assert not is_azure_document_intelligence_model("mistral-document-ai-2505")
class TestDocIntelligenceApiBaseResolution:
def test_generic_azure_ai_base_does_not_hijack_doc_intelligence(self, monkeypatch):
"""The generic Azure base must not overwrite Rust-owned DI resolution."""
monkeypatch.setenv("AZURE_AI_API_BASE", _AZURE_AI_API_BASE)
monkeypatch.delenv("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT", raising=False)
prepared = _prepare("azure_ai/doc-intelligence/prebuilt-layout", None)
assert prepared.api_base is None
def test_explicit_api_base_is_honoured_for_doc_intelligence(self, monkeypatch):
"""A caller-supplied api_base must always win, even for doc-intelligence."""
monkeypatch.setenv("AZURE_AI_API_BASE", _AZURE_AI_API_BASE)
custom = "https://my-di.cognitiveservices.azure.com"
prepared = _prepare("azure_ai/doc-intelligence/prebuilt-layout", custom)
assert prepared.api_base == custom
def test_generic_azure_ai_base_still_applies_to_mistral_ocr(self, monkeypatch):
"""Non doc-intelligence azure_ai models keep using AZURE_AI_API_BASE."""
monkeypatch.setenv("AZURE_AI_API_BASE", _AZURE_AI_API_BASE)
prepared = _prepare("azure_ai/mistral-document-ai-2505", None)
assert prepared.api_base == _AZURE_AI_API_BASE

View file

@ -20,7 +20,7 @@ import orjson
import pytest
from starlette.datastructures import FormData
from litellm.ocr.main import convert_file_document_to_url_document, get_mime_type
from litellm.ocr.input import convert_file_document_to_url_document, get_mime_type
class TestGetMimeType:

View file

@ -2,37 +2,7 @@
Tests for the OCR `req_format` option in the SDK request path.
"""
import pytest
import litellm
from litellm.rust_bridge import ocr as rust_ocr_bridge
from litellm.rust_bridge.ocr import LiteLLMOcrRequest
DOCUMENT = {"type": "document_url", "document_url": "https://example.com/doc.pdf"}
def _request(
optional_params: dict[str, object], model: str = "azure_ai/doc-intelligence/prebuilt-layout"
) -> LiteLLMOcrRequest:
return LiteLLMOcrRequest(
model=model,
document=DOCUMENT,
api_key="fake-key",
api_base=None,
custom_llm_provider=None,
extra_headers=None,
timeout=60.0,
kwargs=optional_params,
)
@pytest.mark.parametrize("optional_params", [{}, {"req_format": "litellm"}])
def test_rust_ocr_serves_default_format(optional_params):
assert rust_ocr_bridge.supported(_request(optional_params)) is True
def test_rust_ocr_serves_native_format_for_document_intelligence():
assert rust_ocr_bridge.supported(_request({"req_format": "native"})) is True
def test_rust_ocr_response_retains_provider_native_response():
@ -50,34 +20,3 @@ def test_rust_ocr_response_retains_provider_native_response():
assert response.get_provider_native_response() == provider_response
assert response.model_dump().get("provider_native_response") is None
@pytest.mark.parametrize("model", ["cohere/cohere-parse", "azure_ai/cohere-parse"])
def test_rust_ocr_skipped_for_unsupported_models(model):
assert rust_ocr_bridge.supported(_request({}, model)) is False
@pytest.mark.asyncio
async def test_native_format_rejected_for_provider_without_support_as_bad_request():
with pytest.raises(litellm.BadRequestError, match="not supported for provider") as exc_info:
await litellm.aocr(
model="mistral/mistral-ocr-latest",
document=DOCUMENT,
api_key="fake-key",
req_format="native",
)
assert exc_info.value.status_code == 400
@pytest.mark.asyncio
async def test_unknown_format_rejected_for_provider_without_support_as_bad_request():
with pytest.raises(litellm.BadRequestError, match="Invalid `req_format`") as exc_info:
await litellm.aocr(
model="mistral/mistral-ocr-latest",
document=DOCUMENT,
api_key="fake-key",
req_format="raw",
)
assert exc_info.value.status_code == 400

File diff suppressed because it is too large Load diff

View file

@ -1,6 +1,8 @@
"""Tests for unified guardrail."""
import logging
from types import SimpleNamespace
from typing import Final
import pytest
@ -19,14 +21,14 @@ from litellm.llms.base_llm.guardrail_translation.utils import (
openai_messages_without_system,
openai_messages_without_tool,
)
from litellm.llms.base_llm.ocr.transformation import OCRPage, OCRResponse
from litellm.llms.mistral.ocr.guardrail_translation.handler import OCRHandler
from litellm.llms.openai.chat.guardrail_translation.handler import (
OpenAIChatCompletionsHandler,
)
from litellm.llms.openai.responses.guardrail_translation.handler import (
OpenAIResponsesHandler,
)
from litellm.llms.base_llm.ocr.transformation import OCRPage, OCRResponse
from litellm.llms.mistral.ocr.guardrail_translation.handler import OCRHandler
from litellm.proxy._experimental.mcp_server.guardrail_translation.handler import (
MCPGuardrailTranslationHandler,
)
@ -644,6 +646,64 @@ class TestUnifiedLLMGuardrails:
class TestOCRGuardrailE2E:
"""End-to-end tests: UnifiedLLMGuardrails -> OCRHandler."""
@pytest.mark.asyncio
@pytest.mark.parametrize("call_type", [CallTypes.ocr, CallTypes.aocr, CallTypes.aresponses])
async def test_post_call_logging_fallback_is_limited_to_ocr(self, call_type: CallTypes) -> None:
guardrail: Final = RecordingGuardrail()
response: Final = (
TestUnifiedLLMGuardrails.TestResponsesRouteAliases._responses_api_response()
if call_type == CallTypes.aresponses
else OCRResponse(model="mistral-ocr-latest", pages=[OCRPage(index=0, markdown="Scan this page")])
)
result: Final = await UnifiedLLMGuardrails().async_post_call_success_hook(
data={
"guardrail_to_apply": guardrail,
"litellm_logging_obj": SimpleNamespace(call_type=call_type.value),
},
user_api_key_dict=UserAPIKeyAuth(),
response=response,
)
assert result is response
if call_type in (CallTypes.ocr, CallTypes.aocr):
assert len(guardrail.apply_calls) == 1
assert guardrail.apply_calls[0]["inputs"]["texts"] == ["Scan this page"]
else:
assert guardrail.apply_calls == []
@pytest.mark.asyncio
@pytest.mark.parametrize("request_route", [None, "/v1/chat/completions"])
async def test_ocr_logging_fallback_preserves_route_and_response_precedence(
self, request_route: str | None, monkeypatch: pytest.MonkeyPatch
) -> None:
from litellm.types.utils import ModelResponse
_patch_translation_mappings(
monkeypatch,
{
CallTypes.completion: OpenAIChatCompletionsHandler,
CallTypes.acompletion: OpenAIChatCompletionsHandler,
CallTypes.aocr: OCRHandler,
},
)
guardrail: Final = RecordingGuardrail()
response: Final = ModelResponse(choices=[{"message": {"role": "assistant", "content": "Chat output"}}])
result: Final = await guardrail.async_post_call_success_deployment_hook(
request_data={
"guardrails": [guardrail.guardrail_name],
"user_api_key_request_route": request_route,
"litellm_logging_obj": SimpleNamespace(call_type=CallTypes.aocr.value),
},
response=response,
call_type=CallTypes.aocr,
)
assert result is response
assert len(guardrail.apply_calls) == 1
assert guardrail.apply_calls[0]["inputs"]["texts"] == ["Chat output"]
@pytest.mark.asyncio
async def test_pre_call_hook_invokes_ocr_handler_for_input(self):
"""

View file

@ -0,0 +1,125 @@
from collections.abc import Mapping
from typing import Final
from unittest.mock import Mock
import pytest
import litellm
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.rust_bridge.ocr import LiteLLMOcrRequest
from litellm.rust_bridge.ocr_lifecycle import NATIVE_OCR_LIFECYCLE
@pytest.mark.parametrize("enabled,available", [(True, False), (False, True), (False, False)])
def test_public_selection_requires_supported_native_ocr(enabled: bool, available: bool) -> None:
native: Final = Mock(side_effect=AssertionError("must not admit"))
litellm.rust(enabled)
NATIVE_OCR_LIFECYCLE.override(native if available else None)
try:
with pytest.raises(RuntimeError, match="Rust OCR is unavailable or does not support this request"):
litellm.ocr("mistral/mistral-ocr-latest", {"type": "document_url", "document_url": "https://example.com"})
finally:
NATIVE_OCR_LIFECYCLE.reset()
litellm.rust(None)
assert native.call_count == 0
def test_admitted_failure_is_returned_without_replay() -> None:
failure: Final = RuntimeError("admitted")
native: Final = Mock(side_effect=failure)
litellm.rust(True)
NATIVE_OCR_LIFECYCLE.override(native)
try:
with pytest.raises(RuntimeError) as caught:
litellm.ocr("mistral/mistral-ocr-latest", {"type": "document_url", "document_url": "https://example.com"})
assert caught.value is failure
finally:
NATIVE_OCR_LIFECYCLE.reset()
litellm.rust(None)
assert native.call_count == 1
def test_public_binding_keeps_positional_fields_and_defaults_out_of_native_hook_kwargs() -> None:
document: Final = {"type": "document_url", "document_url": "https://example.com"}
captured: Final = []
def native(
request: LiteLLMOcrRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
asynchronous: bool,
) -> OCRResponse:
captured.append((request, args, kwargs, asynchronous))
return OCRResponse(pages=[], model=request.model)
litellm.rust(True)
NATIVE_OCR_LIFECYCLE.override(native)
try:
response: Final = litellm.ocr("mistral/mistral-ocr-latest", document)
finally:
NATIVE_OCR_LIFECYCLE.reset()
litellm.rust(None)
request, call_args, hook_kwargs, asynchronous = captured[0]
assert response.model == "mistral/mistral-ocr-latest"
assert request.model == "mistral/mistral-ocr-latest"
assert request.document is document
assert call_args == ("mistral/mistral-ocr-latest", document)
assert hook_kwargs == {}
assert asynchronous is False
def test_public_binding_keeps_keyword_model_and_document_in_native_hook_kwargs() -> None:
document: Final = {"type": "document_url", "document_url": "https://example.com"}
captured: Final = []
def native(
request: LiteLLMOcrRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
asynchronous: bool,
) -> OCRResponse:
assert args == ()
captured.append(kwargs)
return OCRResponse(pages=[], model=request.model)
litellm.rust(True)
NATIVE_OCR_LIFECYCLE.override(native)
try:
litellm.ocr(model="mistral/mistral-ocr-latest", document=document)
finally:
NATIVE_OCR_LIFECYCLE.reset()
litellm.rust(None)
assert captured[0]["model"] == "mistral/mistral-ocr-latest"
assert captured[0]["document"] is document
assert "timeout" not in captured[0]
@pytest.mark.parametrize("enabled", [False, True], ids=["legacy", "native"])
def test_public_duplicate_argument_error_does_not_depend_on_native_selection(enabled: bool) -> None:
native: Final = Mock(side_effect=AssertionError("binding errors precede admission"))
document: Final = {"type": "document_url", "document_url": "https://example.com"}
litellm.rust(enabled)
NATIVE_OCR_LIFECYCLE.override(native)
try:
with pytest.raises(TypeError, match=r"ocr\(\) got multiple values for argument 'model'"):
litellm.ocr("mistral/mistral-ocr-latest", document, model="duplicate")
finally:
NATIVE_OCR_LIFECYCLE.reset()
litellm.rust(None)
assert native.call_count == 0
@pytest.mark.parametrize("enabled", [False, True], ids=["legacy", "native"])
def test_public_missing_required_argument_error_does_not_depend_on_native_selection(enabled: bool) -> None:
native: Final = Mock(side_effect=AssertionError("binding errors precede admission"))
litellm.rust(enabled)
NATIVE_OCR_LIFECYCLE.override(native)
try:
with pytest.raises(TypeError, match=r"ocr\(\) missing 1 required positional argument: 'document'"):
litellm.ocr("mistral/mistral-ocr-latest")
finally:
NATIVE_OCR_LIFECYCLE.reset()
litellm.rust(None)
assert native.call_count == 0

View file

@ -4,10 +4,22 @@ This suite covers OCR requests through LiteLLM's compiled Rust extension. OCR be
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 OCR callback inputs, mutations, ordering, context, failure handling, concurrency, and cleanup. `ocr/test_guardrails.py` covers OCR post-call blocking and response replacement. These contract modules call the Rust bridge directly. `ocr/test_dispatch.py` has the single public API dispatch test, covering enabled native dispatch and disabled Python dispatch. `test_ocr.py` is a strict smoke test of the compiled Rust OCR transport
`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 callback and configuration state but does not select a backend. Native contract tests call `litellm.rust_bridge.ocr` directly, while the strict dispatch test explicitly enables and disables Rust and records which OCR entrypoint runs
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 OCR contract modules are non-strict expected failures until the retained callback implementation from #40070 lands. The public dispatch test remains strict. Passing contract cases appear as XPASS so staging coverage stays visible
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

@ -29,11 +29,6 @@ CALLBACK_ATTRIBUTES: Final = (
"_async_success_callback",
"_async_failure_callback",
)
EXPECTED_FAILURE_REASONS: Final = {
"ocr/test_callbacks.py": "requires the OCR callback lifecycle implementation from #40070",
"ocr/test_guardrails.py": "requires the OCR guardrail lifecycle implementation from #40070",
"ocr/test_requests.py": "requires the OCR request and Azure authentication implementation from #40070",
}
def _list_attribute(container: ModuleType, attribute: str) -> list[object]:
@ -76,6 +71,7 @@ async def isolate_ocr_test_state() -> AsyncIterator[None]:
stack.enter_context(_rebound(litellm, "cache", None)) # test-quality-ok: isolate process-global cache
stack.enter_context(_rebound(_CONFIGURATION, "override", None))
executor: Final = ThreadPoolExecutor(thread_name_prefix="rust-ocr-test-logging")
stack.enter_context(_rebound(litellm_logging, "executor", executor))
stack.enter_context(_rebound(utils, "executor", executor))
stack.enter_context(_rebound(thread_pool_executor, "executor", executor))
try:
@ -95,14 +91,6 @@ def recording_server() -> Generator[RecordingServer]:
def pytest_collection_modifyitems(items: list[pytest.Item]) -> None:
for item in items:
if "test_litellm_rust" not in item.path.parts:
continue
relative_path: Final = "/".join(item.path.parts[item.path.parts.index("test_litellm_rust") + 1 :])
reason: Final = EXPECTED_FAILURE_REASONS.get(relative_path)
if reason is not None:
item.add_marker(pytest.mark.xfail(reason=reason, strict=False))
if not _parse_env_bool(os.environ.get("LITELLM_RUST")):
skip: Final = pytest.mark.skip(reason="requires LITELLM_RUST=1 and a compiled Rust extension")
for item in items:

View file

@ -41,7 +41,7 @@ def test_native_ocr_pre_call_callback_receives_transformed_provider_request(ocr_
observations: Final = []
class Observe(CustomLogger):
def log_pre_api_call(self, model, _messages, kwargs):
def log_pre_api_call(self, model, messages, kwargs):
observations.append((model, copy.deepcopy(kwargs["additional_args"])))
call_native_ocr_with_callbacks(ocr_server, [Observe()], pages=[0])
@ -64,13 +64,13 @@ def test_native_ocr_pre_call_body_edit_reaches_next_callback_and_provider(
observed: Final = []
class Edit(CustomLogger):
def log_pre_api_call(self, model, _messages, kwargs):
def log_pre_api_call(self, model, messages, kwargs):
request_body(kwargs)["include_image_base64"] = True
if raise_after_edit:
raise RuntimeError("pre-call callback failed")
class Observe(CustomLogger):
def log_pre_api_call(self, model, _messages, kwargs):
def log_pre_api_call(self, model, messages, kwargs):
observed.append(copy.deepcopy(request_body(kwargs)))
call_native_ocr_with_callbacks(ocr_server, [Edit(), Observe()], include_image_base64=False)
@ -83,11 +83,11 @@ def test_native_ocr_pre_call_header_edit_reaches_next_callback_and_provider(ocr_
observed: Final = []
class Edit(CustomLogger):
def log_pre_api_call(self, model, _messages, kwargs):
def log_pre_api_call(self, model, messages, kwargs):
request_headers(kwargs)["x-audit-tag"] = "reviewed"
class Observe(CustomLogger):
def log_pre_api_call(self, model, _messages, kwargs):
def log_pre_api_call(self, model, messages, kwargs):
observed.append(dict(request_headers(kwargs)))
call_native_ocr_with_callbacks(ocr_server, [Edit(), Observe()])
@ -96,6 +96,29 @@ def test_native_ocr_pre_call_header_edit_reaches_next_callback_and_provider(ocr_
assert ocr_server.requests[0].headers["x-audit-tag"] == "reviewed"
def test_native_ocr_pre_call_header_rebinding_does_not_replace_execution_root(ocr_server: RecordingServer) -> None:
retained: Final = []
observed: Final = []
class RetainMutateAndRebind(CustomLogger):
def log_pre_api_call(self, model, messages, kwargs):
headers = request_headers(kwargs)
retained.append(headers)
kwargs["additional_args"]["headers"] = {"x-rebound": "not-sent"}
headers["x-retained"] = "sent"
class ObserveRebinding(CustomLogger):
def log_pre_api_call(self, model, messages, kwargs):
observed.append(dict(request_headers(kwargs)))
call_native_ocr_with_callbacks(ocr_server, [RetainMutateAndRebind(), ObserveRebinding()])
assert observed == [{"x-rebound": "not-sent"}]
assert retained[0]["x-retained"] == "sent"
assert ocr_server.requests[0].headers["x-retained"] == "sent"
assert "x-rebound" not in ocr_server.requests[0].headers
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"])
async def test_native_ocr_pre_call_nested_document_edit_updates_caller_callback_and_provider_references(
@ -107,12 +130,12 @@ async def test_native_ocr_pre_call_nested_document_edit_updates_caller_callback_
aliases: Final = []
class Retain(CustomLogger):
def log_pre_api_call(self, model, _messages, kwargs):
def log_pre_api_call(self, model, messages, kwargs):
aliases.append(request_body(kwargs)["document"] is original)
retained.append(request_body(kwargs)["document"])
class Edit(CustomLogger):
def log_pre_api_call(self, model, _messages, kwargs):
def log_pre_api_call(self, model, messages, kwargs):
original["document_url"] = replacement_url
arguments: Final = {
@ -143,7 +166,7 @@ def test_native_ocr_pre_call_document_replacement_does_not_mutate_original_docum
retained: Final = []
class RetainAndReplace(CustomLogger):
def log_pre_api_call(self, model, _messages, kwargs):
def log_pre_api_call(self, model, messages, kwargs):
body = request_body(kwargs)
retained.append(body["document"])
body["document"] = replacement
@ -165,11 +188,11 @@ def test_native_ocr_pre_call_body_rebinding_is_visible_to_callbacks_but_not_prov
observed: Final = []
class Rebind(CustomLogger):
def log_pre_api_call(self, model, _messages, kwargs):
def log_pre_api_call(self, model, messages, kwargs):
kwargs["additional_args"]["complete_input_dict"] = {"replacement": True}
class Observe(CustomLogger):
def log_pre_api_call(self, model, _messages, kwargs):
def log_pre_api_call(self, model, messages, kwargs):
observed.append(request_body(kwargs))
call_native_ocr_with_callbacks(ocr_server, [Rebind(), Observe()])
@ -182,11 +205,11 @@ def test_native_ocr_callback_retained_body_observes_later_callback_mutation(ocr_
queued: Final = []
class QueuePayload(CustomLogger):
def log_pre_api_call(self, model, _messages, kwargs):
def log_pre_api_call(self, model, messages, kwargs):
queued.append(request_body(kwargs))
class Edit(CustomLogger):
def log_pre_api_call(self, model, _messages, kwargs):
def log_pre_api_call(self, model, messages, kwargs):
request_body(kwargs)["queued-edit"] = True
call_native_ocr_with_callbacks(ocr_server, [QueuePayload(), Edit()])
@ -200,7 +223,7 @@ def test_native_ocr_success_callback_receives_state_added_by_pre_call_callback(o
finished: Final = threading.Event()
class Stash(CustomLogger):
def log_pre_api_call(self, model, _messages, kwargs):
def log_pre_api_call(self, model, messages, kwargs):
kwargs["test-token"] = token
def log_success_event(self, kwargs, response_obj, start_time, end_time):
@ -280,7 +303,7 @@ async def test_native_aocr_failure_callbacks_receive_state_added_by_pre_call_cal
observed: Final = []
class TrackInFlightRequest(CustomLogger):
def log_pre_api_call(self, model, _messages, kwargs):
def log_pre_api_call(self, model, messages, kwargs):
kwargs["request-token"] = token
def log_failure_event(self, kwargs, response_obj, start_time, end_time):
@ -364,7 +387,7 @@ async def test_native_azure_ocr_resolves_token_before_pre_call_on_caller_context
return "caller-token"
class Edit(CustomLogger):
def log_pre_api_call(self, model, _messages, kwargs):
def log_pre_api_call(self, model, messages, kwargs):
assert request_headers(kwargs)["Authorization"] == "Bearer caller-token"
observations.append("pre_call")
request_headers(kwargs)["Authorization"] = "Bearer edited"

View file

@ -1,11 +1,9 @@
from typing import Final
from unittest.mock import Mock
import pytest
import litellm
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.ocr import main as ocr_main
from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec
from tests.test_litellm_rust.support.requests import OCR_DOCUMENT, OCR_MODEL, OCR_RESPONSE
@ -18,18 +16,8 @@ def ocr_server(recording_server: RecordingServer) -> RecordingServer:
return recording_server
@pytest.mark.parametrize("rust_enabled", [True, False], ids=["enabled", "disabled"])
def test_public_ocr_dispatches_according_to_rust_setting(
ocr_server: RecordingServer,
monkeypatch: pytest.MonkeyPatch,
rust_enabled: bool,
) -> None:
rust_call: Final = Mock(wraps=ocr_main.rust_ocr_bridge.ocr)
python_call: Final = Mock(wraps=ocr_main.base_llm_http_handler.ocr)
monkeypatch.setattr(ocr_main.rust_ocr_bridge, "ocr", rust_call)
monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", python_call)
litellm.rust(rust_enabled)
def test_public_ocr_uses_native_route_when_enabled(ocr_server: RecordingServer) -> None:
litellm.rust(True)
response: Final = litellm.ocr(
model=OCR_MODEL,
document=OCR_DOCUMENT,
@ -39,6 +27,20 @@ def test_public_ocr_dispatches_according_to_rust_setting(
assert isinstance(response, OCRResponse)
assert response.pages[0].markdown == "native OCR response"
assert rust_call.call_count == int(rust_enabled)
assert python_call.call_count == int(not rust_enabled)
assert len(ocr_server.requests) == 1
assert not ocr_server.requests[0].headers.get("user-agent", "").startswith("python-httpx")
def test_public_ocr_fails_before_network_when_native_is_disabled(ocr_server: RecordingServer) -> None:
litellm.rust(False)
ocr_server.expected_requests = 0
with pytest.raises(RuntimeError, match="Rust OCR is unavailable"):
litellm.ocr(
model=OCR_MODEL,
document=OCR_DOCUMENT,
api_key="test-key",
api_base=ocr_server.base_url,
)
assert ocr_server.requests == []

View file

@ -0,0 +1,769 @@
import asyncio
import datetime
import gc
import sys
import threading
import weakref
from collections.abc import Coroutine
from contextvars import ContextVar
from typing import Final
import pytest
import litellm
from litellm._logging import trace_id_var
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from tests.test_litellm_rust.support.callback_recorder import RecordingLogger, drain_logging
from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec
from tests.test_litellm_rust.support.requests import OCR_RESPONSE, call_aocr, call_ocr
pytestmark = pytest.mark.requires_rust_extension
@pytest.mark.asyncio
@pytest.mark.parametrize("phase", ["deployment", "failure"])
async def test_cancellation_during_failure_obeys_phase_policy(ocr_server: RecordingServer, phase: str) -> None:
ocr_server.enqueue(ResponseSpec(body={"message": "provider failure"}, status=500))
entered: Final = asyncio.Event()
observed: Final = []
class Observer(CustomLogger):
async def async_post_call_failure_deployment_hook(self, request_data, exception, call_type, **kwargs):
if phase == "deployment":
entered.set()
await asyncio.Event().wait()
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
observed.append(kwargs["exception"])
if phase == "failure":
entered.set()
await asyncio.Event().wait()
observer: Final = Observer()
litellm.callbacks.append(observer)
task: Final = asyncio.create_task(call_aocr(ocr_server, callbacks=[observer]))
await asyncio.wait_for(entered.wait(), 5)
task.cancel()
if phase == "deployment":
with pytest.raises(litellm.InternalServerError) as caught:
await task
assert observed == [caught.value]
else:
with pytest.raises(asyncio.CancelledError):
await task
assert len(observed) == 1
assert isinstance(observed[0], litellm.InternalServerError)
@pytest.fixture
def ocr_server(recording_server: RecordingServer) -> RecordingServer:
recording_server.default_response = ResponseSpec(body=OCR_RESPONSE)
return recording_server
@pytest.mark.asyncio
async def test_proxy_metadata_remains_python_owned(ocr_server: RecordingServer) -> None:
from litellm.proxy._types import UserAPIKeyAuth
recorder: Final = RecordingLogger()
auth: Final = UserAPIKeyAuth(user_id="ocr-user")
response: Final = await call_aocr(
ocr_server, callbacks=[recorder], metadata={"user_api_key_auth": auth}, shared_session=object()
)
events: Final = await recorder.wait_for_async("async_log_success_event")
assert response.pages[0].markdown == "native OCR response"
assert events[0].kwargs["litellm_params"]["metadata"]["user_api_key_auth"].user_id == "ocr-user"
assert "metadata" not in ocr_server.requests[0].body
@pytest.mark.asyncio
async def test_response_replacement_finalized_before_dispatch_in_caller_task(ocr_server: RecordingServer) -> None:
caller: Final = asyncio.current_task()
context: Final = ContextVar("lifecycle-test", default="before")
observations: Final = []
recorder: Final = RecordingLogger()
class Replace(CustomLogger):
async def async_pre_call_deployment_hook(self, kwargs, call_type):
context.set("pre")
observations.append(("pre", asyncio.current_task(), context.get()))
return {**kwargs, "pages": [2]}
async def async_post_call_success_deployment_hook(self, request_data, response, call_type):
observations.append(("post", asyncio.current_task(), context.get()))
return response.model_copy(update={"model": "replaced"})
litellm.callbacks.append(Replace())
response: Final = await call_aocr(ocr_server, callbacks=[recorder], litellm_call_id="native-final")
events: Final = await recorder.wait_for_async("async_log_success_event")
assert observations == [("pre", caller, "pre"), ("post", caller, "pre")]
assert context.get() == "pre"
assert ocr_server.requests[0].body["pages"] == [2]
assert response.model == "replaced"
assert events[0].response is response
assert response._hidden_params["litellm_call_id"] == "native-final"
assert "response_cost" in response._hidden_params
@pytest.mark.asyncio
async def test_deployment_hook_replaces_complete_routing_request(ocr_server: RecordingServer) -> None:
ocr_server.enqueue(ResponseSpec(body=OCR_RESPONSE, delay=0.05))
original: Final = {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}
replacement: Final = {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}
observed: Final = []
class Replace(CustomLogger):
async def async_pre_call_deployment_hook(self, kwargs, call_type):
return {
**kwargs,
"model": "azure_ai/mistral-ocr-latest",
"custom_llm_provider": "azure_ai",
"document": replacement,
"api_key": "replacement-key",
"api_base": ocr_server.base_url,
"extra_headers": {"x-deployment": "replacement"},
"timeout": 2,
"pages": [2],
}
class Observe(Logging):
def pre_call(self, input, api_key, additional_args):
observed.append((additional_args["complete_input_dict"]["document"], api_key))
litellm.callbacks.append(Replace())
logger: Final = Observe(
model="mistral-ocr-latest",
messages=[],
stream=False,
call_type="aocr",
start_time=datetime.datetime.now(),
litellm_call_id="deployment-routing",
function_id="deployment-routing",
)
response: Final = await call_aocr(
ocr_server,
document=original,
timeout=0.001,
litellm_logging_obj=logger,
)
assert response.pages[0].markdown == "native OCR response"
assert observed == [(replacement, "replacement-key")]
assert observed[0][0] is replacement
assert replacement == original
assert replacement is not original
assert original == {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}
assert ocr_server.requests[0].path == "/providers/mistral/azure/ocr"
assert ocr_server.requests[0].headers["authorization"] == "Bearer replacement-key"
assert ocr_server.requests[0].headers["x-deployment"] == "replacement"
assert ocr_server.requests[0].body["document"] == replacement
assert ocr_server.requests[0].body["pages"] == [2]
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
async def test_metadata_failure_dispatches_only_failure_and_releases_logger(
ocr_server: RecordingServer, asynchronous: bool
) -> None:
failure: Final = RuntimeError("metadata failed")
seen: Final = []
class FailingMetadata(Logging):
def _response_cost_calculator(self, *args, **kwargs):
raise failure
def success_handler(self, *args, **kwargs):
seen.append("success")
def failure_handler(self, exception, *args, **kwargs):
seen.append(("sync", exception))
async def async_failure_handler(self, exception, *args, **kwargs):
seen.append(("async", exception))
async def invoke():
logger: Final = FailingMetadata(
model="mistral-ocr-latest",
messages=[],
stream=False,
call_type="aocr" if asynchronous else "ocr",
start_time=datetime.datetime.now(),
litellm_call_id="metadata",
function_id="metadata",
)
reference: Final = weakref.ref(logger)
with pytest.raises(RuntimeError) as caught:
await call_aocr(ocr_server, litellm_logging_obj=logger) if asynchronous else call_ocr(
ocr_server, litellm_logging_obj=logger
)
assert caught.value is failure
failure.__traceback__ = None
return reference
reference: Final = await invoke()
await drain_logging()
gc.collect()
assert seen == ([("sync", failure), ("async", failure)] if asynchronous else [("sync", failure)])
assert reference() is None
assert len(ocr_server.requests) == 1
@pytest.mark.asyncio
async def test_mapped_failure_identity_and_deployment_snapshot(ocr_server: RecordingServer) -> None:
ocr_server.enqueue(ResponseSpec(body={"message": "unavailable"}, status=500))
recorder: Final = RecordingLogger()
snapshots: Final = []
class Observe(CustomLogger):
async def async_post_call_failure_deployment_hook(self, request_data, exception, call_type, **kwargs):
snapshots.append(exception)
exception.status_code = 418
litellm.callbacks.append(Observe())
with pytest.raises(litellm.InternalServerError) as caught:
await call_aocr(ocr_server, callbacks=[recorder])
failures: Final = tuple(event for event in recorder.events if "failure" in event.name)
assert [event.name for event in failures] == ["log_failure_event", "async_log_failure_event"]
assert all(event.kwargs["exception"] is caught.value for event in failures)
assert caught.value.status_code == 500
assert snapshots[0] is not caught.value
assert snapshots[0].status_code == 418
assert len(ocr_server.requests) == 1
@pytest.mark.asyncio
@pytest.mark.parametrize("phase", ["pre", "http", "post"])
async def test_cancellation_cleans_up_in_caller_task_without_terminal_dispatch(
ocr_server: RecordingServer, phase: str
) -> None:
entered: Final = asyncio.Event()
recorder: Final = RecordingLogger()
class Pause(CustomLogger):
async def async_pre_call_deployment_hook(self, kwargs, call_type):
if phase == "pre":
entered.set()
await asyncio.Event().wait()
async def async_post_call_success_deployment_hook(self, request_data, response, call_type):
if phase == "post":
entered.set()
await asyncio.Event().wait()
litellm.callbacks.append(Pause())
if phase == "http":
ocr_server.enqueue(ResponseSpec(body=OCR_RESPONSE, delay=0.2))
if phase == "pre":
ocr_server.expected_requests = 0
restored: Final = []
async def invoke():
trace_id_var.set("parent")
try:
await call_aocr(ocr_server, callbacks=[recorder], litellm_trace_id="native-call")
finally:
restored.append(trace_id_var.get())
task: Final = asyncio.create_task(invoke())
if phase == "http":
await ocr_server.wait_for_requests(1)
else:
await asyncio.wait_for(entered.wait(), 5)
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
await drain_logging()
assert restored == ["parent"]
assert not any("success" in name or "failure" in name for name in recorder.names)
@pytest.mark.asyncio
@pytest.mark.parametrize("blocked", [False, True])
async def test_deferred_logging_requires_release_and_runs_at_most_once(
ocr_server: RecordingServer, blocked: bool
) -> None:
recorder: Final = RecordingLogger()
logger: Final = Logging(
model="mistral-ocr-latest",
messages=[],
stream=False,
call_type="aocr",
start_time=datetime.datetime.now(),
litellm_call_id="deferred",
function_id="deferred",
dynamic_async_success_callbacks=[recorder],
)
logger._defer_async_logging = True
response: Final = await call_aocr(ocr_server, litellm_logging_obj=logger)
await drain_logging()
assert "async_log_success_event" not in recorder.names
ProxyBaseLLMRequestProcessing._flush_deferred_async_logging(logger, blocked)
ProxyBaseLLMRequestProcessing._flush_deferred_async_logging(logger, blocked)
await drain_logging()
events: Final = tuple(event for event in recorder.events if event.name == "async_log_success_event")
assert len(events) == int(not blocked)
if events:
assert events[0].response is response
@pytest.mark.asyncio
@pytest.mark.parametrize("failure", [RuntimeError("native enqueue failed"), asyncio.CancelledError("cancelled")])
async def test_deferred_release_handles_enqueue_failure_once_without_replay(
ocr_server: RecordingServer, monkeypatch: pytest.MonkeyPatch, failure: BaseException
) -> None:
import inspect
from litellm.litellm_core_utils import logging_worker
attempts: Final[list[Coroutine[object, object, object]]] = []
diagnostics: Final = []
class FailingWorker:
def ensure_initialized_and_enqueue(self, coroutine: Coroutine[object, object, object]) -> None:
attempts.append(coroutine)
raise failure
recorder: Final = RecordingLogger()
logger: Final = Logging(
model="mistral-ocr-latest",
messages=[],
stream=False,
call_type="aocr",
start_time=datetime.datetime.now(),
litellm_call_id="release-failure",
function_id="release-failure",
dynamic_async_success_callbacks=[recorder],
)
logger._defer_async_logging = True
response: Final = await call_aocr(ocr_server, litellm_logging_obj=logger)
monkeypatch.setattr(logging_worker, "GLOBAL_LOGGING_WORKER", FailingWorker())
monkeypatch.setattr(sys, "unraisablehook", lambda event: diagnostics.append(event.exc_value))
if isinstance(failure, asyncio.CancelledError):
with pytest.raises(asyncio.CancelledError, match="cancelled") as caught:
ProxyBaseLLMRequestProcessing._flush_deferred_async_logging(logger, False)
assert caught.value is failure
assert diagnostics == []
else:
ProxyBaseLLMRequestProcessing._flush_deferred_async_logging(logger, False)
assert diagnostics == [failure]
ProxyBaseLLMRequestProcessing._flush_deferred_async_logging(logger, False)
assert len(attempts) == 1
assert inspect.getcoroutinestate(attempts[0]) == inspect.CORO_CLOSED
assert response.pages[0].markdown == "native OCR response"
assert len(ocr_server.requests) == 1
assert not any("success" in name or "failure" in name for name in recorder.names)
@pytest.mark.asyncio
async def test_abandoned_deferred_logging_is_collectable(ocr_server: RecordingServer) -> None:
async def invoke():
logger: Final = Logging(
model="mistral-ocr-latest",
messages=[],
stream=False,
call_type="aocr",
start_time=datetime.datetime.now(),
litellm_call_id="abandoned",
function_id="abandoned",
)
logger._defer_async_logging = True
await call_aocr(ocr_server, litellm_logging_obj=logger)
return weakref.ref(logger)
reference: Final = await invoke()
await drain_logging()
gc.collect()
assert reference() is None
def test_sync_success_uses_executor_and_copied_caller_context(ocr_server: RecordingServer) -> None:
context: Final = ContextVar("sync-lifecycle", default="missing")
context.set("caller")
thread: Final = threading.current_thread()
finished: Final = threading.Event()
observations: Final = []
class Observe(CustomLogger):
def log_success_event(self, kwargs, response_obj, start_time, end_time):
observations.append((threading.current_thread(), context.get(), response_obj))
finished.set()
response: Final = call_ocr(ocr_server, callbacks=[Observe()])
assert finished.wait(5)
assert observations[0][0] is not thread
assert observations[0][1] == "caller"
assert observations[0][2] is response
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
async def test_invalid_response_runs_post_call_before_failure(ocr_server: RecordingServer, asynchronous: bool) -> None:
ocr_server.enqueue(ResponseSpec(body={"pages": "invalid"}))
events: Final = []
class Observe(Logging):
def pre_call(self, *args, **kwargs):
events.append("pre")
return super().pre_call(*args, **kwargs)
def post_call(self, *args, **kwargs):
events.append(("post", kwargs["original_response"]))
return super().post_call(*args, **kwargs)
def success_handler(self, *args, **kwargs):
events.append("success")
def failure_handler(self, exception, *args, **kwargs):
events.append(("failure", exception))
async def async_failure_handler(self, exception, *args, **kwargs):
events.append(("async_failure", exception))
logger: Final = Observe(
model="mistral-ocr-latest",
messages=[],
stream=False,
call_type="aocr" if asynchronous else "ocr",
start_time=datetime.datetime.now(),
litellm_call_id="invalid",
function_id="invalid",
)
with pytest.raises(litellm.APIConnectionError) as caught:
await call_aocr(ocr_server, litellm_logging_obj=logger) if asynchronous else call_ocr(
ocr_server, litellm_logging_obj=logger
)
assert events[0] == "pre"
assert events[1] == ("post", '{"pages": "invalid"}')
assert events[2] == ("failure", caught.value)
if asynchronous:
assert events[3] == ("async_failure", caught.value)
assert "success" not in events
@pytest.mark.asyncio
async def test_failing_terminal_handler_preserves_public_failure_and_runs_async_handler(
ocr_server: RecordingServer,
) -> None:
ocr_server.enqueue(ResponseSpec(body={"message": "provider failure"}, status=500))
failures: Final = []
class BrokenHandler(Logging):
def failure_handler(self, exception, *args, **kwargs):
failures.append(exception)
raise RuntimeError("handler failed")
async def async_failure_handler(self, exception, *args, **kwargs):
failures.append(exception)
logger: Final = BrokenHandler(
model="mistral-ocr-latest",
messages=[],
stream=False,
call_type="aocr",
start_time=datetime.datetime.now(),
litellm_call_id="broken",
function_id="broken",
)
with pytest.raises(litellm.InternalServerError) as caught:
await call_aocr(ocr_server, litellm_logging_obj=logger)
assert failures == [caught.value, caught.value]
assert len(ocr_server.requests) == 1
@pytest.mark.asyncio
async def test_nested_native_calls_preserve_context_and_dispatch_each_outcome(ocr_server: RecordingServer) -> None:
ocr_server.expected_requests = 2
recorder: Final = RecordingLogger()
outcomes: Final = []
class Nested(CustomLogger):
async def async_pre_call_deployment_hook(self, kwargs, call_type):
if kwargs.get("litellm_call_id") == "outer":
outcomes.append(await call_aocr(ocr_server, callbacks=[recorder], litellm_call_id="inner"))
litellm.callbacks.append(Nested())
outcomes.append(await call_aocr(ocr_server, callbacks=[recorder], litellm_call_id="outer"))
events: Final = await recorder.wait_for_async("async_log_success_event", count=2)
assert [event.kwargs["litellm_call_id"] for event in events] == ["inner", "outer"]
assert events[0].response is outcomes[0]
assert events[1].response is outcomes[1]
assert len(ocr_server.requests) == 2
def test_sync_pre_call_can_make_nested_native_request(ocr_server: RecordingServer) -> None:
ocr_server.expected_requests = 2
observed: Final = []
class Nested(CustomLogger):
def log_pre_api_call(self, model, messages, kwargs):
if kwargs["litellm_call_id"] == "outer-sync":
observed.append(call_ocr(ocr_server, litellm_call_id="inner-sync"))
response: Final = call_ocr(ocr_server, callbacks=[Nested()], litellm_call_id="outer-sync")
assert observed[0].pages[0].markdown == response.pages[0].markdown
assert len(ocr_server.requests) == 2
@pytest.mark.asyncio
async def test_retained_argument_aliases_and_body_roots_survive_envelope_replacement(
ocr_server: RecordingServer,
) -> None:
pages: Final = [0]
document: Final = {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}
opaque: Final = object()
observed: Final = []
class Observe(Logging):
def pre_call(self, input, api_key, additional_args):
body: Final = additional_args["complete_input_dict"]
headers: Final = additional_args["headers"]
observed.append((body["document"] is document, body["pages"] is pages))
pages.append(2)
headers["x-retained"] = "yes"
additional_args["complete_input_dict"] = {"discarded": True}
additional_args["headers"] = {}
observed.append((body, headers))
def post_call(self, original_response, additional_args):
observed.append(
(additional_args["complete_input_dict"] is observed[2][0], additional_args["headers"] is observed[2][1])
)
class Deployment(CustomLogger):
async def async_pre_call_deployment_hook(self, kwargs, call_type):
observed.append(("model" in kwargs, "document" in kwargs, kwargs["opaque"] is opaque))
litellm.callbacks.append(Deployment())
logger: Final = Observe(
model="mistral-ocr-latest",
messages=[],
stream=False,
call_type="aocr",
start_time=datetime.datetime.now(),
litellm_call_id="roots",
function_id="roots",
)
response: Final = await litellm.aocr(
"mistral/mistral-ocr-latest",
document,
api_key="test-key",
api_base=ocr_server.base_url,
pages=pages,
opaque=opaque,
litellm_logging_obj=logger,
)
assert response.pages[0].markdown == "native OCR response"
assert observed[0] == (False, False, True)
assert observed[1] == (True, True)
assert observed[3] == (True, True)
assert ocr_server.requests[0].body["pages"] == [0, 2]
assert ocr_server.requests[0].headers["x-retained"] == "yes"
def test_unstarted_native_coroutine_releases_input_without_reading_file(ocr_server: RecordingServer) -> None:
from litellm.ocr.main import _public_request
from litellm.rust_bridge import _native
ocr_server.expected_requests = 0
effects: Final = []
class File:
def read(self):
effects.append("read")
return b"abc"
def create():
file: Final = File()
kwargs: Final = {"model": "mistral/mistral-ocr-latest", "document": {"type": "file", "file": file}}
coroutine: Final = _native._ocr_lifecycle(_public_request("aocr", (), kwargs), (), kwargs, True)
file.owner = coroutine
coroutine.close()
return weakref.ref(file)
reference: Final = create()
gc.collect()
assert reference() is None
assert effects == []
@pytest.mark.asyncio
async def test_file_read_happens_after_deployment_hook_in_caller_task(ocr_server: RecordingServer) -> None:
effects: Final = []
caller: Final = asyncio.current_task()
class File:
def read(self):
effects.append(("read", asyncio.current_task()))
return b"abc"
class Deployment(CustomLogger):
async def async_pre_call_deployment_hook(self, kwargs, call_type):
await asyncio.sleep(0)
effects.append(("hook", asyncio.current_task()))
litellm.callbacks.append(Deployment())
await call_aocr(ocr_server, document={"type": "file", "file": File()})
assert effects == [("hook", caller), ("read", caller)]
@pytest.mark.asyncio
async def test_failure_callbacks_continue_within_both_families(ocr_server: RecordingServer) -> None:
ocr_server.enqueue(ResponseSpec(body={"message": "failed"}, status=500))
observed: Final = []
class Broken(CustomLogger):
def log_failure_event(self, kwargs, response_obj, start_time, end_time):
observed.append(("broken-sync", kwargs["exception"]))
raise RuntimeError("sync observer")
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
observed.append(("broken-async", kwargs["exception"]))
raise RuntimeError("async observer")
class Following(CustomLogger):
def log_failure_event(self, kwargs, response_obj, start_time, end_time):
observed.append(("following-sync", kwargs["exception"]))
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
observed.append(("following-async", kwargs["exception"]))
with pytest.raises(litellm.InternalServerError) as caught:
await call_aocr(ocr_server, callbacks=[Broken(), Following()])
assert [name for name, _ in observed] == ["broken-sync", "following-sync", "broken-async", "following-async"]
assert all(error is caught.value for _, error in observed)
@pytest.mark.asyncio
async def test_cancelling_native_transport_closes_connection_before_return() -> None:
received: Final = asyncio.Event()
disconnected: Final = asyncio.Event()
async def provider(reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None:
headers: Final = await reader.readuntil(b"\r\n\r\n")
length: Final = next(
int(line.split(b":", 1)[1])
for line in headers.split(b"\r\n")
if line.lower().startswith(b"content-length:")
)
await reader.readexactly(length)
received.set()
assert await reader.read() == b""
disconnected.set()
writer.close()
await writer.wait_closed()
server: Final = await asyncio.start_server(provider, "127.0.0.1", 0)
async with server:
port: Final = server.sockets[0].getsockname()[1]
task: Final = asyncio.create_task(
litellm.aocr(
model="mistral/mistral-ocr-latest",
document={"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"},
api_key="test-key",
api_base=f"http://127.0.0.1:{port}",
)
)
await asyncio.wait_for(received.wait(), 5)
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
await asyncio.wait_for(disconnected.wait(), 1)
@pytest.mark.asyncio
@pytest.mark.parametrize("model", ["reducto/parse-v3", "reducto/parse-legacy"])
async def test_reducto_lifecycle_retains_upload_parse_and_post_call_boundaries(
ocr_server: RecordingServer, model: str
) -> None:
ocr_server.expected_requests = 2
ocr_server.enqueue(ResponseSpec(body={"file_id": "reducto://uploaded.pdf"}))
ocr_server.enqueue(ResponseSpec(body={"result": {"chunks": [{"content": "parsed"}]}}))
boundaries: Final = []
recorder: Final = RecordingLogger()
class Observe(Logging):
def post_call(self, *args, **kwargs):
boundaries.append(tuple(request.path for request in ocr_server.requests))
return super().post_call(*args, **kwargs)
logger: Final = Observe(
model=model,
messages=[],
stream=False,
call_type="aocr",
start_time=datetime.datetime.now(),
litellm_call_id="upload",
function_id="upload",
dynamic_async_success_callbacks=[recorder],
)
response: Final = await call_aocr(ocr_server, model=model, litellm_logging_obj=logger)
events: Final = await recorder.wait_for_async("async_log_success_event")
assert boundaries == [("/upload", "/parse")]
assert b"abc" in ocr_server.requests[0].raw_body
assert "multipart/form-data" in ocr_server.requests[0].headers["content-type"]
assert ocr_server.requests[1].body["input" if model.endswith("v3") else "document_url"] == "reducto://uploaded.pdf"
assert response.pages[0].markdown == "parsed"
assert events[0].response is response
@pytest.mark.asyncio
async def test_document_intelligence_post_call_runs_before_polling(ocr_server: RecordingServer) -> None:
ocr_server.expected_requests = 2
ocr_server.enqueue(
ResponseSpec(
body={"status": "running"},
status=202,
headers={"Operation-Location": f"{ocr_server.base_url}/operations/1", "Retry-After": "0"},
)
)
ocr_server.enqueue(ResponseSpec(body={"status": "succeeded", "analyzeResult": {"pages": []}}))
boundaries: Final = []
class Observe(Logging):
def post_call(self, *args, **kwargs):
boundaries.append(tuple(request.method for request in ocr_server.requests))
return super().post_call(*args, **kwargs)
logger: Final = Observe(
model="azure_ai/doc-intelligence/prebuilt-read",
messages=[],
stream=False,
call_type="aocr",
start_time=datetime.datetime.now(),
litellm_call_id="poll",
function_id="poll",
)
response: Final = await call_aocr(
ocr_server, model="azure_ai/doc-intelligence/prebuilt-read", litellm_logging_obj=logger
)
assert boundaries == [("POST",)]
assert [request.method for request in ocr_server.requests] == ["POST", "GET"]
assert ocr_server.requests[1].path == "/operations/1"
assert response.pages == []
@pytest.mark.asyncio
async def test_vertex_deepseek_public_lifecycle_normalizes_before_success(ocr_server: RecordingServer) -> None:
ocr_server.enqueue(
ResponseSpec(body={"choices": [{"message": {"content": "recognized"}}], "usage": {"prompt_tokens": 1}})
)
recorder: Final = RecordingLogger()
response: Final = await call_aocr(
ocr_server,
model="vertex_ai/deepseek-ocr-maas",
document={"type": "document_url", "document_url": "gs://bucket/document.pdf"},
vertex_project="project-1",
vertex_location="europe-west4",
callbacks=[recorder],
)
events: Final = await recorder.wait_for_async("async_log_success_event")
assert response.pages[0].markdown == "recognized"
assert events[0].response is response
assert (
ocr_server.requests[0].path
== "/v1/projects/project-1/locations/europe-west4/endpoints/openapi/chat/completions"
)

View file

@ -1,3 +1,4 @@
from pathlib import Path
from typing import Final
import pytest
@ -79,6 +80,22 @@ def test_native_ocr_prepares_file_document_like_python(ocr_server: RecordingServ
}
def test_native_ocr_reads_sdk_path_input(ocr_server: RecordingServer, tmp_path: Path) -> None:
document_path: Final = tmp_path / "document.pdf"
document_path.write_bytes(b"%PDF-1.4")
response: Final = call_native_ocr(
ocr_server,
document={"type": "file", "file": document_path},
)
assert response.pages[0].markdown == "native OCR response"
assert ocr_server.requests[0].body["document"] == {
"type": "document_url",
"document_url": "data:application/pdf;base64,JVBERi0xLjQ=",
}
def test_native_ocr_sends_pages_and_image_options(ocr_server: RecordingServer) -> None:
call_native_ocr(ocr_server, pages=[0, 2], include_image_base64=True)
@ -149,7 +166,7 @@ def test_native_ocr_normalizes_provider_response_model_and_usage(ocr_server: Rec
assert response.usage_info.pages_processed == 1
def test_native_ocr_maps_provider_400_without_exposing_response_body(ocr_server: RecordingServer) -> None:
def test_native_ocr_maps_provider_400_with_public_provider_details(ocr_server: RecordingServer) -> None:
ocr_server.enqueue(ResponseSpec(body={"message": "invalid OCR request"}, status=400))
with pytest.raises(litellm.BadRequestError) as caught:
@ -158,13 +175,23 @@ def test_native_ocr_maps_provider_400_without_exposing_response_body(ocr_server:
assert caught.value.status_code == 400
assert caught.value.model == "mistral-ocr-latest"
assert caught.value.llm_provider == "mistral"
assert "invalid OCR request" not in str(caught.value)
assert "invalid OCR request" in str(caught.value)
def test_native_ocr_raises_transport_error_when_request_exceeds_timeout(ocr_server: RecordingServer) -> None:
def test_native_ocr_rejects_unknown_response_format_before_provider_request(ocr_server: RecordingServer) -> None:
ocr_server.expected_requests = 0
with pytest.raises(litellm.BadRequestError, match="Invalid `req_format`"):
call_native_ocr(ocr_server, req_format="raw")
assert ocr_server.requests == []
def test_ocr_raises_public_timeout_when_request_exceeds_timeout(ocr_server: RecordingServer) -> None:
litellm.rust(True)
ocr_server.enqueue(ResponseSpec(body=OCR_RESPONSE, delay=0.2))
with pytest.raises(RuntimeError, match="OCR transport failed"):
with pytest.raises(litellm.Timeout):
call_native_ocr(ocr_server, timeout=0.01)
assert len(ocr_server.requests) == 1
@ -301,13 +328,10 @@ async def test_native_azure_ocr_token_provider_failure_prevents_pre_call_callbac
@pytest.mark.parametrize(
"configuration",
[
{"azure_ad_token": "oidc/assertion", "client_id": "client", "tenant_id": "tenant"},
{"model": "azure_ai/doc-intelligence/prebuilt-read"},
],
ids=["oidc-assertion", "document-intelligence-model"],
[{"azure_ad_token": "oidc/assertion", "client_id": "client", "tenant_id": "tenant"}],
ids=["invalid-oidc-assertion"],
)
def test_native_azure_ocr_rejects_unsupported_configuration_before_token_or_callbacks(
def test_public_azure_ocr_maps_invalid_oidc_configuration_before_token_or_request(
ocr_server: RecordingServer,
isolated_azure_auth: None,
configuration: dict[str, object],
@ -327,10 +351,10 @@ def test_native_azure_ocr_rejects_unsupported_configuration_before_token_or_call
"callbacks": [recorder],
**configuration,
}
with pytest.raises(NotImplementedError):
with pytest.raises(litellm.APIConnectionError):
call_native_ocr(ocr_server, **arguments)
assert calls == []
assert recorder.events == ()
assert "log_pre_api_call" not in recorder.names
assert ocr_server.requests == []

View file

@ -81,7 +81,7 @@ class RecordingLogger(CustomLogger):
await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), timeout=timeout)
return tuple(event for event in self.events if event.name == name)
def log_pre_api_call(self, model, _messages, kwargs):
def log_pre_api_call(self, model, messages, kwargs):
self._record("log_pre_api_call", kwargs)
def log_success_event(self, kwargs, response_obj, start_time, end_time):

View file

@ -58,7 +58,9 @@ def recording_service() -> Iterator[RecordingServer]:
def _handle(self) -> None:
content_length: Final = int(self.headers.get("Content-Length", "0"))
raw_body: Final = self.rfile.read(content_length) if content_length else b""
body: Final = json.loads(raw_body) if raw_body else None
body: Final = (
json.loads(raw_body) if raw_body and self.headers.get_content_type() == "application/json" else None
)
requests.append(
RecordedRequest(
method=self.command,
@ -84,6 +86,7 @@ def recording_service() -> Iterator[RecordingServer]:
pass
do_POST = _handle
do_GET = _handle
def log_message(self, format: str, *args: object) -> None:
pass

View file

@ -2,7 +2,6 @@ from typing import Final
import litellm
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.rust_bridge import ocr as native_ocr
from tests.test_litellm_rust.support.recording_server import RecordingServer
OCR_DOCUMENT: Final = {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}
@ -36,11 +35,11 @@ async def call_aocr(server: RecordingServer, **kwargs: object) -> OCRResponse:
def call_native_ocr(server: RecordingServer, **kwargs: object) -> OCRResponse:
return native_ocr.ocr(ocr_arguments(server, **kwargs))
return call_ocr(server, **kwargs)
async def call_native_aocr(server: RecordingServer, **kwargs: object) -> OCRResponse:
return await native_ocr.aocr(ocr_arguments(server, **kwargs))
return await call_aocr(server, **kwargs)
def request_body(kwargs: dict[str, object]) -> dict[str, object]:

View file

@ -2,6 +2,7 @@ import json
import threading
from collections.abc import Generator
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from io import BytesIO
from typing import Final
import pytest
@ -99,6 +100,38 @@ def test_native_ocr_with_compiled_rust_extension(
}
@pytest.mark.parametrize(
"file_input,mime_type,expected_type,expected_field,expected_uri",
[
(b"abc", "application/pdf", "document_url", "document_url", "data:application/pdf;base64,YWJj"),
(BytesIO(b"abc"), "image/png", "image_url", "image_url", "data:image/png;base64,YWJj"),
],
)
def test_native_lifecycle_core_encodes_python_file_input(
ocr_server,
file_input,
mime_type,
expected_type,
expected_field,
expected_uri,
):
server, requests = ocr_server
litellm.rust(True)
response = litellm.ocr(
model="mistral/mistral-ocr-latest",
document={"type": "file", "file": file_input, "mime_type": mime_type},
api_key="test-key",
api_base=f"http://127.0.0.1:{server.server_port}",
opaque_extension=object(),
)
assert response.pages[0].markdown == "native OCR response"
assert requests[0]["body"]["document"] == {
"type": expected_type,
expected_field: expected_uri,
}
assert "opaque_extension" not in requests[0]["body"]
@pytest.mark.parametrize("asynchronous", [False, True])
@pytest.mark.parametrize("model", ["mistral/mistral-ocr-latest", "azure_ai/doc-intelligence/prebuilt-read"])
@pytest.mark.asyncio
@ -145,24 +178,20 @@ async def test_native_public_ocr_matches_python(model, asynchronous):
server: Final = ThreadingHTTPServer(("127.0.0.1", 0), Handler)
thread: Final = Thread(target=server.serve_forever, daemon=True)
thread.start()
responses: Final = []
try:
for enabled in (False, True):
litellm.rust(enabled)
arguments: Final = {
"model": model,
"document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"},
"api_key": "test-key",
"api_base": f"http://127.0.0.1:{server.server_port}",
"pages": [0, 2],
"timeout": 3.0,
}
response: Final = await litellm.aocr(**arguments) if asynchronous else litellm.ocr(**arguments)
responses.append(response.model_dump())
assert len(calls) == 2
assert calls[0] == calls[1]
for key in ("model", "pages", "object"):
assert responses[0][key] == responses[1][key]
litellm.rust(True)
arguments: Final = {
"model": model,
"document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"},
"api_key": "test-key",
"api_base": f"http://127.0.0.1:{server.server_port}",
"pages": [0, 2],
"timeout": 3.0,
}
response: Final = await litellm.aocr(**arguments) if asynchronous else litellm.ocr(**arguments)
response_data: Final = response.model_dump()
assert len(calls) == 1
assert response_data["object"] == "ocr"
finally:
server.shutdown()
server.server_close()
@ -224,7 +253,7 @@ async def test_native_ocr_enforces_request_deadline_without_fallback(ocr_server,
"num_retries": 0,
}
started = time.monotonic()
with pytest.raises(litellm.APIConnectionError):
with pytest.raises(litellm.Timeout):
await asyncio.wait_for(
litellm.aocr(**arguments) if asynchronous else asyncio.to_thread(litellm.ocr, **arguments),
timeout=3,