This commit is contained in:
Yujong Lee 2026-09-06 17:28:01 -07:00
parent f854c38f36
commit f1553ae9c3
18 changed files with 1576 additions and 6 deletions

View file

@ -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"

View file

@ -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"] }

View file

@ -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";

View file

@ -1,2 +1,3 @@
pub mod transformation;
pub mod transport;
pub mod types;

View file

@ -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<u8>, Vec<u8>)>,
pub body: Vec<u8>,
pub timeout_seconds: f64,
}
pub struct Response {
pub status: u16,
pub headers: Vec<(Vec<u8>, Vec<u8>)>,
pub content: Vec<u8>,
}
pub async fn send(request: Request) -> Result<Response, Error> {
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<Result<reqwest::Client, reqwest::Error>> = 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;

View file

@ -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")
);
}

View file

@ -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]

View file

@ -37,6 +37,37 @@ fn run_sync_on<T, F>(
where
T: Serialize + Send + 'static,
F: Future<Output = Result<T, Error>> + 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<T, F>(
py: Python<'_>,
future: F,
map_error: fn(Error) -> PyErr,
) -> PyResult<T>
where
T: Send + 'static,
F: Future<Output = Result<T, Error>> + Send + 'static,
{
run_sync_value_on(
py,
pyo3_async_runtimes::tokio::get_runtime(),
future,
map_error,
)
}
fn run_sync_value_on<T, F>(
py: Python<'_>,
runtime: &Runtime,
future: F,
map_error: fn(Error) -> PyErr,
) -> PyResult<T>
where
T: Send + 'static,
F: Future<Output = Result<T, Error>> + 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<T, F>(
@ -75,7 +105,7 @@ fn map_core_result<T>(result: Result<T, Error>, map_error: fn(Error) -> PyErr) -
}
}
async fn catch_future_panic<T, F>(future: F) -> PyResult<Result<T, Error>>
pub(crate) async fn catch_future_panic<T, F>(future: F) -> PyResult<Result<T, Error>>
where
F: Future<Output = Result<T, Error>>,
{

View file

@ -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",

View file

@ -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" }
);
}
}
}

View file

@ -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)?;

View file

@ -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<Py<PyAny>> {
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<Py<PyAny>> {
invoke(
boundary,
BoundaryMethod::Prepare.resolve(asynchronous),
PyTuple::empty(boundary.py()),
)
}
fn encode(boundary: &Bound<'_, PyAny>, roots: &Bound<'_, PyAny>) -> PyResult<Request> {
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<Self::Output> {
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<Bound<'py, PyAny>> {
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<Py<PyAny>> {
invoke(
boundary,
BoundaryMethod::Finish.resolve(asynchronous),
PyTuple::new(boundary.py(), [wire])?,
)
}
#[pyfunction]
fn ocr_retained(boundary: &Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
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<Py<PyAny>> {
static DRIVER: PyOnceLock<Py<PyAny>> = 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)?)
}

View file

@ -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),
)
})
}

View file

@ -5,7 +5,7 @@ use pyo3::types::{PyDict, PyTuple};
static AWAIT_CALL: PyOnceLock<Py<PyAny>> = PyOnceLock::new();
#[derive(Clone, Copy)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum InvocationMode {
Direct,
Await,

View file

@ -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(

View file

@ -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()

View file

@ -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))

View file

@ -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"