mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
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:
parent
c264758eb8
commit
6661462d5a
12 changed files with 911 additions and 99 deletions
|
|
@ -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");
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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));
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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));
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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!({
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
})
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue