diff --git a/.github/workflows/test-rust.yml b/.github/workflows/test-rust.yml index 9b8b132df62..799812fd03e 100644 --- a/.github/workflows/test-rust.yml +++ b/.github/workflows/test-rust.yml @@ -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 diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index a72a55d8239..00b9c64ea1e 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -1482,10 +1482,12 @@ name = "litellm-python-interop" version = "0.1.0" dependencies = [ "pyo3", + "pyo3-async-runtimes", "pythonize", "rstest", "serde", "serde_json", + "tokio", ] [[package]] diff --git a/litellm-rust/README.md b/litellm-rust/README.md index c76b33ef5a2..b554c6f8f8c 100644 --- a/litellm-rust/README.md +++ b/litellm-rust/README.md @@ -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. diff --git a/litellm-rust/crates/ai-gateway/src/io/ocr.rs b/litellm-rust/crates/ai-gateway/src/io/ocr.rs index ce8e1abc470..2de6ca68ac7 100644 --- a/litellm-rust/crates/ai-gateway/src/io/ocr.rs +++ b/litellm-rust/crates/ai-gateway/src/io/ocr.rs @@ -1 +1 @@ -pub use crate::ocr::{OcrRequest, ocr, ocr_provider_supported}; +pub use crate::ocr::{OcrRequest, ocr, ocr_provider_supported, ocr_with_observer}; diff --git a/litellm-rust/crates/ai-gateway/src/ocr/handler.rs b/litellm-rust/crates/ai-gateway/src/ocr/handler.rs index 6c6e12724cd..0aac26996e1 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/handler.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/handler.rs @@ -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( request: PreparedOcrRequest, + context: &CallLifecycleContext, hooks: &OcrLifecycleHooks, -) -> Result { + observer: &mut Observer, +) -> Result +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 diff --git a/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs b/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs index ff9b47b4ccc..6e710020f5d 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs @@ -127,6 +127,7 @@ impl OcrLifecycleHooks { }; Ok(ProviderOcrRequest { model, + custom_llm_provider, config, url, body, diff --git a/litellm-rust/crates/ai-gateway/src/ocr/mod.rs b/litellm-rust/crates/ai-gateway/src/ocr/mod.rs index 4820c393bd7..07f2a5bed29 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/mod.rs @@ -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 { + 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( + request: OcrRequest<'_>, + options: &RequestOptions, + context: &LiteLlmRequestContext, + hooks: RequestHooks, + observer: &mut Observer, +) -> Result +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 } diff --git a/litellm-rust/crates/ai-gateway/src/ocr/types.rs b/litellm-rust/crates/ai-gateway/src/ocr/types.rs index 9be628a2c81..4e26522de6b 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/types.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/types.rs @@ -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, diff --git a/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs b/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs index b9dfe8dc4f8..40f6c10e5af 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs @@ -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() }, diff --git a/litellm-rust/crates/ai-gateway/tests/ocr_lifecycle.rs b/litellm-rust/crates/ai-gateway/tests/ocr_lifecycle.rs index f68ac32333a..d3836a8c8c6 100644 --- a/litellm-rust/crates/ai-gateway/tests/ocr_lifecycle.rs +++ b/litellm-rust/crates/ai-gateway/tests/ocr_lifecycle.rs @@ -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>>, + raw_response: Option, + 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 { + 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 { + 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]; diff --git a/litellm-rust/crates/core/src/hook_contracts.rs b/litellm-rust/crates/core/src/hook_contracts.rs new file mode 100644 index 00000000000..7fcde3d2d51 --- /dev/null +++ b/litellm-rust/crates/core/src/hook_contracts.rs @@ -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; + )* + } + }; +} diff --git a/litellm-rust/crates/core/src/lib.rs b/litellm-rust/crates/core/src/lib.rs index daf6ea89b79..5dc26dd13e1 100644 --- a/litellm-rust/crates/core/src/lib.rs +++ b/litellm-rust/crates/core/src/lib.rs @@ -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; diff --git a/litellm-rust/crates/core/src/provider_callbacks/handler.rs b/litellm-rust/crates/core/src/provider_callbacks/handler.rs new file mode 100644 index 00000000000..08706a244d2 --- /dev/null +++ b/litellm-rust/crates/core/src/provider_callbacks/handler.rs @@ -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, + pub api_base: String, + pub headers: BTreeMap, +} + +pub struct ProviderAttemptContext { + pub call_id: String, + pub trace_id: Option, + pub attempt: u32, +} + +#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] +pub async fn send_provider_request( + request: RequestBuilder, + input: ProviderRequest, + context: ProviderAttemptContext, + observer: &mut Observer, +) -> Result +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: &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 { + 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 { + 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 { + 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"]); + } + } + } +} diff --git a/litellm-rust/crates/core/src/provider_callbacks/mod.rs b/litellm-rust/crates/core/src/provider_callbacks/mod.rs new file mode 100644 index 00000000000..f13e70da2ec --- /dev/null +++ b/litellm-rust/crates/core/src/provider_callbacks/mod.rs @@ -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, + }, +} + +#[derive(Clone, Serialize)] +pub struct ProviderPreCall { + pub provider: String, + pub model: String, + pub call_id: String, + pub trace_id: Option, + pub attempt: u32, + pub started_at: f64, + pub request: BTreeMap, + pub api_base: String, + pub headers: BTreeMap, +} + +#[derive(Serialize)] +pub struct ProviderPostCall { + pub provider: String, + pub model: String, + pub call_id: String, + pub trace_id: Option, + pub attempt: u32, + pub started_at: f64, + pub response: Value, + pub status_code: u16, + pub headers: BTreeMap, + pub ended_at: f64, +} + +#[derive(Serialize)] +pub struct ProviderError { + pub provider: String, + pub model: String, + pub call_id: String, + pub trace_id: Option, + pub attempt: u32, + pub started_at: f64, + pub message: String, + pub stage: &'static str, + pub committed: bool, + pub status_code: Option, + pub ended_at: f64, +} + +#[derive(Serialize)] +pub struct ProviderStreamEvent { + pub provider: String, + pub model: String, + pub call_id: String, + pub trace_id: Option, + 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, + 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, + pub event: Option, + pub response_id: Option, + pub sequence: Option, + pub message: Option, +} + +#[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 { + Ok(CallbackDecision::Unchanged) + } + + async fn post_call( + &mut self, + _input: &ProviderPostCall, + ) -> Result { + 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::(json!({"action": "unchanged"})).unwrap(), + CallbackDecision::Unchanged + ); + assert_eq!( + serde_json::from_value::( + json!({"action": "replace", "payload": {"masked": true}}) + ) + .unwrap(), + CallbackDecision::Replace { + payload: json!({"masked": true}) + } + ); + assert_eq!( + serde_json::from_value::(json!({ + "action": "reject", + "message": "blocked", + "status_code": 400 + })) + .unwrap(), + CallbackDecision::Reject { + message: "blocked".to_string(), + status_code: Some(400) + } + ); + } +} diff --git a/litellm-rust/crates/core/src/providers/azure_ai/ocr/transformation.rs b/litellm-rust/crates/core/src/providers/azure_ai/ocr/transformation.rs index d15c032f0bc..3cecbbbab19 100644 --- a/litellm-rust/crates/core/src/providers/azure_ai/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/providers/azure_ai/ocr/transformation.rs @@ -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, diff --git a/litellm-rust/crates/python-bridge/src/callback_bindings.rs b/litellm-rust/crates/python-bridge/src/callback_bindings.rs new file mode 100644 index 00000000000..f7180d84823 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/callback_bindings.rs @@ -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::::new; + let _session_constructor = PythonSession::::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), + Async(PythonProviderSession), +} + +impl PythonProviderObserver { + pub(crate) fn new( + adapter: Option>, + context: PythonCallContext<'_>, + ) -> PyResult { + 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::>()? + .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, + py: Python<'_>, +) -> PyResult> { + let module = py.import("litellm.rust_bridge._native")?; + let runtime = module + .getattr("__python_callback_runtime__")? + .extract::>()? + .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 { + 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 { + 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, + } + } +} diff --git a/litellm-rust/crates/python-bridge/src/constants.rs b/litellm-rust/crates/python-bridge/src/constants.rs new file mode 100644 index 00000000000..9308ba185ca --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/constants.rs @@ -0,0 +1 @@ +pub(crate) const OCR_CALLBACK_CAPACITY: usize = 1024; diff --git a/litellm-rust/crates/python-bridge/src/execution.rs b/litellm-rust/crates/python-bridge/src/execution.rs index f3648158cf6..cc037a124ae 100644 --- a/litellm-rust/crates/python-bridge/src/execution.rs +++ b/litellm-rust/crates/python-bridge/src/execution.rs @@ -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( +pub(crate) struct PythonCallContext<'py> { + pub(crate) py: Python<'py>, + pub(crate) asynchronous: bool, +} + +pub(crate) fn run_sync( py: Python<'_>, - future: F, - map_error: fn(Error) -> PyErr, + future: impl Future> + Send + 'static, + map_error: fn(E) -> PyErr, ) -> PyResult> where T: Serialize + Send + 'static, - F: Future> + Send + 'static, + E: Send + 'static, { run_sync_on( py, @@ -28,15 +32,15 @@ where ) } -fn run_sync_on( +fn run_sync_on( py: Python<'_>, runtime: &Runtime, - future: F, - map_error: fn(Error) -> PyErr, + future: impl Future> + Send + 'static, + map_error: fn(E) -> PyErr, ) -> PyResult> where T: Serialize + Send + 'static, - F: Future> + 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( +pub(crate) fn run_async( py: Python<'_>, - future: F, - map_error: fn(Error) -> PyErr, + future: impl Future> + Send + 'static, + map_error: fn(E) -> PyErr, ) -> PyResult> where T: Serialize + Send + 'static, - F: Future> + 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(result: Result, map_error: fn(Error) -> PyErr) -> PyResult { +fn map_result(result: Result, map_error: fn(E) -> PyErr) -> PyResult { match result { Ok(value) => Ok(value), Err(error) => Err( @@ -75,9 +79,9 @@ fn map_core_result(result: Result, map_error: fn(Error) -> PyErr) - } } -async fn catch_future_panic(future: F) -> PyResult> +async fn catch_future_panic(future: F) -> PyResult> where - F: Future>, + F: Future>, { AssertUnwindSafe(future) .catch_unwind() @@ -85,9 +89,9 @@ where .map_err(panic_to_pyerr) } -async fn wait_for_sync_result(future: F) -> PyResult> +async fn wait_for_sync_result(future: F) -> PyResult> where - F: Future>, + F: Future>, { 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| { diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 5bca907a3cb..aeb05f8b343 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -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>, ) -> PyResult> { 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) -> 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::()?; 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"); diff --git a/litellm-rust/crates/python-bridge/src/python_hook_bindings.rs b/litellm-rust/crates/python-bridge/src/python_hook_bindings.rs new file mode 100644 index 00000000000..915572895d9 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/python_hook_bindings.rs @@ -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 { + context: C, + $( + $method: ::litellm_python_interop::callback_runtime::Callback< + $input, $output, $crate::callback_return_mode!($mode), + >, + )* + } + + impl $session + 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 { + use ::pyo3::types::PyAnyMethods as _; + Ok(Self { + context, + $( + $method: ::litellm_python_interop::callback_runtime::Callback::new( + adapter.getattr(stringify!($method))?, + )?, + )* + }) + } + } + + impl $hooks for $session + 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 + } + )* + } + }; +} diff --git a/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs b/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs index 068e7f66421..f59f5cfd761 100644 --- a/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs +++ b/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs @@ -21,6 +21,8 @@ fn prepare_transcription( input: AudioTranscriptionInputs, options: NativeRequestOptions, context: NativeRequestContext, + _callback_adapter: Option>, + _python_context: crate::execution::PythonCallContext<'_>, ) -> PyResult> + Send + 'static> { let provider_supported = litellm_core::audio_transcription::transcription_provider_supported( options.provider("bedrock"), diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs index f1a13b2ee46..880d3a13f81 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs @@ -23,6 +23,8 @@ fn prepare_chat_completions( input: ChatCompletionsInputs, options: NativeRequestOptions, context: NativeRequestContext, + _callback_adapter: Option>, + _python_context: crate::execution::PythonCallContext<'_>, ) -> PyResult> + Send + 'static> { let context: LiteLlmRequestContext = context.into(); let messages = required_value("messages", input.messages, Value::is_array, "list")?; diff --git a/litellm-rust/crates/python-bridge/src/routes/definition.rs b/litellm-rust/crates/python-bridge/src/routes/definition.rs index c2350845cbd..c9ff25d8965 100644 --- a/litellm-rust/crates/python-bridge/src/routes/definition.rs +++ b/litellm-rust/crates/python-bridge/src/routes/definition.rs @@ -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::PyResult> { - 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::PyResult> { - 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::PyResult> { - 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::PyResult> { - 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>, + _python_context: crate::execution::PythonCallContext<'_>, ) -> PyResult> + 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)", ), ]; diff --git a/litellm-rust/crates/python-bridge/src/routes/messages.rs b/litellm-rust/crates/python-bridge/src/routes/messages.rs index 2ac4f222a86..bf6b2cc382d 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages.rs @@ -19,6 +19,8 @@ fn prepare_messages( input: MessagesInputs, options: NativeRequestOptions, context: NativeRequestContext, + _callback_adapter: Option>, + _python_context: crate::execution::PythonCallContext<'_>, ) -> PyResult> + Send + 'static> { let provider_supported = litellm_core::messages::messages_provider_supported(options.provider("anthropic")); diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr.rs b/litellm-rust/crates/python-bridge/src/routes/ocr.rs index 1ed77211c00..be95d4db55d 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr.rs @@ -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>, + python_context: crate::execution::PythonCallContext<'_>, ) -> PyResult> + 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 }) diff --git a/litellm-rust/crates/python-bridge/tests/callbacks/mod.rs b/litellm-rust/crates/python-bridge/tests/callbacks/mod.rs new file mode 100644 index 00000000000..ebed1db4b10 --- /dev/null +++ b/litellm-rust/crates/python-bridge/tests/callbacks/mod.rs @@ -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 { + Hook(E), + Provider { + error: ProviderFailure, + observer_error: Option, + }, + } + + struct PreparedCall(BeforeSend); + struct ReadyCall(Request); + + impl PreparedCall { + async fn finish_hooks( + self, + hooks: &mut H, + ) -> Result { + hooks.pre_api_call(&self.0).await?; + Ok(ReadyCall(self.0.body)) + } + } + + impl ReadyCall { + fn send(self, calls: &AtomicUsize, fail: bool) -> Result { + 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( + hooks: &mut H, + calls: &AtomicUsize, + fail: bool, + ) -> Result> { + 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 { + 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, +} + +#[pymethods] +impl Harness { + #[pyo3(signature = (adapter, fail=false))] + fn execute<'py>( + &self, + py: Python<'py>, + adapter: &Bound<'py, PyAny>, + fail: bool, + ) -> PyResult> { + 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> { + 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> { + 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> { + 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> { + 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, +} diff --git a/litellm-rust/crates/python-bridge/tests/callbacks/test_callbacks.py b/litellm-rust/crates/python-bridge/tests/callbacks/test_callbacks.py new file mode 100644 index 00000000000..5783468b382 --- /dev/null +++ b/litellm-rust/crates/python-bridge/tests/callbacks/test_callbacks.py @@ -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)) diff --git a/litellm-rust/crates/python-interop/Cargo.toml b/litellm-rust/crates/python-interop/Cargo.toml index 9da6af6e2e2..f5201938faa 100644 --- a/litellm-rust/crates/python-interop/Cargo.toml +++ b/litellm-rust/crates/python-interop/Cargo.toml @@ -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 diff --git a/litellm-rust/crates/python-interop/src/callback_runtime/invoke.py b/litellm-rust/crates/python-interop/src/callback_runtime/invoke.py new file mode 100644 index 00000000000..e4453478d00 --- /dev/null +++ b/litellm-rust/crates/python-interop/src/callback_runtime/invoke.py @@ -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() diff --git a/litellm-rust/crates/python-interop/src/callback_runtime/mod.rs b/litellm-rust/crates/python-interop/src/callback_runtime/mod.rs new file mode 100644 index 00000000000..96144e9e37b --- /dev/null +++ b/litellm-rust/crates/python-interop/src/callback_runtime/mod.rs @@ -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: Send { + fn invoke( + &mut self, + callable: Py, + payload: Py, + ) -> impl Future>> + Send; +} + +pub struct Callback { + callable: Py, + signature: PhantomData (O, M)>, +} + +impl Callback +where + I: Serialize + Sync, + O: DeserializeOwned + Send, + M: ReturnMode, +{ + pub fn new(callable: Bound<'_, PyAny>) -> PyResult { + 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>(&mut self, context: &mut C, input: &I) -> PyResult { + 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, + direct: Py, + capacity: Arc, +} + +#[derive(Clone)] +pub struct CallbackRuntime(Arc); + +impl CallbackRuntime { + pub fn new(module: &Bound<'_, PyModule>, max_in_flight: NonZeroUsize) -> PyResult { + 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 { + Ok(AsyncContext { + runtime: self.clone(), + locals: pyo3_async_runtimes::tokio::get_current_locals(py)?, + interrupted: false, + }) + } + + pub fn sync_context(&self, py: Python<'_>) -> PyResult { + Ok(SyncContext { + runtime: self.clone(), + context: py + .import("contextvars")? + .call_method0("copy_context")? + .unbind(), + caller: thread::current().id(), + }) + } + + fn admit(&self) -> PyResult { + Arc::clone(&self.0.capacity) + .try_acquire_owned() + .map_err(|_| PyRuntimeError::new_err("callback capacity exhausted")) + } +} + +pub struct SyncContext { + runtime: CallbackRuntime, + context: Py, + caller: ThreadId, +} + +impl CallbackContext for SyncContext { + async fn invoke(&mut self, callable: Py, payload: Py) -> PyResult> { + 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 CallbackContext for AsyncContext { + async fn invoke(&mut self, callable: Py, payload: Py) -> PyResult> { + 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, + invocation: Option>, +} + +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))); + } + }); + } +} diff --git a/litellm-rust/crates/python-interop/src/lib.rs b/litellm-rust/crates/python-interop/src/lib.rs index 2e562bdae70..c75dbf077ab 100644 --- a/litellm-rust/crates/python-interop/src/lib.rs +++ b/litellm-rust/crates/python-interop/src/lib.rs @@ -1,3 +1,4 @@ +pub mod callback_runtime; mod gil; mod marshal; diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index 65f6328aafe..695b2420f52 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -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, + ), ) diff --git a/litellm/rust_bridge/bindings.py b/litellm/rust_bridge/bindings.py index ab7b7c296b9..b6c04596635 100644 --- a/litellm/rust_bridge/bindings.py +++ b/litellm/rust_bridge/bindings.py @@ -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) diff --git a/litellm/rust_bridge/callback_adapters.py b/litellm/rust_bridge/callback_adapters.py new file mode 100644 index 00000000000..1e08404e14c --- /dev/null +++ b/litellm/rust_bridge/callback_adapters.py @@ -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()) diff --git a/litellm/rust_bridge/callbacks.py b/litellm/rust_bridge/callbacks.py new file mode 100644 index 00000000000..b17edabf8b4 --- /dev/null +++ b/litellm/rust_bridge/callbacks.py @@ -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: ... diff --git a/litellm/rust_bridge/loader.py b/litellm/rust_bridge/loader.py index 022c38f5a85..16e298fe230 100644 --- a/litellm/rust_bridge/loader.py +++ b/litellm/rust_bridge/loader.py @@ -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) diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index 9512f5ed00a..93245601362 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -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]], diff --git a/litellm/rust_bridge/ocr_callbacks.py b/litellm/rust_bridge/ocr_callbacks.py new file mode 100644 index 00000000000..4a191694593 --- /dev/null +++ b/litellm/rust_bridge/ocr_callbacks.py @@ -0,0 +1,3 @@ +from .callback_adapters import PreCallArguments, ProviderLoggingAdapter + +__all__ = ("PreCallArguments", "ProviderLoggingAdapter") diff --git a/litellm/rust_bridge/protocols.py b/litellm/rust_bridge/protocols.py index 07489679671..4eb522d1707 100644 --- a/litellm/rust_bridge/protocols.py +++ b/litellm/rust_bridge/protocols.py @@ -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: ... diff --git a/litellm/rust_bridge/request.py b/litellm/rust_bridge/request.py index cb23583cf17..27c72023ce9 100644 --- a/litellm/rust_bridge/request.py +++ b/litellm/rust_bridge/request.py @@ -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) diff --git a/litellm/rust_bridge/responses_websocket.py b/litellm/rust_bridge/responses_websocket.py index 8b7d4fe6a23..80cef256428 100644 --- a/litellm/rust_bridge/responses_websocket.py +++ b/litellm/rust_bridge/responses_websocket.py @@ -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, ) diff --git a/litellm/rust_bridge/runtime.py b/litellm/rust_bridge/runtime.py index 0efd4f2e29b..9dcdcb311d1 100644 --- a/litellm/rust_bridge/runtime.py +++ b/litellm/rust_bridge/runtime.py @@ -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, *, diff --git a/litellm/rust_bridge/transcription.py b/litellm/rust_bridge/transcription.py index 3aa5fddea1a..4590e3df775 100644 --- a/litellm/rust_bridge/transcription.py +++ b/litellm/rust_bridge/transcription.py @@ -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, diff --git a/tests/rust-python-harness/strategies/trace_parity/sdk/execution.py b/tests/rust-python-harness/strategies/trace_parity/sdk/execution.py index 5bd0358eeb2..d9a821ccb77 100644 --- a/tests/rust-python-harness/strategies/trace_parity/sdk/execution.py +++ b/tests/rust-python-harness/strategies/trace_parity/sdk/execution.py @@ -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)) diff --git a/tests/rust-python-harness/strategies/trace_parity/sdk/ocr/case.py b/tests/rust-python-harness/strategies/trace_parity/sdk/ocr/case.py index fe214f45339..ea1f9e9b669 100644 --- a/tests/rust-python-harness/strategies/trace_parity/sdk/ocr/case.py +++ b/tests/rust-python-harness/strategies/trace_parity/sdk/ocr/case.py @@ -16,6 +16,8 @@ COMMON_MAPPINGS: Final = ( mapping(rust_span="validate_environment", python_frame=r"(? 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") diff --git a/tests/test_litellm/ocr/callback_support.py b/tests/test_litellm/ocr/callback_support.py new file mode 100644 index 00000000000..cae1638be14 --- /dev/null +++ b/tests/test_litellm/ocr/callback_support.py @@ -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 diff --git a/tests/test_litellm/ocr/live_callback_smoke.py b/tests/test_litellm/ocr/live_callback_smoke.py new file mode 100644 index 00000000000..ffa8ce1b114 --- /dev/null +++ b/tests/test_litellm/ocr/live_callback_smoke.py @@ -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)) diff --git a/tests/test_litellm/ocr/sdk_callback_contract.py b/tests/test_litellm/ocr/sdk_callback_contract.py new file mode 100644 index 00000000000..e5e02de17aa --- /dev/null +++ b/tests/test_litellm/ocr/sdk_callback_contract.py @@ -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() + ) diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index cbafa9a5a01..3dd5f672ba9 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -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): diff --git a/tests/test_litellm/responses/test_rust_bridge_websocket.py b/tests/test_litellm/responses/test_rust_bridge_websocket.py index 93d56f012c8..1ec1a8a9621 100644 --- a/tests/test_litellm/responses/test_rust_bridge_websocket.py +++ b/tests/test_litellm/responses/test_rust_bridge_websocket.py @@ -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 diff --git a/tests/test_litellm/rust_bridge/sdk_callback_wheel_test.py b/tests/test_litellm/rust_bridge/sdk_callback_wheel_test.py new file mode 100644 index 00000000000..f007e67e8a1 --- /dev/null +++ b/tests/test_litellm/rust_bridge/sdk_callback_wheel_test.py @@ -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() diff --git a/tests/test_litellm/rust_bridge/test_bindings.py b/tests/test_litellm/rust_bridge/test_bindings.py index cd562ee91c1..704aec1e776 100644 --- a/tests/test_litellm/rust_bridge/test_bindings.py +++ b/tests/test_litellm/rust_bridge/test_bindings.py @@ -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 diff --git a/tests/test_litellm/rust_bridge/test_callback_adapters.py b/tests/test_litellm/rust_bridge/test_callback_adapters.py new file mode 100644 index 00000000000..f5cdb6cf128 --- /dev/null +++ b/tests/test_litellm/rust_bridge/test_callback_adapters.py @@ -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", + ] diff --git a/tests/test_litellm/rust_bridge/test_chat_completions.py b/tests/test_litellm/rust_bridge/test_chat_completions.py index 85e3e050a46..7999be282ea 100644 --- a/tests/test_litellm/rust_bridge/test_chat_completions.py +++ b/tests/test_litellm/rust_bridge/test_chat_completions.py @@ -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: diff --git a/tests/test_litellm/rust_bridge/test_loader.py b/tests/test_litellm/rust_bridge/test_loader.py new file mode 100644 index 00000000000..327b2accebd --- /dev/null +++ b/tests/test_litellm/rust_bridge/test_loader.py @@ -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"})) diff --git a/tests/test_litellm/rust_bridge/test_runtime.py b/tests/test_litellm/rust_bridge/test_runtime.py index 4b6e5c1290b..2342977ad2a 100644 --- a/tests/test_litellm/rust_bridge/test_runtime.py +++ b/tests/test_litellm/rust_bridge/test_runtime.py @@ -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) diff --git a/tests/test_litellm/test_audio_transcription_rust_bridge.py b/tests/test_litellm/test_audio_transcription_rust_bridge.py index 85303fed5ce..87ed412d379 100644 --- a/tests/test_litellm/test_audio_transcription_rust_bridge.py +++ b/tests/test_litellm/test_audio_transcription_rust_bridge.py @@ -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"}