This commit is contained in:
yujonglee 2026-09-06 08:05:11 +00:00 committed by GitHub
commit d2fc8b2438
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
59 changed files with 3765 additions and 135 deletions

View file

@ -10,6 +10,10 @@ on:
- ".github/scripts/smoke_test_native_wheel.py"
- ".github/scripts/verify_linux_native_wheel.py"
- "tests/test_litellm/rust_bridge/native_route_wheel_test.py"
- "tests/test_litellm/rust_bridge/sdk_callback_wheel_test.py"
- "tests/test_litellm/ocr/**"
- "litellm/ocr/**"
- "litellm/rust_bridge/**"
- ".github/workflows/test-rust.yml"
pull_request:
branches:
@ -25,6 +29,10 @@ on:
- ".github/scripts/smoke_test_native_wheel.py"
- ".github/scripts/verify_linux_native_wheel.py"
- "tests/test_litellm/rust_bridge/native_route_wheel_test.py"
- "tests/test_litellm/rust_bridge/sdk_callback_wheel_test.py"
- "tests/test_litellm/ocr/**"
- "litellm/ocr/**"
- "litellm/rust_bridge/**"
- ".github/workflows/test-rust.yml"
permissions:
@ -133,3 +141,9 @@ jobs:
- name: Test native route wheel
run: python tests/test_litellm/rust_bridge/native_route_wheel_test.py dist/*.whl
- name: Test installed SDK callback parity and fallback
run: |
uv venv /tmp/ocr-callback-sdk --python python
uv pip install --python /tmp/ocr-callback-sdk/bin/python dist/*.whl
/tmp/ocr-callback-sdk/bin/python tests/test_litellm/rust_bridge/sdk_callback_wheel_test.py

View file

@ -1482,10 +1482,12 @@ name = "litellm-python-interop"
version = "0.1.0"
dependencies = [
"pyo3",
"pyo3-async-runtimes",
"pythonize",
"rstest",
"serde",
"serde_json",
"tokio",
]
[[package]]

View file

@ -28,7 +28,7 @@ coverage and production evidence.
## Native request boundary
Native HTTP routes and Responses WebSocket connections accept
`native(request, *, options, context)`. The request carries only endpoint payload.
`native(request, *, options, context, callback_adapter=None)`. The request carries only endpoint payload.
`NativeRequestOptions` carries credentials, typed provider configuration, routing,
headers, query parameters, and timeout. `NativeRequestContext` carries call identity,
attribution, and typed capability facts separately from the provider payload.

View file

@ -1 +1 @@
pub use crate::ocr::{OcrRequest, ocr, ocr_provider_supported};
pub use crate::ocr::{OcrRequest, ocr, ocr_provider_supported, ocr_with_observer};

View file

@ -1,20 +1,39 @@
use litellm_core::call_lifecycle::CallLifecycleContext;
use litellm_core::error::Error;
use litellm_core::http_utils::http_request;
use litellm_core::ocr::transformation::OcrResponseHandling;
use litellm_core::provider_callbacks::ProviderAttemptObserver;
use litellm_core::provider_callbacks::handler::{
ProviderAttemptContext, ProviderRequest, send_provider_request,
};
use serde_json::Value;
use super::common_utils::{poll_document_intelligence, truncate_error_body};
use super::common_utils::poll_document_intelligence;
use super::hooks::OcrLifecycleHooks;
use super::types::PreparedOcrRequest;
use crate::client::http_client;
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(crate) async fn execute_ocr_provider_call(
pub(crate) async fn execute_ocr_provider_call<Observer>(
request: PreparedOcrRequest,
context: &CallLifecycleContext,
hooks: &OcrLifecycleHooks,
) -> Result<Value, Error> {
observer: &mut Observer,
) -> Result<Value, Error>
where
Observer: ProviderAttemptObserver,
Observer::Error: std::fmt::Display,
{
let request = hooks.prepare_provider_request(request).await?;
let mut request_builder = http_client().post(&request.url).json(&request.body);
let provider_request = ProviderRequest {
provider: request.custom_llm_provider.clone(),
model: request.model.clone(),
body: serde_json::from_value(request.body).map_err(|error| {
Error::InvalidRequest(format!("OCR provider request must be an object: {error}"))
})?,
api_base: request.url.clone(),
headers: request.upstream_headers.iter().cloned().collect(),
};
let mut request_builder = http_client().post(&request.url);
for (key, value) in &request.upstream_headers {
request_builder = request_builder.header(key, value);
}
@ -22,16 +41,24 @@ pub(crate) async fn execute_ocr_provider_call(
request_builder = request_builder.timeout(duration);
}
let response = http_request(request_builder)
.await
.map_err(|err| Error::Network(err.to_string()))?;
let response = send_provider_request(
request_builder,
provider_request,
ProviderAttemptContext {
call_id: context.litellm_call_id.clone(),
trace_id: None,
attempt: 1,
},
observer,
)
.await?;
let status = response.status();
let status = response.status;
if request.config.response_handling() == OcrResponseHandling::AzureDocumentIntelligencePoll
&& status.as_u16() == 202
{
let operation_url = response
.headers()
.headers
.get("operation-location")
.and_then(|value| value.to_str().ok())
.map(str::to_string)
@ -58,19 +85,7 @@ pub(crate) async fn execute_ocr_provider_call(
.into_json());
}
let text = response
.text()
.await
.map_err(|err| Error::Network(err.to_string()))?;
if !status.is_success() {
return Err(Error::Http {
status: status.as_u16(),
body: truncate_error_body(&text),
});
}
let response_json: Value = serde_json::from_str(&text)
let response_json: Value = serde_json::from_str(&response.body)
.map_err(|err| Error::InvalidResponse(format!("invalid OCR response JSON: {err}")))?;
Ok(request

View file

@ -127,6 +127,7 @@ impl OcrLifecycleHooks {
};
Ok(ProviderOcrRequest {
model,
custom_llm_provider,
config,
url,
body,

View file

@ -1,6 +1,7 @@
use crate::integrations::types::RequestHooks;
use litellm_core::Error;
use litellm_core::call_lifecycle::CallLifecycle;
use litellm_core::provider_callbacks::{NoopProviderAttemptObserver, ProviderAttemptObserver};
use litellm_core::request_context::LiteLlmRequestContext;
use litellm_core::request_options::RequestOptions;
use serde_json::Value;
@ -23,14 +24,41 @@ pub async fn ocr(
context: &LiteLlmRequestContext,
hooks: RequestHooks,
) -> Result<Value, Error> {
ocr_with_observer(
request,
options,
context,
hooks,
&mut NoopProviderAttemptObserver,
)
.await
}
#[tracing::instrument(
name = "ocr",
target = "litellm::function_trace",
level = "trace",
skip_all
)]
pub async fn ocr_with_observer<Observer>(
request: OcrRequest<'_>,
options: &RequestOptions,
context: &LiteLlmRequestContext,
hooks: RequestHooks,
observer: &mut Observer,
) -> Result<Value, Error>
where
Observer: ProviderAttemptObserver,
Observer::Error: std::fmt::Display,
{
let PreparedOcrCall {
request,
context,
context: lifecycle,
hooks,
} = prepare_ocr_call(request, options.clone(), context, hooks);
CallLifecycle::default()
.run(context, request, &hooks, |request| {
execute_ocr_provider_call(request, &hooks)
.run(lifecycle.clone(), request, &hooks, |request| {
execute_ocr_provider_call(request, &lifecycle, &hooks, observer)
})
.await
}

View file

@ -25,6 +25,7 @@ pub(crate) struct PreparedOcrRequest {
pub(crate) struct ProviderOcrRequest {
pub(crate) model: String,
pub(crate) custom_llm_provider: String,
pub(crate) config: &'static dyn OcrProviderConfig,
pub(crate) url: String,
pub(crate) body: Value,

View file

@ -54,19 +54,19 @@ pub async fn run(
let request = MessagesRequest {
model: provider_model,
body,
options: RequestOptions {
api_key: (deployment.litellm_params.api_key.as_deref()).map(|value| value.to_string()),
api_base: (deployment.litellm_params.api_base.as_deref())
.map(|value| value.to_string()),
custom_llm_provider: (custom_llm_provider).map(|value| value.to_string()),
extra_headers,
timeout: None,
..Default::default()
},
};
let options = RequestOptions {
api_key: (deployment.litellm_params.api_key.as_deref()).map(|value| value.to_string()),
api_base: (deployment.litellm_params.api_base.as_deref()).map(|value| value.to_string()),
custom_llm_provider: (custom_llm_provider).map(|value| value.to_string()),
extra_headers,
timeout: None,
..Default::default()
};
if request.body.get("stream").and_then(Value::as_bool) == Some(true) {
return messages_stream(
request,
&options,
&LiteLlmRequestContext {
..Default::default()
},
@ -77,6 +77,7 @@ pub async fn run(
let response = messages(
request,
&options,
&LiteLlmRequestContext {
..Default::default()
},

View file

@ -12,13 +12,260 @@ use litellm_ai_gateway::integrations::custom_guardrail::{
use litellm_ai_gateway::integrations::custom_logger::{
CallbackTiming, CallbackValue, CustomLogger, LogFuture, ModelCallDetails,
};
use litellm_ai_gateway::ocr::{OcrRequest, ocr};
use litellm_ai_gateway::ocr::{OcrRequest, ocr, ocr_with_observer};
use litellm_core::error::Error;
use litellm_core::provider_callbacks::{
CallbackDecision, ProviderAttemptObserver, ProviderError, ProviderPostCall, ProviderPreCall,
};
use serde_json::{Map, Value, json};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
struct ProviderObserver {
events: Arc<Mutex<Vec<&'static str>>>,
raw_response: Option<String>,
rejected_callback: Option<&'static str>,
decision: Option<&'static str>,
}
impl ProviderAttemptObserver for ProviderObserver {
type Error = &'static str;
async fn pre_call(&mut self, input: &ProviderPreCall) -> Result<CallbackDecision, Self::Error> {
assert_eq!(input.model, "mistral-ocr-4-1");
assert_eq!(input.call_id, "observer-test");
assert_eq!(
input.request["document"]["document_url"],
"https://example.com/document.pdf"
);
assert!(input.api_base.ends_with("/v1/ocr"));
assert!(
input
.headers
.values()
.any(|value| value == "Bearer test-key")
);
self.events.lock().unwrap().push("pre");
match (self.rejected_callback, self.decision) {
(Some("pre"), _) => Err("observer failure"),
(_, Some("replace_pre")) => Ok(CallbackDecision::Replace {
payload: Value::Object(
input
.request
.iter()
.map(|(key, value)| (key.clone(), value.clone()))
.chain(std::iter::once((
"callback_replaced".to_string(),
json!(true),
)))
.collect(),
),
}),
(_, Some("reject_pre")) => Ok(CallbackDecision::Reject {
message: "callback rejected request".to_string(),
status_code: Some(400),
}),
_ => Ok(CallbackDecision::Unchanged),
}
}
async fn post_call(
&mut self,
input: &ProviderPostCall,
) -> Result<CallbackDecision, Self::Error> {
self.events.lock().unwrap().push("post");
self.raw_response = input.response.as_str().map(str::to_string);
match (self.rejected_callback, self.decision) {
(Some("post"), _) => Err("observer failure"),
(_, Some("replace_post")) => Ok(CallbackDecision::Replace {
payload: json!({"pages":[{"index":0,"markdown":"masked"}]}),
}),
_ => Ok(CallbackDecision::Unchanged),
}
}
async fn error(&mut self, input: &ProviderError) -> Result<(), Self::Error> {
assert!(input.committed);
assert!(!input.message.is_empty());
self.events.lock().unwrap().push("error");
if self.rejected_callback == Some("error") {
Err("observer failure")
} else {
Ok(())
}
}
}
fn observer_request() -> OcrRequest<'static> {
OcrRequest {
model: "mistral/mistral-ocr-4-1",
document: json!({"type":"document_url","document_url":"https://example.com/document.pdf"}),
optional_params: Map::new(),
}
}
fn observer_options(api_base: &str) -> RequestOptions {
RequestOptions {
api_key: Some("test-key".into()),
api_base: Some(api_base.into()),
custom_llm_provider: Some("mistral".into()),
timeout: Some(Duration::from_secs(2)),
..Default::default()
}
}
fn observer_context() -> LiteLlmRequestContext {
LiteLlmRequestContext {
litellm_call_id: Some("observer-test".into()),
..Default::default()
}
}
async fn observer_case(status: u16, body: &'static str, decision: Option<&'static str>) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let url = format!("http://{}/v1", listener.local_addr().unwrap());
let events = Arc::new(Mutex::new(Vec::new()));
let provider_events = Arc::clone(&events);
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let request = read_http_request(&mut socket).await;
assert!(request.starts_with("POST /v1/ocr "));
assert_eq!(
request.contains(r#""callback_replaced":true"#),
decision == Some("replace_pre")
);
provider_events.lock().unwrap().push("http");
let response = format!(
"HTTP/1.1 {status} Test\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{body}",
body.len()
);
socket.write_all(response.as_bytes()).await.unwrap();
});
let mut observer = ProviderObserver {
events: Arc::clone(&events),
raw_response: None,
rejected_callback: None,
decision,
};
let result = ocr_with_observer(
observer_request(),
&observer_options(&url),
&observer_context(),
RequestHooks::default(),
&mut observer,
)
.await;
tokio::time::timeout(Duration::from_secs(2), server)
.await
.unwrap()
.unwrap();
if status != 200 {
assert!(matches!(result, Err(Error::Http { status: actual, .. }) if actual == status));
assert_eq!(*events.lock().unwrap(), ["pre", "http", "error"]);
assert_eq!(observer.raw_response, None);
} else {
assert_eq!(*events.lock().unwrap(), ["pre", "http", "post"]);
assert_eq!(observer.raw_response.as_deref(), Some(body));
if body == "invalid-json" {
assert!(matches!(result, Err(Error::InvalidResponse(_))));
} else {
assert_eq!(
result.unwrap()["pages"][0]["markdown"],
if decision == Some("replace_post") {
"masked"
} else {
"ok"
}
);
}
}
}
#[tokio::test]
async fn provider_observers_surround_http_and_can_replace_request_or_response() {
observer_case(200, r#"{"pages":[{"index":0,"markdown":"ok"}]}"#, None).await;
observer_case(200, "invalid-json", None).await;
observer_case(401, r#"{"error":"rejected"}"#, None).await;
observer_case(
200,
r#"{"pages":[{"index":0,"markdown":"ok"}]}"#,
Some("replace_pre"),
)
.await;
observer_case(
200,
r#"{"pages":[{"index":0,"markdown":"ok"}]}"#,
Some("replace_post"),
)
.await;
}
#[tokio::test]
async fn provider_callback_rejection_stops_before_http() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let url = format!("http://{}/v1", listener.local_addr().unwrap());
let events = Arc::new(Mutex::new(Vec::new()));
let mut observer = ProviderObserver {
events: Arc::clone(&events),
raw_response: None,
rejected_callback: None,
decision: Some("reject_pre"),
};
let result = ocr_with_observer(
observer_request(),
&observer_options(&url),
&observer_context(),
RequestHooks::default(),
&mut observer,
)
.await;
assert!(
matches!(result, Err(Error::InvalidRequest(message)) if message == "callback rejected request")
);
assert_eq!(*events.lock().unwrap(), ["pre"]);
assert!(
tokio::time::timeout(Duration::from_millis(50), listener.accept())
.await
.is_err()
);
}
#[tokio::test]
async fn invalid_ocr_preparation_does_not_call_observers_or_provider() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let url = format!("http://{}/v1", listener.local_addr().unwrap());
let events = Arc::new(Mutex::new(Vec::new()));
let mut observer = ProviderObserver {
events: Arc::clone(&events),
raw_response: None,
rejected_callback: None,
decision: None,
};
let request = OcrRequest {
document: json!(42),
..observer_request()
};
assert!(
ocr_with_observer(
request,
&observer_options(&url),
&observer_context(),
RequestHooks::default(),
&mut observer
)
.await
.is_err()
);
assert!(events.lock().unwrap().is_empty());
assert!(
tokio::time::timeout(Duration::from_millis(50), listener.accept())
.await
.is_err()
);
}
async fn read_http_headers(socket: &mut TcpStream) -> String {
let mut request = Vec::new();
let mut buffer = [0_u8; 1024];

View file

@ -0,0 +1,36 @@
use serde::{Serialize, de::DeserializeOwned};
pub trait Hook {
type Input: Serialize + Send + Sync;
type Output: DeserializeOwned + Send;
}
#[macro_export]
macro_rules! define_hooks {
(
$visibility:vis trait $hooks:ident;
{ $($method:ident: $marker:ident($input:ty) -> $output:ty = $mode:ident;)* }
) => {
$(
$visibility struct $marker;
impl $crate::hook_contracts::Hook for $marker {
type Input = $input;
type Output = $output;
}
)*
$visibility trait $hooks: Send {
type Error: Send;
$(
fn $method<'a>(
&'a mut self,
input: &'a <$marker as $crate::hook_contracts::Hook>::Input,
) -> impl ::std::future::Future<
Output = Result<<$marker as $crate::hook_contracts::Hook>::Output, Self::Error>,
> + Send + 'a;
)*
}
};
}

View file

@ -5,11 +5,13 @@ pub mod chat_completions;
pub mod constants;
pub mod eligibility;
pub mod error;
pub mod hook_contracts;
pub mod http_utils;
pub mod messages;
#[cfg(any(feature = "observability", test))]
pub mod observability;
pub mod ocr;
pub mod provider_callbacks;
pub mod providers;
pub mod realtime;
pub mod responses;

View file

@ -0,0 +1,322 @@
use std::collections::BTreeMap;
use std::time::{SystemTime, UNIX_EPOCH};
use reqwest::{RequestBuilder, StatusCode, header::HeaderMap};
use serde_json::Value;
use crate::Error;
use crate::http_utils::{http_request, truncate_error_body};
use crate::provider_callbacks::{
CallbackDecision, ProviderAttemptObserver, ProviderError, ProviderPostCall, ProviderPreCall,
};
pub struct ProviderHttpResponse {
pub status: StatusCode,
pub headers: HeaderMap,
pub body: String,
}
pub struct ProviderRequest {
pub provider: String,
pub model: String,
pub body: BTreeMap<String, Value>,
pub api_base: String,
pub headers: BTreeMap<String, String>,
}
pub struct ProviderAttemptContext {
pub call_id: String,
pub trace_id: Option<String>,
pub attempt: u32,
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub async fn send_provider_request<Observer>(
request: RequestBuilder,
input: ProviderRequest,
context: ProviderAttemptContext,
observer: &mut Observer,
) -> Result<ProviderHttpResponse, Error>
where
Observer: ProviderAttemptObserver,
Observer::Error: std::fmt::Display,
{
let event = ProviderPreCall {
provider: input.provider,
model: input.model,
call_id: context.call_id,
trace_id: context.trace_id,
attempt: context.attempt,
started_at: epoch_seconds(),
request: input.body,
api_base: input.api_base,
headers: input.headers,
};
let body = match observer.pre_call(&event).await.map_err(callback_error)? {
CallbackDecision::Unchanged => Value::Object(event.request.clone().into_iter().collect()),
CallbackDecision::Replace { payload } => payload,
CallbackDecision::Reject { message, .. } => return Err(Error::InvalidRequest(message)),
};
let response = match http_request(request.json(&body)).await {
Ok(response) => response,
Err(error) => {
let mapped = transport_error(error);
notify_error(observer, &event, &mapped, "provider_request", true).await?;
return Err(mapped);
}
};
let status = response.status();
let headers = response.headers().clone();
let body = match response.text().await {
Ok(body) => body,
Err(error) => {
let mapped = transport_error(error);
notify_error(observer, &event, &mapped, "response_body", true).await?;
return Err(mapped);
}
};
if !status.is_success() {
let error = Error::Http {
status: status.as_u16(),
body: truncate_error_body(&body),
};
notify_error(observer, &event, &error, "provider_response", true).await?;
return Err(error);
}
let post_call = ProviderPostCall {
provider: event.provider.clone(),
model: event.model.clone(),
call_id: event.call_id.clone(),
trace_id: event.trace_id.clone(),
attempt: event.attempt,
started_at: event.started_at,
response: Value::String(body.clone()),
status_code: status.as_u16(),
headers: header_values(&headers),
ended_at: epoch_seconds(),
};
let body = match observer
.post_call(&post_call)
.await
.map_err(callback_error)?
{
CallbackDecision::Unchanged => body,
CallbackDecision::Replace {
payload: Value::String(replacement),
} => replacement,
CallbackDecision::Replace { payload } => {
serde_json::to_string(&payload).map_err(|error| {
Error::InvalidResponse(format!("callback response is invalid: {error}"))
})?
}
CallbackDecision::Reject { message, .. } => return Err(Error::InvalidResponse(message)),
};
Ok(ProviderHttpResponse {
status,
headers,
body,
})
}
async fn notify_error<Observer>(
observer: &mut Observer,
context: &ProviderPreCall,
error: &Error,
stage: &'static str,
committed: bool,
) -> Result<(), Error>
where
Observer: ProviderAttemptObserver,
Observer::Error: std::fmt::Display,
{
let event = ProviderError {
provider: context.provider.clone(),
model: context.model.clone(),
call_id: context.call_id.clone(),
trace_id: context.trace_id.clone(),
attempt: context.attempt,
started_at: context.started_at,
message: error.to_string(),
stage,
committed,
status_code: match error {
Error::Http { status, .. } => Some(*status),
_ => None,
},
ended_at: epoch_seconds(),
};
observer.error(&event).await.map_err(callback_error)
}
fn header_values(headers: &HeaderMap) -> BTreeMap<String, String> {
headers
.iter()
.filter_map(|(name, value)| {
value
.to_str()
.ok()
.map(|value| (name.to_string(), value.to_string()))
})
.collect()
}
fn epoch_seconds() -> f64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|duration| duration.as_secs_f64())
.unwrap_or(0.0)
}
fn callback_error(error: impl std::fmt::Display) -> Error {
Error::InvalidResponse(format!("provider callback failed: {error}"))
}
fn transport_error(error: reqwest::Error) -> Error {
Error::Network(if error.is_timeout() {
"Request timed out".into()
} else {
error.to_string()
})
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
use std::time::Duration;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
struct Observer {
events: Vec<&'static str>,
reject: bool,
}
impl ProviderAttemptObserver for Observer {
type Error = std::convert::Infallible;
async fn pre_call(
&mut self,
event: &ProviderPreCall,
) -> Result<CallbackDecision, Self::Error> {
assert_eq!(event.provider, "test-provider");
assert_eq!(event.model, "test-model");
assert_eq!(event.call_id, "call-1");
assert_eq!(event.trace_id.as_deref(), Some("trace-1"));
assert_eq!(event.attempt, 3);
assert_eq!(event.request["input"], "private");
self.events.push("pre");
Ok(if self.reject {
CallbackDecision::Reject {
message: "blocked".into(),
status_code: Some(400),
}
} else {
CallbackDecision::Replace {
payload: json!({"masked": true}),
}
})
}
async fn post_call(
&mut self,
event: &ProviderPostCall,
) -> Result<CallbackDecision, Self::Error> {
assert_eq!(event.response, json!("raw-response"));
assert_eq!(event.status_code, 200);
assert_eq!(event.attempt, 3);
assert_eq!(event.trace_id.as_deref(), Some("trace-1"));
assert!(event.ended_at >= event.started_at);
self.events.push("post");
Ok(CallbackDecision::Replace {
payload: json!("redacted-response"),
})
}
async fn error(&mut self, event: &ProviderError) -> Result<(), Self::Error> {
assert_eq!(event.status_code, Some(429));
assert_eq!(event.stage, "provider_response");
assert_eq!(event.attempt, 3);
assert_eq!(event.trace_id.as_deref(), Some("trace-1"));
assert!(event.committed);
self.events.push("error");
Ok(())
}
}
#[tokio::test]
async fn shared_attempt_preserves_context_decisions_and_provider_errors() {
let client = reqwest::Client::builder()
.timeout(Duration::from_secs(2))
.build()
.unwrap();
for (status, reject) in [(200, false), (429, false), (200, true)] {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let url = format!(
"http://{}/provider-operation",
listener.local_addr().unwrap()
);
let server = tokio::spawn(async move {
if reject {
assert!(
tokio::time::timeout(Duration::from_millis(50), listener.accept())
.await
.is_err()
);
return;
}
let (mut socket, _) = listener.accept().await.unwrap();
let mut request = Vec::new();
let mut buffer = [0; 1024];
while !request.ends_with(b"{\"masked\":true}") {
let count = socket.read(&mut buffer).await.unwrap();
assert!(count > 0);
request.extend_from_slice(&buffer[..count]);
}
let request = String::from_utf8(request).unwrap();
assert!(request.starts_with("POST /provider-operation "));
assert!(!request.contains("private"));
socket.write_all(format!(
"HTTP/1.1 {status} Test\r\ncontent-length: 12\r\nconnection: close\r\n\r\nraw-response"
).as_bytes()).await.unwrap();
});
let mut observer = Observer {
events: Vec::new(),
reject,
};
let result = send_provider_request(
client.post(&url),
ProviderRequest {
provider: "test-provider".into(),
model: "test-model".into(),
body: BTreeMap::from([("input".into(), json!("private"))]),
api_base: url,
headers: BTreeMap::new(),
},
ProviderAttemptContext {
call_id: "call-1".into(),
trace_id: Some("trace-1".into()),
attempt: 3,
},
&mut observer,
)
.await;
tokio::time::timeout(Duration::from_secs(2), server)
.await
.unwrap()
.unwrap();
if reject {
assert!(
matches!(result, Err(Error::InvalidRequest(message)) if message == "blocked")
);
assert_eq!(observer.events, ["pre"]);
} else if status == 429 {
assert!(matches!(result, Err(Error::Http { status: 429, .. })));
assert_eq!(observer.events, ["pre", "error"]);
} else {
assert_eq!(result.unwrap().body, "redacted-response");
assert_eq!(observer.events, ["pre", "post"]);
}
}
}
}

View file

@ -0,0 +1,207 @@
pub mod handler;
use std::collections::BTreeMap;
use std::convert::Infallible;
use serde::{Deserialize, Serialize};
use serde_json::Value;
#[derive(Debug, Deserialize, PartialEq)]
#[serde(tag = "action", rename_all = "snake_case")]
pub enum CallbackDecision {
Unchanged,
Replace {
payload: Value,
},
Reject {
message: String,
status_code: Option<u16>,
},
}
#[derive(Clone, Serialize)]
pub struct ProviderPreCall {
pub provider: String,
pub model: String,
pub call_id: String,
pub trace_id: Option<String>,
pub attempt: u32,
pub started_at: f64,
pub request: BTreeMap<String, Value>,
pub api_base: String,
pub headers: BTreeMap<String, String>,
}
#[derive(Serialize)]
pub struct ProviderPostCall {
pub provider: String,
pub model: String,
pub call_id: String,
pub trace_id: Option<String>,
pub attempt: u32,
pub started_at: f64,
pub response: Value,
pub status_code: u16,
pub headers: BTreeMap<String, String>,
pub ended_at: f64,
}
#[derive(Serialize)]
pub struct ProviderError {
pub provider: String,
pub model: String,
pub call_id: String,
pub trace_id: Option<String>,
pub attempt: u32,
pub started_at: f64,
pub message: String,
pub stage: &'static str,
pub committed: bool,
pub status_code: Option<u16>,
pub ended_at: f64,
}
#[derive(Serialize)]
pub struct ProviderStreamEvent {
pub provider: String,
pub model: String,
pub call_id: String,
pub trace_id: Option<String>,
pub attempt: u32,
pub started_at: f64,
pub event: Value,
pub sequence: u64,
}
#[derive(Serialize)]
pub struct ProviderStreamClose {
pub provider: String,
pub model: String,
pub call_id: String,
pub trace_id: Option<String>,
pub attempt: u32,
pub started_at: f64,
pub outcome: String,
pub ended_at: f64,
}
#[derive(Serialize)]
pub struct SessionEvent {
pub session_id: String,
pub call_id: String,
pub trace_id: Option<String>,
pub event: Option<Value>,
pub response_id: Option<String>,
pub sequence: Option<u64>,
pub message: Option<String>,
}
#[macro_export]
macro_rules! provider_attempt_observer_catalog {
($consumer:path, $($options:tt)*) => {
$consumer! {
$($options)*
{
pre_call: PreCall($crate::provider_callbacks::ProviderPreCall) -> $crate::provider_callbacks::CallbackDecision = direct;
post_call: PostCall($crate::provider_callbacks::ProviderPostCall) -> $crate::provider_callbacks::CallbackDecision = direct;
error: Error($crate::provider_callbacks::ProviderError) -> () = direct;
}
}
};
}
#[macro_export]
macro_rules! streaming_observer_catalog {
($consumer:path, $($options:tt)*) => {
$consumer! {
$($options)*
{
pre_call: StreamingPreCall($crate::provider_callbacks::ProviderPreCall) -> $crate::provider_callbacks::CallbackDecision = direct;
post_call: StreamingPostCall($crate::provider_callbacks::ProviderPostCall) -> $crate::provider_callbacks::CallbackDecision = direct;
error: StreamingError($crate::provider_callbacks::ProviderError) -> () = direct;
stream_event: StreamEvent($crate::provider_callbacks::ProviderStreamEvent) -> $crate::provider_callbacks::CallbackDecision = direct;
stream_close: StreamClose($crate::provider_callbacks::ProviderStreamClose) -> () = direct;
}
}
};
}
#[macro_export]
macro_rules! session_observer_catalog {
($consumer:path, $($options:tt)*) => {
$consumer! {
$($options)*
{
before_connect: BeforeConnect($crate::provider_callbacks::SessionEvent) -> $crate::provider_callbacks::CallbackDecision = awaitable;
connected: Connected($crate::provider_callbacks::SessionEvent) -> () = awaitable;
before_send: BeforeSend($crate::provider_callbacks::SessionEvent) -> $crate::provider_callbacks::CallbackDecision = awaitable;
after_receive: AfterReceive($crate::provider_callbacks::SessionEvent) -> $crate::provider_callbacks::CallbackDecision = awaitable;
response_complete: ResponseComplete($crate::provider_callbacks::SessionEvent) -> () = awaitable;
response_error: ResponseError($crate::provider_callbacks::SessionEvent) -> () = awaitable;
error: SessionError($crate::provider_callbacks::SessionEvent) -> () = awaitable;
close: Close($crate::provider_callbacks::SessionEvent) -> () = awaitable;
}
}
};
}
provider_attempt_observer_catalog!(crate::define_hooks, pub trait ProviderAttemptObserver;);
streaming_observer_catalog!(crate::define_hooks, pub trait StreamingObserver;);
session_observer_catalog!(crate::define_hooks, pub trait SessionObserver;);
pub struct NoopProviderAttemptObserver;
impl ProviderAttemptObserver for NoopProviderAttemptObserver {
type Error = Infallible;
async fn pre_call(&mut self, _input: &ProviderPreCall) -> Result<CallbackDecision, Infallible> {
Ok(CallbackDecision::Unchanged)
}
async fn post_call(
&mut self,
_input: &ProviderPostCall,
) -> Result<CallbackDecision, Infallible> {
Ok(CallbackDecision::Unchanged)
}
async fn error(&mut self, _input: &ProviderError) -> Result<(), Infallible> {
Ok(())
}
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::CallbackDecision;
#[test]
fn callback_decisions_have_a_tagged_wire_contract() {
assert_eq!(
serde_json::from_value::<CallbackDecision>(json!({"action": "unchanged"})).unwrap(),
CallbackDecision::Unchanged
);
assert_eq!(
serde_json::from_value::<CallbackDecision>(
json!({"action": "replace", "payload": {"masked": true}})
)
.unwrap(),
CallbackDecision::Replace {
payload: json!({"masked": true})
}
);
assert_eq!(
serde_json::from_value::<CallbackDecision>(json!({
"action": "reject",
"message": "blocked",
"status_code": 400
}))
.unwrap(),
CallbackDecision::Reject {
message: "blocked".to_string(),
status_code: Some(400)
}
);
}
}

View file

@ -603,6 +603,12 @@ impl OcrProviderConfig for AzureDocumentIntelligenceOcrConfig {
transform_document_intelligence_response(model, response_json, false)
}
#[tracing::instrument(
name = "transform_ocr_response",
target = "litellm::function_trace",
level = "trace",
skip_all
)]
fn transform_ocr_response_with_params(
&self,
model: &str,

View file

@ -0,0 +1,113 @@
use std::num::NonZeroUsize;
use litellm_core::provider_callbacks::{
CallbackDecision, ProviderAttemptObserver, ProviderError, ProviderPostCall, ProviderPreCall,
};
use litellm_python_interop::callback_runtime::{AsyncContext, CallbackRuntime, SyncContext};
use pyo3::prelude::*;
use crate::constants::OCR_CALLBACK_CAPACITY;
use crate::execution::PythonCallContext;
litellm_core::provider_attempt_observer_catalog!(crate::bind_python_hooks,
pub(crate) struct PythonProviderSession;
trait ProviderAttemptObserver;
);
litellm_core::streaming_observer_catalog!(crate::bind_python_hooks,
pub struct PythonStreamingSession;
trait litellm_core::provider_callbacks::StreamingObserver;
);
litellm_core::session_observer_catalog!(crate::bind_python_hooks,
pub struct PythonSession;
trait litellm_core::provider_callbacks::SessionObserver;
);
#[pyclass(frozen)]
struct PythonCallbackRuntime(CallbackRuntime);
pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
let _streaming_constructor = PythonStreamingSession::<SyncContext>::new;
let _session_constructor = PythonSession::<AsyncContext>::new;
let capacity = NonZeroUsize::new(OCR_CALLBACK_CAPACITY)
.expect("Python callback capacity is a positive constant");
module.add(
"__python_callback_runtime__",
PythonCallbackRuntime(CallbackRuntime::new(module, capacity)?),
)
}
pub(crate) enum PythonProviderObserver {
Disabled,
Sync(PythonProviderSession<SyncContext>),
Async(PythonProviderSession<AsyncContext>),
}
impl PythonProviderObserver {
pub(crate) fn new(
adapter: Option<Py<PyAny>>,
context: PythonCallContext<'_>,
) -> PyResult<Self> {
let Some(adapter) = adapter else {
return Ok(Self::Disabled);
};
let py = context.py;
let module = py.import("litellm.rust_bridge._native")?;
let runtime = module
.getattr("__python_callback_runtime__")?
.extract::<PyRef<'_, PythonCallbackRuntime>>()?
.0
.clone();
if context.asynchronous {
Ok(Self::Async(PythonProviderSession::new(
adapter.bind(py),
runtime.async_context(py)?,
)?))
} else {
Ok(Self::Sync(PythonProviderSession::new(
adapter.bind(py),
runtime.sync_context(py)?,
)?))
}
}
}
pub(crate) fn python_async_session(
adapter: Py<PyAny>,
py: Python<'_>,
) -> PyResult<PythonSession<AsyncContext>> {
let module = py.import("litellm.rust_bridge._native")?;
let runtime = module
.getattr("__python_callback_runtime__")?
.extract::<PyRef<'_, PythonCallbackRuntime>>()?
.0
.clone();
PythonSession::new(adapter.bind(py), runtime.async_context(py)?)
}
impl ProviderAttemptObserver for PythonProviderObserver {
type Error = PyErr;
async fn pre_call(&mut self, input: &ProviderPreCall) -> PyResult<CallbackDecision> {
match self {
Self::Disabled => Ok(CallbackDecision::Unchanged),
Self::Sync(session) => session.pre_call(input).await,
Self::Async(session) => session.pre_call(input).await,
}
}
async fn post_call(&mut self, input: &ProviderPostCall) -> PyResult<CallbackDecision> {
match self {
Self::Disabled => Ok(CallbackDecision::Unchanged),
Self::Sync(session) => session.post_call(input).await,
Self::Async(session) => session.post_call(input).await,
}
}
async fn error(&mut self, input: &ProviderError) -> PyResult<()> {
match self {
Self::Disabled => Ok(()),
Self::Sync(session) => session.error(input).await,
Self::Async(session) => session.error(input).await,
}
}
}

View file

@ -0,0 +1 @@
pub(crate) const OCR_CALLBACK_CAPACITY: usize = 1024;

View file

@ -3,7 +3,6 @@ use std::panic::AssertUnwindSafe;
use std::time::Duration;
use futures_util::FutureExt;
use litellm_core::error::Error;
use litellm_python_interop::{Pythonized, panic_to_pyerr, release_gil};
use pyo3::exceptions::PyRuntimeError;
use pyo3::prelude::*;
@ -11,14 +10,19 @@ use serde::Serialize;
use tokio::runtime::{Handle, Runtime};
use tokio::time::{self, MissedTickBehavior};
pub(crate) fn run_sync<T, F>(
pub(crate) struct PythonCallContext<'py> {
pub(crate) py: Python<'py>,
pub(crate) asynchronous: bool,
}
pub(crate) fn run_sync<T, E>(
py: Python<'_>,
future: F,
map_error: fn(Error) -> PyErr,
future: impl Future<Output = Result<T, E>> + Send + 'static,
map_error: fn(E) -> PyErr,
) -> PyResult<Py<PyAny>>
where
T: Serialize + Send + 'static,
F: Future<Output = Result<T, Error>> + Send + 'static,
E: Send + 'static,
{
run_sync_on(
py,
@ -28,15 +32,15 @@ where
)
}
fn run_sync_on<T, F>(
fn run_sync_on<T, E>(
py: Python<'_>,
runtime: &Runtime,
future: F,
map_error: fn(Error) -> PyErr,
future: impl Future<Output = Result<T, E>> + Send + 'static,
map_error: fn(E) -> PyErr,
) -> PyResult<Py<PyAny>>
where
T: Serialize + Send + 'static,
F: Future<Output = Result<T, Error>> + Send + 'static,
E: Send + 'static,
{
if Handle::try_current().is_ok() {
return Err(PyRuntimeError::new_err(
@ -45,27 +49,27 @@ where
}
let result = release_gil(py, move || runtime.block_on(wait_for_sync_result(future)))?;
let result = map_core_result(result, map_error)?;
let result = map_result(result, map_error)?;
Pythonized(result).into_pyobject(py).map(Bound::unbind)
}
pub(crate) fn run_async<T, F>(
pub(crate) fn run_async<T, E>(
py: Python<'_>,
future: F,
map_error: fn(Error) -> PyErr,
future: impl Future<Output = Result<T, E>> + Send + 'static,
map_error: fn(E) -> PyErr,
) -> PyResult<Bound<'_, PyAny>>
where
T: Serialize + Send + 'static,
F: Future<Output = Result<T, Error>> + Send + 'static,
E: Send + 'static,
{
pyo3_async_runtimes::tokio::future_into_py(py, async move {
let result = catch_future_panic(future).await?;
let result = map_core_result(result, map_error)?;
let result = map_result(result, map_error)?;
Ok(Pythonized(result))
})
}
fn map_core_result<T>(result: Result<T, Error>, map_error: fn(Error) -> PyErr) -> PyResult<T> {
fn map_result<T, E>(result: Result<T, E>, map_error: fn(E) -> PyErr) -> PyResult<T> {
match result {
Ok(value) => Ok(value),
Err(error) => Err(
@ -75,9 +79,9 @@ fn map_core_result<T>(result: Result<T, Error>, map_error: fn(Error) -> PyErr) -
}
}
async fn catch_future_panic<T, F>(future: F) -> PyResult<Result<T, Error>>
async fn catch_future_panic<T, E, F>(future: F) -> PyResult<Result<T, E>>
where
F: Future<Output = Result<T, Error>>,
F: Future<Output = Result<T, E>>,
{
AssertUnwindSafe(future)
.catch_unwind()
@ -85,9 +89,9 @@ where
.map_err(panic_to_pyerr)
}
async fn wait_for_sync_result<T, F>(future: F) -> PyResult<Result<T, Error>>
async fn wait_for_sync_result<T, E, F>(future: F) -> PyResult<Result<T, E>>
where
F: Future<Output = Result<T, Error>>,
F: Future<Output = Result<T, E>>,
{
let future = catch_future_panic(future);
tokio::pin!(future);
@ -108,12 +112,14 @@ where
mod tests {
use std::ffi::CString;
use std::future::poll_fn;
use std::process::Command;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, mpsc};
use std::task::Poll;
use std::thread;
use std::time::Instant;
use litellm_core::error::Error;
use pyo3::panic::PanicException;
use pyo3::types::{PyDict, PyModule};
use serde::Serializer;
@ -385,6 +391,31 @@ asyncio.run(exercise())
#[test]
fn async_result_delivery_does_not_stall_tokio_workers() {
const CHILD_PROCESS: &str = "LITELLM_ASYNC_RESULT_DELIVERY_TEST_CHILD";
if std::env::var_os(CHILD_PROCESS).is_none() {
let output =
Command::new(std::env::current_exe().expect("test executable should exist"))
.arg("--exact")
.arg(
thread::current()
.name()
.expect("test thread should be named"),
)
.arg("--nocapture")
.env(CHILD_PROCESS, "1")
.env("TOKIO_WORKER_THREADS", "1")
.output()
.expect("isolated result-delivery test should start");
assert!(
output.status.success(),
"isolated result-delivery test failed:\n{}\n{}",
String::from_utf8_lossy(&output.stdout),
String::from_utf8_lossy(&output.stderr),
);
return;
}
Python::initialize();
ASYNC_PROBE_COMPLETED.store(0, Ordering::SeqCst);
Python::attach(|py| {

View file

@ -1,12 +1,21 @@
mod callback_bindings;
#[cfg(test)]
#[path = "../tests/callbacks/mod.rs"]
mod callback_tests;
mod constants;
mod diagnostics;
mod errors;
mod execution;
#[cfg(feature = "trace-parity")]
mod function_trace;
mod marshal;
mod python_hook_bindings;
mod routes;
use std::sync::atomic::{AtomicU64, Ordering};
use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection as RustResponsesWebSocketConnection;
use litellm_core::provider_callbacks::{CallbackDecision, SessionEvent, SessionObserver};
use litellm_core::responses::types::ResponsesWebSocketRequest;
use pyo3::prelude::*;
use pyo3::types::PyAny;
@ -14,6 +23,8 @@ use pyo3::types::PyAny;
use crate::errors::core_error_to_pyerr;
use crate::marshal::{NativeRequestContext, NativeRequestOptions};
static NEXT_WEBSOCKET_SESSION_ID: AtomicU64 = AtomicU64::new(1);
#[derive(FromPyObject)]
struct WebSocketConnectRequest {
url: String,
@ -27,13 +38,14 @@ struct ResponsesWebSocketConnection {
#[pymethods]
impl ResponsesWebSocketConnection {
#[classmethod]
#[pyo3(signature = (request, *, options, context))]
#[pyo3(signature = (request, *, options, context, callback_adapter=None))]
fn connect<'py>(
_cls: &Bound<'py, pyo3::types::PyType>,
py: Python<'py>,
request: WebSocketConnectRequest,
options: NativeRequestOptions,
context: NativeRequestContext,
callback_adapter: Option<Py<PyAny>>,
) -> PyResult<Bound<'py, PyAny>> {
let provider_supported = litellm_core::responses::websocket::native_websocket_supported(
options.provider("openai"),
@ -43,11 +55,54 @@ impl ResponsesWebSocketConnection {
return Err(crate::errors::RustBridgeDeclined::new_err(reason));
}
let options: litellm_core::request_options::RequestOptions = options.into();
let call_id = context.litellm_call_id.clone().unwrap_or_default();
let session_id = format!(
"responses-websocket-{}",
NEXT_WEBSOCKET_SESSION_ID.fetch_add(1, Ordering::Relaxed)
);
let mut observer = callback_adapter
.map(|adapter| crate::callback_bindings::python_async_session(adapter, py))
.transpose()?;
let request = ResponsesWebSocketRequest { url: request.url };
pyo3_async_runtimes::tokio::future_into_py(py, async move {
let inner = RustResponsesWebSocketConnection::connect(request, &options, &context)
if let Some(observer) = observer.as_mut() {
let decision = observer
.before_connect(&session_event(&session_id, &call_id, None))
.await?;
match decision {
CallbackDecision::Unchanged => {}
CallbackDecision::Replace { .. } => {
return Err(pyo3::exceptions::PyValueError::new_err(
"before_connect cannot replace WebSocket setup",
));
}
CallbackDecision::Reject { message, .. } => {
return Err(pyo3::exceptions::PyValueError::new_err(message));
}
}
}
let inner = match RustResponsesWebSocketConnection::connect(request, &options, &context)
.await
.map_err(core_error_to_pyerr)?;
{
Ok(inner) => inner,
Err(error) => {
if let Some(observer) = observer.as_mut() {
observer
.error(&session_event(
&session_id,
&call_id,
Some(error.to_string()),
))
.await?;
}
return Err(core_error_to_pyerr(error));
}
};
if let Some(observer) = observer.as_mut() {
observer
.connected(&session_event(&session_id, &call_id, None))
.await?;
}
Ok(ResponsesWebSocketConnection { inner })
})
}
@ -74,6 +129,18 @@ impl ResponsesWebSocketConnection {
}
}
fn session_event(session_id: &str, call_id: &str, message: Option<String>) -> SessionEvent {
SessionEvent {
session_id: session_id.to_string(),
call_id: call_id.to_string(),
trace_id: None,
event: None,
response_id: None,
sequence: None,
message,
}
}
#[pyfunction]
#[pyo3(signature = (_model, custom_llm_provider, *, context))]
fn responses_websocket_decline(
@ -95,6 +162,8 @@ mod _native {
#[pymodule_init]
fn init(module: &Bound<'_, PyModule>) -> PyResult<()> {
super::errors::register(module)?;
litellm_python_interop::callback_runtime::register(module)?;
super::callback_bindings::register(module)?;
super::routes::register(module)?;
module.add_class::<super::ResponsesWebSocketConnection>()?;
module.add_function(wrap_pyfunction!(
@ -229,6 +298,16 @@ mod tests {
let code = CString::new(
r#"
import asyncio
import sys
import types
litellm_module = types.ModuleType('litellm')
rust_bridge_module = types.ModuleType('litellm.rust_bridge')
litellm_module.rust_bridge = rust_bridge_module
rust_bridge_module._native = native
sys.modules['litellm'] = litellm_module
sys.modules['litellm.rust_bridge'] = rust_bridge_module
sys.modules['litellm.rust_bridge._native'] = native
async def exercise():
for request, request_options, request_context, field in (
@ -248,7 +327,56 @@ async def exercise():
else:
raise AssertionError('invalid WebSocket input reached execution')
connection = await native.ResponsesWebSocketConnection.connect(Request(url=url), options=options, context=context)
events = []
class Adapter:
async def before_connect(self, event):
events.append(('before_connect', event))
return {'action': 'unchanged'}
async def connected(self, event):
events.append(('connected', event))
async def before_send(self, event):
return {'action': 'unchanged'}
async def after_receive(self, event):
return {'action': 'unchanged'}
async def response_complete(self, event):
pass
async def response_error(self, event):
pass
async def error(self, event):
events.append(('error', event))
async def close(self, event):
pass
class RejectingAdapter(Adapter):
async def before_connect(self, event):
return {'action': 'reject', 'message': 'blocked', 'status_code': 400}
try:
await native.ResponsesWebSocketConnection.connect(
Request(url=url), options=options, context=context, callback_adapter=RejectingAdapter()
)
except ValueError as error:
assert str(error) == 'blocked'
else:
raise AssertionError('rejected WebSocket setup reached execution')
connection = await native.ResponsesWebSocketConnection.connect(
Request(url=url),
options=options,
context=replace(context, litellm_call_id='call-1'),
callback_adapter=Adapter(),
)
assert [name for name, _ in events] == ['before_connect', 'connected']
assert events[0][1]['call_id'] == 'call-1'
assert events[0][1]['session_id'] == events[1][1]['session_id']
assert type(connection) is native.ResponsesWebSocketConnection
await connection.send_text("from-python")
assert await connection.recv_text() == "from-server"
@ -256,6 +384,9 @@ async def exercise():
assert await connection.recv_text() is None
asyncio.run(asyncio.wait_for(exercise(), timeout=5))
sys.modules.pop('litellm.rust_bridge._native')
sys.modules.pop('litellm.rust_bridge')
sys.modules.pop('litellm')
"#,
)
.expect("Python source should not contain null bytes");

View file

@ -0,0 +1,64 @@
#[macro_export]
macro_rules! callback_return_mode {
(direct) => {
::litellm_python_interop::callback_runtime::Direct
};
(awaitable) => {
::litellm_python_interop::callback_runtime::Awaitable
};
}
#[macro_export]
macro_rules! bind_python_hooks {
(
$visibility:vis struct $session:ident;
trait $hooks:path;
{ $($method:ident: $marker:ident($input:ty) -> $output:ty = $mode:ident;)* }
) => {
$visibility struct $session<C> {
context: C,
$(
$method: ::litellm_python_interop::callback_runtime::Callback<
$input, $output, $crate::callback_return_mode!($mode),
>,
)*
}
impl<C> $session<C>
where
$(C: ::litellm_python_interop::callback_runtime::CallbackContext<
$crate::callback_return_mode!($mode),
>,)*
{
$visibility fn new(
adapter: &::pyo3::Bound<'_, ::pyo3::PyAny>,
context: C,
) -> ::pyo3::PyResult<Self> {
use ::pyo3::types::PyAnyMethods as _;
Ok(Self {
context,
$(
$method: ::litellm_python_interop::callback_runtime::Callback::new(
adapter.getattr(stringify!($method))?,
)?,
)*
})
}
}
impl<C> $hooks for $session<C>
where
$(C: ::litellm_python_interop::callback_runtime::CallbackContext<
$crate::callback_return_mode!($mode),
>,)*
{
type Error = ::pyo3::PyErr;
$(
async fn $method(&mut self, input: &$input) -> ::pyo3::PyResult<$output> {
self.$method.call(&mut self.context, input).await
}
)*
}
};
}

View file

@ -21,6 +21,8 @@ fn prepare_transcription(
input: AudioTranscriptionInputs,
options: NativeRequestOptions,
context: NativeRequestContext,
_callback_adapter: Option<Py<PyAny>>,
_python_context: crate::execution::PythonCallContext<'_>,
) -> PyResult<impl Future<Output = Result<Value, Error>> + Send + 'static> {
let provider_supported = litellm_core::audio_transcription::transcription_provider_supported(
options.provider("bedrock"),

View file

@ -23,6 +23,8 @@ fn prepare_chat_completions(
input: ChatCompletionsInputs,
options: NativeRequestOptions,
context: NativeRequestContext,
_callback_adapter: Option<Py<PyAny>>,
_python_context: crate::execution::PythonCallContext<'_>,
) -> PyResult<impl Future<Output = Result<ChatCompletionsResponse, Error>> + Send + 'static> {
let context: LiteLlmRequestContext = context.into();
let messages = required_value("messages", input.messages, Value::is_array, "list")?;

View file

@ -12,26 +12,28 @@ macro_rules! bridge_route {
$(, extra = [$($extra:ident),* $(,)?])? $(,)?
) => {
#[pyfunction]
#[pyo3(signature = (request, *, options, context))]
#[pyo3(signature = (request, *, options, context, callback_adapter=None))]
fn $sync_name(
py: pyo3::Python<'_>,
request: $inputs,
options: $crate::marshal::NativeRequestOptions,
context: $crate::marshal::NativeRequestContext,
callback_adapter: Option<pyo3::Py<pyo3::PyAny>>,
) -> pyo3::PyResult<pyo3::Py<pyo3::PyAny>> {
let future = $prepare(request, options, context)?;
let future = $prepare(request, options, context, callback_adapter, $crate::execution::PythonCallContext {py, asynchronous: false})?;
$crate::execution::run_sync(py, future, $map_error)
}
#[pyfunction]
#[pyo3(signature = (request, *, options, context))]
#[pyo3(signature = (request, *, options, context, callback_adapter=None))]
fn $async_name(
py: pyo3::Python<'_>,
request: $inputs,
options: $crate::marshal::NativeRequestOptions,
context: $crate::marshal::NativeRequestContext,
callback_adapter: Option<pyo3::Py<pyo3::PyAny>>,
) -> pyo3::PyResult<pyo3::Bound<'_, pyo3::PyAny>> {
let future = $prepare(request, options, context)?;
let future = $prepare(request, options, context, callback_adapter, $crate::execution::PythonCallContext {py, asynchronous: true})?;
$crate::execution::run_async(py, future, $map_error)
}
@ -48,26 +50,28 @@ macro_rules! bridge_route {
use super::{$inputs, $map_error, $prepare};
#[pyfunction]
#[pyo3(signature = (request, *, options, context))]
#[pyo3(signature = (request, *, options, context, callback_adapter=None))]
fn $sync_name(
py: pyo3::Python<'_>,
request: $inputs,
options: $crate::marshal::NativeRequestOptions,
context: $crate::marshal::NativeRequestContext,
callback_adapter: Option<pyo3::Py<pyo3::PyAny>>,
) -> pyo3::PyResult<pyo3::Py<pyo3::PyAny>> {
let future = $prepare(request, options, context)?;
let future = $prepare(request, options, context, callback_adapter, $crate::execution::PythonCallContext {py, asynchronous: false})?;
$crate::execution::run_sync(py, $crate::function_trace::capture(future), $map_error)
}
#[pyfunction]
#[pyo3(signature = (request, *, options, context))]
#[pyo3(signature = (request, *, options, context, callback_adapter=None))]
fn $async_name(
py: pyo3::Python<'_>,
request: $inputs,
options: $crate::marshal::NativeRequestOptions,
context: $crate::marshal::NativeRequestContext,
callback_adapter: Option<pyo3::Py<pyo3::PyAny>>,
) -> pyo3::PyResult<pyo3::Bound<'_, pyo3::PyAny>> {
let future = $prepare(request, options, context)?;
let future = $prepare(request, options, context, callback_adapter, $crate::execution::PythonCallContext {py, asynchronous: true})?;
$crate::execution::run_async(py, $crate::function_trace::capture(future), $map_error)
}
@ -155,6 +159,8 @@ mod tests {
inputs: EchoInputs,
_options: crate::marshal::NativeRequestOptions,
_context: crate::marshal::NativeRequestContext,
_callback_adapter: Option<Py<PyAny>>,
_python_context: crate::execution::PythonCallContext<'_>,
) -> PyResult<impl Future<Output = Result<String, Error>> + Send + 'static> {
FUTURE_DROPPED.store(false, Ordering::SeqCst);
let drop_guard = (inputs.value == "pending").then_some(DropGuard);
@ -195,17 +201,25 @@ mod tests {
let module = PyModule::new(py, "routes").expect("module should be created");
crate::routes::register(&module).expect("routes should register");
let routes = [
("ocr", "aocr", "(request, *, options, context)"),
(
"ocr",
"aocr",
"(request, *, options, context, callback_adapter=None)",
),
(
"transcription",
"atranscription",
"(request, *, options, context)",
"(request, *, options, context, callback_adapter=None)",
),
(
"messages",
"amessages",
"(request, *, options, context, callback_adapter=None)",
),
("messages", "amessages", "(request, *, options, context)"),
(
"chat_completions",
"achat_completions",
"(request, *, options, context)",
"(request, *, options, context, callback_adapter=None)",
),
];

View file

@ -19,6 +19,8 @@ fn prepare_messages(
input: MessagesInputs,
options: NativeRequestOptions,
context: NativeRequestContext,
_callback_adapter: Option<Py<PyAny>>,
_python_context: crate::execution::PythonCallContext<'_>,
) -> PyResult<impl Future<Output = Result<AnthropicMessagesResponse, Error>> + Send + 'static> {
let provider_supported =
litellm_core::messages::messages_provider_supported(options.provider("anthropic"));

View file

@ -1,8 +1,9 @@
use crate::callback_bindings::PythonProviderObserver;
use crate::errors::ocr_error_to_pyerr;
use crate::marshal::{NativeRequestContext, NativeRequestOptions};
use litellm_ai_gateway::integrations::types::RequestHooks;
use litellm_ai_gateway::io::ocr::OcrRequest;
use litellm_ai_gateway::io::ocr::ocr as run_route;
use litellm_ai_gateway::io::ocr::ocr_with_observer as run_route;
use litellm_core::Error;
use litellm_core::request_context::LiteLlmRequestContext;
use pyo3::prelude::*;
@ -22,6 +23,8 @@ fn prepare_ocr(
input: OcrInputs,
options: NativeRequestOptions,
context: NativeRequestContext,
callback_adapter: Option<Py<PyAny>>,
python_context: crate::execution::PythonCallContext<'_>,
) -> PyResult<impl Future<Output = Result<Value, Error>> + Send + 'static> {
let context: LiteLlmRequestContext = context.into();
let provider_supported = litellm_ai_gateway::io::ocr::ocr_provider_supported(
@ -33,6 +36,7 @@ fn prepare_ocr(
return Err(crate::errors::RustBridgeDeclined::new_err(reason));
}
let document = input.document;
let mut observer = PythonProviderObserver::new(callback_adapter, python_context)?;
Ok(async move {
run_route(
OcrRequest {
@ -46,6 +50,7 @@ fn prepare_ocr(
callbacks: Vec::new(),
guardrails: Vec::new(),
},
&mut observer,
)
.await
})

View file

@ -0,0 +1,432 @@
use std::collections::BTreeMap;
use std::ffi::{CStr, CString};
use std::num::NonZeroUsize;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use litellm_core::provider_callbacks::{
ProviderError, ProviderPostCall, ProviderPreCall, ProviderStreamClose, ProviderStreamEvent,
SessionEvent, SessionObserver, StreamingObserver,
};
use litellm_python_interop::callback_runtime::CallbackRuntime;
use pyo3::exceptions::PyRuntimeError;
use pyo3::prelude::*;
const PYTHON_TEST_SOURCE: &CStr = pyo3::ffi::c_str!(include_str!("test_callbacks.py"));
macro_rules! provider_catalog {
($consumer:path, $($options:tt)*) => {
$consumer! {
$($options)*
{
pre_request: PreRequest(crate::callback_tests::domain::Request)
-> crate::callback_tests::domain::Request = awaitable;
pre_api_call: PreApiCall(crate::callback_tests::domain::BeforeSend)
-> () = direct;
post_response: PostResponse(crate::callback_tests::domain::Response)
-> crate::callback_tests::domain::Response = awaitable;
failure: Failure(crate::callback_tests::domain::ProviderFailure)
-> () = direct;
}
}
};
}
macro_rules! direct_catalog {
($consumer:path, $($options:tt)*) => {
$consumer! {
$($options)*
{
transform: Transform(crate::callback_tests::domain::Request)
-> crate::callback_tests::domain::Request = direct;
}
}
};
}
mod domain {
use serde::{Deserialize, Serialize};
use super::*;
#[derive(Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct Request {
pub text: String,
}
#[derive(Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct Response {
pub text: String,
}
#[derive(Serialize)]
pub struct BeforeSend {
pub body: Request,
}
#[derive(Serialize)]
pub struct ProviderFailure {
pub message: String,
}
provider_catalog!(litellm_core::define_hooks, pub trait ProviderHooks;);
direct_catalog!(litellm_core::define_hooks, pub trait DirectHooks;);
pub enum CallError<E> {
Hook(E),
Provider {
error: ProviderFailure,
observer_error: Option<E>,
},
}
struct PreparedCall(BeforeSend);
struct ReadyCall(Request);
impl PreparedCall {
async fn finish_hooks<H: ProviderHooks>(
self,
hooks: &mut H,
) -> Result<ReadyCall, H::Error> {
hooks.pre_api_call(&self.0).await?;
Ok(ReadyCall(self.0.body))
}
}
impl ReadyCall {
fn send(self, calls: &AtomicUsize, fail: bool) -> Result<Response, ProviderFailure> {
calls.fetch_add(1, Ordering::SeqCst);
if fail {
return Err(ProviderFailure {
message: "provider failed".into(),
});
}
Ok(Response {
text: format!("processed:{}", self.0.text),
})
}
}
pub async fn execute<H: ProviderHooks>(
hooks: &mut H,
calls: &AtomicUsize,
fail: bool,
) -> Result<Response, CallError<H::Error>> {
let request = hooks
.pre_request(&Request {
text: "input".into(),
})
.await
.map_err(CallError::Hook)?;
let ready = PreparedCall(BeforeSend { body: request })
.finish_hooks(hooks)
.await
.map_err(CallError::Hook)?;
let response = match ready.send(calls, fail) {
Ok(response) => response,
Err(error) => {
let observer_error = hooks.failure(&error).await.err();
return Err(CallError::Provider {
error,
observer_error,
});
}
};
hooks
.post_response(&response)
.await
.map_err(CallError::Hook)
}
}
use domain::{DirectHooks, ProviderHooks};
provider_catalog!(crate::bind_python_hooks,
struct PythonProviderSession;
trait domain::ProviderHooks;
);
direct_catalog!(crate::bind_python_hooks,
struct PythonDirectSession;
trait domain::DirectHooks;
);
fn map_error(error: domain::CallError<PyErr>) -> PyErr {
match error {
domain::CallError::Hook(error) => error,
domain::CallError::Provider {
error,
observer_error,
} => Python::attach(|py| {
let exception = PyRuntimeError::new_err(error.message);
if let Some(observer_error) = observer_error {
exception
.value(py)
.setattr("observer_error", observer_error.value(py))
.expect("exception should retain its observer error");
}
exception
}),
}
}
#[pyclass(frozen)]
struct Harness {
runtime: CallbackRuntime,
calls: Arc<AtomicUsize>,
}
#[pymethods]
impl Harness {
#[pyo3(signature = (adapter, fail=false))]
fn execute<'py>(
&self,
py: Python<'py>,
adapter: &Bound<'py, PyAny>,
fail: bool,
) -> PyResult<Bound<'py, PyAny>> {
let mut session = PythonProviderSession::new(adapter, self.runtime.async_context(py)?)?;
let calls = Arc::clone(&self.calls);
crate::execution::run_async(
py,
async move { domain::execute(&mut session, &calls, fail).await },
map_error,
)
}
fn sync(&self, py: Python<'_>, adapter: &Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
let mut session = PythonDirectSession::new(adapter, self.runtime.sync_context(py)?)?;
crate::execution::run_sync(
py,
async move {
session
.transform(&domain::Request {
text: "input".into(),
})
.await
},
std::convert::identity,
)
}
fn interrupt<'py>(
&self,
py: Python<'py>,
adapter: &Bound<'py, PyAny>,
stop: Bound<'py, PyAny>,
) -> PyResult<Bound<'py, PyAny>> {
let mut session = PythonProviderSession::new(adapter, self.runtime.async_context(py)?)?;
let stop = pyo3_async_runtimes::into_future_with_locals(
&pyo3_async_runtimes::tokio::get_current_locals(py)?,
stop,
)?;
crate::execution::run_async(
py,
async move {
let request = domain::Request {
text: "input".into(),
};
tokio::select! {
result = session.pre_request(&request) => { result?; },
result = stop => { result?; },
}
session.pre_request(&request).await
},
std::convert::identity,
)
}
fn calls(&self) -> usize {
self.calls.load(Ordering::SeqCst)
}
fn streaming(&self, py: Python<'_>, adapter: &Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
let mut session = crate::callback_bindings::PythonStreamingSession::new(
adapter,
self.runtime.sync_context(py)?,
)?;
crate::execution::run_sync(
py,
async move {
let pre = provider_pre_call();
assert!(matches!(
session.pre_call(&pre).await?,
litellm_core::provider_callbacks::CallbackDecision::Unchanged
));
assert!(matches!(
session.post_call(&provider_post_call()).await?,
litellm_core::provider_callbacks::CallbackDecision::Unchanged
));
assert!(matches!(
session.stream_event(&provider_stream_event()).await?,
litellm_core::provider_callbacks::CallbackDecision::Unchanged
));
session.stream_close(&provider_stream_close()).await?;
session.error(&provider_error()).await
},
std::convert::identity,
)
}
fn session<'py>(
&self,
py: Python<'py>,
adapter: &Bound<'py, PyAny>,
) -> PyResult<Bound<'py, PyAny>> {
let mut session =
crate::callback_bindings::PythonSession::new(adapter, self.runtime.async_context(py)?)?;
crate::execution::run_async(
py,
async move {
let event = session_event();
assert!(matches!(
session.before_connect(&event).await?,
litellm_core::provider_callbacks::CallbackDecision::Unchanged
));
session.connected(&event).await?;
assert!(matches!(
session.before_send(&event).await?,
litellm_core::provider_callbacks::CallbackDecision::Unchanged
));
assert!(matches!(
session.after_receive(&event).await?,
litellm_core::provider_callbacks::CallbackDecision::Unchanged
));
session.response_complete(&event).await?;
session.response_error(&event).await?;
session.error(&event).await?;
session.close(&event).await
},
std::convert::identity,
)
}
}
fn provider_pre_call() -> ProviderPreCall {
ProviderPreCall {
provider: "test".to_string(),
model: "model".to_string(),
call_id: "call".to_string(),
trace_id: Some("trace".to_string()),
attempt: 2,
started_at: 1.0,
request: BTreeMap::new(),
api_base: "https://provider.test".to_string(),
headers: BTreeMap::new(),
}
}
fn provider_post_call() -> ProviderPostCall {
ProviderPostCall {
provider: "test".to_string(),
model: "model".to_string(),
call_id: "call".to_string(),
trace_id: Some("trace".to_string()),
attempt: 2,
started_at: 1.0,
response: serde_json::json!({}),
status_code: 200,
headers: BTreeMap::new(),
ended_at: 2.0,
}
}
fn provider_stream_event() -> ProviderStreamEvent {
ProviderStreamEvent {
provider: "test".to_string(),
model: "model".to_string(),
call_id: "call".to_string(),
trace_id: Some("trace".to_string()),
attempt: 2,
started_at: 1.0,
event: serde_json::json!({"type": "delta"}),
sequence: 1,
}
}
fn provider_stream_close() -> ProviderStreamClose {
ProviderStreamClose {
provider: "test".to_string(),
model: "model".to_string(),
call_id: "call".to_string(),
trace_id: Some("trace".to_string()),
attempt: 2,
started_at: 1.0,
outcome: "completed".to_string(),
ended_at: 2.0,
}
}
fn provider_error() -> ProviderError {
ProviderError {
provider: "test".to_string(),
model: "model".to_string(),
call_id: "call".to_string(),
trace_id: Some("trace".to_string()),
attempt: 2,
started_at: 1.0,
message: "retrying".to_string(),
stage: "provider_response",
committed: true,
status_code: Some(429),
ended_at: 2.0,
}
}
fn session_event() -> SessionEvent {
SessionEvent {
session_id: "session".to_string(),
call_id: "call".to_string(),
trace_id: Some("trace".to_string()),
event: Some(serde_json::json!({"type": "response.create"})),
response_id: Some("response".to_string()),
sequence: Some(1),
message: None,
}
}
fn run_python_test(name: &str, capacity: usize) {
Python::initialize();
Python::attach(|py| {
let module = PyModule::new(py, "callback_test").expect("test module should load");
litellm_python_interop::callback_runtime::register(&module).expect("shim should register");
let runtime = CallbackRuntime::new(&module, NonZeroUsize::new(capacity).unwrap())
.expect("runtime should initialize");
let harness = Harness {
runtime,
calls: Arc::new(AtomicUsize::new(0)),
};
let tests = PyModule::from_code(
py,
PYTHON_TEST_SOURCE,
c"callback_tests.py",
&CString::new(format!("callback_tests_{name}")).unwrap(),
)
.expect("test definitions should load");
tests
.call_method1("run", (name, harness))
.expect("callback contract should hold");
});
}
macro_rules! python_tests {
($($name:ident: $capacity:literal,)*) => {
$(
#[test]
fn $name() { run_python_test(stringify!($name), $capacity); }
)*
};
}
python_tests! {
transforms_and_context: 64,
callback_errors: 64,
retained_callbacks: 64,
registration_and_return_contracts: 64,
provider_failure: 64,
cancellation_and_admission: 1,
interrupted_session: 1,
synchronous_callbacks: 1,
callback_catalogs: 64,
}

View file

@ -0,0 +1,422 @@
import asyncio
import contextvars
import gc
import inspect
import threading
import weakref
from collections.abc import Awaitable, Callable
from dataclasses import dataclass, replace
from typing import Final, Protocol
Payload = dict[str, str]
BeforeSend = dict[str, Payload]
class Harness(Protocol):
def execute(self, adapter: object, fail: bool = False) -> Awaitable[Payload]: ...
def sync(self, adapter: object) -> Payload: ...
def interrupt(self, adapter: object, stop: Awaitable[object]) -> Awaitable[Payload]: ...
def calls(self) -> int: ...
def streaming(self, adapter: object) -> None: ...
def session(self, adapter: object) -> Awaitable[None]: ...
REQUEST_CONTEXT: Final = contextvars.ContextVar("request_context", default="default")
@dataclass(frozen=True, slots=True)
class Adapter:
pre_request: Callable[[Payload], object]
pre_api_call: Callable[[BeforeSend], object]
post_response: Callable[[Payload], object]
failure: Callable[[Payload], object]
class Recorder:
def __init__(self) -> None:
self.context = REQUEST_CONTEXT.get()
self.loop = asyncio.get_running_loop()
self.thread = threading.get_ident()
self.events: tuple[str, ...] = ()
def record(self, event: str) -> None:
assert REQUEST_CONTEXT.get() == self.context
assert asyncio.get_running_loop() is self.loop
assert threading.get_ident() == self.thread
self.events += (event,)
async def pre_request(self, request: Payload) -> Payload:
self.record("pre_request")
await asyncio.sleep(0)
return {"text": request["text"].upper()}
def pre_api_call(self, event: BeforeSend) -> None:
self.record("pre_api_call")
event["body"]["text"] = "observer mutation"
async def post_response(self, response: Payload) -> Payload:
self.record("post_response")
await asyncio.sleep(0)
return {"text": response["text"] + "!"}
def failure(self, error: Payload) -> None:
self.record("failure")
assert error == {"message": "provider failed"}
def adapter(self) -> Adapter:
return Adapter(self.pre_request, self.pre_api_call, self.post_response, self.failure)
async def transforms_and_context(harness: Harness) -> None:
async def run_one(index: int) -> None:
token: Final = REQUEST_CONTEXT.set(f"request-{index}")
try:
recorder: Final = Recorder()
result: Final = await harness.execute(recorder.adapter())
assert result == {"text": "processed:INPUT!"}
assert recorder.events == ("pre_request", "pre_api_call", "post_response")
finally:
REQUEST_CONTEXT.reset(token)
await asyncio.gather(*(run_one(index) for index in range(32)))
assert harness.calls() == 32
class CallbackAbort(BaseException):
pass
async def callback_errors(harness: Harness) -> None:
async def check(original: BaseException) -> None:
async def reject(request: Payload) -> Payload:
raise original
try:
await harness.execute(replace(Recorder().adapter(), pre_request=reject))
except BaseException as error:
assert error is original
assert error.__traceback__ is not None
frames: Final = inspect.getinnerframes(error.__traceback__)
assert "reject" in tuple(frame.function for frame in frames)
else:
raise AssertionError("callback failure was lost")
await check(ValueError("rejected"))
await check(CallbackAbort("abort"))
await check(asyncio.CancelledError())
assert harness.calls() == 0
async def retained_callbacks(harness: Harness) -> None:
def start() -> tuple[Awaitable[Payload], weakref.ReferenceType[Recorder]]:
recorder: Final = Recorder()
return harness.execute(recorder.adapter()), weakref.ref(recorder)
future, reference = start()
gc.collect()
assert reference() is not None
assert await future == {"text": "processed:INPUT!"}
gc.collect()
assert reference() is None
async def registration_and_return_contracts(harness: Harness) -> None:
adapter: Final = Recorder().adapter()
for invalid in (object(), replace(adapter, pre_api_call=None)):
try:
harness.execute(invalid)
except (AttributeError, TypeError):
pass
else:
raise AssertionError("invalid adapter accepted")
async def wrong_shape(request: Payload) -> int:
return 42
async def extra_field(request: Payload) -> Payload:
return {"text": "input", "unknown": "rejected"}
def not_awaitable(request: Payload) -> Payload:
return request
async def unfinished() -> None:
pass
coroutine: Final = unfinished()
def wrong_direct(event: BeforeSend) -> object:
return coroutine
def wrong_observer(event: BeforeSend) -> int:
return 42
for invalid_adapter, expected in (
(replace(adapter, pre_request=wrong_shape), "typed contract"),
(replace(adapter, pre_request=extra_field), "typed contract"),
(replace(adapter, pre_request=not_awaitable), "non-awaitable"),
(replace(adapter, pre_api_call=wrong_direct), "direct hook returned an awaitable"),
(replace(adapter, pre_api_call=wrong_observer), "typed contract"),
):
try:
await harness.execute(invalid_adapter)
except TypeError as error:
assert expected in str(error)
else:
raise AssertionError("invalid callback result accepted")
assert inspect.getcoroutinestate(coroutine) == inspect.CORO_CLOSED
assert harness.calls() == 0
def future_value(request: Payload) -> asyncio.Future[Payload]:
result: Final[asyncio.Future[Payload]] = asyncio.get_running_loop().create_future()
result.set_result(request)
return result
assert await harness.execute(replace(adapter, pre_request=future_value)) == {"text": "processed:input!"}
assert harness.calls() == 1
async def provider_failure(harness: Harness) -> None:
recorder: Final = Recorder()
original: Final = LookupError("observer failed")
def fail_observer(error: Payload) -> None:
recorder.failure(error)
raise original
try:
await harness.execute(replace(recorder.adapter(), failure=fail_observer), fail=True)
except RuntimeError as error:
assert str(error) == "provider failed"
assert getattr(error, "observer_error") is original
else:
raise AssertionError("provider error was lost")
assert recorder.events == ("pre_request", "pre_api_call", "failure")
assert harness.calls() == 1
async def interrupted_session(harness: Harness) -> None:
started: Final = asyncio.Event()
stopped: Final = asyncio.Event()
cleaned: Final = asyncio.Event()
async def block(request: Payload) -> Payload:
started.set()
try:
await asyncio.Future[None]()
finally:
cleaned.set()
return request
task: Final = asyncio.ensure_future(
harness.interrupt(replace(Recorder().adapter(), pre_request=block), stopped.wait())
)
await started.wait()
stopped.set()
try:
await task
except RuntimeError as error:
assert str(error) == "callback session was cancelled"
else:
raise AssertionError("interrupted session was reused")
await cleaned.wait()
assert harness.calls() == 0
async def cancellation_and_admission(harness: Harness) -> None:
started: Final = asyncio.Event()
cleaning: Final = asyncio.Event()
release: Final = asyncio.Event()
finished: Final = asyncio.Event()
recorder: Final = Recorder()
async def block(request: Payload) -> Payload:
started.set()
try:
await asyncio.Future[None]()
finally:
cleaning.set()
await release.wait()
finished.set()
return request
async def assert_full() -> None:
try:
await harness.execute(Recorder().adapter())
except RuntimeError as error:
assert str(error) == "callback capacity exhausted"
else:
raise AssertionError("capacity released before callback completion")
task: Final = asyncio.ensure_future(harness.execute(replace(recorder.adapter(), pre_request=block)))
await started.wait()
await assert_full()
task.cancel()
try:
await task
except asyncio.CancelledError:
pass
else:
raise AssertionError("request ignored cancellation")
await cleaning.wait()
await assert_full()
assert harness.calls() == 0
assert recorder.events == ()
release.set()
await finished.wait()
assert await harness.execute(Recorder().adapter()) == {"text": "processed:INPUT!"}
assert harness.calls() == 1
@dataclass(frozen=True, slots=True)
class DirectAdapter:
transform: Callable[[Payload], object]
async def synchronous_callbacks(harness: Harness) -> None:
def on_caller_thread() -> None:
token: Final = REQUEST_CONTEXT.set("synchronous")
caller: Final = threading.get_ident()
def transform(request: Payload) -> Payload:
assert REQUEST_CONTEXT.get() == "synchronous"
assert threading.get_ident() == caller
try:
asyncio.get_running_loop()
except RuntimeError:
pass
else:
raise AssertionError("sync callback invented an event loop")
return {"text": request["text"].upper()}
try:
assert harness.sync(DirectAdapter(transform)) == {"text": "INPUT"}
finally:
REQUEST_CONTEXT.reset(token)
await asyncio.to_thread(on_caller_thread)
original: Final = ValueError("sync failure")
def reject(request: Payload) -> Payload:
raise original
try:
harness.sync(DirectAdapter(reject))
except ValueError as error:
assert error is original
else:
raise AssertionError("sync callback error was lost")
def reenter(request: Payload) -> Payload:
return harness.sync(DirectAdapter(lambda value: value))
try:
harness.sync(DirectAdapter(reenter))
except RuntimeError as error:
assert "Tokio context" in str(error)
else:
raise AssertionError("synchronous re-entry was accepted")
async def unfinished() -> None:
pass
coroutine: Final = unfinished()
try:
harness.sync(DirectAdapter(lambda value: coroutine))
except TypeError as error:
assert str(error) == "direct hook returned an awaitable"
else:
raise AssertionError("sync callback accepted an awaitable")
assert inspect.getcoroutinestate(coroutine) == inspect.CORO_CLOSED
class CatalogAdapter:
def __init__(self) -> None:
self.events: tuple[str, ...] = ()
def _record(self, name: str) -> None:
self.events += (name,)
def pre_call(self, payload: object) -> dict[str, str]:
self._record("pre_call")
return {"action": "unchanged"}
def post_call(self, payload: object) -> dict[str, str]:
self._record("post_call")
return {"action": "unchanged"}
def error(self, payload: object) -> None:
self._record("error")
def stream_event(self, payload: object) -> dict[str, str]:
self._record("stream_event")
return {"action": "unchanged"}
def stream_close(self, payload: object) -> None:
self._record("stream_close")
async def before_connect(self, payload: object) -> dict[str, str]:
self._record("before_connect")
return {"action": "unchanged"}
async def connected(self, payload: object) -> None:
self._record("connected")
async def before_send(self, payload: object) -> dict[str, str]:
self._record("before_send")
return {"action": "unchanged"}
async def after_receive(self, payload: object) -> dict[str, str]:
self._record("after_receive")
return {"action": "unchanged"}
async def response_complete(self, payload: object) -> None:
self._record("response_complete")
async def response_error(self, payload: object) -> None:
self._record("response_error")
async def close(self, payload: object) -> None:
self._record("close")
class SessionCatalogAdapter(CatalogAdapter):
async def error(self, payload: object) -> None:
self._record("error")
async def callback_catalogs(harness: Harness) -> None:
streaming: Final = CatalogAdapter()
harness.streaming(streaming)
assert streaming.events == ("pre_call", "post_call", "stream_event", "stream_close", "error")
session: Final = SessionCatalogAdapter()
await harness.session(session)
assert session.events == (
"before_connect",
"connected",
"before_send",
"after_receive",
"response_complete",
"response_error",
"error",
"close",
)
TESTS: Final = (
transforms_and_context,
callback_errors,
retained_callbacks,
registration_and_return_contracts,
provider_failure,
cancellation_and_admission,
interrupted_session,
synchronous_callbacks,
callback_catalogs,
)
def run(name: str, harness: Harness) -> None:
test: Final = next(test for test in TESTS if test.__name__ == name)
asyncio.run(asyncio.wait_for(test(harness), timeout=10))

View file

@ -7,8 +7,10 @@ repository.workspace = true
[dependencies]
pyo3.workspace = true
pyo3-async-runtimes.workspace = true
pythonize.workspace = true
serde.workspace = true
tokio = { workspace = true, features = ["sync"] }
[dev-dependencies]
rstest.workspace = true

View file

@ -0,0 +1,49 @@
import asyncio
import inspect
from collections.abc import Awaitable, Callable
from typing import Final, cast
def invoke_direct(callback: Callable[[object], object], payload: object) -> object:
result: Final = callback(payload)
if inspect.isawaitable(result):
if inspect.iscoroutine(result):
result.close()
raise TypeError("direct hook returned an awaitable")
return result
class Invocation:
def __init__(
self,
callback: Callable[[object], object],
payload: object,
returns_awaitable: bool,
admission: object,
) -> None:
self.callback = callback
self.payload = payload
self.returns_awaitable = returns_awaitable
self.admission: object | None = admission
self.task: asyncio.Task[object] | None = None
self.cancelled = False
async def run(self) -> object:
self.task = asyncio.current_task()
try:
if self.cancelled:
raise asyncio.CancelledError()
if not self.returns_awaitable:
return invoke_direct(self.callback, self.payload)
result: Final = self.callback(self.payload)
if not inspect.isawaitable(result):
raise TypeError("awaitable hook returned a non-awaitable")
return await cast(Awaitable[object], result)
finally:
self.admission = None
self.task = None
def cancel(self) -> None:
self.cancelled = True
if self.task is not None:
self.task.cancel()

View file

@ -0,0 +1,229 @@
use std::ffi::CStr;
use std::marker::PhantomData;
use std::num::NonZeroUsize;
use std::sync::Arc;
use std::thread::{self, ThreadId};
use pyo3::exceptions::{PyRuntimeError, PyTypeError, PyValueError};
use pyo3::prelude::*;
use pyo3_async_runtimes::TaskLocals;
use serde::{Serialize, de::DeserializeOwned};
use tokio::sync::{OwnedSemaphorePermit, Semaphore};
use crate::{Pythonized, from_py};
const PYTHON_RUNTIME_SOURCE: &CStr = pyo3::ffi::c_str!(include_str!("invoke.py"));
pub struct Direct;
pub struct Awaitable;
pub trait ReturnMode: Send {
const AWAITABLE: bool;
}
impl ReturnMode for Direct {
const AWAITABLE: bool = false;
}
impl ReturnMode for Awaitable {
const AWAITABLE: bool = true;
}
pub trait CallbackContext<M: ReturnMode>: Send {
fn invoke(
&mut self,
callable: Py<PyAny>,
payload: Py<PyAny>,
) -> impl Future<Output = PyResult<Py<PyAny>>> + Send;
}
pub struct Callback<I, O, M> {
callable: Py<PyAny>,
signature: PhantomData<fn(I) -> (O, M)>,
}
impl<I, O, M> Callback<I, O, M>
where
I: Serialize + Sync,
O: DeserializeOwned + Send,
M: ReturnMode,
{
pub fn new(callable: Bound<'_, PyAny>) -> PyResult<Self> {
if !callable.is_callable() {
return Err(PyTypeError::new_err("hook binding must be callable"));
}
Ok(Self {
callable: callable.unbind(),
signature: PhantomData,
})
}
pub async fn call<C: CallbackContext<M>>(&mut self, context: &mut C, input: &I) -> PyResult<O> {
let (callable, payload) = Python::attach(|py| {
Ok::<_, PyErr>((
self.callable.clone_ref(py),
Pythonized(input).into_pyobject(py)?.unbind(),
))
})?;
let result = context.invoke(callable, payload).await?;
Python::attach(|py| {
from_py(result.bind(py))
.map_err(|_| PyTypeError::new_err("hook result does not match its typed contract"))
})
}
}
pub fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
let shim = PyModule::from_code(
module.py(),
PYTHON_RUNTIME_SOURCE,
c"litellm_callbacks.py",
c"_litellm_callbacks",
)?;
module.add("__callback_runtime__", shim)
}
struct RuntimeState {
invocation: Py<PyAny>,
direct: Py<PyAny>,
capacity: Arc<Semaphore>,
}
#[derive(Clone)]
pub struct CallbackRuntime(Arc<RuntimeState>);
impl CallbackRuntime {
pub fn new(module: &Bound<'_, PyModule>, max_in_flight: NonZeroUsize) -> PyResult<Self> {
if max_in_flight.get() > Semaphore::MAX_PERMITS {
return Err(PyValueError::new_err(
"callback capacity exceeds the runtime limit",
));
}
let shim = module.getattr("__callback_runtime__")?;
Ok(Self(Arc::new(RuntimeState {
invocation: shim.getattr("Invocation")?.unbind(),
direct: shim.getattr("invoke_direct")?.unbind(),
capacity: Arc::new(Semaphore::new(max_in_flight.get())),
})))
}
pub fn async_context(&self, py: Python<'_>) -> PyResult<AsyncContext> {
Ok(AsyncContext {
runtime: self.clone(),
locals: pyo3_async_runtimes::tokio::get_current_locals(py)?,
interrupted: false,
})
}
pub fn sync_context(&self, py: Python<'_>) -> PyResult<SyncContext> {
Ok(SyncContext {
runtime: self.clone(),
context: py
.import("contextvars")?
.call_method0("copy_context")?
.unbind(),
caller: thread::current().id(),
})
}
fn admit(&self) -> PyResult<OwnedSemaphorePermit> {
Arc::clone(&self.0.capacity)
.try_acquire_owned()
.map_err(|_| PyRuntimeError::new_err("callback capacity exhausted"))
}
}
pub struct SyncContext {
runtime: CallbackRuntime,
context: Py<PyAny>,
caller: ThreadId,
}
impl CallbackContext<Direct> for SyncContext {
async fn invoke(&mut self, callable: Py<PyAny>, payload: Py<PyAny>) -> PyResult<Py<PyAny>> {
if thread::current().id() != self.caller {
return Err(PyRuntimeError::new_err(
"synchronous callbacks must run on the caller thread",
));
}
let _permit = self.runtime.admit()?;
Python::attach(|py| {
self.context
.call_method1(py, "run", (&self.runtime.0.direct, callable, payload))
})
}
}
pub struct AsyncContext {
runtime: CallbackRuntime,
locals: TaskLocals,
interrupted: bool,
}
#[pyclass(frozen)]
struct Admission {
_permit: OwnedSemaphorePermit,
}
impl<M: ReturnMode> CallbackContext<M> for AsyncContext {
async fn invoke(&mut self, callable: Py<PyAny>, payload: Py<PyAny>) -> PyResult<Py<PyAny>> {
if self.interrupted {
return Err(PyRuntimeError::new_err("callback session was cancelled"));
}
let permit = self.runtime.admit()?;
let (mut cancellation, future) = Python::attach(|py| {
let invocation = self.runtime.0.invocation.call1(
py,
(
callable,
payload,
M::AWAITABLE,
Admission { _permit: permit },
),
)?;
let cancellation = CancelOnDrop {
event_loop: self.locals.event_loop(py).unbind(),
invocation: Some(invocation.clone_ref(py)),
};
let coroutine = invocation.call_method0(py, "run")?;
let future = pyo3_async_runtimes::into_future_with_locals(
&self.locals,
coroutine.clone_ref(py).into_bound(py),
);
match future {
Ok(future) => Ok((cancellation, future)),
Err(error) => {
coroutine.call_method0(py, "close")?;
Err(error)
}
}
})?;
self.interrupted = true;
let result = future.await;
cancellation.invocation = None;
self.interrupted = false;
result
}
}
struct CancelOnDrop {
event_loop: Py<PyAny>,
invocation: Option<Py<PyAny>>,
}
impl Drop for CancelOnDrop {
fn drop(&mut self) {
let Some(invocation) = self.invocation.take() else {
return;
};
Python::attach(|py| {
let result = invocation.getattr(py, "cancel").and_then(|cancel| {
self.event_loop
.call_method1(py, "call_soon_threadsafe", (cancel,))
});
if let Err(error) = result {
error.write_unraisable(py, Some(invocation.bind(py)));
}
});
}
}

View file

@ -1,3 +1,4 @@
pub mod callback_runtime;
mod gil;
mod marshal;

View file

@ -30,6 +30,7 @@ from litellm.llms.base_llm.ocr.transformation import (
)
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.rust_bridge import ocr as rust_ocr_bridge
from litellm.rust_bridge.callback_adapters import ProviderLoggingAdapter
from litellm.rust_bridge.request import (
NativeRequestCapabilities,
NativeRequestOptions,
@ -240,27 +241,8 @@ def _prepare_rust_ocr_call(
api_base=prepared_request.api_base,
litellm_params=prepared_request.litellm_params,
)
resolved_complete_url: Final = provider_config.get_complete_url(
api_base=prepared_request.api_base,
model=prepared_request.model,
optional_params=prepared_request.optional_params,
litellm_params=prepared_request.litellm_params,
)
rust_api_base: Final = _rust_bridge_api_base(prepared_request, resolve_api_key)
rust_optional_params: Final = _rust_bridge_optional_params(prepared_request, resolve_api_key)
prepared_request.litellm_logging_obj.pre_call(
input="OCR document processing",
api_key=resolved_api_key,
additional_args={
"complete_input_dict": {
"model": prepared_request.model,
"document": prepared_request.document,
**rust_optional_params,
},
"api_base": resolved_complete_url,
"headers": resolved_headers,
},
)
return PreparedNativeCall(
request=rust_ocr_bridge.NativeOCRRequest(
model=prepared_request.model,
@ -292,6 +274,11 @@ def _prepare_rust_ocr_call(
native_response_format=(prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native"),
),
),
callback_adapter=ProviderLoggingAdapter(
prepared_request.litellm_logging_obj,
"OCR document processing",
resolved_api_key,
),
)

View file

@ -4,7 +4,7 @@ from collections.abc import Callable
from types import ModuleType
from typing import Final, Generic, TypeVar, cast # noqa: TID251 # PyO3 module boundary
from litellm.rust_bridge.loader import get_native_bridge
from litellm.rust_bridge.loader import get_native_bridge, module_route_ready
from litellm.rust_bridge.protocols import NativeModule
BindingT = TypeVar("BindingT")
@ -31,8 +31,12 @@ class NativeBinding(Generic[BindingT]):
self,
select: Callable[[NativeModule], BindingT],
*,
route: str = "",
required_capabilities: frozenset[str] = frozenset({"callbacks"}),
module_loader: Callable[[], ModuleType | None] | None = None,
) -> None:
self._route: Final = route
self._required_capabilities: Final = required_capabilities
self._select: Final = select
self._module_loader: Final = module_loader
self._override: BindingT | None | _Unset = _UNSET
@ -48,7 +52,11 @@ class NativeBinding(Generic[BindingT]):
value: Final = self._select(module)
except AttributeError:
return None
return value if callable(value) else None
if not callable(value):
return None
if self._route and not module_route_ready(native, self._route, self._required_capabilities):
return None
return value
def override(self, value: BindingT | None) -> None:
self._override = value
@ -56,6 +64,9 @@ class NativeBinding(Generic[BindingT]):
def reset(self) -> None:
self._override = _UNSET
def is_overridden(self) -> bool:
return not isinstance(self._override, _Unset)
_DECLINED: Final = NativeBinding(lambda native: native.RustBridgeDeclined)
_UPSTREAM: Final = NativeBinding(lambda native: native.RustUpstreamError)

View file

@ -0,0 +1,147 @@
import json
from collections.abc import Mapping, MutableMapping
from dataclasses import dataclass
from typing import Final, Protocol
from pydantic import BaseModel, ConfigDict, JsonValue
from typing_extensions import ReadOnly, TypedDict
from .callbacks import CallbackDecision, CallbackUnchanged, SessionCallbackHandle
class PreCallArguments(TypedDict):
complete_input_dict: ReadOnly[Mapping[str, JsonValue]]
api_base: ReadOnly[str]
headers: ReadOnly[Mapping[str, str]]
class ProviderLogging(Protocol):
@property
def model_call_details(self) -> MutableMapping[str, object]: ... # mutable-ok: legacy logger stores provider events
def pre_call(self, *, input: object, api_key: str | None, additional_args: PreCallArguments) -> None: ...
def post_call(self, *, original_response: str, input: object, api_key: str | None) -> None: ...
class ProviderEvent(BaseModel):
model_config = ConfigDict(extra="forbid")
provider: str
model: str
call_id: str
trace_id: str | None = None
attempt: int
started_at: float
class ProviderPreCall(ProviderEvent):
request: Mapping[str, JsonValue]
api_base: str
headers: Mapping[str, str]
class ProviderPostCall(ProviderEvent):
response: JsonValue
status_code: int
headers: Mapping[str, str]
ended_at: float
class ProviderError(ProviderEvent):
message: str
stage: str
committed: bool
status_code: int | None
ended_at: float
class StreamEvent(ProviderEvent):
event: JsonValue
sequence: int
class StreamClose(ProviderEvent):
outcome: str
ended_at: float
class SessionEvent(BaseModel):
model_config = ConfigDict(extra="forbid")
session_id: str
call_id: str
trace_id: str | None = None
event: JsonValue | None = None
response_id: str | None = None
sequence: int | None = None
message: str | None = None
def _unchanged() -> CallbackUnchanged:
return {"action": "unchanged"} # mutable-ok: callback protocol requires a concrete decision payload
@dataclass(frozen=True, slots=True)
class ProviderLoggingAdapter:
logging_obj: ProviderLogging
input: object
api_key: str | None
def pre_call(self, payload: object, /) -> CallbackDecision:
event: Final = ProviderPreCall.model_validate(payload)
additional_args: Final[PreCallArguments] = {
"complete_input_dict": event.request,
"api_base": event.api_base,
"headers": event.headers,
}
self.logging_obj.pre_call(input=self.input, api_key=self.api_key, additional_args=additional_args)
return _unchanged()
def post_call(self, payload: object, /) -> CallbackDecision:
event: Final = ProviderPostCall.model_validate(payload)
response: Final = event.response if isinstance(event.response, str) else json.dumps(event.response)
self.logging_obj.post_call(original_response=response, input=self.input, api_key=self.api_key)
return _unchanged()
def error(self, payload: object, /) -> None:
event: Final = ProviderError.model_validate(payload)
self.logging_obj.model_call_details["provider_error"] = event.model_dump()
def stream_event(self, payload: object, /) -> CallbackDecision:
event: Final = StreamEvent.model_validate(payload)
self.logging_obj.model_call_details["provider_stream_event"] = event.model_dump()
return _unchanged()
def stream_close(self, payload: object, /) -> None:
event: Final = StreamClose.model_validate(payload)
self.logging_obj.model_call_details["provider_stream_close"] = event.model_dump()
@dataclass(frozen=True, slots=True)
class SessionCallbackAdapter:
callback: SessionCallbackHandle
def before_connect(self, payload: object, /) -> CallbackDecision:
return self.callback.before_connect(SessionEvent.model_validate(payload).model_dump())
def connected(self, payload: object, /) -> None:
self.callback.connected(SessionEvent.model_validate(payload).model_dump())
def before_send(self, payload: object, /) -> CallbackDecision:
return self.callback.before_send(SessionEvent.model_validate(payload).model_dump())
def after_receive(self, payload: object, /) -> CallbackDecision:
return self.callback.after_receive(SessionEvent.model_validate(payload).model_dump())
def response_complete(self, payload: object, /) -> None:
self.callback.response_complete(SessionEvent.model_validate(payload).model_dump())
def response_error(self, payload: object, /) -> None:
self.callback.response_error(SessionEvent.model_validate(payload).model_dump())
def error(self, payload: object, /) -> None:
self.callback.error(SessionEvent.model_validate(payload).model_dump())
def close(self, payload: object, /) -> None:
self.callback.close(SessionEvent.model_validate(payload).model_dump())

View file

@ -0,0 +1,64 @@
from typing import Literal, Protocol, TypeAlias
from typing_extensions import ReadOnly, TypedDict
class CallbackUnchanged(TypedDict):
action: ReadOnly[Literal["unchanged"]]
class CallbackReplace(TypedDict):
action: ReadOnly[Literal["replace"]]
payload: ReadOnly[object]
class CallbackReject(TypedDict):
action: ReadOnly[Literal["reject"]]
message: ReadOnly[str]
status_code: ReadOnly[int | None]
CallbackDecision: TypeAlias = CallbackUnchanged | CallbackReplace | CallbackReject
class ProviderAttemptCallbackHandle(Protocol):
"""Observe the provider operation inside one native call.
Successful operations receive ``pre_call`` and ``post_call`` once. Failed
operations receive ``pre_call`` and ``error`` once. Outer SDK success and
failure callbacks remain owned by Python after endpoint dispatch completes.
"""
def pre_call(self, payload: object, /) -> CallbackDecision: ...
def post_call(self, payload: object, /) -> CallbackDecision: ...
def error(self, payload: object, /) -> None: ...
class OneShotCallbackHandle(ProviderAttemptCallbackHandle, Protocol):
pass
class StreamingCallbackHandle(ProviderAttemptCallbackHandle, Protocol):
def stream_event(self, payload: object, /) -> CallbackDecision: ...
def stream_close(self, payload: object, /) -> None: ...
class SessionCallbackHandle(Protocol):
def before_connect(self, payload: object, /) -> CallbackDecision: ...
def connected(self, payload: object, /) -> None: ...
def before_send(self, payload: object, /) -> CallbackDecision: ...
def after_receive(self, payload: object, /) -> CallbackDecision: ...
def response_complete(self, payload: object, /) -> None: ...
def response_error(self, payload: object, /) -> None: ...
def error(self, payload: object, /) -> None: ...
def close(self, payload: object, /) -> None: ...

View file

@ -2,8 +2,9 @@
from __future__ import annotations
from collections.abc import Mapping, Set
from types import ModuleType
from typing import Final
from typing import Final, cast # noqa: TID251 # runtime typing constructs
_BRIDGE_SENTINEL: Final = object()
_cached_bridge: ModuleType | None | object = _BRIDGE_SENTINEL
@ -33,3 +34,21 @@ def reset_native_bridge_cache() -> None:
def native_bridge_available() -> bool:
"""Whether the packaged Rust extension is importable."""
return get_native_bridge() is not None
def native_route_ready(route: str, required_capabilities: frozenset[str] = frozenset()) -> bool:
native: Final = get_native_bridge()
if native is None:
return False
return module_route_ready(native, route, required_capabilities)
def module_route_ready(native: ModuleType, route: str, required_capabilities: frozenset[str]) -> bool:
ready_endpoints: Final = getattr(native, "ready_endpoints", None)
if not isinstance(ready_endpoints, Mapping):
return False
typed_endpoints: Final = cast( # cast-ok: native metadata was validated as a Mapping
Mapping[object, object], ready_endpoints
)
capabilities: Final = typed_endpoints.get(route)
return isinstance(capabilities, Set) and required_capabilities.issubset(capabilities)

View file

@ -41,6 +41,21 @@ def load_rust_aocr() -> RustAocr | None:
return _OCR.asynchronous.load()
def supports_callback_adapter(*, asynchronous: bool = False) -> bool:
binding = _OCR.asynchronous if asynchronous else _OCR.sync
if binding.is_overridden():
return False
from litellm.rust_bridge import get_native_bridge
from litellm.rust_bridge.loader import native_route_ready
native: Final = get_native_bridge()
return (
native is not None
and hasattr(native, "__python_callback_runtime__")
and native_route_ready("ocr", frozenset({"callbacks"}))
)
def dispatch_ocr(
*,
prepare: Callable[[], PreparedNativeCall[NativeOCRRequest]],

View file

@ -0,0 +1,3 @@
from .callback_adapters import PreCallArguments, ProviderLoggingAdapter
__all__ = ("PreCallArguments", "ProviderLoggingAdapter")

View file

@ -3,6 +3,7 @@ from __future__ import annotations
from collections.abc import Awaitable, Mapping, Sequence
from typing import Protocol
from .callbacks import SessionCallbackHandle
from .request import (
NativeChatCompletionsRequest,
NativeFunction,
@ -53,6 +54,7 @@ class RustResponsesWebSocketConnection(Protocol):
*,
options: NativeRequestOptions,
context: NativeRequestContext,
callback_adapter: SessionCallbackHandle | None = None,
) -> RustResponsesWebSocket: ...

View file

@ -5,6 +5,8 @@ from dataclasses import dataclass, replace
from types import MappingProxyType
from typing import Generic, Protocol, TypeVar
from .callbacks import OneShotCallbackHandle
@dataclass(frozen=True, slots=True)
class NativeBedrockOptions:
@ -73,7 +75,7 @@ def vertex_options(params: Mapping[str, object]) -> NativeVertexOptions:
)
from typing_extensions import ReadOnly, TypedDict
from typing_extensions import ReadOnly, TypedDict, TypeVar
class NativePreCallDetails(TypedDict):
@ -164,27 +166,45 @@ def with_capabilities(
RequestT = TypeVar("RequestT")
RequestContraT = TypeVar("RequestContraT", contravariant=True)
ResultT = TypeVar("ResultT", covariant=True)
CallbackT = TypeVar("CallbackT", default=OneShotCallbackHandle)
CallbackContraT = TypeVar("CallbackContraT", contravariant=True, default=OneShotCallbackHandle)
@dataclass(frozen=True, slots=True)
class PreparedNativeCall(Generic[RequestT]):
class PreparedNativeCall(Generic[RequestT, CallbackT]):
request: RequestT
options: NativeRequestOptions = NativeRequestOptions()
context: NativeRequestContext = NativeRequestContext()
callback_adapter: CallbackT | None = None
class NativeFunction(Protocol[RequestContraT, ResultT]):
class NativeFunction(Protocol[RequestContraT, ResultT, CallbackContraT]):
def __call__(
self,
request: RequestContraT,
*,
options: NativeRequestOptions,
context: NativeRequestContext,
callback_adapter: CallbackContraT | None = None,
) -> ResultT: ...
def call_native(native: NativeFunction[RequestT, ResultT], prepared: PreparedNativeCall[RequestT]) -> ResultT:
return native(prepared.request, options=prepared.options, context=prepared.context)
def call_native(
native: NativeFunction[RequestT, ResultT, CallbackT],
prepared: PreparedNativeCall[RequestT, CallbackT],
) -> ResultT:
if prepared.callback_adapter is None:
return native(
prepared.request,
options=prepared.options,
context=prepared.context,
)
return native(
prepared.request,
options=prepared.options,
context=prepared.context,
callback_adapter=prepared.callback_adapter,
)
@dataclass(frozen=True, slots=True)

View file

@ -10,6 +10,7 @@ import httpx
from websockets.exceptions import ConnectionClosedOK
from litellm.rust_bridge.bindings import UNCHANGED, Unchanged
from litellm.rust_bridge.callbacks import SessionCallbackHandle
from litellm.rust_bridge.configuration import rust_enabled
from litellm.rust_bridge.protocols import (
RustResponsesWebSocket,
@ -94,11 +95,12 @@ async def connect(
timeout: float | httpx.Timeout | None,
model: str = "responses websocket",
provider: str = "openai",
callback_adapter: SessionCallbackHandle | None = None,
fallback: Callable[[], Awaitable[Connection | None]] = async_none,
context: NativeRequestContext | None = None,
) -> Connection | None:
return await _RESPONSES_WEBSOCKET.ainvoke(
prepare=lambda: PreparedNativeCall(
prepare=lambda: PreparedNativeCall[NativeResponsesWebSocketRequest, SessionCallbackHandle](
NativeResponsesWebSocketRequest(
url=url,
),
@ -115,6 +117,7 @@ async def connect(
requires_connection=True,
),
),
callback_adapter=callback_adapter,
),
call=lambda connection_type, request: call_native(connection_type.connect, request),
preflight=lambda: assess_route(_PREFLIGHT, model, provider),
@ -132,6 +135,7 @@ async def open_connection(
timeout: float | httpx.Timeout | None,
model: str,
provider: str,
callback_adapter: SessionCallbackHandle | None = None,
fallback: Callable[[], AbstractAsyncContextManager[Connection]],
context: NativeRequestContext | None = None,
) -> AsyncGenerator[Connection]:
@ -146,6 +150,7 @@ async def open_connection(
timeout=timeout,
model=model,
provider=provider,
callback_adapter=callback_adapter,
fallback=python_connection,
context=context,
)

View file

@ -90,7 +90,7 @@ class EndpointBinding(Generic[BindingT]):
select: Callable[[NativeModule], SelectedT],
enabled: RustEnablement,
) -> EndpointBinding[SelectedT]:
binding: Final = NativeBinding(select)
binding: Final = NativeBinding(select, route=route)
return EndpointBinding(
route=route,
load=binding.load,
@ -108,6 +108,9 @@ class EndpointBinding(Generic[BindingT]):
raise RuntimeError("only native Rust bridges support binding resets")
self._native_binding.reset()
def is_overridden(self) -> bool:
return self._native_binding is not None and self._native_binding.is_overridden()
def _attempt(
self,
*,

View file

@ -45,7 +45,7 @@ from litellm.secret_managers.main import get_secret_str
from litellm.types.utils import FileTypes, TranscriptionResponse
_TRANSCRIPTION: Final[EndpointDispatch[RustTranscription, RustAtranscription]] = EndpointDispatch.native(
route="audio transcription",
route="transcription",
sync=lambda native: native.transcription,
asynchronous=lambda native: native.atranscription,
enabled=always_enabled,

View file

@ -2,6 +2,7 @@ from __future__ import annotations
import asyncio
from collections.abc import Awaitable
from contextlib import ExitStack
from pathlib import Path
from typing import Final, Protocol, cast
@ -49,14 +50,40 @@ def _entrypoint(spec: RouteSpec, engine: Engine, *, asynchronous: bool) -> SdkCa
def _collect(
function: SdkCall, kwargs: dict[str, object], engine: Engine, *, asynchronous: bool
function: SdkCall, kwargs: dict[str, object], engine: Engine, *, route: str, asynchronous: bool
) -> tuple[FunctionTraceEvent, ...]:
if engine == "rust":
return native_trace_events(_invoke(function, kwargs, asynchronous=asynchronous))
import litellm
with profile_python(Path(litellm.__file__).parent, threads=True) as profiler:
_invoke(function, kwargs, asynchronous=asynchronous)
with ExitStack() as stack:
if route == "transcription":
from litellm.rust_bridge import get_native_bridge
from litellm.rust_bridge.transcription import (
RustAtranscription,
RustRouteDecline,
RustTranscription,
configure_rust_transcription,
)
native: Final = get_native_bridge()
if native is None:
raise RuntimeError("native transcription is required for diagnostic trace parity")
# This required-native SDK route is injected only for diagnostic comparison.
# Normal callback readiness remains empty before and after this scope.
configure_rust_transcription(
transcription=cast(RustTranscription, native.transcription),
atranscription=cast(RustAtranscription, native.atranscription),
decline=cast(RustRouteDecline, native.transcription_decline),
)
stack.callback(
configure_rust_transcription,
transcription=None,
atranscription=None,
decline=None,
)
with profile_python(Path(litellm.__file__).parent, threads=True) as profiler:
_invoke(function, kwargs, asynchronous=asynchronous)
return tuple(profiler.events)
@ -141,6 +168,7 @@ def collect_trace(
function,
_native_kwargs(spec.route, kwargs) if engine == "rust" else kwargs,
engine,
route=spec.route,
asynchronous=asynchronous,
)
provider.take_requests(len(fixture.provider_responses))

View file

@ -16,6 +16,8 @@ COMMON_MAPPINGS: Final = (
mapping(rust_span="validate_environment", python_frame=r"(?<!_)validate_environment$"),
mapping(rust_span="complete_url", python_frame=r"get_complete_url$"),
mapping(rust_span="transform_ocr_request", python_frame=r"(?<!async_)transform_ocr_request$"),
# Rust wraps the HTTP attempt to invoke observers; Python has no matching frame.
mapping(rust_span="send_provider_request"),
mapping(rust_span="http_request", python_frame=r"AsyncHTTPHandler\.post$|HTTPHandler\.post$"),
)
@ -48,7 +50,7 @@ AZURE_COMMON_MAPPINGS: Final = (
r"|MistralOCRConfig\.transform_ocr_request$"
),
),
COMMON_MAPPINGS[-1],
*COMMON_MAPPINGS[-2:],
)
AZURE_SYNC_MAPPINGS: Final = (
*AZURE_COMMON_MAPPINGS,
@ -201,7 +203,7 @@ VERTEX_COMMON_MAPPINGS: Final = (
r"|MistralOCRConfig\.transform_ocr_request$"
),
),
COMMON_MAPPINGS[-1],
*COMMON_MAPPINGS[-2:],
)
VERTEX_SYNC_MAPPINGS: Final = (
*VERTEX_COMMON_MAPPINGS,
@ -228,6 +230,8 @@ DEEPSEEK_COMMON_MAPPINGS: Final = (
rust_span="transform_ocr_request",
python_frame=r"VertexAIDeepSeekOCRConfig\.transform_ocr_request$",
),
# Rust wraps the HTTP attempt to invoke observers; Python has no matching frame.
mapping(rust_span="send_provider_request"),
mapping(rust_span="http_request", python_frame=r"AsyncHTTPHandler\.post$|HTTPHandler\.post$"),
mapping(
rust_span="transform_ocr_response",
@ -264,6 +268,8 @@ DOCUMENT_INTELLIGENCE_COMMON_MAPPINGS: Final = (
mapping(
rust_span="transform_ocr_request", python_frame=r"AzureDocumentIntelligenceOCRConfig\.transform_ocr_request$"
),
# Rust wraps the HTTP attempt to invoke observers; Python has no matching frame.
mapping(rust_span="send_provider_request"),
mapping(rust_span="http_request", python_frame=r"AsyncHTTPHandler\.post$|HTTPHandler\.post$"),
mapping(
rust_span="poll_document_intelligence",

View file

@ -23,7 +23,7 @@ MAPPINGS: Final = (
mapping(rust_span="transform_transcription_request"),
mapping(
rust_span="execute_audio_transcription_provider_call",
python_frame=r"BedrockAudioTranscriptionRustDispatch\.(?:async_)?audio_transcriptions$",
python_frame=r"rust_bridge/request\.py:\d+ call_native$",
),
mapping(rust_span="transform_transcription_response"),
mapping(rust_span="http_request"),

View file

@ -85,7 +85,12 @@ class ExplodingAsyncMessages:
self.calls = 0
async def __call__(
self, request: NativeMessagesRequest, *, options: object, context: NativeRequestContext
self,
request: NativeMessagesRequest,
*,
options: object,
context: NativeRequestContext,
callback_adapter: object | None = None,
) -> dict[str, object]:
self.calls += 1
raise AssertionError("bridge must not be called")
@ -96,7 +101,12 @@ class RaisingAsyncMessages:
self.calls = 0
async def __call__(
self, request: NativeMessagesRequest, *, options: object, context: NativeRequestContext
self,
request: NativeMessagesRequest,
*,
options: object,
context: NativeRequestContext,
callback_adapter: object | None = None,
) -> dict[str, object]:
self.calls += 1
raise RuntimeError("upstream request failed with status 400: bad request")

View file

@ -0,0 +1,201 @@
import asyncio
import contextvars
import inspect
import json
import threading
from collections.abc import Generator, Mapping
from contextlib import contextmanager
from dataclasses import dataclass
from datetime import datetime
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from typing import Final, TypedDict
from typing_extensions import ReadOnly
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.llms.base_llm.ocr.transformation import OCRResponse
MODEL: Final = "mistral/mistral-ocr-4-1"
DOCUMENT: Final = {"type": "document_url", "document_url": "https://example.com/document.pdf"}
RESPONSE: Final = (
'{"pages":[{"index":0,"markdown":"callback-test"}],"model":"mistral-ocr-4-1","usage_info":{"pages_processed":1}}'
)
REQUEST_CONTEXT: Final = contextvars.ContextVar("ocr_callback_test_context", default="unset")
@dataclass(frozen=True, slots=True)
class CallbackEvent:
name: str
model: object
call_id: object
original_response: object
input: object
additional_args: str
response_type: str
start_time: datetime | None
end_time: datetime | None
context: str
native_provider_hook: bool
class CallbackRecorder(CustomLogger):
def __init__(self, asynchronous: bool, raises: str = "", name: str = "recorder", expected_calls: int = 1) -> None:
super().__init__() # pyright: ignore[reportUnknownMemberType] # CustomLogger exposes untyped keyword arguments
self.asynchronous = asynchronous
self.name = name
self.expected_calls = expected_calls
self.raises = raises
self.events: tuple[CallbackEvent, ...] = ()
self.done = threading.Event()
self.lock = threading.Lock()
def record(
self,
name: str,
kwargs: Mapping[str, object],
response: object = None,
start_time: datetime | None = None,
end_time: datetime | None = None,
) -> None:
event: Final = CallbackEvent(
name,
kwargs.get("model"),
kwargs.get("litellm_call_id"),
kwargs.get("original_response"),
kwargs.get("input"),
json.dumps(kwargs.get("additional_args"), sort_keys=True, default=str),
type(response).__name__,
start_time,
end_time,
REQUEST_CONTEXT.get(),
any(frame.filename == "litellm_callbacks.py" for frame in inspect.stack(context=0)),
)
with self.lock:
self.events += (event,)
terminals: Final = tuple(
item.name for item in self.events if "success" in item.name or "failure" in item.name
)
if len(terminals) >= self.expected_calls * (2 if self.asynchronous and "failure" in name else 1):
self.done.set()
if name == self.raises:
raise ValueError("intentional observer failure")
def log_pre_api_call(self, model: str, messages: object, kwargs: Mapping[str, object]) -> None:
assert model == kwargs.get("model")
assert isinstance(messages, list)
self.record("pre", kwargs)
def log_post_api_call(
self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime | None
) -> None:
self.record("post", kwargs, response_obj, start_time, end_time)
def log_success_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime
) -> None:
self.record("success", kwargs, response_obj, start_time, end_time)
def log_failure_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime
) -> None:
self.record("failure", kwargs, response_obj, start_time, end_time)
async def async_log_success_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime
) -> None:
self.record("async_success", kwargs, response_obj, start_time, end_time)
async def async_log_failure_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime
) -> None:
self.record("async_failure", kwargs, response_obj, start_time, end_time)
async def wait(self) -> tuple[CallbackEvent, ...]:
assert await asyncio.to_thread(self.done.wait, 5), tuple(event.name for event in self.events)
return self.events
class OcrUpstream(ThreadingHTTPServer):
daemon_threads = True
def __init__(self, status: int, body: str, stall: bool = False) -> None:
super().__init__(("127.0.0.1", 0), OcrHandler)
self.status = status
self.body = body.encode()
self.stall = stall
self.release = threading.Event()
self.started = threading.Event()
self.requests: tuple[tuple[str, str], ...] = ()
self.lock = threading.Lock()
@property
def api_base(self) -> str:
return f"http://127.0.0.1:{self.server_port}/v1"
class OcrHandler(BaseHTTPRequestHandler):
def do_POST(self) -> None:
assert isinstance(self.server, OcrUpstream)
body: Final = self.rfile.read(int(self.headers["content-length"])).decode()
with self.server.lock:
self.server.requests += ((self.path, body),)
self.server.started.set()
if self.server.stall:
self.server.release.wait(5)
return
self.send_response(self.server.status)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(self.server.body)))
self.end_headers()
self.wfile.write(self.server.body)
def log_message(self, format: str, *args: object) -> None:
pass
@contextmanager
def ocr_upstream(status: int = 200, body: str = RESPONSE, stall: bool = False) -> Generator[OcrUpstream, None, None]:
with OcrUpstream(status, body, stall) as server:
thread: Final = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
try:
yield server
finally:
server.release.set()
server.shutdown()
thread.join(timeout=5)
assert not thread.is_alive()
def assert_provider_request(server: OcrUpstream) -> None:
assert len(server.requests) == 1
path, body = server.requests[0]
assert path == "/v1/ocr"
assert json.loads(body) == {"model": "mistral-ocr-4-1", "document": DOCUMENT}
class OcrArguments(TypedDict):
model: ReadOnly[str]
document: ReadOnly[dict[str, str]]
api_key: ReadOnly[str]
api_base: ReadOnly[str]
callbacks: ReadOnly[list[CustomLogger]]
num_retries: ReadOnly[int]
timeout: ReadOnly[float]
def verify_installed_package() -> None:
root: Final = Path(litellm.__file__).resolve()
assert "site-packages" in root.parts, f"SDK must come from the installed wheel: {root}"
async def call_ocr(arguments: OcrArguments, asynchronous: bool) -> tuple[OCRResponse | None, Exception | None]:
try:
response: Final = (
await litellm.aocr(**arguments) if asynchronous else await asyncio.to_thread(litellm.ocr, **arguments)
)
return OCRResponse.model_validate(response), None
except Exception as error:
return None, error

View file

@ -0,0 +1,66 @@
import asyncio
import os
import sys
from typing import Final
from unittest.mock import patch
from callback_support import CallbackRecorder, OcrArguments, call_ocr, verify_installed_package
import litellm
from litellm.rust_bridge import get_native_bridge
async def smoke(*, baseline: bool) -> None:
verify_installed_package()
litellm.rust(True)
for asynchronous, rejected in ((False, False), (True, False), (True, True)):
await smoke_case(asynchronous, rejected, baseline=baseline)
async def smoke_case(asynchronous: bool, rejected: bool, *, baseline: bool) -> None:
recorder: Final = CallbackRecorder(asynchronous, name=f"live-{asynchronous}-{rejected}")
arguments: Final[OcrArguments] = {
"model": "mistral/mistral-ocr-latest",
"document": {
"type": "document_url",
"document_url": "https://www.w3.org/WAI/ER/tests/xhtml/testfiles/resources/pdf/dummy.pdf",
},
"api_key": "invalid-smoke-key" if rejected else os.environ["MISTRAL_API_KEY"],
"callbacks": [recorder],
"api_base": "https://api.mistral.ai/v1",
"num_retries": 0,
"timeout": 60,
}
response, error = await call_ocr(arguments, asynchronous)
assert (error is not None) is rejected
if error is not None and not baseline:
assert isinstance(error, litellm.AuthenticationError) and error.status_code == 401
if response is not None:
assert len(response.pages) == 1
outcome: Final = (
f"pages={len(response.pages)}"
if response is not None
else f"{type(error).__name__} status={getattr(error, 'status_code', None)}"
)
events: Final = await recorder.wait()
names: Final = tuple(event.name for event in events)
if rejected:
assert names[0] == "pre" and sorted(names[1:]) == ["async_failure", "failure"]
else:
assert names == ("pre", *(() if baseline else ("post",)), "async_success" if asynchronous else "success")
if not baseline:
assert all(event.native_provider_hook for event in events if event.name in ("pre", "post"))
sys.stdout.write(f"async={asynchronous} rejected={rejected} {outcome} events={names}\n")
sys.stdout.flush()
if __name__ == "__main__":
if "--run-live" not in sys.argv:
raise SystemExit("Pass --run-live and set MISTRAL_API_KEY to make real, billable provider calls")
if "--baseline" in sys.argv:
asyncio.run(smoke(baseline=True))
else:
native: Final = get_native_bridge()
assert native is not None
with patch.object(native, "ready_endpoints", {"ocr": frozenset({"callbacks"})}, create=True):
asyncio.run(smoke(baseline=False))

View file

@ -0,0 +1,280 @@
import asyncio
import json
import sys
from collections import Counter
from collections.abc import Mapping
from dataclasses import dataclass
from datetime import datetime
from typing import Final
from unittest.mock import patch
from callback_support import (
DOCUMENT,
MODEL,
REQUEST_CONTEXT,
RESPONSE,
CallbackRecorder,
OcrArguments,
assert_provider_request,
call_ocr,
ocr_upstream,
verify_installed_package,
)
from pydantic import TypeAdapter
import litellm
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.rust_bridge import get_native_bridge
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.ocr import supports_callback_adapter
from litellm.rust_bridge.callback_adapters import PreCallArguments
@dataclass(frozen=True, slots=True)
class Outcome:
response: str | None
error_type: str | None
status: int | None
events: tuple[str, ...]
response_types: tuple[str, ...]
provider_body: str
async def exercise(asynchronous: bool, rust: bool, case: str, *, native_expected: bool | None = None) -> Outcome:
litellm.rust(rust)
litellm.logging_callback_manager._reset_all_callbacks() # pyright: ignore[reportPrivateUsage] # isolate SDK cases in this dedicated test process
native: Final = rust if native_expected is None else native_expected
if native:
assert supports_callback_adapter(asynchronous=asynchronous), "native callback bridge must be installed"
status: Final = int(case) if case.isdigit() else 500 if case in ("raise_failure", "raise_async_failure") else 200
body: Final = "invalid-json" if case == "malformed" else RESPONSE if status == 200 else '{"error":"rejected"}'
successful: Final = status == 200 and case not in ("malformed", "timeout")
raises: Final = case.removeprefix("raise_") if case.startswith("raise_") else ""
recorder: Final = CallbackRecorder(asynchronous, raises=raises, name=f"first-{asynchronous}-{rust}-{case}")
follower: Final = CallbackRecorder(asynchronous, name=f"second-{asynchronous}-{rust}-{case}")
context: Final = f"request-{asynchronous}-{rust}-{case}"
token: Final = REQUEST_CONTEXT.set(context)
try:
with ocr_upstream(status, body, stall=case == "timeout") as upstream:
params: Final[OcrArguments] = {
"model": MODEL,
"document": DOCUMENT,
"api_key": "test-key",
"api_base": upstream.api_base,
"callbacks": [recorder, follower],
"num_retries": 0,
"timeout": 0.2 if case == "timeout" else 5,
}
result, error = await call_ocr(params, asynchronous)
assert (error is None) is successful
if error is not None:
if case.isdigit():
assert getattr(error, "status_code", None) == status
if case == "timeout":
assert isinstance(error, litellm.Timeout)
if result is not None:
assert result.pages[0].markdown == "callback-test"
await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), 5)
events: Final = tuple(event for event in await recorder.wait() if native or event.name != "post")
following: Final = tuple(event for event in await follower.wait() if native or event.name != "post")
names: Final = tuple(event.name for event in events)
prefix: Final = ("pre", "post") if native and status == 200 and case != "timeout" else ("pre",)
assert names[: len(prefix)] == prefix, (asynchronous, rust, case, native, names, prefix)
terminal: Final = "success" if successful else "failure"
expected: Final = (
("async_success",)
if asynchronous and successful
else ((terminal, f"async_{terminal}") if asynchronous else (terminal,))
)
assert sorted(names[len(prefix) :]) == sorted(expected), (asynchronous, rust, case, native, names, expected)
assert sorted(event.name for event in following) == sorted(names)
assert len({event.call_id for event in events}) == 1
assert isinstance(events[0].call_id, str) and events[0].call_id
assert all(event.model == "mistral-ocr-4-1" for event in events)
assert events[0].input == "OCR document processing"
pre_arguments: Final = TypeAdapter(PreCallArguments).validate_json(events[0].additional_args)
assert pre_arguments["complete_input_dict"] == {"model": "mistral-ocr-4-1", "document": DOCUMENT}
assert pre_arguments["api_base"] == upstream.api_base + "/ocr"
assert {key.lower(): value for key, value in pre_arguments["headers"].items()}[
"authorization"
] == "Bearer test-key"
for event in events[: len(prefix)]:
assert event.context == context
assert event.native_provider_hook is native
if "post" in prefix:
assert events[1].original_response == body
assert events[1].response_type == "NoneType"
assert events[1].start_time is not None and events[1].end_time is None
for event in events[len(prefix) :]:
assert event.start_time is not None and event.end_time is not None
assert event.end_time >= event.start_time
assert_provider_request(upstream)
return Outcome(
result.model_dump_json() if result is not None else None,
type(error).__name__ if error is not None else None,
getattr(error, "status_code", None),
tuple(sorted(event.name for event in events if event.name != "post")),
tuple(sorted(event.response_type for event in events if event.name != "post")),
json.dumps(json.loads(upstream.requests[0][1]), sort_keys=True),
)
finally:
REQUEST_CONTEXT.reset(token)
async def verify_parity() -> None:
for asynchronous in (False, True):
for case in (
"success",
"malformed",
"401",
"429",
"500",
"timeout",
"raise_pre",
"raise_post",
"raise_success",
"raise_async_success",
"raise_failure",
"raise_async_failure",
):
await verify_case(asynchronous, case)
await verify_concurrency()
for rust in (False, True):
await verify_delayed_terminal(rust)
for rust in (False, True):
await verify_without_loggers(rust)
async def verify_case(asynchronous: bool, case: str) -> None:
python: Final = await exercise(asynchronous, False, case)
native: Final = await exercise(asynchronous, True, case)
assert python == native, (asynchronous, case, python, native)
sys.stdout.write(f"PASS async={asynchronous} case={case}\n")
sys.stdout.flush()
async def verify_concurrency() -> None:
litellm.rust(True)
litellm.logging_callback_manager._reset_all_callbacks() # pyright: ignore[reportPrivateUsage] # isolate the concurrent SDK scenario
recorder: Final = CallbackRecorder(True, name="concurrent", expected_calls=32)
with ocr_upstream() as upstream:
async def one(index: int) -> None:
token: Final = REQUEST_CONTEXT.set(f"concurrent-{index}")
try:
result: Final = await litellm.aocr(
model=MODEL,
document=DOCUMENT,
api_key="test-key",
api_base=upstream.api_base,
callbacks=[recorder],
num_retries=0,
timeout=10,
)
assert result.pages[0].markdown == "callback-test"
finally:
REQUEST_CONTEXT.reset(token)
await asyncio.gather(*(one(index) for index in range(32)))
await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), 5)
events: Final = await recorder.wait()
assert len(upstream.requests) == 32
assert Counter(event.name for event in events) == {"pre": 32, "post": 32, "async_success": 32}
assert len({event.call_id for event in events}) == 32
for index in range(32):
assert tuple(event.name for event in events if event.context == f"concurrent-{index}") == (
"pre",
"post",
"async_success",
)
assert all(event.native_provider_hook for event in events if event.name in ("pre", "post"))
class DelayedRecorder(CallbackRecorder):
def __init__(self) -> None:
super().__init__(True, name="delayed")
self.started = asyncio.Event()
self.release = asyncio.Event()
async def async_log_success_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime
) -> None:
self.started.set()
await self.release.wait()
await super().async_log_success_event(kwargs, response_obj, start_time, end_time)
async def verify_delayed_terminal(rust: bool) -> None:
litellm.rust(rust)
litellm.logging_callback_manager._reset_all_callbacks() # pyright: ignore[reportPrivateUsage] # isolate delayed delivery
recorder: Final = DelayedRecorder()
with ocr_upstream() as upstream:
result: Final = await asyncio.wait_for(
litellm.aocr(
model=MODEL,
document=DOCUMENT,
api_key="test-key",
api_base=upstream.api_base,
callbacks=[recorder],
num_retries=0,
),
5,
)
assert result.pages[0].markdown == "callback-test"
await asyncio.wait_for(recorder.started.wait(), 5)
prefix: Final = ("pre", "post") if rust else ("pre",)
assert tuple(event.name for event in recorder.events if rust or event.name != "post") == prefix
recorder.release.set()
await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), 5)
events: Final = await recorder.wait()
assert tuple(event.name for event in events if rust or event.name != "post") == (*prefix, "async_success"), (
events
)
async def verify_without_loggers(rust: bool) -> None:
litellm.rust(rust)
litellm.logging_callback_manager._reset_all_callbacks() # pyright: ignore[reportPrivateUsage] # exercise the no-logger configuration
with ocr_upstream() as upstream:
result: Final = await litellm.aocr(
model=MODEL, document=DOCUMENT, api_key="test-key", api_base=upstream.api_base, num_retries=0
)
assert result.pages[0].markdown == "callback-test"
assert_provider_request(upstream)
await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), 5)
def assert_native_unavailable() -> None:
assert not supports_callback_adapter()
for binding in (
NativeBinding(lambda native: native.ocr, route="ocr"),
NativeBinding(lambda native: native.chat_completions, route="chat_completions"),
NativeBinding(lambda native: native.messages, route="messages"),
NativeBinding(lambda native: native.transcription, route="transcription"),
NativeBinding(lambda native: native.ResponsesWebSocketConnection, route="responses_websocket"),
):
assert binding.load() is None
async def verify_unavailable() -> None:
assert_native_unavailable()
for asynchronous in (False, True):
for case in ("success", "401"):
await exercise(asynchronous, True, case, native_expected=False)
async def verify_foundation() -> None:
assert_native_unavailable()
native: Final = get_native_bridge()
assert native is not None
with patch.object(native, "ready_endpoints", {"ocr": frozenset({"callbacks"})}, create=True):
await verify_parity()
assert_native_unavailable()
if __name__ == "__main__":
if "--installed" in sys.argv:
verify_installed_package()
asyncio.run(
verify_unavailable() if "--without-native" in sys.argv or "--unready" in sys.argv else verify_foundation()
)

View file

@ -11,6 +11,7 @@ import pytest
import litellm
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.rust_bridge.callback_adapters import ProviderLoggingAdapter
from litellm.rust_bridge import configuration
from litellm.rust_bridge.request import (
NativeOCRRequest,
@ -52,6 +53,7 @@ class RecordingBridge:
def __init__(self) -> None:
self.calls: list[dict[str, object]] = []
self.callback_adapter: object | None = None
def __call__(
self,
@ -59,7 +61,9 @@ class RecordingBridge:
*,
options: NativeRequestOptions,
context: NativeRequestContext,
callback_adapter: object | None = None,
) -> dict[str, object]:
self.callback_adapter = callback_adapter
self.calls.append(
{
"model": request.model,
@ -81,6 +85,7 @@ class RecordingAsyncBridge:
def __init__(self) -> None:
self.calls: list[dict[str, object]] = []
self.callback_adapter: object | None = None
async def __call__(
self,
@ -88,7 +93,9 @@ class RecordingAsyncBridge:
*,
options: NativeRequestOptions,
context: NativeRequestContext,
callback_adapter: object | None = None,
) -> dict[str, object]:
self.callback_adapter = callback_adapter
self.calls.append(
{
"model": request.model,
@ -112,6 +119,7 @@ class RaisingBridge:
*,
options: NativeRequestOptions,
context: NativeRequestContext,
callback_adapter: object | None = None,
) -> dict[str, object]:
raise RuntimeError("bridge failed")
@ -123,6 +131,7 @@ class RaisingAsyncBridge:
*,
options: NativeRequestOptions,
context: NativeRequestContext,
callback_adapter: object | None = None,
) -> dict[str, object]:
raise RuntimeError("bridge failed")
@ -394,6 +403,7 @@ def test_load_rust_ocr_uses_compiled_extension(monkeypatch):
fake_module = types.ModuleType("litellm.rust_bridge._native")
fake_module.ocr = lambda **kwargs: dict(FAKE_OCR_RESPONSE) # type: ignore[attr-defined]
fake_module.aocr = lambda **kwargs: dict(FAKE_OCR_RESPONSE) # type: ignore[attr-defined]
fake_module.ready_endpoints = {"ocr": {"callbacks"}} # type: ignore[attr-defined]
monkeypatch.setattr(
importlib.import_module("litellm.rust_bridge.bindings"),
"get_native_bridge",
@ -686,7 +696,7 @@ def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint():
assert bridge.calls[0]["api_base"] == "https://document-intelligence.example.com"
def test_run_rust_ocr_runs_pre_call_logging():
def test_run_rust_ocr_passes_provider_logging_adapter():
logging_obj = RecordingLogging()
bridge = RecordingBridge()
litellm.rust(True)
@ -704,17 +714,11 @@ def test_run_rust_ocr_runs_pre_call_logging():
resolve_api_key=lambda _name: None,
)
assert logging_obj.pre_call_kwargs is not None
assert logging_obj.pre_call_kwargs["input"] == "OCR document processing"
additional_args = logging_obj.pre_call_kwargs["additional_args"]
complete_input = additional_args["complete_input_dict"]
assert complete_input["document"] == DOCUMENT
assert complete_input["include_image_base64"] is True
assert additional_args["api_base"] == "https://api.mistral.ai/v1/ocr"
assert additional_args["headers"] == {
"Authorization": "Bearer sk-test",
"x-trace-id": "trace-1",
}
adapter = bridge.callback_adapter
assert isinstance(adapter, ProviderLoggingAdapter)
assert adapter.logging_obj is logging_obj
assert adapter.input == "OCR document processing"
assert adapter.api_key == "sk-test"
def test_ocr_routes_azure_ai_to_rust_when_enabled(fake_bridge):

View file

@ -1,8 +1,11 @@
from __future__ import annotations
from typing import cast
import pytest
from litellm.rust_bridge import configuration, responses_websocket
from litellm.rust_bridge.callbacks import SessionCallbackHandle
from litellm.rust_bridge.request import NativeRequestContext, NativeResponsesWebSocketRequest
@ -34,6 +37,7 @@ class _FakeNativeBridge:
*,
options: object,
context: NativeRequestContext,
callback_adapter: object | None = None,
) -> _FakeNativeConnection:
return _FakeNativeConnection()
@ -97,6 +101,38 @@ async def test_enabled_bridge_connects_and_adapts_socket(
await connection.close()
@pytest.mark.asyncio
async def test_connection_forwards_session_callback_adapter() -> None:
configuration.rust(True)
received: list[object] = []
callback_adapter = cast(SessionCallbackHandle, object())
class Native:
@classmethod
async def connect(
cls,
request: NativeResponsesWebSocketRequest,
*,
options: object,
context: NativeRequestContext,
callback_adapter: object | None = None,
) -> _FakeNativeConnection:
received.append(callback_adapter)
return _FakeNativeConnection()
responses_websocket.set_rust_responses_websocket(connection=Native)
connection = await responses_websocket.connect(
url="wss://example.test/responses",
headers={},
timeout=None,
callback_adapter=callback_adapter,
)
assert connection is not None
assert received == [callback_adapter]
class _FailingNativeBridge:
@classmethod
async def connect(
@ -105,6 +141,7 @@ class _FailingNativeBridge:
*,
options: object,
context: NativeRequestContext,
callback_adapter: object | None = None,
) -> _FakeNativeConnection:
raise RuntimeError("connection failed")
@ -131,7 +168,7 @@ async def test_connection_dispatch_cleans_up_without_reconnecting(native, sessio
class Native:
@classmethod
async def connect(cls, request, *, options, context):
async def connect(cls, request, *, options, context, callback_adapter=None):
connections.append("native")
assert options.custom_llm_provider == "azure"
return native_socket

View file

@ -0,0 +1,49 @@
import importlib.util
import os
import subprocess
import sys
import tempfile
from pathlib import Path
from typing import Final
import litellm
def main() -> None:
assert "site-packages" in Path(litellm.__file__).resolve().parts, "install the reviewed wheel first"
spec: Final = importlib.util.find_spec("litellm.rust_bridge._native")
assert spec is not None and spec.origin is not None, "wheel must contain the native extension"
native: Final = Path(spec.origin)
hidden: Final = native.with_suffix(native.suffix + ".disabled")
assert not hidden.exists()
script: Final = Path(__file__).resolve().parents[1] / "ocr" / "sdk_callback_contract.py"
environment: Final = {
**{key: value for key, value in os.environ.items() if key not in ("PYTHONPATH", "PYTHONHOME")},
"LITELLM_LOCAL_MODEL_COST_MAP": "True",
"NO_PROXY": "127.0.0.1,localhost",
}
with tempfile.TemporaryDirectory(prefix="ocr-sdk-callbacks-") as directory:
for flags in (("--unready",), ()):
subprocess.run(
[sys.executable, str(script), "--installed", *flags],
cwd=directory,
env=environment,
check=True,
timeout=180,
)
native.rename(hidden)
try:
subprocess.run(
[sys.executable, str(script), "--installed", "--without-native"],
cwd=directory,
env=environment,
check=True,
timeout=60,
)
finally:
hidden.rename(native)
sys.stdout.write("Installed-wheel callback foundation, unready-route, and unavailable-extension checks passed\n")
if __name__ == "__main__":
main()

View file

@ -111,3 +111,52 @@ def test_selectors_are_checked_by_type_checker(tmp_path: Path, expression: str,
assert [(item["rule"], item["range"]["start"]["line"]) for item in diagnostics] == [
(expected_rule, len(source.read_text().splitlines()) - 1)
]
@pytest.mark.parametrize("export", (None, 3))
def test_missing_execution_export_does_not_inspect_readiness(export):
from types import ModuleType
native = ModuleType("test_native")
lookups = []
def missing(name: str):
lookups.append(name)
raise AttributeError(name)
native.__getattr__ = missing
if export is not None:
native.messages = export
binding = bindings.NativeBinding(lambda module: module.messages, route="messages", module_loader=lambda: native)
assert binding.load() is None
assert "ready_endpoints" not in lookups
def test_discovery_reuses_one_module_and_does_not_cache_binding():
from types import ModuleType
native = ModuleType("test_native")
native.ready_endpoints = {"messages": frozenset({"callbacks"})}
def first():
return "first"
def second():
return "second"
native.messages = first
loads = []
def load():
loads.append(native)
return native
binding = bindings.NativeBinding(lambda module: module.messages, route="messages", module_loader=load)
assert binding.load() is first
assert len(loads) == 1
native.messages = second
assert binding.load() is second
assert len(loads) == 2
native.ready_endpoints = {"messages": frozenset()}
assert binding.load() is None
assert len(loads) == 3

View file

@ -0,0 +1,130 @@
from dataclasses import dataclass, field
from typing import Final
from litellm.rust_bridge.callback_adapters import ProviderLoggingAdapter, SessionCallbackAdapter
from litellm.rust_bridge.callbacks import CallbackDecision
@dataclass
class RecordingLogging:
model_call_details: dict[str, object] = field(default_factory=dict)
calls: list[tuple[str, object]] = field(default_factory=list)
def pre_call(self, *, input: object, api_key: str | None, additional_args: object) -> None:
self.calls.append(("pre", (input, api_key, additional_args)))
def post_call(self, *, original_response: str, input: object, api_key: str | None) -> None:
self.calls.append(("post", (original_response, input, api_key)))
def provider_event(**updates: object) -> dict[str, object]:
return {
"provider": "mistral",
"model": "mistral-ocr-latest",
"call_id": "call-1",
"trace_id": "trace-1",
"attempt": 1,
"started_at": 10.0,
**updates,
}
def test_provider_logging_adapter_preserves_provider_lifecycle() -> None:
logging: Final = RecordingLogging()
adapter: Final = ProviderLoggingAdapter(logging, "OCR document processing", "secret")
assert adapter.pre_call(
provider_event(request={"model": "mistral-ocr-latest"}, api_base="https://provider.test", headers={})
) == {"action": "unchanged"}
assert adapter.post_call(provider_event(response={"pages": []}, status_code=200, headers={}, ended_at=11.0)) == {
"action": "unchanged"
}
adapter.error(
provider_event(
message="retryable",
stage="provider_response",
committed=True,
status_code=429,
ended_at=11.0,
)
)
assert adapter.stream_event(provider_event(event={"type": "delta"}, sequence=1)) == {"action": "unchanged"}
adapter.stream_close(provider_event(outcome="completed", ended_at=12.0))
assert [name for name, _ in logging.calls] == ["pre", "post"]
assert logging.model_call_details["provider_stream_event"] == {
**provider_event(),
"event": {"type": "delta"},
"sequence": 1,
}
assert logging.model_call_details["provider_stream_close"] == {
**provider_event(),
"outcome": "completed",
"ended_at": 12.0,
}
assert logging.model_call_details["provider_error"] == {
**provider_event(),
"message": "retryable",
"stage": "provider_response",
"committed": True,
"status_code": 429,
"ended_at": 11.0,
}
@dataclass
class RecordingSession:
events: list[str] = field(default_factory=list)
def before_connect(self, payload: object, /) -> CallbackDecision:
self.events.append("before_connect")
return {"action": "unchanged"}
def connected(self, payload: object, /) -> None:
self.events.append("connected")
def before_send(self, payload: object, /) -> CallbackDecision:
self.events.append("before_send")
return {"action": "reject", "message": "drop frame", "status_code": None}
def after_receive(self, payload: object, /) -> CallbackDecision:
self.events.append("after_receive")
return {"action": "replace", "payload": {"type": "masked"}}
def response_complete(self, payload: object, /) -> None:
self.events.append("response_complete")
def response_error(self, payload: object, /) -> None:
self.events.append("response_error")
def error(self, payload: object, /) -> None:
self.events.append("error")
def close(self, payload: object, /) -> None:
self.events.append("close")
def test_session_adapter_preserves_frame_decisions_and_order() -> None:
callback: Final = RecordingSession()
adapter: Final = SessionCallbackAdapter(callback)
event: Final = {"session_id": "session-1", "call_id": "call-1", "event": {"type": "response.create"}}
assert adapter.before_connect(event) == {"action": "unchanged"}
adapter.connected(event)
assert adapter.before_send(event) == {"action": "reject", "message": "drop frame", "status_code": None}
assert adapter.after_receive(event) == {"action": "replace", "payload": {"type": "masked"}}
adapter.response_complete(event)
adapter.response_error(event)
adapter.error(event)
adapter.close(event)
assert callback.events == [
"before_connect",
"connected",
"before_send",
"after_receive",
"response_complete",
"response_error",
"error",
"close",
]

View file

@ -101,7 +101,7 @@ class _RecordingCall:
self.error = error
self.calls: list[dict] = []
def __call__(self, request, *, options, context):
def __call__(self, request, *, options, context, callback_adapter=None):
self.calls.append({"request": request, "options": options, "context": context})
if self.error is not None:
raise self.error
@ -109,8 +109,14 @@ class _RecordingCall:
class _RecordingAsyncCall(_RecordingCall):
async def __call__(self, request, *, options, context):
return _RecordingCall.__call__(self, request, options=options, context=context)
async def __call__(self, request, *, options, context, callback_adapter=None):
return _RecordingCall.__call__(
self,
request,
options=options,
context=context,
callback_adapter=callback_adapter,
)
def _accepts(**overrides) -> bool:

View file

@ -0,0 +1,47 @@
from __future__ import annotations
from collections.abc import Generator
from types import ModuleType
from typing import Final
import pytest
from litellm.rust_bridge import loader
@pytest.fixture(autouse=True)
def reset_loader_cache() -> Generator[None]:
loader.reset_native_bridge_cache()
yield
loader.reset_native_bridge_cache()
@pytest.mark.parametrize(
("ready_endpoints", "expected"),
(
pytest.param(None, False, id="missing-registry"),
pytest.param({"messages"}, False, id="mutable-registry"),
pytest.param(frozenset(), False, id="unregistered"),
pytest.param({"messages": frozenset()}, True, id="registered"),
),
)
def test_native_route_requires_explicit_readiness_registry(
monkeypatch: pytest.MonkeyPatch,
ready_endpoints: object,
expected: bool,
) -> None:
native: Final = ModuleType("litellm.rust_bridge._native")
if ready_endpoints is not None:
native.ready_endpoints = ready_endpoints
monkeypatch.setattr(loader, "get_native_bridge", lambda: native)
assert loader.native_route_ready("messages") is expected
def test_native_route_requires_declared_capabilities(monkeypatch: pytest.MonkeyPatch) -> None:
native: Final = ModuleType("litellm.rust_bridge._native")
native.ready_endpoints = {"messages": frozenset({"callbacks"})}
monkeypatch.setattr(loader, "get_native_bridge", lambda: native)
assert loader.native_route_ready("messages", frozenset({"callbacks"}))
assert not loader.native_route_ready("messages", frozenset({"streaming_callbacks"}))

View file

@ -7,7 +7,7 @@ from typing import Final
import pytest
from litellm.exceptions import APIError, AuthenticationError, InternalServerError, RateLimitError
from litellm.rust_bridge import bindings, runtime
from litellm.rust_bridge import bindings, loader, runtime
class RustBridgeDeclined(Exception):
@ -336,7 +336,11 @@ def test_native_endpoint_applies_partial_overrides_and_reset(monkeypatch: pytest
monkeypatch.setattr(
bindings,
"get_native_bridge",
lambda: SimpleNamespace(chat_completions=native_sync, achat_completions=native_async),
lambda: SimpleNamespace(
chat_completions=native_sync,
achat_completions=native_async,
ready_endpoints={"test": frozenset({"callbacks"})},
),
)
endpoint: Final[runtime.EndpointDispatch[object, object]] = runtime.EndpointDispatch.native(
route="test",
@ -439,7 +443,9 @@ async def test_preflight_runs_after_binding_selection_before_preparation(
assert events == (
["load", "preflight", "prepare", "native"]
if available and accepted
else ["load", "preflight", "python"] if available else ["load", "python"]
else ["load", "preflight", "python"]
if available
else ["load", "python"]
)
@ -458,3 +464,50 @@ def test_preflight_failure_is_not_a_native_decline() -> None:
error_context=context(),
preflight=preflight,
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"route",
(
"ocr",
"chat_completions",
"messages",
"responses_websocket",
"transcription",
),
)
@pytest.mark.parametrize("capabilities", (None, frozenset(), frozenset({"streaming_callbacks"})))
async def test_unready_routes_never_prepare_or_call_native(
monkeypatch: pytest.MonkeyPatch, route: str, capabilities: frozenset[str] | None
) -> None:
def unexpected(*_args: object) -> object:
pytest.fail("unready native route must not prepare or execute")
native: Final = SimpleNamespace(
ready_endpoints={} if capabilities is None else {route: capabilities},
chat_completions=unexpected,
)
monkeypatch.setattr(loader, "get_native_bridge", lambda: native)
monkeypatch.setattr(bindings, "get_native_bridge", lambda: native)
endpoint: Final = runtime.EndpointBinding.native(
route=route, select=lambda native: native.chat_completions, enabled=runtime.always_enabled
)
arguments: Final = {
"prepare": unexpected,
"preflight": unexpected,
"call": unexpected,
"adapt": unexpected,
"error_context": runtime.BridgeErrorContext(provider="test", model="test-model"),
}
assert not endpoint.can_attempt()
assert endpoint.invoke(**arguments, fallback=lambda: "python") == "python"
with pytest.raises(RuntimeError, match=f"native {route} endpoint is unavailable"):
endpoint.require(**arguments)
async def fallback() -> str:
return "python"
assert await endpoint.ainvoke(**arguments, fallback=fallback) == "python"
with pytest.raises(RuntimeError, match=f"native {route} endpoint is unavailable"):
await endpoint.arequire(**arguments)

View file

@ -34,6 +34,7 @@ class SyncBridge:
*,
options: NativeRequestOptions,
context: NativeRequestContext,
callback_adapter: object | None = None,
) -> dict[str, object]:
self.calls.append(
{
@ -53,6 +54,7 @@ class AsyncBridge:
*,
options: NativeRequestOptions,
context: NativeRequestContext,
callback_adapter: object | None = None,
) -> dict[str, object]:
return {"text": "async"}
@ -135,7 +137,7 @@ async def test_dispatch_async_path_requires_bridge(monkeypatch: pytest.MonkeyPat
def test_bedrock_transcription_uses_rust_only_path() -> None:
rust_bridge.configure_rust_transcription(
transcription=lambda request, *, options, context: {"text": "rust"},
transcription=lambda request, *, options, context, callback_adapter=None: {"text": "rust"},
atranscription=None,
)
try:
@ -152,7 +154,11 @@ def test_bedrock_transcription_uses_rust_only_path() -> None:
@pytest.mark.asyncio
async def test_bedrock_atranscription_uses_rust_only_path() -> None:
async def rust_response(
request: NativeTranscriptionRequest, *, options: object, context: NativeRequestContext
request: NativeTranscriptionRequest,
*,
options: object,
context: NativeRequestContext,
callback_adapter: object | None = None,
) -> dict[str, object]:
return {"text": "rust"}