mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
move io to core
This commit is contained in:
parent
5a7ea7ef71
commit
75548d30f5
30 changed files with 1319 additions and 716 deletions
1
litellm-rust/Cargo.lock
generated
1
litellm-rust/Cargo.lock
generated
|
|
@ -2006,6 +2006,7 @@ dependencies = [
|
|||
name = "litellm-python-bridge"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"bytes",
|
||||
"criterion",
|
||||
"futures-util",
|
||||
"litellm-auth",
|
||||
|
|
|
|||
|
|
@ -211,10 +211,10 @@ mod tests {
|
|||
lifecycle.accept::<crate::ocr::Error>(Ok(()));
|
||||
}
|
||||
let selected = crate::ocr::Error::InvalidRequest("provider".into());
|
||||
assert_eq!(
|
||||
assert!(matches!(
|
||||
lifecycle.accept(Err(HostFailure::Error(selected.clone()))),
|
||||
Some(selected)
|
||||
);
|
||||
Some(crate::ocr::Error::InvalidRequest(message)) if message == "provider"
|
||||
));
|
||||
lifecycle.accept::<crate::ocr::Error>(Ok(()));
|
||||
for phase in [
|
||||
HostPhase::DeploymentFailure,
|
||||
|
|
@ -222,11 +222,10 @@ mod tests {
|
|||
HostPhase::AsyncFailure,
|
||||
] {
|
||||
assert_eq!(lifecycle.phase(), phase);
|
||||
assert_eq!(
|
||||
assert!(
|
||||
lifecycle.accept(Err(HostFailure::Error(crate::ocr::Error::InvalidRequest(
|
||||
"callback".into()
|
||||
)))),
|
||||
None
|
||||
)))).is_none()
|
||||
);
|
||||
}
|
||||
assert_eq!(lifecycle.phase(), HostPhase::Complete);
|
||||
|
|
@ -236,10 +235,10 @@ mod tests {
|
|||
fn cancellation_skips_terminal_dispatch() {
|
||||
let mut lifecycle = HostLifecycle::new(true);
|
||||
let error = crate::ocr::Error::InvalidRequest("cancelled".into());
|
||||
assert_eq!(
|
||||
assert!(matches!(
|
||||
lifecycle.accept(Err(HostFailure::Cancelled(error.clone()))),
|
||||
Some(error)
|
||||
);
|
||||
Some(crate::ocr::Error::InvalidRequest(message)) if message == "cancelled"
|
||||
));
|
||||
assert_eq!(lifecycle.phase(), HostPhase::Complete);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -99,7 +99,6 @@ impl BaseOcrConfig for AzureAICohereParseConfig {
|
|||
CohereParseConfig.transform_ocr_response(model, raw_response, request_format)
|
||||
}
|
||||
|
||||
|
||||
fn validate_request_body(&self, body: &Value) -> Result<(), crate::ocr::Error> {
|
||||
let document = crate::ocr::prepare::body_document(body)?;
|
||||
validate_document(&document)?;
|
||||
|
|
|
|||
|
|
@ -541,7 +541,6 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig {
|
|||
) -> Result<DocumentIntelligenceRequest, crate::ocr::Error> {
|
||||
build_request(document)
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
impl AzureDocumentIntelligenceOCRConfig {
|
||||
|
|
@ -813,11 +812,11 @@ mod tests {
|
|||
&base,
|
||||
json!({"pages":[2,0,0,1],"features":["keyValuePairs","languages"], "future_option": {"nested":null}, "extra_body":{"provider_option":false}}),
|
||||
);
|
||||
request.document = serde_json::from_value(json!({
|
||||
request.document = serde_json::from_value::<OcrDocument>(json!({
|
||||
"type":"document_url",
|
||||
"document_url":"https://example.com/document.pdf"
|
||||
}))
|
||||
.unwrap();
|
||||
.unwrap().into();
|
||||
|
||||
perform_ocr(request).await.unwrap();
|
||||
server.await.unwrap();
|
||||
|
|
|
|||
|
|
@ -100,7 +100,6 @@ impl BaseOcrConfig for AzureAIOCRConfig {
|
|||
MistralOCRConfig.transform_ocr_response(model, raw_response, request_format)
|
||||
}
|
||||
|
||||
|
||||
fn validate_request_body(&self, body: &Value) -> Result<(), crate::ocr::Error> {
|
||||
validate_inline_document(&crate::ocr::prepare::body_document(body)?)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -349,7 +349,7 @@ mod tests {
|
|||
|
||||
#[tokio::test]
|
||||
async fn composed_body_preserves_native_document_fields_and_untyped_overrides() {
|
||||
let mut request = crate::ocr::test_support::wire_request(
|
||||
let request = crate::ocr::test_support::wire_request(
|
||||
"cohere/parse",
|
||||
"https://example.com",
|
||||
json!({
|
||||
|
|
@ -361,10 +361,10 @@ mod tests {
|
|||
}
|
||||
}),
|
||||
);
|
||||
request.document = serde_json::from_value(json!({
|
||||
let request = request.with_document(serde_json::from_value(json!({
|
||||
"type":"image_url","image_url":"https://example.com/original.png"
|
||||
}))
|
||||
.unwrap();
|
||||
.unwrap());
|
||||
let request = crate::ocr::prepare::prepare_request(request);
|
||||
let http = CohereParseConfig
|
||||
.prepare_request(&request, &crate::ocr::test_support::ocr_client())
|
||||
|
|
@ -461,12 +461,11 @@ mod tests {
|
|||
"pages":[{"markdown":{"images":[{"image_base64":42}]}}]
|
||||
}))
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
assert!(matches!(
|
||||
normalize_response("parse", response).unwrap_err(),
|
||||
crate::ocr::Error::ResponseField {
|
||||
path: "pages[0].markdown.images[0].image_base64".into()
|
||||
}
|
||||
);
|
||||
crate::ocr::Error::ResponseField { path }
|
||||
if path == "pages[0].markdown.images[0].image_base64"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
@ -504,13 +503,10 @@ mod tests {
|
|||
"https://example.com",
|
||||
json!({"output_format":null,"req_format":null}),
|
||||
);
|
||||
let request = crate::ocr::types::LiteLLMOcrRequest {
|
||||
document: serde_json::from_value(
|
||||
let request = request.with_document(serde_json::from_value(
|
||||
json!({"type":"image_url","image_url":"https://example.com/a.png"}),
|
||||
)
|
||||
.unwrap(),
|
||||
..request
|
||||
};
|
||||
.unwrap());
|
||||
assert_eq!(
|
||||
request.response_format().unwrap(),
|
||||
crate::ocr::types::OcrResponseFormat::Litellm
|
||||
|
|
@ -677,10 +673,10 @@ mod tests {
|
|||
json!({"type":"image_url","image_url":""}),
|
||||
json!({"type":"image_url","image_url":"data:application/pdf;base64,YQ=="}),
|
||||
] {
|
||||
assert_eq!(
|
||||
assert!(matches!(
|
||||
validate_document(&serde_json::from_value(value).unwrap()),
|
||||
Err(crate::ocr::Error::CohereImageOnly)
|
||||
);
|
||||
));
|
||||
}
|
||||
assert!(serde_json::from_value::<CohereOptions>(json!({"output_format":"html"})).is_err());
|
||||
for format in ["markdown", "blocks"] {
|
||||
|
|
|
|||
|
|
@ -206,12 +206,10 @@ mod tests {
|
|||
#[test]
|
||||
fn explicit_null_model_does_not_use_the_missing_model_default() {
|
||||
let response = serde_json::from_value(json!({"model":null})).unwrap();
|
||||
assert_eq!(
|
||||
assert!(matches!(
|
||||
normalize_response("fallback", response).unwrap_err(),
|
||||
crate::ocr::Error::ResponseField {
|
||||
path: "model".into()
|
||||
}
|
||||
);
|
||||
crate::ocr::Error::ResponseField { path } if path == "model"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
@ -241,10 +239,10 @@ mod tests {
|
|||
false,
|
||||
)
|
||||
.unwrap_err();
|
||||
assert_eq!(
|
||||
assert!(matches!(
|
||||
error,
|
||||
crate::ocr::Error::ResponseField { path: path.into() }
|
||||
);
|
||||
crate::ocr::Error::ResponseField { path: actual } if actual == path
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -738,8 +738,7 @@ mod tests {
|
|||
"result":{"chunks":[]}
|
||||
}))])
|
||||
.await;
|
||||
let mut request = wire_request(model, &base, options);
|
||||
request.document = request.document.with_source(source.into());
|
||||
let request = crate::ocr::test_support::with_source(wire_request(model, &base, options), source);
|
||||
|
||||
perform_ocr(request).await.unwrap();
|
||||
server.await.unwrap();
|
||||
|
|
@ -857,8 +856,8 @@ mod tests {
|
|||
#[case("data:application/pdf;base64,INVALID!")]
|
||||
#[tokio::test]
|
||||
async fn rejects_invalid_document_sources_before_network(#[case] source: &str) {
|
||||
let mut request = wire_request("reducto/parse-v3", "http://127.0.0.1:1", json!({}));
|
||||
request.document = request.document.with_source(source.into());
|
||||
let request = crate::ocr::test_support::with_source(
|
||||
wire_request("reducto/parse-v3", "http://127.0.0.1:1", json!({})), source);
|
||||
assert!(perform_ocr(request).await.is_err());
|
||||
}
|
||||
|
||||
|
|
@ -913,8 +912,8 @@ mod tests {
|
|||
async fn facade_omits_native_response_by_default_and_preserves_auth_priority() {
|
||||
let raw = json!({"job_id":"job-1","result":{"chunks":[]}});
|
||||
let (base, seen, server) = mock_server(vec![MockResponse::json(raw)]).await;
|
||||
let mut request = wire_request("reducto/parse-v3", &base, json!({}));
|
||||
request.document = request.document.with_source("reducto://ready.pdf".into());
|
||||
let mut request = crate::ocr::test_support::with_source(
|
||||
wire_request("reducto/parse-v3", &base, json!({})), "reducto://ready.pdf");
|
||||
request.transport.extra_headers = vec![("authorization".into(), "Bearer existing".into())];
|
||||
|
||||
let response = perform_ocr(request).await.unwrap();
|
||||
|
|
|
|||
|
|
@ -192,7 +192,6 @@ impl BaseOcrConfig for VertexAIDeepSeekOCRConfig {
|
|||
.collect(),
|
||||
})
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
pub(crate) fn normalize_response(
|
||||
|
|
@ -612,7 +611,7 @@ mod tests {
|
|||
"usage":{"prompt_tokens":1}
|
||||
}))])
|
||||
.await;
|
||||
let mut request = wire_request(
|
||||
let request = wire_request(
|
||||
"vertex_ai/deepseek-ocr-maas",
|
||||
&base,
|
||||
json!({
|
||||
|
|
@ -623,9 +622,7 @@ mod tests {
|
|||
"extra_body":{"provider_option":"value"}
|
||||
}),
|
||||
);
|
||||
request.document = request
|
||||
.document
|
||||
.with_source("gs://bucket/document.pdf".into());
|
||||
let request = crate::ocr::test_support::with_source(request, "gs://bucket/document.pdf");
|
||||
|
||||
let response = perform_ocr(request).await.unwrap();
|
||||
server.await.unwrap();
|
||||
|
|
|
|||
|
|
@ -109,7 +109,6 @@ impl BaseOcrConfig for VertexAIOCRConfig {
|
|||
MistralOCRConfig.transform_ocr_response(model, raw_response, request_format)
|
||||
}
|
||||
|
||||
|
||||
fn validate_request_body(&self, body: &Value) -> Result<(), crate::ocr::Error> {
|
||||
validate_inline_document(&crate::ocr::prepare::body_document(body)?)
|
||||
}
|
||||
|
|
@ -339,8 +338,8 @@ mod tests {
|
|||
options.clone(),
|
||||
);
|
||||
let vertex = wire_request("vertex_ai/mistral-ocr-maas", "https://vertex.test", options);
|
||||
let direct = crate::ocr::prepare::prepare_request(direct);
|
||||
let vertex = crate::ocr::prepare::prepare_request(vertex);
|
||||
let direct = crate::ocr::prepare::prepare_request(crate::ocr::test_support::resolved_request(direct));
|
||||
let vertex = crate::ocr::prepare::prepare_request(crate::ocr::test_support::resolved_request(vertex));
|
||||
let direct_http = MistralOCRConfig
|
||||
.prepare_request(&direct, &client)
|
||||
.await
|
||||
|
|
|
|||
|
|
@ -2,11 +2,28 @@ use base64::{Engine, engine::general_purpose::STANDARD};
|
|||
use data_url::mime::Mime;
|
||||
use data_url::{DataUrl, DataUrlError, forgiving_base64::DecodeError};
|
||||
use reqwest::Url;
|
||||
use std::io::Read;
|
||||
use std::path::Path;
|
||||
|
||||
use super::types::{OcrConnection, OcrDocument};
|
||||
use crate::constants::{OCR_INLINE_MAX_BYTES, OCR_MAX_FETCH_REDIRECTS};
|
||||
use crate::media::{DownloadPolicy, MediaFetcher};
|
||||
|
||||
pub(crate) fn read_path_document(
|
||||
path: &Path,
|
||||
mime_type: Option<&str>,
|
||||
) -> Result<OcrDocument, super::Error> {
|
||||
let mut bytes = Vec::new();
|
||||
std::fs::File::open(path)
|
||||
.and_then(|file| file.take(OCR_INLINE_MAX_BYTES as u64 + 1).read_to_end(&mut bytes))
|
||||
.map_err(|source| super::Error::FileRead {
|
||||
path: path.to_owned(),
|
||||
source: std::sync::Arc::new(source),
|
||||
})?;
|
||||
let name = path.file_name().map(|name| name.to_string_lossy());
|
||||
encode_file_document(&bytes, name.as_deref(), mime_type)
|
||||
}
|
||||
|
||||
pub fn encode_file_document(
|
||||
bytes: &[u8],
|
||||
file_name: Option<&str>,
|
||||
|
|
@ -182,6 +199,27 @@ fn map_media_error(error: crate::media::Error) -> crate::ocr::Error {
|
|||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn path_preparation_preserves_io_causes_and_enforces_the_inline_limit() {
|
||||
let path = std::env::temp_dir().join(format!("ocr-document-{:032x}.png", rand::random::<u128>()));
|
||||
let error = read_path_document(&path, None).unwrap_err();
|
||||
let super::super::Error::FileRead { path: failed_path, source } = error else {
|
||||
panic!("missing path must produce a typed file error");
|
||||
};
|
||||
assert_eq!(failed_path, path);
|
||||
assert_eq!(source.kind(), std::io::ErrorKind::NotFound);
|
||||
let file = std::fs::File::create(&path).unwrap();
|
||||
file.set_len(OCR_INLINE_MAX_BYTES as u64 + 1).unwrap();
|
||||
let oversized = read_path_document(&path, None);
|
||||
std::fs::write(&path, b"image bytes").unwrap();
|
||||
let image = read_path_document(&path, None).unwrap();
|
||||
let overridden = read_path_document(&path, Some("application/pdf")).unwrap();
|
||||
std::fs::remove_file(&path).unwrap();
|
||||
assert!(matches!(oversized, Err(super::super::Error::InlineDocumentTooLarge)));
|
||||
assert!(matches!(image, OcrDocument::ImageUrl { image_url, .. } if image_url == "data:image/png;base64,aW1hZ2UgYnl0ZXM="));
|
||||
assert!(matches!(overridden, OcrDocument::DocumentUrl { document_url, .. } if document_url == "data:application/pdf;base64,aW1hZ2UgYnl0ZXM="));
|
||||
}
|
||||
|
||||
fn document(source: &str) -> OcrDocument {
|
||||
OcrDocument::DocumentUrl {
|
||||
document_url: source.into(),
|
||||
|
|
@ -248,10 +286,10 @@ mod tests {
|
|||
#[test]
|
||||
fn file_encoding_enforces_decoded_size_limit() {
|
||||
let bytes = vec![b'a'; OCR_INLINE_MAX_BYTES + 1];
|
||||
assert_eq!(
|
||||
assert!(matches!(
|
||||
encode_file_document(&bytes, None, None),
|
||||
Err(crate::ocr::Error::InlineDocumentTooLarge)
|
||||
);
|
||||
));
|
||||
let document = encode_file_document(&bytes[..OCR_INLINE_MAX_BYTES], None, None).unwrap();
|
||||
let inline = InlineDocument::parse(document.source()).unwrap().unwrap();
|
||||
assert_eq!(
|
||||
|
|
@ -282,10 +320,10 @@ mod tests {
|
|||
] {
|
||||
let inline = InlineDocument::parse(source).unwrap().unwrap();
|
||||
assert_eq!(inline.decode(expected.len()).unwrap(), expected);
|
||||
assert_eq!(
|
||||
assert!(matches!(
|
||||
inline.decode(expected.len() - 1),
|
||||
Err(crate::ocr::Error::InlineDocumentTooLarge)
|
||||
);
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
|
||||
#[derive(Clone, Debug, thiserror::Error)]
|
||||
pub enum Error {
|
||||
#[error("upstream OCR error ({status}): {body}")]
|
||||
Provider {
|
||||
|
|
@ -8,6 +8,14 @@ pub enum Error {
|
|||
},
|
||||
#[error("File is empty or could not be read")]
|
||||
EmptyFile,
|
||||
#[error("Failed to read OCR file {}: {source}", path.display())]
|
||||
FileRead {
|
||||
path: std::path::PathBuf,
|
||||
#[source]
|
||||
source: std::sync::Arc<std::io::Error>,
|
||||
},
|
||||
#[error("OCR document preparation task failed: {0}")]
|
||||
DocumentTask(#[source] std::sync::Arc<tokio::task::JoinError>),
|
||||
#[error("Invalid MIME type: {0}")]
|
||||
InvalidMimeType(String),
|
||||
#[error(
|
||||
|
|
|
|||
|
|
@ -2,13 +2,13 @@ use std::sync::Arc;
|
|||
|
||||
use super::OcrClient;
|
||||
use super::hooks::{OcrHooks, OcrLifecycleHooks, OcrPostCallRequest};
|
||||
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, PreparedOcrRequest};
|
||||
use super::types::{ResolvedOcrRequest, LiteLLMOcrResponse, PreparedOcrRequest};
|
||||
use crate::call_lifecycle::{CallLifecycle, CallLifecycleContext};
|
||||
use crate::llms::base_llm::ocr::transformation::OcrResponseContext;
|
||||
|
||||
pub(crate) async fn perform_ocr_request(
|
||||
client: &OcrClient,
|
||||
request: LiteLLMOcrRequest,
|
||||
request: ResolvedOcrRequest,
|
||||
) -> Result<LiteLLMOcrResponse, super::Error> {
|
||||
request.response_format()?;
|
||||
let context = CallLifecycleContext::new(
|
||||
|
|
@ -43,7 +43,7 @@ pub(crate) struct PreparedOcrCall {
|
|||
impl PreparedOcrCall {
|
||||
pub(crate) async fn prepare(
|
||||
client: OcrClient,
|
||||
request: LiteLLMOcrRequest,
|
||||
request: ResolvedOcrRequest,
|
||||
) -> Result<Self, super::Error> {
|
||||
let request = super::prepare::prepare_request(request);
|
||||
let http = request.config.prepare_request(&request, &client).await?;
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ use std::future::Future;
|
|||
use std::pin::Pin;
|
||||
use std::sync::Arc;
|
||||
|
||||
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrDocument};
|
||||
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrDocument, ResolvedOcrRequest};
|
||||
use crate::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming};
|
||||
use serde::Serialize;
|
||||
use serde_json::Value;
|
||||
|
|
@ -75,19 +75,19 @@ pub(crate) struct OcrLifecycleHooks {
|
|||
pub provider_name: String,
|
||||
}
|
||||
|
||||
impl CallLifecycleHooks<LiteLLMOcrRequest, LiteLLMOcrRequest, LiteLLMOcrResponse>
|
||||
impl CallLifecycleHooks<ResolvedOcrRequest, ResolvedOcrRequest, LiteLLMOcrResponse>
|
||||
for OcrLifecycleHooks
|
||||
{
|
||||
type Error = super::Error;
|
||||
type PreCallFuture<'a> = OcrHookFuture<'a, LiteLLMOcrRequest>;
|
||||
type DuringCallFuture<'a> = OcrHookFuture<'a, LiteLLMOcrRequest>;
|
||||
type PreCallFuture<'a> = OcrHookFuture<'a, ResolvedOcrRequest>;
|
||||
type DuringCallFuture<'a> = OcrHookFuture<'a, ResolvedOcrRequest>;
|
||||
type SuccessFuture<'a> = OcrLogFuture<'a>;
|
||||
type FailureFuture<'a> = OcrLogFuture<'a>;
|
||||
|
||||
fn async_pre_call_hook<'a>(
|
||||
&'a self,
|
||||
_context: &'a CallLifecycleContext,
|
||||
request: LiteLLMOcrRequest,
|
||||
request: ResolvedOcrRequest,
|
||||
) -> Self::PreCallFuture<'a> {
|
||||
Box::pin(async move {
|
||||
if !self.hooks.intercepts_requests() {
|
||||
|
|
@ -118,7 +118,7 @@ impl CallLifecycleHooks<LiteLLMOcrRequest, LiteLLMOcrRequest, LiteLLMOcrResponse
|
|||
fn async_during_call_hook<'a>(
|
||||
&'a self,
|
||||
_context: &'a CallLifecycleContext,
|
||||
request: LiteLLMOcrRequest,
|
||||
request: ResolvedOcrRequest,
|
||||
) -> Self::DuringCallFuture<'a> {
|
||||
Box::pin(async move { Ok(request) })
|
||||
}
|
||||
|
|
|
|||
|
|
@ -9,7 +9,8 @@ use super::hooks::{
|
|||
OcrDuringCallRequest, OcrHookFuture, OcrHooks, OcrLogFuture, OcrPostCallRequest,
|
||||
OcrPreCallRequest,
|
||||
};
|
||||
use super::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrClient};
|
||||
use super::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrClient, OcrDocumentInput, OcrFileContent};
|
||||
use super::types::ResolvedOcrRequest;
|
||||
use crate::call_lifecycle::host::{
|
||||
HostCall, HostCallFuture, HostCallStep, HostFailure, HostLifecycle, HostPhase,
|
||||
};
|
||||
|
|
@ -50,6 +51,7 @@ impl OcrAdmission {
|
|||
#[derive(Clone, Debug)]
|
||||
pub enum OcrHostOperation {
|
||||
ProjectRequest,
|
||||
ReadDocument,
|
||||
Lifecycle(HostPhase),
|
||||
ConstructResponse(Arc<LiteLLMOcrResponse>),
|
||||
MapFailure(super::Error),
|
||||
|
|
@ -82,6 +84,7 @@ impl OcrHostOperation {
|
|||
|
||||
pub enum OcrHostResult {
|
||||
Request(Result<(Box<LiteLLMOcrRequest>, bool), super::Error>),
|
||||
Document(Result<OcrFileContent, super::Error>),
|
||||
Lifecycle(Result<(), HostFailure<super::Error>>),
|
||||
AzureAdToken(Result<ResolvedCredential, litellm_auth::Error>),
|
||||
PreCall(Result<OcrPreCallRequest, super::Error>),
|
||||
|
|
@ -179,6 +182,7 @@ impl OcrCall {
|
|||
if self.lifecycle.phase() == HostPhase::Execute {
|
||||
if self.execution.request.is_none()
|
||||
&& self.execution.execution.is_none()
|
||||
&& self.execution.preparation.is_none()
|
||||
&& !self.execution.completed
|
||||
{
|
||||
self.projecting = true;
|
||||
|
|
@ -322,6 +326,8 @@ struct PendingOperation {
|
|||
struct OcrExecution {
|
||||
client: Option<OcrClient>,
|
||||
request: Option<LiteLLMOcrRequest>,
|
||||
preparation: Option<tokio::task::JoinHandle<Result<ResolvedOcrRequest, super::Error>>>,
|
||||
reading: bool,
|
||||
operations_tx: mpsc::UnboundedSender<PendingOperation>,
|
||||
operations_rx: mpsc::UnboundedReceiver<PendingOperation>,
|
||||
pending_result: Option<oneshot::Sender<OcrHostResult>>,
|
||||
|
|
@ -337,6 +343,8 @@ impl OcrExecution {
|
|||
Self {
|
||||
client: Some(client),
|
||||
request: None,
|
||||
preparation: None,
|
||||
reading: false,
|
||||
operations_tx,
|
||||
operations_rx,
|
||||
pending_result: None,
|
||||
|
|
@ -356,11 +364,35 @@ impl OcrExecution {
|
|||
"OCR call cannot be resumed after completion".into(),
|
||||
));
|
||||
}
|
||||
let result = if self.reading {
|
||||
let Some(OcrHostResult::Document(content)) = result else {
|
||||
return Err(super::Error::InvalidRequest("OCR document read result is required".into()));
|
||||
};
|
||||
self.reading = false;
|
||||
let content = content?;
|
||||
let request = self.request.take().expect("pending document read has a request");
|
||||
let OcrDocumentInput::HostReader { mime_type } = &request.document else {
|
||||
unreachable!("only host readers request document reads");
|
||||
};
|
||||
let document = OcrDocumentInput::Bytes {
|
||||
bytes: content.bytes,
|
||||
file_name: content.file_name,
|
||||
mime_type: mime_type.clone(),
|
||||
};
|
||||
self.request = Some(request.with_document(document));
|
||||
None
|
||||
} else {
|
||||
result
|
||||
};
|
||||
match (self.pending_result.take(), result) {
|
||||
(Some(sender), Some(result)) => sender.send(result).map_err(|_| {
|
||||
super::Error::InvalidRequest("OCR host operation was abandoned".into())
|
||||
})?,
|
||||
(None, None) if self.execution.is_none() => self.start(),
|
||||
(None, None) if self.execution.is_none() => {
|
||||
if let Some(operation) = self.prepare().await? {
|
||||
return Ok(OcrCallStep::Host(operation));
|
||||
}
|
||||
}
|
||||
(Some(sender), None) => {
|
||||
self.pending_result = Some(sender);
|
||||
return Err(super::Error::InvalidRequest(
|
||||
|
|
@ -394,12 +426,41 @@ impl OcrExecution {
|
|||
}
|
||||
}
|
||||
|
||||
fn start(&mut self) {
|
||||
async fn prepare(&mut self) -> Result<Option<OcrHostOperation>, super::Error> {
|
||||
if self.preparation.is_none() {
|
||||
let request = self.request.take().expect("admitted OCR call has a request");
|
||||
if matches!(request.document, OcrDocumentInput::HostReader { .. }) {
|
||||
self.request = Some(request);
|
||||
self.reading = true;
|
||||
return Ok(Some(OcrHostOperation::ReadDocument));
|
||||
}
|
||||
if let OcrDocumentInput::Document(document) = &request.document {
|
||||
let document = document.clone();
|
||||
self.start(request.with_document(document));
|
||||
return Ok(None);
|
||||
}
|
||||
self.preparation = Some(tokio::task::spawn_blocking(move || {
|
||||
let document = match &request.document {
|
||||
OcrDocumentInput::Path { path, mime_type } => {
|
||||
super::document::read_path_document(path, mime_type.as_deref())?
|
||||
}
|
||||
OcrDocumentInput::Bytes { bytes, file_name, mime_type } => {
|
||||
super::document::encode_file_document(bytes, file_name.as_deref(), mime_type.as_deref())?
|
||||
}
|
||||
_ => unreachable!("only native file inputs require preparation"),
|
||||
};
|
||||
Ok(request.with_document(document))
|
||||
}));
|
||||
}
|
||||
let result = self.preparation.as_mut().expect("document preparation started").await;
|
||||
self.preparation = None;
|
||||
let request = result.map_err(|error| super::Error::DocumentTask(Arc::new(error)))??;
|
||||
self.start(request);
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
fn start(&mut self, mut request: ResolvedOcrRequest) {
|
||||
let client = self.client.take().expect("admitted OCR call has a client");
|
||||
let mut request = self
|
||||
.request
|
||||
.take()
|
||||
.expect("admitted OCR call has a request");
|
||||
let intercepts_requests = request.hooks.intercepts_requests();
|
||||
if self.azure_ad_token_provider {
|
||||
request.azure_ad_token_provider = Some(TokenProviderHandle::new(Arc::new(
|
||||
|
|
@ -427,6 +488,11 @@ impl OcrExecution {
|
|||
|
||||
async fn stop(&mut self) {
|
||||
self.cancel();
|
||||
if let Some(preparation) = self.preparation.as_mut() {
|
||||
let _ = preparation.await;
|
||||
}
|
||||
self.preparation = None;
|
||||
self.request = None;
|
||||
if let Some(execution) = self.execution.as_mut() {
|
||||
let _ = execution.await;
|
||||
}
|
||||
|
|
@ -439,6 +505,9 @@ impl Drop for OcrExecution {
|
|||
if let Some(execution) = &self.execution {
|
||||
execution.abort();
|
||||
}
|
||||
if let Some(preparation) = &self.preparation {
|
||||
preparation.abort();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -580,6 +649,9 @@ impl OcrHost for NoopOcrHost {
|
|||
OcrHostOperation::ProjectRequest => OcrHostResult::Request(Err(
|
||||
super::Error::InvalidRequest("OCR host has no request projection".into()),
|
||||
)),
|
||||
OcrHostOperation::ReadDocument => OcrHostResult::Document(Err(
|
||||
super::Error::InvalidRequest("OCR host has no document reader".into()),
|
||||
)),
|
||||
OcrHostOperation::Lifecycle(_)
|
||||
| OcrHostOperation::ConstructResponse(_)
|
||||
| OcrHostOperation::MapFailure(_)
|
||||
|
|
@ -615,6 +687,9 @@ impl OcrHost for OcrHookHost {
|
|||
OcrHostOperation::ProjectRequest => OcrHostResult::Request(Err(
|
||||
super::Error::InvalidRequest("OCR hook host has no request projection".into()),
|
||||
)),
|
||||
OcrHostOperation::ReadDocument => OcrHostResult::Document(Err(
|
||||
super::Error::InvalidRequest("OCR hook host has no document reader".into()),
|
||||
)),
|
||||
OcrHostOperation::Success {
|
||||
context,
|
||||
response,
|
||||
|
|
@ -671,6 +746,104 @@ mod tests {
|
|||
OcrDecline, OcrDocument, OcrHost, OcrHostOperation, OcrHostResult,
|
||||
};
|
||||
|
||||
#[tokio::test]
|
||||
async fn sdk_paths_are_read_during_execution_and_hooks_observe_normalized_documents() {
|
||||
struct CaptureDocument(Arc<Mutex<Vec<OcrDocument>>>);
|
||||
impl OcrHooks for CaptureDocument {
|
||||
fn intercepts_requests(&self) -> bool { true }
|
||||
fn pre_call(&self, request: OcrPreCallRequest) -> OcrHookFuture<'_, OcrPreCallRequest> {
|
||||
self.0.lock().unwrap().push(request.document.clone());
|
||||
Box::pin(async { Ok(request) })
|
||||
}
|
||||
}
|
||||
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
|
||||
let path = std::env::temp_dir().join(format!("ocr-sdk-{:032x}.pdf", rand::random::<u128>()));
|
||||
let documents = Arc::new(Mutex::new(Vec::new()));
|
||||
let request = LiteLLMOcrRequest::from_inputs(
|
||||
"mistral/model".into(), path.clone(), None, Default::default(),
|
||||
crate::ocr::OcrConnectionInputs { api_base: Some(base), api_key: Some("test-key".into()), ..Default::default() },
|
||||
).unwrap().with_host_hooks(Arc::new(CaptureDocument(documents.clone())), None);
|
||||
std::fs::write(&path, b"sdk document").unwrap();
|
||||
let result = perform_ocr(request).await;
|
||||
std::fs::remove_file(&path).unwrap();
|
||||
result.unwrap();
|
||||
server.await.unwrap();
|
||||
let expected = json!({"type":"document_url","document_url":"data:application/pdf;base64,c2RrIGRvY3VtZW50"});
|
||||
assert_eq!(serde_json::to_value(&documents.lock().unwrap()[0]).unwrap(), expected);
|
||||
let requests = seen.lock().unwrap();
|
||||
assert_eq!(requests.len(), 1);
|
||||
let body: Value = serde_json::from_str(requests[0].split_once("\r\n\r\n").unwrap().1).unwrap();
|
||||
assert_eq!(body["document"], expected);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn core_requests_a_host_read_once_before_pre_call() {
|
||||
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
|
||||
let request = wire_request("mistral/model", &base, json!({}))
|
||||
.with_document(crate::ocr::OcrDocumentInput::HostReader { mime_type: None })
|
||||
.with_host_hooks(Arc::new(AdmissionSpy { effects: Arc::new(Mutex::new(0)) }), None);
|
||||
let mut request = Some(request);
|
||||
let NativeOutcome::Completed(mut call) = OcrCall::admit(crate::ocr::test_support::ocr_client(), OcrAdmission::all()) else { panic!("admission declined"); };
|
||||
let mut result = None;
|
||||
let mut reads = 0;
|
||||
let mut pre_calls = 0;
|
||||
loop {
|
||||
match call.resume(result.take()).await.unwrap() {
|
||||
OcrCallStep::Host(operation) => {
|
||||
result = Some(match operation {
|
||||
OcrHostOperation::ProjectRequest => {
|
||||
assert_eq!(reads, 0);
|
||||
OcrHostResult::Request(Ok((Box::new(request.take().unwrap()), false)))
|
||||
}
|
||||
OcrHostOperation::ReadDocument => {
|
||||
reads += 1;
|
||||
OcrHostResult::Document(Ok(crate::ocr::OcrFileContent {
|
||||
bytes: bytes::Bytes::from_static(b"image"), file_name: Some("scan.png".into()),
|
||||
}))
|
||||
}
|
||||
OcrHostOperation::PreCall(request) => {
|
||||
assert_eq!(reads, 1);
|
||||
pre_calls += 1;
|
||||
assert!(matches!(&request.document, OcrDocument::ImageUrl { image_url, .. } if image_url == "data:image/png;base64,aW1hZ2U="));
|
||||
OcrHostResult::PreCall(Ok(request))
|
||||
}
|
||||
operation => NoopOcrHost.invoke(operation).await,
|
||||
});
|
||||
}
|
||||
OcrCallStep::Complete(_) => break,
|
||||
}
|
||||
}
|
||||
server.await.unwrap();
|
||||
assert_eq!((reads, pre_calls, seen.lock().unwrap().len()), (1, 1, 1));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn cancellation_acknowledges_blocking_preparation_completion() {
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
let mut execution = super::OcrExecution::new(crate::ocr::test_support::ocr_client());
|
||||
let finished = Arc::new(AtomicBool::new(false));
|
||||
let completed = finished.clone();
|
||||
let (entered_tx, entered_rx) = tokio::sync::oneshot::channel();
|
||||
let (release_tx, release_rx) = std::sync::mpsc::channel();
|
||||
execution.preparation = Some(tokio::task::spawn_blocking(move || {
|
||||
entered_tx.send(()).unwrap();
|
||||
release_rx.recv().unwrap();
|
||||
completed.store(true, Ordering::SeqCst);
|
||||
Err(crate::ocr::Error::EmptyFile)
|
||||
}));
|
||||
entered_rx.await.unwrap();
|
||||
let mut stop = Box::pin(execution.stop());
|
||||
std::future::poll_fn(|cx| {
|
||||
assert!(stop.as_mut().poll(cx).is_pending());
|
||||
std::task::Poll::Ready(())
|
||||
}).await;
|
||||
assert!(!finished.load(Ordering::SeqCst));
|
||||
release_tx.send(()).unwrap();
|
||||
stop.await;
|
||||
assert!(finished.load(Ordering::SeqCst));
|
||||
assert!(execution.preparation.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn request_boundary_selects_mistral_and_rejects_unknown_providers() {
|
||||
let document = OcrDocument::try_from(
|
||||
|
|
@ -1139,7 +1312,7 @@ mod tests {
|
|||
Ok(request)
|
||||
}));
|
||||
}
|
||||
OcrHostOperation::PostCall(_) => panic!("transport should not be reached"),
|
||||
OcrHostOperation::PostCall(_) | OcrHostOperation::ReadDocument => panic!("transport should not be reached"),
|
||||
},
|
||||
Err(error) => break error,
|
||||
Ok(OcrCallStep::Complete(_)) => panic!("failed call completed"),
|
||||
|
|
@ -1299,7 +1472,7 @@ mod tests {
|
|||
OcrHostResult::Lifecycle(Err(HostFailure::Error(selected.clone())))
|
||||
}
|
||||
OcrHostOperation::Failure { error, .. } => {
|
||||
assert_eq!(error, selected);
|
||||
assert!(matches!(error, crate::ocr::Error::InvalidRequest(message) if message == "public metadata failed"));
|
||||
failures.push("sync");
|
||||
OcrHostResult::Lifecycle(Err(HostFailure::Error(
|
||||
crate::ocr::Error::InvalidRequest("failure callback failed".into()),
|
||||
|
|
@ -1325,7 +1498,7 @@ mod tests {
|
|||
}
|
||||
};
|
||||
server.await.unwrap();
|
||||
assert_eq!(error, selected);
|
||||
assert!(matches!(error, crate::ocr::Error::InvalidRequest(message) if message == "public metadata failed"));
|
||||
assert_eq!(failures, ["sync", "async"]);
|
||||
assert_eq!(seen.lock().unwrap().len(), 1);
|
||||
}
|
||||
|
|
@ -1364,7 +1537,7 @@ mod tests {
|
|||
let selected = crate::ocr::Error::InvalidRequest("cancelled".into());
|
||||
assert!(matches!(
|
||||
call.interrupt(HostFailure::Cancelled(selected.clone())).await,
|
||||
Err(error) if error == selected
|
||||
Err(crate::ocr::Error::InvalidRequest(message)) if message == "cancelled"
|
||||
));
|
||||
assert!(
|
||||
call.resume(Some(OcrHostResult::Lifecycle(Ok(()))))
|
||||
|
|
@ -1605,7 +1778,7 @@ mod tests {
|
|||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(matches!(result, Err(error) if error == selected));
|
||||
assert!(matches!(result, Err(crate::ocr::Error::InvalidRequest(message)) if message == "cancelled"));
|
||||
assert!(
|
||||
dropped.load(Ordering::SeqCst),
|
||||
"cancellation returned while provider captures were still alive"
|
||||
|
|
|
|||
|
|
@ -22,8 +22,8 @@ pub use lifecycle::{
|
|||
};
|
||||
pub use provider_config::{get_api_key_env_var, get_health_check_document};
|
||||
pub use types::{
|
||||
LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrCredentialInputs, OcrDocument,
|
||||
OcrPage, OcrPageDimensions, OcrPageImage, OcrTransportConfig, OcrUsageInfo,
|
||||
LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrConnectionInputs, OcrCredentialInputs,
|
||||
OcrDocument, OcrDocumentInput, OcrFileContent, OcrPage, OcrPageDimensions, OcrPageImage, OcrTransportConfig, OcrUsageInfo,
|
||||
};
|
||||
|
||||
#[cfg(test)]
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ use serde_json::Value;
|
|||
|
||||
use super::OcrClient;
|
||||
use super::hooks::OcrDuringCallRequest;
|
||||
use super::types::{LiteLLMOcrRequest, OcrConnection, OcrDocument, PreparedOcrRequest};
|
||||
use super::types::{ResolvedOcrRequest, OcrConnection, OcrDocument, PreparedOcrRequest};
|
||||
|
||||
pub(crate) async fn transform_request_body<B>(
|
||||
client: &OcrClient,
|
||||
|
|
@ -111,7 +111,7 @@ pub(crate) fn credential_env(name: &str) -> Option<String> {
|
|||
std::env::var(name).ok()
|
||||
}
|
||||
|
||||
pub(crate) fn prepare_request(request: LiteLLMOcrRequest) -> PreparedOcrRequest {
|
||||
pub(crate) fn prepare_request(request: ResolvedOcrRequest) -> PreparedOcrRequest {
|
||||
use litellm_auth::{InputSource, Sourced};
|
||||
|
||||
let credentials = request.credentials.clone();
|
||||
|
|
|
|||
|
|
@ -205,10 +205,10 @@ mod tests {
|
|||
#[case("Mistral")]
|
||||
#[case("unknown")]
|
||||
fn invalid_provider_names_are_rejected(#[case] provider: &str) {
|
||||
assert_eq!(
|
||||
assert!(matches!(
|
||||
resolve_provider_config("model", Some(provider)),
|
||||
Err(crate::ocr::Error::InvalidProvider(provider.into()))
|
||||
);
|
||||
Err(crate::ocr::Error::InvalidProvider(value)) if value == provider
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
|||
use tokio::net::TcpListener;
|
||||
|
||||
use crate::ocr::{
|
||||
LiteLLMOcrRequest, LiteLLMOcrResponse, OcrClient, OcrCredentialInputs, OcrDocument,
|
||||
LiteLLMOcrRequest, LiteLLMOcrResponse, OcrClient, OcrConnectionInputs, OcrDocument,
|
||||
};
|
||||
|
||||
pub(crate) fn ocr_client() -> OcrClient {
|
||||
|
|
@ -23,7 +23,7 @@ pub(crate) async fn perform_ocr(
|
|||
}
|
||||
|
||||
pub(crate) fn wire_request(model: &str, base: &str, options: Value) -> LiteLLMOcrRequest {
|
||||
let request = LiteLLMOcrRequest::new(
|
||||
LiteLLMOcrRequest::from_inputs(
|
||||
model.into(),
|
||||
OcrDocument::try_from(
|
||||
json!({"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}),
|
||||
|
|
@ -31,23 +31,28 @@ pub(crate) fn wire_request(model: &str, base: &str, options: Value) -> LiteLLMOc
|
|||
.unwrap(),
|
||||
None,
|
||||
options.as_object().unwrap().clone().into(),
|
||||
OcrConnectionInputs {
|
||||
api_key: Some("test-key".into()),
|
||||
api_base: Some(base.into()),
|
||||
timeout: Some(std::time::Duration::from_secs(2)),
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
let transport = request.transport.clone().with_overrides(
|
||||
Vec::new(),
|
||||
Default::default(),
|
||||
Some(std::time::Duration::from_secs(2)),
|
||||
);
|
||||
request.with_connection_inputs(
|
||||
OcrCredentialInputs::new(
|
||||
Some("test-key".into()),
|
||||
Default::default(),
|
||||
Some(base.into()),
|
||||
Default::default(),
|
||||
),
|
||||
transport,
|
||||
Default::default(),
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
pub(crate) fn resolved_request(request: LiteLLMOcrRequest) -> super::types::ResolvedOcrRequest {
|
||||
let super::OcrDocumentInput::Document(document) = &request.document else {
|
||||
panic!("wire fixture must contain a URL document");
|
||||
};
|
||||
let document = document.clone();
|
||||
request.with_document(document)
|
||||
}
|
||||
|
||||
pub(crate) fn with_source(request: LiteLLMOcrRequest, source: &str) -> LiteLLMOcrRequest {
|
||||
let request = resolved_request(request);
|
||||
let document = request.document.clone().with_source(source.into());
|
||||
request.with_document(document.into())
|
||||
}
|
||||
|
||||
pub(crate) struct MockResponse {
|
||||
|
|
|
|||
|
|
@ -1,8 +1,10 @@
|
|||
use std::collections::BTreeMap;
|
||||
use std::path::PathBuf;
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use bytes::Bytes;
|
||||
use serde_json::{Map, Value};
|
||||
use serde_with::serde_as;
|
||||
|
||||
|
|
@ -66,6 +68,40 @@ impl TryFrom<Value> for OcrDocument {
|
|||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub enum OcrDocumentInput {
|
||||
Document(OcrDocument),
|
||||
Path {
|
||||
path: PathBuf,
|
||||
mime_type: Option<String>,
|
||||
},
|
||||
Bytes {
|
||||
bytes: Bytes,
|
||||
file_name: Option<String>,
|
||||
mime_type: Option<String>,
|
||||
},
|
||||
HostReader {
|
||||
mime_type: Option<String>,
|
||||
},
|
||||
}
|
||||
|
||||
impl From<OcrDocument> for OcrDocumentInput {
|
||||
fn from(document: OcrDocument) -> Self {
|
||||
Self::Document(document)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<PathBuf> for OcrDocumentInput {
|
||||
fn from(path: PathBuf) -> Self {
|
||||
Self::Path { path, mime_type: None }
|
||||
}
|
||||
}
|
||||
|
||||
pub struct OcrFileContent {
|
||||
pub bytes: Bytes,
|
||||
pub file_name: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
pub enum OcrResponseFormat {
|
||||
|
|
@ -143,6 +179,38 @@ fn nonblank(value: Option<String>) -> Option<String> {
|
|||
.filter(|value| !value.is_empty())
|
||||
}
|
||||
|
||||
/// Caller-supplied connection overrides for a [`LiteLLMOcrRequest`], in the
|
||||
/// shape hosts receive them: JSON-ish headers, optional timeout, optional
|
||||
/// credentials, and per-field provenance in `input_sources`.
|
||||
#[derive(Clone, Debug, Default)]
|
||||
pub struct OcrConnectionInputs {
|
||||
pub api_key: Option<String>,
|
||||
pub api_base: Option<String>,
|
||||
pub extra_headers: Map<String, Value>,
|
||||
pub timeout: Option<Duration>,
|
||||
pub input_sources: BTreeMap<String, InputSource>,
|
||||
}
|
||||
|
||||
impl OcrConnectionInputs {
|
||||
fn source(&self, name: &str) -> InputSource {
|
||||
self.input_sources.get(name).copied().unwrap_or_default()
|
||||
}
|
||||
|
||||
fn header_pairs(&self) -> Result<Vec<(String, String)>, super::Error> {
|
||||
self.extra_headers
|
||||
.iter()
|
||||
.map(|(name, value)| {
|
||||
value
|
||||
.as_str()
|
||||
.map(|value| (name.clone(), value.to_string()))
|
||||
.ok_or_else(|| super::Error::RequestField {
|
||||
path: format!("extra_headers.{name}"),
|
||||
})
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct OcrConnection {
|
||||
pub api_key: Option<String>,
|
||||
|
|
@ -199,9 +267,9 @@ pub(crate) struct ResolvedOcrCredentials {
|
|||
pub api_base: Option<Sourced<String>>,
|
||||
}
|
||||
|
||||
pub struct LiteLLMOcrRequest {
|
||||
pub struct LiteLLMOcrRequest<D = OcrDocumentInput> {
|
||||
pub model: String,
|
||||
pub document: OcrDocument,
|
||||
pub document: D,
|
||||
pub credentials: OcrCredentialInputs,
|
||||
pub transport: OcrTransportConfig,
|
||||
pub hooks: Arc<dyn OcrHooks>,
|
||||
|
|
@ -215,7 +283,7 @@ pub struct LiteLLMOcrRequest {
|
|||
impl LiteLLMOcrRequest {
|
||||
pub fn new(
|
||||
model: String,
|
||||
document: OcrDocument,
|
||||
document: impl Into<OcrDocumentInput>,
|
||||
custom_llm_provider: Option<&str>,
|
||||
optional_params: CallArguments,
|
||||
) -> Result<Self, super::Error> {
|
||||
|
|
@ -245,7 +313,7 @@ impl LiteLLMOcrRequest {
|
|||
|
||||
Ok(Self {
|
||||
model,
|
||||
document,
|
||||
document: document.into(),
|
||||
credentials: OcrCredentialInputs::default(),
|
||||
transport,
|
||||
hooks: Arc::new(NoopOcrHooks),
|
||||
|
|
@ -257,6 +325,24 @@ impl LiteLLMOcrRequest {
|
|||
})
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
impl<D> LiteLLMOcrRequest<D> {
|
||||
pub(crate) fn with_document<T>(self, document: T) -> LiteLLMOcrRequest<T> {
|
||||
LiteLLMOcrRequest {
|
||||
model: self.model,
|
||||
document,
|
||||
credentials: self.credentials,
|
||||
transport: self.transport,
|
||||
hooks: self.hooks,
|
||||
litellm_call_id: self.litellm_call_id,
|
||||
optional_params: self.optional_params,
|
||||
input_sources: self.input_sources,
|
||||
azure_ad_token_provider: self.azure_ad_token_provider,
|
||||
config: self.config,
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn response_format(&self) -> Result<OcrResponseFormat, super::Error> {
|
||||
self.optional_params
|
||||
.get("req_format")
|
||||
|
|
@ -297,8 +383,42 @@ impl LiteLLMOcrRequest {
|
|||
..self
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
impl LiteLLMOcrRequest {
|
||||
/// Builds a request from host-shaped inputs in one step: provider
|
||||
/// resolution, optional-param validation, header/timeout overrides and
|
||||
/// sourced credentials. Hosts should prefer this over sequencing
|
||||
/// [`Self::new`], [`OcrTransportConfig::with_overrides`] and
|
||||
/// [`Self::with_connection_inputs`] by hand.
|
||||
pub fn from_inputs(
|
||||
model: String,
|
||||
document: impl Into<OcrDocumentInput>,
|
||||
custom_llm_provider: Option<&str>,
|
||||
optional_params: CallArguments,
|
||||
connection: OcrConnectionInputs,
|
||||
) -> Result<Self, super::Error> {
|
||||
let request = Self::new(model, document, custom_llm_provider, optional_params)?;
|
||||
let transport = request.transport.clone().with_overrides(
|
||||
connection.header_pairs()?,
|
||||
connection.source("extra_headers"),
|
||||
connection.timeout,
|
||||
);
|
||||
let (api_key_source, api_base_source) =
|
||||
(connection.source("api_key"), connection.source("api_base"));
|
||||
let credentials = OcrCredentialInputs::new(
|
||||
connection.api_key,
|
||||
api_key_source,
|
||||
connection.api_base,
|
||||
api_base_source,
|
||||
);
|
||||
Ok(request.with_connection_inputs(credentials, transport, connection.input_sources))
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) type ResolvedOcrRequest = LiteLLMOcrRequest<OcrDocument>;
|
||||
|
||||
pub(crate) struct PreparedOcrRequest {
|
||||
pub model: String,
|
||||
pub document: OcrDocument,
|
||||
|
|
@ -311,7 +431,7 @@ pub(crate) struct PreparedOcrRequest {
|
|||
}
|
||||
|
||||
impl PreparedOcrRequest {
|
||||
pub(crate) fn new(request: LiteLLMOcrRequest, connection: OcrConnection) -> Self {
|
||||
pub(crate) fn new(request: ResolvedOcrRequest, connection: OcrConnection) -> Self {
|
||||
let LiteLLMOcrRequest {
|
||||
model,
|
||||
document,
|
||||
|
|
@ -446,6 +566,84 @@ mod tests {
|
|||
use super::*;
|
||||
use serde_json::json;
|
||||
|
||||
fn document() -> OcrDocument {
|
||||
OcrDocument::try_from(
|
||||
json!({"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}),
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn from_inputs_applies_connection_overrides_with_field_sources() {
|
||||
let request = LiteLLMOcrRequest::from_inputs(
|
||||
"mistral/model".into(),
|
||||
document(),
|
||||
None,
|
||||
Default::default(),
|
||||
OcrConnectionInputs {
|
||||
api_key: Some(" key ".into()),
|
||||
api_base: Some("".into()),
|
||||
extra_headers: json!({"x-a": "1"}).as_object().unwrap().clone(),
|
||||
timeout: Some(Duration::from_secs(7)),
|
||||
input_sources: [
|
||||
("api_key".to_string(), InputSource::Request),
|
||||
("extra_headers".to_string(), InputSource::Request),
|
||||
]
|
||||
.into(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let api_key = request.credentials.api_key.as_ref().unwrap();
|
||||
assert_eq!(api_key.clone().into_value(), "key");
|
||||
assert_eq!(api_key.source(), InputSource::Request);
|
||||
assert!(request.credentials.api_base.is_none());
|
||||
assert_eq!(
|
||||
request.transport.extra_headers,
|
||||
vec![("x-a".to_string(), "1".to_string())]
|
||||
);
|
||||
assert_eq!(request.transport.extra_headers_source, InputSource::Request);
|
||||
assert_eq!(request.transport.timeout, Duration::from_secs(7));
|
||||
assert_eq!(request.input_sources.len(), 2);
|
||||
|
||||
let defaulted = LiteLLMOcrRequest::from_inputs(
|
||||
"mistral/model".into(),
|
||||
document(),
|
||||
None,
|
||||
Default::default(),
|
||||
OcrConnectionInputs::default(),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
defaulted.transport.timeout,
|
||||
OcrTransportConfig::default().timeout
|
||||
);
|
||||
assert_eq!(
|
||||
defaulted.transport.extra_headers_source,
|
||||
InputSource::Deployment
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn from_inputs_rejects_non_string_header_values_by_path() {
|
||||
let Err(error) = LiteLLMOcrRequest::from_inputs(
|
||||
"mistral/model".into(),
|
||||
document(),
|
||||
None,
|
||||
Default::default(),
|
||||
OcrConnectionInputs {
|
||||
extra_headers: json!({"x-a": 1}).as_object().unwrap().clone(),
|
||||
..Default::default()
|
||||
},
|
||||
) else {
|
||||
panic!("non-string header value accepted");
|
||||
};
|
||||
assert!(matches!(
|
||||
error,
|
||||
super::super::Error::RequestField { ref path } if path == "extra_headers.x-a"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn normalized_response_rejects_invalid_shared_fields() {
|
||||
for fields in [
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ extension-module = ["pyo3/extension-module"]
|
|||
panic-test = []
|
||||
|
||||
[dependencies]
|
||||
bytes.workspace = true
|
||||
futures-util.workspace = true
|
||||
litellm-core.workspace = true
|
||||
litellm-auth.workspace = true
|
||||
|
|
|
|||
|
|
@ -7,8 +7,97 @@ use pyo3::types::PyDict;
|
|||
use serde_json::{Map, Value};
|
||||
|
||||
use litellm_auth::InputSource;
|
||||
use litellm_core::call_arguments::ArgumentSpec;
|
||||
use litellm_python_interop::from_py_preserving_errors as from_py;
|
||||
|
||||
use crate::auth::{AZURE_AD_TOKEN_PROVIDER, PythonTokenProvider};
|
||||
use crate::lifecycle::BoundArguments;
|
||||
|
||||
/// Fields every lifecycle route reads from its bound `*args, **kwargs` before
|
||||
/// asking core to build the typed request. Route-specific inputs (for example
|
||||
/// the OCR `document`) are read separately by the route.
|
||||
pub(crate) struct BoundRouteInputs {
|
||||
pub(crate) model: String,
|
||||
pub(crate) custom_llm_provider: Option<String>,
|
||||
pub(crate) api_key: Option<String>,
|
||||
pub(crate) api_base: Option<String>,
|
||||
pub(crate) extra_headers: Map<String, Value>,
|
||||
pub(crate) timeout: Option<Duration>,
|
||||
/// Only the optional params core reports as consumed plus opaque unknown
|
||||
/// values; bound/control fields are never serialized.
|
||||
pub(crate) optional_params: Map<String, Value>,
|
||||
/// Which projected fields were supplied via the proxy request body.
|
||||
pub(crate) input_sources: BTreeMap<String, InputSource>,
|
||||
pub(crate) azure_ad_token_provider: Option<PythonTokenProvider>,
|
||||
/// Consumed params whose values must be redacted from logging.
|
||||
pub(crate) secret_fields: Vec<&'static str>,
|
||||
}
|
||||
|
||||
const CONNECTION_FIELDS: [&str; 3] = ["api_key", "api_base", "extra_headers"];
|
||||
|
||||
impl BoundRouteInputs {
|
||||
/// `consumed` is the route's provider-selected optional param list for the
|
||||
/// already-extracted `model`/`custom_llm_provider`; `bound_fields` are the
|
||||
/// route's positional parameter names that must not be projected.
|
||||
pub(crate) fn extract(
|
||||
py: Python<'_>,
|
||||
arguments: &BoundArguments<'_>,
|
||||
consumed: Vec<ArgumentSpec>,
|
||||
bound_fields: &[&str],
|
||||
) -> PyResult<Self> {
|
||||
let kwargs = arguments.kwargs();
|
||||
let optional_params = project_optional_fields(kwargs, &consumed, bound_fields)?;
|
||||
let input_sources = request_input_sources(
|
||||
kwargs,
|
||||
optional_params
|
||||
.keys()
|
||||
.map(String::as_str)
|
||||
.chain(CONNECTION_FIELDS),
|
||||
)?;
|
||||
Ok(Self {
|
||||
model: arguments.extract("model")?,
|
||||
custom_llm_provider: arguments.optional("custom_llm_provider")?,
|
||||
api_key: arguments.optional("api_key")?,
|
||||
api_base: arguments.optional("api_base")?,
|
||||
extra_headers: arguments
|
||||
.optional::<Bound<'_, PyAny>>("extra_headers")?
|
||||
.map(|value| {
|
||||
from_py(&value).and_then(|value| required_object("extra_headers", value))
|
||||
})
|
||||
.transpose()?
|
||||
.unwrap_or_default(),
|
||||
timeout: bound_timeout(py, arguments)?,
|
||||
optional_params,
|
||||
input_sources,
|
||||
azure_ad_token_provider: kwargs.get_item("azure_ad_token_provider")?.and_then(
|
||||
|provider| PythonTokenProvider::select(provider, AZURE_AD_TOKEN_PROVIDER),
|
||||
),
|
||||
secret_fields: consumed
|
||||
.into_iter()
|
||||
.filter(|spec| spec.secret)
|
||||
.map(|spec| spec.name)
|
||||
.collect(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// Reads `timeout` through `litellm.rust_bridge.timeouts.timeout_to_seconds`
|
||||
/// so Python `httpx.Timeout` objects keep their existing semantics. Unlike the
|
||||
/// value-route [`optional_timeout`], non-finite or negative values are
|
||||
/// rejected rather than silently dropped.
|
||||
fn bound_timeout(py: Python<'_>, arguments: &BoundArguments<'_>) -> PyResult<Option<Duration>> {
|
||||
arguments
|
||||
.optional::<Bound<'_, PyAny>>("timeout")?
|
||||
.map(|value| python_timeout_seconds(py, value.unbind()))
|
||||
.transpose()?
|
||||
.flatten()
|
||||
.map(|seconds| {
|
||||
Duration::try_from_secs_f64(seconds)
|
||||
.map_err(|_| PyValueError::new_err("timeout must be a non-negative finite number"))
|
||||
})
|
||||
.transpose()
|
||||
}
|
||||
|
||||
pub(crate) struct RouteOptions {
|
||||
pub(crate) model: String,
|
||||
pub(crate) api_key: Option<String>,
|
||||
|
|
|
|||
|
|
@ -1,4 +1,3 @@
|
|||
use pyo3::exceptions::PyBaseException;
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::types::PyDict;
|
||||
use serde_json::Value;
|
||||
|
|
@ -148,17 +147,3 @@ pub(super) fn response(py: Python<'_>, response: &LiteLLMOcrResponse) -> PyResul
|
|||
.call1((to_py(py, response)?,))
|
||||
.map(Bound::unbind)
|
||||
}
|
||||
|
||||
pub(super) fn map_failure(
|
||||
py: Python<'_>,
|
||||
error: &Py<PyBaseException>,
|
||||
model: &str,
|
||||
provider: &str,
|
||||
kwargs: &Py<PyDict>,
|
||||
) -> PyResult<Py<PyBaseException>> {
|
||||
Ok(py
|
||||
.import("litellm.rust_bridge.ocr")?
|
||||
.getattr("map_failure")?
|
||||
.call1((error, model, provider, kwargs))?
|
||||
.extract()?)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,96 +1,49 @@
|
|||
use std::io::Read;
|
||||
use std::path::PathBuf;
|
||||
|
||||
use pyo3::exceptions::{PyFileNotFoundError, PyTypeError, PyValueError};
|
||||
use bytes::Bytes;
|
||||
use pyo3::exceptions::PyValueError;
|
||||
use pyo3::gc::{PyTraverseError, PyVisit};
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::pybacked::PyBackedBytes;
|
||||
#[cfg(test)]
|
||||
use pyo3::types::PyDict;
|
||||
use pyo3::types::{PyBytes, PyString};
|
||||
|
||||
use litellm_core::constants::OCR_INLINE_MAX_BYTES;
|
||||
use litellm_core::ocr::{OcrDocument, encode_file_document};
|
||||
use litellm_core::ocr::{OcrDocumentInput, OcrFileContent};
|
||||
|
||||
enum FileBytes {
|
||||
Python(PyBackedBytes),
|
||||
Native(Vec<u8>),
|
||||
pub(super) struct PythonFileReader {
|
||||
reader: Py<PyAny>,
|
||||
name: Option<String>,
|
||||
}
|
||||
|
||||
impl AsRef<[u8]> for FileBytes {
|
||||
fn as_ref(&self) -> &[u8] {
|
||||
match self {
|
||||
Self::Python(bytes) => bytes,
|
||||
Self::Native(bytes) => bytes,
|
||||
}
|
||||
impl PythonFileReader {
|
||||
pub(super) fn read(&self, py: Python<'_>) -> PyResult<OcrFileContent> {
|
||||
let value = py
|
||||
.import("litellm.rust_bridge.ocr")?
|
||||
.getattr("read_document")?
|
||||
.call1((self.reader.bind(py),))?;
|
||||
Ok(OcrFileContent {
|
||||
bytes: extract_bytes(&value)?,
|
||||
file_name: self.name.clone(),
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
visit.call(&self.reader)
|
||||
}
|
||||
}
|
||||
|
||||
fn read_file_input(
|
||||
py: Python<'_>,
|
||||
file: &Bound<'_, PyAny>,
|
||||
) -> PyResult<(FileBytes, Option<String>)> {
|
||||
if file.is_instance_of::<PyString>() {
|
||||
return Err(PyValueError::new_err(
|
||||
"OCR file input does not accept bare str values. Pass bytes, a pathlib.Path, or a file-like object.",
|
||||
));
|
||||
fn extract_bytes(value: &Bound<'_, PyAny>) -> PyResult<Bytes> {
|
||||
if value.is_exact_instance_of::<PyBytes>() {
|
||||
return Ok(Bytes::from_owner(value.extract::<PyBackedBytes>()?));
|
||||
}
|
||||
if file.is_instance(&py.import("os")?.getattr("PathLike")?)? {
|
||||
let path: PathBuf = file.extract()?;
|
||||
let name = path
|
||||
.file_name()
|
||||
.map(|value| value.to_string_lossy().into_owned());
|
||||
let bytes = py
|
||||
.detach(|| {
|
||||
let mut bytes = Vec::new();
|
||||
std::fs::File::open(&path)?
|
||||
.take(OCR_INLINE_MAX_BYTES as u64 + 1)
|
||||
.read_to_end(&mut bytes)?;
|
||||
Ok::<_, std::io::Error>(bytes)
|
||||
})
|
||||
.map_err(|error| {
|
||||
if error.kind() == std::io::ErrorKind::NotFound {
|
||||
PyFileNotFoundError::new_err(format!("File not found: {}", path.display()))
|
||||
} else {
|
||||
error.into()
|
||||
}
|
||||
})?;
|
||||
return Ok((FileBytes::Native(bytes), name));
|
||||
}
|
||||
if file.is_instance_of::<PyBytes>() {
|
||||
return Ok((FileBytes::Python(file.extract()?), None));
|
||||
}
|
||||
let reader = file
|
||||
.getattr_opt("read")?
|
||||
.filter(|value| value.is_callable());
|
||||
let Some(reader) = reader else {
|
||||
return Err(PyValueError::new_err(format!(
|
||||
"Unsupported file input type: {}. Expected pathlib.Path, bytes, or a file-like object.",
|
||||
file.get_type(),
|
||||
)));
|
||||
};
|
||||
let name = file
|
||||
.getattr_opt("name")?
|
||||
.filter(|value| !value.is_none())
|
||||
.map(|value| value.extract::<String>())
|
||||
.transpose()?;
|
||||
let value = reader.call0()?;
|
||||
let bytes = if value.is_instance_of::<PyString>() {
|
||||
FileBytes::Native(value.extract::<String>()?.into_bytes())
|
||||
} else if value.is_instance_of::<PyBytes>() {
|
||||
FileBytes::Python(value.extract()?)
|
||||
} else {
|
||||
return Err(PyTypeError::new_err(format!(
|
||||
"OCR file read must return bytes or str, got {}",
|
||||
value.get_type(),
|
||||
)));
|
||||
};
|
||||
Ok((bytes, name))
|
||||
// Bytes subclasses may retain GC edges that a native Bytes owner cannot traverse.
|
||||
Ok(Bytes::copy_from_slice(value.extract::<PyBackedBytes>()?.as_ref()))
|
||||
}
|
||||
|
||||
|
||||
|
||||
pub(super) struct FileDocumentInput {
|
||||
bytes: FileBytes,
|
||||
name: Option<String>,
|
||||
mime_type: Option<String>,
|
||||
pub input: OcrDocumentInput,
|
||||
pub reader: Option<PythonFileReader>,
|
||||
}
|
||||
|
||||
impl FromPyObject<'_, '_> for FileDocumentInput {
|
||||
|
|
@ -103,123 +56,112 @@ impl FromPyObject<'_, '_> for FileDocumentInput {
|
|||
Err(error) if error.is_instance_of::<pyo3::exceptions::PyKeyError>(py) => None,
|
||||
Err(error) => return Err(error),
|
||||
};
|
||||
let missing = || PyValueError::new_err(
|
||||
"document with type='file' must include a 'file' field containing a pathlib.Path, file-like object, or bytes",
|
||||
);
|
||||
let file = document.get_item("file").map_err(|error| {
|
||||
if error.is_instance_of::<pyo3::exceptions::PyKeyError>(py) {
|
||||
PyValueError::new_err("document with type='file' must include a 'file' field containing a pathlib.Path, file-like object, or bytes")
|
||||
missing()
|
||||
} else {
|
||||
error
|
||||
}
|
||||
})?;
|
||||
if file.is_none() {
|
||||
return Err(missing());
|
||||
}
|
||||
if file.is_instance_of::<PyString>() {
|
||||
return Err(PyValueError::new_err(
|
||||
"document with type='file' must include a 'file' field containing a pathlib.Path, file-like object, or bytes",
|
||||
"OCR file input does not accept bare str values. Pass bytes, a pathlib.Path, or a file-like object.",
|
||||
));
|
||||
}
|
||||
let (bytes, name) = read_file_input(py, &file)?;
|
||||
if file.is_instance(&py.import("os")?.getattr("PathLike")?)? {
|
||||
return Ok(Self {
|
||||
input: OcrDocumentInput::Path {
|
||||
path: file.extract::<PathBuf>()?,
|
||||
mime_type,
|
||||
},
|
||||
reader: None,
|
||||
});
|
||||
}
|
||||
if file.is_instance_of::<PyBytes>() {
|
||||
return Ok(Self {
|
||||
input: OcrDocumentInput::Bytes {
|
||||
bytes: extract_bytes(&file)?,
|
||||
file_name: None,
|
||||
mime_type,
|
||||
},
|
||||
reader: None,
|
||||
});
|
||||
}
|
||||
let reader = file.getattr_opt("read")?.filter(|value| value.is_callable());
|
||||
let Some(reader) = reader else {
|
||||
return Err(PyValueError::new_err(format!(
|
||||
"Unsupported file input type: {}. Expected pathlib.Path, bytes, or a file-like object.",
|
||||
file.get_type(),
|
||||
)));
|
||||
};
|
||||
let name = file
|
||||
.getattr_opt("name")?
|
||||
.filter(|value| !value.is_none())
|
||||
.map(|value| value.extract::<String>())
|
||||
.transpose()?;
|
||||
Ok(Self {
|
||||
bytes,
|
||||
name,
|
||||
mime_type,
|
||||
input: OcrDocumentInput::HostReader { mime_type },
|
||||
reader: Some(PythonFileReader { reader: reader.unbind(), name }),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn file_document(py: Python<'_>, document: FileDocumentInput) -> PyResult<OcrDocument> {
|
||||
py.detach(|| {
|
||||
encode_file_document(
|
||||
document.bytes.as_ref(),
|
||||
document.name.as_deref(),
|
||||
document.mime_type.as_deref(),
|
||||
)
|
||||
})
|
||||
.map_err(|error| PyValueError::new_err(error.to_string()))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use pyo3::exceptions::PyTypeError;
|
||||
use pyo3::types::PyDict;
|
||||
|
||||
#[test]
|
||||
fn extraction_validates_required_file_and_optional_mime_type() {
|
||||
fn projection_validates_fields_without_consuming_readers_or_opening_paths() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
for expression in [c"{}", c"{'file': None}"] {
|
||||
let document = py.eval(expression, None, None).unwrap();
|
||||
let error = document.extract::<FileDocumentInput>().err().unwrap();
|
||||
assert!(error.is_instance_of::<PyValueError>(py));
|
||||
assert!(error.to_string().contains("must include a 'file' field"));
|
||||
}
|
||||
for expression in [
|
||||
c"{'file': b'abc', 'mime_type': None}",
|
||||
c"{'file': b'abc', 'mime_type': 7}",
|
||||
] {
|
||||
let document = py.eval(expression, None, None).unwrap();
|
||||
let error = document.extract::<FileDocumentInput>().err().unwrap();
|
||||
for expression in [c"{'file': b'abc', 'mime_type': None}", c"{'file': b'abc', 'mime_type': 7}"] {
|
||||
let error = py.eval(expression, None, None).unwrap().extract::<FileDocumentInput>().err().unwrap();
|
||||
assert!(error.is_instance_of::<PyTypeError>(py));
|
||||
}
|
||||
let document = py.eval(c"{'file': b'abc'}", None, None).unwrap();
|
||||
let input: FileDocumentInput = document.extract().unwrap();
|
||||
assert_eq!(input.bytes.as_ref(), b"abc");
|
||||
assert_eq!(input.name, None);
|
||||
assert_eq!(input.mime_type, None);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn extraction_validates_mime_type_before_consuming_file() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = PyDict::new(py);
|
||||
py.run(
|
||||
c"class Reader:
|
||||
py.run(c"from pathlib import Path
|
||||
failure = KeyError('reader failed')
|
||||
class Reader:
|
||||
def __init__(self):
|
||||
self.reads = 0
|
||||
def read(self):
|
||||
self.reads += 1
|
||||
return b'abc'
|
||||
raise failure
|
||||
reader = Reader()
|
||||
document = {'file': reader, 'mime_type': 7}",
|
||||
Some(&locals),
|
||||
Some(&locals),
|
||||
)
|
||||
.unwrap();
|
||||
document = {'file': reader}
|
||||
path_document = {'file': Path('/nonexistent/ocr-projection-test.pdf')}
|
||||
", Some(&locals), Some(&locals)).unwrap();
|
||||
let document = locals.get_item("document").unwrap().unwrap();
|
||||
let error = document.extract::<FileDocumentInput>().err().unwrap();
|
||||
assert!(error.is_instance_of::<PyTypeError>(py));
|
||||
let reads: usize = locals
|
||||
.get_item("reader")
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.getattr("reads")
|
||||
.unwrap()
|
||||
.extract()
|
||||
.unwrap();
|
||||
assert_eq!(reads, 0);
|
||||
let input: FileDocumentInput = document.extract().unwrap();
|
||||
assert_eq!(locals.get_item("reader").unwrap().unwrap().getattr("reads").unwrap().extract::<usize>().unwrap(), 0);
|
||||
assert!(input.reader.is_some());
|
||||
let path: FileDocumentInput = locals.get_item("path_document").unwrap().unwrap().extract().unwrap();
|
||||
assert!(matches!(path.input, OcrDocumentInput::Path { .. }));
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn extraction_preserves_reader_key_error_identity() {
|
||||
fn exact_python_bytes_transfer_without_copying_and_outlive_the_input() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = PyDict::new(py);
|
||||
py.run(
|
||||
c"failure = KeyError('reader failed')
|
||||
class Reader:
|
||||
def read(self):
|
||||
raise failure
|
||||
document = {'file': Reader()}",
|
||||
Some(&locals),
|
||||
Some(&locals),
|
||||
)
|
||||
.unwrap();
|
||||
let document = locals.get_item("document").unwrap().unwrap();
|
||||
let error = document.extract::<FileDocumentInput>().err().unwrap();
|
||||
assert!(
|
||||
error
|
||||
.value(py)
|
||||
.is(locals.get_item("failure").unwrap().unwrap())
|
||||
);
|
||||
let (bytes, pointer) = Python::attach(|py| {
|
||||
let value = PyBytes::new(py, b"document bytes");
|
||||
let pointer = value.as_bytes().as_ptr() as usize;
|
||||
(extract_bytes(value.as_any()).unwrap(), pointer)
|
||||
});
|
||||
assert_eq!(bytes.as_ptr() as usize, pointer);
|
||||
assert_eq!(bytes.as_ref(), b"document bytes");
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,170 +1,252 @@
|
|||
use litellm_core::ocr::Error;
|
||||
use litellm_core::transport::Error as TransportError;
|
||||
use pyo3::exceptions::{PyRuntimeError, PyValueError};
|
||||
use pyo3::exceptions::{PyBaseException, PyException};
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::types::PyDict;
|
||||
|
||||
use crate::errors::{RustUpstreamError, core_error_to_pyerr};
|
||||
|
||||
pub(super) fn to_pyerr(error: Error) -> PyErr {
|
||||
if let Error::Provider {
|
||||
status,
|
||||
body,
|
||||
headers,
|
||||
} = error
|
||||
{
|
||||
let mapped = attach_status(RustUpstreamError::new_err((status, body)), Some(status));
|
||||
return Python::attach(|py| -> PyResult<PyErr> {
|
||||
let headers =
|
||||
pyo3::types::PyDict::from_sequence(&headers.into_pyobject(py)?.into_any())?;
|
||||
mapped.value(py).setattr("headers", headers)?;
|
||||
Ok(mapped)
|
||||
})
|
||||
.unwrap_or_else(|error| error);
|
||||
}
|
||||
let (mapped, status) = match error {
|
||||
Error::MissingDocumentUrl => (
|
||||
PyValueError::new_err(Error::MissingDocumentUrl.to_string()),
|
||||
Some(500),
|
||||
),
|
||||
error @ Error::MissingField(_) => (PyValueError::new_err(error.to_string()), None),
|
||||
error if error.is_request() => (
|
||||
PyValueError::new_err(format!("invalid request: {error}")),
|
||||
Some(400),
|
||||
),
|
||||
error if error.is_response() => (
|
||||
PyRuntimeError::new_err(format!("invalid response: {error}")),
|
||||
None,
|
||||
),
|
||||
error @ (Error::InvalidRequest(_) | Error::Params(_) | Error::Headers(_)) => {
|
||||
(PyValueError::new_err(error.to_string()), Some(400))
|
||||
}
|
||||
Error::Transport(TransportError::Http { status, body }) => {
|
||||
(RustUpstreamError::new_err((status, body)), Some(status))
|
||||
}
|
||||
other => (core_error_to_pyerr(other), None),
|
||||
};
|
||||
attach_status(mapped, status)
|
||||
enum Public {
|
||||
BadRequest,
|
||||
Authentication,
|
||||
NotFound,
|
||||
Timeout,
|
||||
RateLimit,
|
||||
InternalServer,
|
||||
BadGateway,
|
||||
ServiceUnavailable,
|
||||
Api(u16),
|
||||
ApiConnection,
|
||||
}
|
||||
|
||||
fn attach_status(error: PyErr, status: Option<u16>) -> PyErr {
|
||||
if let Some(status) = status {
|
||||
Python::attach(|py| {
|
||||
let value = error.value(py);
|
||||
value.setattr("status_code", status).ok();
|
||||
value.setattr("message", value.to_string()).ok();
|
||||
});
|
||||
impl Public {
|
||||
fn class_name(&self) -> &'static str {
|
||||
match self {
|
||||
Self::BadRequest => "BadRequestError",
|
||||
Self::Authentication => "AuthenticationError",
|
||||
Self::NotFound => "NotFoundError",
|
||||
Self::Timeout => "Timeout",
|
||||
Self::RateLimit => "RateLimitError",
|
||||
Self::InternalServer => "InternalServerError",
|
||||
Self::BadGateway => "BadGatewayError",
|
||||
Self::ServiceUnavailable => "ServiceUnavailableError",
|
||||
Self::Api(_) => "APIError",
|
||||
Self::ApiConnection => "APIConnectionError",
|
||||
}
|
||||
}
|
||||
error
|
||||
|
||||
fn from_status(status: u16) -> Self {
|
||||
match status {
|
||||
400 | 422 => Self::BadRequest,
|
||||
401 => Self::Authentication,
|
||||
404 => Self::NotFound,
|
||||
408 | 504 => Self::Timeout,
|
||||
429 => Self::RateLimit,
|
||||
500 => Self::InternalServer,
|
||||
502 => Self::BadGateway,
|
||||
503 => Self::ServiceUnavailable,
|
||||
other => Self::Api(other),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct Failure {
|
||||
public: Public,
|
||||
detail: String,
|
||||
headers: Vec<(String, String)>,
|
||||
}
|
||||
|
||||
fn classify(error: Error) -> Failure {
|
||||
match error {
|
||||
Error::Provider {
|
||||
status,
|
||||
body,
|
||||
headers,
|
||||
} => Failure {
|
||||
public: Public::from_status(status),
|
||||
detail: body,
|
||||
headers,
|
||||
},
|
||||
Error::Transport(TransportError::Http { status, body }) => Failure {
|
||||
public: Public::from_status(status),
|
||||
detail: body,
|
||||
headers: Vec::new(),
|
||||
},
|
||||
error if error.is_request() => Failure {
|
||||
public: Public::BadRequest,
|
||||
detail: error.to_string(),
|
||||
headers: Vec::new(),
|
||||
},
|
||||
error => Failure {
|
||||
public: Public::ApiConnection,
|
||||
detail: error.to_string(),
|
||||
headers: Vec::new(),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
fn exception_provider(provider: &str) -> String {
|
||||
let mut chars = provider.chars();
|
||||
match chars.next() {
|
||||
Some(first) => format!("{}{}Exception", first.to_uppercase(), chars.as_str()),
|
||||
None => "Exception".into(),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn public_exception(
|
||||
py: Python<'_>,
|
||||
error: Error,
|
||||
model: &str,
|
||||
provider: &str,
|
||||
) -> PyResult<PyErr> {
|
||||
if let Error::FileRead { path, source } = error {
|
||||
let errno = source.raw_os_error().unwrap_or(match source.kind() {
|
||||
std::io::ErrorKind::NotFound => 2,
|
||||
std::io::ErrorKind::PermissionDenied => 13,
|
||||
_ => 5,
|
||||
});
|
||||
return Ok(PyErr::from_value(py.import("builtins")?.getattr("OSError")?.call1((
|
||||
errno, source.to_string(), path,
|
||||
))?));
|
||||
}
|
||||
raise_public(py, classify(error), model, provider, None)
|
||||
}
|
||||
|
||||
pub(super) fn public_host_exception(
|
||||
py: Python<'_>,
|
||||
error: &Py<PyBaseException>,
|
||||
model: &str,
|
||||
provider: &str,
|
||||
) -> PyResult<PyErr> {
|
||||
let value = error.bind(py);
|
||||
if !value.is_instance_of::<PyException>() || is_public(value)? {
|
||||
return Ok(PyErr::from_value(value.clone().into_any()));
|
||||
}
|
||||
let failure = Failure {
|
||||
public: Public::ApiConnection,
|
||||
detail: value.str()?.to_string(),
|
||||
headers: Vec::new(),
|
||||
};
|
||||
raise_public(py, failure, model, provider, Some(value))
|
||||
}
|
||||
|
||||
fn is_public(value: &Bound<'_, PyBaseException>) -> PyResult<bool> {
|
||||
let types = value
|
||||
.py()
|
||||
.import("litellm")?
|
||||
.getattr("LITELLM_EXCEPTION_TYPES")?;
|
||||
for public in types.try_iter()? {
|
||||
if value.is_instance(&public?)? {
|
||||
return Ok(true);
|
||||
}
|
||||
}
|
||||
Ok(false)
|
||||
}
|
||||
|
||||
fn raise_public(
|
||||
py: Python<'_>,
|
||||
failure: Failure,
|
||||
model: &str,
|
||||
provider: &str,
|
||||
context: Option<&Bound<'_, PyBaseException>>,
|
||||
) -> PyResult<PyErr> {
|
||||
let class = failure.public.class_name();
|
||||
let message = format!(
|
||||
"{}{} - {}",
|
||||
match failure.public {
|
||||
Public::RateLimit | Public::Api(_) | Public::ApiConnection => format!("{class}: "),
|
||||
_ => String::new(),
|
||||
},
|
||||
exception_provider(provider),
|
||||
failure.detail
|
||||
);
|
||||
let kwargs = PyDict::new(py);
|
||||
kwargs.set_item("message", message)?;
|
||||
kwargs.set_item("llm_provider", provider)?;
|
||||
kwargs.set_item("model", model)?;
|
||||
if let Public::Api(status) = failure.public {
|
||||
kwargs.set_item("status_code", status)?;
|
||||
}
|
||||
let value = py
|
||||
.import("litellm")?
|
||||
.getattr(class)?
|
||||
.call((), Some(&kwargs))?;
|
||||
if !failure.headers.is_empty() {
|
||||
let headers = PyDict::new(py);
|
||||
for (name, header) in &failure.headers {
|
||||
headers.set_item(name, header)?;
|
||||
}
|
||||
value.setattr("litellm_response_headers", headers)?;
|
||||
}
|
||||
if let Some(context) = context {
|
||||
value.setattr("__context__", context)?;
|
||||
}
|
||||
Ok(PyErr::from_value(value))
|
||||
}
|
||||
|
||||
pub(super) fn to_pyerr(error: Error) -> PyErr {
|
||||
Python::attach(|py| public_exception(py, error, "", "").unwrap_or_else(|error| error))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use pyo3::exceptions::PyValueError;
|
||||
|
||||
#[test]
|
||||
fn provider_error_retains_headers_at_the_python_boundary() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let error = to_pyerr(Error::Provider {
|
||||
status: 429,
|
||||
body: "rate limited".into(),
|
||||
headers: vec![("retry-after".into(), "17".into())],
|
||||
});
|
||||
assert!(error.is_instance_of::<RustUpstreamError>(py));
|
||||
assert_eq!(
|
||||
error
|
||||
.value(py)
|
||||
.getattr("args")
|
||||
.unwrap()
|
||||
.extract::<(u16, String)>()
|
||||
.unwrap(),
|
||||
(429, "rate limited".into())
|
||||
);
|
||||
assert_eq!(
|
||||
error
|
||||
.value(py)
|
||||
.getattr("headers")
|
||||
.unwrap()
|
||||
.get_item("retry-after")
|
||||
.unwrap()
|
||||
.extract::<String>()
|
||||
.unwrap(),
|
||||
"17"
|
||||
);
|
||||
});
|
||||
fn classified(error: Error) -> (&'static str, String, Vec<(String, String)>) {
|
||||
let failure = classify(error);
|
||||
(failure.public.class_name(), failure.detail, failure.headers)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn preserves_python_validation_and_provider_details() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let mapped = to_pyerr(Error::MissingDocumentUrl);
|
||||
assert!(mapped.is_instance_of::<pyo3::exceptions::PyValueError>(py));
|
||||
assert_eq!(mapped.value(py).to_string(), "Document URL is required");
|
||||
assert_eq!(
|
||||
mapped
|
||||
.value(py)
|
||||
.getattr("status_code")
|
||||
.unwrap()
|
||||
.extract::<u16>()
|
||||
.unwrap(),
|
||||
500
|
||||
);
|
||||
let mapped = to_pyerr(Error::Transport(litellm_core::transport::Error::Http {
|
||||
status: 429,
|
||||
body: r#"{"message":"rate limited"}"#.to_string(),
|
||||
fn provider_status_selects_public_class_and_keeps_details() {
|
||||
let (class, detail, headers) = classified(Error::Provider {
|
||||
status: 429,
|
||||
body: "rate limited".into(),
|
||||
headers: vec![("retry-after".into(), "17".into())],
|
||||
});
|
||||
assert_eq!(class, "RateLimitError");
|
||||
assert_eq!(detail, "rate limited");
|
||||
assert_eq!(headers, [("retry-after".to_string(), "17".to_string())]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn status_table_matches_legacy_openai_mapping() {
|
||||
for (status, class) in [
|
||||
(400, "BadRequestError"),
|
||||
(401, "AuthenticationError"),
|
||||
(404, "NotFoundError"),
|
||||
(408, "Timeout"),
|
||||
(422, "BadRequestError"),
|
||||
(429, "RateLimitError"),
|
||||
(500, "InternalServerError"),
|
||||
(502, "BadGatewayError"),
|
||||
(503, "ServiceUnavailableError"),
|
||||
(504, "Timeout"),
|
||||
(418, "APIError"),
|
||||
] {
|
||||
let (mapped, _, _) = classified(Error::Transport(TransportError::Http {
|
||||
status,
|
||||
body: "x".into(),
|
||||
}));
|
||||
assert!(mapped.is_instance_of::<RustUpstreamError>(py));
|
||||
let args: (u16, String) = mapped
|
||||
.value(py)
|
||||
.getattr("args")
|
||||
.and_then(|args| args.extract())
|
||||
.expect("OCR failures retain status and unprefixed provider message");
|
||||
assert_eq!(args, (429, r#"{"message":"rate limited"}"#.to_string()));
|
||||
|
||||
let mapped = to_pyerr(Error::InvalidRequest("invalid format".into()));
|
||||
assert!(mapped.is_instance_of::<PyValueError>(py));
|
||||
assert_eq!(
|
||||
mapped
|
||||
.value(py)
|
||||
.getattr("status_code")
|
||||
.unwrap()
|
||||
.extract::<u16>()
|
||||
.unwrap(),
|
||||
400
|
||||
);
|
||||
});
|
||||
assert_eq!(mapped, class, "{status}");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn typed_request_and_response_failures_keep_python_contracts() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let request = to_pyerr(Error::RequestField {
|
||||
path: "document.type".into(),
|
||||
});
|
||||
assert!(request.is_instance_of::<PyValueError>(py));
|
||||
assert_eq!(
|
||||
request
|
||||
.value(py)
|
||||
.getattr("status_code")
|
||||
.unwrap()
|
||||
.extract::<u16>()
|
||||
.unwrap(),
|
||||
400
|
||||
);
|
||||
assert_eq!(
|
||||
request.value(py).to_string(),
|
||||
"invalid request: invalid OCR request field: document.type"
|
||||
);
|
||||
let response = to_pyerr(Error::EmptyContent);
|
||||
assert!(response.is_instance_of::<PyRuntimeError>(py));
|
||||
assert_eq!(
|
||||
response.value(py).to_string(),
|
||||
"invalid response: OCR response is missing non-empty content"
|
||||
);
|
||||
assert!(!response.value(py).hasattr("status_code").unwrap());
|
||||
});
|
||||
fn request_shape_and_statusless_failures_use_their_public_families() {
|
||||
assert_eq!(classified(Error::RequestFormat).0, "BadRequestError");
|
||||
let (class, detail, _) = classified(Error::MissingAzureAiCredentials);
|
||||
assert_eq!(class, "APIConnectionError");
|
||||
assert!(detail.contains("Missing Azure AI credentials"));
|
||||
let (class, detail, _) = classified(Error::Transport(TransportError::Network(
|
||||
"connection reset".into(),
|
||||
)));
|
||||
assert_eq!(class, "APIConnectionError");
|
||||
assert!(detail.contains("connection reset"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn exception_provider_matches_legacy_capitalisation() {
|
||||
assert_eq!(exception_provider("mistral"), "MistralException");
|
||||
assert_eq!(exception_provider("azure_ai"), "Azure_aiException");
|
||||
assert_eq!(exception_provider(""), "Exception");
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,42 +1,24 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use pyo3::gc::{PyTraverseError, PyVisit};
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::types::{PyDict, PyTuple};
|
||||
use pyo3::types::PyDict;
|
||||
|
||||
use litellm_auth::ResolvedCredential;
|
||||
use litellm_core::ocr::hooks::{OcrDuringCallRequest, OcrPostCallRequest};
|
||||
use litellm_core::ocr::{OcrAdmission, OcrCall, OcrClient, OcrHostOperation, OcrHostResult};
|
||||
use litellm_core::ocr::{OcrCall, OcrHostOperation, OcrHostResult};
|
||||
use litellm_python_interop::{
|
||||
from_py_preserving_errors as from_py, to_py_preserving_errors as to_py,
|
||||
};
|
||||
|
||||
use super::callbacks;
|
||||
use super::document::PythonFileReader;
|
||||
use super::errors::to_pyerr as ocr_error_to_pyerr;
|
||||
use super::project::{admitted_call, project};
|
||||
use super::project::project;
|
||||
use super::{ASYNC_SIGNATURE, SIGNATURE, callbacks, errors};
|
||||
use crate::auth::PythonTokenProvider;
|
||||
use crate::lifecycle::{
|
||||
OperationClass, PythonCallState, PythonRoute, Signature, missing_state, now, run_call,
|
||||
};
|
||||
use crate::lifecycle::{OperationClass, PythonCallState, PythonRoute, missing_state, now};
|
||||
|
||||
const SIGNATURE: Signature = Signature {
|
||||
name: "ocr",
|
||||
parameters: &[
|
||||
"model",
|
||||
"document",
|
||||
"api_key",
|
||||
"api_base",
|
||||
"timeout",
|
||||
"custom_llm_provider",
|
||||
"extra_headers",
|
||||
],
|
||||
required: 2,
|
||||
};
|
||||
|
||||
const ASYNC_SIGNATURE: Signature = Signature {
|
||||
name: "aocr",
|
||||
..SIGNATURE
|
||||
};
|
||||
|
||||
struct PythonOcrHost {
|
||||
pub(super) struct PythonOcrHost {
|
||||
state: PythonCallState,
|
||||
projected: Option<ProjectedOcrHost>,
|
||||
}
|
||||
|
|
@ -48,6 +30,8 @@ struct ProjectedOcrHost {
|
|||
azure_ad_token_provider: Option<PythonTokenProvider>,
|
||||
pre_call: Option<callbacks::OcrLoggingFields>,
|
||||
payload: Option<CapturedOcrPayload>,
|
||||
reader: Option<PythonFileReader>,
|
||||
reader_failed: bool,
|
||||
}
|
||||
|
||||
struct CapturedOcrPayload {
|
||||
|
|
@ -63,6 +47,13 @@ impl CapturedOcrPayload {
|
|||
}
|
||||
|
||||
impl PythonOcrHost {
|
||||
pub(super) fn new(state: PythonCallState) -> Self {
|
||||
Self {
|
||||
state,
|
||||
projected: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn projected(&self) -> PyResult<&ProjectedOcrHost> {
|
||||
self.projected.as_ref().ok_or_else(missing_state)
|
||||
}
|
||||
|
|
@ -87,9 +78,15 @@ impl PythonOcrHost {
|
|||
azure_ad_token_provider: projected.azure_ad_token_provider,
|
||||
pre_call: None,
|
||||
payload: None,
|
||||
reader: projected.reader,
|
||||
reader_failed: false,
|
||||
});
|
||||
Ok(OcrHostResult::Request(Ok((
|
||||
Box::new(projected.request),
|
||||
Box::new(
|
||||
projected
|
||||
.request
|
||||
.with_host_hooks(Arc::new(BridgeOcrHooks), None),
|
||||
),
|
||||
has_token_provider,
|
||||
))))
|
||||
}
|
||||
|
|
@ -161,20 +158,21 @@ impl PythonOcrHost {
|
|||
}
|
||||
|
||||
fn map_failure(&mut self, py: Python<'_>, error: litellm_core::ocr::Error) -> PyResult<()> {
|
||||
if self.state.error.is_none() {
|
||||
self.state.retain_error(py, ocr_error_to_pyerr(error));
|
||||
}
|
||||
if self.state.end.is_none() {
|
||||
self.state.end = Some(now(py)?);
|
||||
}
|
||||
let error = self.state.error.as_ref().ok_or_else(missing_state)?;
|
||||
let (model, provider) = match &self.projected {
|
||||
Some(projected) => (projected.model.as_str(), projected.provider),
|
||||
None => ("", ""),
|
||||
};
|
||||
let mapped = callbacks::map_failure(py, error, model, provider, &self.state.kwargs)?;
|
||||
self.state
|
||||
.retain_error(py, PyErr::from_value(mapped.into_bound(py).into_any()));
|
||||
let mapped = match self.state.error.take() {
|
||||
Some(host_error) if self.projected.as_ref().is_some_and(|host| host.reader_failed) => {
|
||||
PyErr::from_value(host_error.into_bound(py).into_any())
|
||||
}
|
||||
Some(host_error) => errors::public_host_exception(py, &host_error, model, provider)?,
|
||||
None => errors::public_exception(py, error, model, provider)?,
|
||||
};
|
||||
self.state.retain_error(py, mapped);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
|
@ -211,6 +209,12 @@ impl PythonRoute for PythonOcrHost {
|
|||
fn invoke(&mut self, py: Python<'_>, operation: OcrHostOperation) -> PyResult<OcrHostResult> {
|
||||
Ok(match operation {
|
||||
OcrHostOperation::ProjectRequest => self.project(py)?,
|
||||
OcrHostOperation::ReadDocument => {
|
||||
let reader = self.projected_mut()?.reader.take().ok_or_else(missing_state)?;
|
||||
let result = reader.read(py);
|
||||
self.projected_mut()?.reader_failed = result.is_err();
|
||||
OcrHostResult::Document(Ok(result?))
|
||||
}
|
||||
OcrHostOperation::AcquireAzureAdToken => {
|
||||
OcrHostResult::AzureAdToken(Ok(self.acquire_azure_ad_token(py)?))
|
||||
}
|
||||
|
|
@ -250,6 +254,9 @@ impl PythonRoute for PythonOcrHost {
|
|||
if let Some(provider) = &projected.azure_ad_token_provider {
|
||||
provider.traverse(visit)?;
|
||||
}
|
||||
if let Some(reader) = &projected.reader {
|
||||
reader.traverse(visit)?;
|
||||
}
|
||||
if let Some(payload) = &projected.payload {
|
||||
payload.traverse(visit)?;
|
||||
}
|
||||
|
|
@ -257,69 +264,10 @@ impl PythonRoute for PythonOcrHost {
|
|||
}
|
||||
}
|
||||
|
||||
pub(super) struct BridgeOcrHooks;
|
||||
struct BridgeOcrHooks;
|
||||
|
||||
impl litellm_core::ocr::hooks::OcrHooks for BridgeOcrHooks {
|
||||
fn intercepts_requests(&self) -> bool {
|
||||
true
|
||||
}
|
||||
}
|
||||
|
||||
fn call(
|
||||
py: Python<'_>,
|
||||
args: Bound<'_, PyTuple>,
|
||||
kwargs: Option<Bound<'_, PyDict>>,
|
||||
asynchronous: bool,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
let kwargs = kwargs.unwrap_or_else(|| PyDict::new(py));
|
||||
let signature = if asynchronous {
|
||||
&ASYNC_SIGNATURE
|
||||
} else {
|
||||
&SIGNATURE
|
||||
};
|
||||
signature.bind(&args, &kwargs)?;
|
||||
let client = OcrClient::shared().map_err(ocr_error_to_pyerr)?;
|
||||
let call = admitted_call(OcrCall::admit(
|
||||
client,
|
||||
OcrAdmission {
|
||||
asynchronous,
|
||||
..OcrAdmission::all()
|
||||
},
|
||||
))?;
|
||||
let host = PythonOcrHost {
|
||||
state: PythonCallState::new(
|
||||
py,
|
||||
args.unbind(),
|
||||
kwargs.copy()?.unbind(),
|
||||
asynchronous,
|
||||
signature.name,
|
||||
)?,
|
||||
projected: None,
|
||||
};
|
||||
run_call(py, call, host)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
#[pyo3(signature = (*args, **kwargs))]
|
||||
fn ocr(
|
||||
py: Python<'_>,
|
||||
args: Bound<'_, PyTuple>,
|
||||
kwargs: Option<Bound<'_, PyDict>>,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
call(py, args, kwargs, false)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
#[pyo3(signature = (*args, **kwargs))]
|
||||
fn aocr(
|
||||
py: Python<'_>,
|
||||
args: Bound<'_, PyTuple>,
|
||||
kwargs: Option<Bound<'_, PyDict>>,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
call(py, args, kwargs, true)
|
||||
}
|
||||
|
||||
pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
super::super::add_function(module, wrap_pyfunction!(ocr, module)?)?;
|
||||
super::super::add_function(module, wrap_pyfunction!(aocr, module)?)
|
||||
}
|
||||
|
|
@ -1,11 +1,126 @@
|
|||
mod callbacks;
|
||||
mod document;
|
||||
mod errors;
|
||||
mod lifecycle;
|
||||
mod host;
|
||||
mod project;
|
||||
|
||||
use litellm_core::ocr::{NativeOutcome, OcrAdmission, OcrCall, OcrClient};
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::types::{PyDict, PyTuple};
|
||||
|
||||
use self::errors::to_pyerr as ocr_error_to_pyerr;
|
||||
use self::host::PythonOcrHost;
|
||||
use crate::errors::RustBridgeDeclined;
|
||||
use crate::lifecycle::{PythonCallState, Signature, run_call};
|
||||
|
||||
const SIGNATURE: Signature = Signature {
|
||||
name: "ocr",
|
||||
parameters: &[
|
||||
"model",
|
||||
"document",
|
||||
"api_key",
|
||||
"api_base",
|
||||
"timeout",
|
||||
"custom_llm_provider",
|
||||
"extra_headers",
|
||||
],
|
||||
required: 2,
|
||||
};
|
||||
|
||||
const ASYNC_SIGNATURE: Signature = Signature {
|
||||
name: "aocr",
|
||||
..SIGNATURE
|
||||
};
|
||||
|
||||
fn admitted_call(outcome: NativeOutcome<OcrCall>) -> PyResult<OcrCall> {
|
||||
match outcome {
|
||||
NativeOutcome::Completed(call) => Ok(call),
|
||||
NativeOutcome::Declined(reason) => Err(RustBridgeDeclined::new_err(format!(
|
||||
"native OCR admission declined: {reason:?}"
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
fn call(
|
||||
py: Python<'_>,
|
||||
args: Bound<'_, PyTuple>,
|
||||
kwargs: Option<Bound<'_, PyDict>>,
|
||||
asynchronous: bool,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
let kwargs = kwargs.unwrap_or_else(|| PyDict::new(py));
|
||||
let signature = if asynchronous {
|
||||
&ASYNC_SIGNATURE
|
||||
} else {
|
||||
&SIGNATURE
|
||||
};
|
||||
signature.bind(&args, &kwargs)?;
|
||||
let client = OcrClient::shared().map_err(ocr_error_to_pyerr)?;
|
||||
let call = admitted_call(OcrCall::admit(
|
||||
client,
|
||||
OcrAdmission {
|
||||
asynchronous,
|
||||
..OcrAdmission::all()
|
||||
},
|
||||
))?;
|
||||
let host = PythonOcrHost::new(PythonCallState::new(
|
||||
py,
|
||||
args.unbind(),
|
||||
kwargs.copy()?.unbind(),
|
||||
asynchronous,
|
||||
signature.name,
|
||||
)?);
|
||||
run_call(py, call, host)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
#[pyo3(signature = (*args, **kwargs))]
|
||||
fn ocr(
|
||||
py: Python<'_>,
|
||||
args: Bound<'_, PyTuple>,
|
||||
kwargs: Option<Bound<'_, PyDict>>,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
call(py, args, kwargs, false)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
#[pyo3(signature = (*args, **kwargs))]
|
||||
fn aocr(
|
||||
py: Python<'_>,
|
||||
args: Bound<'_, PyTuple>,
|
||||
kwargs: Option<Bound<'_, PyDict>>,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
call(py, args, kwargs, true)
|
||||
}
|
||||
|
||||
pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
lifecycle::register(module)
|
||||
super::add_function(module, wrap_pyfunction!(ocr, module)?)?;
|
||||
super::add_function(module, wrap_pyfunction!(aocr, module)?)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use litellm_core::ocr::{Error, OcrDecline};
|
||||
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn typed_initial_decline_uses_bridge_decline_contract() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let Err(error) = admitted_call(NativeOutcome::Declined(OcrDecline::HostOperations))
|
||||
else {
|
||||
panic!("unsupported host operations should decline admission");
|
||||
};
|
||||
assert!(error.is_instance_of::<RustBridgeDeclined>(py));
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn post_admission_error_does_not_use_bridge_decline_contract() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let error = ocr_error_to_pyerr(Error::InvalidRequest("callback result".into()));
|
||||
assert!(!error.is_instance_of::<RustBridgeDeclined>(py));
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,153 +1,113 @@
|
|||
use std::collections::BTreeMap;
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
//! `ProjectRequest` host step: read the bound `ocr()` arguments and hand core a
|
||||
//! typed [`LiteLLMOcrRequest`]. Generic argument handling lives in
|
||||
//! [`crate::marshal::BoundRouteInputs`]; this module only owns what is specific
|
||||
//! to OCR: the `document` argument and the core constructor call.
|
||||
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::types::PyDict;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use litellm_auth::InputSource;
|
||||
use litellm_core::ocr::{
|
||||
LiteLLMOcrRequest, NativeOutcome, OcrCall, OcrCredentialInputs, OcrDocument,
|
||||
consumed_optional_params,
|
||||
LiteLLMOcrRequest, OcrConnectionInputs, OcrDocument, consumed_optional_params,
|
||||
};
|
||||
use litellm_python_interop::from_py_preserving_errors as from_py;
|
||||
|
||||
use super::errors::to_pyerr as ocr_error_to_pyerr;
|
||||
use super::lifecycle::BridgeOcrHooks;
|
||||
use crate::auth::{AZURE_AD_TOKEN_PROVIDER, PythonTokenProvider};
|
||||
use crate::errors::RustBridgeDeclined;
|
||||
use super::document::{FileDocumentInput, PythonFileReader};
|
||||
use crate::auth::PythonTokenProvider;
|
||||
use crate::lifecycle::BoundArguments;
|
||||
use crate::marshal::{project_optional_fields, python_timeout_seconds, request_input_sources};
|
||||
use crate::marshal::BoundRouteInputs;
|
||||
|
||||
/// Positional parameters of `ocr()` that are never projected into
|
||||
/// `optional_params`.
|
||||
const BOUND_FIELDS: &[&str] = &["model", "document", "timeout", "input_sources"];
|
||||
|
||||
pub(super) struct ProjectedOcrCall {
|
||||
pub request: LiteLLMOcrRequest,
|
||||
pub azure_ad_token_provider: Option<PythonTokenProvider>,
|
||||
pub secret_fields: Vec<&'static str>,
|
||||
pub reader: Option<PythonFileReader>,
|
||||
}
|
||||
|
||||
fn project_document(py: Python<'_>, document: &Bound<'_, PyAny>) -> PyResult<Value> {
|
||||
enum ProjectedDocument {
|
||||
File(FileDocumentInput),
|
||||
Url(serde_json::Value),
|
||||
}
|
||||
|
||||
impl ProjectedDocument {
|
||||
fn into_native(self) -> Result<FileDocumentInput, litellm_core::ocr::Error> {
|
||||
match self {
|
||||
Self::File(file) => Ok(file),
|
||||
Self::Url(value) => Ok(FileDocumentInput {
|
||||
input: OcrDocument::try_from(value)?.into(),
|
||||
reader: None,
|
||||
}),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn project_document(document: &Bound<'_, PyAny>) -> PyResult<ProjectedDocument> {
|
||||
let kind: String = document.get_item("type")?.extract()?;
|
||||
if kind != "file" {
|
||||
return from_py(document);
|
||||
return Ok(ProjectedDocument::Url(from_py(document)?));
|
||||
}
|
||||
let encoded = super::document::file_document(py, document.extract()?)?;
|
||||
serde_json::to_value(encoded)
|
||||
.map_err(|error| pyo3::exceptions::PyValueError::new_err(error.to_string()))
|
||||
document.extract().map(ProjectedDocument::File)
|
||||
}
|
||||
|
||||
fn header_pairs(headers: Option<Map<String, Value>>) -> Result<Vec<(String, String)>, litellm_core::ocr::Error> {
|
||||
headers
|
||||
.unwrap_or_default()
|
||||
.into_iter()
|
||||
.map(|(name, value)| {
|
||||
value
|
||||
.as_str()
|
||||
.map(|value| (name.clone(), value.to_string()))
|
||||
.ok_or_else(|| litellm_core::ocr::Error::RequestField {
|
||||
path: format!("extra_headers.{name}"),
|
||||
})
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn source_for(sources: &BTreeMap<String, InputSource>, name: &str) -> InputSource {
|
||||
sources.get(name).copied().unwrap_or_default()
|
||||
}
|
||||
|
||||
pub(super) fn project(py: Python<'_>, arguments: &BoundArguments<'_>) -> PyResult<ProjectedOcrCall> {
|
||||
let kwargs: &Bound<'_, PyDict> = arguments.kwargs();
|
||||
let model: String = arguments.extract("model")?;
|
||||
let custom_llm_provider: Option<String> = arguments.optional("custom_llm_provider")?;
|
||||
let document = project_document(py, &arguments.required("document")?)?;
|
||||
let api_key: Option<String> = arguments.optional("api_key")?;
|
||||
let specs = consumed_optional_params(&model, custom_llm_provider.as_deref())
|
||||
.map_err(ocr_error_to_pyerr)?;
|
||||
let optional_params = project_optional_fields(kwargs, &specs, BOUND_FIELDS)?;
|
||||
let input_sources = request_input_sources(
|
||||
kwargs,
|
||||
optional_params
|
||||
.keys()
|
||||
.map(String::as_str)
|
||||
.chain(["api_key", "api_base", "extra_headers"]),
|
||||
)?;
|
||||
let azure_ad_token_provider = kwargs
|
||||
.get_item("azure_ad_token_provider")?
|
||||
.and_then(|provider| PythonTokenProvider::select(provider, AZURE_AD_TOKEN_PROVIDER));
|
||||
let api_base: Option<String> = arguments.optional("api_base")?;
|
||||
let extra_headers: Option<Map<String, Value>> = arguments
|
||||
.get("extra_headers")?
|
||||
.filter(|value| !value.is_none())
|
||||
.map(|value| from_py(&value))
|
||||
.transpose()?;
|
||||
let timeout = arguments
|
||||
.get("timeout")?
|
||||
.filter(|value| !value.is_none())
|
||||
.map(|value| python_timeout_seconds(py, value.unbind()))
|
||||
.transpose()?
|
||||
.flatten()
|
||||
.map(|seconds| {
|
||||
Duration::try_from_secs_f64(seconds).map_err(|_| {
|
||||
litellm_core::ocr::Error::RequestField {
|
||||
path: "timeout_seconds".into(),
|
||||
}
|
||||
})
|
||||
})
|
||||
.transpose()
|
||||
.map_err(ocr_error_to_pyerr)?;
|
||||
|
||||
let request = (|| {
|
||||
let core = LiteLLMOcrRequest::new(
|
||||
model,
|
||||
OcrDocument::try_from(document)?,
|
||||
custom_llm_provider.as_deref(),
|
||||
optional_params.into(),
|
||||
)?;
|
||||
let transport = core.transport.clone().with_overrides(
|
||||
header_pairs(extra_headers)?,
|
||||
source_for(&input_sources, "extra_headers"),
|
||||
timeout,
|
||||
);
|
||||
let credentials = OcrCredentialInputs::new(
|
||||
/// Pure core assembly; every failure here is a typed `ocr::Error`.
|
||||
fn build_request(
|
||||
inputs: BoundRouteInputs,
|
||||
document: ProjectedDocument,
|
||||
) -> Result<ProjectedOcrCall, litellm_core::ocr::Error> {
|
||||
let document = document.into_native()?;
|
||||
let BoundRouteInputs {
|
||||
model,
|
||||
custom_llm_provider,
|
||||
api_key,
|
||||
api_base,
|
||||
extra_headers,
|
||||
timeout,
|
||||
optional_params,
|
||||
input_sources,
|
||||
azure_ad_token_provider,
|
||||
secret_fields,
|
||||
} = inputs;
|
||||
let request = LiteLLMOcrRequest::from_inputs(
|
||||
model,
|
||||
document.input,
|
||||
custom_llm_provider.as_deref(),
|
||||
optional_params.into(),
|
||||
OcrConnectionInputs {
|
||||
api_key,
|
||||
source_for(&input_sources, "api_key"),
|
||||
api_base,
|
||||
source_for(&input_sources, "api_base"),
|
||||
);
|
||||
Ok::<_, litellm_core::ocr::Error>(
|
||||
core.with_connection_inputs(credentials, transport, input_sources)
|
||||
.with_host_hooks(Arc::new(BridgeOcrHooks), None),
|
||||
)
|
||||
})()
|
||||
.map_err(ocr_error_to_pyerr)?;
|
||||
|
||||
extra_headers,
|
||||
timeout,
|
||||
input_sources,
|
||||
},
|
||||
)?;
|
||||
Ok(ProjectedOcrCall {
|
||||
request,
|
||||
azure_ad_token_provider,
|
||||
secret_fields: specs
|
||||
.into_iter()
|
||||
.filter(|spec| spec.secret)
|
||||
.map(|spec| spec.name)
|
||||
.collect(),
|
||||
secret_fields,
|
||||
reader: document.reader,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn admitted_call(outcome: NativeOutcome<OcrCall>) -> PyResult<OcrCall> {
|
||||
match outcome {
|
||||
NativeOutcome::Completed(call) => Ok(call),
|
||||
NativeOutcome::Declined(reason) => Err(RustBridgeDeclined::new_err(format!(
|
||||
"native OCR admission declined: {reason:?}"
|
||||
))),
|
||||
}
|
||||
pub(super) fn project(
|
||||
py: Python<'_>,
|
||||
arguments: &BoundArguments<'_>,
|
||||
) -> PyResult<ProjectedOcrCall> {
|
||||
let model: String = arguments.extract("model")?;
|
||||
let custom_llm_provider: Option<String> = arguments.optional("custom_llm_provider")?;
|
||||
let document = project_document(&arguments.required("document")?)?;
|
||||
let consumed = consumed_optional_params(&model, custom_llm_provider.as_deref())
|
||||
.map_err(ocr_error_to_pyerr)?;
|
||||
let inputs = BoundRouteInputs::extract(py, arguments, consumed, BOUND_FIELDS)?;
|
||||
build_request(inputs, document).map_err(ocr_error_to_pyerr)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use litellm_core::ocr::Error;
|
||||
use litellm_core::ocr::OcrDecline;
|
||||
use pyo3::exceptions::{PyKeyError, PyTypeError, PyValueError};
|
||||
use pyo3::exceptions::{PyKeyError, PyTypeError};
|
||||
use pyo3::types::PyDict;
|
||||
|
||||
use super::*;
|
||||
|
||||
|
|
@ -171,29 +131,7 @@ mod tests {
|
|||
}
|
||||
|
||||
#[test]
|
||||
fn typed_initial_decline_uses_bridge_decline_contract() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let Err(error) = admitted_call(NativeOutcome::Declined(OcrDecline::HostOperations))
|
||||
else {
|
||||
panic!("unsupported host operations should decline admission");
|
||||
};
|
||||
assert!(error.is_instance_of::<RustBridgeDeclined>(py));
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn post_admission_error_does_not_use_bridge_decline_contract() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let error = ocr_error_to_pyerr(Error::InvalidRequest("callback result".into()));
|
||||
assert!(error.is_instance_of::<PyValueError>(py));
|
||||
assert!(!error.is_instance_of::<RustBridgeDeclined>(py));
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn file_documents_are_encoded_and_other_documents_pass_through() {
|
||||
fn file_documents_remain_raw_and_other_documents_pass_through() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let file = py
|
||||
|
|
@ -203,13 +141,11 @@ mod tests {
|
|||
None,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
project_document(py, &file).unwrap(),
|
||||
serde_json::json!({
|
||||
"type": "document_url",
|
||||
"document_url": "data:application/pdf;base64,JVBERi0xLjQ=",
|
||||
})
|
||||
);
|
||||
assert!(matches!(
|
||||
project_document(&file).unwrap().into_native().unwrap().input,
|
||||
litellm_core::ocr::OcrDocumentInput::Bytes { bytes, mime_type, .. }
|
||||
if bytes == b"%PDF-1.4"[..] && mime_type.as_deref() == Some("application/pdf")
|
||||
));
|
||||
|
||||
let original = py
|
||||
.eval(
|
||||
|
|
@ -218,13 +154,11 @@ mod tests {
|
|||
None,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
project_document(py, &original).unwrap(),
|
||||
serde_json::json!({
|
||||
"type": "document_url",
|
||||
"document_url": "https://example.com/a.pdf",
|
||||
})
|
||||
);
|
||||
assert!(matches!(
|
||||
project_document(&original).unwrap().into_native().unwrap().input,
|
||||
litellm_core::ocr::OcrDocumentInput::Document(OcrDocument::DocumentUrl { document_url, .. })
|
||||
if document_url == "https://example.com/a.pdf"
|
||||
));
|
||||
});
|
||||
}
|
||||
|
||||
|
|
@ -235,12 +169,7 @@ mod tests {
|
|||
let document = py
|
||||
.eval(c"{'type': 'mystery', 'mystery': 'x'}", None, None)
|
||||
.unwrap();
|
||||
let wire_document = project_document(py, &document).unwrap();
|
||||
assert_eq!(
|
||||
wire_document,
|
||||
serde_json::json!({"type": "mystery", "mystery": "x"})
|
||||
);
|
||||
let error = OcrDocument::try_from(wire_document).unwrap_err();
|
||||
let error = project_document(&document).unwrap().into_native().err().unwrap();
|
||||
assert!(error.to_string().contains("document"));
|
||||
});
|
||||
}
|
||||
|
|
@ -251,15 +180,15 @@ mod tests {
|
|||
Python::attach(|py| {
|
||||
let missing = py.eval(c"{}", None, None).unwrap();
|
||||
assert!(
|
||||
project_document(py, &missing)
|
||||
.unwrap_err()
|
||||
project_document(&missing)
|
||||
.err().unwrap()
|
||||
.is_instance_of::<PyKeyError>(py)
|
||||
);
|
||||
|
||||
let non_string = py.eval(c"{'type': 1}", None, None).unwrap();
|
||||
assert!(
|
||||
project_document(py, &non_string)
|
||||
.unwrap_err()
|
||||
project_document(&non_string)
|
||||
.err().unwrap()
|
||||
.is_instance_of::<PyTypeError>(py)
|
||||
);
|
||||
|
||||
|
|
@ -274,7 +203,7 @@ document = Document()
|
|||
",
|
||||
);
|
||||
let error =
|
||||
project_document(py, &locals.get_item("document").unwrap().unwrap()).unwrap_err();
|
||||
project_document(&locals.get_item("document").unwrap().unwrap()).err().unwrap();
|
||||
assert!(
|
||||
error
|
||||
.value(py)
|
||||
|
|
@ -303,8 +232,8 @@ document = Document()
|
|||
",
|
||||
);
|
||||
let document = locals.get_item("document").unwrap().unwrap();
|
||||
let wire = project_document(py, &document).unwrap();
|
||||
assert_eq!(wire["type"], "document_url");
|
||||
let projected = project_document(&document).unwrap().into_native().unwrap();
|
||||
assert!(matches!(projected.input, litellm_core::ocr::OcrDocumentInput::Bytes { .. }));
|
||||
let reads: Vec<String> = document.getattr("reads").unwrap().extract().unwrap();
|
||||
assert_eq!(reads, ["type", "mime_type", "file"]);
|
||||
});
|
||||
|
|
|
|||
|
|
@ -1,12 +1,11 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Coroutine, Mapping
|
||||
from collections.abc import Callable, Coroutine, Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm.llms.base_llm.ocr.transformation import PROVIDER_NATIVE_RESPONSE_KEY, OCRResponse
|
||||
from litellm.rust_bridge.bindings import NativeBinding
|
||||
|
||||
|
|
@ -32,6 +31,15 @@ NATIVE_AOCR: Final = NativeBinding("aocr", validate=_as_aocr)
|
|||
_NATIVE_RESPONSE: Final = TypeAdapter(Mapping[str, object])
|
||||
|
||||
|
||||
def read_document(reader: Callable[[], object]) -> bytes:
|
||||
value: Final = reader()
|
||||
if isinstance(value, str):
|
||||
return value.encode("utf-8")
|
||||
if isinstance(value, bytes):
|
||||
return value
|
||||
raise TypeError(f"OCR file read must return bytes or str, got {type(value)}")
|
||||
|
||||
|
||||
def build_response(response: Mapping[str, object]) -> OCRResponse:
|
||||
provider_native_response: Final = response.get(PROVIDER_NATIVE_RESPONSE_KEY)
|
||||
normalized: Final = OCRResponse.model_validate(
|
||||
|
|
@ -40,32 +48,3 @@ def build_response(response: Mapping[str, object]) -> OCRResponse:
|
|||
if isinstance(provider_native_response, Mapping):
|
||||
normalized.set_provider_native_response(_NATIVE_RESPONSE.validate_python(provider_native_response))
|
||||
return normalized
|
||||
|
||||
|
||||
class ExceptionMapper(Protocol):
|
||||
def __call__(
|
||||
self,
|
||||
*,
|
||||
model: str,
|
||||
custom_llm_provider: str | None,
|
||||
original_exception: Exception,
|
||||
completion_kwargs: dict[str, object],
|
||||
extra_kwargs: dict[str, object],
|
||||
) -> Exception: ...
|
||||
|
||||
|
||||
def map_failure(error: Exception, model: str, provider: str, kwargs: Mapping[str, object]) -> Exception:
|
||||
mapper: Final = cast(
|
||||
ExceptionMapper, litellm.exception_type
|
||||
) # cast-ok: bounded adapter for the public exception mapper
|
||||
try:
|
||||
return mapper(
|
||||
model=model,
|
||||
custom_llm_provider=provider,
|
||||
original_exception=error,
|
||||
completion_kwargs=dict(kwargs), # mutable-ok: exception mapper requires owned kwargs
|
||||
extra_kwargs=dict(kwargs), # mutable-ok: exception mapper requires owned kwargs
|
||||
)
|
||||
except Exception as public_error:
|
||||
public_error.__context__ = error
|
||||
return public_error
|
||||
|
|
|
|||
|
|
@ -1,5 +1,10 @@
|
|||
import json
|
||||
import asyncio
|
||||
import base64
|
||||
import threading
|
||||
import weakref
|
||||
import gc
|
||||
from pathlib import Path
|
||||
from collections.abc import Generator
|
||||
from datetime import datetime, timezone
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
|
|
@ -332,6 +337,126 @@ def test_native_lifecycle_core_encodes_python_file_input(
|
|||
assert requests[0]["body"]["opaque_extension"] == {"nested": [None, False, 0]}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("asynchronous", [False, True])
|
||||
@pytest.mark.parametrize("source", ["path", "pathlike", "reader", "text", "bytes"])
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_file_sources_preserve_sdk_behavior(ocr_server, tmp_path, asynchronous, source):
|
||||
server, requests = ocr_server
|
||||
path: Final = tmp_path / "scan.png"
|
||||
content: Final = b"document bytes"
|
||||
path.write_bytes(content)
|
||||
|
||||
class CustomPath:
|
||||
def __fspath__(self) -> str:
|
||||
return str(path)
|
||||
|
||||
def __str__(self) -> str:
|
||||
return "/not/the/document.pdf"
|
||||
|
||||
class Reader:
|
||||
name = "scan.png"
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.stream = BytesIO(b"prefix" + content)
|
||||
self.stream.seek(6)
|
||||
self.calls = 0
|
||||
|
||||
def read(self) -> bytes | str:
|
||||
assert threading.get_ident() == caller_thread
|
||||
if asynchronous:
|
||||
assert asyncio.current_task() is caller_task
|
||||
self.calls += 1
|
||||
value: Final = self.stream.read()
|
||||
return value.decode() if source == "text" else value
|
||||
|
||||
caller_thread: Final = threading.get_ident()
|
||||
caller_task: Final = asyncio.current_task()
|
||||
reader: Final = Reader()
|
||||
file_input: Final = {
|
||||
"path": path,
|
||||
"pathlike": CustomPath(),
|
||||
"reader": reader,
|
||||
"text": reader,
|
||||
"bytes": content,
|
||||
}[source]
|
||||
document: Final = {"type": "file", "file": file_input, **({"mime_type": "image/png"} if source == "bytes" else {})}
|
||||
litellm.rust(True)
|
||||
kwargs: Final = {
|
||||
"model": "mistral/mistral-ocr-latest",
|
||||
"document": document,
|
||||
"api_key": "test-key",
|
||||
"api_base": f"http://127.0.0.1:{server.server_port}",
|
||||
}
|
||||
response: Final = await litellm.aocr(**kwargs) if asynchronous else litellm.ocr(**kwargs)
|
||||
assert response.pages[0].markdown == "native OCR response"
|
||||
assert len(requests) == 1
|
||||
assert requests[0]["body"]["document"] == {
|
||||
"type": "image_url",
|
||||
"image_url": "data:image/png;base64," + base64.b64encode(content).decode(),
|
||||
}
|
||||
assert reader.calls == (1 if source in {"reader", "text"} else 0)
|
||||
assert not reader.stream.closed
|
||||
|
||||
|
||||
@pytest.mark.parametrize("asynchronous", [False, True])
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_file_failures_preserve_identity_and_do_not_send(ocr_server, tmp_path, asynchronous):
|
||||
from litellm.rust_bridge import _native
|
||||
|
||||
server, requests = ocr_server
|
||||
failure: Final = KeyError("reader failed")
|
||||
|
||||
class Reader:
|
||||
def __init__(self) -> None:
|
||||
self.calls = 0
|
||||
|
||||
def read(self) -> bytes:
|
||||
self.calls += 1
|
||||
raise failure
|
||||
|
||||
reader: Final = Reader()
|
||||
path: Final = tmp_path / "missing.pdf"
|
||||
kwargs: Final = {
|
||||
"model": "mistral/mistral-ocr-latest",
|
||||
"api_key": "test-key",
|
||||
"api_base": f"http://127.0.0.1:{server.server_port}",
|
||||
}
|
||||
for source, error in ((reader, KeyError), (path, FileNotFoundError)):
|
||||
with pytest.raises(error) as caught:
|
||||
if asynchronous:
|
||||
await _native.aocr(document={"type": "file", "file": source}, **kwargs)
|
||||
else:
|
||||
_native.ocr(document={"type": "file", "file": source}, **kwargs)
|
||||
if source is reader:
|
||||
assert caught.value is failure
|
||||
else:
|
||||
assert caught.value.filename == str(path)
|
||||
assert reader.calls == 1
|
||||
assert requests == []
|
||||
|
||||
|
||||
def test_unstarted_native_file_call_does_not_read_and_releases_reader():
|
||||
from litellm.rust_bridge import _native
|
||||
|
||||
class Reader:
|
||||
def read(self) -> bytes:
|
||||
raise AssertionError("unstarted call consumed its document")
|
||||
|
||||
def create_call():
|
||||
reader: Final = Reader()
|
||||
pending: Final = _native.aocr(
|
||||
model="mistral/mistral-ocr-latest", document={"type": "file", "file": reader}
|
||||
)
|
||||
reader.pending = pending
|
||||
return pending, weakref.ref(reader)
|
||||
|
||||
pending, reference = create_call()
|
||||
pending.close()
|
||||
del pending
|
||||
gc.collect()
|
||||
assert reference() is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize("asynchronous", [False, True])
|
||||
@pytest.mark.parametrize("model", ["mistral/mistral-ocr-latest", "azure_ai/doc-intelligence/prebuilt-read"])
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue