This commit is contained in:
Yujong Lee 2026-09-12 07:40:04 -07:00
parent 3b9fc16f6b
commit 43a19d81ab
40 changed files with 1096 additions and 1122 deletions

3
.gitignore vendored
View file

@ -147,3 +147,6 @@ crash.*.log
ui/litellm-dashboard/out/
litellm.log
.coverage-rust
coverage-rust.xml

View file

@ -2,6 +2,6 @@ litellm-core is the LiteLLM SDK in Rust — it makes the LLM call. Each top-leve
A route module owns everything the call needs: types, the provider template trait, provider transforms (under `providers/`), provider/auth/URL resolution, and the handler that performs the HTTP call. Handlers belong here, never in a host crate.
Not here: serving HTTP (axum routes, extractors), config file reading, rollout state, databases, or callback dispatch. Env reads are limited to credential fallback in a route's `prepare.rs`.
Not here: serving HTTP (axum routes, extractors), config file reading, rollout state, databases, or host-specific callback execution. Core owns lifecycle sequencing and callback payload construction; hosts execute the selected integrations. Env reads are limited to credential fallback in a route's `prepare.rs`.
Routes (messages, ocr, realtime) and providers (anthropic, mistral, openai) are modules, not crates.

View file

@ -1,66 +0,0 @@
# CLAUDE.md
Rules for `litellm-rust/crates/core`.
## Responsibility
`core` is the LiteLLM SDK in Rust: it makes the LLM call. Every top-level
LiteLLM call has a public entrypoint here, named after the route
(`messages::messages()` is the Rust equivalent of `litellm.messages()`), and
calling it returns a typed non-streaming response.
Allowed:
- The public entrypoint for a route, plus its `<route>_stream` variant when the
route supports streaming.
- Provider resolution, auth header construction, URL building, and the provider
HTTP call (shared reused client, connect + request timeouts).
- Shared request/response structs.
- Typed errors with stable, non-sensitive messages.
- Deterministic validation helpers.
- Serialization helpers that intentionally mirror Python output shape.
- Route templates that match Python base config responsibilities, such as
`messages::transformation::AnthropicMessagesProviderConfig`.
Not allowed:
- Serving HTTP: axum routers, extractors, and other transport concerns.
- Filesystem, database, or cache access.
- Config file reading or rollout state; the host resolves those and passes them
in. Env reads are limited to credential fallback in a route's `prepare.rs`.
- Logging callbacks, tracing spans, spend writes, or customer callbacks.
- Provider-specific branching that belongs in `providers`.
- Panics for user/provider-controlled input.
## Typed Contracts (core rule)
Trait and function boundaries MUST be strongly typed. No stringly-typed JSON
(`&str` / `String` / `Vec<String>` / bare `serde_json::Value`) as a transform
input or output. Parse wire bytes into typed structs/enums at the host edge;
`core` and `providers` operate only on those types (e.g. `RealtimeEvent`,
`RealtimeTransformResult`, `OcrRequestData`). A `type`-style discriminator is a
typed field on a struct, not a raw string threaded through the API.
## Structure
Use route names directly under `src/`: `messages`, `ocr`, future
`chat_completions`, `embeddings`, and similar top-level LiteLLM calls. Do not
invent broad names like `engine` for route contracts.
`src/messages` is the reference shape for a route module:
```
mod.rs pub async fn messages(..) (+ messages_stream)
types.rs request/response types
transformation.rs the provider template trait
prepare.rs provider resolution, auth headers, URL
handler.rs the provider call
client.rs the shared reqwest client
```
## Parity Rules
- Every shared type used by a provider transform needs unit tests for
serialization shape.
- If Python parity requires always emitting a `null` field instead of omitting
it, document that in code and pin it with a test.
- Error enums should preserve enough detail for Python/HTTP hosts to map errors
consistently without exposing document contents or upstream bodies.

View file

@ -63,3 +63,6 @@ pub(crate) const REDUCTO_API_KEY_ENV: &str = "REDUCTO_API_KEY";
pub(crate) const REDUCTO_ID_PREFIX: &str = "reducto://";
pub(crate) const AZURE_AI_OCR_PATH: &str = "/providers/mistral/azure/ocr";
pub(crate) const MISTRAL_OCR_API_BASE: &str = "https://api.mistral.ai/v1";
pub(crate) const COHERE_PARSE_API_BASE: &str = "https://api.cohere.com";
pub(crate) const COHERE_API_KEY_ENV: &str = "COHERE_API_KEY";

View file

@ -52,6 +52,15 @@ pub enum Error {
Unsupported(&'static str),
}
impl Error {
pub const fn http_status_code(&self) -> Option<u16> {
match self {
Self::InvalidRequest(_) => Some(400),
_ => None,
}
}
}
#[derive(Debug, ThisError)]
pub(crate) enum MediaError {
#[error("media URL rejected by network policy")]

View file

@ -0,0 +1,131 @@
use super::super::OcrAdapter;
use crate::Error;
use crate::ocr::OcrClient;
use crate::ocr::codecs::cohere::{
CohereParams, CohereResponse, transform_request, transform_response, validate_document,
};
use crate::ocr::document::{inline_remote_document, validate_inline_document};
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};
use crate::providers::azure_ai::auth::AzureAuthInputs;
use crate::url_utils::ApiUrl;
const AZURE_AI_API_BASE_ENV: &str = "AZURE_AI_API_BASE";
pub(crate) struct AzureCohereAdapter;
impl OcrAdapter for AzureCohereAdapter {
type ProviderResponse = CohereResponse;
const PROVIDER: OcrProvider = OcrProvider::AzureAi;
async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, OcrError> {
let params = super::super::super::wire::decode_request_value::<CohereParams>(
serde_json::Value::Object(request.optional_params.clone()),
"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 base = request
.connection
.api_base
.clone()
.or_else(|| credential_env(AZURE_AI_API_BASE_ENV))
.filter(|base| !base.trim().is_empty())
.ok_or_else(|| {
Error::Auth(
"Missing Azure AI API Base - Set AZURE_AI_API_BASE or pass api_base".into(),
)
})?;
let headers =
super::validate_ai_environment(&request.connection, &config, &credential_env).await?;
validate_document(&request.document)?;
let remote = request.document.source().starts_with("http://")
|| request.document.source().starts_with("https://");
let document = inline_remote_document(
client.document_fetcher(),
request.document.clone(),
&request.connection,
)
.await?;
let body = transform_request(&request.model, document, params)?;
transform_request_body(
client,
request,
&complete_url(&base)?,
&headers,
!remote,
body,
|body| {
validate_document(&body.document)?;
validate_inline_document(&body.document)
},
)
.await
}
fn transform_ocr_response(
&self,
request: &LiteLLMOcrRequest,
response: Self::ProviderResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
transform_response(&request.model, response)
}
}
fn complete_url(base: &str) -> Result<String, OcrError> {
let mut url = reqwest::Url::parse(base).map_err(|_| invalid_api_base())?;
if !matches!(url.scheme(), "http" | "https") {
return Err(invalid_api_base().into());
}
let path = url.path().trim_end_matches('/').to_string();
if path.ends_with("/v2/parse") {
url.set_path(&path);
return Ok(url.into());
}
url.set_path(path.strip_suffix("/models").unwrap_or(&path));
ApiUrl::parse(url.as_str())
.and_then(|url| url.complete_path(&["providers", "cohere", "v2", "parse"]))
.map(|url| url.into_string())
.map_err(|_| invalid_api_base().into())
}
fn invalid_api_base() -> OcrRequestError {
OcrRequestError::RequestField {
path: "api_base".into(),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn completes_foundry_urls_without_duplicate_paths_and_preserves_queries() {
for suffix in [
"",
"/models",
"/providers/cohere/v2",
"/providers/cohere/v2/parse",
] {
assert_eq!(
complete_url(&format!("https://example.com{suffix}?tenant=a")).unwrap(),
"https://example.com/providers/cohere/v2/parse?tenant=a"
);
}
assert_eq!(
complete_url("https://example.com/v2/parse?tenant=a").unwrap(),
"https://example.com/v2/parse?tenant=a"
);
assert!(complete_url("relative/path").is_err());
}
}

View file

@ -61,7 +61,7 @@ impl OcrAdapter for AzureDocumentIntelligenceAdapter {
url: &str,
headers: &[(String, String)],
request: &LiteLLMOcrRequest,
) -> Result<Vec<u8>, OcrError> {
) -> Result<crate::ocr::wire::DecodedOcrResponse<Self::ProviderResponse>, OcrError> {
polling::read_operation_response(
client.polling_http(),
response,
@ -72,7 +72,6 @@ impl OcrAdapter for AzureDocumentIntelligenceAdapter {
&request.hooks,
)
.await
.map(|decoded| decoded.text.into_bytes())
}
}

View file

@ -44,7 +44,7 @@ pub(super) async fn read_operation_response(
}
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
poll_operation(http_client, operation, headers, connection, native, hooks).await
}
async fn poll_operation(
@ -53,6 +53,7 @@ async fn poll_operation(
headers: &[(String, String)],
connection: &OcrConnection,
native: bool,
hooks: &Arc<dyn OcrHooks>,
) -> Result<DecodedOcrResponse<AzureDocumentIntelligenceOperation>, OcrError> {
let deadline = Instant::now()
.checked_add(connection.poll_timeout)
@ -88,7 +89,10 @@ async fn poll_operation(
.await
.map_err(|_| OcrPollingError::PollTimeout)??;
match &decoded.data.status {
Some(OperationStatus::Succeeded) => return Ok(decoded),
Some(OperationStatus::Succeeded) => {
crate::ocr::handler::post_call(hooks, decoded.text.as_bytes()).await?;
return Ok(decoded);
}
Some(OperationStatus::Running | OperationStatus::NotStarted) => {
tokio::time::timeout_at(deadline, tokio::time::sleep(Duration::from_secs(retry)))
.await

View file

@ -92,7 +92,7 @@ fn get_complete_url(
})
}
async fn validate_environment(
pub(in crate::ocr::adapters) async fn validate_environment(
connection: &OcrConnection,
config: &AzureAuthInputs,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),

View file

@ -1,3 +1,4 @@
mod cohere;
mod document_intelligence;
mod mistral;
@ -10,8 +11,10 @@ use crate::ocr::error::OcrError;
use crate::ocr::types::OcrConnection;
use crate::providers::azure_ai::auth::{AzureAuthInputs, AzureAuthService};
pub(crate) use cohere::AzureCohereAdapter;
pub(crate) use document_intelligence::AzureDocumentIntelligenceAdapter;
pub(crate) use mistral::AzureMistralAdapter;
pub(super) use mistral::validate_environment as validate_ai_environment;
async fn resolve_entra(
config: &AzureAuthInputs,
@ -22,6 +25,10 @@ async fn resolve_entra(
.get_or_init(AzureAuthService::default)
.get_azure_ad_token(config, env_lookup)
.await
.or_else(|error| match error {
crate::AuthError::EmptyAzureToken => Ok(None),
other => Err(other),
})
.map(|credential| {
credential.map(|credential| {
let source = credential.source();

View file

@ -0,0 +1,123 @@
use super::OcrAdapter;
use crate::Error;
use crate::constants::{COHERE_API_KEY_ENV, COHERE_PARSE_API_BASE};
use crate::ocr::OcrClient;
use crate::ocr::codecs::cohere::{
CohereParams, CohereResponse, transform_request, transform_response, validate_document,
};
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};
use crate::url_utils::ApiUrl;
pub(crate) struct CohereAdapter;
impl OcrAdapter for CohereAdapter {
type ProviderResponse = CohereResponse;
const PROVIDER: OcrProvider = OcrProvider::Cohere;
async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, OcrError> {
let params = super::super::wire::decode_request_value::<CohereParams>(
serde_json::Value::Object(request.optional_params.clone()),
"optional_params",
)?;
let headers = validate_environment(&request.connection, &credential_env)?;
let url = complete_url(
request
.connection
.api_base
.as_deref()
.unwrap_or(COHERE_PARSE_API_BASE),
)?;
let body = transform_request(&request.model, request.document.clone(), params)?;
transform_request_body(client, request, &url, &headers, true, body, |body| {
validate_document(&body.document)
})
.await
}
fn transform_ocr_response(
&self,
request: &LiteLLMOcrRequest,
response: Self::ProviderResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
transform_response(&request.model, response)
}
}
fn complete_url(base: &str) -> Result<String, OcrError> {
let parsed = reqwest::Url::parse(base).map_err(|_| invalid_api_base())?;
if !matches!(parsed.scheme(), "http" | "https") {
return Err(invalid_api_base().into());
}
ApiUrl::parse(base)
.and_then(|url| url.complete_path(&["v2", "parse"]))
.map(|url| url.into_string())
.map_err(|_| invalid_api_base().into())
}
fn invalid_api_base() -> OcrRequestError {
OcrRequestError::RequestField {
path: "api_base".into(),
}
}
fn validate_environment(
connection: &OcrConnection,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Vec<(String, String)>, OcrError> {
if crate::http_utils::has_header(&connection.extra_headers, "authorization") {
return Ok(connection.extra_headers.clone());
}
let key = connection
.api_key
.as_deref()
.map(str::trim)
.filter(|key| !key.is_empty())
.map(str::to_string)
.or_else(|| env_lookup(COHERE_API_KEY_ENV).filter(|key| !key.trim().is_empty()))
.ok_or_else(|| {
Error::Auth("Missing COHERE_API_KEY - set it in the environment or pass api_key".into())
})?;
Ok(
std::iter::once(("Authorization".into(), format!("Bearer {key}")))
.chain(connection.extra_headers.clone())
.collect(),
)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn completes_provider_urls_without_duplicate_paths_and_preserves_queries() {
for suffix in ["", "/v2", "/v2/parse"] {
assert_eq!(
complete_url(&format!("https://example.com{suffix}?tenant=a")).unwrap(),
"https://example.com/v2/parse?tenant=a"
);
}
}
#[test]
fn rejects_invalid_urls_and_blank_keys() {
assert!(complete_url("relative/path").is_err());
assert!(complete_url("ftp://example.com").is_err());
assert!(matches!(
validate_environment(
&OcrConnection {
api_key: Some(" ".into()),
..Default::default()
},
&|_| None,
),
Err(OcrError::Public(Error::Auth(_)))
));
}
}

View file

@ -8,11 +8,13 @@ use super::registry::OcrProvider;
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
mod azure;
mod cohere;
mod mistral;
mod reducto;
mod vertex;
pub(crate) use azure::{AzureDocumentIntelligenceAdapter, AzureMistralAdapter};
pub(crate) use azure::{AzureCohereAdapter, AzureDocumentIntelligenceAdapter, AzureMistralAdapter};
pub(crate) use cohere::CohereAdapter;
pub(crate) use mistral::MistralAdapter;
pub(crate) use reducto::{ReductoLegacyAdapter, ReductoV3Adapter};
pub(crate) use vertex::{VertexDeepSeekAdapter, VertexMistralAdapter};
@ -54,11 +56,16 @@ pub(crate) trait OcrAdapter: Send + Sync + Sized + 'static {
_url: &str,
_headers: &[(String, String)],
request: &LiteLLMOcrRequest,
) -> impl Future<Output = Result<Vec<u8>, OcrError>> + Send {
) -> impl Future<
Output = Result<super::wire::DecodedOcrResponse<Self::ProviderResponse>, OcrError>,
> + Send {
async move {
let bytes = super::client::read_response_bytes(response).await?;
super::handler::post_call(&request.hooks, &bytes).await?;
Ok(bytes)
Ok(super::wire::decode_response(
&bytes,
request.response_format()? == super::types::OcrResponseFormat::Native,
)?)
}
}
}
@ -66,6 +73,8 @@ pub(crate) trait OcrAdapter: Send + Sync + Sized + 'static {
macro_rules! for_each_ocr_adapter {
($callback:ident) => {
$callback! {
Cohere, $crate::ocr::adapters::CohereAdapter, $crate::ocr::adapters::CohereAdapter, Cohere;
AzureCohere, $crate::ocr::adapters::AzureCohereAdapter, $crate::ocr::adapters::AzureCohereAdapter, AzureAi;
Mistral, $crate::ocr::adapters::MistralAdapter, $crate::ocr::adapters::MistralAdapter, Mistral;
AzureMistral, $crate::ocr::adapters::AzureMistralAdapter, $crate::ocr::adapters::AzureMistralAdapter, AzureAi;
AzureDocumentIntelligence, $crate::ocr::adapters::AzureDocumentIntelligenceAdapter, $crate::ocr::adapters::AzureDocumentIntelligenceAdapter, AzureAi;

View file

@ -0,0 +1,235 @@
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value, json};
use crate::ocr::document::InlineDocument;
use crate::ocr::error::{OcrRequestError, OcrResponseError};
use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument};
#[derive(Clone, Copy, Debug, Default, Deserialize, Serialize)]
#[serde(rename_all = "lowercase")]
pub(crate) enum OutputFormat {
#[default]
Markdown,
Blocks,
}
#[derive(Deserialize)]
pub(crate) struct CohereParams {
#[serde(default)]
pub output_format: OutputFormat,
}
#[derive(Deserialize, Serialize)]
pub(crate) struct CohereRequest {
pub model: String,
pub document: OcrDocument,
pub output_format: OutputFormat,
}
pub(crate) fn validate_document(document: &OcrDocument) -> Result<(), OcrRequestError> {
let OcrDocument::ImageUrl { image_url, .. } = document else {
return Err(OcrRequestError::CohereImageOnly);
};
if image_url.is_empty() {
return Err(OcrRequestError::CohereImageOnly);
}
if let Some(inline) = InlineDocument::parse(image_url)? {
if !inline.mime_type().type_.eq_ignore_ascii_case("image") {
return Err(OcrRequestError::CohereImageOnly);
}
inline.decode(crate::constants::OCR_INLINE_MAX_BYTES)?;
}
Ok(())
}
#[derive(Deserialize)]
pub(crate) struct CohereResponse {
#[serde(default)]
pages: Vec<CoherePage>,
meta: Option<CohereMeta>,
}
#[derive(Deserialize)]
struct CoherePage {
index: Option<i64>,
markdown: Option<CohereMarkdown>,
blocks: Option<Vec<Map<String, Value>>>,
}
#[derive(Deserialize)]
struct CohereMarkdown {
#[serde(default)]
content: String,
images: Option<Vec<Map<String, Value>>>,
}
#[derive(Deserialize)]
struct CohereMeta {
billed_units: Option<CohereBilledUnits>,
}
#[derive(Deserialize)]
struct CohereBilledUnits {
pages: Option<i64>,
}
pub(crate) fn transform_response(
model: &str,
response: CohereResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
let pages_processed = response
.meta
.and_then(|meta| meta.billed_units)
.and_then(|units| units.pages)
.map(Ok)
.unwrap_or_else(|| {
i64::try_from(response.pages.len()).map_err(|_| OcrResponseError::NumericRange("pages"))
})?;
let pages = response
.pages
.into_iter()
.enumerate()
.map(|(position, page)| {
let index = page.index.map(Ok).unwrap_or_else(|| {
i64::try_from(position).map_err(|_| OcrResponseError::NumericRange("page index"))
})?;
let (content, images) = page
.markdown
.map(|markdown| {
let images =
markdown
.images
.filter(|images| !images.is_empty())
.map(|images| {
images
.into_iter()
.map(|mut image| {
if let Some(Value::Object(bbox)) =
image.get("bounding_box").cloned()
{
image.insert("bbox".into(), Value::Object(bbox));
}
Value::Object(image)
})
.collect::<Vec<_>>()
});
(markdown.content, images)
})
.unwrap_or_default();
let mut normalized = json!({"index": index, "markdown": content, "images": images});
if let Some(blocks) = page.blocks {
normalized["blocks"] = json!(blocks);
}
Ok(normalized)
})
.collect::<Result<Vec<_>, OcrResponseError>>()?;
Ok(LiteLLMOcrResponse {
pages,
model: model.into(),
document_annotation: None,
usage_info: Some(json!({"pages_processed": pages_processed})),
object: "ocr".into(),
extra_fields: Map::new(),
provider_native_response: None,
})
}
pub(crate) fn transform_request(
model: &str,
document: OcrDocument,
params: CohereParams,
) -> Result<CohereRequest, OcrRequestError> {
validate_document(&document)?;
Ok(CohereRequest {
model: model.into(),
document,
output_format: params.output_format,
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn response_normalizes_markdown_images_blocks_and_billed_pages() {
let response = serde_json::from_value(json!({
"pages": [
{"index": 4, "markdown": {"content": "receipt", "images": [{"id":"image", "bounding_box":{"top_left_x":1}, "description":"scan"}]}},
{"blocks": [{"type":"text","text":"total"}]}
],
"meta": {"billed_units":{"pages":3}}
})).unwrap();
let normalized = transform_response("parse-v5.0", response).unwrap();
assert_eq!(normalized.pages[0]["index"], 4);
assert_eq!(normalized.pages[0]["markdown"], "receipt");
assert_eq!(normalized.pages[0]["images"][0]["bbox"]["top_left_x"], 1);
assert_eq!(normalized.pages[0]["images"][0]["description"], "scan");
assert_eq!(normalized.pages[1]["index"], 1);
assert_eq!(normalized.pages[1]["markdown"], "");
assert_eq!(normalized.pages[1]["blocks"][0]["text"], "total");
assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 3);
}
#[test]
fn response_defaults_and_invalid_fields() {
for value in [
json!({}),
json!({"meta":null}),
json!({"pages":[],"meta":{"billed_units":null}}),
] {
let normalized =
transform_response("parse", serde_json::from_value(value).unwrap()).unwrap();
assert!(normalized.pages.is_empty());
assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 0);
}
for value in [
json!({"pages":null}),
json!({"pages":[{"markdown":"text"}]}),
json!({"pages":[{"index":"bad"}]}),
] {
assert!(serde_json::from_value::<CohereResponse>(value).is_err());
}
let normalized = transform_response(
"parse",
serde_json::from_value(json!({"pages":[{"markdown":null}]})).unwrap(),
)
.unwrap();
assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 1);
assert!(normalized.pages[0]["images"].is_null());
}
#[test]
fn request_requires_image_and_supported_output_format() {
for value in [
json!({"type":"document_url","document_url":"https://example.com/a.pdf"}),
json!({"type":"image_url","image_url":""}),
json!({"type":"image_url","image_url":"data:application/pdf;base64,YQ=="}),
] {
assert_eq!(
validate_document(&serde_json::from_value(value).unwrap()),
Err(OcrRequestError::CohereImageOnly)
);
}
assert!(serde_json::from_value::<CohereParams>(json!({"output_format":"html"})).is_err());
for format in ["markdown", "blocks"] {
assert!(
serde_json::from_value::<CohereParams>(json!({"output_format":format})).is_ok()
);
}
let request = transform_request(
"parse-v5.0",
serde_json::from_value(json!({
"type":"image_url",
"image_url":"https://example.com/image.png"
}))
.unwrap(),
serde_json::from_value(json!({})).unwrap(),
)
.unwrap();
assert_eq!(
serde_json::to_value(request).unwrap()["output_format"],
"markdown"
);
}
}

View file

@ -1,3 +1,4 @@
pub(crate) mod cohere;
pub(crate) mod deepseek;
pub(crate) mod document_intelligence;
pub(crate) mod mistral;

View file

@ -4,6 +4,10 @@ use crate::error::TransportError;
#[derive(Debug, Clone, PartialEq, Eq, Error)]
pub enum OcrRequestError {
#[error(
"Cohere Parse only accepts `image_url` documents; document_url and PDF inputs are not supported"
)]
CohereImageOnly,
#[error("Invalid `req_format`. Expected 'native' or 'litellm'.")]
RequestFormat,
#[error("invalid OCR request field: {path}")]

View file

@ -75,10 +75,10 @@ impl PreparedOcrCall {
($( $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?;
let decoded = $instance.read_response(&self.client, response, &url, &headers, &self.request).await?;
Ok(OcrProviderResponse {
request: self.request,
bytes,
data: OcrProviderData::$variant(decoded),
})
}, )+
}
@ -106,12 +106,14 @@ fn request_headers(request: &reqwest::Request) -> Result<Vec<(String, String)>,
macro_rules! provider_data {
($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => {
enum OcrProviderData {
$( $variant(super::wire::DecodedOcrResponse<<$adapter as OcrAdapter>::ProviderResponse>), )+
}
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)?;
match self.data {
$( OcrProviderData::$variant(decoded) => {
let response = $instance.transform_ocr_response(&self.request, decoded.data)?;
Ok(LiteLLMOcrResponse { provider_native_response: decoded.native, ..response })
}, )+
@ -123,7 +125,7 @@ macro_rules! provider_data {
pub(crate) struct OcrProviderResponse {
request: LiteLLMOcrRequest,
bytes: Vec<u8>,
data: OcrProviderData,
}
pub(crate) async fn post_call(hooks: &Arc<dyn OcrHooks>, bytes: &[u8]) -> Result<(), Error> {

View file

@ -23,6 +23,7 @@ super::adapters::for_each_ocr_adapter!(define_adapter_types);
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum OcrProvider {
Cohere,
Mistral,
AzureAi,
Reducto,
@ -32,6 +33,7 @@ pub(crate) enum OcrProvider {
impl OcrProvider {
pub(crate) const fn as_str(self) -> &'static str {
match self {
Self::Cohere => "cohere",
Self::Mistral => "mistral",
Self::AzureAi => "azure_ai",
Self::Reducto => "reducto",
@ -50,6 +52,7 @@ pub(crate) fn resolve_wire_adapter(
custom_llm_provider: OcrProvider::Mistral.as_str(),
});
let typed_provider = match provider.custom_llm_provider {
"cohere" => OcrProvider::Cohere,
"mistral" => OcrProvider::Mistral,
"azure_ai" => OcrProvider::AzureAi,
"reducto" => OcrProvider::Reducto,
@ -57,10 +60,17 @@ pub(crate) fn resolve_wire_adapter(
value => return Err(Error::InvalidProvider(value.to_string())),
};
let adapter = match typed_provider {
OcrProvider::Cohere => OcrAdapterKind::Cohere,
OcrProvider::Mistral => OcrAdapterKind::Mistral,
OcrProvider::AzureAi if is_document_intelligence_model(provider.model) => {
OcrAdapterKind::AzureDocumentIntelligence
}
OcrProvider::AzureAi
if provider.model.to_ascii_lowercase().contains("cohere")
&& provider.model.to_ascii_lowercase().contains("parse") =>
{
OcrAdapterKind::AzureCohere
}
OcrProvider::AzureAi => OcrAdapterKind::AzureMistral,
OcrProvider::Reducto if provider.model.eq_ignore_ascii_case("parse-legacy") => {
OcrAdapterKind::ReductoLegacy

View file

@ -89,6 +89,7 @@ pub fn consumed_optional_param_names(
let (_, adapter) = super::registry::resolve_wire_adapter(model, custom_llm_provider)?;
let provider_fields: &[&str] = match adapter {
OcrAdapterKind::Cohere | OcrAdapterKind::AzureCohere => &["output_format"],
OcrAdapterKind::Mistral | OcrAdapterKind::AzureMistral | OcrAdapterKind::VertexMistral => {
MISTRAL_OPTION_FIELDS
}
@ -98,9 +99,9 @@ pub fn consumed_optional_param_names(
OcrAdapterKind::VertexDeepSeek => DEEPSEEK_OPTION_FIELDS,
};
let auth_fields: &[&str] = match adapter {
OcrAdapterKind::AzureMistral | OcrAdapterKind::AzureDocumentIntelligence => {
AZURE_AUTH_OPTION_FIELDS
}
OcrAdapterKind::AzureMistral
| OcrAdapterKind::AzureDocumentIntelligence
| OcrAdapterKind::AzureCohere => AZURE_AUTH_OPTION_FIELDS,
OcrAdapterKind::VertexMistral | OcrAdapterKind::VertexDeepSeek => VERTEX_AUTH_OPTION_FIELDS,
_ => &[],
};

View file

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

View file

@ -180,8 +180,17 @@ impl super::hooks::OcrHooks for SubmissionBoundary {
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}"#));
match self.request_count.lock().unwrap().len() {
1 => assert_eq!(request.original_response, json!(r#"{"submitted":true}"#)),
2 => assert!(
request
.original_response
.as_str()
.unwrap()
.contains("succeeded")
),
count => panic!("unexpected callback after {count} requests"),
}
Ok(request)
})
}

View file

@ -73,7 +73,18 @@ pub(crate) fn ocr_error_to_pyerr(err: Error) -> PyErr {
Error::Network(message) if message.contains("timed out") => {
ocr_upstream_error(408, message)
}
other => core_error_to_pyerr(other),
other => {
let status = other.http_status_code();
let error = core_error_to_pyerr(other);
if let Some(status) = status {
Python::attach(|py| {
let value = error.value(py);
value.setattr("status_code", status).ok();
value.setattr("message", value.to_string()).ok();
});
}
error
}
}
}
@ -111,6 +122,18 @@ mod ocr_error_tests {
.and_then(|args| args.extract())
.expect("OCR failures retain status and unprefixed provider message");
assert_eq!(args, (429, r#"{"message":"rate limited"}"#.to_string()));
let mapped = ocr_error_to_pyerr(Error::InvalidRequest("invalid format".into()));
assert!(mapped.is_instance_of::<PyValueError>(py));
assert_eq!(
mapped
.value(py)
.getattr("status_code")
.unwrap()
.extract::<u16>()
.unwrap(),
400
);
});
}
}

View file

@ -60,9 +60,14 @@ where
E: Send + 'static,
F: Future<Output = Result<T, E>> + Send + 'static,
{
let result = run_sync_value_on(py, runtime, async move {
map_core_result(future.await, map_error)
})?;
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)?;
Pythonized(result).into_pyobject(py).map(Bound::unbind)
}
@ -76,8 +81,9 @@ where
E: Send + 'static,
F: Future<Output = Result<T, E>> + Send + 'static,
{
run_async_value(py, async move {
let result = map_core_result(future.await, map_error)?;
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)?;
Ok(Pythonized(result))
})
}

View file

@ -65,7 +65,7 @@ impl ResponsesWebSocketConnection {
}
}
#[pymodule(gil_used = true)]
#[pymodule(gil_used = false)]
mod _native {
use pyo3::prelude::*;

View file

@ -10,7 +10,9 @@ use litellm_core::ocr::wire::{OcrWireRequest, consumed_optional_param_names, dec
use litellm_core::ocr::{
NativeOutcome, OcrAdmission, OcrCall, OcrClient, OcrHostOperation, OcrHostResult,
};
use litellm_python_interop::{from_py, to_py};
use litellm_python_interop::{
from_py_preserving_errors as from_py, to_py_preserving_errors as to_py,
};
use crate::errors::{RustBridgeDeclined, ocr_error_to_pyerr};
use crate::lifecycle::{PythonCallState, PythonRoute, missing_state, now, run_call};
@ -82,9 +84,14 @@ impl PythonOcrHost {
fn python_pre_call(
&mut self,
py: Python<'_>,
request: OcrDuringCallRequest,
mut request: OcrDuringCallRequest,
) -> PyResult<OcrDuringCallRequest> {
let pre_call = self.pre_call.as_ref().ok_or_else(missing_state)?;
if let Some(body) = request.body.as_object_mut() {
for name in &request.retained_fields {
body.remove(name);
}
}
let body = to_py(py, &request.body)?
.into_bound(py)
.cast_into::<PyDict>()?;
@ -134,11 +141,9 @@ impl PythonOcrHost {
.iter()
.map(|(name, value)| Ok((name.extract::<String>()?, value.extract::<String>()?)))
.collect::<PyResult<Vec<_>>>()?;
Ok(OcrDuringCallRequest {
body: from_py(&body)?,
headers,
..request
})
request.body = from_py(&body)?;
request.headers = headers;
Ok(request)
}
fn python_post_call(
@ -413,6 +418,13 @@ fn _ocr_lifecycle(
kwargs: Bound<'_, PyDict>,
asynchronous: bool,
) -> PyResult<Py<PyAny>> {
if let Ok(gil_enabled) = py.import("sys")?.getattr("_is_gil_enabled")
&& !gil_enabled.call0()?.is_truthy()?
{
return Err(pyo3::exceptions::PyRuntimeError::new_err(
"native OCR requires the Python GIL",
));
}
let client = OcrClient::shared().map_err(ocr_error_to_pyerr)?;
let call = admitted_call(OcrCall::admit(
client,

View file

@ -2,4 +2,6 @@ mod gil;
mod marshal;
pub use gil::{release_count, release_gil};
pub use marshal::{Pythonized, from_py, panic_to_pyerr, to_py};
pub use marshal::{
Pythonized, from_py, from_py_preserving_errors, panic_to_pyerr, to_py, to_py_preserving_errors,
};

View file

@ -1,12 +1,20 @@
use std::any::Any;
use std::panic::{AssertUnwindSafe, catch_unwind};
use pyo3::exceptions::PyValueError;
use pyo3::panic::PanicException;
use pyo3::prelude::*;
use serde::Serialize;
use serde::de::DeserializeOwned;
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()))
}
pub fn from_py_preserving_errors<T>(value: &Bound<'_, PyAny>) -> PyResult<T>
where
T: DeserializeOwned,
{
@ -17,14 +25,18 @@ pub fn to_py<T>(py: Python<'_>, value: &T) -> PyResult<Py<PyAny>>
where
T: Serialize + ?Sized,
{
pythonize_bound(py, value).map(Bound::unbind)
pythonize::pythonize(py, value)
.map(Bound::unbind)
.map_err(|error| PyValueError::new_err(error.to_string()))
}
fn pythonize_bound<'py, T>(py: Python<'py>, value: &T) -> PyResult<Bound<'py, PyAny>>
pub fn to_py_preserving_errors<T>(py: Python<'_>, value: &T) -> PyResult<Py<PyAny>>
where
T: Serialize + ?Sized,
{
pythonize::pythonize(py, value).map_err(PyErr::from)
pythonize::pythonize(py, value)
.map(Bound::unbind)
.map_err(PyErr::from)
}
pub struct Pythonized<T>(pub T);
@ -38,7 +50,9 @@ where
type Error = PyErr;
fn into_pyobject(self, py: Python<'py>) -> PyResult<Self::Output> {
catch_unwind(AssertUnwindSafe(|| pythonize_bound(py, &self.0))).map_err(panic_to_pyerr)?
catch_unwind(AssertUnwindSafe(|| pythonize::pythonize(py, &self.0)))
.map_err(panic_to_pyerr)?
.map_err(|error| PyValueError::new_err(error.to_string()))
}
}
@ -112,7 +126,15 @@ value = Broken()
Some(&locals),
)
.unwrap();
let error = from_py::<i64>(&locals.get_item("value").unwrap().unwrap()).unwrap_err();
let value = locals.get_item("value").unwrap().unwrap();
let legacy_error = from_py::<i64>(&value).unwrap_err();
assert!(legacy_error.is_instance_of::<PyValueError>(py));
assert!(
!legacy_error
.value(py)
.is(locals.get_item("failure").unwrap().unwrap())
);
let error = from_py_preserving_errors::<i64>(&value).unwrap_err();
assert!(
error
.value(py)

View file

@ -1,241 +0,0 @@
from __future__ import annotations
import asyncio
import contextvars
import datetime
from collections.abc import Awaitable, Callable, Coroutine
from concurrent.futures import Future
from typing import Protocol
class CompletionLogging(Protocol):
def success_handler(
self,
result: object,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
def failure_handler(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
def async_success_handler(
self,
result: object,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> Coroutine[object, object, None]: ...
def async_failure_handler(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> Awaitable[None]: ...
def handle_sync_success_callbacks_for_async_calls(
self,
result: object,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
class CompletionExecutor(Protocol):
def submit(
self,
function: Callable[..., object],
/,
*args: object,
) -> Future[object]: ...
class Completion(Protocol):
def success(
self,
result: object,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
def failure(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
async def async_failure(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
class PythonCompletion:
def __init__(
self,
logging_obj: CompletionLogging,
executor: CompletionExecutor | None,
*,
async_call: bool,
internal_call: bool,
completion_with_fallbacks: bool,
) -> None:
self._logging_obj = logging_obj
self._executor = executor
self._async_call = async_call
self._internal_call = internal_call
self._completion_with_fallbacks = completion_with_fallbacks
def success(
self,
result: object,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None:
if not self._async_call:
assert self._executor is not None
context = contextvars.copy_context()
self._executor.submit(
context.run,
self._logging_obj.success_handler,
result,
start_time,
end_time,
)
return
if not self._internal_call:
if getattr(self._logging_obj, "_defer_async_logging", False):
def enqueue_deferred_logging() -> None:
asyncio.create_task(self._dispatch_async_success(result, start_time, end_time))
setattr( # noqa: B010 # optional legacy logger field is absent from narrow test doubles
self._logging_obj,
"_enqueue_deferred_logging",
enqueue_deferred_logging,
)
else:
asyncio.create_task(self._dispatch_async_success(result, start_time, end_time))
self._logging_obj.handle_sync_success_callbacks_for_async_calls(
result=result,
start_time=start_time,
end_time=end_time,
)
def failure(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None:
if self._async_call and self._internal_call:
return
self._logging_obj.failure_handler(exception, traceback_exception, start_time, end_time)
async def async_failure(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None:
if not self._async_call or self._internal_call:
return
await self._logging_obj.async_failure_handler(exception, traceback_exception, start_time, end_time)
async def _dispatch_async_success(
self,
result: object,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None:
if self._completion_with_fallbacks:
return
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue( # pyright: ignore[reportUnknownMemberType] # legacy worker lacks generic coroutine annotations
async_coroutine=self._logging_obj.async_success_handler(
result=result,
start_time=start_time,
end_time=end_time,
)
)
self._logging_obj.handle_sync_success_callbacks_for_async_calls(
result=result,
start_time=start_time,
end_time=end_time,
)
class CallCompletion:
def __init__(self, implementation: Completion) -> None:
self._python_implementation: Completion | None = implementation
self._implementation: Completion | None = implementation
self._attached = False
@property
def python_implementation(self) -> Completion:
assert self._python_implementation is not None
return self._python_implementation
def attach(self, implementation: Completion) -> bool:
if self._attached:
return False
self._implementation = implementation
self._attached = True
return True
def release(self) -> None:
self._python_implementation = None
self._implementation = None
def success(
self,
result: object,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None:
implementation = self._implementation
assert implementation is not None
implementation.success(result, start_time, end_time)
def failure(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None:
implementation = self._implementation
assert implementation is not None
implementation.failure(exception, traceback_exception, start_time, end_time)
async def async_failure(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None:
implementation = self._implementation
assert implementation is not None
await implementation.async_failure(
exception,
traceback_exception,
start_time,
end_time,
)

View file

@ -3,7 +3,9 @@ import mimetypes
import os
import re
from io import IOBase
from typing import Any, Final
from typing import Final, Literal, Protocol
from typing_extensions import NotRequired, ReadOnly, TypedDict
from litellm._logging import verbose_logger
@ -31,7 +33,17 @@ def get_mime_type(file_path: str) -> str:
return guessed or "application/octet-stream"
def convert_file_document_to_url_document(document: dict[str, Any]) -> dict[str, str]:
class FileReader(Protocol):
def read(self) -> bytes | str: ...
class FileDocument(TypedDict):
type: ReadOnly[Literal["file"]]
file: ReadOnly[bytes | os.PathLike[str] | FileReader]
mime_type: ReadOnly[NotRequired[str]]
def convert_file_document_to_url_document(document: FileDocument) -> dict[str, str]:
file_input: Final = document.get("file")
if file_input is None:
raise ValueError(

View file

@ -3181,10 +3181,11 @@ 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)
if getattr(logging_obj, "call_type", None) in ("ocr", "aocr"):
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 OCR 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

@ -8,7 +8,6 @@ from typing import Final, Protocol, cast # noqa: TID251 # validates dynamicall
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
@ -44,8 +43,6 @@ 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()
@ -99,12 +96,6 @@ def call_azure_ad_token_provider(provider: object) -> str:
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
)

View file

@ -7,6 +7,7 @@ import ast
import asyncio
import base64
import binascii
import contextvars
import copy
import datetime
import hashlib
@ -80,7 +81,6 @@ from litellm.constants import (
PROVIDERS_THAT_AUTHENTICATE_ON_PROVIDER_INFO,
TOOL_CHOICE_OBJECT_TOKEN_COUNT,
)
from litellm.litellm_core_utils.call_completion import CallCompletion, PythonCompletion
from litellm.litellm_core_utils.core_helpers import normalize_drop_params
from litellm.litellm_core_utils.fallback_generalizations import (
match_capability_generalizations,
@ -1196,6 +1196,79 @@ def function_setup(
raise e
def _dispatch_success_logging(
logging_obj: LiteLLMLoggingObject,
result: object,
start_time: datetime.datetime,
end_time: datetime.datetime,
is_completion_with_fallbacks: bool,
is_litellm_internal_call: bool,
) -> None:
if not is_litellm_internal_call:
if getattr(logging_obj, "_defer_async_logging", False):
def _enqueue_deferred_logging() -> None:
asyncio.create_task(
_client_async_logging_helper(
logging_obj=logging_obj,
result=result,
start_time=start_time,
end_time=end_time,
is_completion_with_fallbacks=is_completion_with_fallbacks,
)
)
logging_obj._enqueue_deferred_logging = _enqueue_deferred_logging
else:
asyncio.create_task(
_client_async_logging_helper(
logging_obj=logging_obj,
result=result,
start_time=start_time,
end_time=end_time,
is_completion_with_fallbacks=is_completion_with_fallbacks,
)
)
logging_obj.handle_sync_success_callbacks_for_async_calls(
result=result,
start_time=start_time,
end_time=end_time,
)
async def _client_async_logging_helper(
logging_obj: LiteLLMLoggingObject,
result,
start_time,
end_time,
is_completion_with_fallbacks: bool,
):
if (
is_completion_with_fallbacks is False
): # don't log the parent event litellm.completion_with_fallbacks as a 'log_success_event', this will lead to double logging the same call - https://github.com/BerriAI/litellm/issues/7477
print_verbose(
f"Async Wrapper: Completed Call, calling async_success_handler: {logging_obj.async_success_handler}"
)
################################################
# Async Logging Worker
################################################
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
async_coroutine=logging_obj.async_success_handler(result=result, start_time=start_time, end_time=end_time)
)
################################################
# Sync Logging Worker
################################################
logging_obj.handle_sync_success_callbacks_for_async_calls(
result=result,
start_time=start_time,
end_time=end_time,
)
def _get_wrapper_num_retries(kwargs: dict[str, Any], exception: Exception) -> tuple[int | None, dict[str, Any]]:
"""
Get the number of retries from the kwargs and the retry policy.
@ -1464,7 +1537,6 @@ def client(original_function):
start_time: Final = datetime.datetime.now()
result = None
logging_obj: LiteLLMLoggingObject | None = kwargs.get("litellm_logging_obj", None)
completion: CallCompletion | None = None
# only set litellm_call_id if its not in kwargs
if "litellm_call_id" not in kwargs:
@ -1480,17 +1552,7 @@ def client(original_function):
# Type assertion: logging_obj is guaranteed to be non-None after function_setup
assert logging_obj is not None, "logging_obj should not be None after function_setup"
from litellm.litellm_core_utils.thread_pool_executor import executor as logging_executor
completion = CallCompletion(
PythonCompletion(
logging_obj,
logging_executor,
async_call=False,
internal_call=False,
completion_with_fallbacks=False,
)
)
## LOAD CREDENTIALS
load_credentials_from_list(kwargs)
kwargs["litellm_logging_obj"] = logging_obj
@ -1588,12 +1650,7 @@ def client(original_function):
except Exception as e:
print_verbose(f"Error while checking max token limit: {e}")
# MODEL CALL
invocation_kwargs: Final = (
{**kwargs, "_litellm_call_completion": completion}
if original_function.__name__ == CallTypes.ocr.value
else kwargs
)
result = original_function(*args, **invocation_kwargs)
result = original_function(*args, **kwargs)
end_time = datetime.datetime.now()
if _is_streaming_request(
kwargs=kwargs,
@ -1647,8 +1704,8 @@ def client(original_function):
kwargs=kwargs,
)
update_response_metadata = getattr(sys.modules[__name__], "update_response_metadata")
update_response_metadata(
_update_response_metadata: Final = getattr(sys.modules[__name__], "update_response_metadata")
_update_response_metadata(
result=result,
logging_obj=logging_obj,
model=model,
@ -1656,8 +1713,21 @@ def client(original_function):
start_time=start_time,
end_time=end_time,
)
# LOG SUCCESS - handle streaming success logging in the _next_ object, remove `handle_success` once it's deprecated
verbose_logger.info("Wrapper: Completed Call, calling success_handler")
completion.success(result, start_time, end_time)
# Copy the current context to propagate it to the background thread
# This is essential for OpenTelemetry span context propagation
ctx: Final = contextvars.copy_context()
executor: Final = getattr(sys.modules[__name__], "executor")
executor.submit(
ctx.run,
logging_obj.success_handler,
result,
start_time,
end_time,
)
# RETURN RESULT
return result
except Exception as e:
call_type = original_function.__name__
@ -1731,14 +1801,11 @@ def client(original_function):
end_time = datetime.datetime.now()
# LOG FAILURE - handle streaming failure logging in the _next_ object, remove `handle_failure` once it's deprecated
if completion is not None:
completion.failure(e, traceback_exception, start_time, end_time)
elif logging_obj:
logging_obj.failure_handler(e, traceback_exception, start_time, end_time)
if logging_obj:
logging_obj.failure_handler(
e, traceback_exception, start_time, end_time
) # DO NOT MAKE THREADED - router retry fallback relies on this!
raise e
finally:
if completion is not None:
completion.release()
@wraps(original_function)
async def wrapper_async(*args, **kwargs):
@ -1747,7 +1814,6 @@ def client(original_function):
result = None
_update_response_metadata: Final = getattr(sys.modules[__name__], "update_response_metadata")
logging_obj: LiteLLMLoggingObject | None = kwargs.get("litellm_logging_obj", None)
completion: CallCompletion | None = None
LLMCachingHandler: Final = _get_cached_llm_caching_handler()
_llm_caching_handler: Final[LLMCachingHandler] = LLMCachingHandler(
original_function=original_function,
@ -1773,15 +1839,7 @@ def client(original_function):
# Type assertion: logging_obj is guaranteed to be non-None after function_setup
assert logging_obj is not None, "logging_obj should not be None after function_setup"
completion = CallCompletion(
PythonCompletion(
logging_obj,
None,
async_call=True,
internal_call=_is_litellm_internal_call,
completion_with_fallbacks=is_completion_with_fallbacks,
)
)
modified_kwargs: Final = await async_pre_call_deployment_hook(kwargs, call_type)
if modified_kwargs is not None:
kwargs = modified_kwargs
@ -1866,12 +1924,7 @@ def client(original_function):
# MODEL CALL
try:
invocation_kwargs: Final = (
{**kwargs, "_litellm_call_completion": completion}
if original_function.__name__ == CallTypes.aocr.value
else kwargs
)
result = await original_function(*args, **invocation_kwargs)
result = await original_function(*args, **kwargs)
except Exception as deployment_error:
_deployment_call_end_time = datetime.datetime.now() # noqa: DTZ005 # matches the naive datetimes this whole function already times start_time/end_time with
try:
@ -1940,7 +1993,14 @@ def client(original_function):
and _caching_handler_response is not None
and _caching_handler_response.final_embedding_cached_response is not None
):
completion.success(result, start_time, end_time)
_dispatch_success_logging(
logging_obj=logging_obj,
result=result,
start_time=start_time,
end_time=end_time,
is_completion_with_fallbacks=is_completion_with_fallbacks,
is_litellm_internal_call=_is_litellm_internal_call,
)
return _llm_caching_handler._combine_cached_embedding_response_with_api_result(
_caching_handler_response=_caching_handler_response,
embedding_response=result,
@ -1956,7 +2016,14 @@ def client(original_function):
start_time=start_time,
end_time=end_time,
)
completion.success(result, start_time, end_time)
_dispatch_success_logging(
logging_obj=logging_obj,
result=result,
start_time=start_time,
end_time=end_time,
is_completion_with_fallbacks=is_completion_with_fallbacks,
is_litellm_internal_call=_is_litellm_internal_call,
)
return result
except Exception as e:
@ -1964,12 +2031,17 @@ def client(original_function):
# Reuse the timestamp taken right when the deployment call itself failed, before
# the failure hook ran, so a slow callback doesn't inflate the reported duration.
end_time = _deployment_call_end_time if _deployment_call_end_time is not None else datetime.datetime.now() # noqa: DTZ005 # matches the naive datetimes this whole function already times start_time/end_time with
if completion is not None:
completion.failure(e, traceback_exception, start_time, end_time)
await completion.async_failure(e, traceback_exception, start_time, end_time)
elif logging_obj and not _is_litellm_internal_call:
logging_obj.failure_handler(e, traceback_exception, start_time, end_time)
await logging_obj.async_failure_handler(e, traceback_exception, start_time, end_time)
if logging_obj and not _is_litellm_internal_call:
try:
logging_obj.failure_handler(
e, traceback_exception, start_time, end_time
) # DO NOT MAKE THREADED - router retry fallback relies on this!
except Exception as e:
raise e
try:
await logging_obj.async_failure_handler(e, traceback_exception, start_time, end_time)
except Exception as e:
raise e
call_type = original_function.__name__
num_retries, kwargs = _get_wrapper_num_retries(kwargs=kwargs, exception=e)
@ -2039,8 +2111,6 @@ def client(original_function):
raise e
finally:
if completion is not None:
completion.release()
# Restore trace_id/session_id contextvars to their pre-call value once
# this call (in this asyncio Task) is fully done - see
# request_correlation_in_logs. Unlike wrapper()'s sync path, it's safe to

View file

@ -1,402 +0,0 @@
import asyncio
import contextvars
import datetime
import weakref
from collections.abc import Callable, Coroutine
from concurrent.futures import Future, ThreadPoolExecutor
from threading import get_ident
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import pytest
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.utils import client
class RecordingExecutor:
def __init__(self) -> None:
self.submissions: list[tuple[Callable[..., object], tuple[object, ...]]] = []
def submit(self, function: Callable[..., object], *args: object) -> Future[object]:
self.submissions.append((function, args))
future: Final[Future[object]] = Future()
future.set_result(function(*args))
return future
class RecordingCompletion:
def __init__(self) -> None:
self.successes: list[object] = []
self.failures: list[Exception] = []
def success(
self,
result: object,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None:
self.successes.append(result)
def failure(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None:
self.failures.append(exception)
async def async_failure(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None:
self.failures.append(exception)
class RecordingLogging:
def __init__(self, marker: contextvars.ContextVar[str], observed: list[tuple[object, str]]) -> None:
self._marker = marker
self._observed = observed
self._defer_async_logging = False
self._enqueue_deferred_logging: Callable[[], None] | None = None
def success_handler(
self,
result: object,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None:
self._observed.append((result, self._marker.get()))
def failure_handler(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
async def async_success_handler(
self,
result: object,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
async def async_failure_handler(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
def handle_sync_success_callbacks_for_async_calls(
self,
result: object,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
def test_python_completion_preserves_sync_context_and_response_identity() -> None:
marker: Final[contextvars.ContextVar[str]] = contextvars.ContextVar("marker")
marker.set("request-context")
executor: Final = RecordingExecutor()
response: Final = object()
observed: Final[list[tuple[object, str]]] = []
logging_obj: Final = RecordingLogging(marker, observed)
completion: Final = PythonCompletion(
logging_obj,
executor,
async_call=False,
internal_call=False,
completion_with_fallbacks=False,
)
now: Final = datetime.datetime.now(datetime.timezone.utc)
completion.success(response, now, now)
assert observed == [(response, "request-context")]
assert len(executor.submissions) == 1
def test_sync_wrapper_dispatches_with_logging_executor_and_caller_context(monkeypatch: pytest.MonkeyPatch) -> None:
marker: Final[contextvars.ContextVar[str]] = contextvars.ContextVar("wrapper-context", default="missing")
marker.set("request-context")
caller_thread: Final = get_ident()
response: Final = object()
observed: Final[list[tuple[object, str, int]]] = []
logging_obj: Final = MagicMock()
def record_success(result: object, start_time: datetime.datetime, end_time: datetime.datetime) -> None:
observed.append((result, marker.get(), get_ident()))
def ocr(**kwargs: object) -> object:
return response
logging_obj.success_handler.side_effect = record_success
monkeypatch.setattr("litellm.utils.function_setup", MagicMock(return_value=(logging_obj, {})))
monkeypatch.setattr("litellm.utils.load_credentials_from_list", MagicMock())
wrapped: Final = client(ocr)
with ThreadPoolExecutor(max_workers=1) as executor:
monkeypatch.setattr(thread_pool_executor, "executor", executor)
result: Final = wrapped()
assert result is response
assert len(observed) == 1
assert observed[0][0] is response
assert observed[0][1] == "request-context"
assert observed[0][2] != caller_thread
@pytest.mark.asyncio
async def test_call_completion_attaches_once_and_forwards_final_objects() -> None:
python_completion: Final = RecordingCompletion()
native_completion: Final = RecordingCompletion()
ignored_completion: Final = RecordingCompletion()
completion: Final = CallCompletion(python_completion)
response: Final = object()
error: Final = ValueError("mapped failure")
now: Final = datetime.datetime.now(datetime.timezone.utc)
assert completion.python_implementation is python_completion
assert completion.attach(native_completion)
assert not completion.attach(ignored_completion)
completion.success(response, now, now)
completion.failure(error, "traceback", now, now)
await completion.async_failure(error, "traceback", now, now)
assert native_completion.successes == [response]
assert native_completion.failures == [error, error]
assert python_completion.successes == []
assert ignored_completion.successes == []
@pytest.mark.asyncio
async def test_python_completion_retains_deferred_success_arguments(monkeypatch: pytest.MonkeyPatch) -> None:
response: Final = object()
logging_obj: Final = MagicMock()
logging_obj._defer_async_logging = True
logging_obj.async_success_handler = AsyncMock()
completion: Final = CallCompletion(
PythonCompletion(
logging_obj,
RecordingExecutor(),
async_call=True,
internal_call=False,
completion_with_fallbacks=False,
)
)
worker: Final = MagicMock()
monkeypatch.setattr("litellm.litellm_core_utils.logging_worker.GLOBAL_LOGGING_WORKER", worker)
scheduled: Final[list[Coroutine[object, object, None]]] = []
monkeypatch.setattr("asyncio.create_task", scheduled.append)
now: Final = datetime.datetime.now(datetime.timezone.utc)
completion.success(response, now, now)
completion.release()
logging_obj._enqueue_deferred_logging()
assert len(scheduled) == 1
await scheduled[0]
await worker.ensure_initialized_and_enqueue.call_args.kwargs["async_coroutine"]
logging_obj.async_success_handler.assert_awaited_once_with(result=response, start_time=now, end_time=now)
@pytest.mark.asyncio
async def test_async_ocr_wrapper_injects_completion_after_fresh_deployment_kwargs(
monkeypatch: pytest.MonkeyPatch,
) -> None:
native_completion: Final = RecordingCompletion()
original_response: Final = object()
replacement_response: Final = object()
shared_metadata: Final = {"request": "shared"}
replacement_kwargs: Final[dict[str, object]] = {"metadata": shared_metadata}
hook_input: dict[str, object] | None = None
async def fresh_kwargs(kwargs: dict[str, object], call_type: str) -> dict[str, object]:
nonlocal hook_input
hook_input = kwargs
return replacement_kwargs
async def aocr(**kwargs: object) -> object:
completion = kwargs.get("_litellm_call_completion")
assert isinstance(completion, CallCompletion)
assert completion.attach(native_completion)
assert kwargs["metadata"] is shared_metadata
return original_response
async def replace_response(request_data: dict[str, object], response: object, call_type: object) -> object:
assert response is original_response
return replacement_response
monkeypatch.setattr("litellm.utils.async_pre_call_deployment_hook", fresh_kwargs)
monkeypatch.setattr("litellm.utils.async_post_call_success_deployment_hook", replace_response)
monkeypatch.setattr("litellm.utils.function_setup", MagicMock(return_value=(MagicMock(), {})))
monkeypatch.setattr("litellm.utils.load_credentials_from_list", MagicMock())
wrapped: Final = client(aocr)
result: Final = await wrapped()
assert result is replacement_response
assert native_completion.successes == [replacement_response]
assert hook_input is not None
assert "_litellm_call_completion" not in hook_input
assert "_litellm_call_completion" not in replacement_kwargs
assert replacement_kwargs["metadata"] is shared_metadata
@pytest.mark.asyncio
async def test_async_ocr_wrapper_sends_final_failure_to_attached_completion(
monkeypatch: pytest.MonkeyPatch,
) -> None:
native_completion: Final = RecordingCompletion()
mapped_error: Final = ValueError("mapped OCR failure")
async def aocr(**kwargs: object) -> object:
completion = kwargs.get("_litellm_call_completion")
assert isinstance(completion, CallCompletion)
assert completion.attach(native_completion)
raise mapped_error
monkeypatch.setattr("litellm.utils.function_setup", MagicMock(return_value=(MagicMock(), {})))
monkeypatch.setattr("litellm.utils.load_credentials_from_list", MagicMock())
wrapped: Final = client(aocr)
with pytest.raises(ValueError, match="mapped OCR failure") as caught:
await wrapped()
assert caught.value is mapped_error
assert native_completion.failures == [mapped_error, mapped_error]
@pytest.mark.asyncio
async def test_async_ocr_wrapper_reports_metadata_failure_without_success(
monkeypatch: pytest.MonkeyPatch,
) -> None:
native_completion: Final = RecordingCompletion()
response: Final = object()
metadata_error: Final = ValueError("metadata failure")
async def aocr(**kwargs: object) -> object:
completion = kwargs.get("_litellm_call_completion")
assert isinstance(completion, CallCompletion)
assert completion.attach(native_completion)
return response
def fail_metadata(**kwargs: object) -> None:
raise metadata_error
monkeypatch.setattr("litellm.utils.function_setup", MagicMock(return_value=(MagicMock(), {})))
monkeypatch.setattr("litellm.utils.load_credentials_from_list", MagicMock())
monkeypatch.setattr("litellm.utils.update_response_metadata", fail_metadata)
wrapped: Final = client(aocr)
with pytest.raises(ValueError, match="metadata failure") as caught:
await wrapped()
assert caught.value is metadata_error
assert native_completion.successes == []
assert native_completion.failures == [metadata_error, metadata_error]
@pytest.mark.parametrize("asynchronous", [False, True])
@pytest.mark.asyncio
async def test_wrapper_completion_stays_separate_from_provider_options(
monkeypatch: pytest.MonkeyPatch, asynchronous: bool
) -> None:
native_completion: Final = RecordingCompletion()
response: Final = OCRResponse(model="mistral-ocr-latest", pages=[])
metadata: Final = {"request": "shared"}
pages: Final = [0, 2]
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 aocr(*, _litellm_call_completion: CallCompletion, **kwargs: object) -> OCRResponse:
return ocr(_litellm_call_completion=_litellm_call_completion, **kwargs)
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"},
"api_key": "test-key",
"metadata": metadata,
"pages": pages,
}
result: Final = await wrapped(**arguments) if asynchronous else wrapped(**arguments)
assert result is response
assert native_completion.successes == [response]
assert native_completion.failures == []
assert arguments["metadata"] is metadata
assert "_litellm_call_completion" not in arguments
@pytest.mark.parametrize(
("asynchronous", "exit_path"),
[(False, "success"), (True, "success"), (False, "callback_error"), (True, "callback_error"), (True, "cancelled")],
)
@pytest.mark.asyncio
async def test_wrapper_releases_completion_resources_on_every_exit(
monkeypatch: pytest.MonkeyPatch, asynchronous: bool, exit_path: str
) -> None:
retained: Final[list[tuple[CallCompletion, weakref.ReferenceType[object]]]] = []
callback_error: Final = RuntimeError("failure callback failed")
implementation: Final = MagicMock(spec=RecordingCompletion)
implementation.failure.side_effect = callback_error
implementation.async_failure = AsyncMock()
response: Final = object()
def ocr(*, _litellm_call_completion: CallCompletion, **kwargs: object) -> object:
retained.append((_litellm_call_completion, weakref.ref(_litellm_call_completion.python_implementation)))
assert _litellm_call_completion.attach(implementation)
if exit_path == "cancelled":
raise asyncio.CancelledError
if exit_path == "callback_error":
raise ValueError("provider failed")
return response
async def aocr(*, _litellm_call_completion: CallCompletion, **kwargs: object) -> object:
return ocr(_litellm_call_completion=_litellm_call_completion, **kwargs)
monkeypatch.setattr("litellm.utils.function_setup", MagicMock(return_value=(MagicMock(), {})))
monkeypatch.setattr("litellm.utils.load_credentials_from_list", MagicMock())
wrapped: Final = client(aocr if asynchronous else ocr)
if exit_path == "success":
result: Final = await wrapped() if asynchronous else wrapped()
assert result is response
assert implementation.success.call_args.args[0] is response
elif exit_path == "cancelled":
with pytest.raises(asyncio.CancelledError):
await wrapped()
implementation.success.assert_not_called()
implementation.failure.assert_not_called()
else:
with pytest.raises(RuntimeError, match="failure callback failed") as caught:
await wrapped() if asynchronous else wrapped()
assert caught.value is callback_error
implementation.async_failure.assert_not_called()
assert len(retained) == 1
assert retained[0][1]() is None

View file

@ -1,9 +1,5 @@
import base64
import json
import pytest
import litellm
from litellm.llms.azure_ai.ocr.cohere_parse_transformation import AzureAICohereParseConfig
from litellm.llms.azure_ai.ocr.common_utils import get_azure_ai_ocr_config
from litellm.llms.azure_ai.ocr.document_intelligence.transformation import AzureDocumentIntelligenceOCRConfig
@ -12,27 +8,6 @@ from litellm.llms.azure_ai.ocr.transformation import AzureAIOCRConfig
MODEL = "azure_ai/Cohere-parse-v5"
API_BASE = "https://resource.services.ai.azure.com"
PARSE_URL = f"{API_BASE}/providers/cohere/v2/parse"
IMAGE_URL = "https://example.com/receipt.png"
PNG_BYTES = base64.b64decode(
"iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNkYPhfDwAChwGA60e6kgAAAABJRU5ErkJggg=="
)
PNG_DATA_URI = f"data:image/png;base64,{base64.b64encode(PNG_BYTES).decode()}"
def _parse_response() -> dict:
return {
"id": "882bf973-9dfa-4d02-9d30-709247008efd",
"pages": [{"index": 0, "type": "markdown", "markdown": {"content": "# Receipt\n\nTotal Due: $4.00"}}],
"meta": {"api_version": {"version": "2"}, "billed_units": {"pages": 1}},
}
@pytest.fixture()
def disable_aiohttp_transport(monkeypatch):
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
litellm.in_memory_llm_clients_cache.flush_cache()
yield
litellm.in_memory_llm_clients_cache.flush_cache()
@pytest.mark.parametrize(
@ -95,90 +70,3 @@ def test_validate_environment_requires_api_base(monkeypatch) -> None:
with pytest.raises(ValueError, match="AZURE_AI_API_BASE"):
AzureAICohereParseConfig().validate_environment(headers={}, model="Cohere-parse-v5", api_key="key")
@pytest.mark.asyncio
async def test_aocr_inlines_remote_image_and_posts_to_foundry(disable_aiohttp_transport, respx_mock):
respx_mock.get(IMAGE_URL).respond(content=PNG_BYTES, headers={"Content-Type": "image/png"})
route = respx_mock.post(PARSE_URL).respond(json=_parse_response())
response = await litellm.aocr(
model=MODEL,
document={"type": "image_url", "image_url": IMAGE_URL},
api_base=API_BASE,
api_key="azure-key",
)
request = route.calls.last.request
assert request.headers["Authorization"] == "Bearer azure-key"
assert json.loads(request.content) == {
"model": "Cohere-parse-v5",
"document": {"type": "image_url", "image_url": PNG_DATA_URI},
"output_format": "markdown",
}
assert response.pages[0].markdown == "# Receipt\n\nTotal Due: $4.00"
assert response.usage_info.pages_processed == 1
@pytest.mark.asyncio
async def test_aocr_passes_data_uri_through_without_fetching(disable_aiohttp_transport, respx_mock):
route = respx_mock.post(PARSE_URL).respond(json=_parse_response())
await litellm.aocr(
model=MODEL,
document={"type": "image_url", "image_url": PNG_DATA_URI},
api_base=API_BASE,
api_key="azure-key",
output_format="blocks",
)
body = json.loads(route.calls.last.request.content)
assert body["document"]["image_url"] == PNG_DATA_URI
assert body["output_format"] == "blocks"
def test_ocr_sync_inlines_remote_image(respx_mock):
respx_mock.get(IMAGE_URL).respond(content=PNG_BYTES, headers={"Content-Type": "image/png"})
route = respx_mock.post(PARSE_URL).respond(json=_parse_response())
response = litellm.ocr(
model=MODEL,
document={"type": "image_url", "image_url": IMAGE_URL},
api_base=API_BASE,
api_key="azure-key",
)
assert json.loads(route.calls.last.request.content)["document"]["image_url"] == PNG_DATA_URI
assert response.pages[0].markdown == "# Receipt\n\nTotal Due: $4.00"
@pytest.mark.asyncio
async def test_aocr_rejects_pdf_before_calling_foundry(disable_aiohttp_transport, respx_mock):
route = respx_mock.post(PARSE_URL).respond(json=_parse_response())
with pytest.raises(litellm.BadRequestError, match="only accepts `image_url` documents") as exc_info:
await litellm.aocr(
model=MODEL,
document={"type": "document_url", "document_url": "https://example.com/doc.pdf"},
api_base=API_BASE,
api_key="azure-key",
)
assert exc_info.value.llm_provider == "azure_ai"
assert not route.called
@pytest.mark.asyncio
async def test_ahealth_check_ocr_sends_an_image_to_the_foundry_cohere_parse_deployment(
disable_aiohttp_transport, respx_mock
):
route = respx_mock.post(PARSE_URL).respond(json=_parse_response())
result = await litellm.ahealth_check(
model_params={"model": MODEL, "api_base": API_BASE, "api_key": "test-key"}, mode="ocr"
)
document = json.loads(route.calls.last.request.content)["document"]
assert document["type"] == "image_url"
assert document["image_url"].startswith("data:image/png;base64,")
assert "error" not in result

View file

@ -1,8 +1,11 @@
import json
from typing import Final
from unittest.mock import Mock
import httpx
import pytest
import litellm
from litellm.llms.cohere.ocr.transformation import CohereParseConfig
PARSE_URL = "https://api.cohere.com/v2/parse"
MODEL = "cohere/parse-v5.0"
@ -57,173 +60,38 @@ def _blocks_response() -> dict:
}
@pytest.fixture()
def disable_aiohttp_transport(monkeypatch):
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
litellm.in_memory_llm_clients_cache.flush_cache()
yield
litellm.in_memory_llm_clients_cache.flush_cache()
@pytest.mark.asyncio
async def test_aocr_sends_markdown_parse_request_and_normalizes_pages(disable_aiohttp_transport, respx_mock):
route = respx_mock.post(PARSE_URL).respond(json=_markdown_response())
response = await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key")
request = route.calls.last.request
assert request.headers["Authorization"] == "Bearer test-key"
assert json.loads(request.content) == {
"model": "parse-v5.0",
"document": IMAGE_DOCUMENT,
"output_format": "markdown",
}
assert response.object == "ocr"
assert [page.index for page in response.pages] == [0, 1]
assert response.pages[0].markdown == "# Receipt\n\nTotal Due: $4.00"
assert response.pages[1].markdown == "Page two"
assert response.pages[1].images is None
image = response.pages[0].images[0]
assert image.bbox == BOUNDING_BOX
assert image.model_extra["description"] == "A parking receipt"
assert image.model_extra["bounding_box_normalized"]["bottom_right_x"] == 1
assert response.usage_info.pages_processed == 2
assert response.get_provider_native_response() is None
@pytest.mark.asyncio
async def test_aocr_usage_prefers_billed_units_over_page_count(disable_aiohttp_transport, respx_mock):
respx_mock.post(PARSE_URL).respond(json=_markdown_response(billed_pages=3))
response = await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key")
assert response.usage_info.pages_processed == 3
@pytest.mark.asyncio
async def test_aocr_usage_falls_back_to_page_count_without_meta(disable_aiohttp_transport, respx_mock):
respx_mock.post(PARSE_URL).respond(json=_markdown_response(billed_pages=None))
response = await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key")
assert response.usage_info.pages_processed == 2
@pytest.mark.asyncio
async def test_aocr_blocks_output_format_forwards_param_and_keeps_blocks(disable_aiohttp_transport, respx_mock):
route = respx_mock.post(PARSE_URL).respond(json=_blocks_response())
response = await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key", output_format="blocks")
assert json.loads(route.calls.last.request.content)["output_format"] == "blocks"
assert response.pages[0].markdown == ""
assert response.pages[0].model_extra["blocks"] == [{"type": "text", "text": "Total Due: $4.00"}]
assert response.usage_info.pages_processed == 1
@pytest.mark.asyncio
async def test_aocr_native_format_carries_provider_payload(disable_aiohttp_transport, respx_mock):
payload = _markdown_response()
route = respx_mock.post(PARSE_URL).respond(json=payload)
response = await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key", req_format="native")
assert "req_format" not in json.loads(route.calls.last.request.content)
assert response.get_provider_native_response() == payload
assert response.pages[0].markdown == "# Receipt\n\nTotal Due: $4.00"
@pytest.mark.asyncio
async def test_aocr_rejects_unknown_output_format_before_calling_provider(disable_aiohttp_transport, respx_mock):
route = respx_mock.post(PARSE_URL).respond(json=_markdown_response())
with pytest.raises(litellm.BadRequestError, match="Invalid `output_format`: 'html'") as exc_info:
await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key", output_format="html")
assert exc_info.value.status_code == 400
assert not route.called
@pytest.mark.asyncio
@pytest.mark.parametrize(
"document",
[
{"type": "document_url", "document_url": "https://example.com/doc.pdf"},
{"type": "image_url", "image_url": "data:application/pdf;base64,JVBERi0="},
{"type": "image_url", "image_url": ""},
],
)
async def test_aocr_rejects_non_image_documents_before_calling_provider(
disable_aiohttp_transport, respx_mock, document
):
route = respx_mock.post(PARSE_URL).respond(json=_markdown_response())
with pytest.raises(litellm.BadRequestError, match="only accepts `image_url` documents") as exc_info:
await litellm.aocr(model=MODEL, document=document, api_key="test-key")
assert exc_info.value.status_code == 400
assert not route.called
@pytest.mark.asyncio
@pytest.mark.parametrize(
"api_base, expected_url",
[
("https://gateway.example.com", "https://gateway.example.com/v2/parse"),
("https://gateway.example.com/cohere/", "https://gateway.example.com/cohere/v2/parse"),
("https://gateway.example.com/v2", "https://gateway.example.com/v2/parse"),
("https://gateway.example.com/v2/parse", "https://gateway.example.com/v2/parse"),
],
)
async def test_aocr_posts_to_api_base_variants(disable_aiohttp_transport, respx_mock, api_base, expected_url):
route = respx_mock.post(expected_url).respond(json=_markdown_response())
await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key", api_base=api_base)
assert route.called
@pytest.mark.asyncio
async def test_aocr_surfaces_provider_error_with_its_status_and_message(disable_aiohttp_transport, respx_mock):
respx_mock.post(PARSE_URL).respond(
status_code=400, json={"id": "83b0d95e", "message": "output_format must be `blocks` or `markdown`"}
@pytest.mark.parametrize("output_format", ["markdown", "blocks"])
def test_transform_cohere_request_filters_options(output_format: str) -> None:
config: Final = CohereParseConfig()
params: Final = config.map_ocr_params(
{"output_format": output_format, "req_format": "native", "unknown": True}, {}, "parse-v5.0"
)
with pytest.raises(litellm.BadRequestError, match="output_format must be") as exc_info:
await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key")
assert exc_info.value.status_code == 400
request: Final = config.transform_ocr_request("parse-v5.0", IMAGE_DOCUMENT, params, {})
assert request.data == {"model": "parse-v5.0", "document": IMAGE_DOCUMENT, "output_format": output_format}
@pytest.mark.asyncio
async def test_aocr_reads_api_key_from_environment(disable_aiohttp_transport, respx_mock, monkeypatch):
monkeypatch.setenv("COHERE_API_KEY", "env-key")
route = respx_mock.post(PARSE_URL).respond(json=_markdown_response())
await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT)
assert route.calls.last.request.headers["Authorization"] == "Bearer env-key"
@pytest.mark.parametrize("native", [False, True])
def test_transform_cohere_response_keeps_images_and_native_payload(native: bool) -> None:
payload: Final = _markdown_response(3)
response: Final = CohereParseConfig().transform_ocr_response(
"parse-v5.0", httpx.Response(200, json=payload), Mock(), {"req_format": "native" if native else "litellm"}
)
assert response.pages[0].markdown == "# Receipt\n\nTotal Due: $4.00"
assert response.pages[0].images[0].bbox == BOUNDING_BOX
assert response.pages[0].images[0].model_extra["description"] == "A parking receipt"
assert response.pages[1].images is None
assert response.usage_info.pages_processed == 3
assert response.get_provider_native_response() == (payload if native else None)
@pytest.mark.asyncio
async def test_aocr_without_api_key_names_the_env_var(disable_aiohttp_transport, respx_mock, monkeypatch):
monkeypatch.delenv("COHERE_API_KEY", raising=False)
monkeypatch.setattr(litellm, "cohere_key", None)
route = respx_mock.post(PARSE_URL).respond(json=_markdown_response())
with pytest.raises(Exception, match="Missing COHERE_API_KEY"):
await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT)
assert not route.called
def test_transform_cohere_blocks() -> None:
response: Final = CohereParseConfig().transform_ocr_response(
"parse-v5.0", httpx.Response(200, json=_blocks_response()), Mock()
)
assert response.pages[0].model_extra["blocks"] == [{"type": "text", "text": "Total Due: $4.00"}]
assert response.pages[0].markdown == ""
@pytest.mark.asyncio
async def test_ahealth_check_ocr_sends_an_image_cohere_parse_accepts(disable_aiohttp_transport, respx_mock):
route = respx_mock.post(PARSE_URL).respond(json=_markdown_response())
result = await litellm.ahealth_check(model_params={"model": MODEL, "api_key": "test-key"}, mode="ocr")
document = json.loads(route.calls.last.request.content)["document"]
assert document["type"] == "image_url"
assert document["image_url"].startswith("data:image/png;base64,")
assert "error" not in result
def test_transform_cohere_rejects_unsupported_output_format() -> None:
with pytest.raises(litellm.UnsupportedParamsError, match="output_format"):
CohereParseConfig().map_ocr_params({"output_format": "html"}, {}, "parse-v5.0")

View file

@ -15,7 +15,7 @@ Streaming: CSW.__anext__ stores args on logging_obj at stream end.
"""
import asyncio
from typing import Any
from typing import Any, Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -297,6 +297,38 @@ async def test_no_flag_fires_create_task_normally():
# ---------------------------------------------------------------------------
@pytest.mark.parametrize("call_type", ["ocr", "aocr", "completion", "acompletion", "embedding", "responses"])
@pytest.mark.parametrize("exception_raised", [False, True])
def test_native_pending_logging_is_released_only_for_ocr(call_type: str, exception_raised: bool) -> None:
pending: Final = MagicMock()
enqueue: Final = MagicMock()
logger: Final = MagicMock(
call_type=call_type,
_native_pending_logging=pending,
_enqueue_deferred_logging=enqueue,
)
ProxyBaseLLMRequestProcessing._flush_deferred_async_logging(
logging_obj=logger,
exception_raised=exception_raised,
)
ProxyBaseLLMRequestProcessing._flush_deferred_async_logging(
logging_obj=logger,
exception_raised=exception_raised,
)
if call_type in ("ocr", "aocr"):
pending.release.assert_called_once_with(not exception_raised)
assert logger._native_pending_logging is None
else:
pending.release.assert_not_called()
assert logger._native_pending_logging is pending
if exception_raised:
enqueue.assert_not_called()
else:
enqueue.assert_called_once_with()
def test_flush_deferred_async_logging_fires_on_success():
"""
Happy path: with no exception, the production flush helper invokes the

View file

@ -10,11 +10,11 @@ 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:
@pytest.mark.parametrize("enabled", [True, False])
def test_public_selection_requires_available_native_ocr(enabled: bool) -> None:
native: Final = Mock(side_effect=AssertionError("must not admit"))
litellm.rust(enabled)
NATIVE_OCR_LIFECYCLE.override(native if available else None)
NATIVE_OCR_LIFECYCLE.override(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"})
@ -96,7 +96,7 @@ def test_public_binding_keeps_keyword_model_and_document_in_native_hook_kwargs()
assert "timeout" not in captured[0]
@pytest.mark.parametrize("enabled", [False, True], ids=["legacy", "native"])
@pytest.mark.parametrize("enabled", [False, True], ids=["flag-disabled", "flag-enabled"])
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"}
@ -111,7 +111,7 @@ def test_public_duplicate_argument_error_does_not_depend_on_native_selection(ena
assert native.call_count == 0
@pytest.mark.parametrize("enabled", [False, True], ids=["legacy", "native"])
@pytest.mark.parametrize("enabled", [False, True], ids=["flag-disabled", "flag-enabled"])
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)
@ -123,3 +123,27 @@ def test_public_missing_required_argument_error_does_not_depend_on_native_select
NATIVE_OCR_LIFECYCLE.reset()
litellm.rust(None)
assert native.call_count == 0
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
@pytest.mark.parametrize("enabled", [False, True, None])
async def test_public_ocr_ignores_rust_flag(
monkeypatch: pytest.MonkeyPatch, asynchronous: bool, enabled: bool | None
) -> None:
from unittest.mock import AsyncMock
monkeypatch.setenv("LITELLM_RUST", "0")
response: Final = OCRResponse(pages=[], model="mistral-ocr-latest")
native: Final = AsyncMock(return_value=response) if asynchronous else Mock(return_value=response)
litellm.rust(enabled)
NATIVE_OCR_LIFECYCLE.override(native)
try:
if asynchronous:
assert await litellm.aocr("mistral/mistral-ocr-latest", {}) is response
else:
assert litellm.ocr("mistral/mistral-ocr-latest", {}) is response
assert native.call_count == 1
finally:
NATIVE_OCR_LIFECYCLE.reset()
litellm.rust(None)

View file

@ -1,5 +1,6 @@
import asyncio
import contextlib
import contextvars
import json
import logging
import os
@ -7,6 +8,7 @@ import queue
import threading
from datetime import datetime, timedelta, timezone
from collections.abc import Iterator
from concurrent.futures import ThreadPoolExecutor
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
@ -60,6 +62,36 @@ from litellm.utils import (
# Adds the parent directory to the system path
def test_non_ocr_wrapper_preserves_logging_executor_and_context(monkeypatch: pytest.MonkeyPatch) -> None:
marker: Final = contextvars.ContextVar("non-ocr-logging-context", default="missing")
token: Final = marker.set("caller-context")
caller_thread: Final = threading.get_ident()
response: Final = object()
logger: Final = MagicMock()
observed: Final = queue.Queue[tuple[object, str, int]]()
def record_success(result: object, start_time: datetime, end_time: datetime) -> None:
observed.put((result, marker.get(), threading.get_ident()))
def embedding(**kwargs: object) -> object:
return response
logger.success_handler.side_effect = record_success
monkeypatch.setattr("litellm.utils.function_setup", MagicMock(return_value=(logger, {})))
try:
with ThreadPoolExecutor(max_workers=1) as executor:
monkeypatch.setattr("litellm.utils.executor", executor)
result: Final = client(embedding)()
logged_response, context, worker_thread = observed.get_nowait()
assert result is response
assert logged_response is response
assert context == "caller-context"
assert worker_thread != caller_thread
assert observed.empty()
finally:
marker.reset(token)
def test_cloudflare_model_info_includes_rpm(local_model_cost_map: None) -> None:
assert litellm.get_model_info("cloudflare/@cf/meta/llama-3.1-8b-instruct-fp8")["rpm"] == 300
assert litellm.get_model_info("cloudflare/@cf/moonshotai/kimi-k2.6")["rpm"] == 20

View file

@ -0,0 +1,141 @@
from typing import Final
import pytest
import litellm
from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec
pytestmark = pytest.mark.requires_rust_extension
MODELS: Final = ("cohere/parse-v5.0", "azure_ai/Cohere-parse-v5.0")
IMAGE: Final = {"type": "image_url", "image_url": "data:image/png;base64,YWJj"}
BOX: Final = {"top_left_x": 0, "top_left_y": 0, "bottom_right_x": 32, "bottom_right_y": 32}
PAYLOAD: Final = {
"pages": [
{
"index": 4,
"markdown": {"content": "receipt", "images": [{"id": "image", "bounding_box": BOX, "description": "scan"}]},
},
{"markdown": {"content": "page two"}},
],
"meta": {"billed_units": {"pages": 3}},
}
@pytest.mark.asyncio
@pytest.mark.parametrize("model", MODELS)
@pytest.mark.parametrize("asynchronous", [False, True])
async def test_public_cohere_request_and_normalization(
recording_server: RecordingServer, model: str, asynchronous: bool
) -> None:
recording_server.enqueue(ResponseSpec(body=PAYLOAD))
args: Final = {
"model": model,
"document": IMAGE,
"api_base": recording_server.base_url,
"api_key": "test-key",
"req_format": "native",
"unrecognized": True,
}
response: Final = await litellm.aocr(**args) if asynchronous else litellm.ocr(**args)
request: Final = recording_server.requests[0]
assert request.path == ("/providers/cohere/v2/parse" if model.startswith("azure_ai/") else "/v2/parse")
assert request.headers["authorization"] == "Bearer test-key"
assert request.body == {"model": model.split("/", 1)[1], "document": IMAGE, "output_format": "markdown"}
assert [page.index for page in response.pages] == [4, 1]
assert response.pages[0].markdown == "receipt"
assert response.pages[0].images[0].bbox == BOX
assert response.pages[0].images[0].model_extra["description"] == "scan"
assert response.pages[1].images is None
assert response.usage_info.pages_processed == 3
assert response.get_provider_native_response() == PAYLOAD
@pytest.mark.asyncio
@pytest.mark.parametrize("model", MODELS)
async def test_public_cohere_blocks_and_usage_fallback(recording_server: RecordingServer, model: str) -> None:
blocks: Final = [{"type": "text", "text": "total"}]
recording_server.enqueue(ResponseSpec(body={"pages": [{"blocks": blocks}]}))
response: Final = await litellm.aocr(
model=model, document=IMAGE, api_base=recording_server.base_url, api_key="test-key", output_format="blocks"
)
assert recording_server.requests[0].body["output_format"] == "blocks"
assert response.pages[0].model_extra["blocks"] == blocks
assert response.pages[0].markdown == ""
assert response.usage_info.pages_processed == 1
assert response.get_provider_native_response() is None
@pytest.mark.asyncio
@pytest.mark.parametrize("model", MODELS)
@pytest.mark.parametrize(
"document",
[
{"type": "document_url", "document_url": "https://example.com/file.pdf"},
{"type": "image_url", "image_url": "data:application/pdf;base64,YQ=="},
{"type": "image_url", "image_url": ""},
],
)
async def test_public_cohere_rejects_non_images_before_network(
recording_server: RecordingServer, model: str, document: dict[str, str]
) -> None:
recording_server.expected_requests = 0
with pytest.raises(litellm.BadRequestError, match="only accepts `image_url`"):
await litellm.aocr(model=model, document=document, api_base=recording_server.base_url, api_key="test-key")
@pytest.mark.asyncio
@pytest.mark.parametrize("model", MODELS)
async def test_public_cohere_rejects_unknown_format(recording_server: RecordingServer, model: str) -> None:
recording_server.expected_requests = 0
with pytest.raises(litellm.BadRequestError, match="output_format"):
await litellm.aocr(
model=model, document=IMAGE, api_base=recording_server.base_url, api_key="test-key", output_format="html"
)
@pytest.mark.asyncio
@pytest.mark.parametrize("model", MODELS)
async def test_public_cohere_provider_failure(recording_server: RecordingServer, model: str) -> None:
recording_server.enqueue(ResponseSpec(status=400, body={"message": "output_format must be blocks or markdown"}))
with pytest.raises(litellm.BadRequestError, match="output_format must be") as caught:
await litellm.aocr(model=model, document=IMAGE, api_base=recording_server.base_url, api_key="test-key")
assert caught.value.status_code == 400
@pytest.mark.asyncio
@pytest.mark.parametrize("model", MODELS)
async def test_public_cohere_health_check(recording_server: RecordingServer, model: str) -> None:
recording_server.enqueue(ResponseSpec(body=PAYLOAD))
response: Final = await litellm.ahealth_check(
model_params={"model": model, "api_key": "test-key", "api_base": recording_server.base_url}, mode="ocr"
)
assert "error" not in response
assert recording_server.requests[0].body["document"]["image_url"].startswith("data:image/png;base64,")
@pytest.mark.asyncio
@pytest.mark.parametrize("suffix", ["", "/cohere/", "/v2", "/v2/parse"])
async def test_public_cohere_url_variants(recording_server: RecordingServer, suffix: str) -> None:
recording_server.enqueue(ResponseSpec(body=PAYLOAD))
await litellm.aocr(model=MODELS[0], document=IMAGE, api_base=recording_server.base_url + suffix, api_key="test-key")
assert recording_server.requests[0].path == ("/cohere/v2/parse" if suffix == "/cohere/" else "/v2/parse")
@pytest.mark.asyncio
async def test_public_cohere_environment_key_and_remote_url(
recording_server: RecordingServer, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setenv("COHERE_API_KEY", "env-key")
recording_server.enqueue(ResponseSpec(body=PAYLOAD))
document: Final = {"type": "image_url", "image_url": "https://example.com/receipt.png"}
await litellm.aocr(model=MODELS[0], document=document, api_base=recording_server.base_url)
assert recording_server.requests[0].headers["authorization"] == "Bearer env-key"
assert recording_server.requests[0].body["document"] == document
@pytest.mark.asyncio
async def test_public_cohere_missing_key(recording_server: RecordingServer, monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.delenv("COHERE_API_KEY", raising=False)
recording_server.expected_requests = 0
with pytest.raises(Exception, match="Missing COHERE_API_KEY"):
await litellm.aocr(model=MODELS[0], document=IMAGE, api_base=recording_server.base_url)

View file

@ -1,3 +1,4 @@
import sys
from typing import Final
import pytest
@ -16,8 +17,26 @@ def ocr_server(recording_server: RecordingServer) -> RecordingServer:
return recording_server
def test_public_ocr_uses_native_route_when_enabled(ocr_server: RecordingServer) -> None:
litellm.rust(True)
def test_native_ocr_rejects_disabled_gil_before_provider_call(
ocr_server: RecordingServer, monkeypatch: pytest.MonkeyPatch
) -> None:
ocr_server.expected_requests = 0
monkeypatch.setattr(sys, "_is_gil_enabled", lambda: False, raising=False)
with pytest.raises(RuntimeError, match="native OCR requires the Python GIL"):
litellm.ocr(
model=OCR_MODEL,
document=OCR_DOCUMENT,
api_key="test-key",
api_base=ocr_server.base_url,
)
assert not ocr_server.requests
@pytest.mark.parametrize("enabled", [False, True, None])
def test_public_ocr_uses_native_route_independently_of_flag(ocr_server: RecordingServer, enabled: bool | None) -> None:
litellm.rust(enabled)
response: Final = litellm.ocr(
model=OCR_MODEL,
document=OCR_DOCUMENT,
@ -29,18 +48,3 @@ def test_public_ocr_uses_native_route_when_enabled(ocr_server: RecordingServer)
assert response.pages[0].markdown == "native OCR response"
assert len(ocr_server.requests) == 1
assert not ocr_server.requests[0].headers.get("user-agent", "").startswith("python-httpx")
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

@ -1,6 +1,7 @@
import asyncio
import datetime
import gc
import json
import sys
import threading
import weakref
@ -711,7 +712,9 @@ async def test_reducto_lifecycle_retains_upload_parse_and_post_call_boundaries(
@pytest.mark.asyncio
async def test_document_intelligence_post_call_runs_before_polling(ocr_server: RecordingServer) -> None:
async def test_document_intelligence_post_call_observes_submission_and_final_result(
ocr_server: RecordingServer,
) -> None:
ocr_server.expected_requests = 2
ocr_server.enqueue(
ResponseSpec(
@ -725,7 +728,7 @@ async def test_document_intelligence_post_call_runs_before_polling(ocr_server: R
class Observe(Logging):
def post_call(self, *args, **kwargs):
boundaries.append(tuple(request.method for request in ocr_server.requests))
boundaries.append((tuple(request.method for request in ocr_server.requests), kwargs["original_response"]))
return super().post_call(*args, **kwargs)
logger: Final = Observe(
@ -740,7 +743,9 @@ async def test_document_intelligence_post_call_runs_before_polling(ocr_server: R
response: Final = await call_aocr(
ocr_server, model="azure_ai/doc-intelligence/prebuilt-read", litellm_logging_obj=logger
)
assert boundaries == [("POST",)]
assert [methods for methods, _ in boundaries] == [("POST",), ("POST", "GET")]
assert json.loads(boundaries[0][1])["status"] == "running"
assert json.loads(boundaries[1][1])["status"] == "succeeded"
assert [request.method for request in ocr_server.requests] == ["POST", "GET"]
assert ocr_server.requests[1].path == "/operations/1"
assert response.pages == []