From f1553ae9c3ce9830fb9c82fe2ad3c1c0236e04d0 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sun, 6 Sep 2026 17:28:01 -0700 Subject: [PATCH] wip --- litellm-rust/Cargo.lock | 23 + litellm-rust/Cargo.toml | 1 + litellm-rust/crates/core/src/constants.rs | 2 + litellm-rust/crates/core/src/ocr/mod.rs | 1 + litellm-rust/crates/core/src/ocr/transport.rs | 77 +++ .../crates/core/src/ocr/transport/tests.rs | 44 ++ litellm-rust/crates/python-bridge/Cargo.toml | 3 +- .../crates/python-bridge/src/execution.rs | 36 +- litellm-rust/crates/python-bridge/src/lib.rs | 4 +- .../python-bridge/src/routes/bindings.rs | 157 +++++ .../crates/python-bridge/src/routes/mod.rs | 3 + .../python-bridge/src/routes/ocr_retained.rs | 144 ++++ .../python-bridge/tests/ocr_retained.rs | 243 +++++++ .../crates/python-interop/src/callback.rs | 2 +- litellm/ocr/main.py | 50 ++ litellm/rust_bridge/ocr.py | 4 + litellm/rust_bridge/ocr_retained.py | 146 ++++ .../ocr/retained_boundary_fixture.py | 642 ++++++++++++++++++ 18 files changed, 1576 insertions(+), 6 deletions(-) create mode 100644 litellm-rust/crates/core/src/ocr/transport.rs create mode 100644 litellm-rust/crates/core/src/ocr/transport/tests.rs create mode 100644 litellm-rust/crates/python-bridge/src/routes/bindings.rs create mode 100644 litellm-rust/crates/python-bridge/src/routes/ocr_retained.rs create mode 100644 litellm-rust/crates/python-bridge/tests/ocr_retained.rs create mode 100644 litellm/rust_bridge/ocr_retained.py create mode 100644 tests/test_litellm/ocr/retained_boundary_fixture.py diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 43b9ec1aac2..5234f2e7f66 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -1473,6 +1473,7 @@ dependencies = [ "pyo3-async-runtimes", "serde", "serde_json", + "strum", "tokio", "tokio-tungstenite", "tracing", @@ -2367,6 +2368,28 @@ version = "1.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6ce2be8dc25455e1f91df71bfa12ad37d7af1092ae736f3a6cd0e37bc7810596" +[[package]] +name = "strum" +version = "0.26.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8fec0f0aef304996cf250b31b5a10dee7980c85da9d759361292b8bca5a18f06" +dependencies = [ + "strum_macros", +] + +[[package]] +name = "strum_macros" +version = "0.26.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4c6bee85a5a24955dc440386795aa378cd9cf82acd5f764469152d2270e581be" +dependencies = [ + "heck", + "proc-macro2", + "quote", + "rustversion", + "syn 2.0.119", +] + [[package]] name = "subtle" version = "2.6.1" diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 82de7f40069..41c2128b500 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -33,6 +33,7 @@ rustls-native-certs = "0.8" serde = { version = "1.0", features = ["derive"] } serde_json = { version = "1.0", features = ["float_roundtrip"] } sha2 = "0.10" +strum = { version = "0.26", features = ["derive"] } subtle = "2" thiserror = "2.0" tokio = { version = "1", features = ["rt-multi-thread", "macros", "time", "net"] } diff --git a/litellm-rust/crates/core/src/constants.rs b/litellm-rust/crates/core/src/constants.rs index fc81f4fa029..b54df02d72b 100644 --- a/litellm-rust/crates/core/src/constants.rs +++ b/litellm-rust/crates/core/src/constants.rs @@ -32,6 +32,8 @@ pub(crate) const CHAT_COMPLETIONS_CONNECT_TIMEOUT_SECS: u64 = 10; pub(crate) const AUDIO_TRANSCRIPTION_TIMEOUT_SECS: u64 = 600; +pub(crate) const OCR_CONNECT_TIMEOUT_SECS: u64 = 10; + /// `object` field every non-streaming chat completion response carries. pub const CHAT_COMPLETION_OBJECT: &str = "chat.completion"; diff --git a/litellm-rust/crates/core/src/ocr/mod.rs b/litellm-rust/crates/core/src/ocr/mod.rs index ec2fbb969a6..c2e707160b6 100644 --- a/litellm-rust/crates/core/src/ocr/mod.rs +++ b/litellm-rust/crates/core/src/ocr/mod.rs @@ -1,2 +1,3 @@ pub mod transformation; +pub mod transport; pub mod types; diff --git a/litellm-rust/crates/core/src/ocr/transport.rs b/litellm-rust/crates/core/src/ocr/transport.rs new file mode 100644 index 00000000000..9501645cb1d --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/transport.rs @@ -0,0 +1,77 @@ +use std::sync::OnceLock; +use std::time::Duration; + +use reqwest::header::{HeaderMap, HeaderName, HeaderValue}; + +use crate::constants::OCR_CONNECT_TIMEOUT_SECS; +use crate::error::Error; + +pub struct Request { + pub url: String, + pub headers: Vec<(Vec, Vec)>, + pub body: Vec, + pub timeout_seconds: f64, +} + +pub struct Response { + pub status: u16, + pub headers: Vec<(Vec, Vec)>, + pub content: Vec, +} + +pub async fn send(request: Request) -> Result { + let timeout = Duration::try_from_secs_f64(request.timeout_seconds) + .ok() + .filter(|timeout| !timeout.is_zero()) + .ok_or_else(|| Error::InvalidRequest("OCR timeout must be positive and finite".into()))?; + let mut headers = HeaderMap::new(); + for (name, value) in request.headers { + let name = HeaderName::from_bytes(&name) + .map_err(|_| Error::InvalidRequest("invalid OCR header name".into()))?; + let value = HeaderValue::from_bytes(&value) + .map_err(|_| Error::InvalidRequest("invalid OCR header value".into()))?; + headers.append(name, value); + } + + static CLIENT: OnceLock> = OnceLock::new(); + let client = CLIENT + .get_or_init(|| { + reqwest::Client::builder() + .connect_timeout(Duration::from_secs(OCR_CONNECT_TIMEOUT_SECS)) + .redirect(reqwest::redirect::Policy::none()) + .no_gzip() + .no_brotli() + .no_deflate() + .no_zstd() + .build() + }) + .as_ref() + .map_err(|_| Error::Network("could not initialize OCR HTTP client".into()))?; + let response = client + .post(request.url) + .headers(headers) + .body(request.body) + .timeout(timeout) + .send() + .await + .map_err(|_| Error::Network("OCR transport failed".into()))?; + let status = response.status().as_u16(); + let headers = response + .headers() + .iter() + .map(|(name, value)| (name.as_str().as_bytes().to_vec(), value.as_bytes().to_vec())) + .collect(); + let content = response + .bytes() + .await + .map_err(|_| Error::Network("could not read OCR response".into()))? + .to_vec(); + Ok(Response { + status, + headers, + content, + }) +} + +#[cfg(test)] +mod tests; diff --git a/litellm-rust/crates/core/src/ocr/transport/tests.rs b/litellm-rust/crates/core/src/ocr/transport/tests.rs new file mode 100644 index 00000000000..4a77700e9f9 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/transport/tests.rs @@ -0,0 +1,44 @@ +use super::*; + +fn request() -> Request { + Request { + url: "unknown://private-document?secret=credential".into(), + headers: vec![], + body: vec![0, 255], + timeout_seconds: 1.0, + } +} + +#[tokio::test] +async fn rejects_invalid_timeouts_and_headers_without_echoing_wire_values() { + for timeout_seconds in [0.0, -1.0, f64::NAN, f64::INFINITY, f64::MAX] { + let result = send(Request { + timeout_seconds, + ..request() + }) + .await; + assert!( + matches!(result, Err(Error::InvalidRequest(message)) if message == "OCR timeout must be positive and finite") + ); + } + for (headers, expected) in [ + ( + vec![(b"private\nname".to_vec(), b"secret".to_vec())], + "invalid OCR header name", + ), + ( + vec![(b"x-proof".to_vec(), b"private\nvalue".to_vec())], + "invalid OCR header value", + ), + ] { + let result = send(Request { + headers, + ..request() + }) + .await; + assert!(matches!(result, Err(Error::InvalidRequest(message)) if message == expected)); + } + assert!( + matches!(send(request()).await, Err(Error::Network(message)) if message == "OCR transport failed") + ); +} diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index bda09a7d840..ab9abe4e3ae 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -7,7 +7,7 @@ repository.workspace = true [lib] name = "_native" -crate-type = ["cdylib"] +crate-type = ["cdylib", "rlib"] [features] default = ["abi3"] @@ -30,6 +30,7 @@ pyo3.workspace = true pyo3-async-runtimes.workspace = true serde.workspace = true serde_json.workspace = true +strum.workspace = true tokio.workspace = true [dev-dependencies] diff --git a/litellm-rust/crates/python-bridge/src/execution.rs b/litellm-rust/crates/python-bridge/src/execution.rs index f3648158cf6..5b423a32e7d 100644 --- a/litellm-rust/crates/python-bridge/src/execution.rs +++ b/litellm-rust/crates/python-bridge/src/execution.rs @@ -37,6 +37,37 @@ fn run_sync_on( where T: Serialize + Send + 'static, F: Future> + Send + 'static, +{ + let result = run_sync_value_on(py, runtime, future, map_error)?; + Pythonized(result).into_pyobject(py).map(Bound::unbind) +} + +pub(crate) fn run_sync_value( + py: Python<'_>, + future: F, + map_error: fn(Error) -> PyErr, +) -> PyResult +where + T: Send + 'static, + F: Future> + Send + 'static, +{ + run_sync_value_on( + py, + pyo3_async_runtimes::tokio::get_runtime(), + future, + map_error, + ) +} + +fn run_sync_value_on( + py: Python<'_>, + runtime: &Runtime, + future: F, + map_error: fn(Error) -> PyErr, +) -> PyResult +where + T: Send + 'static, + F: Future> + Send + 'static, { if Handle::try_current().is_ok() { return Err(PyRuntimeError::new_err( @@ -45,8 +76,7 @@ where } let result = release_gil(py, move || runtime.block_on(wait_for_sync_result(future)))?; - let result = map_core_result(result, map_error)?; - Pythonized(result).into_pyobject(py).map(Bound::unbind) + map_core_result(result, map_error) } pub(crate) fn run_async( @@ -75,7 +105,7 @@ fn map_core_result(result: Result, map_error: fn(Error) -> PyErr) - } } -async fn catch_future_panic(future: F) -> PyResult> +pub(crate) async fn catch_future_panic(future: F) -> PyResult> where F: Future>, { diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 384f0be5a1b..c00319eb83e 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -63,7 +63,7 @@ impl ResponsesWebSocketConnection { } #[pymodule(gil_used = false)] -mod _native { +pub mod _native { use pyo3::prelude::*; #[pymodule_init] @@ -98,6 +98,8 @@ mod tests { "RustUpstreamError", "ocr", "aocr", + "ocr_retained", + "aocr_retained", "transcription", "atranscription", "messages", diff --git a/litellm-rust/crates/python-bridge/src/routes/bindings.rs b/litellm-rust/crates/python-bridge/src/routes/bindings.rs new file mode 100644 index 00000000000..529bda63142 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/bindings.rs @@ -0,0 +1,157 @@ +use litellm_python_interop::InvocationMode; +use strum::{Display, EnumIter, IntoStaticStr}; + +/// A method name paired inseparably with the mode used to drive it. +/// +/// Keeping the two together is what prevents a `("prepare", Await)`-style +/// mismatch. Callers never assemble these by hand; they ask a catalog entry +/// ([`BoundaryMethod`] or [`LoggingMethod`]) to `resolve` one. +pub(crate) struct MethodBinding { + pub(crate) name: &'static str, + pub(crate) mode: InvocationMode, +} + +/// OCR route boundary methods: the Python operations invoked with +/// [`litellm_python_interop::PreparedCall`] to carry a request through +/// preparation, encoding, transport and finalization. +/// +/// The `asynchronous` flag drives both the method *name* and its invocation +/// mode together, so a sync/async pair cannot desynchronize. +#[derive(Debug, Clone, Copy, PartialEq, Eq, EnumIter, Display, IntoStaticStr)] +pub(crate) enum BoundaryMethod { + Prepare, + Encode, + Finish, +} + +impl BoundaryMethod { + pub(crate) fn resolve(self, asynchronous: bool) -> MethodBinding { + match (self, asynchronous) { + (Self::Prepare, true) => MethodBinding { + name: "aprepare", + mode: InvocationMode::Await, + }, + (Self::Prepare, false) => MethodBinding { + name: "prepare", + mode: InvocationMode::Direct, + }, + (Self::Encode, _) => MethodBinding { + name: "encode", + mode: InvocationMode::Direct, + }, + (Self::Finish, true) => MethodBinding { + name: "afinish", + mode: InvocationMode::Await, + }, + (Self::Finish, false) => MethodBinding { + name: "finish", + mode: InvocationMode::Direct, + }, + } + } +} + +/// Bound methods on the Python `Logging` object that the bridge invokes. +/// +/// Sync hooks (`pre_call`, `success_handler`, `failure_handler`) run inline +/// even on async routes; async hooks return a coroutine for the caller's loop +/// to drive. This mirrors the split LiteLLM keeps in Python between +/// `dynamic_success_callbacks` / `dynamic_async_success_callbacks`. +#[derive(Debug, Clone, Copy, PartialEq, Eq, EnumIter, Display, IntoStaticStr)] +pub(crate) enum LoggingMethod { + PreCall, + PostCall, + SuccessHandler, + FailureHandler, + AsyncSuccessHandler, + AsyncFailureHandler, +} + +// `LoggingMethod` is the forward-looking catalog for the retained-callback +// foundation. Nothing in production consumes it yet (the hooks them are +// exercised by the `component_contract` fixtures via a direct `PreparedCall`), +// so `resolve` is only reached from this module's tests for now. +#[cfg_attr(not(test), allow(dead_code))] +impl LoggingMethod { + pub(crate) fn resolve(self) -> MethodBinding { + match self { + Self::PreCall => MethodBinding { + name: "pre_call", + mode: InvocationMode::Direct, + }, + Self::PostCall => MethodBinding { + name: "post_call", + mode: InvocationMode::Direct, + }, + Self::SuccessHandler => MethodBinding { + name: "success_handler", + mode: InvocationMode::Direct, + }, + Self::FailureHandler => MethodBinding { + name: "failure_handler", + mode: InvocationMode::Direct, + }, + Self::AsyncSuccessHandler => MethodBinding { + name: "async_success_handler", + mode: InvocationMode::Await, + }, + Self::AsyncFailureHandler => MethodBinding { + name: "async_failure_handler", + mode: InvocationMode::Await, + }, + } + } +} + +#[cfg(test)] +mod tests { + use strum::IntoEnumIterator; + + use super::*; + + #[test] + fn boundary_methods_pair_name_and_mode_consistently() { + for method in BoundaryMethod::iter() { + let sync = method.resolve(false); + let asynchronous = method.resolve(true); + + assert!(!sync.name.is_empty()); + assert!(!asynchronous.name.is_empty()); + + if method == BoundaryMethod::Encode { + // `encode` has no async variant; it is always a direct call. + assert_eq!(sync.name, asynchronous.name); + assert_eq!(sync.mode, InvocationMode::Direct); + assert_eq!(asynchronous.mode, InvocationMode::Direct); + } else { + assert_eq!(sync.mode, InvocationMode::Direct); + assert_eq!(asynchronous.mode, InvocationMode::Await); + let async_name = asynchronous.name.strip_prefix('a').unwrap(); + assert_eq!(sync.name, async_name); + } + } + } + + #[test] + fn logging_methods_are_direct_unless_async() { + let mut names = Vec::new(); + for method in LoggingMethod::iter() { + let binding = method.resolve(); + assert!(!binding.name.is_empty()); + assert!(!names.contains(&binding.name), "duplicate method name"); + names.push(binding.name); + + let is_async = matches!( + method, + LoggingMethod::AsyncSuccessHandler | LoggingMethod::AsyncFailureHandler + ); + assert_eq!( + binding.mode == InvocationMode::Await, + is_async, + "{} must be {}", + binding.name, + if is_async { "Await" } else { "Direct" } + ); + } + } +} diff --git a/litellm-rust/crates/python-bridge/src/routes/mod.rs b/litellm-rust/crates/python-bridge/src/routes/mod.rs index 7e81f2ffe9b..05b11935956 100644 --- a/litellm-rust/crates/python-bridge/src/routes/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/mod.rs @@ -7,12 +7,15 @@ mod definition; mod gateway_messages; mod audio_transcription; +mod bindings; mod chat_completions; mod messages; mod ocr; +mod ocr_retained; pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { ocr::register(module)?; + ocr_retained::register(module)?; audio_transcription::register(module)?; messages::register(module)?; chat_completions::register(module)?; diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr_retained.rs b/litellm-rust/crates/python-bridge/src/routes/ocr_retained.rs new file mode 100644 index 00000000000..ec35d57418b --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/ocr_retained.rs @@ -0,0 +1,144 @@ +use litellm_core::ocr::transport::{self, Request, Response}; +use litellm_python_interop::{InvocationOutcome, PreparedCall}; +use pyo3::prelude::*; +use pyo3::sync::PyOnceLock; +use pyo3::types::{PyBytes, PyList, PyTuple}; + +use super::bindings::{BoundaryMethod, MethodBinding}; +use crate::errors::core_error_to_pyerr; +use crate::execution::{catch_future_panic, run_sync_value}; + +fn invoke( + boundary: &Bound<'_, PyAny>, + binding: MethodBinding, + args: Bound<'_, PyTuple>, +) -> PyResult> { + let call = PreparedCall::new( + binding.mode, + boundary.getattr(binding.name)?.unbind(), + args.unbind(), + None, + ); + match call.invoke(boundary.py())? { + InvocationOutcome::Returned(value) | InvocationOutcome::Awaitable(value) => Ok(value), + } +} + +#[pyfunction] +fn prepare(boundary: &Bound<'_, PyAny>, asynchronous: bool) -> PyResult> { + invoke( + boundary, + BoundaryMethod::Prepare.resolve(asynchronous), + PyTuple::empty(boundary.py()), + ) +} + +fn encode(boundary: &Bound<'_, PyAny>, roots: &Bound<'_, PyAny>) -> PyResult { + type ByteHeaders<'py> = Vec<(Bound<'py, PyBytes>, Bound<'py, PyBytes>)>; + let encoded = invoke( + boundary, + BoundaryMethod::Encode.resolve(false), + PyTuple::new(boundary.py(), [roots])?, + )?; + let (url, headers, body, timeout_seconds): (String, ByteHeaders<'_>, Bound<'_, PyBytes>, f64) = + encoded.extract(boundary.py())?; + Ok(Request { + url, + headers: headers + .into_iter() + .map(|(name, value)| (name.as_bytes().to_vec(), value.as_bytes().to_vec())) + .collect(), + body: body.as_bytes().to_vec(), + timeout_seconds, + }) +} + +struct Wire(Response); + +impl<'py> IntoPyObject<'py> for Wire { + type Target = PyTuple; + type Output = Bound<'py, PyTuple>; + type Error = PyErr; + + fn into_pyobject(self, py: Python<'py>) -> PyResult { + let headers = PyList::new( + py, + self.0 + .headers + .iter() + .map(|(name, value)| (PyBytes::new(py, name), PyBytes::new(py, value))), + )?; + (self.0.status, headers, PyBytes::new(py, &self.0.content)).into_pyobject(py) + } +} + +#[pyfunction] +fn send<'py>( + boundary: &Bound<'py, PyAny>, + roots: &Bound<'py, PyAny>, +) -> PyResult> { + let request = encode(boundary, roots)?; + pyo3_async_runtimes::tokio::future_into_py(boundary.py(), async move { + let response = catch_future_panic(transport::send(request)) + .await? + .map_err(core_error_to_pyerr)?; + Ok(Wire(response)) + }) +} + +#[pyfunction] +fn finish( + boundary: &Bound<'_, PyAny>, + wire: &Bound<'_, PyAny>, + asynchronous: bool, +) -> PyResult> { + invoke( + boundary, + BoundaryMethod::Finish.resolve(asynchronous), + PyTuple::new(boundary.py(), [wire])?, + ) +} + +#[pyfunction] +fn ocr_retained(boundary: &Bound<'_, PyAny>) -> PyResult> { + let py = boundary.py(); + let roots = prepare(boundary, false)?; + let request = encode(boundary, roots.bind(py))?; + let response = run_sync_value(py, transport::send(request), core_error_to_pyerr)?; + let wire = Wire(response).into_pyobject(py)?; + finish(boundary, wire.as_any(), false) +} + +#[pyfunction] +fn aocr_retained(boundary: &Bound<'_, PyAny>) -> PyResult> { + static DRIVER: PyOnceLock> = PyOnceLock::new(); + let py = boundary.py(); + let driver = DRIVER.get_or_try_init(py, || { + PyModule::from_code( + py, + c"async def drive(boundary, prepare, send, finish): + roots = await prepare(boundary, True) + wire = await send(boundary, roots) + return await finish(boundary, wire, True) +", + c"ocr_retained_driver.py", + c"_ocr_retained_driver", + )? + .getattr("drive") + .map(Bound::unbind) + })?; + driver.call1( + py, + ( + boundary, + wrap_pyfunction!(prepare, py)?, + wrap_pyfunction!(send, py)?, + wrap_pyfunction!(finish, py)?, + ), + ) +} + +pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { + module.add_function(wrap_pyfunction!(ocr_retained, module)?)?; + module.add_function(wrap_pyfunction!(aocr_retained, module)?) +} diff --git a/litellm-rust/crates/python-bridge/tests/ocr_retained.rs b/litellm-rust/crates/python-bridge/tests/ocr_retained.rs new file mode 100644 index 00000000000..c5e0408ea9b --- /dev/null +++ b/litellm-rust/crates/python-bridge/tests/ocr_retained.rs @@ -0,0 +1,243 @@ +use pyo3::prelude::*; +use pyo3::types::PyDict; + +#[test] +#[ignore = "requires repo Python"] +fn retained_real_production_boundary_differential_and_lifecycle() -> PyResult<()> { + Python::initialize(); + Python::attach(|py| { + let module = pyo3::wrap_pymodule!(_native::_native)(py).into_bound(py); + let globals = PyDict::new(py); + globals.set_item("native", module)?; + let fixture = std::ffi::CString::new(include_str!(concat!( + env!("CARGO_MANIFEST_DIR"), + "/../../../tests/test_litellm/ocr/retained_boundary_fixture.py" + )))?; + py.run(&fixture, Some(&globals), Some(&globals)) + }) +} + +#[test] +fn retained_routes_preserve_callbacks_context_wire_and_ownership() -> PyResult<()> { + Python::initialize(); + Python::attach(|py| { + let module = pyo3::wrap_pymodule!(_native::_native)(py).into_bound(py); + let globals = PyDict::new(py); + globals.set_item("native", module)?; + py.run( + cr" +import asyncio +import contextvars +import gc +import http.server +import inspect +import threading +import weakref + +started = threading.Event() +release = threading.Event() +requests = [] + +class Handler(http.server.BaseHTTPRequestHandler): + def log_message(self, *args): + pass + + def do_POST(self): + body = self.rfile.read(int(self.headers['Content-Length'])) + requests.append((self.path, self.headers.get_all('X-Proof'), body)) + if body == b'hold': + started.set() + release.wait(5) + self.send_response(429) + self.send_header('X-Reply', 'one') + self.send_header('X-Reply', 'two') + self.send_header('Content-Length', '3') + self.end_headers() + try: + self.wfile.write(b'\x00\xffR') + except (BrokenPipeError, ConnectionResetError): + pass + +server = http.server.ThreadingHTTPServer(('127.0.0.1', 0), Handler) +thread = threading.Thread(target=server.serve_forever, daemon=True) +thread.start() +url = 'http://127.0.0.1:%s/ocr' % server.server_port +context = contextvars.ContextVar('proof', default='unset') + +class Graph(dict): + pass + +class Boundary: + def __init__(self, *, asynchronous=False, nested=False, failure=None, hold=False): + self.asynchronous = asynchronous + self.nested = nested + self.failure = failure + self.hold = hold + self.events = [] + self.thread = threading.get_ident() + self.task = asyncio.current_task() if asynchronous else None + self.result = object() + self.error = LookupError('original callback error') + + def phase(self, name): + assert threading.get_ident() == self.thread + if self.asynchronous: + assert asyncio.current_task() is self.task + assert context.get() == ('initial' if name == 'prepare' else 'prepared') + self.events.append(name) + if self.failure == name: + raise self.error + if not self.nested and not self.hold: + child = Boundary(nested=True) + assert native.ocr_retained(child) is child.result + assert child.events == ['prepare', 'encode', 'finish'] + + def prepare(self): + self.phase('prepare') + headers = Graph({'X-Proof': 'original'}) + document = object() + body = Graph(document=document, alias=document) + body['cycle'] = body + self.refs = (weakref.ref(headers), weakref.ref(body)) + self.view = {'headers': headers, 'body': body} + headers['X-Proof'] = 'mutated' + self.view['headers'] = {'replacement': True} + self.view['body'] = {'replacement': True} + return (headers, url, body, None) + + async def aprepare(self): + await asyncio.sleep(0) + roots = self.prepare() + context.set('prepared') + return roots + + def encode(self, roots): + self.phase('encode') + headers, target, body, files = roots + assert headers is self.refs[0]() and body is self.refs[1]() + assert headers['X-Proof'] == 'mutated' + assert body['document'] is body['alias'] and body['cycle'] is body + assert files is None + assert self.view == {'headers': {'replacement': True}, 'body': {'replacement': True}} + return (target, [(b'X-Proof', b'mutated'), (b'X-Proof', b'duplicate')], + b'hold' if self.hold else b'\x00\xffQ', 3.0) + + def finish(self, wire): + self.phase('finish') + assert type(wire) is tuple and len(wire) == 3 + status, headers, content = wire + assert status == 429 + assert type(headers) is list + assert all(type(pair) is tuple and all(type(v) is bytes for v in pair) for pair in headers) + assert [v for k, v in headers if k == b'x-reply'] == [b'one', b'two'] + assert type(content) is bytes and content == b'\x00\xffR' + assert all(ref() is not None for ref in self.refs) + return self.result + + async def afinish(self, wire): + await asyncio.sleep(0) + return self.finish(wire) + +def collected(boundary): + gc.collect() + assert all(ref() is None for ref in boundary.refs) + assert boundary.view == {'headers': {'replacement': True}, 'body': {'replacement': True}} + +def check_error(boundary, error, phase): + assert error is boundary.error + names = [] + traceback = error.__traceback__ + while traceback: + names.append(traceback.tb_frame.f_code.co_name) + traceback = traceback.tb_next + assert phase in names and 'phase' in names + assert boundary.events == ['prepare', 'encode', 'finish'][:['prepare', 'encode', 'finish'].index(phase) + 1] + +async def exercise(): + context.set('initial') + boundary = Boundary(asynchronous=True) + pending = native.aocr_retained(boundary) + assert inspect.iscoroutine(pending) + assert boundary.events == [] + assert await pending is boundary.result + assert context.get() == 'prepared' + assert boundary.events == ['prepare', 'encode', 'finish'] + collected(boundary) + + unused = Boundary(asynchronous=True) + ref = weakref.ref(unused) + pending = native.aocr_retained(unused) + assert unused.events == [] + del unused + assert ref() is not None + pending.close() + del pending + gc.collect() + assert ref() is None + + for phase in ('prepare', 'encode', 'finish'): + context.set('initial') + boundary = Boundary(asynchronous=True, nested=True, failure=phase) + try: + await native.aocr_retained(boundary) + except LookupError as error: + check_error(boundary, error, phase) + else: + raise AssertionError('callback error was swallowed') + boundary.error.__traceback__ = None + if phase != 'prepare': + collected(boundary) + + context.set('initial') + boundary = Boundary(asynchronous=True, hold=True) + async def cancellable(): + boundary.task = asyncio.current_task() + await native.aocr_retained(boundary) + task = asyncio.create_task(cancellable()) + assert await asyncio.to_thread(started.wait, 2) + gc.collect() + assert all(ref() is not None for ref in boundary.refs) + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + else: + raise AssertionError('cancellation was swallowed') + assert boundary.events == ['prepare', 'encode'] + del task + await asyncio.sleep(0) + collected(boundary) + release.set() + +try: + boundary = Boundary() + assert native.ocr_retained(boundary) is boundary.result + assert boundary.events == ['prepare', 'encode', 'finish'] + collected(boundary) + for phase in ('prepare', 'encode', 'finish'): + boundary = Boundary(nested=True, failure=phase) + try: + native.ocr_retained(boundary) + except LookupError as error: + check_error(boundary, error, phase) + else: + raise AssertionError('callback error was swallowed') + boundary.error.__traceback__ = None + if phase != 'prepare': + collected(boundary) + asyncio.run(asyncio.wait_for(exercise(), 15)) + assert requests + assert all(path == '/ocr' and headers == ['mutated', 'duplicate'] and body in (b'\x00\xffQ', b'hold') + for path, headers, body in requests) +finally: + release.set() + server.shutdown() + server.server_close() + thread.join(5) +", + Some(&globals), + Some(&globals), + ) + }) +} diff --git a/litellm-rust/crates/python-interop/src/callback.rs b/litellm-rust/crates/python-interop/src/callback.rs index 6c73f7c9a5d..bc22195f0d0 100644 --- a/litellm-rust/crates/python-interop/src/callback.rs +++ b/litellm-rust/crates/python-interop/src/callback.rs @@ -5,7 +5,7 @@ use pyo3::types::{PyDict, PyTuple}; static AWAIT_CALL: PyOnceLock> = PyOnceLock::new(); -#[derive(Clone, Copy)] +#[derive(Clone, Copy, Debug, PartialEq, Eq)] pub enum InvocationMode { Direct, Await, diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index df3f9d2096b..2eeef48cb8d 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -31,6 +31,7 @@ from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.rust_bridge import ocr as rust_ocr_bridge from litellm.rust_bridge.bindings import native_exception_types from litellm.rust_bridge.configuration import rust_enabled +from litellm.rust_bridge import ocr_retained as rust_ocr_retained_bridge from litellm.types.router import GenericLiteLLMParams from litellm.utils import ProviderConfigManager, client @@ -309,10 +310,42 @@ def _map_rust_ocr_error( ) +def _retained_ocr_boundary(prepared: _PreparedOCRRequest) -> rust_ocr_retained_bridge.OCRRetainedBoundary: + return rust_ocr_retained_bridge.OCRRetainedBoundary( + handler=base_llm_http_handler, + model=prepared.model, + document=prepared.document, + optional_params=prepared.optional_params, + logging_obj=prepared.litellm_logging_obj, + api_key=prepared.api_key, + api_base=prepared.api_base, + headers=prepared.extra_headers, + provider_config=prepared.provider_config, + litellm_params=prepared.litellm_params, + custom_llm_provider=prepared.custom_llm_provider, + timeout=prepared.effective_timeout, + ) + + def _run_rust_ocr( prepared_request: _PreparedOCRRequest, resolve_api_key: Callable[[str], str | None], + *, + load_retained: Callable[[], rust_ocr_retained_bridge.RustOCRRetained | None] = ( + rust_ocr_retained_bridge.load_rust_ocr_retained + ), ) -> OCRResponse | None: + if prepared_request.custom_llm_provider == "mistral" and not rust_ocr_bridge.has_rust_ocr_override(): + retained: Final = load_retained() + if ( + retained is None + or rust_ocr_retained_bridge.retained_timeout_seconds(prepared_request.effective_timeout) is None + ): + return None + retained_response: Final = retained(_retained_ocr_boundary(prepared_request)) + if retained_response is None: + raise ValueError("Retained OCR returned no response after preparation") + return retained_response if rust_ocr_bridge.load_rust_ocr() is None: return None prepared: Final = _prepare_rust_ocr_call( @@ -340,7 +373,24 @@ def _run_rust_ocr( async def _run_rust_aocr( prepared_request: _PreparedOCRRequest, resolve_api_key: Callable[[str], str | None], + *, + load_retained: Callable[[], rust_ocr_retained_bridge.RustAOCRRetained | None] = ( + rust_ocr_retained_bridge.load_rust_aocr_retained + ), ) -> OCRResponse | None: + if prepared_request.custom_llm_provider == "mistral" and not rust_ocr_bridge.has_rust_ocr_override( + asynchronous=True + ): + retained: Final = load_retained() + if ( + retained is None + or rust_ocr_retained_bridge.retained_timeout_seconds(prepared_request.effective_timeout) is None + ): + return None + retained_response: Final = await retained(_retained_ocr_boundary(prepared_request)) + if retained_response is None: + raise ValueError("Retained OCR returned no response after preparation") + return retained_response if rust_ocr_bridge.load_rust_aocr() is None: return None prepared: Final = _prepare_rust_ocr_call( diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index b7fdb5a98ef..05c09528df6 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -53,6 +53,10 @@ _OCR: Final = NativeBinding("ocr", validate=_as_ocr) _AOCR: Final = NativeBinding("aocr", validate=_as_aocr) +def has_rust_ocr_override(*, asynchronous: bool = False) -> bool: + return _rust_aocr_impl is not None if asynchronous else _rust_ocr_impl is not None + + def load_rust_ocr() -> RustOcr | None: return _OCR.load() diff --git a/litellm/rust_bridge/ocr_retained.py b/litellm/rust_bridge/ocr_retained.py new file mode 100644 index 00000000000..cc431aabade --- /dev/null +++ b/litellm/rust_bridge/ocr_retained.py @@ -0,0 +1,146 @@ +"""Python preparation, encoding, and response transforms for retained OCR calls.""" + +from __future__ import annotations + +import math +from collections.abc import Awaitable, Callable +from dataclasses import dataclass, field +from typing import Final, cast # noqa: TID251 # native callables and legacy header types require boundary casts + +import httpx + +import litellm +from litellm.litellm_core_utils.litellm_logging import Logging +from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig, DocumentType, OCRResponse +from litellm.llms.custom_httpx.http_handler import ( + AsyncHTTPHandler, + HTTPHandler, + _get_httpx_client, + get_async_httpx_client, +) +from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler +from litellm.rust_bridge.timeouts import timeout_to_seconds + +OCRRoots = tuple[dict[str, object], str, dict[str, object], None] +OCRWire = tuple[int, list[tuple[bytes, bytes]], bytes] +OCREncoded = tuple[str, list[tuple[bytes, bytes]], bytes, float] + + +def retained_timeout_seconds(timeout: float | httpx.Timeout) -> float | None: + seconds: Final = timeout_to_seconds(timeout) + return seconds if seconds is not None and math.isfinite(seconds) and seconds > 0 else None + + +@dataclass(kw_only=True, slots=True) +class OCRRetainedBoundary: + handler: BaseLLMHTTPHandler + model: str + document: DocumentType + optional_params: dict[str, object] + logging_obj: Logging + api_key: str | None + api_base: str | None + headers: dict[str, object] | None + provider_config: BaseOCRConfig + litellm_params: dict[str, object] + custom_llm_provider: str + timeout: float | httpx.Timeout + client: HTTPHandler | AsyncHTTPHandler | None = None + request: httpx.Request | None = field(default=None, init=False) + + def prepare(self) -> OCRRoots: + roots: Final = self.handler._prepare_ocr_request( + model=self.model, + document=self.document, + optional_params=self.optional_params, + logging_obj=self.logging_obj, + api_key=self.api_key, + api_base=self.api_base, + headers=self.headers, + provider_config=self.provider_config, + litellm_params=self.litellm_params, + ) + if not isinstance(self.client, HTTPHandler): + self.client = _get_httpx_client() + return roots + + async def aprepare(self) -> OCRRoots: + roots: Final = await self.handler._async_prepare_ocr_request( + model=self.model, + document=self.document, + optional_params=self.optional_params, + logging_obj=self.logging_obj, + api_key=self.api_key, + api_base=self.api_base, + headers=self.headers, + provider_config=self.provider_config, + litellm_params=self.litellm_params, + ) + if not isinstance(self.client, AsyncHTTPHandler): + self.client = get_async_httpx_client(llm_provider=litellm.LlmProviders(self.custom_llm_provider)) + return roots + + def encode(self, roots: OCRRoots) -> OCREncoded: + headers, url, data, _files = roots + seconds: Final = retained_timeout_seconds(self.timeout) + if seconds is None: + raise ValueError("Retained OCR requires a positive finite read timeout") + if self.client is None: + raise RuntimeError("Retained OCR must be prepared before encoding") + try: + self.request = self.client.client.build_request( + "POST", + url, + headers=cast(dict[str, str], headers), + json=data, + timeout=self.timeout, + ) + return str(self.request.url), self.request.headers.raw, self.request.read(), seconds + except Exception as e: # noqa: BLE001 # match the Python OCR handler's encoding error mapping + raise self.handler._handle_error(e=e, provider_config=self.provider_config) + + def _response(self, wire: OCRWire) -> httpx.Response: + if self.request is None: + raise RuntimeError("Retained OCR must be encoded before finishing") + status, headers, content = wire + try: + response: Final = httpx.Response(status, headers=headers, content=content, request=self.request) + response.raise_for_status() + except Exception as e: # noqa: BLE001 # match the Python OCR handler's response error mapping + raise self.handler._handle_error(e=e, provider_config=self.provider_config) + return response + + def finish(self, wire: OCRWire) -> OCRResponse: + return self.handler._transform_ocr_response( + provider_config=self.provider_config, + model=self.model, + response=self._response(wire), + logging_obj=self.logging_obj, + optional_params=self.optional_params, + ) + + async def afinish(self, wire: OCRWire) -> OCRResponse: + return await self.provider_config.async_transform_ocr_response( + model=self.model, + raw_response=self._response(wire), + logging_obj=self.logging_obj, + optional_params=self.optional_params, + ) + + +RustOCRRetained = Callable[[OCRRetainedBoundary], OCRResponse] +RustAOCRRetained = Callable[[OCRRetainedBoundary], Awaitable[OCRResponse]] + + +def load_rust_ocr_retained() -> RustOCRRetained | None: + from litellm.rust_bridge import get_native_bridge + + native: Final = get_native_bridge() + return cast(RustOCRRetained | None, getattr(native, "ocr_retained", None)) + + +def load_rust_aocr_retained() -> RustAOCRRetained | None: + from litellm.rust_bridge import get_native_bridge + + native: Final = get_native_bridge() + return cast(RustAOCRRetained | None, getattr(native, "aocr_retained", None)) diff --git a/tests/test_litellm/ocr/retained_boundary_fixture.py b/tests/test_litellm/ocr/retained_boundary_fixture.py new file mode 100644 index 00000000000..c0561e1b88a --- /dev/null +++ b/tests/test_litellm/ocr/retained_boundary_fixture.py @@ -0,0 +1,642 @@ +"""Executed by the PyO3 retained test with the actual built module in `native`.""" + +import asyncio +import contextvars +import gc +import importlib +import inspect +import json +import threading +import unittest +import weakref +from copy import deepcopy +from datetime import datetime +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from types import ModuleType +from unittest.mock import Mock, patch + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.litellm_logging import Logging +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler +from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler +from litellm.llms.mistral.ocr.transformation import MistralOCRConfig +from litellm.rust_bridge.ocr_retained import OCREncoded, OCRRetainedBoundary, OCRRoots + +native: ModuleType = globals()["native"] +ocr_main = importlib.import_module("litellm.ocr.main") +context = contextvars.ContextVar("retained-real-boundary", default="unset") +MODEL = "mistral-ocr-latest" +RESPONSE = {"pages": [{"index": 0, "markdown": "local OCR"}], "model": MODEL, "usage_info": {"pages_processed": 1}} + + +class Graph(dict): + pass + + +class Header(str): + pass + + +class PreCallAbort(BaseException): + pass + + +class CopiedDocumentBoundary(OCRRetainedBoundary): + def prepare(self) -> OCRRoots: + self.document = deepcopy(self.document) + return super().prepare() + + async def aprepare(self) -> OCRRoots: + self.document = deepcopy(self.document) + return await super().aprepare() + + +class ReboundBodyBoundary(OCRRetainedBoundary): + def encode(self, roots: OCRRoots) -> OCREncoded: + headers, url, _body, files = roots + view = self.logging_obj.model_call_details["additional_args"] + return super().encode((headers, url, view["complete_input_dict"], files)) + + +class ReboundHeadersBoundary(OCRRetainedBoundary): + def encode(self, roots: OCRRoots) -> OCREncoded: + _headers, url, body, files = roots + view = self.logging_obj.model_call_details["additional_args"] + return super().encode((view["headers"], url, body, files)) + + +class Callback(CustomLogger): + def __init__(self, action): + super().__init__() + self.action = action + self.failures = [] + self.calls = 0 + + def log_pre_api_call(self, model, messages, kwargs): + self.calls += 1 + try: + return self.action(kwargs["additional_args"]) + except AssertionError as error: + self.failures.append(str(error)) + raise + + +class Server: + def __init__(self): + self.requests = [] + self.started = threading.Event() + self.release = threading.Event() + self.finished = threading.Event() + owner = self + + class Handler(BaseHTTPRequestHandler): + def log_message(self, *args): + pass + + def do_POST(self): + body = self.rfile.read(int(self.headers["Content-Length"])) + owner.requests.append((self.path, sorted((k.lower(), v) for k, v in self.headers.items()), body)) + if self.path == "/blocked/v1/ocr": + owner.started.set() + if not owner.release.wait(10): + owner.finished.set() + return + failed = self.path == "/error/v1/ocr" + payload = b'{"message":"controlled HTTP failure"}' if failed else json.dumps(RESPONSE).encode() + self.send_response(429 if failed else 200) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(payload))) + self.end_headers() + try: + self.wfile.write(payload) + except (BrokenPipeError, ConnectionResetError): + pass + finally: + if self.path == "/blocked/v1/ocr": + owner.finished.set() + + self.httpd = ThreadingHTTPServer(("127.0.0.1", 0), Handler) + self.thread = threading.Thread(target=self.httpd.serve_forever, daemon=True) + self.thread.start() + self.url = f"http://127.0.0.1:{self.httpd.server_port}" + + def close(self): + self.release.set() + self.httpd.shutdown() + self.httpd.server_close() + self.thread.join(5) + assert not self.thread.is_alive() + + +def inputs(server, callbacks=(), *, document=None, optional=None, path="", client=None): + return { + "model": MODEL, + "document": Graph(type="document_url", document_url="https://example.test/original.pdf") + if document is None + else document, + "optional_params": {} if optional is None else optional, + "logging_obj": Logging( + model=MODEL, + messages=[], + stream=False, + call_type="ocr", + start_time=datetime.now(), + litellm_call_id="retained-real", + function_id="retained-real", + dynamic_input_callbacks=list(callbacks), + supports_correlation_logging=False, + ), + "api_key": "local-test-key", + "api_base": server.url + path, + "headers": {"X-Proof": "original"}, + "provider_config": MistralOCRConfig(), + "litellm_params": {}, + "custom_llm_provider": "mistral", + "timeout": 5.0, + "client": client, + } + + +def invoke(mode, kwargs, *, boundary_factory=OCRRetainedBoundary): + handler = BaseLLMHTTPHandler() + if mode == "python-sync": + return handler.ocr(**kwargs) + if mode == "python-async": + return handler.async_ocr(**kwargs) + if mode == "native-sync": + return native.ocr_retained(boundary_factory(handler=handler, **kwargs)) + if mode == "native-async": + return native.aocr_retained(boundary_factory(handler=handler, **kwargs)) + prepared = ocr_main._PreparedOCRRequest( + **{ + k: kwargs[k] + for k in ( + "model", + "document", + "api_key", + "api_base", + "custom_llm_provider", + "provider_config", + "optional_params", + "litellm_params", + ) + }, + extra_headers=kwargs["headers"], + effective_timeout=kwargs["timeout"], + litellm_logging_obj=kwargs["logging_obj"], + ) + if mode == "sdk-sync": + return ocr_main._run_rust_ocr(prepared, lambda _: None, load_retained=lambda: native.ocr_retained) + assert mode == "sdk-async" + return ocr_main._run_rust_aocr(prepared, lambda _: None, load_retained=lambda: native.aocr_retained) + + +class RealBoundaryTests(unittest.TestCase): + def setUp(self): + self.server = Server() + self.addCleanup(self.server.close) + self.sync_client = HTTPHandler(timeout=5.0) + self.addCleanup(self.sync_client.close) + + def check_callbacks(self, callbacks): + for callback in callbacks: + self.assertEqual(callback.failures, []) + self.assertEqual(callback.calls, 1) + + async def differential(self, mode, *, boundary_factory=OCRRetainedBoundary): + context.set("caller") + thread = threading.get_ident() + task = asyncio.current_task() + document = Graph(type="document_url", document_url="https://example.test/original.pdf") + nested = Graph(values=[1]) + optional = {"unknown_python_json": {7: ("tuple", 2)}, "nested": nested} + retained = {} + events = [] + + def phase(name, expected): + self.assertEqual(threading.get_ident(), thread) + self.assertIs(asyncio.current_task(), task) + self.assertEqual(context.get(), expected) + events.append(name) + + def mutate(view): + phase("mutate", "caller") + body, headers = view["complete_input_dict"], view["headers"] + self.assertIs(body["document"], document, "caller document identity was not retained") + self.assertIs(body["nested"], nested) + self.assertIs(body["unknown_python_json"], optional["unknown_python_json"]) + retained.update(body=body, headers=headers, view=view) + headers["X-Proof"] = "in-place" + body["body_mutation"] = True + nested["values"].append(2) + view["headers"] = {"X-Proof": "must-not-send"} + view["complete_input_dict"] = {"document": {"document_url": "must-not-send"}} + context.set("mutated") + child_callback = Callback(lambda _: events.append("reentry")) + child = invoke("native-sync", inputs(self.server, [child_callback], path="/child", client=self.sync_client)) + self.assertEqual(child.pages[0].markdown, "local OCR") + self.check_callbacks([child_callback]) + return {"headers": {"X-Proof": "ignored-return"}, "complete_input_dict": {"invalid": object()}} + + def mutate_then_raise(view): + phase("raise", "mutated") + retained["body"]["before_error"] = True + retained["headers"]["X-Before-Error"] = "yes" + context.set("caught") + raise RuntimeError("intentional non-blocking pre_call error") + + def closure_only(view): + phase("later", "caught") + self.assertIsNot(view["complete_input_dict"], retained["body"]) + self.assertIsNot(view["headers"], retained["headers"]) + document["document_url"] = "https://example.test/closure.pdf" + context.set("later") + + callbacks = [Callback(action) for action in (mutate, mutate_then_raise, closure_only)] + async_client = AsyncHTTPHandler(timeout=5.0) + try: + kwargs = inputs( + self.server, + callbacks, + document=document, + optional=optional, + client=async_client if mode.endswith("async") else self.sync_client, + ) + before = len(self.server.requests) + pending = invoke(mode, kwargs, boundary_factory=boundary_factory) + if mode.endswith("async"): + self.assertTrue(inspect.iscoroutine(pending)) + self.assertEqual(events, []) + self.assertEqual(len(self.server.requests), before) + self.assertNotIn("additional_args", kwargs["logging_obj"].model_call_details) + response = await pending + else: + response = pending + self.check_callbacks(callbacks) + self.assertEqual(events, ["mutate", "reentry", "raise", "later"]) + self.assertEqual(context.get(), "later") + self.assertEqual(len(self.server.requests), before + 2) + wire = self.server.requests[-1] + self.assertEqual(wire[0], "/v1/ocr") + expected = ( + b'{"model":"mistral-ocr-latest","document":{"type":"document_url",' + b'"document_url":"https://example.test/closure.pdf"},"unknown_python_json":{"7":["tuple",2]},' + b'"nested":{"values":[1,2]},"body_mutation":true,"before_error":true}' + ) + self.assertEqual(wire[2], expected, "wire body must encode retained mutations after pre_call") + self.assertIn(("x-proof", "in-place"), wire[1], "wire headers must use retained execution headers") + self.assertIn(("x-before-error", "yes"), wire[1]) + self.assertIn(("authorization", "Bearer local-test-key"), wire[1]) + self.assertIs(retained["body"]["document"], document) + retained["body"]["nested"]["values"].append(3) + retained["headers"]["X-After-Return"] = "usable" + self.assertEqual(optional["nested"]["values"], [1, 2, 3]) + self.assertEqual(wire[2], expected) + self.assertEqual(response.pages[0].markdown, "local OCR") + return wire, response.model_dump() + finally: + await async_client.close() + + def test_differential_callbacks_wire_and_sdk_dispatch(self): + async def exercise(): + baseline = await self.differential("python-sync") + for mode in ("python-async", "native-sync", "native-async", "sdk-sync", "sdk-async"): + with self.subTest(mode=mode): + self.assertEqual(await self.differential(mode), baseline) + + asyncio.run(exercise()) + + def check_negative_control(self, boundary_factory, failure, expected_requests): + async def exercise(): + baseline = await self.differential("python-sync") + for mode in ("native-sync", "native-async"): + with self.subTest(mode=mode): + symbol = "aocr_retained" if mode.endswith("async") else "ocr_retained" + self.assertTrue(inspect.isbuiltin(getattr(native, symbol))) + self.assertEqual(await self.differential(mode), baseline) + before = len(self.server.requests) + with self.assertRaisesRegex(AssertionError, failure): + await self.differential(mode, boundary_factory=boundary_factory) + self.assertEqual(len(self.server.requests), before + expected_requests) + self.assertEqual(self.server.requests[-1][0], "/v1/ocr") + + asyncio.run(exercise()) + + def test_negative_control_copied_caller_document(self): + self.check_negative_control(CopiedDocumentBoundary, "caller document identity was not retained", 1) + + def test_negative_control_rebound_logging_body(self): + self.check_negative_control(ReboundBodyBoundary, "wire body must encode retained mutations after pre_call", 2) + + def test_negative_control_rebound_logging_headers(self): + self.check_negative_control(ReboundHeadersBoundary, "wire headers must use retained execution headers", 2) + + def test_differential_callback_retained_mutation_after_post_received(self): + async def suspended(mode): + self.server.started.clear() + self.server.release.clear() + self.server.finished.clear() + retained = {} + + def retain(view): + retained.update(body=view["complete_input_dict"], headers=view["headers"], view=view) + view["complete_input_dict"] = {"replacement": True} + view["headers"] = {"X-Proof": "must-not-send"} + + callback = Callback(retain) + client = AsyncHTTPHandler(timeout=5.0) + kwargs = inputs(self.server, [callback], path="/blocked", client=client) + logging = kwargs["logging_obj"] + logging.log_raw_request_response = True + before = len(self.server.requests) + task = asyncio.create_task(invoke(mode, kwargs)) + try: + self.assertTrue(await asyncio.to_thread(self.server.started.wait, 5)) + self.assertFalse(task.done()) + self.assertFalse(self.server.release.is_set()) + self.assertFalse(self.server.finished.is_set()) + self.check_callbacks([callback]) + self.assertEqual(len(self.server.requests), before + 1) + path, headers, body = self.server.requests[-1] + wire = (path, tuple(headers), body) + expected = ( + b'{"model":"mistral-ocr-latest","document":{"type":"document_url",' + b'"document_url":"https://example.test/original.pdf"}}' + ) + self.assertEqual(path, "/blocked/v1/ocr") + self.assertEqual(body, expected) + self.assertIn(("x-proof", "original"), headers) + self.assertNotIn(("x-proof", "must-not-send"), headers) + logged_body = logging.model_call_details["raw_request_typed_dict"]["raw_request_body"] + self.assertIs(logged_body, retained["body"]) + self.assertIs(logged_body["document"], kwargs["document"]) + self.assertIs(logging.model_call_details["additional_args"], retained["view"]) + self.assertIsNot(retained["view"]["complete_input_dict"], retained["body"]) + self.assertIsNot(retained["view"]["headers"], retained["headers"]) + + retained["body"]["after_encoding"] = True + retained["body"]["document"]["document_url"] = "https://example.test/while-blocked.pdf" + retained["headers"]["X-Proof"] = "while-blocked" + self.assertTrue(logged_body["after_encoding"]) + self.assertEqual(kwargs["document"]["document_url"], "https://example.test/while-blocked.pdf") + self.assertEqual(logged_body["document"]["document_url"], "https://example.test/while-blocked.pdf") + self.assertEqual(retained["headers"]["X-Proof"], "while-blocked") + self.assertEqual(retained["view"]["complete_input_dict"], {"replacement": True}) + self.assertEqual(retained["view"]["headers"], {"X-Proof": "must-not-send"}) + self.assertFalse(task.done()) + self.assertFalse(self.server.finished.is_set()) + self.assertEqual((path, tuple(headers), body), wire) + + self.server.release.set() + response = await asyncio.wait_for(task, 5) + self.assertEqual(response.pages[0].markdown, "local OCR") + self.assertIs(logging.model_call_details["raw_request_typed_dict"]["raw_request_body"], logged_body) + self.assertTrue(logged_body["after_encoding"]) + self.assertEqual(len(self.server.requests), before + 1) + received_path, received_headers, received_body = self.server.requests[-1] + self.assertEqual((received_path, tuple(received_headers), received_body), wire) + return wire, logged_body, retained["headers"], response.model_dump() + finally: + self.server.release.set() + if not task.done(): + task.cancel() + await asyncio.gather(task, return_exceptions=True) + await client.close() + if self.server.started.is_set(): + self.assertTrue(await asyncio.to_thread(self.server.finished.wait, 5)) + + async def exercise(): + baseline = await suspended("python-async") + self.assertEqual(await suspended("native-async"), baseline) + + asyncio.run(exercise()) + + def test_public_rust_dispatch_wire_fallback_and_escaping_base_exception(self): + async def exercise(): + for asynchronous in (False, True): + for outcome in ("success", "missing-symbol", "pre-call-abort"): + with self.subTest(asynchronous=asynchronous, outcome=outcome): + symbol = "aocr_retained" if asynchronous else "ocr_retained" + self.assertTrue(inspect.isbuiltin(getattr(native, symbol))) + missing_symbol = ModuleType("native_without_" + symbol) + missing_symbol.__dict__.update( + (name, value) for name, value in vars(native).items() if name != symbol + ) + handler = Mock(wraps=BaseLLMHTTPHandler()) + escaped = PreCallAbort("public pre_call must escape unchanged") + before = len(self.server.requests) + + def mutate(view): + loader.assert_called_once_with() + self.assertEqual(handler.ocr.call_count, int(outcome == "missing-symbol")) + view["headers"]["X-Proof"] = "public-in-place" + view["complete_input_dict"]["public_mutation"] = True + view["complete_input_dict"]["document"]["document_url"] = ( + "https://example.test/public-mutated.pdf" + ) + if outcome == "pre-call-abort": + raise escaped + + callback = Callback(mutate) + kwargs = inputs(self.server, [callback]) + with ( + patch( + "litellm.rust_bridge.get_native_bridge", + return_value=missing_symbol if outcome == "missing-symbol" else native, + ) as loader, + patch.object(ocr_main, "base_llm_http_handler", handler), + ): + public_kwargs = { + **{ + key: kwargs[key] + for key in ( + "model", + "document", + "api_key", + "api_base", + "custom_llm_provider", + "timeout", + ) + }, + "extra_headers": kwargs["headers"], + "litellm_logging_obj": kwargs["logging_obj"], + "rust": True, + } + + async def call(): + if asynchronous: + return await litellm.aocr(**public_kwargs) + return litellm.ocr(**public_kwargs) + + if outcome == "pre-call-abort": + with self.assertRaises(PreCallAbort) as caught: + await call() + self.assertIs(caught.exception, escaped) + else: + response = await call() + self.assertEqual(response.pages[0].markdown, "local OCR") + + loader.assert_called_once_with() + self.check_callbacks([callback]) + self.assertEqual(handler.ocr.call_count, int(outcome == "missing-symbol")) + self.assertEqual(handler.async_ocr.call_count, 0) + prepare = handler._async_prepare_ocr_request if asynchronous else handler._prepare_ocr_request + if outcome == "missing-symbol": + prepare.assert_not_called() + self.assertIs(handler.ocr.call_args.kwargs["logging_obj"], kwargs["logging_obj"]) + self.assertEqual(handler.ocr.call_args.kwargs["aocr"], asynchronous) + else: + prepare.assert_called_once() + self.assertIs(prepare.call_args.kwargs["logging_obj"], kwargs["logging_obj"]) + unused_prepare = ( + handler._prepare_ocr_request if asynchronous else handler._async_prepare_ocr_request + ) + unused_prepare.assert_not_called() + self.assertEqual(len(self.server.requests), before + int(outcome != "pre-call-abort")) + if outcome != "pre-call-abort": + path, headers, body = self.server.requests[-1] + self.assertEqual(path, "/v1/ocr") + self.assertIn(("x-proof", "public-in-place"), headers) + self.assertIn(("authorization", "Bearer local-test-key"), headers) + self.assertEqual( + json.loads(body), + { + "model": MODEL, + "document": { + "type": "document_url", + "document_url": "https://example.test/public-mutated.pdf", + }, + "public_mutation": True, + }, + ) + + asyncio.run(exercise()) + + def lifecycle_inputs(self, outcome, retained): + refs = [] + + def callback(view): + body, headers = view["complete_input_dict"], view["headers"] + body["sentinel"] = Graph(alive=True) + headers["X-Sentinel"] = Header("alive") + refs.extend((weakref.ref(body["sentinel"]), weakref.ref(headers["X-Sentinel"]))) + if outcome == "encoding": + body["not_json"] = object() + if retained is not None: + retained.extend((body, headers)) + view["complete_input_dict"] = {} + view["headers"] = {} + + logger = Callback(callback) + kwargs = inputs( + self.server, + [logger], + optional={"nested": Graph(alive=True)}, + path={"http": "/error", "cancel": "/blocked"}.get(outcome, ""), + client=self.sync_client, + ) + refs.extend(weakref.ref(kwargs[key]) for key in ("document", "logging_obj")) + refs.append(weakref.ref(kwargs["optional_params"]["nested"])) + return kwargs, logger, refs + + async def lifecycle(self, mode, outcome, retained=None): + kwargs, logger, refs = self.lifecycle_inputs(outcome, retained) + before = len(self.server.requests) + async_client = AsyncHTTPHandler(timeout=5.0) + if mode.endswith("async"): + kwargs["client"] = async_client + try: + try: + pending = invoke(mode, kwargs) + response = await pending if mode.endswith("async") else pending + except BaseLLMException as error: + self.assertIn(outcome, ("encoding", "http")) + self.assertEqual(error.status_code, 500 if outcome == "encoding" else 429) + signature = (type(error), error.status_code, str(error)) + else: + self.assertEqual(outcome, "success") + self.assertEqual(response.pages[0].markdown, "local OCR") + signature = response.model_dump() + self.check_callbacks([logger]) + self.assertEqual(len(refs), 5) + self.assertEqual(len(self.server.requests), before + (outcome != "encoding")) + return refs, signature + finally: + await async_client.close() + + def test_collection_after_success_encoding_failure_and_http_error(self): + async def exercise(): + for outcome in ("success", "encoding", "http"): + baseline = None + for mode in ("python-sync", "python-async", "native-sync", "native-async"): + with self.subTest(mode=mode, outcome=outcome): + refs, signature = await self.lifecycle(mode, outcome) + gc.collect() + self.assertTrue(all(ref() is None for ref in refs), "request graph leaked after " + outcome) + if baseline is None: + baseline = signature + self.assertEqual(signature, baseline) + + asyncio.run(exercise()) + + def test_callback_retained_graph_remains_usable_then_collects(self): + async def exercise(): + for mode in ("native-sync", "native-async"): + with self.subTest(mode=mode): + retained = [] + refs, _ = await self.lifecycle(mode, "success", retained) + gc.collect() + self.assertIsNone(refs[1]()) + self.assertTrue(all(refs[index]() is not None for index in (0, 2, 3, 4))) + retained[0]["document"]["after_return"] = "usable" + retained[0]["sentinel"]["alive"] = "still usable" + retained[1]["X-After-Return"] = "usable" + self.assertEqual(refs[0]()["after_return"], "usable") + self.assertEqual(refs[3]()["alive"], "still usable") + retained.clear() + gc.collect() + self.assertTrue(all(ref() is None for ref in refs)) + + asyncio.run(exercise()) + + def test_collection_after_cancellation_during_blocked_transport(self): + async def cancel(): + kwargs, logger, refs = self.lifecycle_inputs("cancel", None) + client = AsyncHTTPHandler(timeout=5.0) + kwargs["client"] = client + task = asyncio.create_task(invoke("native-async", kwargs)) + try: + self.assertTrue(await asyncio.to_thread(self.server.started.wait, 5)) + self.check_callbacks([logger]) + self.assertEqual(len(refs), 5) + gc.collect() + self.assertTrue(all(ref() is not None for ref in refs)) + task.cancel() + with self.assertRaises(asyncio.CancelledError): + await task + self.assertFalse(self.server.release.is_set()) + return refs + finally: + if not task.done(): + task.cancel() + await asyncio.gather(task, return_exceptions=True) + await client.close() + + async def exercise(): + refs = await cancel() + barrier = asyncio.get_running_loop().create_future() + asyncio.get_running_loop().call_soon(barrier.set_result, None) + await barrier + gc.collect() + self.assertTrue(all(ref() is None for ref in refs), "cancelled native call retained the request graph") + self.server.release.set() + self.assertTrue(await asyncio.to_thread(self.server.finished.wait, 5)) + self.assertEqual(len(self.server.requests), 1) + + asyncio.run(exercise()) + + +result = unittest.TextTestRunner(verbosity=2).run(unittest.defaultTestLoader.loadTestsFromTestCase(RealBoundaryTests)) +assert result.wasSuccessful(), "real retained OCR boundary tests failed"