This commit is contained in:
Yujong Lee 2026-09-15 20:10:07 -07:00
parent 6d66be8bef
commit b9fca28c6e
17 changed files with 442 additions and 569 deletions

View file

@ -160,7 +160,7 @@ fn redact(
pub(super) fn response(py: Python<'_>, response: &LiteLLMOcrResponse) -> PyResult<Py<PyAny>> {
py.import("litellm.rust_bridge.ocr")?
.getattr("_response")?
.getattr("build_response")?
.call1((to_py(py, response)?,))
.map(Bound::unbind)
}
@ -172,7 +172,7 @@ pub(super) fn map_failure(
provider: &str,
) -> PyResult<Py<PyBaseException>> {
Ok(py
.import("litellm.rust_bridge.ocr_lifecycle")?
.import("litellm.rust_bridge.ocr")?
.getattr("map_failure")?
.call1((error, request, provider))?
.extract()?)

View file

@ -29,7 +29,7 @@ enum OcrHostData {
struct ProjectedOcrHost {
fields: ProjectedOcrFields,
pre_call: Option<callbacks::OcrLoggingFields>,
retained_fields: Option<Py<PyDict>>,
callback_inputs: Option<Py<PyDict>>,
body: Option<Py<PyDict>>,
headers: Option<Py<PyDict>>,
}
@ -55,20 +55,10 @@ impl PythonOcrHost {
request: OcrPreCallRequest,
) -> PyResult<OcrPreCallRequest> {
let kwargs = self.state.kwargs.bind(py);
let retained_fields = PyDict::new(py);
for name in request
.optional_params
.as_object()
.ok_or_else(missing_state)?
.keys()
{
if let Some(value) = kwargs.get_item(name)? {
retained_fields.set_item(name, value)?;
}
}
retained_fields.set_item("document", &self.projected()?.fields.document)?;
let callback_inputs = kwargs.copy()?;
callback_inputs.set_item("document", &self.projected()?.fields.document)?;
let projected = self.projected_mut()?;
projected.retained_fields = Some(retained_fields.unbind());
projected.callback_inputs = Some(callback_inputs.unbind());
projected.pre_call = Some((&request).into());
Ok(request)
}
@ -112,7 +102,7 @@ impl PythonOcrHost {
let body = to_py(py, &request.body)?
.into_bound(py)
.cast_into::<PyDict>()?;
if let Some(retained) = &self.projected()?.retained_fields {
if let Some(retained) = &self.projected()?.callback_inputs {
for name in &request.retained_fields {
if let Some(value) = retained.bind(py).get_item(name)? {
body.set_item(name, value)?;
@ -202,7 +192,7 @@ impl PythonRoute for PythonOcrHost {
self.data = OcrHostData::Projected(Box::new(ProjectedOcrHost {
fields: projected.fields,
pre_call: None,
retained_fields: None,
callback_inputs: None,
body: None,
headers: None,
}));
@ -265,7 +255,7 @@ impl PythonRoute for PythonOcrHost {
if let Some(provider) = &projected.fields.azure_ad_token_provider {
provider.traverse(visit)?;
}
visit.call(&projected.retained_fields)?;
visit.call(&projected.callback_inputs)?;
visit.call(&projected.body)?;
visit.call(&projected.headers)
}
@ -282,14 +272,18 @@ impl litellm_core::ocr::hooks::OcrHooks for BridgeOcrHooks {
}
}
#[pyfunction]
fn _ocr_lifecycle(
fn call(
py: Python<'_>,
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
kwargs: Option<Bound<'_, PyDict>>,
asynchronous: bool,
) -> PyResult<Py<PyAny>> {
let kwargs = kwargs.unwrap_or_else(|| PyDict::new(py));
let name = if asynchronous { "aocr" } else { "ocr" };
let request = py
.import("litellm.rust_bridge.ocr")?
.getattr("bind_request")?
.call1((name, &args, &kwargs))?;
let client = OcrClient::shared().map_err(ocr_error_to_pyerr)?;
let call = admitted_call(OcrCall::admit(
client,
@ -304,7 +298,7 @@ fn _ocr_lifecycle(
args.unbind(),
kwargs.copy()?.unbind(),
asynchronous,
if asynchronous { "aocr" } else { "ocr" },
name,
)?,
data: OcrHostData::Unprojected {
request: request.unbind(),
@ -313,6 +307,27 @@ fn _ocr_lifecycle(
run_call(py, call, host)
}
pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
module.add_function(wrap_pyfunction!(_ocr_lifecycle, module)?)
#[pyfunction]
#[pyo3(signature = (*args, **kwargs))]
fn ocr(
py: Python<'_>,
args: Bound<'_, PyTuple>,
kwargs: Option<Bound<'_, PyDict>>,
) -> PyResult<Py<PyAny>> {
call(py, args, kwargs, false)
}
#[pyfunction]
#[pyo3(signature = (*args, **kwargs))]
fn aocr(
py: Python<'_>,
args: Bound<'_, PyTuple>,
kwargs: Option<Bound<'_, PyDict>>,
) -> PyResult<Py<PyAny>> {
call(py, args, kwargs, true)
}
pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
super::super::add_function(module, wrap_pyfunction!(ocr, module)?)?;
super::super::add_function(module, wrap_pyfunction!(aocr, module)?)
}

View file

@ -4,11 +4,16 @@ mod errors;
mod lifecycle;
mod project;
mod request;
mod value;
#[cfg(feature = "trace-parity")]
mod trace;
use pyo3::prelude::*;
pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
value::register(module)?;
lifecycle::register(module)
}
#[cfg(feature = "trace-parity")]
pub(super) fn register_trace(module: &Bound<'_, PyModule>) -> PyResult<()> {
trace::register(module)
}

View file

@ -36,12 +36,7 @@ pub(super) struct PythonOcrInput<'a, 'py> {
pub kwargs: &'a Bound<'py, PyDict>,
}
struct PythonOcrFields<'a, 'py> {
request: &'a Bound<'py, PyAny>,
kwargs: &'a Bound<'py, PyDict>,
}
impl<'py> PythonOcrFields<'_, 'py> {
impl<'py> PythonOcrInput<'_, 'py> {
fn lookup(&self, name: &str) -> PyResult<Bound<'py, PyAny>> {
match self.kwargs.get_item(name)? {
Some(value) => Ok(value),
@ -86,35 +81,16 @@ impl<'py> PythonOcrFields<'_, 'py> {
}
}
enum ProjectedDocument {
File { wire: Value, retained: Py<PyAny> },
Other { wire: Value, retained: Py<PyAny> },
}
impl ProjectedDocument {
fn project(py: Python<'_>, document: &Bound<'_, PyAny>) -> PyResult<Self> {
let kind: String = document.get_item("type")?.extract()?;
if kind != "file" {
return Ok(Self::Other {
wire: from_py(document)?,
retained: document.clone().unbind(),
});
}
let input = document.extract()?;
let encoded = super::document::file_document(py, input)?;
let wire = serde_json::to_value(encoded)
.map_err(|error| pyo3::exceptions::PyValueError::new_err(error.to_string()))?;
Ok(Self::File {
retained: to_py(py, &wire)?,
wire,
})
}
fn into_parts(self) -> (Value, Py<PyAny>) {
match self {
Self::File { wire, retained } | Self::Other { wire, retained } => (wire, retained),
}
fn project_document(py: Python<'_>, document: &Bound<'_, PyAny>) -> PyResult<(Value, Py<PyAny>)> {
let kind: String = document.get_item("type")?.extract()?;
if kind != "file" {
return Ok((from_py(document)?, document.clone().unbind()));
}
let encoded = super::document::file_document(py, document.extract()?)?;
let wire = serde_json::to_value(encoded)
.map_err(|error| pyo3::exceptions::PyValueError::new_err(error.to_string()))?;
let retained = to_py(py, &wire)?;
Ok((wire, retained))
}
impl TryFrom<PythonOcrInput<'_, '_>> for ProjectedOcrCall {
@ -125,11 +101,10 @@ impl TryFrom<PythonOcrInput<'_, '_>> for ProjectedOcrCall {
let kwargs = input.kwargs;
let py = request.py();
let boundary_request = request.clone().unbind();
let arguments = PythonOcrFields { request, kwargs };
let arguments = input;
let model = arguments.model()?;
let custom_llm_provider = arguments.custom_llm_provider()?;
let (wire_document, retained_document) =
ProjectedDocument::project(py, &arguments.document()?)?.into_parts();
let (wire_document, retained_document) = project_document(py, &arguments.document()?)?;
let api_key = arguments.api_key()?;
let specs = consumed_optional_params(&model, custom_llm_provider.as_deref())
.map_err(ocr_error_to_pyerr)?;
@ -215,15 +190,8 @@ mod tests {
fn arguments<'a, 'py>(
request: &'a Bound<'py, PyAny>,
kwargs: &'a Bound<'py, PyDict>,
) -> PythonOcrFields<'a, 'py> {
PythonOcrFields { request, kwargs }
}
fn project_document(
py: Python<'_>,
document: &Bound<'_, PyAny>,
) -> PyResult<(Value, Py<PyAny>)> {
ProjectedDocument::project(py, document).map(ProjectedDocument::into_parts)
) -> PythonOcrInput<'a, 'py> {
PythonOcrInput { request, kwargs }
}
fn stub_timeout_conversion(py: Python<'_>) {

View file

@ -85,7 +85,7 @@ fn ocr(
) -> PyResult<Py<PyAny>> {
run_sync(
py,
prepare_ocr(OcrInputs {
crate::function_trace::capture(prepare_ocr(OcrInputs {
model,
document,
api_key,
@ -95,7 +95,7 @@ fn ocr(
optional_params,
input_sources,
timeout_seconds,
})?,
})?),
ocr_error_to_pyerr,
)
}
@ -117,7 +117,7 @@ fn aocr(
) -> PyResult<Bound<'_, PyAny>> {
run_async(
py,
prepare_ocr(OcrInputs {
crate::function_trace::capture(prepare_ocr(OcrInputs {
model,
document,
api_key,
@ -127,7 +127,7 @@ fn aocr(
optional_params,
input_sources,
timeout_seconds,
})?,
})?),
ocr_error_to_pyerr,
)
}
@ -136,83 +136,3 @@ pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
super::super::add_function(module, wrap_pyfunction!(ocr, module)?)?;
super::super::add_function(module, wrap_pyfunction!(aocr, module)?)
}
#[cfg(feature = "trace-parity")]
mod trace {
use super::{OcrInputs, Value, ocr_error_to_pyerr, prepare_ocr, run_async, run_sync};
use pyo3::prelude::*;
#[pyfunction]
#[pyo3(signature = (model, document, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, input_sources=None, timeout_seconds=None))]
#[allow(clippy::too_many_arguments)]
fn ocr(
py: Python<'_>,
model: String,
#[pyo3(from_py_with = litellm_python_interop::from_py)] document: Value,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
#[pyo3(from_py_with = litellm_python_interop::from_py)] extra_headers: Option<Value>,
#[pyo3(from_py_with = litellm_python_interop::from_py)] optional_params: Option<Value>,
#[pyo3(from_py_with = litellm_python_interop::from_py)] input_sources: Option<Value>,
timeout_seconds: Option<f64>,
) -> PyResult<Py<PyAny>> {
run_sync(
py,
crate::function_trace::capture(prepare_ocr(OcrInputs {
model,
document,
api_key,
api_base,
custom_llm_provider,
extra_headers,
optional_params,
input_sources,
timeout_seconds,
})?),
ocr_error_to_pyerr,
)
}
#[pyfunction]
#[pyo3(signature = (model, document, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, input_sources=None, timeout_seconds=None))]
#[allow(clippy::too_many_arguments)]
fn aocr(
py: Python<'_>,
model: String,
#[pyo3(from_py_with = litellm_python_interop::from_py)] document: Value,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
#[pyo3(from_py_with = litellm_python_interop::from_py)] extra_headers: Option<Value>,
#[pyo3(from_py_with = litellm_python_interop::from_py)] optional_params: Option<Value>,
#[pyo3(from_py_with = litellm_python_interop::from_py)] input_sources: Option<Value>,
timeout_seconds: Option<f64>,
) -> PyResult<Bound<'_, PyAny>> {
run_async(
py,
crate::function_trace::capture(prepare_ocr(OcrInputs {
model,
document,
api_key,
api_base,
custom_llm_provider,
extra_headers,
optional_params,
input_sources,
timeout_seconds,
})?),
ocr_error_to_pyerr,
)
}
pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
super::super::super::add_function(module, wrap_pyfunction!(ocr, module)?)?;
super::super::super::add_function(module, wrap_pyfunction!(aocr, module)?)
}
}
#[cfg(feature = "trace-parity")]
pub(super) fn register_trace(module: &Bound<'_, PyModule>) -> PyResult<()> {
trace::register(module)
}

View file

@ -9,11 +9,7 @@ fn sync_and_async_route_signatures_match_the_python_contract() {
let module = PyModule::new(py, "routes").expect("module should be created");
register(&module).expect("routes should register");
let routes = [
(
"ocr",
"aocr",
"(model, document, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, input_sources=None, timeout_seconds=None)",
),
("ocr", "aocr", "(*args, **kwargs)"),
(
"transcription",
"atranscription",
@ -98,22 +94,20 @@ fn sync_and_async_routes_apply_the_same_input_validation() {
.expect("kwargs should accept extra_headers");
let document = PyDict::new(py);
for (sync_name, async_name) in [("ocr", "aocr"), ("transcription", "atranscription")] {
let sync_error = module
.getattr(sync_name)
.and_then(|function| function.call(("model", &document), Some(&kwargs)))
.expect_err("sync route should reject non-dict extra_headers");
let async_error = module
.getattr(async_name)
.and_then(|function| function.call(("model", &document), Some(&kwargs)))
.expect_err("async route should reject non-dict extra_headers");
let sync_error = module
.getattr("transcription")
.and_then(|function| function.call(("model", &document), Some(&kwargs)))
.expect_err("sync route should reject non-dict extra_headers");
let async_error = module
.getattr("atranscription")
.and_then(|function| function.call(("model", &document), Some(&kwargs)))
.expect_err("async route should reject non-dict extra_headers");
assert_eq!(
sync_error.to_string(),
"ValueError: extra_headers must be a dict"
);
assert_eq!(async_error.to_string(), sync_error.to_string());
}
assert_eq!(
sync_error.to_string(),
"ValueError: extra_headers must be a dict"
);
assert_eq!(async_error.to_string(), sync_error.to_string());
});
}
@ -162,15 +156,13 @@ fn route_input_validation_preserves_left_to_right_order() {
let invalid_payload =
PyModule::new(py, "invalid_payload").expect("invalid payload should be created");
for name in ["ocr", "transcription"] {
let error = module
.getattr(name)
.and_then(|function| {
function.call(("model", &invalid_payload), Some(&headers_kwargs))
})
.expect_err("payload should be validated before headers");
assert!(!error.to_string().contains("extra_headers"));
}
let error = module
.getattr("transcription")
.and_then(|function| {
function.call(("model", &invalid_payload), Some(&headers_kwargs))
})
.expect_err("payload should be validated before headers");
assert!(!error.to_string().contains("extra_headers"));
});
}

View file

@ -1,59 +1,25 @@
from collections.abc import Awaitable, Callable, Coroutine, Mapping
from collections.abc import Awaitable, Callable, Coroutine
from typing import Final, cast # noqa: TID251 # native binding selects a sync result or an async awaitable
import httpx
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.ocr import legacy
from litellm.ocr.legacy import convert_file_document_to_url_document, get_mime_type
from litellm.rust_bridge.bindings import native_exception_types
from litellm.rust_bridge.configuration import rust_ocr_enabled
from litellm.rust_bridge.ocr import LiteLLMOcrRequest
from litellm.rust_bridge.ocr_lifecycle import select
from litellm.rust_bridge.ocr import NATIVE_AOCR, NATIVE_OCR, bind_request
__all__ = ("aocr", "convert_file_document_to_url_document", "get_mime_type", "ocr")
def _bind_request(
model: str,
document: Mapping[str, object],
api_key: str | None = None,
api_base: str | None = None,
timeout: float | httpx.Timeout | None = None,
custom_llm_provider: str | None = None,
extra_headers: dict[str, object] | None = None,
**kwargs: object, # kwargs-ok: public OCR accepts provider-specific options
) -> LiteLLMOcrRequest:
return LiteLLMOcrRequest(
model=model,
document=document,
api_key=api_key,
api_base=api_base,
timeout=timeout,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
kwargs=kwargs,
)
def _public_request(name: str, args: tuple[object, ...], kwargs: dict[str, object]) -> LiteLLMOcrRequest:
try:
return _bind_request(*args, **kwargs) # pyright: ignore[reportArgumentType] # Python binds the public arguments before native validation
except TypeError as error:
raise TypeError(str(error).replace("_bind_request()", f"{name}()")) from None
def ocr(
*args: object,
**kwargs: object, # kwargs-ok: preserve the public OCR call shape
) -> OCRResponse | Coroutine[object, object, OCRResponse]:
request: Final = _public_request("ocr", args, kwargs)
native: Final = select(request) if rust_ocr_enabled() else None
bind_request("ocr", args, kwargs)
native: Final = NATIVE_OCR.load() if rust_ocr_enabled() and not kwargs.get("aocr") else None
if native is not None:
try:
return cast( # cast-ok: False selects the synchronous result
OCRResponse, native(request, args, kwargs, False)
)
return native(*args, **kwargs)
except _decline_types():
pass
fallback: Final = cast( # cast-ok: forward the original call shape through the legacy @client decorator
@ -63,13 +29,11 @@ def ocr(
async def aocr(*args: object, **kwargs: object) -> OCRResponse: # kwargs-ok: preserve the public OCR call shape
request: Final = _public_request("aocr", args, kwargs)
native: Final = select(request) if rust_ocr_enabled() else None
bind_request("aocr", args, kwargs)
native: Final = NATIVE_AOCR.load() if rust_ocr_enabled() and not kwargs.get("aocr") else None
if native is not None:
try:
return await cast( # cast-ok: True selects the asynchronous result
Awaitable[OCRResponse], native(request, args, kwargs, True)
)
return await native(*args, **kwargs)
except _decline_types():
pass
fallback: Final = cast( # cast-ok: forward the original call shape through the legacy @client decorator

View file

@ -1,135 +1,109 @@
from asyncio import Future
from collections.abc import Coroutine, Mapping, Sequence
from typing import Literal, Never, TypeAlias, final
from collections.abc import Coroutine, Mapping
import httpx
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.rust_bridge.ocr import LiteLLMOcrRequest
_InputSource: TypeAlias = Literal["request", "deployment", "environment"]
class RustBridgeDeclined(Exception): ...
class RustUpstreamError(Exception): ...
def ocr(
model: str,
document: object,
api_key: str | None = None,
api_base: str | None = None,
custom_llm_provider: str | None = None,
extra_headers: Mapping[str, object] | None = None,
optional_params: Mapping[str, object] | None = None,
input_sources: Mapping[str, _InputSource] | None = None,
timeout_seconds: float | None = None,
) -> dict[str, object]: ...
document: Mapping[str, object],
api_key: str | None = ...,
api_base: str | None = ...,
timeout: float | httpx.Timeout | None = ...,
custom_llm_provider: str | None = ...,
extra_headers: dict[str, object] | None = ...,
**kwargs: object,
) -> OCRResponse: ...
def aocr(
model: str,
document: object,
api_key: str | None = None,
api_base: str | None = None,
custom_llm_provider: str | None = None,
extra_headers: Mapping[str, object] | None = None,
optional_params: Mapping[str, object] | None = None,
input_sources: Mapping[str, _InputSource] | None = None,
timeout_seconds: float | None = None,
) -> Future[dict[str, object]]: ...
_OCR_MAX_FILE_BYTES: int
def _ocr_upload_document(
file_content: bytes,
file_name: str | None = None,
content_type: str | None = None,
) -> dict[str, str]: ...
def _ocr_file_document(document: Mapping[str, object]) -> dict[str, str]: ...
def _ocr_mime_type(file_name: str) -> str: ...
def _ocr_lifecycle(
request: LiteLLMOcrRequest,
args: tuple[object, ...],
kwargs: dict[str, object],
asynchronous: bool,
) -> OCRResponse | Coroutine[object, object, OCRResponse]: ...
document: Mapping[str, object],
api_key: str | None = ...,
api_base: str | None = ...,
timeout: float | httpx.Timeout | None = ...,
custom_llm_provider: str | None = ...,
extra_headers: dict[str, object] | None = ...,
**kwargs: object,
) -> Coroutine[object, object, OCRResponse]: ...
def transcription(
model: str,
audio: object,
api_key: str | None = None,
api_base: str | None = None,
custom_llm_provider: str | None = None,
extra_headers: Mapping[str, object] | None = None,
optional_params: Mapping[str, object] | None = None,
timeout_seconds: float | None = None,
api_key: str | None = ...,
api_base: str | None = ...,
custom_llm_provider: str | None = ...,
extra_headers: object = ...,
optional_params: object = ...,
timeout_seconds: float | None = ...,
) -> dict[str, object]: ...
def atranscription(
model: str,
audio: object,
api_key: str | None = None,
api_base: str | None = None,
custom_llm_provider: str | None = None,
extra_headers: Mapping[str, object] | None = None,
optional_params: Mapping[str, object] | None = None,
timeout_seconds: float | None = None,
api_key: str | None = ...,
api_base: str | None = ...,
custom_llm_provider: str | None = ...,
extra_headers: object = ...,
optional_params: object = ...,
timeout_seconds: float | None = ...,
) -> Future[dict[str, object]]: ...
def messages(
model: str,
body: Mapping[str, object],
api_key: str | None = None,
api_base: str | None = None,
custom_llm_provider: str | None = None,
extra_headers: Mapping[str, object] | None = None,
timeout_seconds: float | None = None,
body: object,
api_key: str | None = ...,
api_base: str | None = ...,
custom_llm_provider: str | None = ...,
extra_headers: object = ...,
timeout_seconds: float | None = ...,
) -> dict[str, object]: ...
def amessages(
model: str,
body: Mapping[str, object],
api_key: str | None = None,
api_base: str | None = None,
custom_llm_provider: str | None = None,
extra_headers: Mapping[str, object] | None = None,
timeout_seconds: float | None = None,
body: object,
api_key: str | None = ...,
api_base: str | None = ...,
custom_llm_provider: str | None = ...,
extra_headers: object = ...,
timeout_seconds: float | None = ...,
) -> Future[dict[str, object]]: ...
def chat_completions_decline(
model: str,
messages: Sequence[object],
optional_params: Mapping[str, object] | None = None,
custom_llm_provider: str | None = None,
) -> str | None: ...
def chat_completions(
model: str,
messages: Sequence[object],
optional_params: Mapping[str, object] | None = None,
api_key: str | None = None,
api_base: str | None = None,
custom_llm_provider: str | None = None,
extra_headers: Mapping[str, object] | None = None,
timeout_seconds: float | None = None,
messages: object,
optional_params: object = ...,
api_key: str | None = ...,
api_base: str | None = ...,
custom_llm_provider: str | None = ...,
extra_headers: object = ...,
timeout_seconds: float | None = ...,
) -> dict[str, object]: ...
def achat_completions(
model: str,
messages: Sequence[object],
optional_params: Mapping[str, object] | None = None,
api_key: str | None = None,
api_base: str | None = None,
custom_llm_provider: str | None = None,
extra_headers: Mapping[str, object] | None = None,
timeout_seconds: float | None = None,
messages: object,
optional_params: object = ...,
api_key: str | None = ...,
api_base: str | None = ...,
custom_llm_provider: str | None = ...,
extra_headers: object = ...,
timeout_seconds: float | None = ...,
) -> Future[dict[str, object]]: ...
def chat_completions_decline(
model: str,
messages: object,
optional_params: object = ...,
custom_llm_provider: str | None = ...,
) -> str | None: ...
@final
class ResponsesWebSocketConnection:
def __new__(cls, _uninstantiable: Never, /) -> Never: ...
@classmethod
def connect(
cls,
url: str,
headers: Mapping[str, str] | None = None,
timeout_seconds: float | None = None,
cls, url: str, headers: object = ..., timeout_seconds: float | None = ...
) -> Future[ResponsesWebSocketConnection]: ...
def send_text(self, text: str) -> Future[None]: ...
def recv_text(self) -> Future[str | None]: ...
def close(self) -> Future[None]: ...
@final
class TokenCounter:
def __new__(cls, tokenizer_json: str) -> TokenCounter: ...
def __init__(self, tokenizer_json: str) -> None: ...
@staticmethod
def from_cl100k_ranks(rank_file: str) -> TokenCounter: ...
@staticmethod
@ -137,25 +111,3 @@ class TokenCounter:
def acount_request(self, body: bytes) -> Future[dict[str, object]]: ...
def gil_stats() -> dict[str, int]: ...
__all__ = [
"_OCR_MAX_FILE_BYTES",
"ResponsesWebSocketConnection",
"RustBridgeDeclined",
"RustUpstreamError",
"TokenCounter",
"_ocr_file_document",
"_ocr_lifecycle",
"_ocr_mime_type",
"_ocr_upload_document",
"achat_completions",
"amessages",
"aocr",
"atranscription",
"chat_completions",
"chat_completions_decline",
"gil_stats",
"messages",
"ocr",
"transcription",
]

View file

@ -1,17 +1,16 @@
"""Thin Python wrapper for the native Rust OCR bridge."""
from __future__ import annotations
from collections.abc import Awaitable, Mapping
from collections.abc import Coroutine, Mapping
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final, Protocol, cast # noqa: TID251 # native extension exposes dynamically typed callables
from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables
import httpx
from pydantic import TypeAdapter
import litellm
from litellm.llms.base_llm.ocr.transformation import PROVIDER_NATIVE_RESPONSE_KEY, OCRResponse
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.timeouts import timeout_to_seconds as _timeout_to_seconds
@dataclass(frozen=True, slots=True)
@ -24,39 +23,34 @@ class LiteLLMOcrRequest:
custom_llm_provider: str | None
extra_headers: dict[str, object] | None
kwargs: Mapping[str, object]
input_sources: Mapping[str, str] | None = None
def _bind_request(
model: str,
document: Mapping[str, object],
api_key: str | None = None,
api_base: str | None = None,
timeout: float | httpx.Timeout | None = None,
custom_llm_provider: str | None = None,
extra_headers: dict[str, object] | None = None,
**kwargs: object, # kwargs-ok: public OCR accepts provider-specific options
) -> LiteLLMOcrRequest:
return LiteLLMOcrRequest(model, document, api_key, api_base, timeout, custom_llm_provider, extra_headers, kwargs)
def bind_request(name: str, args: tuple[object, ...], kwargs: dict[str, object]) -> LiteLLMOcrRequest:
try:
return _bind_request(*args, **kwargs) # pyright: ignore[reportArgumentType] # Python binds arguments before native validation
except TypeError as error:
raise TypeError(str(error).replace("_bind_request()", f"{name}()")) from None
class RustOcr(Protocol):
def __call__(
self,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
input_sources: dict[str, str],
timeout_seconds: float | None,
) -> dict[str, object]:
raise NotImplementedError
def __call__(self, *args: object, **kwargs: object) -> OCRResponse: ...
class RustAocr(Protocol):
def __call__(
self,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
input_sources: dict[str, str],
timeout_seconds: float | None,
) -> Awaitable[dict[str, object]]:
raise NotImplementedError
def __call__(self, *args: object, **kwargs: object) -> Coroutine[object, object, OCRResponse]: ...
def _as_ocr(value: object) -> RustOcr | None:
@ -67,79 +61,45 @@ def _as_aocr(value: object) -> RustAocr | None:
return cast(RustAocr, value) if callable(value) else None
_OCR: Final = NativeBinding("ocr", validate=_as_ocr)
_AOCR: Final = NativeBinding("aocr", validate=_as_aocr)
NATIVE_OCR: Final = NativeBinding("ocr", validate=_as_ocr)
NATIVE_AOCR: Final = NativeBinding("aocr", validate=_as_aocr)
_NATIVE_RESPONSE: Final = TypeAdapter(Mapping[str, object])
def load_rust_ocr() -> RustOcr | None:
return _OCR.load()
def load_rust_aocr() -> RustAocr | None:
return _AOCR.load()
def _response(response: Mapping[str, object]) -> OCRResponse:
def build_response(response: Mapping[str, object]) -> OCRResponse:
provider_native_response: Final = response.get(PROVIDER_NATIVE_RESPONSE_KEY)
normalized: Final = OCRResponse.model_validate(
MappingProxyType({key: value for key, value in response.items() if key != PROVIDER_NATIVE_RESPONSE_KEY})
)
if isinstance(provider_native_response, Mapping):
normalized.set_provider_native_response(provider_native_response)
normalized.set_provider_native_response(_NATIVE_RESPONSE.validate_python(provider_native_response))
return normalized
def ocr(
*,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout: float | httpx.Timeout | None,
input_sources: Mapping[str, str] | None = None,
) -> dict[str, object] | None:
rust_ocr: Final = load_rust_ocr()
if rust_ocr is None:
return None
return rust_ocr(
model=model,
document=document,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
input_sources=dict(input_sources or {}), # mutable-ok: native boundary requires a concrete dict
timeout_seconds=_timeout_to_seconds(timeout),
)
class ExceptionMapper(Protocol):
def __call__(
self,
*,
model: str,
custom_llm_provider: str | None,
original_exception: Exception,
completion_kwargs: dict[str, object],
extra_kwargs: dict[str, object],
) -> Exception: ...
async def aocr(
*,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout: float | httpx.Timeout | None,
input_sources: Mapping[str, str] | None = None,
) -> dict[str, object] | None:
rust_aocr: Final = load_rust_aocr()
if rust_aocr is None:
return None
return await rust_aocr(
model=model,
document=document,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
input_sources=dict(input_sources or {}), # mutable-ok: native boundary requires a concrete dict
timeout_seconds=_timeout_to_seconds(timeout),
)
def map_failure(error: Exception, request: LiteLLMOcrRequest, request_provider: str) -> Exception:
mapper: Final = cast(
ExceptionMapper, litellm.exception_type
) # cast-ok: bounded adapter for the public exception mapper
try:
return mapper(
model=request.model.removeprefix(f"{request_provider}/"),
custom_llm_provider=request_provider,
original_exception=error,
completion_kwargs=dict(request.kwargs), # mutable-ok: exception mapper requires owned kwargs
extra_kwargs=dict(request.kwargs), # mutable-ok: exception mapper requires owned kwargs
)
except Exception as public_error:
public_error.__context__ = error
return public_error

View file

@ -1,67 +0,0 @@
from __future__ import annotations
from collections.abc import Awaitable, Mapping, Sequence
from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables
import litellm
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.ocr import LiteLLMOcrRequest
class NativeOcrLifecycle(Protocol):
def __call__(
self,
request: LiteLLMOcrRequest,
args: Sequence[object],
kwargs: Mapping[str, object],
asynchronous: bool,
) -> OCRResponse | Awaitable[OCRResponse]: ...
class ExceptionMapper(Protocol):
def __call__(
self,
*,
model: str,
custom_llm_provider: str | None,
original_exception: Exception,
completion_kwargs: dict[str, object],
extra_kwargs: dict[str, object],
) -> Exception: ...
def _binding(value: object) -> NativeOcrLifecycle | None:
if not callable(value):
return None
return cast("NativeOcrLifecycle", value) # cast-ok: callable validated at the native binding boundary
NATIVE_OCR_LIFECYCLE: Final = NativeBinding("_ocr_lifecycle", validate=_binding)
def select(request: LiteLLMOcrRequest) -> NativeOcrLifecycle | None:
if request.kwargs.get("aocr"):
return None
return NATIVE_OCR_LIFECYCLE.load()
def arguments(request: LiteLLMOcrRequest) -> Mapping[str, object]:
return request.kwargs
def map_failure(error: Exception, request: LiteLLMOcrRequest, request_provider: str) -> Exception:
mapper: Final = cast( # cast-ok: bounded adapter for the legacy public exception mapper
ExceptionMapper, litellm.exception_type
)
try:
return mapper(
model=request.model.removeprefix(f"{request_provider}/"),
custom_llm_provider=request_provider,
original_exception=error,
completion_kwargs=dict(arguments(request)), # mutable-ok: exception mapper requires owned kwargs
extra_kwargs=dict(request.kwargs), # mutable-ok: exception mapper requires owned kwargs
)
except Exception as public_error:
public_error.__context__ = error
return public_error

View file

@ -7,7 +7,106 @@ from .....shared.parity.recorded_http import HttpHeader, RecordedHttpResponse
from ...models import RouteFixture, RouteSpec, TraceScenario, TraceSuite
def _fixture(model: str, document: dict[str, str] | None = None) -> RouteFixture:
SUCCESS_CALLBACK_SYNC_MAPPING: Final = mapping(
rust_span="success_callback",
python_frame=r"BoundedLoggingThreadPoolExecutor\.submit$",
)
SUCCESS_CALLBACK_ASYNC_MAPPING: Final = mapping(
rust_span="success_callback",
python_frame=r"Logging\.async_success_handler$",
)
FAILURE_CALLBACK_MAPPING: Final = mapping(
rust_span="failure_callback",
python_frame=r"Logging\.(?:async_)?failure_handler$",
)
IGNORED_SUCCESS_CALLBACK_MAPPING: Final = mapping(rust_span="success_callback")
SYNC_MAPPINGS: Final = (
*COMMON_MAPPINGS,
mapping(rust_span="execute_ocr_provider_call", python_frame=r"BaseLLMHTTPHandler\.ocr$"),
mapping(span="python_transform_ocr_response_wrapper", python_frame=r"BaseLLMHTTPHandler\._transform_ocr_response$"),
mapping(
rust_span="transform_ocr_response",
python_frame=r"MistralOCRConfig\.transform_ocr_response$",
),
)
ASYNC_MAPPINGS: Final = (
*COMMON_MAPPINGS,
mapping(span="python_ocr_wrapper", python_frame=r"BaseLLMHTTPHandler\.ocr$"),
mapping(rust_span="execute_ocr_provider_call", python_frame=r"BaseLLMHTTPHandler\.async_ocr$"),
mapping(
rust_span="transform_ocr_response",
python_frame=r"MistralOCRConfig\.transform_ocr_response$",
),
)
PUBLIC_RUST_DISPATCH_MAPPINGS: Final = (
mapping(span="public_sdk_entrypoint", python_frame=r"ocr/main\.py:\d+ a?ocr$"),
mapping(span="public_request", python_frame=r"rust_bridge/ocr\.py:\d+ bind_request$"),
mapping(span="bind_request", python_frame=r"rust_bridge/ocr\.py:\d+ _bind_request$"),
mapping(span="rust_ocr_enabled", python_frame=r"rust_bridge/configuration\.py:\d+ rust_ocr_enabled$"),
mapping(span="load_native_bridge", python_frame=r"rust_bridge/bindings\.py:\d+ NativeBinding\.load$"),
mapping(span="native_call_setup", python_frame=r"rust_bridge/lifecycle\.py:\d+ setup$"),
mapping(span="native_response", python_frame=r"rust_bridge/ocr\.py:\d+ build_response$"),
mapping(span="native_call_finalize", python_frame=r"rust_bridge/lifecycle\.py:\d+ finalize$"),
mapping(
span="native_success_bookkeeping",
python_frame=r"rust_bridge/lifecycle\.py:\d+ success_bookkeeping$",
),
*(mapping(rust_span=item.rust) for item in SYNC_MAPPINGS if item.rust is not None),
)
CALLBACK_SUCCESS_SYNC_MAPPINGS: Final = (*SYNC_MAPPINGS, SUCCESS_CALLBACK_SYNC_MAPPING)
CALLBACK_SUCCESS_ASYNC_MAPPINGS: Final = (*ASYNC_MAPPINGS, SUCCESS_CALLBACK_ASYNC_MAPPING)
CALLBACK_FAILURE_SYNC_MAPPINGS: Final = (
*COMMON_MAPPINGS,
mapping(rust_span="execute_ocr_provider_call", python_frame=r"BaseLLMHTTPHandler\.ocr$"),
FAILURE_CALLBACK_MAPPING,
)
CALLBACK_FAILURE_ASYNC_MAPPINGS: Final = (
*COMMON_MAPPINGS,
mapping(span="python_ocr_wrapper", python_frame=r"BaseLLMHTTPHandler\.ocr$"),
mapping(rust_span="execute_ocr_provider_call", python_frame=r"BaseLLMHTTPHandler\.async_ocr$"),
FAILURE_CALLBACK_MAPPING,
)
AZURE_COMMON_MAPPINGS: Final = (
*COMMON_MAPPINGS[:7],
mapping(
rust_span="transform_ocr_request",
python_frame=(
r"AzureAIOCRConfig\.(?:async_)?transform_ocr_request$"
r"|MistralOCRConfig\.transform_ocr_request$"
),
),
COMMON_MAPPINGS[-1],
)
AZURE_SYNC_MAPPINGS: Final = (
*AZURE_COMMON_MAPPINGS,
mapping(rust_span="execute_ocr_provider_call", python_frame=r"BaseLLMHTTPHandler\.ocr$"),
mapping(
span="python_transform_ocr_response_wrapper",
python_frame=r"BaseLLMHTTPHandler\._transform_ocr_response$",
),
mapping(
rust_span="transform_ocr_response",
python_frame=r"MistralOCRConfig\.transform_ocr_response$",
),
)
AZURE_ASYNC_MAPPINGS: Final = (
*AZURE_COMMON_MAPPINGS,
mapping(span="python_ocr_wrapper", python_frame=r"BaseLLMHTTPHandler\.ocr$"),
mapping(rust_span="execute_ocr_provider_call", python_frame=r"BaseLLMHTTPHandler\.async_ocr$"),
mapping(
rust_span="transform_ocr_response",
python_frame=r"MistralOCRConfig\.transform_ocr_response$",
),
)
def _fixture(engine: Engine, model: str, document: dict[str, str] | None = None) -> RouteFixture:
response: Final = json.dumps(
{
"pages": [{"index": 0, "markdown": "hello"}],

View file

@ -17,7 +17,7 @@ from litellm.llms.custom_httpx import llm_http_handler
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.ocr.legacy import _prepare_ocr_request
from litellm.rust_bridge import bindings, configuration
from litellm.rust_bridge.ocr_lifecycle import NATIVE_OCR_LIFECYCLE
from litellm.rust_bridge.ocr import NATIVE_AOCR, NATIVE_OCR
@pytest.fixture
@ -45,7 +45,8 @@ async def provider(monkeypatch: pytest.MonkeyPatch) -> AsyncGenerator[Mock]:
monkeypatch.setattr(llm_http_handler, "_get_httpx_client", lambda: sync_handler)
monkeypatch.setattr(llm_http_handler, "get_async_httpx_client", lambda llm_provider: async_handler)
yield handler
NATIVE_OCR_LIFECYCLE.reset()
NATIVE_OCR.reset()
NATIVE_AOCR.reset()
configuration.reset_rust_configuration()
@ -60,7 +61,9 @@ async def test_python_request_response_and_callbacks(
if dispatch != "disabled":
monkeypatch.setenv("LITELLM_RUST", "1")
NATIVE_OCR_LIFECYCLE.override(Mock(side_effect=Declined()) if dispatch == "declined" else None)
(NATIVE_AOCR if mode == "async" else NATIVE_OCR).override(
Mock(side_effect=Declined()) if dispatch == "declined" else None
)
main: Final = importlib.import_module("litellm.ocr.main")
monkeypatch.setattr(main, "native_exception_types", lambda: (Declined, RuntimeError))
logger: Final = Mock(spec=CustomLogger)

View file

@ -7,7 +7,7 @@ from litellm.rust_bridge import ocr as rust_ocr_bridge
def test_rust_ocr_response_retains_provider_native_response():
provider_response = {"status": "succeeded", "analyzeResult": {"content": "native"}}
response = rust_ocr_bridge._response(
response = rust_ocr_bridge.build_response(
{
"pages": [],
"model": "prebuilt-layout",

View file

@ -16,6 +16,9 @@ from pathlib import Path
from socket import socket as Socket
from typing import Final
import litellm
from litellm.llms.base_llm.ocr.transformation import OCRResponse
REQUEST_STARTED: Final = threading.Event()
REQUEST_CANCELLED: Final = threading.Event()
@ -145,12 +148,13 @@ def route_kwargs(route: str, api_base: str, outcome: str) -> dict[str, object]:
"timeout_seconds": 3.0,
}
if route == "ocr":
return common | {
return {key: value for key, value in common.items() if key != "timeout_seconds"} | {
"model": "mistral-ocr-latest",
"document": {"type": "document_url", "document_url": "https://example.com/document.pdf"},
"api_key": "sk-native",
"custom_llm_provider": "mistral",
"optional_params": {"include_image_base64": True},
"include_image_base64": True,
"timeout": 3.0,
}
if route == "transcription":
return common | {
@ -186,6 +190,10 @@ def route_kwargs(route: str, api_base: str, outcome: str) -> dict[str, object]:
def assert_success(route: str, response: object) -> None:
if route == "ocr":
assert isinstance(response, OCRResponse)
assert response.pages[0].markdown == "native-ocr"
return
if not isinstance(response, dict):
raise TypeError(f"{route} returned {type(response).__name__}, expected dict")
actual: Final = success_value(route, response)
@ -206,7 +214,7 @@ def azure_ocr_kwargs(api_base: str) -> dict[str, object]:
"x-test-outcome": "success",
"x-test-route": "azure_ocr",
},
"optional_params": {"azure_ad_token": "prepared-azure-token"},
"azure_ad_token": "prepared-azure-token",
}
@ -218,7 +226,8 @@ def azure_di_kwargs(api_base: str) -> dict[str, object]:
"api_base": api_base,
"custom_llm_provider": "azure_ai",
"extra_headers": {"x-test-outcome": "success", "x-test-route": "azure_di"},
"optional_params": {"req_format": "native", "pages": [0, 2]},
"req_format": "native",
"pages": [0, 2],
}
@ -233,7 +242,11 @@ def success_value(route: str, response: dict[object, object]) -> object:
def assert_rate_limit(native: object, route: str, error: BaseException) -> None:
if route in {"ocr", "chat_completions"}:
if route == "ocr":
assert isinstance(error, litellm.RateLimitError)
assert error.status_code == 429
return
if route == "chat_completions":
upstream_error: Final = native.RustUpstreamError
if not isinstance(error, upstream_error) or error.args[0] != 429:
raise AssertionError(f"{route} returned the wrong 429 error: {error!r}")
@ -248,13 +261,13 @@ def exercise_sync(native: object, api_base: str) -> None:
assert_success(route, function(**route_kwargs(route, api_base, "success")))
try:
function(**route_kwargs(route, api_base, "429"))
except (RuntimeError, native.RustUpstreamError) as error:
except (RuntimeError, native.RustUpstreamError, litellm.RateLimitError) as error:
assert_rate_limit(native, route, error)
else:
raise AssertionError(f"{route} accepted a 429 response")
assert_success("ocr", native.ocr(**azure_ocr_kwargs(api_base)))
di_response: Final = native.ocr(**azure_di_kwargs(api_base))
assert di_response["provider_native_response"]["status"] == "succeeded"
assert di_response.get_provider_native_response()["status"] == "succeeded"
async def exercise_async(native: object, api_base: str) -> None:
@ -263,13 +276,13 @@ async def exercise_async(native: object, api_base: str) -> None:
assert_success(route, await function(**route_kwargs(route, api_base, "success")))
try:
await function(**route_kwargs(route, api_base, "429"))
except (RuntimeError, native.RustUpstreamError) as error:
except (RuntimeError, native.RustUpstreamError, litellm.RateLimitError) as error:
assert_rate_limit(native, route, error)
else:
raise AssertionError(f"a{route} accepted a 429 response")
assert_success("ocr", await native.aocr(**azure_ocr_kwargs(api_base)))
di_response: Final = await native.aocr(**azure_di_kwargs(api_base))
assert di_response["provider_native_response"]["status"] == "succeeded"
assert di_response.get_provider_native_response()["status"] == "succeeded"
async def exercise_async_concurrency(native: object, api_base: str) -> None:

View file

@ -1,4 +1,4 @@
from collections.abc import Generator, Mapping
from collections.abc import Generator
from typing import Final
from unittest.mock import AsyncMock, Mock
@ -8,8 +8,7 @@ import litellm
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.ocr import legacy
from litellm.rust_bridge import bindings, configuration
from litellm.rust_bridge.ocr import LiteLLMOcrRequest
from litellm.rust_bridge.ocr_lifecycle import NATIVE_OCR_LIFECYCLE
from litellm.rust_bridge.ocr import NATIVE_AOCR, NATIVE_OCR
@pytest.fixture(autouse=True)
@ -17,7 +16,8 @@ def isolated_ocr_configuration(monkeypatch: pytest.MonkeyPatch) -> Generator[Non
monkeypatch.delenv("LITELLM_RUST", raising=False)
configuration.reset_rust_configuration()
yield
NATIVE_OCR_LIFECYCLE.reset()
NATIVE_OCR.reset()
NATIVE_AOCR.reset()
configuration.reset_rust_configuration()
@ -27,7 +27,7 @@ async def test_unavailable_native_uses_legacy(monkeypatch: pytest.MonkeyPatch, a
response: Final = OCRResponse(pages=[], model="mistral-ocr-latest")
fallback: Final = AsyncMock(return_value=response) if asynchronous else Mock(return_value=response)
monkeypatch.setattr(legacy, "aocr" if asynchronous else "ocr", fallback)
NATIVE_OCR_LIFECYCLE.override(None)
(NATIVE_AOCR if asynchronous else NATIVE_OCR).override(None)
document: Final = {"type": "document_url", "document_url": "https://example.com"}
result: Final = (
@ -44,13 +44,13 @@ def test_admitted_failure_is_returned_without_replay() -> None:
failure: Final = RuntimeError("admitted")
native: Final = Mock(side_effect=failure)
litellm.rust(True)
NATIVE_OCR_LIFECYCLE.override(native)
NATIVE_OCR.override(native)
try:
with pytest.raises(RuntimeError) as caught:
litellm.ocr("mistral/mistral-ocr-latest", {"type": "document_url", "document_url": "https://example.com"})
assert caught.value is failure
finally:
NATIVE_OCR_LIFECYCLE.reset()
NATIVE_OCR.reset()
litellm.rust(None)
assert native.call_count == 1
@ -59,52 +59,40 @@ def test_public_binding_keeps_positional_fields_and_defaults_out_of_native_hook_
document: Final = {"type": "document_url", "document_url": "https://example.com"}
captured: Final = []
def native(
request: LiteLLMOcrRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
asynchronous: bool,
) -> OCRResponse:
captured.append((request, args, kwargs, asynchronous))
return OCRResponse(pages=[], model=request.model)
def native(*args: object, **kwargs: object) -> OCRResponse:
captured.append((args, kwargs))
return OCRResponse(pages=[], model=args[0])
litellm.rust(True)
NATIVE_OCR_LIFECYCLE.override(native)
NATIVE_OCR.override(native)
try:
response: Final = litellm.ocr("mistral/mistral-ocr-latest", document)
finally:
NATIVE_OCR_LIFECYCLE.reset()
NATIVE_OCR.reset()
litellm.rust(None)
request, call_args, hook_kwargs, asynchronous = captured[0]
call_args, hook_kwargs = captured[0]
assert response.model == "mistral/mistral-ocr-latest"
assert request.model == "mistral/mistral-ocr-latest"
assert request.document is document
assert call_args[1] is document
assert call_args == ("mistral/mistral-ocr-latest", document)
assert hook_kwargs == {}
assert asynchronous is False
def test_public_binding_keeps_keyword_model_and_document_in_native_hook_kwargs() -> None:
document: Final = {"type": "document_url", "document_url": "https://example.com"}
captured: Final = []
def native(
request: LiteLLMOcrRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
asynchronous: bool,
) -> OCRResponse:
def native(*args: object, **kwargs: object) -> OCRResponse:
assert args == ()
captured.append(kwargs)
return OCRResponse(pages=[], model=request.model)
return OCRResponse(pages=[], model=kwargs["model"])
litellm.rust(True)
NATIVE_OCR_LIFECYCLE.override(native)
NATIVE_OCR.override(native)
try:
litellm.ocr(model="mistral/mistral-ocr-latest", document=document)
finally:
NATIVE_OCR_LIFECYCLE.reset()
NATIVE_OCR.reset()
litellm.rust(None)
assert captured[0]["model"] == "mistral/mistral-ocr-latest"
@ -117,12 +105,12 @@ def test_public_duplicate_argument_error_does_not_depend_on_native_selection(ena
native: Final = Mock(side_effect=AssertionError("binding errors precede admission"))
document: Final = {"type": "document_url", "document_url": "https://example.com"}
litellm.rust(enabled)
NATIVE_OCR_LIFECYCLE.override(native)
NATIVE_OCR.override(native)
try:
with pytest.raises(TypeError, match=r"ocr\(\) got multiple values for argument 'model'"):
litellm.ocr("mistral/mistral-ocr-latest", document, model="duplicate")
finally:
NATIVE_OCR_LIFECYCLE.reset()
NATIVE_OCR.reset()
litellm.rust(None)
assert native.call_count == 0
@ -131,12 +119,12 @@ def test_public_duplicate_argument_error_does_not_depend_on_native_selection(ena
def test_public_missing_required_argument_error_does_not_depend_on_native_selection(enabled: bool) -> None:
native: Final = Mock(side_effect=AssertionError("binding errors precede admission"))
litellm.rust(enabled)
NATIVE_OCR_LIFECYCLE.override(native)
NATIVE_OCR.override(native)
try:
with pytest.raises(TypeError, match=r"ocr\(\) missing 1 required positional argument: 'document'"):
litellm.ocr("mistral/mistral-ocr-latest")
finally:
NATIVE_OCR_LIFECYCLE.reset()
NATIVE_OCR.reset()
litellm.rust(None)
assert native.call_count == 0
@ -177,7 +165,7 @@ async def test_native_is_enabled_by_default(
monkeypatch.setenv("LITELLM_RUST", environment)
response: Final = OCRResponse(pages=[], model="mistral-ocr-latest")
native: Final = AsyncMock(return_value=response) if asynchronous else Mock(return_value=response)
NATIVE_OCR_LIFECYCLE.override(native)
(NATIVE_AOCR if asynchronous else NATIVE_OCR).override(native)
fallback: Final = Mock(side_effect=AssertionError("legacy must not run"))
monkeypatch.setattr(legacy, "aocr" if asynchronous else "ocr", fallback)
@ -204,7 +192,7 @@ async def test_only_native_declines_replay_on_legacy(
) -> None:
failure: Final = Declined("unsupported") if declined else RuntimeError("provider already called")
native: Final = AsyncMock(side_effect=failure) if asynchronous else Mock(side_effect=failure)
NATIVE_OCR_LIFECYCLE.override(native)
(NATIVE_AOCR if asynchronous else NATIVE_OCR).override(native)
import importlib
main: Final = importlib.import_module("litellm.ocr.main")

View file

@ -64,6 +64,45 @@ def ocr_server(recording_server: RecordingServer) -> RecordingServer:
return recording_server
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
async def test_native_bindings_run_callbacks_and_send_their_mutations(
ocr_server: RecordingServer, asynchronous: bool
) -> None:
from litellm.rust_bridge import _native
from tests.test_litellm_rust.support.requests import ocr_arguments
pages: Final = [0]
marker: Final = object()
observed: Final = []
class Observe(Logging):
def pre_call(self, input, api_key, additional_args):
body: Final = additional_args["complete_input_dict"]
assert body["pages"] is pages
observed.append(self.model_call_details["litellm_params"]["metadata"]["marker"])
body["pages"].append(2)
additional_args["headers"]["x-callback"] = "native"
logger: Final = Observe(
model="mistral-ocr-latest",
messages=[],
stream=False,
call_type="aocr" if asynchronous else "ocr",
start_time=datetime.datetime.now(),
litellm_call_id="direct-native",
function_id="direct-native",
)
arguments: Final = ocr_arguments(ocr_server, pages=pages, metadata={"marker": marker}, litellm_logging_obj=logger)
response: Final = await _native.aocr(**arguments) if asynchronous else _native.ocr(**arguments)
assert response.pages[0].markdown == "native OCR response"
assert observed == [marker] and observed[0] is marker
assert len(ocr_server.requests) == 1
assert ocr_server.requests[0].body["pages"] == [0, 2]
assert ocr_server.requests[0].headers["x-callback"] == "native"
@pytest.mark.asyncio
async def test_proxy_metadata_remains_python_owned(ocr_server: RecordingServer) -> None:
from litellm.proxy._types import UserAPIKeyAuth
@ -79,6 +118,21 @@ async def test_proxy_metadata_remains_python_owned(ocr_server: RecordingServer)
assert "metadata" not in ocr_server.requests[0].body
@pytest.mark.parametrize("name", ["ocr", "aocr"])
def test_native_bindings_reject_argument_errors_before_consuming_inputs(name: str) -> None:
from litellm.rust_bridge import _native
class File:
def read(self):
raise AssertionError("invalid binding must not read the file")
native: Final = getattr(_native, name)
with pytest.raises(TypeError, match=rf"{name}\(\) got multiple values for argument 'model'"):
native("mistral/model", {"type": "file", "file": File()}, model="duplicate")
with pytest.raises(TypeError, match=rf"{name}\(\) missing 1 required positional argument: 'document'"):
native("mistral/model")
@pytest.mark.asyncio
async def test_request_level_custom_pricing_reaches_logging_params_and_bills_the_call(
ocr_server: RecordingServer,
@ -623,7 +677,6 @@ async def test_retained_argument_aliases_and_body_roots_survive_envelope_replace
def test_unstarted_native_coroutine_releases_input_without_reading_file(ocr_server: RecordingServer) -> None:
from litellm.ocr.main import _public_request
from litellm.rust_bridge import _native
ocr_server.expected_requests = 0
@ -637,7 +690,7 @@ def test_unstarted_native_coroutine_releases_input_without_reading_file(ocr_serv
def create():
file: Final = File()
kwargs: Final = {"model": "mistral/mistral-ocr-latest", "document": {"type": "file", "file": file}}
coroutine: Final = _native._ocr_lifecycle(_public_request("aocr", (), kwargs), (), kwargs, True)
coroutine: Final = _native.aocr(**kwargs)
file.owner = coroutine
coroutine.close()
return weakref.ref(file)

View file

@ -20,7 +20,6 @@ from litellm.llms.mistral.ocr.transformation import MistralOCRConfig
from litellm.llms.reducto.ocr.transformation import ReductoParseLegacyConfig, ReductoParseV3Config
from litellm.llms.vertex_ai.ocr.deepseek_transformation import VertexAIDeepSeekOCRConfig
from litellm.llms.vertex_ai.ocr.transformation import VertexAIOCRConfig
from litellm.rust_bridge import ocr as rust_ocr_bridge
from tests.test_litellm_rust.support.recording_server import ResponseSpec, recording_service
pytestmark = pytest.mark.requires_rust_extension
@ -168,6 +167,8 @@ def test_native_provider_transforms_match_python_shapes(
options: dict[str, object],
payload: dict[str, object],
) -> None:
from litellm.rust_bridge import _native
model: Final = (
"deepseek-ocr"
if isinstance(config, VertexAIDeepSeekOCRConfig)
@ -197,17 +198,17 @@ def test_native_provider_transforms_match_python_shapes(
).model_dump()
with recording_service() as server:
server.enqueue(ResponseSpec(body=payload))
response: Final = rust_ocr_bridge.ocr(
response: Final = _native.ocr(
model=model,
document={**document},
api_key="test-key",
api_base=server.base_url,
custom_llm_provider=provider,
extra_headers=None,
optional_params={**options, **({"vertex_project": "project"} if provider == "vertex_ai" else {})},
timeout=3,
**{**options, **({"vertex_project": "project"} if provider == "vertex_ai" else {})},
)
assert response == expected_response
assert response.model_dump() == expected_response
assert server.requests[0].body == expected_request
@ -272,24 +273,25 @@ def ocr_server() -> Generator[tuple[ThreadingHTTPServer, list[dict[str, object]]
def test_native_ocr_with_compiled_rust_extension(
ocr_server: tuple[ThreadingHTTPServer, list[dict[str, object]]],
) -> None:
from litellm.rust_bridge import _native
server, requests = ocr_server
address: Final = server.server_address
host: Final = str(address[0])
port: Final = int(address[1])
response: Final = rust_ocr_bridge.ocr(
response: Final = _native.ocr(
model="mistral-ocr-latest",
document={"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"},
api_key="test-key",
api_base=f"http://{host}:{port}",
custom_llm_provider="mistral",
extra_headers=None,
optional_params={},
timeout=None,
)
assert response is not None
assert response["pages"][0]["markdown"] == "native OCR response"
assert response.pages[0].markdown == "native OCR response"
assert len(requests) == 1
assert not requests[0]["headers"].get("user-agent", "").startswith("python-httpx")
assert requests[0]["body"] == {
@ -417,12 +419,18 @@ async def test_native_ocr_failures_do_not_retry_on_python(ocr_server, asynchrono
assert not requests[0]["headers"].get("user-agent", "").startswith("python-httpx")
@pytest.mark.parametrize("custom_provider", ["mistral", "not-a-provider"])
def test_native_ocr_rejects_invalid_input_before_network(ocr_server, custom_provider):
@pytest.mark.parametrize(
"custom_provider,error_type,message",
[
("mistral", litellm.BadRequestError, "document"),
("not-a-provider", litellm.APIConnectionError, "invalid provider"),
],
)
def test_native_ocr_rejects_invalid_input_before_network(ocr_server, custom_provider, error_type, message):
from litellm.rust_bridge import _native
server, requests = ocr_server
with pytest.raises(ValueError, match="Document URL is required"):
with pytest.raises(error_type, match=message):
_native.ocr(
model="mistral-ocr-latest",
custom_llm_provider=custom_provider,