fix(ocr-errors): preserve public error and timeout contracts (#33605)

* fix(ocr-errors): preserve public error and timeout contracts

Classify reqwest timeouts as a typed CoreError::Timeout and add
CoreError::public_status_code() as the single exhaustive mapping from a
typed core error to its public HTTP status. The Python bridge raises a
typed RustOcrError carrying that status instead of a generic
RuntimeError, and litellm.ocr()/aocr() translate it into the matching
public exception so AuthenticationError/401, NotFoundError/404,
BadRequestError/4xx, InternalServerError/5xx and Timeout are preserved
end to end instead of collapsing to APIConnectionError/500.

Reject empty or whitespace-only 200 bodies so they fail loudly rather
than becoming an empty OCR success; invalid JSON already fails. Upstream
error bodies stay bounded and sanitized.

Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>

* fix(ocr-errors): map invalid OCR input to BadRequestError

Invalid caller input (bad document type, non-dict document, unusable
file input) was raised as a plain ValueError inside ocr()/aocr() and
collapsed to APIConnectionError/500 through the generic handler. Route
every OCR failure through one _map_ocr_exception host mapping: typed
RustOcrError keeps its status-based public exception, a plain ValueError
becomes BadRequestError/400, and a pydantic ValidationError (malformed
response, not client input) stays on the generic path.

Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>

* refactor(ocr-errors): data-minimize public errors and preserve status-specific exceptions

Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>

* refactor(ocr-errors): preserve exact unknown status and privatize input error

Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>

* refactor(ocr-errors): exhaustive match mapper and typed error tests

Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>

* fix(ocr-errors): sanitize InvalidRequest public message

Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>

* refactor(ocr-errors): raise typed public union and drop NotFound provider miss

Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>

* fix(ocr-errors): map malformed provider responses to sanitized 500 and hide input-error detail

Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>

---------

Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
Co-authored-by: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-07-16 19:49:26 -07:00 • committed by GitHub
parent c264758eb8
commit 6661462d5a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
12 changed files with 911 additions and 99 deletions

View file

@ -17,8 +17,8 @@ use serde_json::{Map, Value};
mod common_utils;
use common_utils::{
convert_document_url_to_data_uri, has_header, ocr_provider_config, poll_document_intelligence,
string_headers, truncate_error_body, upload_reducto_document,
classify_reqwest_error, convert_document_url_to_data_uri, has_header, ocr_provider_config,
poll_document_intelligence, string_headers, truncate_error_body, upload_reducto_document,
};
/// OCR over large documents can take a while; bound it generously rather than
@ -74,8 +74,12 @@ pub struct OcrRequest<'a> {
/// Async: intended to be awaited directly by the Python bridge's async entrypoint.
pub async fn ocr(request: OcrRequest<'_>) -> CoreResult<Value> {
let model = request.model;
let config = ocr_provider_config(request.custom_llm_provider, model)
.ok_or_else(|| CoreError::InvalidProvider(request.custom_llm_provider.to_string()))?;
let config = ocr_provider_config(request.custom_llm_provider, model).ok_or_else(|| {
CoreError::InvalidProvider(format!(
"no OCR provider '{}' registered for model '{model}'",
request.custom_llm_provider
))
})?;
let env_lookup = |key: &str| std::env::var(key).ok();
let headers = string_headers(request.extra_headers)?;
@ -116,7 +120,7 @@ pub async fn ocr(request: OcrRequest<'_>) -> CoreResult<Value> {
let response = request_builder
.send()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
.map_err(classify_reqwest_error)?;
let status = response.status();
if config.response_handling() == OcrResponseHandling::AzureDocumentIntelligencePoll
@ -141,10 +145,7 @@ pub async fn ocr(request: OcrRequest<'_>) -> CoreResult<Value> {
.into_json());
}
let text = response
.text()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
let text = response.text().await.map_err(classify_reqwest_error)?;
if !status.is_success() {
return Err(CoreError::Http {
@ -153,6 +154,12 @@ pub async fn ocr(request: OcrRequest<'_>) -> CoreResult<Value> {
});
}
if text.trim().is_empty() {
return Err(CoreError::InvalidResponse(
"OCR provider returned an empty success response".to_string(),
));
}
let response_json: Value = serde_json::from_str(&text)
.map_err(|err| CoreError::InvalidResponse(format!("invalid OCR response JSON: {err}")))?;
@ -412,4 +419,144 @@ mod tests {
)
);
}
async fn respond_once(status_line: &'static str, body: String) -> String {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("test listener binds");
let addr = listener.local_addr().expect("listener has local addr");
tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts one request");
let _ = read_http_headers(&mut socket).await;
let response = format!(
"HTTP/1.1 {status_line}\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
body.len(),
body
);
socket
.write_all(response.as_bytes())
.await
.expect("writes response");
});
format!("http://{addr}")
}
async fn run_mistral_ocr(api_base: String, timeout: Duration) -> CoreResult<Value> {
ocr(OcrRequest {
model: "mistral-ocr-latest",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
api_key: Some("sk-test"),
api_base: Some(&api_base),
custom_llm_provider: "mistral",
extra_headers: None,
optional_params: Map::new(),
timeout: Some(timeout),
})
.await
}
#[tokio::test]
async fn ocr_preserves_upstream_error_status() {
for status in [
"401 Unauthorized",
"404 Not Found",
"500 Internal Server Error",
] {
let base = respond_once(status, r#"{"error":"nope"}"#.to_string()).await;
let err = run_mistral_ocr(base, Duration::from_secs(5))
.await
.expect_err("upstream error surfaces");
let expected = status[..3].parse::<u16>().expect("status prefix parses");
match err {
CoreError::Http { status: got, .. } => assert_eq!(got, expected),
other => panic!("expected Http error, got {other:?}"),
}
assert_eq!(err.public_status_code(), Some(expected));
}
}
#[tokio::test]
async fn ocr_bounds_oversized_error_body() {
let base = respond_once("500 Internal Server Error", "x".repeat(6000)).await;
let err = run_mistral_ocr(base, Duration::from_secs(5))
.await
.expect_err("oversized error surfaces");
match err {
CoreError::Http { body, .. } => {
assert!(body.ends_with("... (truncated)"));
assert!(
body.chars().count() < 300,
"body not bounded: {} chars",
body.chars().count()
);
}
other => panic!("expected Http error, got {other:?}"),
}
}
#[tokio::test]
async fn ocr_rejects_invalid_json_success_body() {
let base = respond_once("200 OK", "not json".to_string()).await;
let err = run_mistral_ocr(base, Duration::from_secs(5))
.await
.expect_err("invalid JSON surfaces");
assert!(matches!(err, CoreError::InvalidResponse(_)));
assert_eq!(err.public_status_code(), Some(500));
}
#[tokio::test]
async fn ocr_rejects_empty_success_body() {
let base = respond_once("200 OK", " ".to_string()).await;
let err = run_mistral_ocr(base, Duration::from_secs(5))
.await
.expect_err("empty success surfaces");
match err {
CoreError::InvalidResponse(message) => assert!(message.contains("empty")),
other => panic!("expected InvalidResponse, got {other:?}"),
}
}
#[tokio::test]
async fn ocr_classifies_per_request_timeout() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("test listener binds");
let addr = listener.local_addr().expect("listener has local addr");
tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts one request");
let _ = read_http_headers(&mut socket).await;
tokio::time::sleep(Duration::from_secs(3)).await;
let _ = socket.write_all(b"HTTP/1.1 200 OK\r\n\r\n").await;
});
let err = run_mistral_ocr(format!("http://{addr}"), Duration::from_millis(200))
.await
.expect_err("timeout surfaces");
assert_eq!(err, CoreError::Timeout);
assert_eq!(err.public_status_code(), Some(408));
}
#[tokio::test]
async fn ocr_maps_unregistered_provider_to_invalid_provider() {
let err = ocr(OcrRequest {
model: "some-model",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
api_key: Some("sk-test"),
api_base: None,
custom_llm_provider: "definitely-not-a-provider",
extra_headers: None,
optional_params: Map::new(),
timeout: Some(Duration::from_secs(5)),
})
.await
.expect_err("unregistered provider surfaces");
assert!(matches!(err, CoreError::InvalidProvider(_)));
assert_eq!(err.public_status_code(), Some(400));
assert_eq!(err.public_message(), "Invalid OCR request");
}
}

View file

@ -28,6 +28,14 @@ const AZURE_DOCUMENT_INTELLIGENCE_POLL_TIMEOUT_SECS: u64 = 120;
const DEFAULT_MAX_IMAGE_URL_DOWNLOAD_SIZE_MB: f64 = 50.0;
const MAX_SAFE_FETCH_REDIRECTS: usize = 10;
pub(super) fn classify_reqwest_error(err: reqwest::Error) -> CoreError {
if err.is_timeout() {
CoreError::Timeout
} else {
CoreError::Network(err.to_string())
}
}
pub(super) fn truncate_error_body(body: &str) -> String {
if body.chars().count() <= ERROR_BODY_MAX_CHARS {
return body.to_string();
@ -227,7 +235,7 @@ async fn safe_get_document_url(url: &str) -> CoreResult<(Url, reqwest::Response)
.get(current_url.clone())
.send()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
.map_err(classify_reqwest_error)?;
if !response.status().is_redirection() {
return Ok((current_url, response));
}
@ -268,11 +276,7 @@ async fn read_response_with_limit(
let mut bytes = Vec::new();
let mut bytes_downloaded: u64 = 0;
while let Some(chunk) = response
.chunk()
.await
.map_err(|err| CoreError::Network(err.to_string()))?
{
while let Some(chunk) = response.chunk().await.map_err(classify_reqwest_error)? {
bytes_downloaded += chunk.len() as u64;
enforce_download_size(bytes_downloaded, max_bytes, url)?;
bytes.extend_from_slice(&chunk);
@ -388,12 +392,9 @@ async fn upload_reducto_bytes(
let response = request_builder
.send()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
.map_err(classify_reqwest_error)?;
let status = response.status();
let text = response
.text()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
let text = response.text().await.map_err(classify_reqwest_error)?;
if !status.is_success() {
return Err(CoreError::Http {
status: status.as_u16(),
@ -477,7 +478,7 @@ fn operation_status(response_json: &Value) -> CoreResult<&str> {
let status = response_json
.get("status")
.and_then(Value::as_str)
.ok_or(CoreError::MissingField("status"))?;
.ok_or(CoreError::missing_response_field("status"))?;
match status {
"succeeded" => Ok("succeeded"),
"running" | "notStarted" => Ok("running"),
@ -515,10 +516,7 @@ pub(super) async fn poll_document_intelligence(
));
loop {
if start.elapsed() > timeout {
return Err(CoreError::Network(format!(
"Azure Document Intelligence operation polling timed out after {} seconds",
timeout.as_secs()
)));
return Err(CoreError::Timeout);
}
let mut request_builder = http_client().get(operation_url);
@ -530,13 +528,10 @@ pub(super) async fn poll_document_intelligence(
let response = request_builder
.send()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
.map_err(classify_reqwest_error)?;
let retry_after = retry_after_secs(&response);
let status = response.status();
let text = response
.text()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
let text = response.text().await.map_err(classify_reqwest_error)?;
if !status.is_success() {
return Err(CoreError::Http {
status: status.as_u16(),

View file

@ -21,12 +21,59 @@ pub enum CoreError {
Auth(String),
#[error("OCR request failed with status {status}: {body}")]
Http { status: u16, body: String },
#[error("OCR request timed out")]
Timeout,
#[error("OCR network error: {0}")]
Network(String),
#[error("routing error: {0}")]
Routing(String),
}
impl CoreError {
pub fn unexpected_response_type(value: &serde_json::Value) -> Self {
CoreError::InvalidResponse(format!(
"expected object OCR response, got {}",
json_type_name(value)
))
}
pub fn missing_response_field(field: &'static str) -> Self {
CoreError::InvalidResponse(format!("OCR response missing required field: {field}"))
}
pub fn public_status_code(&self) -> Option<u16> {
match self {
CoreError::Http { status, .. } => Some(*status),
CoreError::Auth(_) => Some(401),
CoreError::InvalidType { .. }
| CoreError::MissingField(_)
| CoreError::InvalidProvider(_)
| CoreError::InvalidRequest(_) => Some(400),
CoreError::Timeout => Some(408),
CoreError::InvalidResponse(_) | CoreError::Routing(_) => Some(500),
CoreError::Network(_) => None,
}
}
pub fn public_message(&self) -> String {
match self {
CoreError::Http { status, .. } => format!("OCR request failed with status {status}"),
CoreError::Network(_) => "OCR request could not reach the provider".to_string(),
CoreError::InvalidResponse(_) => {
"OCR provider returned an invalid response".to_string()
}
CoreError::Routing(_) => "OCR request could not be routed".to_string(),
CoreError::Auth(_) => "OCR request failed provider authentication".to_string(),
CoreError::InvalidRequest(_) | CoreError::InvalidProvider(_) => {
"Invalid OCR request".to_string()
}
CoreError::Timeout | CoreError::InvalidType { .. } | CoreError::MissingField(_) => {
self.to_string()
}
}
}
}
pub fn json_type_name(value: &serde_json::Value) -> &'static str {
match value {
serde_json::Value::Null => "null",
@ -37,3 +84,172 @@ pub fn json_type_name(value: &serde_json::Value) -> &'static str {
serde_json::Value::Object(_) => "object",
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn public_status_code_preserves_public_contracts() {
assert_eq!(
CoreError::Http {
status: 404,
body: "not found".to_string()
}
.public_status_code(),
Some(404)
);
assert_eq!(
CoreError::Auth("bad key".to_string()).public_status_code(),
Some(401)
);
assert_eq!(CoreError::Timeout.public_status_code(), Some(408));
assert_eq!(
CoreError::InvalidRequest("bad".to_string()).public_status_code(),
Some(400)
);
assert_eq!(
CoreError::MissingField("document.type").public_status_code(),
Some(400)
);
assert_eq!(
CoreError::InvalidType {
expected: "object",
actual: "string"
}
.public_status_code(),
Some(400)
);
assert_eq!(
CoreError::InvalidResponse("empty".to_string()).public_status_code(),
Some(500)
);
assert_eq!(
CoreError::Routing("no deployment".to_string()).public_status_code(),
Some(500)
);
assert_eq!(
CoreError::Network("dns".to_string()).public_status_code(),
None
);
assert_eq!(
CoreError::InvalidProvider("mistral".to_string()).public_status_code(),
Some(400)
);
}
#[test]
fn timeout_message_is_data_minimized() {
assert_eq!(CoreError::Timeout.to_string(), "OCR request timed out");
}
#[test]
fn public_message_hides_upstream_body() {
let err = CoreError::Http {
status: 500,
body: "signed-url=https://secret.example/token=abc123 leaked".to_string(),
};
let message = err.public_message();
assert_eq!(message, "OCR request failed with status 500");
assert!(!message.contains("secret"));
assert!(!message.contains("abc123"));
}
#[test]
fn public_message_hides_network_and_response_detail() {
let network = CoreError::Network(
"error sending request for url (https://signed.example/token=xyz)".to_string(),
);
assert_eq!(
network.public_message(),
"OCR request could not reach the provider"
);
assert!(!network.public_message().contains("token=xyz"));
let invalid = CoreError::InvalidResponse("expected value at line 1 column 2".to_string());
assert_eq!(
invalid.public_message(),
"OCR provider returned an invalid response"
);
}
#[test]
fn public_message_hides_auth_detail() {
let err = CoreError::Auth(
"google auth: failed to load service account key /secrets/sa.json token=ya29.abc123"
.to_string(),
);
let message = err.public_message();
assert_eq!(message, "OCR request failed provider authentication");
assert!(!message.contains("sa.json"));
assert!(!message.contains("ya29"));
assert!(!message.contains("token"));
}
#[test]
fn public_message_hides_invalid_request_detail() {
let err = CoreError::InvalidRequest(
"document_url=https://signed.example/doc.pdf?token=SECRET123 \
base64=QUJDREVG header=x-api-key page=42"
.to_string(),
);
let message = err.public_message();
assert_eq!(message, "Invalid OCR request");
assert!(!message.contains("SECRET123"));
assert!(!message.contains("token"));
assert!(!message.contains("base64"));
assert!(!message.contains("QUJDREVG"));
assert!(!message.contains("x-api-key"));
assert!(!message.contains("signed.example"));
assert_eq!(err.public_status_code(), Some(400));
}
#[test]
fn public_message_keeps_only_static_detail() {
assert_eq!(
CoreError::MissingField("document.type").public_message(),
"missing required field: document.type"
);
assert_eq!(
CoreError::InvalidType {
expected: "object",
actual: "string"
}
.public_message(),
"expected object, got string"
);
assert_eq!(
CoreError::Routing("model_list parse failed".to_string()).public_message(),
"OCR request could not be routed"
);
}
#[test]
fn malformed_response_maps_to_sanitized_500() {
let wrong_type = CoreError::unexpected_response_type(&serde_json::json!("boom"));
assert!(matches!(wrong_type, CoreError::InvalidResponse(_)));
assert_eq!(wrong_type.public_status_code(), Some(500));
assert_eq!(
wrong_type.public_message(),
"OCR provider returned an invalid response"
);
let missing = CoreError::missing_response_field("status");
assert!(matches!(missing, CoreError::InvalidResponse(_)));
assert_eq!(missing.public_status_code(), Some(500));
assert_eq!(
missing.public_message(),
"OCR provider returned an invalid response"
);
}
#[test]
fn public_message_hides_invalid_provider_detail() {
let err = CoreError::InvalidProvider(
"no OCR provider 'internal-secret-router' registered for model 'gpt-8'".to_string(),
);
assert_eq!(err.public_message(), "Invalid OCR request");
assert!(!err.public_message().contains("internal-secret-router"));
assert_eq!(err.public_status_code(), Some(400));
}
}

View file

@ -362,14 +362,11 @@ impl OcrProviderConfig for AzureDocumentIntelligenceOcrConfig {
) -> CoreResult<OcrResponseData> {
let response = response_json
.as_object()
.ok_or_else(|| CoreError::InvalidType {
expected: "object",
actual: json_type_name(&response_json),
})?;
.ok_or_else(|| CoreError::unexpected_response_type(&response_json))?;
let status = response
.get("status")
.and_then(Value::as_str)
.ok_or(CoreError::MissingField("status"))?;
.ok_or(CoreError::missing_response_field("status"))?;
if status != "succeeded" {
return Err(CoreError::InvalidResponse(format!(
"Azure Document Intelligence analysis failed with status: {status}"
@ -517,4 +514,24 @@ mod tests {
Some(json!({"pages_processed": 1, "doc_size_bytes": null}))
);
}
#[test]
fn document_intelligence_response_missing_status_is_invalid_response() {
let err = AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG
.transform_ocr_response("prebuilt-layout", json!({"analyzeResult": {"pages": []}}))
.expect_err("missing status should be rejected");
assert!(matches!(err, CoreError::InvalidResponse(_)));
assert_eq!(err.public_status_code(), Some(500));
}
#[test]
fn document_intelligence_response_non_object_is_invalid_response() {
let err = AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG
.transform_ocr_response("prebuilt-layout", json!("boom"))
.expect_err("non-object provider response should be rejected");
assert!(matches!(err, CoreError::InvalidResponse(_)));
assert_eq!(err.public_status_code(), Some(500));
}
}

View file

@ -107,10 +107,7 @@ impl OcrProviderConfig for MistralOcrConfig {
) -> CoreResult<OcrResponseData> {
let response_object = response_json
.as_object()
.ok_or_else(|| CoreError::InvalidType {
expected: "object",
actual: json_type_name(&response_json),
})?;
.ok_or_else(|| CoreError::unexpected_response_type(&response_json))?;
let pages = response_object
.get("pages")
@ -257,6 +254,15 @@ mod tests {
);
}
#[test]
fn transform_ocr_response_rejects_non_object_as_invalid_response() {
let err = transform_ocr_response("mistral-ocr-latest", json!("boom"))
.expect_err("non-object provider response should be rejected");
assert!(matches!(err, CoreError::InvalidResponse(_)));
assert_eq!(err.public_status_code(), Some(500));
}
#[test]
fn transform_ocr_response_normalizes_mistral_json() {
let response = json!({

View file

@ -237,10 +237,7 @@ impl OcrProviderConfig for ReductoParseLegacyConfig {
fn transform_reducto_response(model: &str, response_json: Value) -> CoreResult<OcrResponseData> {
let response = response_json
.as_object()
.ok_or_else(|| CoreError::InvalidType {
expected: "object",
actual: json_type_name(&response_json),
})?;
.ok_or_else(|| CoreError::unexpected_response_type(&response_json))?;
let result = response.get("result").unwrap_or(&response_json);
let usage = response.get("usage").cloned().unwrap_or_else(|| json!({}));
Ok(OcrResponseData {

View file

@ -292,10 +292,7 @@ impl OcrProviderConfig for VertexAiDeepSeekOcrConfig {
) -> CoreResult<OcrResponseData> {
let response = response_json
.as_object()
.ok_or_else(|| CoreError::InvalidType {
expected: "object",
actual: json_type_name(&response_json),
})?;
.ok_or_else(|| CoreError::unexpected_response_type(&response_json))?;
let usage = response.get("usage").cloned();
let content = first_choice_content(&response_json)?;
let mut ocr_data = ocr_data_from_content(content.clone(), usage.clone(), model);
@ -314,10 +311,9 @@ impl OcrProviderConfig for VertexAiDeepSeekOcrConfig {
});
}
let object = ocr_data.as_object().ok_or_else(|| CoreError::InvalidType {
expected: "object",
actual: json_type_name(&ocr_data),
})?;
let object = ocr_data
.as_object()
.ok_or_else(|| CoreError::unexpected_response_type(&ocr_data))?;
let pages = object
.get("pages")
.and_then(Value::as_array)

View file

@ -2,7 +2,7 @@ use std::time::Duration;
use litellm_ai_gateway::io::ocr::{ocr as run_ocr, OcrRequest};
use litellm_core::error::CoreError;
use pyo3::exceptions::{PyRuntimeError, PyValueError};
use pyo3::exceptions::PyValueError;
use pyo3::prelude::*;
use pyo3::types::{PyAny, PyDict};
use serde_json::{Map, Value};
@ -29,15 +29,22 @@ fn json_to_py(py: Python<'_>, value: Value) -> PyResult<Py<PyAny>> {
Ok(json.call_method1("loads", (encoded,))?.unbind())
}
fn core_error_to_pyerr(err: CoreError) -> PyErr {
match err {
CoreError::Auth(message) => PyValueError::new_err(message),
CoreError::InvalidProvider(_)
| CoreError::InvalidRequest(_)
| CoreError::InvalidType { .. }
| CoreError::MissingField(_) => PyValueError::new_err(err.to_string()),
other => PyRuntimeError::new_err(other.to_string()),
}
fn core_error_to_pyerr(py: Python<'_>, err: CoreError) -> PyErr {
let status_code = err.public_status_code();
let message = err.public_message();
build_rust_ocr_error(py, &message, status_code).unwrap_or_else(|import_err| import_err)
}
fn build_rust_ocr_error(
py: Python<'_>,
message: &str,
status_code: Option<u16>,
) -> PyResult<PyErr> {
let exc_type = py
.import("litellm.ocr.rust_bridge")?
.getattr("RustOcrError")?;
let instance = exc_type.call1((message, status_code))?;
Ok(PyErr::from_value(instance))
}
fn optional_object_to_map(
@ -120,7 +127,7 @@ fn ocr(
match result {
Ok(value) => json_to_py(py, value),
Err(err) => Err(core_error_to_pyerr(err)),
Err(err) => Err(core_error_to_pyerr(py, err)),
}
}
@ -159,7 +166,7 @@ fn aocr(
timeout,
})
.await
.map_err(core_error_to_pyerr)?;
.map_err(|err| Python::with_gil(|py| core_error_to_pyerr(py, err)))?;
Python::with_gil(|py| json_to_py(py, value))
})

View file

@ -10,6 +10,7 @@ from io import IOBase
from typing import Any, Coroutine, Union, cast
import httpx
from typing_extensions import Never
import litellm
from litellm._logging import verbose_logger
@ -19,12 +20,112 @@ from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.ocr.rust_bridge import (
RustAocr,
RustOcr,
RustOcrError,
load_rust_aocr,
load_rust_ocr,
)
from litellm.utils import client, filter_out_litellm_params
class _OCRInputError(ValueError):
pass
def _ocr_error_response(status_code: int) -> httpx.Response:
return httpx.Response(
status_code=status_code,
request=httpx.Request(method="POST", url="https://litellm.ai"),
)
def _raise_rust_ocr_exception(
err: RustOcrError, model: str, custom_llm_provider: str | None
) -> Never:
provider = custom_llm_provider or "mistral"
status_code = err.status_code
message = err.message
match status_code:
case None:
raise litellm.APIConnectionError(
message=message, llm_provider=provider, model=model
)
case 400:
raise litellm.BadRequestError(
message=message, model=model, llm_provider=provider
)
case 401:
raise litellm.AuthenticationError(
message=message, llm_provider=provider, model=model
)
case 403:
raise litellm.PermissionDeniedError(
message=message,
llm_provider=provider,
model=model,
response=_ocr_error_response(403),
)
case 404:
raise litellm.NotFoundError(
message=message, model=model, llm_provider=provider
)
case 408:
raise litellm.Timeout(message=message, model=model, llm_provider=provider)
case 422:
raise litellm.UnprocessableEntityError(
message=message,
model=model,
llm_provider=provider,
response=_ocr_error_response(422),
)
case 429:
raise litellm.RateLimitError(
message=message, llm_provider=provider, model=model
)
case 500:
raise litellm.InternalServerError(
message=message, llm_provider=provider, model=model
)
case 502:
raise litellm.BadGatewayError(
message=message, llm_provider=provider, model=model
)
case 503:
raise litellm.ServiceUnavailableError(
message=message, llm_provider=provider, model=model
)
case _:
raise litellm.APIError(
status_code=status_code,
message=message,
llm_provider=provider,
model=model,
)
def _raise_ocr_exception(
e: Exception,
model: str,
custom_llm_provider: str | None,
completion_kwargs: dict[str, object],
kwargs: dict[str, object],
) -> Never:
if isinstance(e, RustOcrError):
_raise_rust_ocr_exception(e, model, custom_llm_provider)
if isinstance(e, _OCRInputError):
raise litellm.BadRequestError(
message="Invalid OCR request",
model=model,
llm_provider=custom_llm_provider or "mistral",
) from e
raise litellm.exception_type(
model=model,
custom_llm_provider=custom_llm_provider,
original_exception=e,
completion_kwargs=completion_kwargs,
extra_kwargs=kwargs,
)
def _timeout_to_seconds(
timeout: Union[float, httpx.Timeout] | None,
) -> float | None:
@ -65,7 +166,7 @@ def _resolve_ocr_call_context(
litellm_call_id = cast(str | None, kwargs.get("litellm_call_id", None))
if not isinstance(document, dict):
raise ValueError(
raise _OCRInputError(
f"document must be a dict with 'type' and URL/file field, got {type(document)}"
)
@ -76,7 +177,7 @@ def _resolve_ocr_call_context(
doc_type = document.get("type")
if doc_type not in ["document_url", "image_url"]:
raise ValueError(
raise _OCRInputError(
f"Invalid document type: {doc_type}. "
"Must be 'document_url', 'image_url', or 'file'"
)
@ -357,12 +458,12 @@ async def aocr(
litellm_logging_obj=litellm_logging_obj,
)
except Exception as e:
raise litellm.exception_type(
_raise_ocr_exception(
e,
model=model,
custom_llm_provider=custom_llm_provider,
original_exception=e,
completion_kwargs=completion_kwargs,
extra_kwargs=kwargs,
kwargs=kwargs,
)
@ -420,7 +521,7 @@ def convert_file_document_to_url_document(document: dict[str, Any]) -> dict[str,
"""
file_input = document.get("file")
if file_input is None:
raise ValueError(
raise _OCRInputError(
"document with type='file' must include a 'file' field containing "
"a pathlib.Path, file-like object, or bytes"
)
@ -436,7 +537,7 @@ def convert_file_document_to_url_document(document: dict[str, Any]) -> dict[str,
# Opening it as a path is an arbitrary local file read on the proxy
# host, which is then base64-encoded and forwarded to the OCR
# provider — an exfiltration primitive.
raise ValueError(
raise _OCRInputError(
"OCR file input does not accept bare str values. Pass bytes, "
"a pathlib.Path, or a file-like object. To OCR a local file "
"from a path, call open(path, 'rb') yourself."
@ -446,7 +547,7 @@ def convert_file_document_to_url_document(document: dict[str, Any]) -> dict[str,
# Python-level type that HTTP form values can't fabricate.
file_path = str(file_input)
if not os.path.isfile(file_path):
raise FileNotFoundError(f"File not found: {file_path}")
raise _OCRInputError("OCR file input path does not exist")
mime_type = get_mime_type(file_path)
file_name = os.path.basename(file_path)
with open(file_path, "rb") as f:
@ -462,19 +563,19 @@ def convert_file_document_to_url_document(document: dict[str, Any]) -> dict[str,
if isinstance(file_bytes, str):
file_bytes = file_bytes.encode("utf-8")
else:
raise ValueError(
raise _OCRInputError(
f"Unsupported file input type: {type(file_input)}. "
"Expected pathlib.Path, bytes, or a file-like object."
)
if not file_bytes:
raise ValueError("File is empty or could not be read")
raise _OCRInputError("File is empty or could not be read")
if "mime_type" in document:
mime_type = document["mime_type"]
if not _MIME_PATTERN.match(mime_type):
raise ValueError(f"Invalid MIME type: {mime_type}")
raise _OCRInputError(f"Invalid MIME type: {mime_type}")
base64_data = base64.b64encode(file_bytes).decode("utf-8")
data_uri = f"data:{mime_type};base64,{base64_data}"
@ -621,10 +722,10 @@ def ocr(
litellm_logging_obj=litellm_logging_obj,
)
except Exception as e:
raise litellm.exception_type(
_raise_ocr_exception(
e,
model=model,
custom_llm_provider=custom_llm_provider,
original_exception=e,
completion_kwargs=completion_kwargs,
extra_kwargs=kwargs,
kwargs=kwargs,
)

View file

@ -13,6 +13,13 @@ from __future__ import annotations
from typing import Awaitable, Final, Protocol, cast
class RustOcrError(Exception):
def __init__(self, message: str, status_code: int | None = None) -> None:
super().__init__(message)
self.message = message
self.status_code = status_code
class RustOcr(Protocol):
"""Signature of the compiled Rust OCR entrypoint."""

View file

@ -200,12 +200,14 @@ class TestConvertFileDocumentToUrlDocument:
convert_file_document_to_url_document({"type": "file"})
def test_should_raise_error_for_nonexistent_pathlib_path(self):
"""Non-existent pathlib.Path should raise FileNotFoundError."""
with pytest.raises(FileNotFoundError, match="File not found"):
"""Non-existent pathlib.Path should raise a path-free input error."""
with pytest.raises(ValueError, match="does not exist") as exc_info:
convert_file_document_to_url_document(
{"type": "file", "file": Path("/nonexistent/path/to/file.pdf")}
)
assert "/nonexistent/path/to/file.pdf" not in str(exc_info.value)
def test_should_raise_error_for_empty_file(self):
"""Empty file should raise ValueError."""
with tempfile.NamedTemporaryFile(suffix=".pdf", delete=False) as f:

View file

@ -3,12 +3,15 @@
import importlib
import builtins
import types
from dataclasses import dataclass
import httpx
import pydantic
import pytest
import litellm
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.ocr.rust_bridge import RustOcrError
# `litellm/__init__.py` does `from .ocr.main import *`, which binds the `ocr`
# function onto `litellm.ocr` and shadows the submodule, so import the modules
@ -36,6 +39,40 @@ class CapturedException(Exception):
pass
@dataclass(slots=True)
class ExceptionTypeCall:
model: str
custom_llm_provider: str | None
original_exception: Exception
completion_kwargs: dict[str, object]
extra_kwargs: dict[str, object]
class ExceptionTypeSpy:
def __init__(self) -> None:
self.calls: list[ExceptionTypeCall] = []
def __call__(
self,
*,
model: str,
custom_llm_provider: str | None,
original_exception: Exception,
completion_kwargs: dict[str, object],
extra_kwargs: dict[str, object],
) -> CapturedException:
self.calls.append(
ExceptionTypeCall(
model=model,
custom_llm_provider=custom_llm_provider,
original_exception=original_exception,
completion_kwargs=completion_kwargs,
extra_kwargs=extra_kwargs,
)
)
return CapturedException("wrapped")
class RecordingBridge:
"""A fake ``RustOcr`` callable that records the args it was handed."""
@ -130,6 +167,44 @@ class RaisingAsyncBridge:
raise RuntimeError("bridge failed")
class RustErrorBridge:
def __init__(self, message: str, status_code: int | None) -> None:
self.message = message
self.status_code = status_code
def __call__(
self,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
) -> dict[str, object]:
raise RustOcrError(self.message, self.status_code)
class RustErrorAsyncBridge:
def __init__(self, message: str, status_code: int | None) -> None:
self.message = message
self.status_code = status_code
async def __call__(
self,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
) -> dict[str, object]:
raise RustOcrError(self.message, self.status_code)
class RecordingLogging:
"""A spy standing in for ``LiteLLMLoggingObj`` to capture ``pre_call``."""
@ -317,6 +392,7 @@ def test_run_rust_ocr_forwards_args_and_wraps_response():
"timeout_seconds": 12.5,
}
def test_run_rust_ocr_runs_pre_call_logging():
"""The Rust shortcut must run pre_call so callbacks and spend tracking fire."""
logging_obj = RecordingLogging()
@ -396,21 +472,16 @@ def test_ocr_routes_azure_ai_to_rust_by_default(fake_bridge):
def test_ocr_exception_type_uses_resolved_provider_context(
monkeypatch: pytest.MonkeyPatch,
):
captured: dict[str, object] = {}
def fake_exception_type(**kwargs: object) -> CapturedException:
captured.update(kwargs)
return CapturedException("wrapped")
monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type)
) -> None:
spy = ExceptionTypeSpy()
monkeypatch.setattr(ocr_main.litellm, "exception_type", spy)
rust_bridge._set_rust_ocr_bridge(ocr=RaisingBridge())
with pytest.raises(CapturedException):
litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
assert captured["model"] == "mistral-ocr-latest"
assert captured["custom_llm_provider"] == "mistral"
assert spy.calls[0].model == "mistral-ocr-latest"
assert spy.calls[0].custom_llm_provider == "mistral"
@pytest.mark.asyncio
@ -438,21 +509,16 @@ async def test_aocr_routes_to_async_rust_by_default(fake_async_bridge):
@pytest.mark.asyncio
async def test_aocr_exception_type_uses_resolved_provider_context(
monkeypatch: pytest.MonkeyPatch,
):
captured: dict[str, object] = {}
def fake_exception_type(**kwargs: object) -> CapturedException:
captured.update(kwargs)
return CapturedException("wrapped")
monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type)
) -> None:
spy = ExceptionTypeSpy()
monkeypatch.setattr(ocr_main.litellm, "exception_type", spy)
rust_bridge._set_rust_ocr_bridge(aocr=RaisingAsyncBridge())
with pytest.raises(CapturedException):
await litellm.aocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
assert captured["model"] == "mistral-ocr-latest"
assert captured["custom_llm_provider"] == "mistral"
assert spy.calls[0].model == "mistral-ocr-latest"
assert spy.calls[0].custom_llm_provider == "mistral"
def test_ocr_forwards_timeout_to_rust(fake_bridge):
@ -480,3 +546,258 @@ def test_ocr_requires_rust_bridge_when_unavailable(monkeypatch):
litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
assert "Rust OCR bridge is required" in str(exc_info.value)
RUST_OCR_ERROR_CASES = [
pytest.param(400, litellm.BadRequestError, 400, id="400_bad_request"),
pytest.param(401, litellm.AuthenticationError, 401, id="401_authentication"),
pytest.param(403, litellm.PermissionDeniedError, 403, id="403_permission_denied"),
pytest.param(404, litellm.NotFoundError, 404, id="404_not_found"),
pytest.param(408, litellm.Timeout, 408, id="408_timeout"),
pytest.param(
422, litellm.UnprocessableEntityError, 422, id="422_unprocessable_entity"
),
pytest.param(429, litellm.RateLimitError, 429, id="429_rate_limit"),
pytest.param(500, litellm.InternalServerError, 500, id="500_internal"),
pytest.param(502, litellm.BadGatewayError, 502, id="502_bad_gateway"),
pytest.param(503, litellm.ServiceUnavailableError, 503, id="503_unavailable"),
pytest.param(None, litellm.APIConnectionError, 500, id="none_connection"),
]
@pytest.mark.parametrize(
("status_code", "expected_exception", "expected_status"), RUST_OCR_ERROR_CASES
)
def test_rust_ocr_error_maps_to_public_exception(
status_code: int | None,
expected_exception: type[Exception],
expected_status: int,
) -> None:
with pytest.raises(expected_exception) as exc_info:
ocr_main._raise_rust_ocr_exception(
RustOcrError("upstream boom", status_code),
model="mistral-ocr-latest",
custom_llm_provider="mistral",
)
exc = exc_info.value
assert exc.status_code == expected_status
assert exc.llm_provider == "mistral"
assert exc.model == "mistral-ocr-latest"
assert "upstream boom" in str(exc)
@pytest.mark.parametrize(
("status_code", "expected_exception", "expected_status"), RUST_OCR_ERROR_CASES
)
def test_ocr_raises_typed_exception_from_rust_error(
status_code: int | None,
expected_exception: type[Exception],
expected_status: int,
) -> None:
rust_bridge._set_rust_ocr_bridge(ocr=RustErrorBridge("upstream boom", status_code))
with pytest.raises(expected_exception) as exc_info:
litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
assert exc_info.value.status_code == expected_status
assert exc_info.value.llm_provider == "mistral"
@pytest.mark.parametrize(
("status_code", "expected_exception", "expected_status"), RUST_OCR_ERROR_CASES
)
@pytest.mark.asyncio
async def test_aocr_raises_typed_exception_from_rust_error(
status_code: int | None,
expected_exception: type[Exception],
expected_status: int,
) -> None:
rust_bridge._set_rust_ocr_bridge(
aocr=RustErrorAsyncBridge("upstream boom", status_code)
)
with pytest.raises(expected_exception) as exc_info:
await litellm.aocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
assert exc_info.value.status_code == expected_status
assert exc_info.value.llm_provider == "mistral"
UNKNOWN_STATUS_CASES = [
pytest.param(409, id="409_conflict"),
pytest.param(451, id="451_legal_reasons"),
pytest.param(504, id="504_gateway_timeout"),
pytest.param(418, id="418_teapot"),
]
@pytest.mark.parametrize("status_code", UNKNOWN_STATUS_CASES)
def test_rust_ocr_error_unknown_status_preserves_exact_status(
status_code: int,
) -> None:
with pytest.raises(litellm.APIError) as exc_info:
ocr_main._raise_rust_ocr_exception(
RustOcrError("upstream boom", status_code),
model="mistral-ocr-latest",
custom_llm_provider="mistral",
)
exc = exc_info.value
assert type(exc) is litellm.APIError
assert exc.status_code == status_code
assert exc.llm_provider == "mistral"
assert exc.model == "mistral-ocr-latest"
def test_rust_ocr_error_message_is_preserved_bounded() -> None:
bounded = "x" * 256 + "... (truncated)"
with pytest.raises(litellm.InternalServerError) as exc_info:
ocr_main._raise_rust_ocr_exception(
RustOcrError(bounded, 500),
model="mistral-ocr-latest",
custom_llm_provider="mistral",
)
assert bounded in str(exc_info.value)
INVALID_OCR_INPUTS = [
pytest.param({"type": "bogus", "document_url": "https://x/y.pdf"}, id="bad_type"),
pytest.param("not-a-dict", id="non_dict_document"),
]
@pytest.mark.parametrize("document", INVALID_OCR_INPUTS)
def test_ocr_invalid_input_raises_bad_request(document: object) -> None:
with pytest.raises(litellm.BadRequestError) as exc_info:
litellm.ocr(model=MODEL, document=document, api_key="sk-test")
assert exc_info.value.status_code == 400
@pytest.mark.parametrize("document", INVALID_OCR_INPUTS)
@pytest.mark.asyncio
async def test_aocr_invalid_input_raises_bad_request(document: object) -> None:
with pytest.raises(litellm.BadRequestError) as exc_info:
await litellm.aocr(model=MODEL, document=document, api_key="sk-test")
assert exc_info.value.status_code == 400
def test_raise_ocr_exception_maps_input_error_to_bad_request() -> None:
with pytest.raises(litellm.BadRequestError) as exc_info:
ocr_main._raise_ocr_exception(
ocr_main._OCRInputError("Invalid document type: bogus"),
model="mistral-ocr-latest",
custom_llm_provider="mistral",
completion_kwargs={},
kwargs={},
)
exc = exc_info.value
assert exc.status_code == 400
assert "Invalid OCR request" in str(exc)
assert "bogus" not in str(exc)
OCR_INPUT_CANARIES = [
"https://signed.example/doc.pdf",
"token=SECRET123",
"QUJDREVGYmFzZTY0",
"page=42",
"application/x-canary-mime",
"/var/secrets/service_account.json",
"sk-canary-secret",
"canary_document_type",
]
def test_raise_ocr_exception_input_error_publishes_generic_message() -> None:
canary = " ".join(OCR_INPUT_CANARIES)
with pytest.raises(litellm.BadRequestError) as exc_info:
ocr_main._raise_ocr_exception(
ocr_main._OCRInputError(canary),
model="mistral-ocr-latest",
custom_llm_provider="mistral",
completion_kwargs={},
kwargs={},
)
message = str(exc_info.value)
assert "Invalid OCR request" in message
for marker in OCR_INPUT_CANARIES:
assert marker not in message
def test_ocr_input_error_public_message_drops_canaries() -> None:
document = {
"type": " ".join(OCR_INPUT_CANARIES),
"document_url": "https://signed.example/doc.pdf?token=SECRET123&page=42",
}
with pytest.raises(litellm.BadRequestError) as exc_info:
litellm.ocr(model=MODEL, document=document, api_key="sk-test")
message = str(exc_info.value)
assert "Invalid OCR request" in message
for marker in OCR_INPUT_CANARIES:
assert marker not in message
@pytest.mark.asyncio
async def test_aocr_input_error_public_message_drops_canaries() -> None:
document = {
"type": " ".join(OCR_INPUT_CANARIES),
"document_url": "https://signed.example/doc.pdf?token=SECRET123&page=42",
}
with pytest.raises(litellm.BadRequestError) as exc_info:
await litellm.aocr(model=MODEL, document=document, api_key="sk-test")
message = str(exc_info.value)
assert "Invalid OCR request" in message
for marker in OCR_INPUT_CANARIES:
assert marker not in message
def test_raise_ocr_exception_keeps_plain_value_error_off_bad_request(
monkeypatch: pytest.MonkeyPatch,
) -> None:
spy = ExceptionTypeSpy()
monkeypatch.setattr(ocr_main.litellm, "exception_type", spy)
internal = ValueError("internal invariant broke")
with pytest.raises(CapturedException):
ocr_main._raise_ocr_exception(
internal,
model="mistral-ocr-latest",
custom_llm_provider="mistral",
completion_kwargs={},
kwargs={},
)
assert spy.calls[0].original_exception is internal
def test_raise_ocr_exception_keeps_validation_error_off_bad_request(
monkeypatch: pytest.MonkeyPatch,
) -> None:
class _Model(pydantic.BaseModel):
value: int
adapter = pydantic.TypeAdapter(_Model)
with pytest.raises(pydantic.ValidationError) as validation_info:
adapter.validate_python({"value": "not-an-int"})
spy = ExceptionTypeSpy()
monkeypatch.setattr(ocr_main.litellm, "exception_type", spy)
with pytest.raises(CapturedException):
ocr_main._raise_ocr_exception(
validation_info.value,
model="mistral-ocr-latest",
custom_llm_provider="mistral",
completion_kwargs={},
kwargs={},
)
assert spy.calls[0].original_exception is validation_info.value