mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
refactor callback
This commit is contained in:
parent
75548d30f5
commit
1771e5bb68
29 changed files with 799 additions and 472 deletions
|
|
@ -223,9 +223,11 @@ mod tests {
|
|||
] {
|
||||
assert_eq!(lifecycle.phase(), phase);
|
||||
assert!(
|
||||
lifecycle.accept(Err(HostFailure::Error(crate::ocr::Error::InvalidRequest(
|
||||
"callback".into()
|
||||
)))).is_none()
|
||||
lifecycle
|
||||
.accept(Err(HostFailure::Error(crate::ocr::Error::InvalidRequest(
|
||||
"callback".into()
|
||||
))))
|
||||
.is_none()
|
||||
);
|
||||
}
|
||||
assert_eq!(lifecycle.phase(), HostPhase::Complete);
|
||||
|
|
|
|||
|
|
@ -816,7 +816,8 @@ mod tests {
|
|||
"type":"document_url",
|
||||
"document_url":"https://example.com/document.pdf"
|
||||
}))
|
||||
.unwrap().into();
|
||||
.unwrap()
|
||||
.into();
|
||||
|
||||
perform_ocr(request).await.unwrap();
|
||||
server.await.unwrap();
|
||||
|
|
|
|||
|
|
@ -361,10 +361,12 @@ mod tests {
|
|||
}
|
||||
}),
|
||||
);
|
||||
let request = request.with_document(serde_json::from_value(json!({
|
||||
"type":"image_url","image_url":"https://example.com/original.png"
|
||||
}))
|
||||
.unwrap());
|
||||
let request = request.with_document(
|
||||
serde_json::from_value(json!({
|
||||
"type":"image_url","image_url":"https://example.com/original.png"
|
||||
}))
|
||||
.unwrap(),
|
||||
);
|
||||
let request = crate::ocr::prepare::prepare_request(request);
|
||||
let http = CohereParseConfig
|
||||
.prepare_request(&request, &crate::ocr::test_support::ocr_client())
|
||||
|
|
@ -503,10 +505,12 @@ mod tests {
|
|||
"https://example.com",
|
||||
json!({"output_format":null,"req_format":null}),
|
||||
);
|
||||
let request = request.with_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());
|
||||
.unwrap(),
|
||||
);
|
||||
assert_eq!(
|
||||
request.response_format().unwrap(),
|
||||
crate::ocr::types::OcrResponseFormat::Litellm
|
||||
|
|
|
|||
|
|
@ -738,7 +738,8 @@ mod tests {
|
|||
"result":{"chunks":[]}
|
||||
}))])
|
||||
.await;
|
||||
let request = crate::ocr::test_support::with_source(wire_request(model, &base, options), source);
|
||||
let request =
|
||||
crate::ocr::test_support::with_source(wire_request(model, &base, options), source);
|
||||
|
||||
perform_ocr(request).await.unwrap();
|
||||
server.await.unwrap();
|
||||
|
|
@ -857,7 +858,9 @@ mod tests {
|
|||
#[tokio::test]
|
||||
async fn rejects_invalid_document_sources_before_network(#[case] source: &str) {
|
||||
let request = crate::ocr::test_support::with_source(
|
||||
wire_request("reducto/parse-v3", "http://127.0.0.1:1", json!({})), source);
|
||||
wire_request("reducto/parse-v3", "http://127.0.0.1:1", json!({})),
|
||||
source,
|
||||
);
|
||||
assert!(perform_ocr(request).await.is_err());
|
||||
}
|
||||
|
||||
|
|
@ -913,7 +916,9 @@ mod tests {
|
|||
let raw = json!({"job_id":"job-1","result":{"chunks":[]}});
|
||||
let (base, seen, server) = mock_server(vec![MockResponse::json(raw)]).await;
|
||||
let mut request = crate::ocr::test_support::with_source(
|
||||
wire_request("reducto/parse-v3", &base, json!({})), "reducto://ready.pdf");
|
||||
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();
|
||||
|
|
@ -984,6 +989,7 @@ mod tests {
|
|||
request: OcrDuringCallRequest,
|
||||
) -> OcrHookFuture<'_, OcrDuringCallRequest> {
|
||||
Box::pin(async move {
|
||||
assert_eq!(request.optional_params["use_cache"], json!(true));
|
||||
assert_eq!(
|
||||
request.body["document_url"],
|
||||
"data:application/pdf;base64,YWJj"
|
||||
|
|
@ -1000,7 +1006,7 @@ mod tests {
|
|||
async fn guardrail_rewrites_document_before_upload() {
|
||||
let (base, seen, server) =
|
||||
mock_server(vec![MockResponse::json(json!({"result":{"chunks":[]}}))]).await;
|
||||
let mut request = wire_request("reducto/parse-v3", &base, json!({}));
|
||||
let mut request = wire_request("reducto/parse-v3", &base, json!({"use_cache":true}));
|
||||
request.hooks = Arc::new(RewriteDocument);
|
||||
|
||||
perform_ocr(request).await.unwrap();
|
||||
|
|
|
|||
|
|
@ -338,8 +338,12 @@ mod tests {
|
|||
options.clone(),
|
||||
);
|
||||
let vertex = wire_request("vertex_ai/mistral-ocr-maas", "https://vertex.test", options);
|
||||
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 = 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
|
||||
|
|
|
|||
|
|
@ -39,9 +39,10 @@ impl OcrClient {
|
|||
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
|
||||
use super::{
|
||||
NativeOutcome, OcrAdmission, OcrCall, OcrCallStep, OcrHookHost, OcrHost,
|
||||
OcrHostOperation, OcrHostResult,
|
||||
OcrHostOperation, OcrHostResult, OcrProjectedRequest,
|
||||
};
|
||||
|
||||
let intercepts_requests = request.hooks.intercepts_requests();
|
||||
let host = OcrHookHost::new(request.hooks.clone());
|
||||
let mut request = Some(request);
|
||||
let NativeOutcome::Completed(mut call) = OcrCall::admit(self.clone(), OcrAdmission::all())
|
||||
|
|
@ -54,14 +55,15 @@ impl OcrClient {
|
|||
loop {
|
||||
match call.resume(result.take()).await? {
|
||||
OcrCallStep::Host(OcrHostOperation::ProjectRequest) => {
|
||||
result = Some(OcrHostResult::Request(Ok((
|
||||
Box::new(request.take().ok_or_else(|| {
|
||||
result = Some(OcrHostResult::Request(Ok(OcrProjectedRequest {
|
||||
request: Box::new(request.take().ok_or_else(|| {
|
||||
crate::ocr::Error::InvalidRequest(
|
||||
"OCR request was already projected".into(),
|
||||
)
|
||||
})?),
|
||||
false,
|
||||
))))
|
||||
intercepts_requests,
|
||||
host_token_provider: false,
|
||||
})))
|
||||
}
|
||||
OcrCallStep::Host(operation) => result = Some(host.invoke(operation).await),
|
||||
OcrCallStep::Complete(response) => return Ok(response),
|
||||
|
|
|
|||
|
|
@ -15,7 +15,10 @@ pub(crate) fn read_path_document(
|
|||
) -> 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))
|
||||
.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),
|
||||
|
|
@ -201,9 +204,14 @@ mod tests {
|
|||
|
||||
#[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 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 {
|
||||
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);
|
||||
|
|
@ -215,9 +223,16 @@ mod tests {
|
|||
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="));
|
||||
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 {
|
||||
|
|
|
|||
|
|
@ -8,6 +8,8 @@ pub enum Error {
|
|||
},
|
||||
#[error("File is empty or could not be read")]
|
||||
EmptyFile,
|
||||
#[error("Host OCR document read failed")]
|
||||
HostDocumentRead,
|
||||
#[error("Failed to read OCR file {}: {source}", path.display())]
|
||||
FileRead {
|
||||
path: std::path::PathBuf,
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ use std::sync::Arc;
|
|||
|
||||
use super::OcrClient;
|
||||
use super::hooks::{OcrHooks, OcrLifecycleHooks, OcrPostCallRequest};
|
||||
use super::types::{ResolvedOcrRequest, LiteLLMOcrResponse, PreparedOcrRequest};
|
||||
use super::types::{LiteLLMOcrResponse, PreparedOcrRequest, ResolvedOcrRequest};
|
||||
use crate::call_lifecycle::{CallLifecycle, CallLifecycleContext};
|
||||
use crate::llms::base_llm::ocr::transformation::OcrResponseContext;
|
||||
|
||||
|
|
|
|||
|
|
@ -22,6 +22,7 @@ pub struct OcrPreCallRequest {
|
|||
pub struct OcrDuringCallRequest {
|
||||
pub model: String,
|
||||
pub custom_llm_provider: String,
|
||||
pub optional_params: Value,
|
||||
pub api_key: Option<String>,
|
||||
pub url: String,
|
||||
pub headers: Vec<(String, String)>,
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
use std::future::Future;
|
||||
use std::pin::Pin;
|
||||
use std::sync::Arc;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
use tokio::sync::{mpsc, oneshot};
|
||||
|
||||
|
|
@ -9,8 +10,8 @@ use super::hooks::{
|
|||
OcrDuringCallRequest, OcrHookFuture, OcrHooks, OcrLogFuture, OcrPostCallRequest,
|
||||
OcrPreCallRequest,
|
||||
};
|
||||
use super::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrClient, OcrDocumentInput, OcrFileContent};
|
||||
use super::types::ResolvedOcrRequest;
|
||||
use super::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrClient, OcrDocumentInput, OcrFileContent};
|
||||
use crate::call_lifecycle::host::{
|
||||
HostCall, HostCallFuture, HostCallStep, HostFailure, HostLifecycle, HostPhase,
|
||||
};
|
||||
|
|
@ -82,8 +83,14 @@ impl OcrHostOperation {
|
|||
}
|
||||
}
|
||||
|
||||
pub struct OcrProjectedRequest {
|
||||
pub request: Box<LiteLLMOcrRequest>,
|
||||
pub intercepts_requests: bool,
|
||||
pub host_token_provider: bool,
|
||||
}
|
||||
|
||||
pub enum OcrHostResult {
|
||||
Request(Result<(Box<LiteLLMOcrRequest>, bool), super::Error>),
|
||||
Request(Result<OcrProjectedRequest, super::Error>),
|
||||
Document(Result<OcrFileContent, super::Error>),
|
||||
Lifecycle(Result<(), HostFailure<super::Error>>),
|
||||
AzureAdToken(Result<ResolvedCredential, litellm_auth::Error>),
|
||||
|
|
@ -160,9 +167,10 @@ impl OcrCall {
|
|||
Some(OcrHostResult::Request(result)) if self.projecting => {
|
||||
self.projecting = false;
|
||||
match result {
|
||||
Ok((request, azure_ad_token_provider)) => {
|
||||
self.execution.request = Some(*request);
|
||||
self.execution.azure_ad_token_provider = azure_ad_token_provider;
|
||||
Ok(projected) => {
|
||||
self.execution.set_request(*projected.request);
|
||||
self.execution.intercepts_requests = projected.intercepts_requests;
|
||||
self.execution.host_token_provider = projected.host_token_provider;
|
||||
}
|
||||
Err(error) => self.accept(Err(HostFailure::Error(error))),
|
||||
}
|
||||
|
|
@ -268,6 +276,15 @@ impl OcrCall {
|
|||
fn accept(&mut self, result: Result<(), HostFailure<super::Error>>) {
|
||||
let cancelled = matches!(&result, Err(HostFailure::Cancelled(_)));
|
||||
if let Some(error) = self.lifecycle.accept(result) {
|
||||
if let Some((_, timing)) = self
|
||||
.execution
|
||||
.terminal
|
||||
.lock()
|
||||
.unwrap_or_else(|error| error.into_inner())
|
||||
.as_mut()
|
||||
{
|
||||
timing.end_time = epoch_seconds();
|
||||
}
|
||||
if cancelled {
|
||||
self.error = Some(error);
|
||||
} else {
|
||||
|
|
@ -333,11 +350,33 @@ struct OcrExecution {
|
|||
pending_result: Option<oneshot::Sender<OcrHostResult>>,
|
||||
execution: Option<tokio::task::JoinHandle<Result<LiteLLMOcrResponse, super::Error>>>,
|
||||
completed: bool,
|
||||
azure_ad_token_provider: bool,
|
||||
intercepts_requests: bool,
|
||||
host_token_provider: bool,
|
||||
terminal: Arc<std::sync::Mutex<Option<(CallLifecycleContext, CallLifecycleTiming)>>>,
|
||||
}
|
||||
|
||||
impl OcrExecution {
|
||||
fn set_request(&mut self, mut request: LiteLLMOcrRequest) {
|
||||
let call_id = request
|
||||
.litellm_call_id
|
||||
.clone()
|
||||
.unwrap_or_else(|| format!("ocr-{:032x}", rand::random::<u128>()));
|
||||
let context = CallLifecycleContext::new(
|
||||
"ocr",
|
||||
request.model.clone(),
|
||||
request.provider_name(),
|
||||
call_id.clone(),
|
||||
);
|
||||
request.litellm_call_id = Some(call_id);
|
||||
let start = epoch_seconds();
|
||||
*self
|
||||
.terminal
|
||||
.lock()
|
||||
.unwrap_or_else(|error| error.into_inner()) =
|
||||
Some((context, CallLifecycleTiming::new(start, start, Vec::new())));
|
||||
self.request = Some(request);
|
||||
}
|
||||
|
||||
fn new(client: OcrClient) -> Self {
|
||||
let (operations_tx, operations_rx) = mpsc::unbounded_channel();
|
||||
Self {
|
||||
|
|
@ -350,7 +389,8 @@ impl OcrExecution {
|
|||
pending_result: None,
|
||||
execution: None,
|
||||
completed: false,
|
||||
azure_ad_token_provider: false,
|
||||
intercepts_requests: false,
|
||||
host_token_provider: false,
|
||||
terminal: Arc::default(),
|
||||
}
|
||||
}
|
||||
|
|
@ -366,11 +406,16 @@ impl OcrExecution {
|
|||
}
|
||||
let result = if self.reading {
|
||||
let Some(OcrHostResult::Document(content)) = result else {
|
||||
return Err(super::Error::InvalidRequest("OCR document read result is required".into()));
|
||||
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 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");
|
||||
};
|
||||
|
|
@ -428,7 +473,10 @@ impl OcrExecution {
|
|||
|
||||
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");
|
||||
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;
|
||||
|
|
@ -444,15 +492,25 @@ impl OcrExecution {
|
|||
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())?
|
||||
}
|
||||
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;
|
||||
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);
|
||||
|
|
@ -461,8 +519,7 @@ impl OcrExecution {
|
|||
|
||||
fn start(&mut self, mut request: ResolvedOcrRequest) {
|
||||
let client = self.client.take().expect("admitted OCR call has a client");
|
||||
let intercepts_requests = request.hooks.intercepts_requests();
|
||||
if self.azure_ad_token_provider {
|
||||
if self.host_token_provider {
|
||||
request.azure_ad_token_provider = Some(TokenProviderHandle::new(Arc::new(
|
||||
OcrAzureAdTokenProvider {
|
||||
operations: self.operations_tx.clone(),
|
||||
|
|
@ -471,7 +528,7 @@ impl OcrExecution {
|
|||
}
|
||||
request.hooks = Arc::new(ProtocolHooks {
|
||||
operations: self.operations_tx.clone(),
|
||||
intercepts_requests,
|
||||
intercepts_requests: self.intercepts_requests,
|
||||
terminal: self.terminal.clone(),
|
||||
});
|
||||
self.execution = Some(tokio::spawn(async move {
|
||||
|
|
@ -517,6 +574,13 @@ struct ProtocolHooks {
|
|||
terminal: Arc<std::sync::Mutex<Option<(CallLifecycleContext, CallLifecycleTiming)>>>,
|
||||
}
|
||||
|
||||
fn epoch_seconds() -> f64 {
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
.as_secs_f64()
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct OcrAzureAdTokenProvider {
|
||||
operations: mpsc::UnboundedSender<PendingOperation>,
|
||||
|
|
@ -743,36 +807,60 @@ mod tests {
|
|||
use crate::ocr::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
|
||||
use crate::ocr::{
|
||||
LiteLLMOcrRequest, NativeOutcome, NoopOcrHost, OcrAdmission, OcrCall, OcrCallStep,
|
||||
OcrDecline, OcrDocument, OcrHost, OcrHostOperation, OcrHostResult,
|
||||
OcrDecline, OcrDocument, OcrHost, OcrHostOperation, OcrHostResult, OcrProjectedRequest,
|
||||
};
|
||||
|
||||
fn projected_request(request: LiteLLMOcrRequest) -> OcrHostResult {
|
||||
OcrHostResult::Request(Ok(OcrProjectedRequest {
|
||||
intercepts_requests: request.hooks.intercepts_requests(),
|
||||
request: Box::new(request),
|
||||
host_token_provider: false,
|
||||
}))
|
||||
}
|
||||
|
||||
#[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 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 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);
|
||||
"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);
|
||||
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();
|
||||
let body: Value =
|
||||
serde_json::from_str(requests[0].split_once("\r\n\r\n").unwrap().1).unwrap();
|
||||
assert_eq!(body["document"], expected);
|
||||
}
|
||||
|
||||
|
|
@ -781,42 +869,80 @@ mod tests {
|
|||
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);
|
||||
.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 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,
|
||||
});
|
||||
while let OcrCallStep::Host(operation) = call.resume(result.take()).await.unwrap() {
|
||||
result = Some(match operation {
|
||||
OcrHostOperation::ProjectRequest => {
|
||||
assert_eq!(reads, 0);
|
||||
projected_request(request.take().unwrap())
|
||||
}
|
||||
OcrCallStep::Complete(_) => break,
|
||||
}
|
||||
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,
|
||||
});
|
||||
}
|
||||
server.await.unwrap();
|
||||
assert_eq!((reads, pre_calls, seen.lock().unwrap().len()), (1, 1, 1));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn sdk_file_failure_dispatches_the_typed_error_once() {
|
||||
struct CaptureFailure(Arc<Mutex<Vec<crate::ocr::Error>>>);
|
||||
impl OcrHooks for CaptureFailure {
|
||||
fn failure<'a>(
|
||||
&'a self,
|
||||
_: &'a CallLifecycleContext,
|
||||
error: &'a crate::ocr::Error,
|
||||
_: &'a CallLifecycleTiming,
|
||||
) -> OcrLogFuture<'a> {
|
||||
self.0.lock().unwrap().push(error.clone());
|
||||
Box::pin(async {})
|
||||
}
|
||||
}
|
||||
let path =
|
||||
std::env::temp_dir().join(format!("missing-ocr-{:032x}.pdf", rand::random::<u128>()));
|
||||
let failures = Arc::new(Mutex::new(Vec::new()));
|
||||
let request = wire_request("mistral/model", "http://127.0.0.1:1", json!({}))
|
||||
.with_document(path.clone().into())
|
||||
.with_host_hooks(Arc::new(CaptureFailure(failures.clone())), None);
|
||||
let error = perform_ocr(request).await.unwrap_err();
|
||||
assert!(
|
||||
matches!(error, crate::ocr::Error::FileRead { path: actual, .. } if actual == path)
|
||||
);
|
||||
let failures = failures.lock().unwrap();
|
||||
assert_eq!(failures.len(), 1);
|
||||
assert!(
|
||||
matches!(&failures[0], crate::ocr::Error::FileRead { source, .. } if source.kind() == std::io::ErrorKind::NotFound)
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn cancellation_acknowledges_blocking_preparation_completion() {
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
|
|
@ -836,7 +962,8 @@ mod tests {
|
|||
std::future::poll_fn(|cx| {
|
||||
assert!(stop.as_mut().poll(cx).is_pending());
|
||||
std::task::Poll::Ready(())
|
||||
}).await;
|
||||
})
|
||||
.await;
|
||||
assert!(!finished.load(Ordering::SeqCst));
|
||||
release_tx.send(()).unwrap();
|
||||
stop.await;
|
||||
|
|
@ -1034,6 +1161,8 @@ mod tests {
|
|||
mut request: OcrDuringCallRequest,
|
||||
) -> OcrHookFuture<'_, OcrDuringCallRequest> {
|
||||
Box::pin(async move {
|
||||
assert_eq!(request.optional_params["pages"], json!([0]));
|
||||
assert_eq!(request.optional_params["extra_body"]["pages"], json!([2]));
|
||||
assert_eq!(request.body["pages"], json!([2]));
|
||||
assert_eq!(request.body.get("future"), Some(&Value::Null));
|
||||
request.body.as_object_mut().unwrap().remove("future");
|
||||
|
|
@ -1287,10 +1416,7 @@ mod tests {
|
|||
result = Some(OcrHostResult::Lifecycle(Ok(())))
|
||||
}
|
||||
OcrHostOperation::ProjectRequest => {
|
||||
result = Some(OcrHostResult::Request(Ok((
|
||||
Box::new(request.take().unwrap()),
|
||||
false,
|
||||
))))
|
||||
result = Some(projected_request(request.take().unwrap()))
|
||||
}
|
||||
OcrHostOperation::AcquireAzureAdToken => {
|
||||
panic!("test request has no token provider")
|
||||
|
|
@ -1312,7 +1438,9 @@ mod tests {
|
|||
Ok(request)
|
||||
}));
|
||||
}
|
||||
OcrHostOperation::PostCall(_) | OcrHostOperation::ReadDocument => 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"),
|
||||
|
|
@ -1345,10 +1473,7 @@ mod tests {
|
|||
let error = loop {
|
||||
match call.resume(result.take()).await {
|
||||
Ok(OcrCallStep::Host(OcrHostOperation::ProjectRequest)) => {
|
||||
result = Some(OcrHostResult::Request(Ok((
|
||||
Box::new(request.take().unwrap()),
|
||||
false,
|
||||
))));
|
||||
result = Some(projected_request(request.take().unwrap()));
|
||||
}
|
||||
Ok(OcrCallStep::Host(operation)) => {
|
||||
if let OcrHostOperation::PostCall(request) = &operation {
|
||||
|
|
@ -1412,7 +1537,7 @@ mod tests {
|
|||
});
|
||||
result = Some(match operation {
|
||||
OcrHostOperation::ProjectRequest => {
|
||||
OcrHostResult::Request(Ok((Box::new(request.take().unwrap()), false)))
|
||||
projected_request(request.take().unwrap())
|
||||
}
|
||||
operation => host.invoke(operation).await,
|
||||
});
|
||||
|
|
@ -1472,7 +1597,9 @@ mod tests {
|
|||
OcrHostResult::Lifecycle(Err(HostFailure::Error(selected.clone())))
|
||||
}
|
||||
OcrHostOperation::Failure { error, .. } => {
|
||||
assert!(matches!(error, crate::ocr::Error::InvalidRequest(message) if message == "public metadata failed"));
|
||||
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()),
|
||||
|
|
@ -1488,7 +1615,7 @@ mod tests {
|
|||
panic!("finalization failure used provider/success dispatch")
|
||||
}
|
||||
OcrHostOperation::ProjectRequest => {
|
||||
OcrHostResult::Request(Ok((Box::new(request.take().unwrap()), false)))
|
||||
projected_request(request.take().unwrap())
|
||||
}
|
||||
operation => host.invoke(operation).await,
|
||||
});
|
||||
|
|
@ -1498,7 +1625,9 @@ mod tests {
|
|||
}
|
||||
};
|
||||
server.await.unwrap();
|
||||
assert!(matches!(error, crate::ocr::Error::InvalidRequest(message) if message == "public metadata failed"));
|
||||
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);
|
||||
}
|
||||
|
|
@ -1525,10 +1654,7 @@ mod tests {
|
|||
match call.resume(result.take()).await.unwrap() {
|
||||
OcrCallStep::Host(OcrHostOperation::PreCall(_)) => break,
|
||||
OcrCallStep::Host(OcrHostOperation::ProjectRequest) => {
|
||||
result = Some(OcrHostResult::Request(Ok((
|
||||
Box::new(request.take().unwrap()),
|
||||
false,
|
||||
))))
|
||||
result = Some(projected_request(request.take().unwrap()))
|
||||
}
|
||||
OcrCallStep::Host(operation) => result = Some(host.invoke(operation).await),
|
||||
OcrCallStep::Complete(_) => panic!("provider executed before pre-call result"),
|
||||
|
|
@ -1751,7 +1877,7 @@ mod tests {
|
|||
_ = entered.notified() => break,
|
||||
step = call.resume(result.take()) => {
|
||||
result = Some(match step.unwrap() {
|
||||
OcrCallStep::Host(OcrHostOperation::ProjectRequest) => OcrHostResult::Request(Ok((Box::new(request.take().unwrap()), false))),
|
||||
OcrCallStep::Host(OcrHostOperation::ProjectRequest) => projected_request(request.take().unwrap()),
|
||||
OcrCallStep::Host(operation) => NoopOcrHost.invoke(operation).await,
|
||||
OcrCallStep::Complete(_) => panic!("pending provider completed"),
|
||||
});
|
||||
|
|
@ -1778,7 +1904,9 @@ mod tests {
|
|||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(matches!(result, Err(crate::ocr::Error::InvalidRequest(message)) if message == "cancelled"));
|
||||
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"
|
||||
|
|
|
|||
|
|
@ -18,12 +18,13 @@ pub use client::{OcrClient, ocr};
|
|||
pub use document::{encode_file_document, mime_type_for_name, upload_mime_type};
|
||||
pub use lifecycle::{
|
||||
NativeOutcome, NativeResult, NoopOcrHost, OcrAdmission, OcrCall, OcrCallStep, OcrDecline,
|
||||
OcrHookHost, OcrHost, OcrHostOperation, OcrHostResult,
|
||||
OcrHookHost, OcrHost, OcrHostOperation, OcrHostResult, OcrProjectedRequest,
|
||||
};
|
||||
pub use provider_config::{get_api_key_env_var, get_health_check_document};
|
||||
pub use types::{
|
||||
LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrConnectionInputs, OcrCredentialInputs,
|
||||
OcrDocument, OcrDocumentInput, OcrFileContent, OcrPage, OcrPageDimensions, OcrPageImage, OcrTransportConfig, OcrUsageInfo,
|
||||
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::{ResolvedOcrRequest, OcrConnection, OcrDocument, PreparedOcrRequest};
|
||||
use super::types::{OcrConnection, OcrDocument, PreparedOcrRequest, ResolvedOcrRequest};
|
||||
|
||||
pub(crate) async fn transform_request_body<B>(
|
||||
client: &OcrClient,
|
||||
|
|
@ -28,6 +28,7 @@ where
|
|||
.during_call(OcrDuringCallRequest {
|
||||
model: request.model.clone(),
|
||||
custom_llm_provider: request.provider_name().into(),
|
||||
optional_params: Value::Object(request.optional_params.clone().into()),
|
||||
api_key: request.connection.api_key.clone(),
|
||||
url: url.into(),
|
||||
headers: headers.to_vec(),
|
||||
|
|
@ -78,6 +79,7 @@ pub(crate) async fn guardrail_document(
|
|||
.during_call(OcrDuringCallRequest {
|
||||
model: request.model.clone(),
|
||||
custom_llm_provider: request.provider_name().into(),
|
||||
optional_params: Value::Object(request.optional_params.clone().into()),
|
||||
api_key: request.connection.api_key.clone(),
|
||||
url: url.into(),
|
||||
headers: headers.to_vec(),
|
||||
|
|
|
|||
|
|
@ -3,8 +3,8 @@ use std::path::PathBuf;
|
|||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use bytes::Bytes;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
use serde_with::serde_as;
|
||||
|
||||
|
|
@ -93,7 +93,10 @@ impl From<OcrDocument> for OcrDocumentInput {
|
|||
|
||||
impl From<PathBuf> for OcrDocumentInput {
|
||||
fn from(path: PathBuf) -> Self {
|
||||
Self::Path { path, mime_type: None }
|
||||
Self::Path {
|
||||
path,
|
||||
mime_type: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -324,7 +327,6 @@ impl LiteLLMOcrRequest {
|
|||
config,
|
||||
})
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
impl<D> LiteLLMOcrRequest<D> {
|
||||
|
|
@ -383,7 +385,6 @@ impl<D> LiteLLMOcrRequest<D> {
|
|||
..self
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
impl LiteLLMOcrRequest {
|
||||
|
|
|
|||
|
|
@ -25,17 +25,12 @@ use bindings::DeploymentHooks;
|
|||
pub(crate) use bindings::PythonLogger;
|
||||
use handle::{Execution, ExecutionBody, ExecutionStep};
|
||||
|
||||
pub(crate) enum OperationClass {
|
||||
Phase(HostPhase),
|
||||
Route,
|
||||
}
|
||||
|
||||
pub(crate) trait PythonRoute: Send + Sync {
|
||||
type Call: NativeCall + 'static;
|
||||
|
||||
fn state(&self) -> &PythonCallState;
|
||||
fn state_mut(&mut self) -> &mut PythonCallState;
|
||||
fn classify(operation: &<Self::Call as NativeCall>::Operation) -> OperationClass;
|
||||
fn phase(operation: &<Self::Call as NativeCall>::Operation) -> Option<HostPhase>;
|
||||
fn lifecycle_result() -> <Self::Call as NativeCall>::Result;
|
||||
fn map_error(error: <Self::Call as NativeCall>::Error) -> PyErr;
|
||||
fn host_error(message: String) -> <Self::Call as NativeCall>::Error;
|
||||
|
|
@ -167,12 +162,8 @@ impl<R: PythonRoute> PythonLifecycle<R> {
|
|||
HostFailure::Cancelled(native)
|
||||
};
|
||||
let state = self.route.state_mut();
|
||||
if state.error.is_none() || (cancelled && phase != Some(HostPhase::DeploymentFailure)) {
|
||||
state.retain_error(py, error);
|
||||
}
|
||||
if state.end.is_none() {
|
||||
state.end = now(py).ok();
|
||||
}
|
||||
state.retain_first_error(py, error, cancelled && phase != Some(HostPhase::DeploymentFailure));
|
||||
let _ = state.finish(py);
|
||||
failure
|
||||
}
|
||||
|
||||
|
|
@ -215,10 +206,7 @@ impl<R: PythonRoute> PythonLifecycle<R> {
|
|||
}
|
||||
HostStep::Ready(NativeCallStep::Host(operation)) => operation,
|
||||
};
|
||||
let phase = match R::classify(&operation) {
|
||||
OperationClass::Phase(phase) => Some(phase),
|
||||
OperationClass::Route => None,
|
||||
};
|
||||
let phase = R::phase(&operation);
|
||||
let result = match phase {
|
||||
Some(phase) => match self.route.state_mut().invoke(py, phase) {
|
||||
Ok(HostStep::Suspend(awaitable)) => {
|
||||
|
|
@ -499,6 +487,19 @@ impl PythonCallState {
|
|||
self.error = Some(error.into_value(py));
|
||||
}
|
||||
|
||||
pub fn retain_first_error(&mut self, py: Python<'_>, error: PyErr, replace: bool) {
|
||||
if self.error.is_none() || replace {
|
||||
self.retain_error(py, error);
|
||||
}
|
||||
}
|
||||
|
||||
pub fn finish(&mut self, py: Python<'_>) -> PyResult<()> {
|
||||
if self.end.is_none() {
|
||||
self.end = Some(now(py)?);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
visit.call(&self.args)?;
|
||||
visit.call(&self.kwargs)?;
|
||||
|
|
@ -694,8 +695,8 @@ mod tests {
|
|||
&mut self.0
|
||||
}
|
||||
|
||||
fn classify(_: &()) -> OperationClass {
|
||||
OperationClass::Route
|
||||
fn phase(_: &()) -> Option<HostPhase> {
|
||||
None
|
||||
}
|
||||
|
||||
fn lifecycle_result() {}
|
||||
|
|
|
|||
|
|
@ -13,6 +13,11 @@ use litellm_python_interop::from_py_preserving_errors as from_py;
|
|||
use crate::auth::{AZURE_AD_TOKEN_PROVIDER, PythonTokenProvider};
|
||||
use crate::lifecycle::BoundArguments;
|
||||
|
||||
pub(crate) struct Projection<Native, Retained> {
|
||||
pub native: Native,
|
||||
pub retained: Retained,
|
||||
}
|
||||
|
||||
/// 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.
|
||||
|
|
|
|||
|
|
@ -3,142 +3,55 @@ use pyo3::types::PyDict;
|
|||
use serde_json::Value;
|
||||
|
||||
use litellm_core::ocr::LiteLLMOcrResponse;
|
||||
use litellm_core::ocr::hooks::OcrPreCallRequest;
|
||||
use litellm_core::ocr::hooks::OcrDuringCallRequest;
|
||||
use litellm_python_interop::to_py_preserving_errors as to_py;
|
||||
|
||||
use super::host::PythonPayload;
|
||||
use crate::lifecycle::PythonLogger;
|
||||
|
||||
pub(super) struct OcrLoggingFields {
|
||||
model: String,
|
||||
custom_llm_provider: String,
|
||||
optional_params: Value,
|
||||
}
|
||||
|
||||
impl From<&OcrPreCallRequest> for OcrLoggingFields {
|
||||
fn from(request: &OcrPreCallRequest) -> Self {
|
||||
Self {
|
||||
model: request.model.clone(),
|
||||
custom_llm_provider: request.custom_llm_provider.clone(),
|
||||
optional_params: request.optional_params.clone(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl PythonLogger {
|
||||
pub(super) fn update_ocr(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
kwargs: &Py<PyDict>,
|
||||
pre_call: &OcrLoggingFields,
|
||||
secret_fields: &[&str],
|
||||
url: &str,
|
||||
) -> PyResult<()> {
|
||||
let update = PyDict::new(py);
|
||||
update.set_item("kwargs", redact(py, kwargs.bind(py), secret_fields)?)?;
|
||||
update.set_item("model", &pre_call.model)?;
|
||||
update.set_item(
|
||||
"optional_params",
|
||||
redact(
|
||||
py,
|
||||
&to_py(py, &pre_call.optional_params)?
|
||||
.into_bound(py)
|
||||
.cast_into::<PyDict>()?,
|
||||
secret_fields,
|
||||
)?,
|
||||
)?;
|
||||
let params = PyDict::new(py);
|
||||
params.set_item(
|
||||
"litellm_call_id",
|
||||
kwargs.bind(py).get_item("litellm_call_id")?,
|
||||
)?;
|
||||
params.set_item("api_base", url)?;
|
||||
for name in ["logger_fn", "litellm_request_debug"] {
|
||||
if let Some(value) = kwargs.bind(py).get_item(name)? {
|
||||
params.set_item(name, value)?;
|
||||
}
|
||||
}
|
||||
for name in custom_pricing_fields(py)? {
|
||||
if let Some(value) = kwargs.bind(py).get_item(&name)?
|
||||
&& !value.is_none()
|
||||
{
|
||||
params.set_item(name, value)?;
|
||||
}
|
||||
}
|
||||
update.set_item("litellm_params", params)?;
|
||||
update.set_item("custom_llm_provider", &pre_call.custom_llm_provider)?;
|
||||
self.object(py)
|
||||
.call_method("update_from_kwargs", (), Some(&update))?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(crate) fn pre_ocr(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
api_key: Option<&str>,
|
||||
body: &Bound<'_, PyDict>,
|
||||
headers: &Bound<'_, PyDict>,
|
||||
url: &str,
|
||||
) -> PyResult<()> {
|
||||
let additional = PyDict::new(py);
|
||||
additional.set_item("complete_input_dict", body)?;
|
||||
additional.set_item("headers", headers)?;
|
||||
additional.set_item("api_base", url)?;
|
||||
let kwargs = PyDict::new(py);
|
||||
kwargs.set_item("input", "OCR document processing")?;
|
||||
kwargs.set_item("api_key", api_key)?;
|
||||
kwargs.set_item("additional_args", &additional)?;
|
||||
self.object(py).call_method("pre_call", (), Some(&kwargs))?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(crate) fn post_ocr(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
original_response: &Value,
|
||||
body: Option<&Py<PyDict>>,
|
||||
headers: Option<&Py<PyDict>>,
|
||||
) -> PyResult<()> {
|
||||
let additional = PyDict::new(py);
|
||||
additional.set_item("complete_input_dict", body)?;
|
||||
additional.set_item("headers", headers)?;
|
||||
let kwargs = PyDict::new(py);
|
||||
kwargs.set_item("original_response", to_py(py, original_response)?)?;
|
||||
kwargs.set_item("additional_args", &additional)?;
|
||||
self.object(py)
|
||||
.call_method("post_call", (), Some(&kwargs))?;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
fn custom_pricing_fields(py: Python<'_>) -> PyResult<Vec<String>> {
|
||||
py.import("litellm.types.utils")?
|
||||
.getattr("CustomPricingLiteLLMParams")?
|
||||
.getattr("model_fields")?
|
||||
.cast_into::<PyDict>()?
|
||||
.keys()
|
||||
.iter()
|
||||
.map(|name| name.extract::<String>())
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn redact(
|
||||
pub(super) fn update_logging(
|
||||
py: Python<'_>,
|
||||
params: &Bound<'_, PyDict>,
|
||||
logger: &PythonLogger,
|
||||
kwargs: &Py<PyDict>,
|
||||
request: &OcrDuringCallRequest,
|
||||
secret_fields: &[&str],
|
||||
) -> PyResult<Py<PyDict>> {
|
||||
let redacted = PyDict::new(py);
|
||||
for (name, value) in params {
|
||||
let name = name.extract::<String>()?;
|
||||
if name == "proxy_server_request" {
|
||||
continue;
|
||||
}
|
||||
if secret_fields.contains(&name.as_str()) {
|
||||
redacted.set_item(name, "****")?;
|
||||
} else {
|
||||
redacted.set_item(name, value)?;
|
||||
}
|
||||
}
|
||||
Ok(redacted.unbind())
|
||||
) -> PyResult<()> {
|
||||
py.import("litellm.rust_bridge.ocr")?
|
||||
.getattr("update_logging")?
|
||||
.call1((
|
||||
logger.object(py),
|
||||
kwargs,
|
||||
&request.model,
|
||||
&request.custom_llm_provider,
|
||||
to_py(py, &request.optional_params)?,
|
||||
secret_fields,
|
||||
&request.url,
|
||||
))?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(super) fn pre_call(
|
||||
py: Python<'_>,
|
||||
logger: &PythonLogger,
|
||||
request: &OcrDuringCallRequest,
|
||||
payload: &PythonPayload,
|
||||
) -> PyResult<()> {
|
||||
py.import("litellm.rust_bridge.ocr")?
|
||||
.getattr("pre_call")?
|
||||
.call1((logger.object(py), request.api_key.as_deref(), &payload.body, &payload.headers, &request.url))?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(super) fn post_call(
|
||||
py: Python<'_>,
|
||||
logger: &PythonLogger,
|
||||
original_response: &Value,
|
||||
payload: &PythonPayload,
|
||||
) -> PyResult<()> {
|
||||
py.import("litellm.rust_bridge.ocr")?
|
||||
.getattr("post_call")?
|
||||
.call1((logger.object(py), to_py(py, original_response)?, &payload.body, &payload.headers))?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(super) fn response(py: Python<'_>, response: &LiteLLMOcrResponse) -> PyResult<Py<PyAny>> {
|
||||
|
|
|
|||
|
|
@ -36,11 +36,11 @@ fn extract_bytes(value: &Bound<'_, PyAny>) -> PyResult<Bytes> {
|
|||
return Ok(Bytes::from_owner(value.extract::<PyBackedBytes>()?));
|
||||
}
|
||||
// Bytes subclasses may retain GC edges that a native Bytes owner cannot traverse.
|
||||
Ok(Bytes::copy_from_slice(value.extract::<PyBackedBytes>()?.as_ref()))
|
||||
Ok(Bytes::copy_from_slice(
|
||||
value.extract::<PyBackedBytes>()?.as_ref(),
|
||||
))
|
||||
}
|
||||
|
||||
|
||||
|
||||
pub(super) struct FileDocumentInput {
|
||||
pub input: OcrDocumentInput,
|
||||
pub reader: Option<PythonFileReader>,
|
||||
|
|
@ -56,9 +56,11 @@ 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 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) {
|
||||
missing()
|
||||
|
|
@ -93,7 +95,9 @@ impl FromPyObject<'_, '_> for FileDocumentInput {
|
|||
reader: None,
|
||||
});
|
||||
}
|
||||
let reader = file.getattr_opt("read")?.filter(|value| value.is_callable());
|
||||
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.",
|
||||
|
|
@ -107,7 +111,10 @@ impl FromPyObject<'_, '_> for FileDocumentInput {
|
|||
.transpose()?;
|
||||
Ok(Self {
|
||||
input: OcrDocumentInput::HostReader { mime_type },
|
||||
reader: Some(PythonFileReader { reader: reader.unbind(), name }),
|
||||
reader: Some(PythonFileReader {
|
||||
reader: reader.unbind(),
|
||||
name,
|
||||
}),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
|
@ -127,12 +134,21 @@ mod tests {
|
|||
let error = document.extract::<FileDocumentInput>().err().unwrap();
|
||||
assert!(error.is_instance_of::<PyValueError>(py));
|
||||
}
|
||||
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();
|
||||
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 locals = PyDict::new(py);
|
||||
py.run(c"from pathlib import Path
|
||||
py.run(
|
||||
c"from pathlib import Path
|
||||
failure = KeyError('reader failed')
|
||||
class Reader:
|
||||
def __init__(self):
|
||||
|
|
@ -143,12 +159,31 @@ class Reader:
|
|||
reader = Reader()
|
||||
document = {'file': reader}
|
||||
path_document = {'file': Path('/nonexistent/ocr-projection-test.pdf')}
|
||||
", Some(&locals), Some(&locals)).unwrap();
|
||||
",
|
||||
Some(&locals),
|
||||
Some(&locals),
|
||||
)
|
||||
.unwrap();
|
||||
let document = locals.get_item("document").unwrap().unwrap();
|
||||
let input: FileDocumentInput = document.extract().unwrap();
|
||||
assert_eq!(locals.get_item("reader").unwrap().unwrap().getattr("reads").unwrap().extract::<usize>().unwrap(), 0);
|
||||
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();
|
||||
let path: FileDocumentInput = locals
|
||||
.get_item("path_document")
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.extract()
|
||||
.unwrap();
|
||||
assert!(matches!(path.input, OcrDocumentInput::Path { .. }));
|
||||
});
|
||||
}
|
||||
|
|
|
|||
|
|
@ -103,9 +103,13 @@ pub(super) fn public_exception(
|
|||
std::io::ErrorKind::PermissionDenied => 13,
|
||||
_ => 5,
|
||||
});
|
||||
return Ok(PyErr::from_value(py.import("builtins")?.getattr("OSError")?.call1((
|
||||
errno, source.to_string(), path,
|
||||
))?));
|
||||
return Ok(PyErr::from_value(
|
||||
py.import("builtins")?.getattr("OSError")?.call1((
|
||||
errno,
|
||||
source.to_string(),
|
||||
path.into_pyobject(py)?.call_method0("__fspath__")?,
|
||||
))?,
|
||||
));
|
||||
}
|
||||
raise_public(py, classify(error), model, provider, None)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,12 +1,11 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use pyo3::gc::{PyTraverseError, PyVisit};
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::types::PyDict;
|
||||
|
||||
use litellm_auth::ResolvedCredential;
|
||||
use litellm_core::call_lifecycle::host::HostPhase;
|
||||
use litellm_core::ocr::hooks::{OcrDuringCallRequest, OcrPostCallRequest};
|
||||
use litellm_core::ocr::{OcrCall, OcrHostOperation, OcrHostResult};
|
||||
use litellm_core::ocr::{OcrCall, OcrHostOperation, OcrHostResult, OcrProjectedRequest};
|
||||
use litellm_python_interop::{
|
||||
from_py_preserving_errors as from_py, to_py_preserving_errors as to_py,
|
||||
};
|
||||
|
|
@ -14,32 +13,50 @@ use litellm_python_interop::{
|
|||
use super::document::PythonFileReader;
|
||||
use super::errors::to_pyerr as ocr_error_to_pyerr;
|
||||
use super::project::project;
|
||||
use super::{ASYNC_SIGNATURE, SIGNATURE, callbacks, errors};
|
||||
use super::{callbacks, errors};
|
||||
use crate::auth::PythonTokenProvider;
|
||||
use crate::lifecycle::{OperationClass, PythonCallState, PythonRoute, missing_state, now};
|
||||
use crate::lifecycle::{PythonCallState, PythonRoute, Signature, missing_state};
|
||||
use crate::marshal::Projection;
|
||||
|
||||
pub(super) struct PythonOcrHost {
|
||||
state: PythonCallState,
|
||||
projected: Option<ProjectedOcrHost>,
|
||||
signature: &'static Signature,
|
||||
retained: Option<OcrRetained>,
|
||||
}
|
||||
|
||||
struct ProjectedOcrHost {
|
||||
model: String,
|
||||
provider: &'static str,
|
||||
secret_fields: Vec<&'static str>,
|
||||
azure_ad_token_provider: Option<PythonTokenProvider>,
|
||||
pre_call: Option<callbacks::OcrLoggingFields>,
|
||||
payload: Option<CapturedOcrPayload>,
|
||||
reader: Option<PythonFileReader>,
|
||||
reader_failed: bool,
|
||||
pub(super) struct OcrRetained {
|
||||
pub model: String,
|
||||
pub provider: &'static str,
|
||||
pub secret_fields: Vec<&'static str>,
|
||||
pub azure_ad_token_provider: Option<PythonTokenProvider>,
|
||||
pub payload: Option<PythonPayload>,
|
||||
pub reader: Option<PythonFileReader>,
|
||||
}
|
||||
|
||||
struct CapturedOcrPayload {
|
||||
body: Py<PyDict>,
|
||||
headers: Py<PyDict>,
|
||||
pub(super) struct PythonPayload {
|
||||
pub body: Py<PyDict>,
|
||||
pub headers: Py<PyDict>,
|
||||
}
|
||||
|
||||
impl CapturedOcrPayload {
|
||||
impl PythonPayload {
|
||||
fn from_request(py: Python<'_>, request: &OcrDuringCallRequest) -> PyResult<Self> {
|
||||
let body = to_py(py, &request.body)?.into_bound(py).cast_into::<PyDict>()?;
|
||||
let headers = PyDict::new(py);
|
||||
for (name, value) in &request.headers {
|
||||
headers.set_item(name, value)?;
|
||||
}
|
||||
Ok(Self { body: body.unbind(), headers: headers.unbind() })
|
||||
}
|
||||
|
||||
fn write_back(&self, py: Python<'_>, mut request: OcrDuringCallRequest) -> PyResult<OcrDuringCallRequest> {
|
||||
request.body = from_py(self.body.bind(py))?;
|
||||
request.headers = self.headers.bind(py)
|
||||
.iter()
|
||||
.map(|(name, value)| Ok((name.extract::<String>()?, value.extract::<String>()?)))
|
||||
.collect::<PyResult<Vec<_>>>()?;
|
||||
Ok(request)
|
||||
}
|
||||
|
||||
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
visit.call(&self.body)?;
|
||||
visit.call(&self.headers)
|
||||
|
|
@ -47,52 +64,36 @@ impl CapturedOcrPayload {
|
|||
}
|
||||
|
||||
impl PythonOcrHost {
|
||||
pub(super) fn new(state: PythonCallState) -> Self {
|
||||
pub(super) fn new(state: PythonCallState, signature: &'static Signature) -> Self {
|
||||
Self {
|
||||
state,
|
||||
projected: None,
|
||||
signature,
|
||||
retained: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn projected(&self) -> PyResult<&ProjectedOcrHost> {
|
||||
self.projected.as_ref().ok_or_else(missing_state)
|
||||
fn retained(&self) -> PyResult<&OcrRetained> {
|
||||
self.retained.as_ref().ok_or_else(missing_state)
|
||||
}
|
||||
|
||||
fn projected_mut(&mut self) -> PyResult<&mut ProjectedOcrHost> {
|
||||
self.projected.as_mut().ok_or_else(missing_state)
|
||||
fn retained_mut(&mut self) -> PyResult<&mut OcrRetained> {
|
||||
self.retained.as_mut().ok_or_else(missing_state)
|
||||
}
|
||||
|
||||
fn project(&mut self, py: Python<'_>) -> PyResult<OcrHostResult> {
|
||||
let signature = if self.state.asynchronous {
|
||||
&ASYNC_SIGNATURE
|
||||
} else {
|
||||
&SIGNATURE
|
||||
};
|
||||
let arguments = signature.bind(self.state.args.bind(py), self.state.kwargs.bind(py))?;
|
||||
let projected = project(py, &arguments)?;
|
||||
let has_token_provider = projected.azure_ad_token_provider.is_some();
|
||||
self.projected = Some(ProjectedOcrHost {
|
||||
model: projected.request.model.clone(),
|
||||
provider: projected.request.provider_name(),
|
||||
secret_fields: projected.secret_fields,
|
||||
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
|
||||
.with_host_hooks(Arc::new(BridgeOcrHooks), None),
|
||||
),
|
||||
has_token_provider,
|
||||
))))
|
||||
let arguments = self.signature.bind(self.state.args.bind(py), self.state.kwargs.bind(py))?;
|
||||
let Projection { native, retained } = project(py, &arguments)?;
|
||||
let host_token_provider = retained.azure_ad_token_provider.is_some();
|
||||
self.retained = Some(retained);
|
||||
Ok(OcrHostResult::Request(Ok(OcrProjectedRequest {
|
||||
request: Box::new(native),
|
||||
intercepts_requests: true,
|
||||
host_token_provider,
|
||||
})))
|
||||
}
|
||||
|
||||
fn acquire_azure_ad_token(&self, py: Python<'_>) -> PyResult<ResolvedCredential> {
|
||||
self.projected()?
|
||||
self.retained()?
|
||||
.azure_ad_token_provider
|
||||
.as_ref()
|
||||
.ok_or_else(missing_state)?
|
||||
|
|
@ -102,42 +103,21 @@ impl PythonOcrHost {
|
|||
fn during_call(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
mut request: OcrDuringCallRequest,
|
||||
request: OcrDuringCallRequest,
|
||||
) -> PyResult<OcrDuringCallRequest> {
|
||||
let projected = self.projected()?;
|
||||
let pre_call = projected.pre_call.as_ref().ok_or_else(missing_state)?;
|
||||
let retained = self.retained()?;
|
||||
let logger = self.state.logger()?;
|
||||
logger.update_ocr(
|
||||
callbacks::update_logging(
|
||||
py,
|
||||
logger,
|
||||
&self.state.kwargs,
|
||||
pre_call,
|
||||
&projected.secret_fields,
|
||||
&request.url,
|
||||
&request,
|
||||
&retained.secret_fields,
|
||||
)?;
|
||||
let body = to_py(py, &request.body)?
|
||||
.into_bound(py)
|
||||
.cast_into::<PyDict>()?;
|
||||
let headers = PyDict::new(py);
|
||||
for (name, value) in &request.headers {
|
||||
headers.set_item(name, value)?;
|
||||
}
|
||||
logger.pre_ocr(
|
||||
py,
|
||||
request.api_key.as_deref(),
|
||||
&body,
|
||||
&headers,
|
||||
&request.url,
|
||||
)?;
|
||||
request.body = from_py(&body)?;
|
||||
request.headers = headers
|
||||
.iter()
|
||||
.map(|(name, value)| Ok((name.extract::<String>()?, value.extract::<String>()?)))
|
||||
.collect::<PyResult<Vec<_>>>()?;
|
||||
let projected = self.projected_mut()?;
|
||||
projected.payload = Some(CapturedOcrPayload {
|
||||
body: body.unbind(),
|
||||
headers: headers.unbind(),
|
||||
});
|
||||
let payload = PythonPayload::from_request(py, &request)?;
|
||||
callbacks::pre_call(py, logger, &request, &payload)?;
|
||||
let request = payload.write_back(py, request)?;
|
||||
self.retained_mut()?.payload = Some(payload);
|
||||
Ok(request)
|
||||
}
|
||||
|
||||
|
|
@ -146,27 +126,24 @@ impl PythonOcrHost {
|
|||
py: Python<'_>,
|
||||
request: OcrPostCallRequest,
|
||||
) -> PyResult<OcrPostCallRequest> {
|
||||
let projected = self.projected()?;
|
||||
let payload = projected.payload.as_ref();
|
||||
self.state.logger()?.post_ocr(
|
||||
let payload = self.retained()?.payload.as_ref().ok_or_else(missing_state)?;
|
||||
callbacks::post_call(
|
||||
py,
|
||||
self.state.logger()?,
|
||||
&request.original_response,
|
||||
payload.map(|payload| &payload.body),
|
||||
payload.map(|payload| &payload.headers),
|
||||
payload,
|
||||
)?;
|
||||
Ok(request)
|
||||
}
|
||||
|
||||
fn map_failure(&mut self, py: Python<'_>, error: litellm_core::ocr::Error) -> PyResult<()> {
|
||||
if self.state.end.is_none() {
|
||||
self.state.end = Some(now(py)?);
|
||||
}
|
||||
let (model, provider) = match &self.projected {
|
||||
Some(projected) => (projected.model.as_str(), projected.provider),
|
||||
self.state.finish(py)?;
|
||||
let (model, provider) = match &self.retained {
|
||||
Some(retained) => (retained.model.as_str(), retained.provider),
|
||||
None => ("", ""),
|
||||
};
|
||||
let mapped = match self.state.error.take() {
|
||||
Some(host_error) if self.projected.as_ref().is_some_and(|host| host.reader_failed) => {
|
||||
Some(host_error) if matches!(error, litellm_core::ocr::Error::HostDocumentRead) => {
|
||||
PyErr::from_value(host_error.into_bound(py).into_any())
|
||||
}
|
||||
Some(host_error) => errors::public_host_exception(py, &host_error, model, provider)?,
|
||||
|
|
@ -188,10 +165,8 @@ impl PythonRoute for PythonOcrHost {
|
|||
&mut self.state
|
||||
}
|
||||
|
||||
fn classify(operation: &OcrHostOperation) -> OperationClass {
|
||||
operation
|
||||
.phase()
|
||||
.map_or(OperationClass::Route, OperationClass::Phase)
|
||||
fn phase(operation: &OcrHostOperation) -> Option<HostPhase> {
|
||||
operation.phase()
|
||||
}
|
||||
|
||||
fn lifecycle_result() -> OcrHostResult {
|
||||
|
|
@ -210,16 +185,24 @@ impl PythonRoute for PythonOcrHost {
|
|||
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?))
|
||||
let reader = self
|
||||
.retained_mut()?
|
||||
.reader
|
||||
.take()
|
||||
.ok_or_else(missing_state)?;
|
||||
match reader.read(py) {
|
||||
Ok(content) => OcrHostResult::Document(Ok(content)),
|
||||
Err(error) if error.is_instance_of::<pyo3::exceptions::PyException>(py) => {
|
||||
self.state.retain_first_error(py, error, false);
|
||||
OcrHostResult::Document(Err(litellm_core::ocr::Error::HostDocumentRead))
|
||||
}
|
||||
Err(error) => return Err(error),
|
||||
}
|
||||
}
|
||||
OcrHostOperation::AcquireAzureAdToken => {
|
||||
OcrHostResult::AzureAdToken(Ok(self.acquire_azure_ad_token(py)?))
|
||||
}
|
||||
OcrHostOperation::PreCall(request) => {
|
||||
self.projected_mut()?.pre_call = Some((&request).into());
|
||||
OcrHostResult::PreCall(Ok(request))
|
||||
}
|
||||
OcrHostOperation::DuringCall(request) => {
|
||||
|
|
@ -229,7 +212,7 @@ impl PythonRoute for PythonOcrHost {
|
|||
OcrHostResult::PostCall(Ok(self.post_call(py, request)?))
|
||||
}
|
||||
OcrHostOperation::ConstructResponse(response) => {
|
||||
self.state.end = Some(now(py)?);
|
||||
self.state.finish(py)?;
|
||||
self.state.response = Some(callbacks::response(py, response.as_ref())?);
|
||||
OcrHostResult::Lifecycle(Ok(()))
|
||||
}
|
||||
|
|
@ -244,30 +227,22 @@ impl PythonRoute for PythonOcrHost {
|
|||
}
|
||||
|
||||
fn cleanup(&mut self) {
|
||||
self.projected = None;
|
||||
self.retained = None;
|
||||
}
|
||||
|
||||
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
let Some(projected) = &self.projected else {
|
||||
let Some(retained) = &self.retained else {
|
||||
return Ok(());
|
||||
};
|
||||
if let Some(provider) = &projected.azure_ad_token_provider {
|
||||
if let Some(provider) = &retained.azure_ad_token_provider {
|
||||
provider.traverse(visit)?;
|
||||
}
|
||||
if let Some(reader) = &projected.reader {
|
||||
if let Some(reader) = &retained.reader {
|
||||
reader.traverse(visit)?;
|
||||
}
|
||||
if let Some(payload) = &projected.payload {
|
||||
if let Some(payload) = &retained.payload {
|
||||
payload.traverse(visit)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
struct BridgeOcrHooks;
|
||||
|
||||
impl litellm_core::ocr::hooks::OcrHooks for BridgeOcrHooks {
|
||||
fn intercepts_requests(&self) -> bool {
|
||||
true
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -68,7 +68,7 @@ fn call(
|
|||
kwargs.copy()?.unbind(),
|
||||
asynchronous,
|
||||
signature.name,
|
||||
)?);
|
||||
)?, signature);
|
||||
run_call(py, call, host)
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -10,54 +10,34 @@ use litellm_core::ocr::{
|
|||
};
|
||||
use litellm_python_interop::from_py_preserving_errors as from_py;
|
||||
|
||||
use super::document::FileDocumentInput;
|
||||
use super::errors::to_pyerr as ocr_error_to_pyerr;
|
||||
use super::document::{FileDocumentInput, PythonFileReader};
|
||||
use crate::auth::PythonTokenProvider;
|
||||
use super::host::OcrRetained;
|
||||
use crate::lifecycle::BoundArguments;
|
||||
use crate::marshal::BoundRouteInputs;
|
||||
use crate::marshal::{BoundRouteInputs, Projection};
|
||||
|
||||
/// 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>,
|
||||
}
|
||||
|
||||
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> {
|
||||
fn project_document(document: &Bound<'_, PyAny>) -> PyResult<Result<FileDocumentInput, litellm_core::ocr::Error>> {
|
||||
let kind: String = document.get_item("type")?.extract()?;
|
||||
if kind != "file" {
|
||||
return Ok(ProjectedDocument::Url(from_py(document)?));
|
||||
let value: serde_json::Value = from_py(document)?;
|
||||
return Ok(OcrDocument::try_from(value).map(|document| FileDocumentInput {
|
||||
input: document.into(),
|
||||
reader: None,
|
||||
}));
|
||||
}
|
||||
document.extract().map(ProjectedDocument::File)
|
||||
document.extract().map(Ok)
|
||||
}
|
||||
|
||||
/// 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()?;
|
||||
document: Result<FileDocumentInput, litellm_core::ocr::Error>,
|
||||
) -> Result<Projection<LiteLLMOcrRequest, OcrRetained>, litellm_core::ocr::Error> {
|
||||
let document = document?;
|
||||
let BoundRouteInputs {
|
||||
model,
|
||||
custom_llm_provider,
|
||||
|
|
@ -83,18 +63,23 @@ fn build_request(
|
|||
input_sources,
|
||||
},
|
||||
)?;
|
||||
Ok(ProjectedOcrCall {
|
||||
request,
|
||||
azure_ad_token_provider,
|
||||
secret_fields,
|
||||
reader: document.reader,
|
||||
Ok(Projection {
|
||||
retained: OcrRetained {
|
||||
model: request.model.clone(),
|
||||
provider: request.provider_name(),
|
||||
azure_ad_token_provider,
|
||||
secret_fields,
|
||||
reader: document.reader,
|
||||
payload: None,
|
||||
},
|
||||
native: request,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn project(
|
||||
py: Python<'_>,
|
||||
arguments: &BoundArguments<'_>,
|
||||
) -> PyResult<ProjectedOcrCall> {
|
||||
) -> PyResult<Projection<LiteLLMOcrRequest, OcrRetained>> {
|
||||
let model: String = arguments.extract("model")?;
|
||||
let custom_llm_provider: Option<String> = arguments.optional("custom_llm_provider")?;
|
||||
let document = project_document(&arguments.required("document")?)?;
|
||||
|
|
@ -142,7 +127,7 @@ mod tests {
|
|||
)
|
||||
.unwrap();
|
||||
assert!(matches!(
|
||||
project_document(&file).unwrap().into_native().unwrap().input,
|
||||
project_document(&file).unwrap().unwrap().input,
|
||||
litellm_core::ocr::OcrDocumentInput::Bytes { bytes, mime_type, .. }
|
||||
if bytes == b"%PDF-1.4"[..] && mime_type.as_deref() == Some("application/pdf")
|
||||
));
|
||||
|
|
@ -155,7 +140,7 @@ mod tests {
|
|||
)
|
||||
.unwrap();
|
||||
assert!(matches!(
|
||||
project_document(&original).unwrap().into_native().unwrap().input,
|
||||
project_document(&original).unwrap().unwrap().input,
|
||||
litellm_core::ocr::OcrDocumentInput::Document(OcrDocument::DocumentUrl { document_url, .. })
|
||||
if document_url == "https://example.com/a.pdf"
|
||||
));
|
||||
|
|
@ -169,7 +154,10 @@ mod tests {
|
|||
let document = py
|
||||
.eval(c"{'type': 'mystery', 'mystery': 'x'}", None, None)
|
||||
.unwrap();
|
||||
let error = project_document(&document).unwrap().into_native().err().unwrap();
|
||||
let error = project_document(&document)
|
||||
.unwrap()
|
||||
.err()
|
||||
.unwrap();
|
||||
assert!(error.to_string().contains("document"));
|
||||
});
|
||||
}
|
||||
|
|
@ -181,14 +169,16 @@ mod tests {
|
|||
let missing = py.eval(c"{}", None, None).unwrap();
|
||||
assert!(
|
||||
project_document(&missing)
|
||||
.err().unwrap()
|
||||
.err()
|
||||
.unwrap()
|
||||
.is_instance_of::<PyKeyError>(py)
|
||||
);
|
||||
|
||||
let non_string = py.eval(c"{'type': 1}", None, None).unwrap();
|
||||
assert!(
|
||||
project_document(&non_string)
|
||||
.err().unwrap()
|
||||
.err()
|
||||
.unwrap()
|
||||
.is_instance_of::<PyTypeError>(py)
|
||||
);
|
||||
|
||||
|
|
@ -202,8 +192,9 @@ class Document:
|
|||
document = Document()
|
||||
",
|
||||
);
|
||||
let error =
|
||||
project_document(&locals.get_item("document").unwrap().unwrap()).err().unwrap();
|
||||
let error = project_document(&locals.get_item("document").unwrap().unwrap())
|
||||
.err()
|
||||
.unwrap();
|
||||
assert!(
|
||||
error
|
||||
.value(py)
|
||||
|
|
@ -232,8 +223,11 @@ document = Document()
|
|||
",
|
||||
);
|
||||
let document = locals.get_item("document").unwrap().unwrap();
|
||||
let projected = project_document(&document).unwrap().into_native().unwrap();
|
||||
assert!(matches!(projected.input, litellm_core::ocr::OcrDocumentInput::Bytes { .. }));
|
||||
let projected = project_document(&document).unwrap().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"]);
|
||||
});
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ from typing import Final, cast # noqa: TID251 # native binding selects a sync
|
|||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.ocr import legacy
|
||||
from litellm.ocr.legacy import convert_file_document_to_url_document, get_mime_type
|
||||
from litellm.rust_bridge.bindings import native_exception_types
|
||||
from litellm.rust_bridge.bindings import native_decline_types
|
||||
from litellm.rust_bridge.configuration import rust_ocr_enabled
|
||||
from litellm.rust_bridge.ocr import NATIVE_AOCR, NATIVE_OCR
|
||||
|
||||
|
|
@ -19,7 +19,7 @@ def ocr(
|
|||
if native is not None:
|
||||
try:
|
||||
return native(*args, **kwargs)
|
||||
except _decline_types():
|
||||
except native_decline_types():
|
||||
pass
|
||||
fallback: Final = cast( # cast-ok: forward the original call shape through the legacy @client decorator
|
||||
Callable[..., OCRResponse | Coroutine[object, object, OCRResponse]], legacy.ocr
|
||||
|
|
@ -32,14 +32,10 @@ async def aocr(*args: object, **kwargs: object) -> OCRResponse: # kwargs-ok: pr
|
|||
if native is not None:
|
||||
try:
|
||||
return await native(*args, **kwargs)
|
||||
except _decline_types():
|
||||
except native_decline_types():
|
||||
pass
|
||||
fallback: Final = cast( # cast-ok: forward the original call shape through the legacy @client decorator
|
||||
Callable[..., Awaitable[OCRResponse]], legacy.aocr
|
||||
)
|
||||
return await fallback(*args, **kwargs)
|
||||
|
||||
|
||||
def _decline_types() -> tuple[type[BaseException], ...]:
|
||||
exception_types: Final = native_exception_types()
|
||||
return (exception_types[0],) if exception_types is not None else ()
|
||||
|
|
|
|||
|
|
@ -29,10 +29,12 @@ def _build_document_from_upload(
|
|||
filename: str | None,
|
||||
content_type: str | None,
|
||||
) -> dict[str, str]:
|
||||
mime_type: Final = content_type.split(";")[0].strip() if content_type else None
|
||||
if not mime_type or mime_type == "application/octet-stream":
|
||||
if filename:
|
||||
mime_type = get_mime_type(filename)
|
||||
supplied_mime: Final = content_type.split(";")[0].strip() if content_type else None
|
||||
mime_type: Final = (
|
||||
get_mime_type(filename)
|
||||
if filename and (not supplied_mime or supplied_mime == "application/octet-stream")
|
||||
else supplied_mime
|
||||
)
|
||||
|
||||
return convert_file_document_to_url_document(
|
||||
{
|
||||
|
|
|
|||
|
|
@ -47,3 +47,8 @@ def native_exception_types() -> tuple[type[BaseException], type[BaseException]]
|
|||
if not isinstance(declined, type) or not isinstance(upstream, type):
|
||||
return None
|
||||
return declined, upstream
|
||||
|
||||
|
||||
def native_decline_types() -> tuple[type[BaseException], ...]:
|
||||
exceptions: Final = native_exception_types()
|
||||
return (exceptions[0],) if exceptions is not None else ()
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable, Coroutine, Mapping
|
||||
from collections.abc import Callable, Coroutine, Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables
|
||||
|
||||
|
|
@ -8,6 +8,84 @@ from pydantic import TypeAdapter
|
|||
|
||||
from litellm.llms.base_llm.ocr.transformation import PROVIDER_NATIVE_RESPONSE_KEY, OCRResponse
|
||||
from litellm.rust_bridge.bindings import NativeBinding
|
||||
from litellm.types.utils import CustomPricingLiteLLMParams
|
||||
|
||||
|
||||
class OcrLoggingProtocol(Protocol):
|
||||
def update_from_kwargs(
|
||||
self,
|
||||
*,
|
||||
kwargs: dict[str, object],
|
||||
model: str,
|
||||
optional_params: dict[str, object],
|
||||
litellm_params: dict[str, object],
|
||||
custom_llm_provider: str,
|
||||
) -> object: ...
|
||||
|
||||
def pre_call(self, *, input: str, api_key: str | None, additional_args: dict[str, object]) -> object: ...
|
||||
|
||||
def post_call(self, *, original_response: object, additional_args: dict[str, object]) -> object: ...
|
||||
|
||||
|
||||
def _redact(params: Mapping[str, object], secret_fields: Sequence[str]) -> dict[str, object]:
|
||||
return {
|
||||
name: "****" if name in secret_fields else value
|
||||
for name, value in params.items()
|
||||
if name != "proxy_server_request"
|
||||
}
|
||||
|
||||
|
||||
def update_logging(
|
||||
logger: OcrLoggingProtocol,
|
||||
kwargs: Mapping[str, object],
|
||||
model: str,
|
||||
custom_llm_provider: str,
|
||||
optional_params: Mapping[str, object],
|
||||
secret_fields: Sequence[str],
|
||||
url: str,
|
||||
) -> None:
|
||||
logger.update_from_kwargs(
|
||||
kwargs=_redact(kwargs, secret_fields),
|
||||
model=model,
|
||||
optional_params=_redact(optional_params, secret_fields),
|
||||
litellm_params={
|
||||
"litellm_call_id": kwargs.get("litellm_call_id"),
|
||||
"api_base": url,
|
||||
**{name: kwargs[name] for name in ("logger_fn", "litellm_request_debug") if name in kwargs},
|
||||
**{
|
||||
name: kwargs[name]
|
||||
for name in CustomPricingLiteLLMParams.model_fields
|
||||
if name in kwargs and kwargs[name] is not None
|
||||
},
|
||||
},
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
|
||||
def pre_call(
|
||||
logger: OcrLoggingProtocol,
|
||||
api_key: str | None,
|
||||
body: dict[str, object],
|
||||
headers: dict[str, str],
|
||||
url: str,
|
||||
) -> None:
|
||||
logger.pre_call(
|
||||
input="OCR document processing",
|
||||
api_key=api_key,
|
||||
additional_args={"complete_input_dict": body, "headers": headers, "api_base": url},
|
||||
)
|
||||
|
||||
|
||||
def post_call(
|
||||
logger: OcrLoggingProtocol,
|
||||
original_response: object,
|
||||
body: dict[str, object],
|
||||
headers: dict[str, str],
|
||||
) -> None:
|
||||
logger.post_call(
|
||||
original_response=original_response,
|
||||
additional_args={"complete_input_dict": body, "headers": headers},
|
||||
)
|
||||
|
||||
|
||||
class RustOcr(Protocol):
|
||||
|
|
|
|||
|
|
@ -33,3 +33,21 @@ def test_binding_validates_native_attribute(
|
|||
binding: Final = bindings.NativeBinding("route", validate=lambda item: item if isinstance(item, int) else None)
|
||||
|
||||
assert binding.load() == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize("available", [False, True])
|
||||
def test_decline_accessor_only_catches_admission_declines(monkeypatch: pytest.MonkeyPatch, available: bool) -> None:
|
||||
class Declined(Exception):
|
||||
pass
|
||||
|
||||
class Upstream(Exception):
|
||||
pass
|
||||
|
||||
native: Final = SimpleNamespace(RustBridgeDeclined=Declined, RustUpstreamError=Upstream) if available else None
|
||||
monkeypatch.setattr(bindings, "get_native_bridge", lambda: native)
|
||||
assert bindings.native_decline_types() == ((Declined,) if available else ())
|
||||
with pytest.raises(Upstream):
|
||||
try:
|
||||
raise Upstream("already dispatched")
|
||||
except bindings.native_decline_types():
|
||||
pytest.fail("upstream failure was allowed to replay")
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ import litellm
|
|||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.ocr import legacy
|
||||
from litellm.rust_bridge import bindings, configuration
|
||||
from litellm.rust_bridge.ocr import NATIVE_AOCR, NATIVE_OCR
|
||||
from litellm.rust_bridge.ocr import NATIVE_AOCR, NATIVE_OCR, post_call, pre_call, update_logging
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
|
|
@ -21,6 +21,72 @@ def isolated_ocr_configuration(monkeypatch: pytest.MonkeyPatch) -> Generator[Non
|
|||
configuration.reset_rust_configuration()
|
||||
|
||||
|
||||
def test_logging_redacts_views_and_preserves_opaque_arguments_and_pricing() -> None:
|
||||
logger: Final = Mock()
|
||||
opaque: Final = object()
|
||||
logger_fn: Final = object()
|
||||
kwargs: Final = {
|
||||
"vertex_credentials": "secret",
|
||||
"proxy_server_request": opaque,
|
||||
"metadata": opaque,
|
||||
"logger_fn": logger_fn,
|
||||
"litellm_request_debug": False,
|
||||
"litellm_call_id": "call-id",
|
||||
"input_cost_per_token": 0,
|
||||
"output_cost_per_token": None,
|
||||
}
|
||||
optional: Final = {"vertex_credentials": "secret", "pages": [1], "proxy_server_request": opaque}
|
||||
update_logging(logger, kwargs, "model", "vertex_ai", optional, ("vertex_credentials",), "https://provider")
|
||||
logger.update_from_kwargs.assert_called_once_with(
|
||||
kwargs={
|
||||
"vertex_credentials": "****",
|
||||
"metadata": opaque,
|
||||
"logger_fn": logger_fn,
|
||||
"litellm_request_debug": False,
|
||||
"litellm_call_id": "call-id",
|
||||
"input_cost_per_token": 0,
|
||||
"output_cost_per_token": None,
|
||||
},
|
||||
model="model",
|
||||
custom_llm_provider="vertex_ai",
|
||||
optional_params={"vertex_credentials": "****", "pages": [1]},
|
||||
litellm_params={
|
||||
"litellm_call_id": "call-id",
|
||||
"api_base": "https://provider",
|
||||
"logger_fn": logger_fn,
|
||||
"litellm_request_debug": False,
|
||||
"input_cost_per_token": 0,
|
||||
},
|
||||
)
|
||||
assert kwargs["vertex_credentials"] == optional["vertex_credentials"] == "secret"
|
||||
assert kwargs["proxy_server_request"] is optional["proxy_server_request"] is opaque
|
||||
assert logger.update_from_kwargs.call_args.kwargs["kwargs"]["metadata"] is opaque
|
||||
assert logger.update_from_kwargs.call_args.kwargs["optional_params"]["pages"] is optional["pages"]
|
||||
|
||||
|
||||
def test_logging_callbacks_receive_captured_payload_roots_and_propagate_errors() -> None:
|
||||
logger: Final = Mock()
|
||||
body: Final[dict[str, object]] = {"document": "original"}
|
||||
headers: Final = {"authorization": "key"}
|
||||
response: Final = object()
|
||||
pre_call(logger, "key", body, headers, "https://provider")
|
||||
post_call(logger, response, body, headers)
|
||||
logger.pre_call.assert_called_once_with(
|
||||
input="OCR document processing",
|
||||
api_key="key",
|
||||
additional_args={"complete_input_dict": body, "headers": headers, "api_base": "https://provider"},
|
||||
)
|
||||
for callback in (logger.pre_call, logger.post_call):
|
||||
assert callback.call_args.kwargs["additional_args"]["complete_input_dict"] is body
|
||||
assert callback.call_args.kwargs["additional_args"]["headers"] is headers
|
||||
assert logger.post_call.call_args.kwargs["original_response"] is response
|
||||
failure: Final = RuntimeError("callback failed")
|
||||
failing_logger: Final = Mock(pre_call=Mock(side_effect=failure))
|
||||
with pytest.raises(RuntimeError) as caught:
|
||||
pre_call(failing_logger, None, body, headers, "https://provider")
|
||||
assert caught.value is failure
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", [False, True])
|
||||
async def test_unavailable_native_uses_legacy(monkeypatch: pytest.MonkeyPatch, asynchronous: bool) -> None:
|
||||
|
|
|
|||
|
|
@ -1,14 +1,15 @@
|
|||
import json
|
||||
import asyncio
|
||||
import base64
|
||||
import gc
|
||||
import json
|
||||
import os
|
||||
import threading
|
||||
import weakref
|
||||
import gc
|
||||
from pathlib import Path
|
||||
from collections.abc import Generator
|
||||
from collections.abc import Coroutine, Generator
|
||||
from datetime import datetime, timezone
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from io import BytesIO
|
||||
from pathlib import Path
|
||||
from typing import Final, Protocol
|
||||
|
||||
import httpx
|
||||
|
|
@ -340,7 +341,12 @@ def test_native_lifecycle_core_encodes_python_file_input(
|
|||
@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):
|
||||
async def test_native_file_sources_preserve_sdk_behavior(
|
||||
ocr_server: tuple[ThreadingHTTPServer, list[dict[str, object]]],
|
||||
tmp_path: Path,
|
||||
asynchronous: bool,
|
||||
source: str,
|
||||
) -> None:
|
||||
server, requests = ocr_server
|
||||
path: Final = tmp_path / "scan.png"
|
||||
content: Final = b"document bytes"
|
||||
|
|
@ -388,9 +394,12 @@ async def test_native_file_sources_preserve_sdk_behavior(ocr_server, tmp_path, a
|
|||
"api_base": f"http://127.0.0.1:{server.server_port}",
|
||||
}
|
||||
response: Final = await litellm.aocr(**kwargs) if asynchronous else litellm.ocr(**kwargs)
|
||||
assert isinstance(response, OCRResponse)
|
||||
assert response.pages[0].markdown == "native OCR response"
|
||||
assert len(requests) == 1
|
||||
assert requests[0]["body"]["document"] == {
|
||||
body: Final = requests[0]["body"]
|
||||
assert isinstance(body, dict)
|
||||
assert body["document"] == {
|
||||
"type": "image_url",
|
||||
"image_url": "data:image/png;base64," + base64.b64encode(content).decode(),
|
||||
}
|
||||
|
|
@ -400,7 +409,11 @@ async def test_native_file_sources_preserve_sdk_behavior(ocr_server, tmp_path, a
|
|||
|
||||
@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):
|
||||
async def test_native_file_failures_preserve_identity_and_do_not_send(
|
||||
ocr_server: tuple[ThreadingHTTPServer, list[dict[str, object]]],
|
||||
tmp_path: Path,
|
||||
asynchronous: bool,
|
||||
) -> None:
|
||||
from litellm.rust_bridge import _native
|
||||
|
||||
server, requests = ocr_server
|
||||
|
|
@ -416,37 +429,47 @@ async def test_native_file_failures_preserve_identity_and_do_not_send(ocr_server
|
|||
|
||||
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}",
|
||||
}
|
||||
|
||||
async def invoke(source: Reader | Path) -> None:
|
||||
if asynchronous:
|
||||
await _native.aocr(
|
||||
model="mistral/mistral-ocr-latest",
|
||||
document={"type": "file", "file": source},
|
||||
api_key="test-key",
|
||||
api_base=f"http://127.0.0.1:{server.server_port}",
|
||||
)
|
||||
else:
|
||||
_native.ocr(
|
||||
model="mistral/mistral-ocr-latest",
|
||||
document={"type": "file", "file": source},
|
||||
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)
|
||||
await invoke(source)
|
||||
if source is reader:
|
||||
assert caught.value is failure
|
||||
else:
|
||||
assert caught.value.filename == str(path)
|
||||
assert isinstance(caught.value, FileNotFoundError)
|
||||
assert getattr(caught.value, "filename") == str(path)
|
||||
assert reader.calls == 1
|
||||
assert requests == []
|
||||
|
||||
|
||||
def test_unstarted_native_file_call_does_not_read_and_releases_reader():
|
||||
def test_unstarted_native_file_call_does_not_read_and_releases_reader() -> None:
|
||||
from litellm.rust_bridge import _native
|
||||
|
||||
class Reader:
|
||||
pending: Coroutine[object, object, OCRResponse] | None = None
|
||||
|
||||
def read(self) -> bytes:
|
||||
raise AssertionError("unstarted call consumed its document")
|
||||
|
||||
def create_call():
|
||||
def create_call() -> tuple[Coroutine[object, object, OCRResponse], weakref.ReferenceType[Reader]]:
|
||||
reader: Final = Reader()
|
||||
pending: Final = _native.aocr(
|
||||
model="mistral/mistral-ocr-latest", document={"type": "file", "file": reader}
|
||||
)
|
||||
pending: Final = _native.aocr(model="mistral/mistral-ocr-latest", document={"type": "file", "file": reader})
|
||||
reader.pending = pending
|
||||
return pending, weakref.ref(reader)
|
||||
|
||||
|
|
@ -457,6 +480,49 @@ def test_unstarted_native_file_call_does_not_read_and_releases_reader():
|
|||
assert reference() is None
|
||||
|
||||
|
||||
@pytest.mark.skipif(not hasattr(os, "mkfifo"), reason="requires Unix named pipes")
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_path_cancellation_waits_for_read_completion(
|
||||
ocr_server: tuple[ThreadingHTTPServer, list[dict[str, object]]], tmp_path: Path
|
||||
) -> None:
|
||||
from litellm.rust_bridge import _native
|
||||
|
||||
server, requests = ocr_server
|
||||
path: Final = tmp_path / "document.pdf"
|
||||
os.mkfifo(path)
|
||||
entered: Final = threading.Event()
|
||||
release: Final = threading.Event()
|
||||
|
||||
def supply_document() -> None:
|
||||
with path.open("wb") as stream:
|
||||
stream.write(b"document")
|
||||
stream.flush()
|
||||
entered.set()
|
||||
release.wait(5)
|
||||
|
||||
writer: Final = threading.Thread(target=supply_document, daemon=True)
|
||||
writer.start()
|
||||
task: Final = asyncio.create_task(
|
||||
_native.aocr(
|
||||
model="mistral/mistral-ocr-latest",
|
||||
document={"type": "file", "file": path},
|
||||
api_key="test-key",
|
||||
api_base=f"http://127.0.0.1:{server.server_port}",
|
||||
)
|
||||
)
|
||||
try:
|
||||
assert await asyncio.to_thread(entered.wait, 3)
|
||||
task.cancel()
|
||||
await asyncio.sleep(0.05)
|
||||
assert not task.done()
|
||||
finally:
|
||||
release.set()
|
||||
await asyncio.to_thread(writer.join, 3)
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await task
|
||||
assert requests == []
|
||||
|
||||
|
||||
@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