feat(ocr): drive callbacks through native lifecycle

This commit is contained in:
Yujong Lee 2026-09-11 08:24:33 -07:00 • committed by yujonglee
parent 23e8cf4825
commit c4c89796e8
18 changed files with 229 additions and 140 deletions

View file

@ -64,6 +64,7 @@ async fn execute_ocr_provider_call<A: OcrAdapter>(
let decoded = adapter
.read_response(client, response, &url, &headers, &request)
.await?;
request.hooks.post_call(&decoded.text).await?;
let response = adapter.transform_ocr_response(&request, decoded.data)?;
Ok(LiteLLMOcrResponse {
provider_native_response: decoded.native,

View file

@ -28,7 +28,7 @@ pub struct OcrDuringCallRequest {
pub body: Value,
}
pub struct OcrPreparedRequest {
pub struct OcrRequestDraft {
pub model: String,
pub url: String,
pub headers: Vec<(String, String)>,
@ -36,12 +36,12 @@ pub struct OcrPreparedRequest {
}
pub trait OcrHooks: Send + Sync {
fn prepared_request(
&self,
request: OcrPreparedRequest,
) -> OcrHookFuture<'_, OcrPreparedRequest> {
fn before_send(&self, request: OcrRequestDraft) -> OcrHookFuture<'_, OcrRequestDraft> {
Box::pin(async move { Ok(request) })
}
fn post_call<'a>(&'a self, _response: &'a str) -> OcrHookFuture<'a, ()> {
Box::pin(async { Ok(()) })
}
fn pre_call(&self, request: OcrPreCallRequest) -> OcrHookFuture<'_, OcrPreCallRequest> {
Box::pin(async move { Ok(request) })
}

View file

@ -3,7 +3,7 @@ use serde_json::{Map, Value};
use super::OcrClient;
use super::error::{OcrError, OcrRequestError};
use super::hooks::{OcrDuringCallRequest, OcrPreparedRequest};
use super::hooks::{OcrDuringCallRequest, OcrRequestDraft};
use super::types::{LiteLLMOcrRequest, OcrDocument};
#[derive(Debug, Deserialize)]
@ -94,9 +94,9 @@ pub(crate) async fn build_http_request<B>(
where
B: Serialize + DeserializeOwned,
{
let prepared = request
let draft = request
.hooks
.prepared_request(OcrPreparedRequest {
.before_send(OcrRequestDraft {
model: request.model.clone(),
url: url.into(),
headers: headers.to_vec(),
@ -105,15 +105,15 @@ where
})?,
})
.await?;
let body: B = super::wire::decode_request_value(prepared.body, "guardrail.body")?;
let body: B = super::wire::decode_request_value(draft.body, "guardrail.body")?;
let builder = client
.provider_http()
.post(&prepared.url)
.post(&draft.url)
.json(&body)
.timeout(request.connection.timeout);
crate::http_utils::with_headers(
builder,
&prepared.headers,
&draft.headers,
crate::http_utils::HeaderPolicy::All,
)
.build()

View file

@ -17,6 +17,7 @@ use serde_json::{Map, Value};
pub struct DecodedOcrResponse<T> {
pub data: T,
pub native: Option<Value>,
pub text: String,
}
#[derive(Deserialize)]
@ -135,7 +136,11 @@ pub fn decode_response<T: DeserializeOwned>(
} else {
None
};
Ok(DecodedOcrResponse { data, native })
Ok(DecodedOcrResponse {
data,
native,
text: String::from_utf8_lossy(bytes).into_owned(),
})
}
pub fn decode_pre_call_result(

View file

@ -3,7 +3,7 @@ use std::sync::{Arc, Mutex};
use serde_json::{Value, json};
use super::OcrClient;
use super::hooks::{OcrHookFuture, OcrHooks, OcrLogFuture, OcrPreCallRequest, OcrPreparedRequest};
use super::hooks::{OcrHookFuture, OcrHooks, OcrLogFuture, OcrPreCallRequest, OcrRequestDraft};
use super::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
use super::wire::{OcrWireRequest, decode_request};
use crate::call_lifecycle::{CallLifecycleContext, CallLifecycleTiming};
@ -123,13 +123,10 @@ struct RecordingHooks {
block: bool,
}
struct EditPreparedRequest;
struct EditRequestDraft;
impl OcrHooks for EditPreparedRequest {
fn prepared_request(
&self,
mut request: OcrPreparedRequest,
) -> OcrHookFuture<'_, OcrPreparedRequest> {
impl OcrHooks for EditRequestDraft {
fn before_send(&self, mut request: OcrRequestDraft) -> OcrHookFuture<'_, OcrRequestDraft> {
Box::pin(async move {
assert_eq!(request.model, "model");
assert!(request.url.ends_with("/v1/ocr"));
@ -148,14 +145,14 @@ impl OcrHooks for EditPreparedRequest {
}
#[tokio::test]
async fn prepared_request_hook_edits_wire_body_and_headers_without_guardrails() {
async fn before_send_hook_edits_wire_body_and_headers_without_guardrails() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let request = wire_request(
"mistral/model",
&base,
json!({"include_image_base64":false}),
)
.with_host_hooks(Arc::new(EditPreparedRequest), None);
.with_host_hooks(Arc::new(EditRequestDraft), None);
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
@ -166,12 +163,9 @@ async fn prepared_request_hook_edits_wire_body_and_headers_without_guardrails()
}
impl OcrHooks for RecordingHooks {
fn prepared_request(
&self,
request: OcrPreparedRequest,
) -> OcrHookFuture<'_, OcrPreparedRequest> {
fn before_send(&self, request: OcrRequestDraft) -> OcrHookFuture<'_, OcrRequestDraft> {
Box::pin(async move {
self.events.lock().unwrap().push("prepared");
self.events.lock().unwrap().push("before_send");
Ok(request)
})
}
@ -235,7 +229,7 @@ async fn lifecycle_orders_hooks_and_emits_one_success() {
server.await.unwrap();
assert_eq!(
*events.lock().unwrap(),
["pre", "during", "prepared", "success"]
["pre", "during", "before_send", "success"]
);
assert_eq!(seen.lock().unwrap().len(), 1);
}
@ -277,7 +271,7 @@ async fn upstream_failure_emits_one_terminal_failure() {
server.await.unwrap();
assert_eq!(
*events.lock().unwrap(),
["pre", "during", "prepared", "failure"]
["pre", "during", "before_send", "failure"]
);
assert_eq!(seen.lock().unwrap().len(), 1);
}
@ -333,7 +327,7 @@ async fn every_adapter_runs_the_complete_lifecycle() {
server.await.unwrap();
assert_eq!(
*events.lock().unwrap(),
["pre", "during", "prepared", "success"],
["pre", "during", "before_send", "success"],
"{model}"
);
assert_eq!(seen.lock().unwrap().len(), 1, "{model}");
@ -343,7 +337,7 @@ async fn every_adapter_runs_the_complete_lifecycle() {
#[derive(Clone, Copy, Debug)]
enum FailureStage {
During,
Prepared,
BeforeSend,
InvalidBody,
Preparation,
Response,
@ -372,17 +366,14 @@ impl OcrHooks for FailingHooks {
})
}
fn prepared_request(
&self,
request: OcrPreparedRequest,
) -> OcrHookFuture<'_, OcrPreparedRequest> {
fn before_send(&self, request: OcrRequestDraft) -> OcrHookFuture<'_, OcrRequestDraft> {
Box::pin(async move {
let request = self.recording.prepared_request(request).await?;
let request = self.recording.before_send(request).await?;
match self.stage {
FailureStage::Prepared => Err(crate::Error::InvalidRequest(
"blocked prepared request".into(),
FailureStage::BeforeSend => Err(crate::Error::InvalidRequest(
"blocked before_send request".into(),
)),
FailureStage::InvalidBody => Ok(OcrPreparedRequest {
FailureStage::InvalidBody => Ok(OcrRequestDraft {
body: json!({"document":null}),
..request
}),
@ -416,7 +407,7 @@ impl OcrHooks for FailingHooks {
#[rstest::rstest]
#[case(FailureStage::During)]
#[case(FailureStage::Prepared)]
#[case(FailureStage::BeforeSend)]
#[case(FailureStage::InvalidBody)]
#[case(FailureStage::Preparation)]
#[case(FailureStage::Response)]
@ -454,7 +445,7 @@ async fn lifecycle_reports_failures_once_at_each_boundary(#[case] stage: Failure
let expected = match stage {
FailureStage::Preparation => vec!["pre", "failure"],
FailureStage::During => vec!["pre", "during", "failure"],
_ => vec!["pre", "during", "prepared", "failure"],
_ => vec!["pre", "during", "before_send", "failure"],
};
assert_eq!(*events.lock().unwrap(), expected);
assert_eq!(

View file

@ -89,9 +89,9 @@ pub(crate) fn ocr_route_error(err: Error) -> BridgeError {
Error::MissingField("document_url" | "image_url") => {
BridgeError::InvalidArgument("Document URL is required".into())
}
Error::Http { status, .. } => BridgeError::Upstream {
Error::Http { status, body } => BridgeError::Upstream {
status: Some(status),
message: "OCR provider request failed".into(),
message: body,
},
other => required_route_error(other),
}
@ -147,7 +147,7 @@ mod ocr_error_tests {
}
#[test]
fn ocr_errors_preserve_status_without_provider_body() {
fn ocr_errors_preserve_provider_body() {
Python::initialize();
Python::attach(|py| {
for field in ["document_url", "image_url"] {
@ -166,7 +166,7 @@ mod ocr_error_tests {
.getattr("args")
.and_then(|args| args.extract())
.expect("OCR failures retain status and unprefixed provider message");
assert_eq!(args, (429, "OCR provider request failed".to_string()));
assert_eq!(args, (429, r#"{"message":"rate limited"}"#.to_string()));
});
}
}

View file

@ -225,7 +225,7 @@ mod tests {
(
"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, logging_obj=None, callback_loop=None, token_provider=None)",
"(model, document, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, input_sources=None, timeout_seconds=None, logging_obj=None, callback_loop=None, token_provider=None, call_completion=None)",
),
(
"transcription",

View file

@ -5,7 +5,7 @@ use pyo3::prelude::*;
use serde_json::Value;
use std::future::Future;
use crate::errors::fallback_route_error;
use crate::errors::required_route_error;
use crate::marshal::{RouteOptions, RouteOptionsInputs, required_value};
fn prepare_messages(
@ -61,5 +61,5 @@ bridge_route! {
timeout_seconds: Option<f64>,
},
prepare = prepare_messages,
errors = fallback_route_error,
errors = required_route_error,
}

View file

@ -42,6 +42,7 @@ fn prepare_ocr(
})
.transpose()?,
error: Mutex::new(None),
execution_body: Mutex::new(None),
}))
})
})
@ -84,6 +85,10 @@ fn prepare_ocr(
timeout_seconds: timeout.map(|value| value.as_secs_f64()),
})
.map_err(ocr_route_error)?;
if let Some(call_completion) = &inputs.call_completion {
callbacks::NativeOcrCompletion::attach(call_completion)
.map_err(BridgeError::Host)?;
}
if let Some(hooks) = &hooks
&& hooks.token_provider.is_some()
{
@ -151,6 +156,7 @@ bridge_route! {
logging_obj: Option<Py<PyAny>>,
callback_loop: Option<Py<PyAny>>,
token_provider: Option<Py<PyAny>>,
call_completion: Option<Py<PyAny>>,
},
prepare = prepare_ocr,
errors = std::convert::identity,

View file

@ -2,7 +2,7 @@ use std::sync::Mutex;
use litellm_core::Error;
use litellm_core::auth::{ResolvedCredential, SecretValue, TokenFuture, TokenProvider};
use litellm_core::ocr::hooks::{OcrHookFuture, OcrHooks, OcrPreparedRequest};
use litellm_core::ocr::hooks::{OcrHookFuture, OcrHooks, OcrRequestDraft};
use litellm_python_interop::{from_py, to_py};
use pyo3::prelude::*;
use pyo3::types::{PyDict, PyTuple};
@ -30,6 +30,77 @@ pub(super) struct PythonOcrHooks {
pub api_key: Option<String>,
pub locals: Option<TaskLocals>,
pub error: Mutex<Option<PyErr>>,
pub execution_body: Mutex<Option<Py<PyDict>>>,
}
#[pyclass]
pub(super) struct NativeOcrCompletion {
python_completion: Py<PyAny>,
}
impl NativeOcrCompletion {
pub(super) fn attach(call_completion: &Py<PyAny>) -> PyResult<()> {
Python::attach(|py| {
let python_completion = call_completion.getattr(py, "python_implementation")?;
let native = Py::new(py, Self { python_completion })?;
let attached: bool = call_completion
.call_method1(py, "attach", (native,))?
.extract(py)?;
if attached {
Ok(())
} else {
Err(pyo3::exceptions::PyRuntimeError::new_err(
"OCR completion was already attached",
))
}
})
}
}
#[pymethods]
impl NativeOcrCompletion {
fn success(
&self,
py: Python<'_>,
result: Py<PyAny>,
start_time: Py<PyAny>,
end_time: Py<PyAny>,
) -> PyResult<()> {
self.python_completion
.call_method1(py, "success", (result, start_time, end_time))?;
Ok(())
}
fn failure(
&self,
py: Python<'_>,
exception: Py<PyAny>,
traceback_exception: String,
start_time: Py<PyAny>,
end_time: Py<PyAny>,
) -> PyResult<()> {
self.python_completion.call_method1(
py,
"failure",
(exception, traceback_exception, start_time, end_time),
)?;
Ok(())
}
fn async_failure(
&self,
py: Python<'_>,
exception: Py<PyAny>,
traceback_exception: String,
start_time: Py<PyAny>,
end_time: Py<PyAny>,
) -> PyResult<Py<PyAny>> {
self.python_completion.call_method1(
py,
"async_failure",
(exception, traceback_exception, start_time, end_time),
)
}
}
impl PythonOcrHooks {
@ -75,7 +146,7 @@ impl PythonOcrHooks {
fn callback(
&self,
py: Python<'_>,
request: &OcrPreparedRequest,
request: &OcrRequestDraft,
) -> PyResult<(Py<PyAny>, Py<PyDict>, Py<PyDict>)> {
let body = to_py(py, &request.body)?
.into_bound(py)
@ -113,19 +184,23 @@ impl PythonOcrHooks {
}
impl OcrHooks for PythonOcrHooks {
fn prepared_request(
&self,
request: OcrPreparedRequest,
) -> OcrHookFuture<'_, OcrPreparedRequest> {
fn before_send(&self, request: OcrRequestDraft) -> OcrHookFuture<'_, OcrRequestDraft> {
Box::pin(async move {
if self.logger.is_none() {
return Ok(request);
}
let result: PyResult<OcrPreparedRequest> = async {
let result: PyResult<OcrRequestDraft> = async {
let (callback, body, headers) = Python::attach(|py| self.callback(py, &request))?;
self.invoke(callback).await?;
Python::attach(|py| {
Ok(OcrPreparedRequest {
*self
.execution_body
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner()) =
Some(body.clone_ref(py));
});
Python::attach(|py| {
Ok(OcrRequestDraft {
body: from_py(body.bind(py).as_any())?,
headers: headers
.bind(py)
@ -143,6 +218,46 @@ impl OcrHooks for PythonOcrHooks {
})
})
}
fn post_call<'a>(&'a self, response: &'a str) -> OcrHookFuture<'a, ()> {
Box::pin(async move {
let Some(logger) = &self.logger else {
return Ok(());
};
let result: PyResult<()> = async {
let callback = Python::attach(|py| {
let kwargs = PyDict::new(py);
kwargs.set_item("api_key", &self.api_key)?;
kwargs.set_item("original_response", response)?;
let additional = PyDict::new(py);
let body = self
.execution_body
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
additional.set_item(
"complete_input_dict",
body.as_ref().ok_or_else(|| {
pyo3::exceptions::PyRuntimeError::new_err("missing OCR execution body")
})?,
)?;
kwargs.set_item("additional_args", additional)?;
Ok::<_, PyErr>(
py.import("functools")?
.getattr("partial")?
.call((logger.bind(py).getattr("post_call")?,), Some(&kwargs))?
.unbind(),
)
})?;
self.invoke(callback).await?;
Ok(())
}
.await;
result.map_err(|error| {
self.retain_error(error);
Error::InvalidRequest("OCR post-call hook failed".into())
})
})
}
}
impl std::fmt::Debug for PythonOcrHooks {

View file

@ -2476,21 +2476,12 @@ class BaseLLMHTTPHandler:
extra_headers=headers,
timeout=timeout,
)
except Exception as rust_error: # noqa: BLE001 # only explicit pre-dispatch declines permit fallback
from litellm.rust_bridge.bindings import native_exception_types, upstream_error_details
exceptions: Final = native_exception_types()
if exceptions is not None and isinstance(rust_error, exceptions[0]):
return None
if exceptions is not None and isinstance(rust_error, exceptions[1]):
status, message = upstream_error_details(rust_error)
raise litellm.APIError(
status_code=status,
message=message,
llm_provider=custom_llm_provider,
model=model,
) from rust_error
raise
except Exception as rust_error: # noqa: BLE001 # rollout-safety fallback: any Rust bridge failure must fall back to the Python path
verbose_logger.debug(
"Rust Anthropic messages bridge raised %s; falling back to Python path",
type(rust_error).__name__,
)
return None
if rust_response is None:
return None

View file

@ -26,7 +26,6 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo
convert_to_model_response_object,
)
from litellm.llms.bedrock.request_metadata import bedrock_request_metadata_is_owned
from litellm.rust_bridge.bindings import native_exception_types, upstream_error_details
from litellm.rust_bridge.configuration import rust_enabled
from litellm.rust_bridge.loader import get_native_bridge
from litellm.rust_bridge.timeouts import timeout_to_seconds
@ -275,6 +274,17 @@ def rust_chat_completions_accepts(
return True
def _rust_bridge_exceptions() -> tuple[type[BaseException], type[BaseException]] | None:
native_bridge: Final = get_native_bridge()
if native_bridge is None:
return None
declined: Final = getattr(native_bridge, "RustBridgeDeclined", None)
upstream: Final = getattr(native_bridge, "RustUpstreamError", None)
if declined is None or upstream is None:
return None
return declined, upstream
def _reraise_or_decline(
rust_error: BaseException,
*,
@ -288,14 +298,20 @@ def _reraise_or_decline(
second attempt bills for it twice. Those surface as an `APIError` carrying
the upstream status, which LiteLLM's exception mapping already understands.
"""
exceptions: Final = native_exception_types()
exceptions: Final = _rust_bridge_exceptions()
if exceptions is None:
raise rust_error
verbose_logger.debug(
"Rust chat completions bridge raised %s; falling back to Python path",
type(rust_error).__name__,
)
return
declined, upstream_failed = exceptions
if isinstance(rust_error, upstream_failed):
status, message = upstream_error_details(rust_error)
args: Final = rust_error.args
status: Final = args[0] if args else 0
message: Final = args[1] if len(args) > 1 else ""
raise APIError(
status_code=status,
status_code=int(status) or 500,
message=f"litellm rust chat completions: {message}",
llm_provider=custom_llm_provider or "",
model=model,

View file

@ -73,6 +73,7 @@ class RustOcr(Protocol):
logging_obj: _OCRLogging | None,
callback_loop: asyncio.AbstractEventLoop | None,
token_provider: object,
call_completion: object,
) -> dict[str, object]:
raise NotImplementedError
@ -92,6 +93,7 @@ class RustAocr(Protocol):
logging_obj: _OCRLogging | None,
callback_loop: asyncio.AbstractEventLoop | None,
token_provider: object,
call_completion: object,
) -> Awaitable[dict[str, object]]:
raise NotImplementedError
@ -347,10 +349,11 @@ def run(
optional_params=dict(marshalled.kwargs), # mutable-ok: PyO3 OCR binding requires a concrete dict
input_sources=marshalled.input_sources,
timeout=marshalled.timeout,
logging_obj=cast(
logging_obj=cast( # cast-ok: client decorator injects Logging
_OCRLogging, request.kwargs["litellm_logging_obj"]
), # cast-ok: client decorator injects Logging
),
token_provider=request.kwargs.get("azure_ad_token_provider"),
call_completion=request.call_completion,
)
except Exception as error:
mapped: Final = _map_error(error, request)
@ -379,10 +382,11 @@ async def arun(
optional_params=dict(marshalled.kwargs), # mutable-ok: PyO3 OCR binding requires a concrete dict
input_sources=marshalled.input_sources,
timeout=marshalled.timeout,
logging_obj=cast(
logging_obj=cast( # cast-ok: client decorator injects Logging
_OCRLogging, request.kwargs["litellm_logging_obj"]
), # cast-ok: client decorator injects Logging
),
token_provider=request.kwargs.get("azure_ad_token_provider"),
call_completion=request.call_completion,
)
except Exception as error:
mapped: Final = _map_error(error, request)
@ -405,6 +409,7 @@ def ocr(
input_sources: Mapping[str, str] | None = None,
logging_obj: _OCRLogging | None = None,
token_provider: object = None,
call_completion: object = None,
) -> dict[str, object] | None:
rust_ocr: Final = load_rust_ocr()
if rust_ocr is None:
@ -422,6 +427,7 @@ def ocr(
logging_obj=logging_obj,
callback_loop=None,
token_provider=token_provider,
call_completion=call_completion,
)
@ -438,6 +444,7 @@ async def aocr(
input_sources: Mapping[str, str] | None = None,
logging_obj: _OCRLogging | None = None,
token_provider: object = None,
call_completion: object = None,
) -> dict[str, object] | None:
rust_aocr: Final = load_rust_aocr()
if rust_aocr is None:
@ -455,4 +462,5 @@ async def aocr(
logging_obj=logging_obj,
callback_loop=asyncio.get_running_loop(),
token_provider=token_provider,
call_completion=call_completion,
)

View file

@ -6,7 +6,7 @@ from enum import Enum
from typing import Final, Generic, NoReturn, TypeAlias, TypeVar
from litellm.exceptions import APIError
from litellm.rust_bridge.bindings import native_exception_types, upstream_error_details
from litellm.rust_bridge.bindings import native_exception_types
NativeT = TypeVar("NativeT")
ResultT = TypeVar("ResultT")
@ -137,9 +137,13 @@ def _required_reason(result: RustDeclined | RustUnavailable) -> str:
def _raise_upstream(error: BaseException, context: BridgeErrorContext) -> NoReturn:
status, message = upstream_error_details(error)
args: Final[tuple[object, ...]] = error.args
status_value: Final = args[0] if args else 0
message_value: Final = args[1] if len(args) > 1 else str(error)
status: Final = status_value if isinstance(status_value, int) else 0
message: Final = message_value if isinstance(message_value, str) else str(message_value)
raise APIError(
status_code=status,
status_code=status or 500,
message=f"litellm rust {context.route}: {message}",
llm_provider=context.provider,
model=context.model,

View file

@ -238,66 +238,17 @@ async def test_gate_invokes_rust_and_marks_response_header():
@pytest.mark.asyncio
async def test_gate_propagates_unclassified_bridge_failure():
async def test_gate_falls_back_to_python_when_bridge_raises():
bridge = RaisingAsyncMessages()
litellm.rust(True)
rust_messages.set_rust_messages(amessages=bridge)
with pytest.raises(RuntimeError, match="upstream request failed"):
await _gate()
response = await _gate()
assert response is None
assert bridge.calls == 1
@pytest.mark.asyncio
@pytest.mark.parametrize("status", [0, 400, 429, 500])
async def test_gate_never_falls_back_after_possible_dispatch(monkeypatch, status):
from types import SimpleNamespace
from litellm.rust_bridge import bindings
class Declined(Exception):
pass
class Upstream(Exception):
pass
error = Upstream(status, "provider failure")
async def bridge(**kwargs):
raise error
monkeypatch.setattr(
bindings, "get_native_bridge", lambda: SimpleNamespace(RustBridgeDeclined=Declined, RustUpstreamError=Upstream)
)
litellm.rust(True)
rust_messages.set_rust_messages(amessages=bridge)
with pytest.raises(litellm.APIError) as caught:
await _gate()
assert caught.value.status_code == (status or 500)
assert caught.value.__cause__ is error
@pytest.mark.asyncio
async def test_gate_falls_back_only_for_explicit_decline(monkeypatch):
from types import SimpleNamespace
from litellm.rust_bridge import bindings
class Declined(Exception):
pass
class Upstream(Exception):
pass
async def bridge(**kwargs):
raise Declined("unsupported before dispatch")
monkeypatch.setattr(
bindings, "get_native_bridge", lambda: SimpleNamespace(RustBridgeDeclined=Declined, RustUpstreamError=Upstream)
)
litellm.rust(True)
rust_messages.set_rust_messages(amessages=bridge)
assert await _gate() is None
@pytest.mark.asyncio
async def test_gate_skips_rust_when_flag_absent():
bridge = ExplodingAsyncMessages()

View file

@ -65,6 +65,7 @@ class RecordingBridge:
logging_obj: object = None,
callback_loop: asyncio.AbstractEventLoop | None = None,
token_provider: object = None,
call_completion: object = None,
) -> dict[str, object]:
self.logging_obj = logging_obj
self.calls.append(
@ -103,6 +104,7 @@ class RecordingAsyncBridge:
logging_obj: object = None,
callback_loop: asyncio.AbstractEventLoop | None = None,
token_provider: object = None,
call_completion: object = None,
) -> dict[str, object]:
self.calls.append(
{
@ -135,6 +137,7 @@ class RaisingBridge:
logging_obj: object = None,
callback_loop: asyncio.AbstractEventLoop | None = None,
token_provider: object = None,
call_completion: object = None,
) -> dict[str, object]:
raise RuntimeError("bridge failed")
@ -154,6 +157,7 @@ class RaisingAsyncBridge:
logging_obj: object = None,
callback_loop: asyncio.AbstractEventLoop | None = None,
token_provider: object = None,
call_completion: object = None,
) -> dict[str, object]:
raise RuntimeError("bridge failed")

View file

@ -55,9 +55,6 @@ class _FakeNative:
def _fake_native_bridge(monkeypatch):
"""Expose the bridge's exception classes without the compiled extension."""
monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative())
from litellm.rust_bridge import bindings
monkeypatch.setattr(bindings, "get_native_bridge", lambda: _FakeNative())
def _hide_native_bridge(monkeypatch):

View file

@ -149,7 +149,7 @@ def test_native_ocr_normalizes_provider_response_model_and_usage(ocr_server: Rec
assert response.usage_info.pages_processed == 1
def test_native_ocr_maps_provider_400_without_exposing_response_body(ocr_server: RecordingServer) -> None:
def test_native_ocr_maps_provider_400_with_response_body(ocr_server: RecordingServer) -> None:
ocr_server.enqueue(ResponseSpec(body={"message": "invalid OCR request"}, status=400))
with pytest.raises(litellm.BadRequestError) as caught:
@ -158,7 +158,7 @@ def test_native_ocr_maps_provider_400_without_exposing_response_body(ocr_server:
assert caught.value.status_code == 400
assert caught.value.model == "mistral-ocr-latest"
assert caught.value.llm_provider == "mistral"
assert "invalid OCR request" not in str(caught.value)
assert "invalid OCR request" in str(caught.value)
def test_native_ocr_raises_transport_error_when_request_exceeds_timeout(ocr_server: RecordingServer) -> None: