mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
perf(ocr): bound responses and reduce native scheduling overhead
This commit is contained in:
parent
612d335f59
commit
d17b9a8b0b
17 changed files with 472 additions and 17 deletions
1
litellm-rust/Cargo.lock
generated
1
litellm-rust/Cargo.lock
generated
|
|
@ -1948,6 +1948,7 @@ dependencies = [
|
|||
"azure_core",
|
||||
"azure_identity",
|
||||
"base64 0.22.1",
|
||||
"bytes",
|
||||
"data-url",
|
||||
"gcp_auth",
|
||||
"mime_guess",
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ license = "MIT"
|
|||
repository = "https://github.com/BerriAI/litellm"
|
||||
|
||||
[workspace.dependencies]
|
||||
bytes = "1"
|
||||
tracing = "0.1"
|
||||
tracing-subscriber = { version = "0.3", default-features = false, features = ["registry", "std"] }
|
||||
litellm-core = { path = "crates/core" }
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ repository.workspace = true
|
|||
autotests = false
|
||||
|
||||
[dependencies]
|
||||
bytes.workspace = true
|
||||
base64.workspace = true
|
||||
azure_core.workspace = true
|
||||
azure_identity.workspace = true
|
||||
|
|
|
|||
|
|
@ -46,6 +46,7 @@ pub const FUNCTION_TRACE_TARGET: &str = "litellm::function_trace";
|
|||
|
||||
pub(crate) const MEDIA_CONNECT_TIMEOUT_SECS: u64 = 10;
|
||||
|
||||
pub(crate) const OCR_RESPONSE_MAX_BYTES: usize = 64 * 1024 * 1024;
|
||||
pub(crate) const OCR_HTTP_TIMEOUT_SECS: u64 = 600;
|
||||
pub(crate) const OCR_CONNECT_TIMEOUT_SECS: u64 = 10;
|
||||
pub const OCR_INLINE_MAX_BYTES: usize = 50 * 1024 * 1024;
|
||||
|
|
|
|||
|
|
@ -24,7 +24,9 @@ pub(super) async fn read_operation_response(
|
|||
hooks: &Arc<dyn OcrHooks>,
|
||||
) -> Result<DecodedOcrResponse<AzureDocumentIntelligenceOperation>, OcrError> {
|
||||
if response.status() != reqwest::StatusCode::ACCEPTED {
|
||||
let bytes = crate::ocr::client::read_response_bytes(response).await?;
|
||||
let bytes =
|
||||
crate::ocr::client::read_response_bytes(response, connection.max_response_bytes)
|
||||
.await?;
|
||||
crate::ocr::handler::post_call(hooks, &bytes).await?;
|
||||
return Ok(crate::ocr::wire::decode_response(&bytes, native)?);
|
||||
}
|
||||
|
|
@ -42,7 +44,8 @@ pub(super) async fn read_operation_response(
|
|||
{
|
||||
return Err(OcrPollingError::PollOrigin.into());
|
||||
}
|
||||
let bytes = crate::ocr::client::read_response_bytes(response).await?;
|
||||
let bytes =
|
||||
crate::ocr::client::read_response_bytes(response, connection.max_response_bytes).await?;
|
||||
crate::ocr::handler::post_call(hooks, &bytes).await?;
|
||||
poll_operation(http_client, operation, headers, connection, native, hooks).await
|
||||
}
|
||||
|
|
@ -84,7 +87,11 @@ async fn poll_operation(
|
|||
.max(1);
|
||||
let decoded = tokio::time::timeout_at(
|
||||
deadline,
|
||||
read_json_response::<AzureDocumentIntelligenceOperation>(response, native),
|
||||
read_json_response::<AzureDocumentIntelligenceOperation>(
|
||||
response,
|
||||
native,
|
||||
connection.max_response_bytes,
|
||||
),
|
||||
)
|
||||
.await
|
||||
.map_err(|_| OcrPollingError::PollTimeout)??;
|
||||
|
|
|
|||
|
|
@ -60,7 +60,9 @@ pub(crate) trait OcrAdapter: Send + Sync + Sized + 'static {
|
|||
Output = Result<super::wire::DecodedOcrResponse<Self::ProviderResponse>, OcrError>,
|
||||
> + Send {
|
||||
async move {
|
||||
let bytes = super::client::read_response_bytes(response).await?;
|
||||
let bytes =
|
||||
super::client::read_response_bytes(response, request.connection.max_response_bytes)
|
||||
.await?;
|
||||
super::handler::post_call(&request.hooks, &bytes).await?;
|
||||
Ok(super::wire::decode_response(
|
||||
&bytes,
|
||||
|
|
|
|||
|
|
@ -93,7 +93,7 @@ pub(super) async fn prepare_document(
|
|||
.map_err(crate::error::TransportError::from)?;
|
||||
let uploaded = crate::ocr::client::read_json_response::<
|
||||
crate::ocr::codecs::reducto::ReductoUploadResponse,
|
||||
>(response, false)
|
||||
>(response, false, connection.max_response_bytes)
|
||||
.await?
|
||||
.data;
|
||||
let file_id = uploaded
|
||||
|
|
|
|||
|
|
@ -1,9 +1,10 @@
|
|||
use std::sync::OnceLock;
|
||||
use std::time::Duration;
|
||||
|
||||
use bytes::{Bytes, BytesMut};
|
||||
use serde::de::DeserializeOwned;
|
||||
|
||||
use super::error::OcrError;
|
||||
use super::error::{OcrError, OcrResponseError};
|
||||
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
|
||||
use super::wire::{DecodedOcrResponse, decode_response};
|
||||
use crate::Error;
|
||||
|
|
@ -128,14 +129,40 @@ pub async fn ocr(request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, Error
|
|||
pub async fn read_json_response<T: DeserializeOwned>(
|
||||
response: reqwest::Response,
|
||||
native: bool,
|
||||
max_response_bytes: usize,
|
||||
) -> Result<DecodedOcrResponse<T>, OcrError> {
|
||||
let bytes = read_response_bytes(response).await?;
|
||||
let bytes = read_response_bytes(response, max_response_bytes).await?;
|
||||
Ok(decode_response(&bytes, native)?)
|
||||
}
|
||||
|
||||
pub(crate) async fn read_response_bytes(response: reqwest::Response) -> Result<Vec<u8>, OcrError> {
|
||||
pub(crate) async fn read_response_bytes(
|
||||
mut response: reqwest::Response,
|
||||
max_response_bytes: usize,
|
||||
) -> Result<Bytes, OcrError> {
|
||||
let status = response.status();
|
||||
let bytes = response.bytes().await.map_err(transport_error)?;
|
||||
let limit = if status.is_success() {
|
||||
max_response_bytes
|
||||
} else {
|
||||
max_response_bytes.min(4 * (crate::constants::UPSTREAM_ERROR_BODY_MAX_CHARS + 1))
|
||||
};
|
||||
if status.is_success()
|
||||
&& response
|
||||
.content_length()
|
||||
.is_some_and(|length| length > limit as u64)
|
||||
{
|
||||
return Err(OcrResponseError::TooLarge { limit }.into());
|
||||
}
|
||||
let mut bytes = BytesMut::new();
|
||||
while let Some(chunk) = response.chunk().await.map_err(transport_error)? {
|
||||
let remaining = limit.saturating_sub(bytes.len());
|
||||
if status.is_success() && chunk.len() > remaining {
|
||||
return Err(OcrResponseError::TooLarge { limit }.into());
|
||||
}
|
||||
bytes.extend_from_slice(&chunk[..chunk.len().min(remaining)]);
|
||||
if !status.is_success() && bytes.len() == limit {
|
||||
break;
|
||||
}
|
||||
}
|
||||
if !status.is_success() {
|
||||
return Err(crate::error::TransportError::Http {
|
||||
status: status.as_u16(),
|
||||
|
|
@ -143,7 +170,7 @@ pub(crate) async fn read_response_bytes(response: reqwest::Response) -> Result<V
|
|||
}
|
||||
.into());
|
||||
}
|
||||
Ok(bytes.to_vec())
|
||||
Ok(bytes.freeze())
|
||||
}
|
||||
|
||||
pub(crate) fn transport_error(error: reqwest::Error) -> Error {
|
||||
|
|
|
|||
|
|
@ -44,6 +44,8 @@ pub enum OcrRequestError {
|
|||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Error)]
|
||||
pub enum OcrResponseError {
|
||||
#[error("OCR response exceeds the size limit of {limit} bytes")]
|
||||
TooLarge { limit: usize },
|
||||
#[error("invalid OCR response field: {path}")]
|
||||
ResponseField { path: String },
|
||||
#[error("OCR response is missing non-empty content")]
|
||||
|
|
|
|||
|
|
@ -68,6 +68,7 @@ pub struct OcrConnection {
|
|||
pub extra_headers_source: InputSource,
|
||||
pub timeout: Duration,
|
||||
pub max_download_bytes: u64,
|
||||
pub max_response_bytes: usize,
|
||||
pub poll_timeout: Duration,
|
||||
}
|
||||
|
||||
|
|
@ -82,6 +83,7 @@ impl Default for OcrConnection {
|
|||
extra_headers_source: InputSource::Deployment,
|
||||
timeout: Duration::from_secs(OCR_HTTP_TIMEOUT_SECS),
|
||||
max_download_bytes: crate::constants::OCR_DOWNLOAD_MAX_BYTES,
|
||||
max_response_bytes: crate::constants::OCR_RESPONSE_MAX_BYTES,
|
||||
poll_timeout: Duration::from_secs(crate::constants::OCR_POLL_TIMEOUT_SECS),
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ use serde::{
|
|||
};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
const COMMON_OPTION_FIELDS: &[&str] = &["req_format", "extra_body"];
|
||||
const COMMON_OPTION_FIELDS: &[&str] = &["req_format", "extra_body", "max_response_bytes"];
|
||||
const MISTRAL_OPTION_FIELDS: &[&str] = &[
|
||||
"pages",
|
||||
"include_image_base64",
|
||||
|
|
@ -140,11 +140,28 @@ pub fn decode_request(wire: OcrWireRequest) -> Result<LiteLLMOcrRequest, Error>
|
|||
})
|
||||
.transpose()?;
|
||||
let defaults = OcrConnection::default();
|
||||
let max_response_bytes = wire
|
||||
.optional_params
|
||||
.get("max_response_bytes")
|
||||
.map(|value| {
|
||||
value
|
||||
.as_u64()
|
||||
.and_then(|value| usize::try_from(value).ok())
|
||||
.filter(|value| *value > 0)
|
||||
.ok_or_else(|| OcrRequestError::RequestField {
|
||||
path: "max_response_bytes".into(),
|
||||
})
|
||||
})
|
||||
.transpose()?
|
||||
.unwrap_or(defaults.max_response_bytes);
|
||||
let request = LiteLLMOcrRequest::new(
|
||||
wire.model,
|
||||
document,
|
||||
wire.custom_llm_provider.as_deref(),
|
||||
wire.optional_params,
|
||||
wire.optional_params
|
||||
.into_iter()
|
||||
.filter(|(name, _)| name != "max_response_bytes")
|
||||
.collect(),
|
||||
)?;
|
||||
let connection = OcrConnection {
|
||||
api_key: nonblank(wire.api_key),
|
||||
|
|
@ -155,6 +172,7 @@ pub fn decode_request(wire: OcrWireRequest) -> Result<LiteLLMOcrRequest, Error>
|
|||
extra_headers_source,
|
||||
timeout: timeout.unwrap_or(defaults.timeout),
|
||||
max_download_bytes: defaults.max_download_bytes,
|
||||
max_response_bytes,
|
||||
poll_timeout: defaults.poll_timeout,
|
||||
};
|
||||
Ok(LiteLLMOcrRequest {
|
||||
|
|
|
|||
|
|
@ -624,3 +624,215 @@ async fn missing_host_result_preserves_pending_operation() {
|
|||
OcrCallStep::Host(OcrHostOperation::Lifecycle(HostPhase::Prepare))
|
||||
));
|
||||
}
|
||||
|
||||
async fn read_bounded_response(
|
||||
response: Vec<u8>,
|
||||
limit: usize,
|
||||
) -> Result<bytes::Bytes, super::error::OcrError> {
|
||||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
let server = tokio::spawn(async move {
|
||||
let (mut socket, _) = listener.accept().await.unwrap();
|
||||
let mut request = [0; 4096];
|
||||
assert!(socket.read(&mut request).await.unwrap() > 0);
|
||||
socket.write_all(&response).await.unwrap();
|
||||
std::future::pending::<()>().await;
|
||||
});
|
||||
let response = reqwest::Client::new()
|
||||
.get(format!("http://{address}"))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
let result = tokio::time::timeout(
|
||||
std::time::Duration::from_secs(2),
|
||||
super::client::read_response_bytes(response, limit),
|
||||
)
|
||||
.await;
|
||||
server.abort();
|
||||
let _ = server.await;
|
||||
result.expect("bounded reads must finish without waiting for the rest of an oversized body")
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn response_limit_accepts_exact_size_and_rejects_declared_and_chunked_overflow() {
|
||||
use super::error::{OcrError, OcrResponseError};
|
||||
|
||||
for response in [
|
||||
"HTTP/1.1 200 OK\r\nContent-Length: 8\r\n\r\nabcdefgh",
|
||||
"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n4\r\nabcd\r\n4\r\nefgh\r\n0\r\n\r\n",
|
||||
] {
|
||||
assert_eq!(
|
||||
read_bounded_response(response.as_bytes().to_vec(), 8)
|
||||
.await
|
||||
.unwrap(),
|
||||
"abcdefgh"
|
||||
);
|
||||
}
|
||||
for response in [
|
||||
"HTTP/1.1 200 OK\r\nContent-Length: 9\r\n\r\n",
|
||||
"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n4\r\nabcd\r\n5\r\nefghi\r\n",
|
||||
] {
|
||||
assert!(matches!(
|
||||
read_bounded_response(response.as_bytes().to_vec(), 8).await,
|
||||
Err(OcrError::Response(OcrResponseError::TooLarge { limit: 8 }))
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn oversized_error_retains_http_status_and_bounded_diagnostics_without_draining() {
|
||||
let prefix = "x".repeat(4 * (crate::constants::UPSTREAM_ERROR_BODY_MAX_CHARS + 1));
|
||||
for headers in ["Content-Length: 1000000", "Transfer-Encoding: chunked"] {
|
||||
let body = if headers.starts_with("Transfer") {
|
||||
format!("{:x}\r\n{prefix}\r\n", prefix.len())
|
||||
} else {
|
||||
prefix.clone()
|
||||
};
|
||||
let response = format!("HTTP/1.1 429 Too Many Requests\r\n{headers}\r\n\r\n{body}");
|
||||
let error = read_bounded_response(response.into_bytes(), 4096)
|
||||
.await
|
||||
.unwrap_err();
|
||||
match error {
|
||||
super::error::OcrError::Transport(crate::error::TransportError::Http {
|
||||
status,
|
||||
body,
|
||||
}) => {
|
||||
assert_eq!(status, 429);
|
||||
assert_eq!(
|
||||
body,
|
||||
format!(
|
||||
"{}... (truncated)",
|
||||
"x".repeat(crate::constants::UPSTREAM_ERROR_BODY_MAX_CHARS)
|
||||
)
|
||||
);
|
||||
}
|
||||
error => panic!("unexpected error: {error}"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn response_limit_is_validated_and_not_forwarded_to_the_provider() {
|
||||
let request = wire_request(
|
||||
"mistral/model",
|
||||
"http://localhost",
|
||||
json!({"max_response_bytes": 123}),
|
||||
);
|
||||
assert_eq!(request.connection.max_response_bytes, 123);
|
||||
assert!(!request.optional_params.contains_key("max_response_bytes"));
|
||||
for value in [
|
||||
json!(0),
|
||||
json!(-1),
|
||||
json!(true),
|
||||
json!("123"),
|
||||
json!(1.5),
|
||||
Value::Null,
|
||||
] {
|
||||
let wire = serde_json::from_value(json!({
|
||||
"model": "mistral/model", "document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"},
|
||||
"optional_params": {"max_response_bytes": value}
|
||||
})).unwrap();
|
||||
let Err(error) = decode_request(wire) else {
|
||||
panic!("invalid response limit accepted")
|
||||
};
|
||||
assert!(error.to_string().contains("max_response_bytes"));
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct PendingToken {
|
||||
entered: Arc<tokio::sync::Notify>,
|
||||
dropped: Arc<std::sync::atomic::AtomicBool>,
|
||||
}
|
||||
|
||||
struct TokenFutureDrop(Arc<std::sync::atomic::AtomicBool>);
|
||||
|
||||
impl Drop for TokenFutureDrop {
|
||||
fn drop(&mut self) {
|
||||
self.0.store(true, std::sync::atomic::Ordering::SeqCst);
|
||||
}
|
||||
}
|
||||
|
||||
impl crate::auth::TokenProvider for PendingToken {
|
||||
fn acquire(&self) -> crate::auth::TokenFuture<'_> {
|
||||
Box::pin(async move {
|
||||
let _guard = TokenFutureDrop(self.dropped.clone());
|
||||
self.entered.notify_one();
|
||||
std::future::pending().await
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn cancellation_waits_for_provider_capture_drop_even_when_acknowledgement_is_cancelled() {
|
||||
use crate::call_lifecycle::host::HostFailure;
|
||||
use std::future::Future;
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
use std::task::Poll;
|
||||
|
||||
for interrupt_acknowledgement in [false, true] {
|
||||
let entered = Arc::new(tokio::sync::Notify::new());
|
||||
let dropped = Arc::new(AtomicBool::new(false));
|
||||
let request = wire_request("azure_ai/mistral-ocr", "https://example.invalid", json!({}));
|
||||
let request = super::LiteLLMOcrRequest {
|
||||
connection: super::OcrConnection {
|
||||
extra_headers: vec![("authorization".into(), "Bearer test-key".into())],
|
||||
..request.connection
|
||||
},
|
||||
azure_ad_token_provider: Some(crate::auth::TokenProviderHandle::new(Arc::new(
|
||||
PendingToken {
|
||||
entered: entered.clone(),
|
||||
dropped: dropped.clone(),
|
||||
},
|
||||
))),
|
||||
..request
|
||||
};
|
||||
let NativeOutcome::Completed(mut call) =
|
||||
OcrCall::admit(super::test_support::ocr_client(), OcrAdmission::all())
|
||||
else {
|
||||
panic!("supported call declined")
|
||||
};
|
||||
let mut request = Some(request);
|
||||
let mut result = None;
|
||||
tokio::time::timeout(std::time::Duration::from_secs(2), async {
|
||||
loop {
|
||||
tokio::select! {
|
||||
_ = 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(operation) => NoopOcrHost.invoke(operation).await,
|
||||
OcrCallStep::Complete(_) => panic!("pending provider completed"),
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
}).await.unwrap();
|
||||
assert!(!dropped.load(Ordering::SeqCst));
|
||||
let selected = crate::Error::InvalidRequest("cancelled".into());
|
||||
if interrupt_acknowledgement {
|
||||
let mut acknowledgement =
|
||||
Box::pin(call.interrupt(HostFailure::Cancelled(selected.clone())));
|
||||
std::future::poll_fn(|cx| {
|
||||
assert!(acknowledgement.as_mut().poll(cx).is_pending());
|
||||
Poll::Ready(())
|
||||
})
|
||||
.await;
|
||||
drop(acknowledgement);
|
||||
assert!(!dropped.load(Ordering::SeqCst));
|
||||
}
|
||||
let result = tokio::time::timeout(
|
||||
std::time::Duration::from_secs(2),
|
||||
call.interrupt(HostFailure::Cancelled(selected.clone())),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(matches!(result, Err(error) if error == selected));
|
||||
assert!(
|
||||
dropped.load(Ordering::SeqCst),
|
||||
"cancellation returned while provider captures were still alive"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,5 +1,7 @@
|
|||
use std::future::Future;
|
||||
use std::panic::AssertUnwindSafe;
|
||||
use std::pin::Pin;
|
||||
use std::task::{Context, Poll, Waker};
|
||||
use std::time::Duration;
|
||||
|
||||
use futures_util::FutureExt;
|
||||
|
|
@ -96,6 +98,24 @@ where
|
|||
pyo3_async_runtimes::tokio::future_into_py(py, async move { catch_future_panic(future).await? })
|
||||
}
|
||||
|
||||
pub(crate) fn poll_async_value<T, F>(py: Python<'_>, future: Pin<&mut F>) -> PyResult<Poll<T>>
|
||||
where
|
||||
T: Send,
|
||||
F: Future<Output = PyResult<T>> + Send,
|
||||
{
|
||||
let result = release_gil(py, || {
|
||||
let _runtime = pyo3_async_runtimes::tokio::get_runtime().enter();
|
||||
std::panic::catch_unwind(AssertUnwindSafe(|| {
|
||||
future.poll(&mut Context::from_waker(Waker::noop()))
|
||||
}))
|
||||
.map_err(panic_to_pyerr)
|
||||
})?;
|
||||
match result {
|
||||
Poll::Ready(result) => result.map(Poll::Ready),
|
||||
Poll::Pending => Ok(Poll::Pending),
|
||||
}
|
||||
}
|
||||
|
||||
fn map_core_result<T, E>(result: Result<T, E>, map_error: fn(E) -> PyErr) -> PyResult<T> {
|
||||
match result {
|
||||
Ok(value) => Ok(value),
|
||||
|
|
@ -223,6 +243,76 @@ mod tests {
|
|||
.expect("result should convert")
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn inline_poll_releases_gil_and_enters_runtime() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let (sender, receiver) = mpsc::sync_channel(1);
|
||||
let worker = thread::spawn(move || Python::attach(|_| sender.send(()).unwrap()));
|
||||
let mut future = Box::pin(async move {
|
||||
receiver.recv_timeout(Duration::from_secs(2)).unwrap();
|
||||
Ok(Handle::try_current().is_ok())
|
||||
});
|
||||
assert_eq!(
|
||||
poll_async_value(py, future.as_mut()).unwrap(),
|
||||
Poll::Ready(true)
|
||||
);
|
||||
worker.join().unwrap();
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn inline_poll_contains_panics_and_preserves_python_errors() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let mut panicking = Box::pin(poll_fn(|_| -> Poll<PyResult<()>> {
|
||||
panic!("inline native panic")
|
||||
}));
|
||||
let error = poll_async_value(py, panicking.as_mut()).unwrap_err();
|
||||
assert!(error.is_instance_of::<PanicException>(py));
|
||||
let original = PyRuntimeError::new_err("inline failure");
|
||||
let identity = original.value(py).clone().unbind();
|
||||
let mut failing = Box::pin(async move { Err::<(), _>(original) });
|
||||
let error = poll_async_value(py, failing.as_mut()).unwrap_err();
|
||||
assert!(error.value(py).is(identity.bind(py)));
|
||||
});
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
fn pending_after_inline_poll(py: Python<'_>) -> PyResult<Bound<'_, PyAny>> {
|
||||
let starts = Arc::new(AtomicUsize::new(0));
|
||||
let observed = Arc::clone(&starts);
|
||||
let mut future = Box::pin(async move {
|
||||
starts.fetch_add(1, Ordering::SeqCst);
|
||||
tokio::time::sleep(Duration::from_millis(5)).await;
|
||||
Ok(starts.load(Ordering::SeqCst))
|
||||
});
|
||||
assert!(poll_async_value(py, future.as_mut())?.is_pending());
|
||||
assert_eq!(observed.load(Ordering::SeqCst), 1);
|
||||
run_async_value(py, future)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn inline_pending_future_resumes_on_tokio_without_restarting() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = PyDict::new(py);
|
||||
locals
|
||||
.set_item(
|
||||
"pending",
|
||||
wrap_pyfunction!(pending_after_inline_poll, py).unwrap(),
|
||||
)
|
||||
.unwrap();
|
||||
py.run(
|
||||
pyo3::ffi::c_str!(
|
||||
"import asyncio\nasync def exercise():\n assert await asyncio.wait_for(pending(), 2) == 1\nasyncio.run(exercise())"
|
||||
),
|
||||
Some(&locals),
|
||||
Some(&locals),
|
||||
).unwrap();
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sync_runner_polls_future_on_the_caller_thread() {
|
||||
Python::initialize();
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
use std::future::Future;
|
||||
use std::pin::Pin;
|
||||
use std::sync::Arc;
|
||||
use std::task::Poll;
|
||||
|
||||
use futures_util::future::{AbortHandle, Abortable};
|
||||
use litellm_core::call_lifecycle::host::{HostFailure, HostPhase, HostStep};
|
||||
|
|
@ -10,7 +11,7 @@ use pyo3::prelude::*;
|
|||
use pyo3::types::{PyDict, PyTuple};
|
||||
use tokio::sync::Mutex;
|
||||
|
||||
use crate::execution::{run_async_value, run_sync_value};
|
||||
use crate::execution::{poll_async_value, run_async_value, run_sync_value};
|
||||
|
||||
mod bindings;
|
||||
mod handle;
|
||||
|
|
@ -129,6 +130,10 @@ impl<R: PythonRoute> PythonLifecycle<R> {
|
|||
Ok(())
|
||||
};
|
||||
if self.route.state().asynchronous {
|
||||
let mut future = Box::pin(future);
|
||||
if let Poll::Ready(()) = poll_async_value(py, future.as_mut())? {
|
||||
return Ok(HostStep::Ready(self.take_native_result()?));
|
||||
}
|
||||
let (abort, registration) = AbortHandle::new_pair();
|
||||
self.native_abort = Some(abort);
|
||||
self.pending = Some(PendingOperation::Native);
|
||||
|
|
@ -755,6 +760,51 @@ mod tests {
|
|||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ready_native_lifecycle_completes_without_scheduling() {
|
||||
let _guard = PYTHON_GLOBALS
|
||||
.lock()
|
||||
.unwrap_or_else(|error| error.into_inner());
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let source = std::ffi::CString::new(include_str!(
|
||||
"../../../../../litellm/rust_bridge/lifecycle.py"
|
||||
))
|
||||
.unwrap();
|
||||
PyModule::from_code(
|
||||
py,
|
||||
&source,
|
||||
pyo3::ffi::c_str!("lifecycle.py"),
|
||||
pyo3::ffi::c_str!("litellm.rust_bridge.lifecycle"),
|
||||
)
|
||||
.unwrap();
|
||||
let route = SyntheticRoute(
|
||||
PythonCallState::new(
|
||||
py,
|
||||
PyTuple::empty(py).unbind(),
|
||||
PyDict::new(py).unbind(),
|
||||
true,
|
||||
"synthetic",
|
||||
)
|
||||
.unwrap(),
|
||||
);
|
||||
let coroutine = run_call(py, SyntheticCall(false), route).unwrap();
|
||||
let completed = coroutine
|
||||
.call_method1(py, "send", (py.None(),))
|
||||
.unwrap_err();
|
||||
assert!(completed.is_instance_of::<pyo3::exceptions::PyStopIteration>(py));
|
||||
assert_eq!(
|
||||
completed
|
||||
.value(py)
|
||||
.getattr("value")
|
||||
.unwrap()
|
||||
.extract::<String>()
|
||||
.unwrap(),
|
||||
"shared lifecycle",
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn python_driver_preserves_inline_await_and_native_ownership() {
|
||||
let _guard = PYTHON_GLOBALS
|
||||
|
|
|
|||
|
|
@ -9,12 +9,28 @@ use litellm_python_interop::to_py_preserving_errors as to_py;
|
|||
|
||||
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(crate) fn update_ocr(
|
||||
pub(super) fn update_ocr(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
kwargs: &Py<PyDict>,
|
||||
pre_call: &OcrPreCallRequest,
|
||||
pre_call: &OcrLoggingFields,
|
||||
url: &str,
|
||||
) -> PyResult<()> {
|
||||
let redact = py
|
||||
|
|
|
|||
|
|
@ -27,7 +27,7 @@ use crate::marshal::{project_optional_fields, python_timeout_seconds, request_in
|
|||
struct PythonOcrHost {
|
||||
state: PythonCallState,
|
||||
request: Option<Py<PyAny>>,
|
||||
pre_call: Option<OcrPreCallRequest>,
|
||||
pre_call: Option<callbacks::OcrLoggingFields>,
|
||||
document: Option<Py<PyAny>>,
|
||||
api_key: Option<Py<PyAny>>,
|
||||
azure_ad_token_provider: Option<PythonTokenProvider>,
|
||||
|
|
@ -68,7 +68,7 @@ impl PythonOcrHost {
|
|||
self.document.as_ref().ok_or_else(missing_state)?,
|
||||
)?;
|
||||
self.retained_fields = Some(retained_fields.unbind());
|
||||
self.pre_call = Some(request.clone());
|
||||
self.pre_call = Some((&request).into());
|
||||
Ok(request)
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -797,3 +797,28 @@ async def test_shared_call_limits_still_reject_before_reading_ocr_file(
|
|||
await call_aocr(ocr_server, **arguments) if asynchronous else call_ocr(ocr_server, **arguments)
|
||||
assert reads == []
|
||||
assert ocr_server.requests == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", [False, True])
|
||||
@pytest.mark.parametrize("extra_bytes", [0, 1])
|
||||
async def test_response_limit_is_enforced_at_the_public_boundary(
|
||||
ocr_server: RecordingServer, asynchronous: bool, extra_bytes: int
|
||||
) -> None:
|
||||
limit: Final = len(json.dumps(OCR_RESPONSE).encode()) - extra_bytes
|
||||
if extra_bytes:
|
||||
with pytest.raises(litellm.APIConnectionError, match="OCR response exceeds the size limit"):
|
||||
await call_aocr(ocr_server, max_response_bytes=limit) if asynchronous else call_ocr(
|
||||
ocr_server, max_response_bytes=limit
|
||||
)
|
||||
else:
|
||||
response: Final = (
|
||||
await call_aocr(ocr_server, max_response_bytes=limit)
|
||||
if asynchronous
|
||||
else call_ocr(ocr_server, max_response_bytes=limit)
|
||||
)
|
||||
assert response.pages[0].markdown == "native OCR response"
|
||||
assert len(ocr_server.requests) == 1
|
||||
body: Final = ocr_server.requests[0].body
|
||||
assert isinstance(body, dict)
|
||||
assert "max_response_bytes" not in body
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue