mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
feat(ocr): drive callbacks through native lifecycle
This commit is contained in:
parent
23e8cf4825
commit
c4c89796e8
18 changed files with 229 additions and 140 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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) })
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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!(
|
||||
|
|
|
|||
|
|
@ -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()));
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue