fix(ocr): bound document fetches and polling

Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-07-18 03:39:24 +00:00
parent 0cecdd3253
commit c79f4891fb
5 changed files with 177 additions and 90 deletions

View file

@ -36,3 +36,13 @@ pub(crate) const VERTEXAI_CREDENTIALS_ENV: &str = "VERTEXAI_CREDENTIALS";
pub(crate) const VERTEX_CREDENTIALS_CACHE_CAPACITY: usize = 64;
pub(crate) const ENV_REFERENCE_PREFIX: &str = "os.environ/";
pub(crate) const DEFAULT_OCR_REQUEST_TIMEOUT_SECS: u64 = 600;
pub(crate) const OCR_ERROR_BODY_MAX_CHARS: usize = 256;
pub(crate) const AZURE_DOCUMENT_INTELLIGENCE_POLL_TIMEOUT_SECS: u64 = 120;
pub(crate) const DEFAULT_MAX_IMAGE_URL_DOWNLOAD_SIZE_MB: f64 = 50.0;
pub(crate) const MAX_SAFE_FETCH_REDIRECTS: usize = 10;

View file

@ -1,13 +1,24 @@
use std::sync::OnceLock;
use std::time::Duration;
const OCR_TIMEOUT_SECS: u64 = 600;
use crate::constants::DEFAULT_OCR_REQUEST_TIMEOUT_SECS;
pub(super) fn http_client() -> &'static reqwest::Client {
static CLIENT: OnceLock<reqwest::Client> = OnceLock::new();
CLIENT.get_or_init(|| {
reqwest::Client::builder()
.timeout(Duration::from_secs(OCR_TIMEOUT_SECS))
.timeout(Duration::from_secs(DEFAULT_OCR_REQUEST_TIMEOUT_SECS))
.build()
.expect("failed to build reqwest client")
})
}
pub(super) fn safe_fetch_client() -> &'static reqwest::Client {
static CLIENT: OnceLock<reqwest::Client> = OnceLock::new();
CLIENT.get_or_init(|| {
reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.timeout(Duration::from_secs(DEFAULT_OCR_REQUEST_TIMEOUT_SECS))
.build()
.expect("failed to build reqwest client")
})

View file

@ -21,12 +21,11 @@ use litellm_core::providers::vertex_ai::ocr::transformation::{
VERTEX_AI_DEEPSEEK_OCR_CONFIG, VERTEX_AI_OCR_CONFIG,
};
use super::client::http_client;
const ERROR_BODY_MAX_CHARS: usize = 256;
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;
use super::client::{http_client, safe_fetch_client};
use crate::constants::{
AZURE_DOCUMENT_INTELLIGENCE_POLL_TIMEOUT_SECS, DEFAULT_MAX_IMAGE_URL_DOWNLOAD_SIZE_MB,
DEFAULT_OCR_REQUEST_TIMEOUT_SECS, MAX_SAFE_FETCH_REDIRECTS, OCR_ERROR_BODY_MAX_CHARS,
};
pub(super) fn classify_reqwest_error(err: reqwest::Error) -> CoreError {
if err.is_timeout() {
@ -37,10 +36,10 @@ pub(super) fn classify_reqwest_error(err: reqwest::Error) -> CoreError {
}
pub(super) fn truncate_error_body(body: &str) -> String {
if body.chars().count() <= ERROR_BODY_MAX_CHARS {
if body.chars().count() <= OCR_ERROR_BODY_MAX_CHARS {
return body.to_string();
}
let truncated: String = body.chars().take(ERROR_BODY_MAX_CHARS).collect();
let truncated: String = body.chars().take(OCR_ERROR_BODY_MAX_CHARS).collect();
format!("{truncated}... (truncated)")
}
@ -221,21 +220,31 @@ fn redirect_location(response: &reqwest::Response, url: &Url) -> CoreResult<Url>
.map_err(|err| CoreError::InvalidResponse(format!("invalid OCR document redirect: {err}")))
}
async fn safe_get_document_url(url: &str) -> CoreResult<(Url, reqwest::Response)> {
let client = reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(|err| CoreError::Network(err.to_string()))?;
async fn safe_fetch_request(url: Url, timeout: Duration) -> CoreResult<reqwest::Response> {
safe_fetch_client()
.get(url)
.timeout(timeout)
.send()
.await
.map_err(classify_reqwest_error)
}
async fn safe_get_document_url(
url: &str,
timeout: Option<Duration>,
) -> CoreResult<(Url, reqwest::Response)> {
let mut current_url = Url::parse(url)
.map_err(|err| CoreError::InvalidRequest(format!("invalid OCR document URL: {err}")))?;
let timeout = timeout.unwrap_or(Duration::from_secs(DEFAULT_OCR_REQUEST_TIMEOUT_SECS));
let deadline = Instant::now() + timeout;
for _ in 0..MAX_SAFE_FETCH_REDIRECTS {
validate_safe_fetch_url(&current_url).await?;
let response = client
.get(current_url.clone())
.send()
.await
.map_err(classify_reqwest_error)?;
let remaining = deadline.saturating_duration_since(Instant::now());
if remaining.is_zero() {
return Err(CoreError::Timeout);
}
let response = safe_fetch_request(current_url.clone(), remaining).await?;
if !response.status().is_redirection() {
return Ok((current_url, response));
}
@ -266,8 +275,8 @@ fn enforce_download_size(content_length: u64, max_bytes: u64, url: &Url) -> Core
async fn read_response_with_limit(
mut response: reqwest::Response,
url: &Url,
max_bytes: u64,
) -> CoreResult<Vec<u8>> {
let max_bytes = max_document_download_bytes();
if let Some(content_length) = response.content_length() {
enforce_download_size(content_length, max_bytes, url)?;
} else {
@ -284,7 +293,10 @@ async fn read_response_with_limit(
Ok(bytes)
}
pub(super) async fn convert_document_url_to_data_uri(document: Value) -> CoreResult<Value> {
pub(super) async fn convert_document_url_to_data_uri(
document: Value,
timeout: Option<Duration>,
) -> CoreResult<Value> {
let Some((field, url)) = document_url_field(&document)? else {
return Ok(document);
};
@ -292,10 +304,12 @@ pub(super) async fn convert_document_url_to_data_uri(document: Value) -> CoreRes
return Ok(document);
}
let (final_url, response) = safe_get_document_url(url).await?;
let (final_url, response) = safe_get_document_url(url, timeout).await?;
let status = response.status();
if !status.is_success() {
let body = response.text().await.unwrap_or_default();
let body = read_response_with_limit(response, &final_url, max_document_download_bytes())
.await
.map(|bytes| String::from_utf8_lossy(&bytes).into_owned())?;
return Err(CoreError::Http {
status: status.as_u16(),
body: truncate_error_body(&body),
@ -310,7 +324,8 @@ pub(super) async fn convert_document_url_to_data_uri(document: Value) -> CoreRes
.filter(|value| !value.is_empty())
.unwrap_or("application/octet-stream")
.to_string();
let bytes = read_response_with_limit(response, &final_url).await?;
let bytes =
read_response_with_limit(response, &final_url, max_document_download_bytes()).await?;
let data_uri = format!(
"data:{content_type};base64,{}",
BASE64_STANDARD.encode(bytes)
@ -510,12 +525,13 @@ pub(super) async fn poll_document_intelligence(
));
}
let start = Instant::now();
let timeout = timeout.unwrap_or(Duration::from_secs(
AZURE_DOCUMENT_INTELLIGENCE_POLL_TIMEOUT_SECS,
));
let deadline = Instant::now() + timeout;
loop {
if start.elapsed() > timeout {
let remaining = deadline.saturating_duration_since(Instant::now());
if remaining.is_zero() {
return Err(CoreError::Timeout);
}
@ -525,6 +541,7 @@ pub(super) async fn poll_document_intelligence(
request_builder = request_builder.header(key, value);
}
}
request_builder = request_builder.timeout(remaining);
let response = request_builder
.send()
.await
@ -544,7 +561,11 @@ pub(super) async fn poll_document_intelligence(
if operation_status(&response_json)? == "succeeded" {
return Ok(response_json);
}
tokio::time::sleep(Duration::from_secs(retry_after)).await;
let remaining = deadline.saturating_duration_since(Instant::now());
if remaining.is_zero() {
return Err(CoreError::Timeout);
}
tokio::time::sleep(Duration::from_secs(retry_after).min(remaining)).await;
}
}
@ -552,6 +573,8 @@ pub(super) async fn poll_document_intelligence(
mod tests {
use super::*;
use serde_json::json;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
#[test]
fn blocks_private_and_metadata_ips() {
@ -576,10 +599,13 @@ mod tests {
#[tokio::test]
async fn convert_document_url_rejects_loopback_fetch() {
let error = convert_document_url_to_data_uri(json!({
"type": "image_url",
"image_url": "http://127.0.0.1/image.png"
}))
let error = convert_document_url_to_data_uri(
json!({
"type": "image_url",
"image_url": "http://127.0.0.1/image.png"
}),
None,
)
.await
.unwrap_err();
@ -597,10 +623,89 @@ mod tests {
"image_url": "data:image/png;base64,abcd"
});
let transformed = convert_document_url_to_data_uri(document.clone())
let transformed = convert_document_url_to_data_uri(document.clone(), None)
.await
.unwrap();
assert_eq!(transformed, document);
}
#[tokio::test]
async fn safe_fetch_request_honors_timeout() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (_socket, _) = listener.accept().await.unwrap();
tokio::time::sleep(Duration::from_secs(1)).await;
});
let error = safe_fetch_request(
Url::parse(&format!("http://{address}/document")).unwrap(),
Duration::from_millis(20),
)
.await
.unwrap_err();
assert_eq!(error, CoreError::Timeout);
server.abort();
}
#[tokio::test]
async fn response_body_limit_applies_to_error_responses() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut request = [0_u8; 1024];
let _ = socket.read(&mut request).await.unwrap();
socket
.write_all(
b"HTTP/1.1 400 Bad Request\r\ncontent-length: 6\r\nconnection: close\r\n\r\nabcdef",
)
.await
.unwrap();
});
let url = Url::parse(&format!("http://{address}/document")).unwrap();
let response = safe_fetch_client().get(url.clone()).send().await.unwrap();
let error = read_response_with_limit(response, &url, 5)
.await
.unwrap_err();
assert!(matches!(error, CoreError::InvalidRequest(_)));
server.await.unwrap();
}
#[tokio::test]
async fn document_intelligence_poll_respects_remaining_deadline() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut request = [0_u8; 1024];
let _ = socket.read(&mut request).await.unwrap();
let body = r#"{"status":"running"}"#;
let response = format!(
"HTTP/1.1 200 OK\r\ncontent-length: {}\r\nretry-after: 60\r\nconnection: close\r\n\r\n{body}",
body.len()
);
socket.write_all(response.as_bytes()).await.unwrap();
});
let operation_url = format!("http://{address}/operations/1");
let original_url = format!("http://{address}/document");
let started = Instant::now();
let error = poll_document_intelligence(
&operation_url,
&original_url,
&[],
Some(Duration::from_millis(50)),
)
.await
.unwrap_err();
assert_eq!(error, CoreError::Timeout);
assert!(started.elapsed() < Duration::from_millis(500));
server.await.unwrap();
}
}

View file

@ -118,7 +118,7 @@ impl OcrLifecycleHooks {
let document = match config.document_preparation() {
OcrDocumentPreparation::None => request.document,
OcrDocumentPreparation::DataUri => {
convert_document_url_to_data_uri(request.document).await?
convert_document_url_to_data_uri(request.document, request.timeout).await?
}
OcrDocumentPreparation::ReductoUpload => {
upload_reducto_document(request.document, &url, &upstream_headers, request.timeout)

View file

@ -38,25 +38,17 @@ def _ocr_error_response(status_code: int) -> httpx.Response:
)
def _raise_rust_ocr_exception(
err: RustOcrError, model: str, custom_llm_provider: str | None
) -> Never:
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
)
raise litellm.APIConnectionError(message=message, llm_provider=provider, model=model)
case 400:
raise litellm.BadRequestError(
message=message, model=model, llm_provider=provider
)
raise litellm.BadRequestError(message=message, model=model, llm_provider=provider)
case 401:
raise litellm.AuthenticationError(
message=message, llm_provider=provider, model=model
)
raise litellm.AuthenticationError(message=message, llm_provider=provider, model=model)
case 403:
raise litellm.PermissionDeniedError(
message=message,
@ -65,9 +57,7 @@ def _raise_rust_ocr_exception(
response=_ocr_error_response(403),
)
case 404:
raise litellm.NotFoundError(
message=message, model=model, llm_provider=provider
)
raise litellm.NotFoundError(message=message, model=model, llm_provider=provider)
case 408:
raise litellm.Timeout(message=message, model=model, llm_provider=provider)
case 422:
@ -78,21 +68,13 @@ def _raise_rust_ocr_exception(
response=_ocr_error_response(422),
)
case 429:
raise litellm.RateLimitError(
message=message, llm_provider=provider, model=model
)
raise litellm.RateLimitError(message=message, llm_provider=provider, model=model)
case 500:
raise litellm.InternalServerError(
message=message, llm_provider=provider, model=model
)
raise litellm.InternalServerError(message=message, llm_provider=provider, model=model)
case 502:
raise litellm.BadGatewayError(
message=message, llm_provider=provider, model=model
)
raise litellm.BadGatewayError(message=message, llm_provider=provider, model=model)
case 503:
raise litellm.ServiceUnavailableError(
message=message, llm_provider=provider, model=model
)
raise litellm.ServiceUnavailableError(message=message, llm_provider=provider, model=model)
case _:
raise litellm.APIError(
status_code=status_code,
@ -166,9 +148,7 @@ def _resolve_ocr_call_context(
litellm_call_id = cast(str | None, kwargs.get("litellm_call_id", None))
if not isinstance(document, dict):
raise _OCRInputError(
f"document must be a dict with 'type' and URL/file field, got {type(document)}"
)
raise _OCRInputError(f"document must be a dict with 'type' and URL/file field, got {type(document)}")
doc_type = document.get("type")
@ -177,10 +157,7 @@ def _resolve_ocr_call_context(
doc_type = document.get("type")
if doc_type not in ["document_url", "image_url"]:
raise _OCRInputError(
f"Invalid document type: {doc_type}. "
"Must be 'document_url', 'image_url', or 'file'"
)
raise _OCRInputError(f"Invalid document type: {doc_type}. Must be 'document_url', 'image_url', or 'file'")
caller_supplied_api_base = api_base is not None
@ -201,10 +178,7 @@ def _resolve_ocr_call_context(
suppress_dynamic_api_base = (
not caller_supplied_api_base
and custom_llm_provider == "azure_ai"
and (
"doc-intelligence" in model.lower()
or "documentintelligence" in model.lower()
)
and ("doc-intelligence" in model.lower() or "documentintelligence" in model.lower())
)
if dynamic_api_base and not suppress_dynamic_api_base:
api_base = dynamic_api_base
@ -213,17 +187,9 @@ def _resolve_ocr_call_context(
forwarded_kwargs = {
**filter_out_litellm_params(kwargs=kwargs),
**{
key: kwargs[key]
for key in _OCR_PUBLIC_PARAMS_RESERVED_BY_LITELLM
if kwargs.get(key) is not None
},
}
optional_params = {
key: value
for key, value in forwarded_kwargs.items()
if key not in _RUST_BRIDGE_INTERNAL_PARAMS
**{key: kwargs[key] for key in _OCR_PUBLIC_PARAMS_RESERVED_BY_LITELLM if kwargs.get(key) is not None},
}
optional_params = {key: value for key, value in forwarded_kwargs.items() if key not in _RUST_BRIDGE_INTERNAL_PARAMS}
verbose_logger.debug(f"OCR optional_params forwarded to Rust: {optional_params}")
@ -455,9 +421,7 @@ async def aocr(
extra_headers=extra_headers,
kwargs=kwargs,
)
completion_kwargs.update(
{"model": model, "custom_llm_provider": custom_llm_provider}
)
completion_kwargs.update({"model": model, "custom_llm_provider": custom_llm_provider})
rust_aocr = load_rust_aocr()
if rust_aocr is None:
@ -584,8 +548,7 @@ def convert_file_document_to_url_document(document: dict[str, Any]) -> dict[str,
file_bytes = file_bytes.encode("utf-8")
else:
raise _OCRInputError(
f"Unsupported file input type: {type(file_input)}. "
"Expected pathlib.Path, bytes, or a file-like object."
f"Unsupported file input type: {type(file_input)}. Expected pathlib.Path, bytes, or a file-like object."
)
if not file_bytes:
@ -721,9 +684,7 @@ def ocr(
extra_headers=extra_headers,
timeout=timeout,
)
completion_kwargs.update(
{"model": model, "custom_llm_provider": custom_llm_provider}
)
completion_kwargs.update({"model": model, "custom_llm_provider": custom_llm_provider})
rust_ocr = load_rust_ocr()
if rust_ocr is None: