mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
wip
This commit is contained in:
parent
f854c38f36
commit
f1553ae9c3
18 changed files with 1576 additions and 6 deletions
23
litellm-rust/Cargo.lock
generated
23
litellm-rust/Cargo.lock
generated
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"] }
|
||||
|
|
|
|||
|
|
@ -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";
|
||||
|
||||
|
|
|
|||
|
|
@ -1,2 +1,3 @@
|
|||
pub mod transformation;
|
||||
pub mod transport;
|
||||
pub mod types;
|
||||
|
|
|
|||
77
litellm-rust/crates/core/src/ocr/transport.rs
Normal file
77
litellm-rust/crates/core/src/ocr/transport.rs
Normal 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;
|
||||
44
litellm-rust/crates/core/src/ocr/transport/tests.rs
Normal file
44
litellm-rust/crates/core/src/ocr/transport/tests.rs
Normal 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")
|
||||
);
|
||||
}
|
||||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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>>,
|
||||
{
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
157
litellm-rust/crates/python-bridge/src/routes/bindings.rs
Normal file
157
litellm-rust/crates/python-bridge/src/routes/bindings.rs
Normal 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" }
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -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)?;
|
||||
|
|
|
|||
144
litellm-rust/crates/python-bridge/src/routes/ocr_retained.rs
Normal file
144
litellm-rust/crates/python-bridge/src/routes/ocr_retained.rs
Normal 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)?)
|
||||
}
|
||||
243
litellm-rust/crates/python-bridge/tests/ocr_retained.rs
Normal file
243
litellm-rust/crates/python-bridge/tests/ocr_retained.rs
Normal 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),
|
||||
)
|
||||
})
|
||||
}
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
146
litellm/rust_bridge/ocr_retained.py
Normal file
146
litellm/rust_bridge/ocr_retained.py
Normal 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))
|
||||
642
tests/test_litellm/ocr/retained_boundary_fixture.py
Normal file
642
tests/test_litellm/ocr/retained_boundary_fixture.py
Normal 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"
|
||||
Loading…
Add table
Reference in a new issue