diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 00b9c64ea1e..3c07435dc23 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -1446,6 +1446,8 @@ dependencies = [ "aws-smithy-runtime-api", "aws-types", "base64", + "futures-channel", + "futures-util", "rand 0.8.7", "reqwest", "rstest", @@ -1454,6 +1456,7 @@ dependencies = [ "sha2 0.10.9", "thiserror 2.0.19", "tokio", + "tokio-tungstenite", "tracing", "tracing-subscriber", ] @@ -1464,7 +1467,6 @@ version = "0.1.0" dependencies = [ "criterion", "futures-util", - "litellm-ai-gateway", "litellm-core", "litellm-python-interop", "pyo3", diff --git a/litellm-rust/crates/ai-gateway/src/client.rs b/litellm-rust/crates/ai-gateway/src/client.rs deleted file mode 100644 index ff2606f0229..00000000000 --- a/litellm-rust/crates/ai-gateway/src/client.rs +++ /dev/null @@ -1,14 +0,0 @@ -use std::sync::OnceLock; -use std::time::Duration; - -const HTTP_CLIENT_TIMEOUT_SECS: u64 = 600; - -pub(crate) fn http_client() -> &'static reqwest::Client { - static CLIENT: OnceLock = OnceLock::new(); - CLIENT.get_or_init(|| { - reqwest::Client::builder() - .timeout(Duration::from_secs(HTTP_CLIENT_TIMEOUT_SECS)) - .build() - .expect("failed to build reqwest client") - }) -} diff --git a/litellm-rust/crates/ai-gateway/src/constants.rs b/litellm-rust/crates/ai-gateway/src/constants.rs index 78af374bf70..ca79e2127d3 100644 --- a/litellm-rust/crates/ai-gateway/src/constants.rs +++ b/litellm-rust/crates/ai-gateway/src/constants.rs @@ -29,9 +29,6 @@ pub(crate) const DEFAULT_FLUSH_INTERVAL_MS: u64 = 500; #[cfg(feature = "server")] pub(crate) const DEFAULT_PROVIDER: &str = "openai"; -pub(crate) const DEFAULT_RESPONSES_WS_CONNECT_TIMEOUT_SECS: u64 = 10; -pub(crate) const DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS: u64 = 300; - /// HTTP path for the non-streaming Anthropic Messages route. #[cfg(feature = "server")] pub(crate) const MESSAGES_ROUTE_PATH: &str = "/v1/messages"; diff --git a/litellm-rust/crates/ai-gateway/src/io/responses_ws.rs b/litellm-rust/crates/ai-gateway/src/io/responses_ws.rs index ee907a17869..f5bbd7112b6 100644 --- a/litellm-rust/crates/ai-gateway/src/io/responses_ws.rs +++ b/litellm-rust/crates/ai-gateway/src/io/responses_ws.rs @@ -1,557 +1,9 @@ -use std::sync::Arc; -use std::time::Duration; +//! Compatibility re-exports: the Responses WebSocket transport moved to +//! `litellm-core` (`litellm_core::responses::connection`) so the python bridge +//! can use it without depending on this crate. Re-exported here so existing +//! gateway imports keep working. -use futures_util::stream::{SplitSink, SplitStream}; -use futures_util::{Sink, SinkExt, Stream, StreamExt}; -use litellm_core::Error; -use litellm_core::http_utils::string_headers; -use litellm_core::providers::openai::responses::transformation::OPENAI_RESPONSES_WS_CONFIG; -use litellm_core::request_context::LiteLlmRequestContext; -use litellm_core::request_options::RequestOptions; -use litellm_core::responses::types::{ResponsesWebSocketRequest, ResponsesWsEvent}; -use litellm_core::responses::websocket::ResponsesWebSocketProviderConfig; -use tokio::net::TcpStream; -use tokio::sync::Mutex; -use tokio_tungstenite::tungstenite::Message; -use tokio_tungstenite::tungstenite::client::IntoClientRequest; -use tokio_tungstenite::tungstenite::http::HeaderValue; -use tokio_tungstenite::tungstenite::http::header::{AUTHORIZATION, HeaderName}; -use tokio_tungstenite::{MaybeTlsStream, WebSocketStream, connect_async}; - -use crate::constants::{ - DEFAULT_RESPONSES_WS_CONNECT_TIMEOUT_SECS, DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS, +pub use litellm_core::responses::connection::{ + ResponsesUpstreamWs, ResponsesWebSocketConnection, ResponsesWebSocketStreaming, + async_responses_websocket, responses_ws, }; - -const OPENAI_API_KEY_ENV: &str = "OPENAI_API_KEY"; -const MISSING_KEY_MESSAGE: &str = "Missing OpenAI API Key - a Responses WebSocket call is being made but no key was passed via params or the OPENAI_API_KEY environment variable"; - -pub type ResponsesUpstreamWs = WebSocketStream>; -type UpstreamTx = SplitSink; -type UpstreamRx = SplitStream; - -#[derive(Clone)] -pub struct ResponsesWebSocketConnection { - socket: Arc>>, -} - -impl ResponsesWebSocketConnection { - pub async fn connect( - input: ResponsesWebSocketRequest, - options: &RequestOptions, - _context: &LiteLlmRequestContext, - ) -> Result { - if !litellm_core::responses::websocket::native_websocket_supported( - options.custom_llm_provider.as_deref().unwrap_or("openai"), - ) { - return Err(Error::Unsupported("unsupported native WebSocket provider")); - } - let headers = string_headers("Responses WebSocket", options.extra_headers.clone())?; - let mut request = input - .url - .as_str() - .into_client_request() - .map_err(|error| Error::Network(error.to_string()))?; - for (name, value) in headers { - let header_name = name - .parse::() - .map_err(|error| Error::InvalidRequest(error.to_string()))?; - let header_value = HeaderValue::from_str(&value) - .map_err(|error| Error::InvalidRequest(error.to_string()))?; - request.headers_mut().insert(header_name, header_value); - } - let connect = connect_async(request); - let result = match options.timeout { - Some(timeout) => tokio::time::timeout(timeout, connect).await.map_err(|_| { - Error::Network("Responses WebSocket connection timed out".to_string()) - })?, - None => connect.await, - }; - let (socket, _) = result.map_err(|error| match error { - tokio_tungstenite::tungstenite::Error::Http(response) => Error::Http { - status: response.status().as_u16(), - body: String::new(), - }, - other => Error::Network(other.to_string()), - })?; - Ok(Self { - socket: Arc::new(Mutex::new(Some(socket))), - }) - } - - pub async fn send_text(&self, text: String) -> Result<(), Error> { - let mut socket = self.socket.lock().await; - let Some(socket) = socket.as_mut() else { - return Err(Error::Network("Responses WebSocket is closed".to_string())); - }; - socket - .send(Message::Text(text)) - .await - .map_err(|error| Error::Network(error.to_string())) - } - - pub async fn recv_text(&self) -> Result, Error> { - let mut socket_guard = self.socket.lock().await; - let Some(socket) = socket_guard.as_mut() else { - return Ok(None); - }; - match socket.next().await { - Some(Ok(Message::Text(text))) => Ok(Some(text)), - Some(Ok(Message::Binary(bytes))) => String::from_utf8(bytes.to_vec()) - .map(Some) - .map_err(|error| Error::InvalidResponse(error.to_string())), - Some(Ok(Message::Close(_))) | None => Ok(None), - Some(Ok(_)) => Ok(None), - Some(Err(error)) => Err(Error::Network(error.to_string())), - } - } - - pub async fn close(&self) -> Result<(), Error> { - let mut socket = self.socket.lock().await; - if let Some(socket) = socket.as_mut() { - socket - .close(None) - .await - .map_err(|error| Error::Network(error.to_string()))?; - } - *socket = None; - Ok(()) - } -} - -pub(crate) fn resolve_api_key(api_key: Option<&str>) -> Result { - api_key - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(str::to_string) - .or_else(|| { - std::env::var(OPENAI_API_KEY_ENV) - .ok() - .filter(|value| !value.trim().is_empty()) - }) - .ok_or_else(|| Error::Auth(MISSING_KEY_MESSAGE.to_string())) -} - -async fn dial_upstream( - model: &str, - api_key: &str, - api_base: Option<&str>, -) -> Result { - let url = OPENAI_RESPONSES_WS_CONFIG.complete_websocket_url(api_base, model); - let mut request = url - .as_str() - .into_client_request() - .map_err(|error| Error::Network(error.to_string()))?; - request.headers_mut().insert( - AUTHORIZATION, - HeaderValue::from_str(&format!("Bearer {api_key}")) - .map_err(|error| Error::Auth(error.to_string()))?, - ); - let result = tokio::time::timeout( - Duration::from_secs(DEFAULT_RESPONSES_WS_CONNECT_TIMEOUT_SECS), - connect_async(request), - ) - .await - .map_err(|_| Error::Network("Responses WebSocket connection timed out".to_string()))?; - result - .map(|(socket, _)| socket) - .map_err(|error| match error { - tokio_tungstenite::tungstenite::Error::Http(response) => Error::Http { - status: response.status().as_u16(), - body: String::new(), - }, - other => Error::Network(other.to_string()), - }) -} - -pub struct ResponsesWebSocketStreaming; - -impl ResponsesWebSocketStreaming { - pub async fn bidirectional_forward( - model: &str, - upstream_tx: UpstreamTx, - upstream_rx: UpstreamRx, - idle_timeout: Option, - observe: impl FnMut(&ResponsesWsEvent) + Send, - client_in: In, - client_out: Out, - ) -> Result<(), Error> - where - In: Stream + Unpin + Send, - Out: Sink + Unpin + Send, - Out::Error: std::fmt::Display, - { - splice( - model, - upstream_tx, - upstream_rx, - idle_timeout, - observe, - client_in, - client_out, - ) - .await - } -} - -pub(crate) async fn splice( - model: &str, - mut upstream_tx: UpstreamTx, - mut upstream_rx: UpstreamRx, - idle_timeout: Option, - mut observe: impl FnMut(&ResponsesWsEvent) + Send, - mut client_in: In, - mut client_out: Out, -) -> Result<(), Error> -where - In: Stream + Unpin + Send, - Out: Sink + Unpin + Send, - Out::Error: std::fmt::Display, -{ - let idle = - idle_timeout.unwrap_or_else(|| Duration::from_secs(DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS)); - loop { - tokio::select! { - event = client_in.next() => { - let Some(event) = event else { break }; - for outbound in OPENAI_RESPONSES_WS_CONFIG - .transform_ws_request(&event, model)? - .events - { - let payload = serde_json::to_string(&outbound) - .map_err(|error| Error::InvalidResponse(error.to_string()))?; - upstream_tx.send(Message::Text(payload)) - .await - .map_err(|error| Error::Network(error.to_string()))?; - } - } - message = upstream_rx.next() => { - let Some(message) = message else { break }; - match message.map_err(|error| Error::Network(error.to_string()))? { - Message::Text(text) => { - let event = serde_json::from_str::(&text) - .map_err(|error| Error::InvalidResponse(error.to_string()))?; - observe(&event); - for outbound in OPENAI_RESPONSES_WS_CONFIG - .transform_ws_response(&event, model)? - .events - { - client_out.send(outbound) - .await - .map_err(|error| Error::Network(error.to_string()))?; - } - } - Message::Close(_) => break, - _ => {} - } - } - _ = tokio::time::sleep(idle) => break, - } - } - Ok(()) -} - -#[allow(clippy::too_many_arguments)] -pub async fn async_responses_websocket( - model: &str, - api_key: Option<&str>, - api_base: Option<&str>, - first_frame: Option, - idle_timeout: Option, - mut observe: impl FnMut(&ResponsesWsEvent) + Send, - client_in: In, - client_out: Out, -) -> Result<(), Error> -where - In: Stream + Unpin + Send, - Out: Sink + Unpin + Send, - Out::Error: std::fmt::Display, -{ - let key = resolve_api_key(api_key)?; - let upstream = dial_upstream(model, &key, api_base).await?; - let (mut upstream_tx, upstream_rx) = upstream.split(); - if let Some(first_frame) = first_frame { - for outbound in OPENAI_RESPONSES_WS_CONFIG - .transform_ws_request(&first_frame, model)? - .events - { - let payload = serde_json::to_string(&outbound) - .map_err(|error| Error::InvalidResponse(error.to_string()))?; - upstream_tx - .send(Message::Text(payload)) - .await - .map_err(|error| Error::Network(error.to_string()))?; - } - } - ResponsesWebSocketStreaming::bidirectional_forward( - model, - upstream_tx, - upstream_rx, - idle_timeout, - &mut observe, - client_in, - client_out, - ) - .await -} - -#[allow(clippy::too_many_arguments)] -pub async fn responses_ws( - model: &str, - api_key: Option<&str>, - api_base: Option<&str>, - first_frame: Option, - idle_timeout: Option, - observe: impl FnMut(&ResponsesWsEvent) + Send, - client_in: In, - client_out: Out, -) -> Result<(), Error> -where - In: Stream + Unpin + Send, - Out: Sink + Unpin + Send, - Out::Error: std::fmt::Display, -{ - async_responses_websocket( - model, - api_key, - api_base, - first_frame, - idle_timeout, - observe, - client_in, - client_out, - ) - .await -} - -#[cfg(test)] -mod tests { - use super::*; - use futures_channel::mpsc; - use futures_util::{SinkExt, StreamExt}; - use litellm_core::responses::types::ResponsesWsEventType; - use serde_json::json; - use tokio::io::AsyncWriteExt; - use tokio::net::TcpListener; - use tokio_tungstenite::accept_async; - - async fn websocket_base() -> (String, tokio::task::JoinHandle<()>) { - let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind"); - let address = listener.local_addr().expect("local address"); - let task = tokio::spawn(async move { - let (stream, _) = listener.accept().await.expect("accept"); - let mut socket = accept_async(stream).await.expect("websocket handshake"); - while let Some(Ok(Message::Text(text))) = socket.next().await { - let request: serde_json::Value = serde_json::from_str(&text).expect("request json"); - let model = request - .get("model") - .and_then(serde_json::Value::as_str) - .or_else(|| { - request - .get("response") - .and_then(serde_json::Value::as_object) - .and_then(|response| { - response.get("model").and_then(serde_json::Value::as_str) - }) - }) - .expect("enforced model"); - socket - .send(Message::Text( - json!({ - "type": "response.created", - "response": { - "id": format!("resp-{model}"), - "model": model, - "extra": "preserved" - } - }) - .to_string(), - )) - .await - .expect("created event"); - socket - .send(Message::Text( - json!({ - "type": "response.completed", - "response": { - "id": format!("resp-{model}"), - "model": model, - "usage": { - "input_tokens": 1, - "output_tokens": 2, - "total_tokens": 3 - } - } - }) - .to_string(), - )) - .await - .expect("completed event"); - } - }); - (format!("http://{address}"), task) - } - - fn event(value: serde_json::Value) -> ResponsesWsEvent { - serde_json::from_value(value).expect("event") - } - - #[test] - fn explicit_nonblank_key_wins() { - assert_eq!( - resolve_api_key(Some(" explicit ")).expect("key"), - "explicit" - ); - } - - #[test] - fn blank_key_is_not_accepted_without_environment_key() { - if std::env::var(OPENAI_API_KEY_ENV).is_err() { - assert!(resolve_api_key(Some(" ")).is_err()); - } - } - - #[tokio::test] - async fn forwards_events_sequentially_and_enforces_model() { - let (api_base, server) = websocket_base().await; - let (client_tx, client_rx) = mpsc::unbounded(); - let (output_tx, mut output_rx) = mpsc::unbounded(); - let (observed_tx, observed_rx) = mpsc::unbounded(); - client_tx - .unbounded_send(event(json!({ - "type": "response.create", - "model": "wrong" - }))) - .expect("first request"); - client_tx - .unbounded_send(event(json!({ - "type": "response.create", - "response": {"model": "also-wrong"} - }))) - .expect("second request"); - - let task = tokio::spawn(async move { - responses_ws( - "authorized-model", - Some("test-key"), - Some(&api_base), - None, - Some(Duration::from_secs(1)), - move |event| { - observed_tx - .unbounded_send(event.clone()) - .expect("observe event"); - }, - client_rx, - output_tx, - ) - .await - }); - - let first = output_rx.next().await.expect("first output"); - let second = output_rx.next().await.expect("second output"); - let third = output_rx.next().await.expect("third output"); - let fourth = output_rx.next().await.expect("fourth output"); - drop(client_tx); - task.await.expect("splice task").expect("successful splice"); - server.await.expect("server task"); - - assert_eq!(first.event_type, ResponsesWsEventType::ResponseCreated); - assert_eq!(first.model(), Some("authorized-model")); - assert_eq!(first.data["response"]["extra"], "preserved"); - assert_eq!(second.event_type, ResponsesWsEventType::ResponseCompleted); - assert_eq!(third.event_type, ResponsesWsEventType::ResponseCreated); - assert_eq!(fourth.event_type, ResponsesWsEventType::ResponseCompleted); - let observed: Vec<_> = observed_rx.collect().await; - assert_eq!(observed.len(), 4); - assert!( - observed - .iter() - .all(|event| event.event_type != ResponsesWsEventType::ResponseCreate) - ); - } - - #[tokio::test] - async fn idle_timeout_ends_without_upstream_events() { - let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind"); - let address = listener.local_addr().expect("address"); - let server = tokio::spawn(async move { - let (stream, _) = listener.accept().await.expect("accept"); - let _socket = accept_async(stream).await.expect("handshake"); - tokio::time::sleep(Duration::from_secs(1)).await; - }); - let (_client_tx, client_rx) = mpsc::unbounded::(); - let (output_tx, mut output_rx) = mpsc::unbounded(); - let result = responses_ws( - "model", - Some("key"), - Some(&format!("http://{address}")), - None, - Some(Duration::from_millis(20)), - |_| {}, - client_rx, - output_tx, - ) - .await; - assert!(result.is_ok()); - assert!(output_rx.next().await.is_none()); - server.abort(); - } - - #[tokio::test] - async fn dial_http_status_is_preserved() { - let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind"); - let address = listener.local_addr().expect("address"); - let server = tokio::spawn(async move { - let (mut stream, _) = listener.accept().await.expect("accept"); - stream - .write_all(b"HTTP/1.1 401 Unauthorized\r\nContent-Length: 0\r\n\r\n") - .await - .expect("response"); - }); - let (_client_tx, client_rx) = mpsc::unbounded::(); - let (output_tx, _output_rx) = mpsc::unbounded(); - let error = responses_ws( - "model", - Some("key"), - Some(&format!("http://{address}")), - None, - Some(Duration::from_millis(20)), - |_| {}, - client_rx, - output_tx, - ) - .await - .expect_err("status error"); - assert!(matches!(error, Error::Http { status: 401, .. })); - server.await.expect("server task"); - } - - #[tokio::test] - async fn dial_http_500_status_is_preserved() { - let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind"); - let address = listener.local_addr().expect("address"); - let server = tokio::spawn(async move { - let (mut stream, _) = listener.accept().await.expect("accept"); - stream - .write_all(b"HTTP/1.1 500 Internal Server Error\r\nContent-Length: 0\r\n\r\n") - .await - .expect("response"); - }); - let (_client_tx, client_rx) = mpsc::unbounded::(); - let (output_tx, _output_rx) = mpsc::unbounded(); - let error = responses_ws( - "model", - Some("key"), - Some(&format!("http://{address}")), - None, - Some(Duration::from_millis(20)), - |_| {}, - client_rx, - output_tx, - ) - .await - .expect_err("status error"); - assert!(matches!(error, Error::Http { status: 500, .. })); - server.await.expect("server task"); - } -} diff --git a/litellm-rust/crates/ai-gateway/src/lib.rs b/litellm-rust/crates/ai-gateway/src/lib.rs index 08fbde564ed..3738854cd23 100644 --- a/litellm-rust/crates/ai-gateway/src/lib.rs +++ b/litellm-rust/crates/ai-gateway/src/lib.rs @@ -13,7 +13,6 @@ //! binary turns on. pub mod audio_transcription; -mod client; pub mod io; pub mod ocr; diff --git a/litellm-rust/crates/ai-gateway/src/ocr/handler.rs b/litellm-rust/crates/ai-gateway/src/ocr/handler.rs index 0aac26996e1..a633f8db16a 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/handler.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/handler.rs @@ -1,99 +1,16 @@ -use litellm_core::call_lifecycle::CallLifecycleContext; -use litellm_core::error::Error; -use litellm_core::ocr::transformation::OcrResponseHandling; -use litellm_core::provider_callbacks::ProviderAttemptObserver; -use litellm_core::provider_callbacks::handler::{ - ProviderAttemptContext, ProviderRequest, send_provider_request, -}; +use litellm_core::Error; +use litellm_core::ocr::observers::OcrObserver; +use litellm_core::ocr::{PreparedOcrRequest, execute_ocr_provider_call as core_execute}; use serde_json::Value; -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, - observer: &mut Observer, -) -> Result -where - Observer: ProviderAttemptObserver, - Observer::Error: std::fmt::Display, -{ + observer: &mut impl OcrObserver, +) -> Result { let request = hooks.prepare_provider_request(request).await?; - 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); - } - if let Some(duration) = request.timeout { - request_builder = request_builder.timeout(duration); - } - - 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; - if request.config.response_handling() == OcrResponseHandling::AzureDocumentIntelligencePoll - && status.as_u16() == 202 - { - let operation_url = response - .headers - .get("operation-location") - .and_then(|value| value.to_str().ok()) - .map(str::to_string) - .ok_or_else(|| { - Error::InvalidResponse( - "Azure Document Intelligence returned 202 but no Operation-Location header found" - .to_string(), - ) - })?; - let response_json = poll_document_intelligence( - &operation_url, - &request.url, - &request.upstream_headers, - request.timeout, - ) - .await?; - return Ok(request - .config - .transform_ocr_response_with_params( - &request.model, - response_json, - &request.optional_params, - )? - .into_json()); - } - - let response_json: Value = serde_json::from_str(&response.body) - .map_err(|err| Error::InvalidResponse(format!("invalid OCR response JSON: {err}")))?; - - Ok(request - .config - .transform_ocr_response_with_params( - &request.model, - response_json, - &request.optional_params, - )? - .into_json()) + core_execute(request, observer).await } diff --git a/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs b/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs index c8fde0bf449..a6c2f59f911 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs @@ -1,28 +1,24 @@ use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming}; use litellm_core::error::Error; -use litellm_core::providers::reducto::ocr::transformation::{ - build_upload_request, extract_document_source, extract_upload_file_id, -}; -use litellm_core::request_context::RequestAttribution; +use litellm_core::ocr::{PreparedOcrRequest, ProviderOcrRequest, prepare_ocr_provider_call}; use serde_json::{Map, Value, json}; use std::future::Future; use std::pin::Pin; -use super::common_utils::{convert_document_url_to_data_uri, string_headers, truncate_error_body}; -use super::types::{PreparedOcrRequest, ProviderOcrRequest}; -use crate::client::http_client; use crate::integrations::custom_guardrail::{ CustomGuardrailRunner, GuardrailContext, GuardrailError, GuardrailRequest, }; use crate::integrations::custom_logger::{ CallType, CallbackTiming, CallbackValue, CustomLoggerRunner, LoggingError, ModelCallDetails, }; -use crate::integrations::types::{StandardLoggingMetadata, StandardLoggingPayload}; +use crate::integrations::types::{ + RequestMetadata, StandardLoggingMetadata, StandardLoggingPayload, +}; pub(crate) struct OcrLifecycleHooks { logger_runner: CustomLoggerRunner, guardrail_runner: CustomGuardrailRunner, - request_metadata: RequestAttribution, + request_metadata: RequestMetadata, } type OcrFuture<'a, T> = Pin> + Send + 'a>>; @@ -32,7 +28,7 @@ impl OcrLifecycleHooks { pub(crate) fn new( logger_runner: CustomLoggerRunner, guardrail_runner: CustomGuardrailRunner, - request_metadata: RequestAttribution, + request_metadata: RequestMetadata, ) -> Self { Self { logger_runner, @@ -77,83 +73,24 @@ impl OcrLifecycleHooks { &self, request: PreparedOcrRequest, ) -> Result { - let config = request.config?; - let env_lookup = |key: &str| std::env::var(key).ok(); - let upstream_headers = config.validate_environment( - string_headers(request.extra_headers)?, - request.api_key.as_deref(), - &env_lookup, - )?; - let url_params = request - .optional_params - .clone() - .into_iter() - .chain(request.vertex.into_map()) - .collect(); - let url = config.complete_url( - request.api_base.as_deref(), - &request.model, - &url_params, - &env_lookup, - )?; - let model = request.model.clone(); - let custom_llm_provider = request.custom_llm_provider.clone(); - let is_reducto = custom_llm_provider == "reducto"; - let document = if is_reducto { - let guarded_document = self - .run_during_call_guardrails(&model, &custom_llm_provider, &url, request.document) - .await?; - upload_reducto_document( - &guarded_document, - request.api_base.as_deref(), - request.timeout, - &upstream_headers, - ) - .await? - } else if config.requires_data_uri_document() { - convert_document_url_to_data_uri(request.document).await? - } else { - request.document - }; - let optional_params = request.optional_params; - let body = config - .transform_ocr_request(&request.model, document, optional_params.clone())? - .data; - let body = if is_reducto { - body - } else { - self.run_during_call_guardrails(&model, &custom_llm_provider, &url, body) - .await? - }; - Ok(ProviderOcrRequest { - model, - custom_llm_provider, - config, - url, - body, - optional_params, - upstream_headers, - timeout: request.timeout, - }) + let provider_request = prepare_ocr_provider_call(request).await?; + self.run_during_call_guardrails(provider_request).await } async fn run_during_call_guardrails( &self, - model: &str, - custom_llm_provider: &str, - url: &str, - body: Value, - ) -> Result { + request: ProviderOcrRequest, + ) -> Result { if self.guardrail_runner.is_empty() { - return Ok(body); + return Ok(request); } let context = guardrail_context(&self.request_metadata); let guardrail_request = GuardrailRequest::new(json!({ - "model": model, - "custom_llm_provider": custom_llm_provider, - "url": url, - "body": body, + "model": request.model(), + "custom_llm_provider": request.custom_llm_provider(), + "url": request.url(), + "body": request.body(), })); let (guardrail_request, _) = self .guardrail_runner @@ -161,6 +98,7 @@ impl OcrLifecycleHooks { .await .map_err(guardrail_error_to_core_error)?; parse_ocr_during_call_guardrail_request(guardrail_request) + .map(|body| request.with_body(body)) } fn standard_logging_payload( @@ -192,63 +130,6 @@ impl OcrLifecycleHooks { } } -async fn upload_reducto_document( - document: &Value, - api_base: Option<&str>, - timeout: Option, - upstream_headers: &[(String, String)], -) -> Result { - let source = extract_document_source(document)?; - let Some(authorization) = upstream_headers - .iter() - .find(|(name, _)| name.eq_ignore_ascii_case("authorization")) - .map(|(_, value)| value.as_str()) - else { - return Err(Error::Auth( - "Reducto upload requires an Authorization header".to_string(), - )); - }; - let Some(upload) = build_upload_request(source, authorization, api_base) else { - return Ok(document.clone()); - }; - let part = reqwest::multipart::Part::bytes(upload.bytes) - .file_name(upload.file_name) - .mime_str(&upload.mime_type) - .map_err(|error| Error::InvalidRequest(error.to_string()))?; - let form = reqwest::multipart::Form::new().part("file", part); - let mut request_builder = http_client().post(upload.url).multipart(form); - for (name, value) in upstream_headers { - if !name.eq_ignore_ascii_case("content-type") - && !name.eq_ignore_ascii_case("content-length") - { - request_builder = request_builder.header(name, value); - } - } - if let Some(timeout) = timeout { - request_builder = request_builder.timeout(timeout); - } - let response = request_builder - .send() - .await - .map_err(|error| Error::Network(error.to_string()))?; - let status = response.status(); - let body = response - .text() - .await - .map_err(|error| Error::Network(error.to_string()))?; - if !status.is_success() { - return Err(Error::Http { - status: status.as_u16(), - body: truncate_error_body(&body), - }); - } - let response_json: Value = serde_json::from_str(&body).map_err(|error| { - Error::InvalidResponse(format!("invalid Reducto upload response JSON: {error}")) - })?; - let file_id = extract_upload_file_id(&response_json)?; - Ok(json!({"type": "document_url", "document_url": file_id})) -} - impl CallLifecycleHooks for OcrLifecycleHooks { type PreCallFuture<'a> = OcrFuture<'a, PreparedOcrRequest>; type DuringCallFuture<'a> = OcrFuture<'a, PreparedOcrRequest>; @@ -329,7 +210,7 @@ impl CallLifecycleHooks for OcrLi } } -fn guardrail_context(metadata: &RequestAttribution) -> GuardrailContext { +fn guardrail_context(metadata: &RequestMetadata) -> GuardrailContext { GuardrailContext { call_type: CallType::Ocr, selected_guardrails: Vec::new(), diff --git a/litellm-rust/crates/ai-gateway/src/ocr/mod.rs b/litellm-rust/crates/ai-gateway/src/ocr/mod.rs index bf21d564ff7..db4a1cd3fdc 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/mod.rs @@ -6,7 +6,6 @@ use litellm_core::request_context::LiteLlmRequestContext; use litellm_core::request_options::RequestOptions; use serde_json::Value; -mod common_utils; mod handler; mod hooks; mod prepare; diff --git a/litellm-rust/crates/ai-gateway/src/ocr/prepare.rs b/litellm-rust/crates/ai-gateway/src/ocr/prepare.rs index 1fe884489ad..33ebc88016b 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/prepare.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/prepare.rs @@ -1,181 +1,46 @@ -use crate::integrations::types::RequestHooks; -use litellm_core::call_lifecycle::CallLifecycleContext; -use litellm_core::request_context::LiteLlmRequestContext; -use litellm_core::request_options::RequestOptions; -use std::sync::atomic::{AtomicU64, Ordering}; -use std::time::{SystemTime, UNIX_EPOCH}; +use litellm_core::ocr::{OcrRequest as CoreOcrRequest, PreparedOcrRequest}; -use litellm_core::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider}; -use serde_json::{Map, Value}; - -use super::common_utils::ocr_provider_config; use super::hooks::OcrLifecycleHooks; -use super::types::{OcrRequest, PreparedOcrRequest}; +use super::types::OcrRequest; use crate::integrations::custom_guardrail::CustomGuardrailRunner; use crate::integrations::custom_logger::CustomLoggerRunner; pub(crate) struct PreparedOcrCall { - pub(crate) context: CallLifecycleContext, pub(crate) request: PreparedOcrRequest, pub(crate) hooks: OcrLifecycleHooks, } -#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] -pub(crate) fn prepare_ocr_call( - request: OcrRequest<'_>, - options: RequestOptions, - context: &LiteLlmRequestContext, - hooks: RequestHooks, -) -> PreparedOcrCall { - let call_id = context - .litellm_call_id - .clone() - .unwrap_or_else(new_ocr_call_id); - let provider_info = - get_custom_llm_provider(request.model, options.custom_llm_provider.as_deref()).unwrap_or( - CustomLlmProvider { - model: request.model, - custom_llm_provider: "mistral", - }, - ); - let model = provider_info.model.to_string(); - let custom_llm_provider = provider_info.custom_llm_provider.to_string(); - let config = ocr_provider_config(&custom_llm_provider, &model) - .ok_or_else(|| litellm_core::Error::InvalidProvider(custom_llm_provider.clone())) - .and_then(|config| { - validate_request_format(config, &request.optional_params, &custom_llm_provider)?; - Ok(config) - }); - let optional_params = match &config { - Ok(config) => { - let supported = config.supported_ocr_params(); - config.map_ocr_params( - &request - .optional_params - .iter() - .filter(|(name, _)| supported.contains(&name.as_str())) - .map(|(name, value)| (name.clone(), value.clone())) - .collect(), - ) - } - Err(_) => request.optional_params, - }; - +pub(crate) fn prepare_ocr_call(request: OcrRequest<'_>) -> PreparedOcrCall { + let OcrRequest { + model, + document, + api_key, + api_base, + custom_llm_provider, + extra_headers, + optional_params, + timeout, + callbacks, + guardrails, + request_metadata, + litellm_call_id, + } = request; PreparedOcrCall { - context: CallLifecycleContext::new( - "ocr", - model.clone(), - custom_llm_provider.clone(), - call_id, - ), - request: PreparedOcrRequest { - config, + request: litellm_core::ocr::prepare_ocr_call(CoreOcrRequest { model, + document, + api_key, + api_base, custom_llm_provider, - document: request.document, - vertex: options.vertex.unwrap_or_default(), - api_key: options.api_key, - api_base: options.api_base, - extra_headers: options.extra_headers, + extra_headers, optional_params, - timeout: options.timeout, - }, + timeout, + litellm_call_id, + }), hooks: OcrLifecycleHooks::new( - CustomLoggerRunner::new(hooks.callbacks), - CustomGuardrailRunner::new(hooks.guardrails), - context.attribution.clone(), + CustomLoggerRunner::new(callbacks), + CustomGuardrailRunner::new(guardrails), + request_metadata, ), } } - -fn validate_request_format( - config: &'static dyn litellm_core::ocr::transformation::OcrProviderConfig, - optional_params: &Map, - provider: &str, -) -> Result<(), litellm_core::Error> { - let Some(format) = optional_params.get("req_format") else { - return Ok(()); - }; - match format.as_str() { - Some("litellm") => Ok(()), - Some("native") if config.supported_ocr_params().contains(&"req_format") => Ok(()), - Some("native") => Err(litellm_core::Error::InvalidRequest(format!( - "`req_format=native` is not supported for provider {provider}" - ))), - _ => Err(litellm_core::Error::InvalidRequest(format!( - "Invalid `req_format`: {format}. Expected `litellm` or `native`" - ))), - } -} - -fn new_ocr_call_id() -> String { - static COUNTER: AtomicU64 = AtomicU64::new(1); - let sequence = COUNTER.fetch_add(1, Ordering::Relaxed); - let timestamp = SystemTime::now() - .duration_since(UNIX_EPOCH) - .map(|duration| duration.as_nanos()) - .unwrap_or(0); - format!("ocr-{timestamp}-{sequence}") -} - -#[cfg(test)] -mod tests { - use crate::integrations::types::RequestHooks; - use litellm_core::error::Error; - use litellm_core::request_context::LiteLlmRequestContext; - use litellm_core::request_options::RequestOptions; - use serde_json::{Map, json}; - - use super::{OcrRequest, prepare_ocr_call}; - - fn base_ocr_request(model: &str) -> OcrRequest<'_> { - OcrRequest { - model, - document: json!({ - "type": "document_url", - "document_url": "https://example.com/doc.pdf" - }), - optional_params: Map::new(), - } - } - - fn request_with_format(format: &str) -> OcrRequest<'_> { - let mut request = base_ocr_request("mistral/mistral-ocr-latest"); - request.optional_params = Map::from_iter([("req_format".to_string(), json!(format))]); - request - } - - #[test] - fn native_format_rejected_for_provider_without_support_as_bad_request() { - let prepared = prepare_ocr_call( - request_with_format("native"), - RequestOptions::default(), - &LiteLlmRequestContext { - ..Default::default() - }, - RequestHooks { - ..Default::default() - }, - ); - assert!( - matches!(prepared.request.config, Err(Error::InvalidRequest(message)) if message.contains("not supported for provider")) - ); - } - - #[test] - fn unknown_format_rejected_for_provider_without_support_as_bad_request() { - let prepared = prepare_ocr_call( - request_with_format("raw"), - RequestOptions::default(), - &LiteLlmRequestContext { - ..Default::default() - }, - RequestHooks { - ..Default::default() - }, - ); - assert!( - matches!(prepared.request.config, Err(Error::InvalidRequest(message)) if message.contains("Invalid `req_format`")) - ); - } -} diff --git a/litellm-rust/crates/ai-gateway/src/ocr/types.rs b/litellm-rust/crates/ai-gateway/src/ocr/types.rs index 4e26522de6b..e96d2df1adb 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/types.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/types.rs @@ -1,35 +1,23 @@ +use std::sync::Arc; use std::time::Duration; -use litellm_core::ocr::transformation::OcrProviderConfig; -use litellm_core::request_options::VertexOptions; use serde_json::{Map, Value}; +use crate::integrations::custom_guardrail::CustomGuardrail; +use crate::integrations::custom_logger::CustomLogger; +use crate::integrations::types::RequestMetadata; + pub struct OcrRequest<'a> { pub model: &'a str, pub document: Value, + pub api_key: Option<&'a str>, + pub api_base: Option<&'a str>, + pub custom_llm_provider: Option<&'a str>, + pub extra_headers: Option>, pub optional_params: Map, -} - -pub(crate) struct PreparedOcrRequest { - pub(crate) config: Result<&'static dyn OcrProviderConfig, litellm_core::Error>, - pub(crate) model: String, - pub(crate) custom_llm_provider: String, - pub(crate) document: Value, - pub(crate) vertex: VertexOptions, - pub(crate) api_key: Option, - pub(crate) api_base: Option, - pub(crate) extra_headers: Option>, - pub(crate) optional_params: Map, - pub(crate) timeout: Option, -} - -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, - pub(crate) optional_params: Map, - pub(crate) upstream_headers: Vec<(String, String)>, - pub(crate) timeout: Option, + pub timeout: Option, + pub callbacks: Vec>, + pub guardrails: Vec>, + pub request_metadata: RequestMetadata, + pub litellm_call_id: Option<&'a str>, } diff --git a/litellm-rust/crates/ai-gateway/tests/ocr_lifecycle.rs b/litellm-rust/crates/ai-gateway/tests/ocr_lifecycle.rs deleted file mode 100644 index d3836a8c8c6..00000000000 --- a/litellm-rust/crates/ai-gateway/tests/ocr_lifecycle.rs +++ /dev/null @@ -1,928 +0,0 @@ -use litellm_ai_gateway::integrations::types::RequestHooks; -use litellm_core::request_context::LiteLlmRequestContext; -use litellm_core::request_context::RequestAttribution; -use litellm_core::request_options::RequestOptions; -use std::sync::{Arc, Mutex}; -use std::time::Duration; - -use litellm_ai_gateway::integrations::custom_guardrail::{ - CustomGuardrail, GuardrailContext, GuardrailDecision, GuardrailError, GuardrailEventHook, - GuardrailFuture, GuardrailRequest, -}; -use litellm_ai_gateway::integrations::custom_logger::{ - CallbackTiming, CallbackValue, CustomLogger, LogFuture, ModelCallDetails, -}; -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]; - loop { - let n = socket.read(&mut buffer).await.expect("reads request"); - if n == 0 { - break; - } - request.extend_from_slice(&buffer[..n]); - if request.windows(4).any(|window| window == b"\r\n\r\n") { - break; - } - } - String::from_utf8(request).expect("request is utf8") -} - -async fn read_http_request(socket: &mut TcpStream) -> String { - let mut request = Vec::new(); - let mut buffer = [0_u8; 1024]; - let header_end = loop { - let n = socket.read(&mut buffer).await.expect("reads request"); - if n == 0 { - break request.len(); - } - request.extend_from_slice(&buffer[..n]); - if let Some(position) = request.windows(4).position(|window| window == b"\r\n\r\n") { - break position + 4; - } - }; - let headers = String::from_utf8_lossy(&request[..header_end]); - let content_length = headers - .lines() - .find_map(|line| { - let (name, value) = line.split_once(':')?; - name.eq_ignore_ascii_case("content-length") - .then(|| value.trim().parse::().ok()) - .flatten() - }) - .unwrap_or(0); - while request.len().saturating_sub(header_end) < content_length { - let n = socket.read(&mut buffer).await.expect("reads body"); - if n == 0 { - break; - } - request.extend_from_slice(&buffer[..n]); - } - String::from_utf8(request).expect("request is utf8") -} - -#[derive(Clone, Debug, PartialEq)] -struct RecordedLogEvent { - hook: &'static str, - model: String, - call_type: String, - user_id: Option, - response_object: Option, - error_kind: Option, -} - -#[derive(Default)] -struct RecordingOcrLogger { - events: Mutex>, -} - -impl RecordingOcrLogger { - fn events(&self) -> Vec { - self.events.lock().unwrap().clone() - } -} - -impl CustomLogger for RecordingOcrLogger { - fn async_log_success_event<'a>( - &'a self, - model_call_details: &'a ModelCallDetails, - response_obj: &'a CallbackValue, - _timing: CallbackTiming, - ) -> LogFuture<'a> { - Box::pin(async move { - self.events.lock().unwrap().push(RecordedLogEvent { - hook: "async_log_success_event", - model: model_call_details.model.clone(), - call_type: model_call_details.call_type.to_string(), - user_id: model_call_details.metadata.user_api_key_user_id.clone(), - response_object: Some(response_obj.object.clone()), - error_kind: None, - }); - Ok(()) - }) - } - - fn async_log_failure_event<'a>( - &'a self, - model_call_details: &'a ModelCallDetails, - response_obj: Option<&'a CallbackValue>, - _timing: CallbackTiming, - ) -> LogFuture<'a> { - Box::pin(async move { - self.events.lock().unwrap().push(RecordedLogEvent { - hook: "async_log_failure_event", - model: model_call_details.model.clone(), - call_type: model_call_details.call_type.to_string(), - user_id: model_call_details.metadata.user_api_key_user_id.clone(), - response_object: response_obj.map(|value| value.object.clone()), - error_kind: model_call_details - .failure_error - .as_ref() - .map(|error| error.kind.clone()), - }); - Ok(()) - }) - } -} - -struct RecordingOcrGuardrail { - hooks: Vec, - events: Mutex>, - block_pre_call: bool, - block_during_call: bool, -} - -impl RecordingOcrGuardrail { - fn new(hooks: Vec) -> Self { - Self { - hooks, - events: Mutex::new(Vec::new()), - block_pre_call: false, - block_during_call: false, - } - } - - fn blocking_pre_call() -> Self { - Self { - hooks: vec![GuardrailEventHook::PreCall], - events: Mutex::new(Vec::new()), - block_pre_call: true, - block_during_call: false, - } - } - - fn blocking_during_call() -> Self { - Self { - hooks: vec![GuardrailEventHook::DuringCall], - events: Mutex::new(Vec::new()), - block_pre_call: false, - block_during_call: true, - } - } - - fn events(&self) -> Vec<&'static str> { - self.events.lock().unwrap().clone() - } -} - -impl CustomGuardrail for RecordingOcrGuardrail { - fn guardrail_name(&self) -> &str { - "recording-ocr-guardrail" - } - - fn supported_event_hooks(&self) -> &[GuardrailEventHook] { - &self.hooks - } - - fn async_pre_call_hook<'a>( - &'a self, - _context: &'a GuardrailContext, - mut request: GuardrailRequest, - ) -> GuardrailFuture<'a> { - Box::pin(async move { - self.events.lock().unwrap().push("async_pre_call_hook"); - if self.block_pre_call { - return Ok(GuardrailDecision::Block(GuardrailError::blocked( - "blocked before provider", - ))); - } - request.data["document"]["guarded_pre"] = json!(true); - Ok(GuardrailDecision::Mask(request)) - }) - } - - fn async_moderation_hook<'a>( - &'a self, - _context: &'a GuardrailContext, - mut request: GuardrailRequest, - ) -> GuardrailFuture<'a> { - Box::pin(async move { - self.events.lock().unwrap().push("async_moderation_hook"); - if self.block_during_call { - return Ok(GuardrailDecision::Block(GuardrailError::blocked( - "blocked before provider", - ))); - } - request.data["body"]["guarded_during"] = json!(true); - Ok(GuardrailDecision::Mask(request)) - }) - } -} - -fn base_ocr_request(model: &str) -> (OcrRequest<'_>, RequestOptions) { - ( - OcrRequest { - model, - document: json!({ - "type": "document_url", - "document_url": "https://example.com/doc.pdf" - }), - optional_params: Map::new(), - }, - RequestOptions { - api_key: Some("sk-test".to_string()), - ..Default::default() - }, - ) -} - -#[tokio::test] -async fn reducto_during_call_guardrail_blocks_before_upload() { - let listener = TcpListener::bind("127.0.0.1:0") - .await - .expect("test listener binds"); - let address = listener.local_addr().expect("listener has local address"); - let api_base = format!("http://{address}"); - let guardrail = Arc::new(RecordingOcrGuardrail::blocking_during_call()); - let (mut request, mut options) = base_ocr_request("reducto/parse-v3"); - options.api_base = Some(&api_base).map(|value| value.to_string()); - request.document = json!({ - "type": "document_url", - "document_url": "data:application/pdf;base64,JVBERi0xLjQ=" - }); - let hooks = RequestHooks { - guardrails: vec![guardrail.clone()], - ..Default::default() - }; - - let error = ocr( - request, - &options, - &LiteLlmRequestContext { - ..Default::default() - }, - hooks, - ) - .await - .expect_err("guardrail blocks upload"); - - assert!(matches!(error, Error::InvalidRequest(_))); - assert_eq!(guardrail.events(), vec!["async_moderation_hook"]); - let accepted = tokio::time::timeout(Duration::from_millis(100), listener.accept()).await; - assert!(accepted.is_err(), "upload socket should not be touched"); -} - -#[tokio::test] -async fn reducto_upload_error_body_is_truncated() { - let listener = TcpListener::bind("127.0.0.1:0") - .await - .expect("test listener binds"); - let address = listener.local_addr().expect("listener has local address"); - let server = tokio::spawn(async move { - let (mut socket, _) = listener.accept().await.expect("accepts upload request"); - let _request = read_http_request(&mut socket).await; - let body = "x".repeat(300); - let response = format!( - "HTTP/1.1 500 Internal Server Error\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", - body.len(), - body - ); - socket - .write_all(response.as_bytes()) - .await - .expect("writes upload response"); - }); - let api_base = format!("http://{address}"); - let (mut request, mut options) = base_ocr_request("reducto/parse-v3"); - options.api_base = Some(&api_base).map(|value| value.to_string()); - request.document = json!({ - "type": "document_url", - "document_url": "data:application/pdf;base64,JVBERi0xLjQ=" - }); - - let error = ocr( - request, - &options, - &LiteLlmRequestContext { - ..Default::default() - }, - RequestHooks::default(), - ) - .await - .expect_err("upload should fail"); - - assert!( - matches!(error, Error::Http { status: 500, body } if body.chars().count() < 300 && body.ends_with("... (truncated)")) - ); - server.await.expect("server task completes"); -} - -#[tokio::test] -async fn ocr_lifecycle_runs_pre_during_and_success_hooks() { - let listener = TcpListener::bind("127.0.0.1:0") - .await - .expect("test listener binds"); - let addr = listener.local_addr().expect("listener has local addr"); - - let server = tokio::spawn(async move { - let (mut socket, _) = listener.accept().await.expect("accepts one request"); - let request = read_http_request(&mut socket).await; - let response_body = r#"{"pages":[{"index":0,"markdown":"ok"}],"model":"mistral-ocr-latest","usage_info":{"pages_processed":1}}"#; - let response = format!( - "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", - response_body.len(), - response_body - ); - socket - .write_all(response.as_bytes()) - .await - .expect("writes response"); - request - }); - - let logger = Arc::new(RecordingOcrLogger::default()); - let guardrail = Arc::new(RecordingOcrGuardrail::new(vec![ - GuardrailEventHook::PreCall, - GuardrailEventHook::DuringCall, - ])); - let response = ocr( - OcrRequest { - model: "mistral-ocr-latest", - document: json!({ - "type": "document_url", - "document_url": "https://example.com/doc.pdf" - }), - optional_params: Map::new(), - }, - &RequestOptions { - api_key: (Some("sk-test")).map(|value| value.to_string()), - api_base: (Some(&format!("http://{addr}"))).map(|value| value.to_string()), - custom_llm_provider: (Some("mistral")).map(|value| value.to_string()), - extra_headers: None, - timeout: Some(Duration::from_secs(5)), - ..Default::default() - }, - &LiteLlmRequestContext { - attribution: RequestAttribution { - user_api_key_user_id: Some("user-1".to_string()), - ..Default::default() - }, - litellm_call_id: (Some("ocr-call-1")).map(|value| value.to_string()), - ..Default::default() - }, - RequestHooks { - callbacks: vec![logger.clone()], - guardrails: vec![guardrail.clone()], - }, - ) - .await - .expect("ocr request succeeds"); - - assert_eq!(response["pages"][0]["markdown"], "ok"); - assert_eq!( - guardrail.events(), - vec!["async_pre_call_hook", "async_moderation_hook"] - ); - assert_eq!( - logger.events(), - vec![RecordedLogEvent { - hook: "async_log_success_event", - model: "mistral-ocr-latest".to_string(), - call_type: "ocr".to_string(), - user_id: Some("user-1".to_string()), - response_object: Some("ocr".to_string()), - error_kind: None, - }] - ); - - let request = server.await.expect("server task completes"); - assert!(request.contains(r#""guarded_pre":true"#), "{request}"); - assert!(request.contains(r#""guarded_during":true"#), "{request}"); -} - -#[tokio::test] -async fn ocr_lifecycle_runs_failure_hook_on_provider_error() { - let listener = TcpListener::bind("127.0.0.1:0") - .await - .expect("test listener binds"); - let addr = listener.local_addr().expect("listener has local addr"); - - let server = tokio::spawn(async move { - let (mut socket, _) = listener.accept().await.expect("accepts one request"); - let _request = read_http_request(&mut socket).await; - let response_body = "provider failed"; - let response = format!( - "HTTP/1.1 500 Internal Server Error\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", - response_body.len(), - response_body - ); - socket - .write_all(response.as_bytes()) - .await - .expect("writes response"); - }); - - let logger = Arc::new(RecordingOcrLogger::default()); - let err = ocr( - OcrRequest { - model: "mistral-ocr-latest", - document: json!({ - "type": "document_url", - "document_url": "https://example.com/doc.pdf" - }), - optional_params: Map::new(), - }, - &RequestOptions { - api_key: (Some("sk-test")).map(|value| value.to_string()), - api_base: (Some(&format!("http://{addr}"))).map(|value| value.to_string()), - custom_llm_provider: (Some("mistral")).map(|value| value.to_string()), - extra_headers: None, - timeout: Some(Duration::from_secs(5)), - ..Default::default() - }, - &LiteLlmRequestContext { - attribution: RequestAttribution::default(), - litellm_call_id: (Some("ocr-call-2")).map(|value| value.to_string()), - ..Default::default() - }, - RequestHooks { - callbacks: vec![logger.clone()], - guardrails: Vec::new(), - }, - ) - .await - .expect_err("provider error propagates"); - - assert!(matches!(err, Error::Http { status: 500, .. })); - server.await.expect("server task completes"); - assert_eq!( - logger.events(), - vec![RecordedLogEvent { - hook: "async_log_failure_event", - model: "mistral-ocr-latest".to_string(), - call_type: "ocr".to_string(), - user_id: None, - response_object: Some("error".to_string()), - error_kind: Some("HttpError".to_string()), - }] - ); -} - -#[tokio::test] -async fn ocr_lifecycle_pre_call_block_skips_provider_socket() { - let listener = TcpListener::bind("127.0.0.1:0") - .await - .expect("test listener binds"); - let addr = listener.local_addr().expect("listener has local addr"); - let logger = Arc::new(RecordingOcrLogger::default()); - let guardrail = Arc::new(RecordingOcrGuardrail::blocking_pre_call()); - - let err = ocr( - OcrRequest { - model: "mistral-ocr-latest", - document: json!({ - "type": "document_url", - "document_url": "https://example.com/doc.pdf" - }), - optional_params: Map::new(), - }, - &RequestOptions { - api_key: (Some("sk-test")).map(|value| value.to_string()), - api_base: (Some(&format!("http://{addr}"))).map(|value| value.to_string()), - custom_llm_provider: (Some("mistral")).map(|value| value.to_string()), - extra_headers: None, - timeout: Some(Duration::from_millis(100)), - ..Default::default() - }, - &LiteLlmRequestContext { - attribution: RequestAttribution::default(), - litellm_call_id: (Some("ocr-call-3")).map(|value| value.to_string()), - ..Default::default() - }, - RequestHooks { - callbacks: vec![logger.clone()], - guardrails: vec![guardrail.clone()], - }, - ) - .await - .expect_err("guardrail blocks request"); - - assert!(matches!(err, Error::InvalidRequest(_))); - assert_eq!(guardrail.events(), vec!["async_pre_call_hook"]); - assert_eq!( - logger.events(), - vec![RecordedLogEvent { - hook: "async_log_failure_event", - model: "mistral-ocr-latest".to_string(), - call_type: "ocr".to_string(), - user_id: None, - response_object: Some("error".to_string()), - error_kind: Some("InvalidRequest".to_string()), - }] - ); - let accepted = tokio::time::timeout(Duration::from_millis(100), listener.accept()).await; - assert!(accepted.is_err(), "provider socket should not be touched"); -} - -#[tokio::test] -async fn ocr_does_not_duplicate_authorization_header_when_header_is_supplied() { - let listener = TcpListener::bind("127.0.0.1:0") - .await - .expect("test listener binds"); - let addr = listener.local_addr().expect("listener has local addr"); - - let server = tokio::spawn(async move { - let (mut socket, _) = listener.accept().await.expect("accepts one request"); - let request = read_http_headers(&mut socket).await; - let response_body = r#"{"pages":[{"index":0,"markdown":"ok"}],"model":"mistral-ocr-latest","usage_info":{"pages_processed":1}}"#; - let response = format!( - "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", - response_body.len(), - response_body - ); - socket - .write_all(response.as_bytes()) - .await - .expect("writes response"); - request - }); - - let mut headers = Map::new(); - headers.insert( - "Authorization".to_string(), - Value::String("Bearer sk-from-python".to_string()), - ); - headers.insert( - "x-trace-id".to_string(), - Value::String("trace-1".to_string()), - ); - - let response = ocr( - OcrRequest { - model: "mistral-ocr-latest", - document: json!({ - "type": "document_url", - "document_url": "https://example.com/doc.pdf" - }), - optional_params: Map::new(), - }, - &RequestOptions { - api_key: (Some("sk-for-rust-fallback")).map(|value| value.to_string()), - api_base: (Some(&format!("http://{addr}"))).map(|value| value.to_string()), - custom_llm_provider: (Some("mistral")).map(|value| value.to_string()), - extra_headers: Some(headers), - timeout: Some(Duration::from_secs(5)), - ..Default::default() - }, - &LiteLlmRequestContext { - attribution: RequestAttribution::default(), - litellm_call_id: None, - ..Default::default() - }, - RequestHooks { - callbacks: Vec::new(), - guardrails: Vec::new(), - }, - ) - .await - .expect("ocr request succeeds"); - - assert_eq!(response["pages"][0]["markdown"], "ok"); - - let request = server.await.expect("server task completes"); - let authorization_count = request - .lines() - .filter(|line| line.to_ascii_lowercase().starts_with("authorization:")) - .count(); - assert_eq!(authorization_count, 1, "{request}"); - assert!( - request.contains("authorization: Bearer sk-from-python") - || request.contains("Authorization: Bearer sk-from-python"), - "{request}" - ); -} - -#[tokio::test] -async fn document_intelligence_poll_uses_resolved_subscription_key() { - let listener = TcpListener::bind("127.0.0.1:0") - .await - .expect("test listener binds"); - let addr = listener.local_addr().expect("listener has local addr"); - let operation_url = format!("http://{addr}/operations/1"); - - let server = tokio::spawn(async move { - let (mut post_socket, _) = listener.accept().await.expect("accepts post request"); - let post_request = read_http_headers(&mut post_socket).await; - let post_response = format!( - "HTTP/1.1 202 Accepted\r\noperation-location: {operation_url}\r\ncontent-length: 0\r\nconnection: close\r\n\r\n" - ); - post_socket - .write_all(post_response.as_bytes()) - .await - .expect("writes post response"); - - let (mut poll_socket, _) = listener.accept().await.expect("accepts poll request"); - let poll_request = read_http_headers(&mut poll_socket).await; - let response_body = r#"{"status":"succeeded","analyzeResult":{"pages":[{"pageNumber":1,"lines":[{"content":"ok"}]}]}}"#; - let poll_response = format!( - "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", - response_body.len(), - response_body - ); - poll_socket - .write_all(poll_response.as_bytes()) - .await - .expect("writes poll response"); - (post_request, poll_request) - }); - - let response = ocr( - OcrRequest { - model: "doc-intelligence/prebuilt-read", - document: json!({ - "type": "document_url", - "document_url": "https://example.com/doc.pdf" - }), - optional_params: Map::new(), - }, - &RequestOptions { - api_key: (Some("di-key")).map(|value| value.to_string()), - api_base: (Some(&format!("http://{addr}"))).map(|value| value.to_string()), - custom_llm_provider: (Some("azure_ai")).map(|value| value.to_string()), - extra_headers: None, - timeout: Some(Duration::from_secs(5)), - ..Default::default() - }, - &LiteLlmRequestContext { - attribution: RequestAttribution::default(), - litellm_call_id: None, - ..Default::default() - }, - RequestHooks { - callbacks: Vec::new(), - guardrails: Vec::new(), - }, - ) - .await - .expect("document intelligence request succeeds"); - - assert_eq!(response["pages"][0]["markdown"], "ok"); - - let (post_request, poll_request) = server.await.expect("server task completes"); - assert!( - post_request - .to_ascii_lowercase() - .contains("ocp-apim-subscription-key: di-key"), - "{post_request}" - ); - assert!( - poll_request - .to_ascii_lowercase() - .contains("ocp-apim-subscription-key: di-key"), - "{poll_request}" - ); -} diff --git a/litellm-rust/crates/core/Cargo.toml b/litellm-rust/crates/core/Cargo.toml index c0de7ff3977..1c67abf0770 100644 --- a/litellm-rust/crates/core/Cargo.toml +++ b/litellm-rust/crates/core/Cargo.toml @@ -7,11 +7,14 @@ repository.workspace = true [dependencies] base64.workspace = true +futures-util.workspace = true rand.workspace = true reqwest.workspace = true serde.workspace = true serde_json.workspace = true thiserror.workspace = true +tokio = { workspace = true, features = ["sync"] } +tokio-tungstenite.workspace = true tracing.workspace = true tracing-subscriber = { workspace = true, optional = true } sha2.workspace = true @@ -35,6 +38,7 @@ bedrock-auth = [ observability = ["dep:tracing-subscriber"] [dev-dependencies] +futures-channel = "0.3" rstest.workspace = true tokio = { workspace = true, features = ["macros", "rt-multi-thread"] } tracing-subscriber.workspace = true diff --git a/litellm-rust/crates/core/src/constants.rs b/litellm-rust/crates/core/src/constants.rs index fc81f4fa029..ccd8c62b43b 100644 --- a/litellm-rust/crates/core/src/constants.rs +++ b/litellm-rust/crates/core/src/constants.rs @@ -32,6 +32,18 @@ pub(crate) const CHAT_COMPLETIONS_CONNECT_TIMEOUT_SECS: u64 = 10; pub(crate) const AUDIO_TRANSCRIPTION_TIMEOUT_SECS: u64 = 600; +/// Full-request timeout ceiling for OCR provider calls, in seconds. The +/// per-request timeout from the caller still overrides this on the request +/// builder. +pub(crate) const OCR_TIMEOUT_SECS: u64 = 600; + +/// Connect timeout for the Responses WebSocket upstream dial, in seconds. +pub(crate) const RESPONSES_WS_CONNECT_TIMEOUT_SECS: u64 = 10; + +/// Idle timeout that ends a Responses WebSocket splice when neither side +/// produces an event, in seconds. +pub(crate) const RESPONSES_WS_IDLE_TIMEOUT_SECS: u64 = 300; + /// `object` field every non-streaming chat completion response carries. pub const CHAT_COMPLETION_OBJECT: &str = "chat.completion"; diff --git a/litellm-rust/crates/core/src/ocr/client.rs b/litellm-rust/crates/core/src/ocr/client.rs new file mode 100644 index 00000000000..2181b58cd49 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/client.rs @@ -0,0 +1,14 @@ +use std::sync::OnceLock; +use std::time::Duration; + +use crate::constants::OCR_TIMEOUT_SECS; + +pub(super) fn http_client() -> &'static reqwest::Client { + static CLIENT: OnceLock = OnceLock::new(); + CLIENT.get_or_init(|| { + reqwest::Client::builder() + .timeout(Duration::from_secs(OCR_TIMEOUT_SECS)) + .build() + .unwrap_or_else(|_| reqwest::Client::new()) + }) +} diff --git a/litellm-rust/crates/ai-gateway/src/ocr/common_utils.rs b/litellm-rust/crates/core/src/ocr/common_utils.rs similarity index 70% rename from litellm-rust/crates/ai-gateway/src/ocr/common_utils.rs rename to litellm-rust/crates/core/src/ocr/common_utils.rs index d2be17260a3..42935f232d0 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/common_utils.rs +++ b/litellm-rust/crates/core/src/ocr/common_utils.rs @@ -3,78 +3,16 @@ use std::time::{Duration, Instant}; use base64::Engine; use base64::engine::general_purpose::STANDARD as BASE64_STANDARD; -use litellm_core::error::Error; -use litellm_core::ocr::transformation::OcrProviderConfig; use reqwest::Url; -use serde_json::{Map, Value}; +use serde_json::Value; -use litellm_core::providers::azure_ai::ocr::transformation::{ - AZURE_AI_OCR_CONFIG, AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG, -}; -use litellm_core::providers::mistral::ocr::transformation::MISTRAL_OCR_CONFIG; -use litellm_core::providers::reducto::ocr::transformation as reducto; -use litellm_core::providers::vertex_ai::ocr::transformation as vertex_ai; -use litellm_core::providers::vertex_ai::ocr::transformation::{ - VERTEX_AI_DEEPSEEK_OCR_CONFIG, VERTEX_AI_OCR_CONFIG, -}; +use super::client::http_client; +use crate::Error; +use crate::http_utils::truncate_error_body; -use crate::client::http_client; - -const ERROR_BODY_MAX_CHARS: usize = 256; -const AZURE_DOCUMENT_INTELLIGENCE_POLL_TIMEOUT_SECS: u64 = 120; const DEFAULT_MAX_IMAGE_URL_DOWNLOAD_SIZE_MB: f64 = 50.0; const MAX_SAFE_FETCH_REDIRECTS: usize = 10; - -pub(super) fn truncate_error_body(body: &str) -> String { - if body.chars().count() <= ERROR_BODY_MAX_CHARS { - return body.to_string(); - } - let truncated: String = body.chars().take(ERROR_BODY_MAX_CHARS).collect(); - format!("{truncated}... (truncated)") -} - -#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] -pub(super) fn ocr_provider_config( - provider: &str, - model: &str, -) -> Option<&'static dyn OcrProviderConfig> { - match provider { - "mistral" => Some(&MISTRAL_OCR_CONFIG), - "reducto" => reducto::config_for_model(model), - "azure_ai" if is_azure_document_intelligence_model(model) => { - Some(&AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG) - } - "azure_ai" => Some(&AZURE_AI_OCR_CONFIG), - "vertex_ai" if vertex_ai::is_deepseek_model(model) => Some(&VERTEX_AI_DEEPSEEK_OCR_CONFIG), - "vertex_ai" => Some(&VERTEX_AI_OCR_CONFIG), - _ => None, - } -} - -fn is_azure_document_intelligence_model(model: &str) -> bool { - let model = model.to_ascii_lowercase(); - model.contains("doc-intelligence") || model.contains("documentintelligence") -} - -pub(super) fn string_headers( - extra_headers: Option>, -) -> Result, Error> { - extra_headers - .unwrap_or_default() - .into_iter() - .map(|(key, value)| { - value - .as_str() - .map(|value| (key.clone(), value.to_string())) - .ok_or_else(|| { - Error::InvalidRequest(format!( - "OCR extra_headers.{key} must be a string, got {}", - litellm_core::error::json_type_name(&value) - )) - }) - }) - .collect() -} +const AZURE_DOCUMENT_INTELLIGENCE_POLL_TIMEOUT_SECS: u64 = 120; fn document_url_field(document: &Value) -> Result, Error> { let Some(object) = document.as_object() else { @@ -336,7 +274,6 @@ fn operation_status(response_json: &Value) -> Result<&str, Error> { } } -#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] pub(super) async fn poll_document_intelligence( operation_url: &str, original_url: &str, @@ -395,10 +332,8 @@ pub(super) async fn poll_document_intelligence( #[cfg(test)] mod tests { - use litellm_core::ocr::transformation::OcrResponseHandling; - use serde_json::json; - use super::*; + use serde_json::json; #[test] fn blocks_private_and_metadata_ips() { @@ -443,87 +378,4 @@ mod tests { assert_eq!(transformed, document); } - - #[test] - fn truncate_error_body_passes_short_strings_through() { - let body = "Unauthorized"; - assert_eq!(truncate_error_body(body), "Unauthorized"); - } - - #[test] - fn truncate_error_body_caps_long_payloads() { - let body = "x".repeat(306); - let truncated = truncate_error_body(&body); - - assert!(truncated.ends_with("... (truncated)")); - let prefix_chars = truncated - .strip_suffix("... (truncated)") - .expect("truncated marker present") - .chars() - .count(); - assert_eq!(prefix_chars, 256); - } - - #[test] - fn truncate_error_body_does_not_split_multibyte_chars() { - let body = "é".repeat(266); - let truncated = truncate_error_body(&body); - assert!(truncated.is_char_boundary(truncated.len())); - } - - #[test] - fn ocr_dispatch_supports_migrated_providers() { - assert!(ocr_provider_config("mistral", "mistral-ocr-latest").is_some()); - assert!( - ocr_provider_config("azure_ai", "pixtral-12b-2409") - .expect("azure ai config resolves") - .requires_data_uri_document() - ); - assert_eq!( - ocr_provider_config("azure_ai", "doc-intelligence/prebuilt-read") - .expect("document intelligence config resolves") - .response_handling(), - OcrResponseHandling::AzureDocumentIntelligencePoll - ); - assert!( - ocr_provider_config("vertex_ai", "deepseek-ocr-maas") - .expect("vertex deepseek config resolves") - .supported_ocr_params() - .contains(&"temperature") - ); - assert!(ocr_provider_config("openai", "gpt-4o").is_none()); - } - - #[test] - fn string_headers_accepts_string_values() { - let headers = json!({ - "x-trace-id": "trace-1" - }) - .as_object() - .unwrap() - .clone(); - - assert_eq!( - string_headers(Some(headers)).expect("string headers accepted"), - vec![("x-trace-id".to_string(), "trace-1".to_string())] - ); - } - - #[test] - fn string_headers_rejects_non_string_values() { - let headers = json!({ - "x-retry-count": 3 - }) - .as_object() - .unwrap() - .clone(); - - let err = string_headers(Some(headers)).expect_err("non-string header rejected"); - assert_eq!( - err, - Error::InvalidRequest( - "OCR extra_headers.x-retry-count must be a string, got number".to_string() - ) - ); - } } diff --git a/litellm-rust/crates/core/src/ocr/handler.rs b/litellm-rust/crates/core/src/ocr/handler.rs new file mode 100644 index 00000000000..49138ddb3e2 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/handler.rs @@ -0,0 +1,117 @@ +use reqwest::{RequestBuilder, StatusCode, header::HeaderMap}; +use serde_json::Value; + +use super::client::http_client; +use super::common_utils::poll_document_intelligence; +use super::observers::{OcrObserver, OcrPostCall, OcrPreCall}; +use super::transformation::OcrResponseHandling; +use super::types::{OcrRequestData, ProviderOcrRequest}; +use crate::Error; +use crate::http_utils::{http_request, truncate_error_body}; + +pub struct OcrHttpResponse { + pub status: StatusCode, + pub headers: HeaderMap, + pub body: String, +} + +#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] +pub async fn send_ocr_request( + request: RequestBuilder, + event: &OcrPreCall, + observer: &mut impl OcrObserver, +) -> Result { + if observer.pre_call(event).await.is_err() { + tracing::warn!("OCR pre-call observer failed"); + } + let response = http_request(request).await.map_err(transport_error)?; + let status = response.status(); + let headers = response.headers().clone(); + let body = response.text().await.map_err(transport_error)?; + if !status.is_success() { + return Err(Error::Http { + status: status.as_u16(), + body: truncate_error_body(&body), + }); + } + let event = OcrPostCall { + original_response: body, + }; + if observer.post_call(&event).await.is_err() { + tracing::warn!("OCR post-call observer failed"); + } + Ok(OcrHttpResponse { + status, + headers, + body: event.original_response, + }) +} + +fn transport_error(error: reqwest::Error) -> Error { + Error::Network(if error.is_timeout() { + "Request timed out".into() + } else { + error.to_string() + }) +} + +#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] +pub async fn execute_ocr_provider_call( + request: ProviderOcrRequest, + observer: &mut impl OcrObserver, +) -> Result { + let mut request_builder = http_client().post(request.url()).json(request.body()); + for (key, value) in &request.upstream_headers { + request_builder = request_builder.header(key, value); + } + if let Some(duration) = request.timeout { + request_builder = request_builder.timeout(duration); + } + + let event = OcrPreCall { + model: request.model().to_string(), + request: OcrRequestData { + data: request.body().clone(), + files: None, + }, + api_base: request.url().to_string(), + headers: request.upstream_headers.iter().cloned().collect(), + }; + let response = send_ocr_request(request_builder, &event, observer).await?; + + let status = response.status; + if request.config.response_handling() == OcrResponseHandling::AzureDocumentIntelligencePoll + && status.as_u16() == 202 + { + let operation_url = response + .headers + .get("operation-location") + .and_then(|value| value.to_str().ok()) + .map(str::to_string) + .ok_or_else(|| { + Error::InvalidResponse( + "Azure Document Intelligence returned 202 but no Operation-Location header found" + .to_string(), + ) + })?; + let response_json = poll_document_intelligence( + &operation_url, + request.url(), + &request.upstream_headers, + request.timeout, + ) + .await?; + return Ok(request + .config + .transform_ocr_response(request.model(), response_json)? + .into_json()); + } + + let response_json: Value = serde_json::from_str(&response.body) + .map_err(|err| Error::InvalidResponse(format!("invalid OCR response JSON: {err}")))?; + + Ok(request + .config + .transform_ocr_response(request.model(), response_json)? + .into_json()) +} diff --git a/litellm-rust/crates/core/src/ocr/mod.rs b/litellm-rust/crates/core/src/ocr/mod.rs index ec2fbb969a6..2f71dddac31 100644 --- a/litellm-rust/crates/core/src/ocr/mod.rs +++ b/litellm-rust/crates/core/src/ocr/mod.rs @@ -1,2 +1,45 @@ +mod client; +mod common_utils; +pub mod handler; +pub mod observers; +pub mod prepare; pub mod transformation; pub mod types; + +pub use handler::execute_ocr_provider_call; +pub use prepare::{prepare_ocr_call, prepare_ocr_provider_call}; +pub use types::{OcrRequest, PreparedOcrRequest, ProviderOcrRequest}; + +use serde_json::Value; + +use crate::Error; +use crate::ocr::observers::{NoopOcrObserver, OcrObserver}; + +#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] +pub async fn ocr(request: OcrRequest<'_>) -> Result { + ocr_with_observer(request, &mut NoopOcrObserver).await +} + +#[tracing::instrument( + name = "ocr", + target = "litellm::function_trace", + level = "trace", + skip_all +)] +pub async fn ocr_with_observer( + request: OcrRequest<'_>, + observer: &mut impl OcrObserver, +) -> Result { + let prepared = prepare_ocr_call(request); + let provider_request = prepare_ocr_provider_call(prepared).await?; + execute_ocr_provider_call(provider_request, observer).await +} + +pub fn ocr_admitted(model: &str, provider: &str, request_format: Option<&str>) -> bool { + common_utils::ocr_provider_config(provider, model).is_some_and(|config| { + request_format != Some("native") || config.supported_ocr_params().contains(&"req_format") + }) +} + +#[cfg(test)] +mod tests; diff --git a/litellm-rust/crates/core/src/ocr/observers.rs b/litellm-rust/crates/core/src/ocr/observers.rs new file mode 100644 index 00000000000..4fe050a235a --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/observers.rs @@ -0,0 +1,48 @@ +use std::collections::BTreeMap; +use std::convert::Infallible; + +use serde::Serialize; + +use super::types::OcrRequestData; + +#[derive(Serialize)] +pub struct OcrPreCall { + pub model: String, + pub request: OcrRequestData, + pub api_base: String, + pub headers: BTreeMap, +} + +#[derive(Serialize)] +pub struct OcrPostCall { + pub original_response: String, +} + +#[macro_export] +macro_rules! ocr_observer_catalog { + ($consumer:path, $($options:tt)*) => { + $consumer! { + $($options)* + { + pre_call: PreCall($crate::ocr::observers::OcrPreCall) -> () = direct; + post_call: PostCall($crate::ocr::observers::OcrPostCall) -> () = direct; + } + } + }; +} + +ocr_observer_catalog!(crate::define_hooks, pub trait OcrObserver;); + +pub struct NoopOcrObserver; + +impl OcrObserver for NoopOcrObserver { + type Error = Infallible; + + async fn pre_call(&mut self, _input: &OcrPreCall) -> Result<(), Infallible> { + Ok(()) + } + + async fn post_call(&mut self, _input: &OcrPostCall) -> Result<(), Infallible> { + Ok(()) + } +} diff --git a/litellm-rust/crates/core/src/ocr/prepare.rs b/litellm-rust/crates/core/src/ocr/prepare.rs new file mode 100644 index 00000000000..83081d3580a --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/prepare.rs @@ -0,0 +1,130 @@ +use std::sync::atomic::{AtomicU64, Ordering}; +use std::time::{SystemTime, UNIX_EPOCH}; + +use crate::Error; +use crate::http_utils::string_headers; +use crate::ocr::common_utils::convert_document_url_to_data_uri; +use crate::ocr::transformation::OcrProviderConfig; +use crate::ocr::types::{OcrRequest, PreparedOcrRequest, ProviderOcrRequest}; +use crate::providers::azure_ai::ocr::transformation::{ + AZURE_AI_OCR_CONFIG, AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG, +}; +use crate::providers::mistral::ocr::transformation::MISTRAL_OCR_CONFIG; +use crate::providers::vertex_ai::ocr::transformation as vertex_ai; +use crate::providers::vertex_ai::ocr::transformation::{ + VERTEX_AI_DEEPSEEK_OCR_CONFIG, VERTEX_AI_OCR_CONFIG, +}; +use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider}; + +#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] +pub(crate) fn ocr_provider_config( + provider: &str, + model: &str, +) -> Option<&'static dyn OcrProviderConfig> { + match provider { + "mistral" => Some(&MISTRAL_OCR_CONFIG), + "azure_ai" if is_azure_document_intelligence_model(model) => { + Some(&AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG) + } + "azure_ai" => Some(&AZURE_AI_OCR_CONFIG), + "vertex_ai" if vertex_ai::is_deepseek_model(model) => Some(&VERTEX_AI_DEEPSEEK_OCR_CONFIG), + "vertex_ai" => Some(&VERTEX_AI_OCR_CONFIG), + _ => None, + } +} + +fn is_azure_document_intelligence_model(model: &str) -> bool { + let model = model.to_ascii_lowercase(); + model.contains("doc-intelligence") || model.contains("documentintelligence") +} + +#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] +pub fn prepare_ocr_call(request: OcrRequest<'_>) -> PreparedOcrRequest { + let call_id = request + .litellm_call_id + .map(str::to_string) + .unwrap_or_else(new_ocr_call_id); + let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider) + .unwrap_or(CustomLlmProvider { + model: request.model, + custom_llm_provider: "mistral", + }); + let model = provider_info.model.to_string(); + let custom_llm_provider = provider_info.custom_llm_provider.to_string(); + let config = ocr_provider_config(&custom_llm_provider, &model) + .ok_or_else(|| Error::InvalidProvider(custom_llm_provider.clone())); + let optional_params = match &config { + Ok(config) => { + let supported = config.supported_ocr_params(); + config.map_ocr_params( + &request + .optional_params + .into_iter() + .filter(|(name, _)| supported.contains(&name.as_str())) + .collect(), + ) + } + Err(_) => request.optional_params, + }; + + PreparedOcrRequest { + config, + model, + custom_llm_provider, + litellm_call_id: call_id, + document: request.document, + api_key: request.api_key.map(str::to_string), + api_base: request.api_base.map(str::to_string), + extra_headers: request.extra_headers, + optional_params, + timeout: request.timeout, + } +} + +#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] +pub async fn prepare_ocr_provider_call( + request: PreparedOcrRequest, +) -> Result { + let config = request.config?; + let env_lookup = |key: &str| std::env::var(key).ok(); + let upstream_headers = config.validate_environment( + string_headers("OCR", request.extra_headers)?, + request.api_key.as_deref(), + &env_lookup, + )?; + let url = config.complete_url( + request.api_base.as_deref(), + &request.model, + &request.optional_params, + &env_lookup, + )?; + let model = request.model.clone(); + let custom_llm_provider = request.custom_llm_provider.clone(); + let document = if config.requires_data_uri_document() { + convert_document_url_to_data_uri(request.document).await? + } else { + request.document + }; + let body = config + .transform_ocr_request(&request.model, document, request.optional_params)? + .data; + Ok(ProviderOcrRequest { + model, + custom_llm_provider, + config, + url, + body, + upstream_headers, + timeout: request.timeout, + }) +} + +fn new_ocr_call_id() -> String { + static COUNTER: AtomicU64 = AtomicU64::new(1); + let sequence = COUNTER.fetch_add(1, Ordering::Relaxed); + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|duration| duration.as_nanos()) + .unwrap_or(0); + format!("ocr-{timestamp}-{sequence}") +} diff --git a/litellm-rust/crates/core/src/ocr/tests.rs b/litellm-rust/crates/core/src/ocr/tests.rs new file mode 100644 index 00000000000..ae5b8e96e7c --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/tests.rs @@ -0,0 +1,362 @@ +use std::sync::{Arc, Mutex}; +use std::time::Duration; + +use serde_json::{Map, Value, json}; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; +use tokio::net::{TcpListener, TcpStream}; + +use crate::Error; +use crate::http_utils::has_header; +use crate::ocr::observers::{OcrObserver, OcrPostCall, OcrPreCall}; +use crate::ocr::prepare::ocr_provider_config; +use crate::ocr::transformation::OcrResponseHandling; +use crate::ocr::{OcrRequest, ocr, ocr_with_observer}; + +struct ProviderObserver { + events: Arc>>, + raw_response: Option, + reject: bool, +} + +impl OcrObserver for ProviderObserver { + type Error = &'static str; + + async fn pre_call(&mut self, input: &OcrPreCall) -> Result<(), Self::Error> { + assert_eq!(input.model, "mistral-ocr-4-1"); + assert_eq!( + input.request.data["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"); + if self.reject { + Err("observer failure") + } else { + Ok(()) + } + } + + async fn post_call(&mut self, input: &OcrPostCall) -> Result<(), Self::Error> { + self.events.lock().unwrap().push("post"); + self.raw_response = Some(input.original_response.clone()); + if self.reject { + Err("observer failure") + } else { + Ok(()) + } + } +} + +fn observer_request(api_base: &str) -> OcrRequest<'_> { + OcrRequest { + model: "mistral/mistral-ocr-4-1", + document: json!({"type":"document_url","document_url":"https://example.com/document.pdf"}), + api_key: Some("test-key"), + api_base: Some(api_base), + custom_llm_provider: Some("mistral"), + extra_headers: None, + optional_params: Map::new(), + timeout: Some(Duration::from_secs(2)), + litellm_call_id: Some("observer-test"), + } +} + +async fn observer_case(status: u16, body: &'static str, reject: bool) { + 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 ")); + 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, + reject, + }; + let result = ocr_with_observer(observer_request(&url), &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"]); + 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"], "ok"); + } + } +} + +#[tokio::test] +async fn provider_observers_surround_http_and_cannot_replace_its_outcome() { + for reject in [false, true] { + observer_case(200, r#"{"pages":[{"index":0,"markdown":"ok"}]}"#, reject).await; + observer_case(200, "invalid-json", reject).await; + observer_case(401, r#"{"error":"rejected"}"#, reject).await; + } +} + +#[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, + reject: false, + }; + let request = OcrRequest { + document: json!(42), + ..observer_request(&url) + }; + assert!(ocr_with_observer(request, &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]; + loop { + let n = socket.read(&mut buffer).await.expect("reads request"); + if n == 0 { + break; + } + request.extend_from_slice(&buffer[..n]); + if request.windows(4).any(|window| window == b"\r\n\r\n") { + break; + } + } + String::from_utf8(request).expect("request is utf8") +} + +async fn read_http_request(socket: &mut TcpStream) -> String { + let mut request = Vec::new(); + let mut buffer = [0_u8; 1024]; + let header_end = loop { + let n = socket.read(&mut buffer).await.expect("reads request"); + if n == 0 { + break request.len(); + } + request.extend_from_slice(&buffer[..n]); + if let Some(position) = request.windows(4).position(|window| window == b"\r\n\r\n") { + break position + 4; + } + }; + let headers = String::from_utf8_lossy(&request[..header_end]); + let content_length = headers + .lines() + .find_map(|line| { + let (name, value) = line.split_once(':')?; + name.eq_ignore_ascii_case("content-length") + .then(|| value.trim().parse::().ok()) + .flatten() + }) + .unwrap_or(0); + while request.len().saturating_sub(header_end) < content_length { + let n = socket.read(&mut buffer).await.expect("reads body"); + if n == 0 { + break; + } + request.extend_from_slice(&buffer[..n]); + } + String::from_utf8(request).expect("request is utf8") +} + +#[test] +fn ocr_dispatch_supports_migrated_providers() { + assert!(ocr_provider_config("mistral", "mistral-ocr-latest").is_some()); + assert!( + ocr_provider_config("azure_ai", "pixtral-12b-2409") + .expect("azure ai config resolves") + .requires_data_uri_document() + ); + assert_eq!( + ocr_provider_config("azure_ai", "doc-intelligence/prebuilt-read") + .expect("document intelligence config resolves") + .response_handling(), + OcrResponseHandling::AzureDocumentIntelligencePoll + ); + assert!( + ocr_provider_config("vertex_ai", "deepseek-ocr-maas") + .expect("vertex deepseek config resolves") + .supported_ocr_params() + .contains(&"temperature") + ); + assert!(ocr_provider_config("openai", "gpt-4o").is_none()); +} + +#[test] +fn auth_header_detection_is_case_insensitive() { + let headers = vec![ + ("x-trace-id".to_string(), "trace-1".to_string()), + ("authorization".to_string(), "Bearer sk-test".to_string()), + ]; + + assert!(has_header(&headers, "authorization")); + + let headers = vec![("Authorization".to_string(), "Bearer sk-test".to_string())]; + assert!(has_header(&headers, "authorization")); + + let headers = vec![("x-trace-id".to_string(), "trace-1".to_string())]; + assert!(!has_header(&headers, "authorization")); +} + +#[tokio::test] +async fn ocr_does_not_duplicate_authorization_header_when_header_is_supplied() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("test listener binds"); + let addr = listener.local_addr().expect("listener has local addr"); + + let server = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.expect("accepts one request"); + let request = read_http_headers(&mut socket).await; + let response_body = r#"{"pages":[{"index":0,"markdown":"ok"}],"model":"mistral-ocr-latest","usage_info":{"pages_processed":1}}"#; + let response = format!( + "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", + response_body.len(), + response_body + ); + socket + .write_all(response.as_bytes()) + .await + .expect("writes response"); + request + }); + + let mut headers = Map::new(); + headers.insert( + "Authorization".to_string(), + Value::String("Bearer sk-from-python".to_string()), + ); + headers.insert( + "x-trace-id".to_string(), + Value::String("trace-1".to_string()), + ); + + let response = ocr(OcrRequest { + model: "mistral-ocr-latest", + document: json!({ + "type": "document_url", + "document_url": "https://example.com/doc.pdf" + }), + api_key: Some("sk-for-rust-fallback"), + api_base: Some(&format!("http://{addr}")), + custom_llm_provider: Some("mistral"), + extra_headers: Some(headers), + optional_params: Map::new(), + timeout: Some(Duration::from_secs(5)), + litellm_call_id: None, + }) + .await + .expect("ocr request succeeds"); + + assert_eq!(response["pages"][0]["markdown"], "ok"); + + let request = server.await.expect("server task completes"); + let authorization_count = request + .lines() + .filter(|line| line.to_ascii_lowercase().starts_with("authorization:")) + .count(); + assert_eq!(authorization_count, 1, "{request}"); + assert!( + request.contains("authorization: Bearer sk-from-python") + || request.contains("Authorization: Bearer sk-from-python"), + "{request}" + ); +} + +#[tokio::test] +async fn document_intelligence_poll_uses_resolved_subscription_key() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("test listener binds"); + let addr = listener.local_addr().expect("listener has local addr"); + let operation_url = format!("http://{addr}/operations/1"); + + let server = tokio::spawn(async move { + let (mut post_socket, _) = listener.accept().await.expect("accepts post request"); + let post_request = read_http_headers(&mut post_socket).await; + let post_response = format!( + "HTTP/1.1 202 Accepted\r\noperation-location: {operation_url}\r\ncontent-length: 0\r\nconnection: close\r\n\r\n" + ); + post_socket + .write_all(post_response.as_bytes()) + .await + .expect("writes post response"); + + let (mut poll_socket, _) = listener.accept().await.expect("accepts poll request"); + let poll_request = read_http_headers(&mut poll_socket).await; + let response_body = r#"{"status":"succeeded","analyzeResult":{"pages":[{"pageNumber":1,"lines":[{"content":"ok"}]}]}}"#; + let poll_response = format!( + "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", + response_body.len(), + response_body + ); + poll_socket + .write_all(poll_response.as_bytes()) + .await + .expect("writes poll response"); + (post_request, poll_request) + }); + + let response = ocr(OcrRequest { + model: "doc-intelligence/prebuilt-read", + document: json!({ + "type": "document_url", + "document_url": "https://example.com/doc.pdf" + }), + api_key: Some("di-key"), + api_base: Some(&format!("http://{addr}")), + custom_llm_provider: Some("azure_ai"), + extra_headers: None, + optional_params: Map::new(), + timeout: Some(Duration::from_secs(5)), + litellm_call_id: None, + }) + .await + .expect("document intelligence request succeeds"); + + assert_eq!(response["pages"][0]["markdown"], "ok"); + + let (post_request, poll_request) = server.await.expect("server task completes"); + assert!( + post_request + .to_ascii_lowercase() + .contains("ocp-apim-subscription-key: di-key"), + "{post_request}" + ); + assert!( + poll_request + .to_ascii_lowercase() + .contains("ocp-apim-subscription-key: di-key"), + "{poll_request}" + ); +} diff --git a/litellm-rust/crates/core/src/ocr/types.rs b/litellm-rust/crates/core/src/ocr/types.rs index 71cdb232a87..7c29ebcec14 100644 --- a/litellm-rust/crates/core/src/ocr/types.rs +++ b/litellm-rust/crates/core/src/ocr/types.rs @@ -1,6 +1,80 @@ +use std::time::Duration; + use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; +use crate::Error; +use crate::call_lifecycle::{CallLifecycleContext, CallLifecycleRequest}; +use crate::ocr::transformation::OcrProviderConfig; + +pub struct OcrRequest<'a> { + pub model: &'a str, + pub document: Value, + pub api_key: Option<&'a str>, + pub api_base: Option<&'a str>, + pub custom_llm_provider: Option<&'a str>, + pub extra_headers: Option>, + pub optional_params: Map, + pub timeout: Option, + pub litellm_call_id: Option<&'a str>, +} + +pub struct PreparedOcrRequest { + pub config: Result<&'static dyn OcrProviderConfig, Error>, + pub model: String, + pub custom_llm_provider: String, + pub litellm_call_id: String, + pub document: Value, + pub api_key: Option, + pub api_base: Option, + pub extra_headers: Option>, + pub optional_params: Map, + pub timeout: Option, +} + +impl CallLifecycleRequest for PreparedOcrRequest { + fn lifecycle_context(&self) -> CallLifecycleContext { + CallLifecycleContext::new( + "ocr", + self.model.clone(), + self.custom_llm_provider.clone(), + self.litellm_call_id.clone(), + ) + } +} + +pub struct ProviderOcrRequest { + pub(super) model: String, + pub(super) custom_llm_provider: String, + pub(super) config: &'static dyn OcrProviderConfig, + pub(super) url: String, + pub(super) body: Value, + pub(super) upstream_headers: Vec<(String, String)>, + pub(super) timeout: Option, +} + +impl ProviderOcrRequest { + pub fn model(&self) -> &str { + &self.model + } + + pub fn custom_llm_provider(&self) -> &str { + &self.custom_llm_provider + } + + pub fn url(&self) -> &str { + &self.url + } + + pub fn body(&self) -> &Value { + &self.body + } + + pub fn with_body(self, body: Value) -> Self { + Self { body, ..self } + } +} + #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] pub struct OcrRequestData { pub data: Value, diff --git a/litellm-rust/crates/core/src/responses/connection.rs b/litellm-rust/crates/core/src/responses/connection.rs new file mode 100644 index 00000000000..568964e23df --- /dev/null +++ b/litellm-rust/crates/core/src/responses/connection.rs @@ -0,0 +1,677 @@ +use std::collections::HashMap; +use std::sync::Arc; +use std::time::Duration; + +use futures_util::stream::{SplitSink, SplitStream}; +use futures_util::{Sink, SinkExt, Stream, StreamExt}; +use tokio::net::TcpStream; +use tokio::runtime::Handle; +use tokio::sync::Mutex; +use tokio_tungstenite::tungstenite::Message; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; +use tokio_tungstenite::tungstenite::http::HeaderValue; +use tokio_tungstenite::tungstenite::http::header::{AUTHORIZATION, HeaderName}; +use tokio_tungstenite::{MaybeTlsStream, WebSocketStream, connect_async}; + +use crate::Error; +use crate::constants::{RESPONSES_WS_CONNECT_TIMEOUT_SECS, RESPONSES_WS_IDLE_TIMEOUT_SECS}; +use crate::providers::openai::responses::transformation::OPENAI_RESPONSES_WS_CONFIG; +use crate::responses::types::ResponsesWsEvent; +use crate::responses::websocket::ResponsesWebSocketProviderConfig; + +const OPENAI_API_KEY_ENV: &str = "OPENAI_API_KEY"; +const MISSING_KEY_MESSAGE: &str = "Missing OpenAI API Key - a Responses WebSocket call is being made but no key was passed via params or the OPENAI_API_KEY environment variable"; + +pub type ResponsesUpstreamWs = WebSocketStream>; +type UpstreamTx = SplitSink; +type UpstreamRx = SplitStream; + +/// A connected Responses WebSocket upstream. +/// +/// The send and receive halves are locked separately so a pending +/// [`recv_text`](Self::recv_text) never blocks +/// [`send_text`](Self::send_text). Dropping the last clone closes the +/// upstream socket: the runtime handle captured at connect time is used to +/// flush a close frame without blocking the dropping thread. +#[derive(Clone)] +pub struct ResponsesWebSocketConnection { + tx: Arc>>, + rx: Arc>>, + runtime: Handle, +} + +impl ResponsesWebSocketConnection { + pub async fn connect_url( + url: &str, + headers: &HashMap, + timeout: Option, + ) -> Result { + let mut request = url + .into_client_request() + .map_err(|error| Error::Network(error.to_string()))?; + for (name, value) in headers { + let header_name = name + .parse::() + .map_err(|error| Error::InvalidRequest(error.to_string()))?; + let header_value = HeaderValue::from_str(value) + .map_err(|error| Error::InvalidRequest(error.to_string()))?; + request.headers_mut().insert(header_name, header_value); + } + let runtime = Handle::try_current().map_err(|error| Error::Network(error.to_string()))?; + let connect = connect_async(request); + let result = match timeout { + Some(timeout) => tokio::time::timeout(timeout, connect).await.map_err(|_| { + Error::Network("Responses WebSocket connection timed out".to_string()) + })?, + None => connect.await, + }; + let (socket, _) = result.map_err(|error| match error { + tokio_tungstenite::tungstenite::Error::Http(response) => Error::Http { + status: response.status().as_u16(), + body: String::new(), + }, + other => Error::Network(other.to_string()), + })?; + let (tx, rx) = socket.split(); + Ok(Self { + tx: Arc::new(Mutex::new(Some(tx))), + rx: Arc::new(Mutex::new(Some(rx))), + runtime, + }) + } + + pub async fn send_text(&self, text: String) -> Result<(), Error> { + let mut sender = self.tx.lock().await; + let Some(sender) = sender.as_mut() else { + return Err(Error::Network("Responses WebSocket is closed".to_string())); + }; + sender + .send(Message::Text(text)) + .await + .map_err(|error| Error::Network(error.to_string())) + } + + pub async fn recv_text(&self) -> Result, Error> { + let mut receiver = self.rx.lock().await; + let Some(receiver) = receiver.as_mut() else { + return Ok(None); + }; + match receiver.next().await { + Some(Ok(Message::Text(text))) => Ok(Some(text)), + Some(Ok(Message::Binary(bytes))) => String::from_utf8(bytes.to_vec()) + .map(Some) + .map_err(|error| Error::InvalidResponse(error.to_string())), + Some(Ok(Message::Close(_))) | None => Ok(None), + Some(Ok(_)) => Ok(None), + Some(Err(error)) => Err(Error::Network(error.to_string())), + } + } + + pub async fn close(&self) -> Result<(), Error> { + let mut sender = self.tx.lock().await; + if let Some(mut sender) = sender.take() { + sender + .send(Message::Close(None)) + .await + .map_err(|error| Error::Network(error.to_string()))?; + } + let mut receiver = self.rx.lock().await; + *receiver = None; + Ok(()) + } +} + +impl Drop for ResponsesWebSocketConnection { + fn drop(&mut self) { + if Arc::strong_count(&self.tx) != 1 || Arc::strong_count(&self.rx) != 1 { + return; + } + let tx = Arc::clone(&self.tx); + let rx = Arc::clone(&self.rx); + self.runtime.spawn(async move { + let mut sender = tx.lock().await; + if let Some(mut sender) = sender.take() { + let _ = sender.send(Message::Close(None)).await; + } + let mut receiver = rx.lock().await; + *receiver = None; + }); + } +} + +fn resolve_api_key(api_key: Option<&str>) -> Result { + api_key + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_string) + .or_else(|| { + std::env::var(OPENAI_API_KEY_ENV) + .ok() + .filter(|value| !value.trim().is_empty()) + }) + .ok_or_else(|| Error::Auth(MISSING_KEY_MESSAGE.to_string())) +} + +async fn dial_upstream( + model: &str, + api_key: &str, + api_base: Option<&str>, +) -> Result { + let url = OPENAI_RESPONSES_WS_CONFIG.complete_websocket_url(api_base, model); + let mut request = url + .as_str() + .into_client_request() + .map_err(|error| Error::Network(error.to_string()))?; + request.headers_mut().insert( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {api_key}")) + .map_err(|error| Error::Auth(error.to_string()))?, + ); + let result = tokio::time::timeout( + Duration::from_secs(RESPONSES_WS_CONNECT_TIMEOUT_SECS), + connect_async(request), + ) + .await + .map_err(|_| Error::Network("Responses WebSocket connection timed out".to_string()))?; + result + .map(|(socket, _)| socket) + .map_err(|error| match error { + tokio_tungstenite::tungstenite::Error::Http(response) => Error::Http { + status: response.status().as_u16(), + body: String::new(), + }, + other => Error::Network(other.to_string()), + }) +} + +pub struct ResponsesWebSocketStreaming; + +impl ResponsesWebSocketStreaming { + pub async fn bidirectional_forward( + model: &str, + upstream_tx: UpstreamTx, + upstream_rx: UpstreamRx, + idle_timeout: Option, + observe: impl FnMut(&ResponsesWsEvent) + Send, + client_in: In, + client_out: Out, + ) -> Result<(), Error> + where + In: Stream + Unpin + Send, + Out: Sink + Unpin + Send, + Out::Error: std::fmt::Display, + { + splice( + model, + upstream_tx, + upstream_rx, + idle_timeout, + observe, + client_in, + client_out, + ) + .await + } +} + +async fn splice( + model: &str, + mut upstream_tx: UpstreamTx, + mut upstream_rx: UpstreamRx, + idle_timeout: Option, + mut observe: impl FnMut(&ResponsesWsEvent) + Send, + mut client_in: In, + mut client_out: Out, +) -> Result<(), Error> +where + In: Stream + Unpin + Send, + Out: Sink + Unpin + Send, + Out::Error: std::fmt::Display, +{ + let idle = idle_timeout.unwrap_or_else(|| Duration::from_secs(RESPONSES_WS_IDLE_TIMEOUT_SECS)); + loop { + tokio::select! { + event = client_in.next() => { + let Some(event) = event else { break }; + for outbound in OPENAI_RESPONSES_WS_CONFIG + .transform_ws_request(&event, model)? + .events + { + let payload = serde_json::to_string(&outbound) + .map_err(|error| Error::InvalidResponse(error.to_string()))?; + upstream_tx.send(Message::Text(payload)) + .await + .map_err(|error| Error::Network(error.to_string()))?; + } + } + message = upstream_rx.next() => { + let Some(message) = message else { break }; + match message.map_err(|error| Error::Network(error.to_string()))? { + Message::Text(text) => { + let event = serde_json::from_str::(&text) + .map_err(|error| Error::InvalidResponse(error.to_string()))?; + observe(&event); + for outbound in OPENAI_RESPONSES_WS_CONFIG + .transform_ws_response(&event, model)? + .events + { + client_out.send(outbound) + .await + .map_err(|error| Error::Network(error.to_string()))?; + } + } + Message::Close(_) => break, + _ => {} + } + } + _ = tokio::time::sleep(idle) => break, + } + } + Ok(()) +} + +#[allow(clippy::too_many_arguments)] +pub async fn async_responses_websocket( + model: &str, + api_key: Option<&str>, + api_base: Option<&str>, + first_frame: Option, + idle_timeout: Option, + mut observe: impl FnMut(&ResponsesWsEvent) + Send, + client_in: In, + client_out: Out, +) -> Result<(), Error> +where + In: Stream + Unpin + Send, + Out: Sink + Unpin + Send, + Out::Error: std::fmt::Display, +{ + let key = resolve_api_key(api_key)?; + let upstream = dial_upstream(model, &key, api_base).await?; + let (mut upstream_tx, upstream_rx) = upstream.split(); + if let Some(first_frame) = first_frame { + for outbound in OPENAI_RESPONSES_WS_CONFIG + .transform_ws_request(&first_frame, model)? + .events + { + let payload = serde_json::to_string(&outbound) + .map_err(|error| Error::InvalidResponse(error.to_string()))?; + upstream_tx + .send(Message::Text(payload)) + .await + .map_err(|error| Error::Network(error.to_string()))?; + } + } + ResponsesWebSocketStreaming::bidirectional_forward( + model, + upstream_tx, + upstream_rx, + idle_timeout, + &mut observe, + client_in, + client_out, + ) + .await +} + +#[allow(clippy::too_many_arguments)] +pub async fn responses_ws( + model: &str, + api_key: Option<&str>, + api_base: Option<&str>, + first_frame: Option, + idle_timeout: Option, + observe: impl FnMut(&ResponsesWsEvent) + Send, + client_in: In, + client_out: Out, +) -> Result<(), Error> +where + In: Stream + Unpin + Send, + Out: Sink + Unpin + Send, + Out::Error: std::fmt::Display, +{ + async_responses_websocket( + model, + api_key, + api_base, + first_frame, + idle_timeout, + observe, + client_in, + client_out, + ) + .await +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::responses::types::ResponsesWsEventType; + use futures_channel::mpsc; + use futures_util::{SinkExt, StreamExt}; + use serde_json::json; + use tokio::io::AsyncWriteExt; + use tokio::net::TcpListener; + use tokio_tungstenite::accept_async; + + async fn websocket_base() -> (String, tokio::task::JoinHandle<()>) { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind"); + let address = listener.local_addr().expect("local address"); + let task = tokio::spawn(async move { + let (stream, _) = listener.accept().await.expect("accept"); + let mut socket = accept_async(stream).await.expect("websocket handshake"); + while let Some(Ok(Message::Text(text))) = socket.next().await { + let request: serde_json::Value = serde_json::from_str(&text).expect("request json"); + let model = request + .get("model") + .and_then(serde_json::Value::as_str) + .or_else(|| { + request + .get("response") + .and_then(serde_json::Value::as_object) + .and_then(|response| { + response.get("model").and_then(serde_json::Value::as_str) + }) + }) + .expect("enforced model"); + socket + .send(Message::Text( + json!({ + "type": "response.created", + "response": { + "id": format!("resp-{model}"), + "model": model, + "extra": "preserved" + } + }) + .to_string(), + )) + .await + .expect("created event"); + socket + .send(Message::Text( + json!({ + "type": "response.completed", + "response": { + "id": format!("resp-{model}"), + "model": model, + "usage": { + "input_tokens": 1, + "output_tokens": 2, + "total_tokens": 3 + } + } + }) + .to_string(), + )) + .await + .expect("completed event"); + } + }); + (format!("http://{address}"), task) + } + + fn event(value: serde_json::Value) -> ResponsesWsEvent { + serde_json::from_value(value).expect("event") + } + + #[test] + fn explicit_nonblank_key_wins() { + assert_eq!( + resolve_api_key(Some(" explicit ")).expect("key"), + "explicit" + ); + } + + #[test] + fn blank_key_is_not_accepted_without_environment_key() { + if std::env::var(OPENAI_API_KEY_ENV).is_err() { + assert!(resolve_api_key(Some(" ")).is_err()); + } + } + + #[tokio::test] + async fn forwards_events_sequentially_and_enforces_model() { + let (api_base, server) = websocket_base().await; + let (client_tx, client_rx) = mpsc::unbounded(); + let (output_tx, mut output_rx) = mpsc::unbounded(); + let (observed_tx, observed_rx) = mpsc::unbounded(); + client_tx + .unbounded_send(event(json!({ + "type": "response.create", + "model": "wrong" + }))) + .expect("first request"); + client_tx + .unbounded_send(event(json!({ + "type": "response.create", + "response": {"model": "also-wrong"} + }))) + .expect("second request"); + + let task = tokio::spawn(async move { + responses_ws( + "authorized-model", + Some("test-key"), + Some(&api_base), + None, + Some(Duration::from_secs(1)), + move |event| { + observed_tx + .unbounded_send(event.clone()) + .expect("observe event"); + }, + client_rx, + output_tx, + ) + .await + }); + + let first = output_rx.next().await.expect("first output"); + let second = output_rx.next().await.expect("second output"); + let third = output_rx.next().await.expect("third output"); + let fourth = output_rx.next().await.expect("fourth output"); + drop(client_tx); + task.await.expect("splice task").expect("successful splice"); + server.await.expect("server task"); + + assert_eq!(first.event_type, ResponsesWsEventType::ResponseCreated); + assert_eq!(first.model(), Some("authorized-model")); + assert_eq!(first.data["response"]["extra"], "preserved"); + assert_eq!(second.event_type, ResponsesWsEventType::ResponseCompleted); + assert_eq!(third.event_type, ResponsesWsEventType::ResponseCreated); + assert_eq!(fourth.event_type, ResponsesWsEventType::ResponseCompleted); + let observed: Vec<_> = observed_rx.collect().await; + assert_eq!(observed.len(), 4); + assert!( + observed + .iter() + .all(|event| event.event_type != ResponsesWsEventType::ResponseCreate) + ); + } + + #[tokio::test] + async fn idle_timeout_ends_without_upstream_events() { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind"); + let address = listener.local_addr().expect("address"); + let server = tokio::spawn(async move { + let (stream, _) = listener.accept().await.expect("accept"); + let _socket = accept_async(stream).await.expect("handshake"); + tokio::time::sleep(Duration::from_secs(1)).await; + }); + let (_client_tx, client_rx) = mpsc::unbounded::(); + let (output_tx, mut output_rx) = mpsc::unbounded(); + let result = responses_ws( + "model", + Some("key"), + Some(&format!("http://{address}")), + None, + Some(Duration::from_millis(20)), + |_| {}, + client_rx, + output_tx, + ) + .await; + assert!(result.is_ok()); + assert!(output_rx.next().await.is_none()); + server.abort(); + } + + #[tokio::test] + async fn dial_http_status_is_preserved() { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind"); + let address = listener.local_addr().expect("address"); + let server = tokio::spawn(async move { + let (mut stream, _) = listener.accept().await.expect("accept"); + stream + .write_all(b"HTTP/1.1 401 Unauthorized\r\nContent-Length: 0\r\n\r\n") + .await + .expect("response"); + }); + let (_client_tx, client_rx) = mpsc::unbounded::(); + let (output_tx, _output_rx) = mpsc::unbounded(); + let error = responses_ws( + "model", + Some("key"), + Some(&format!("http://{address}")), + None, + Some(Duration::from_millis(20)), + |_| {}, + client_rx, + output_tx, + ) + .await + .expect_err("status error"); + assert!(matches!(error, Error::Http { status: 401, .. })); + server.await.expect("server task"); + } + + #[tokio::test] + async fn dial_http_500_status_is_preserved() { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind"); + let address = listener.local_addr().expect("address"); + let server = tokio::spawn(async move { + let (mut stream, _) = listener.accept().await.expect("accept"); + stream + .write_all(b"HTTP/1.1 500 Internal Server Error\r\nContent-Length: 0\r\n\r\n") + .await + .expect("response"); + }); + let (_client_tx, client_rx) = mpsc::unbounded::(); + let (output_tx, _output_rx) = mpsc::unbounded(); + let error = responses_ws( + "model", + Some("key"), + Some(&format!("http://{address}")), + None, + Some(Duration::from_millis(20)), + |_| {}, + client_rx, + output_tx, + ) + .await + .expect_err("status error"); + assert!(matches!(error, Error::Http { status: 500, .. })); + server.await.expect("server task"); + } + + #[tokio::test] + async fn dropping_the_last_connection_closes_the_upstream_socket() { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind"); + let address = listener.local_addr().expect("address"); + let server = tokio::spawn(async move { + let (stream, _) = listener.accept().await.expect("accept"); + let mut socket = accept_async(stream).await.expect("handshake"); + let message = socket.next().await.expect("close frame").expect("frame"); + assert!(matches!(message, Message::Close(_))); + }); + let connection = ResponsesWebSocketConnection::connect_url( + &format!("ws://{address}"), + &HashMap::new(), + None, + ) + .await + .expect("connect"); + drop(connection); + tokio::time::timeout(Duration::from_secs(2), server) + .await + .expect("close frame after drop") + .expect("server task"); + } + + #[tokio::test] + async fn dropping_one_clone_leaves_the_socket_open_until_the_last_drop() { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind"); + let address = listener.local_addr().expect("address"); + let server = tokio::spawn(async move { + let (stream, _) = listener.accept().await.expect("accept"); + let mut socket = accept_async(stream).await.expect("handshake"); + let early = tokio::time::timeout(Duration::from_millis(150), socket.next()).await; + assert!(early.is_err(), "no close frame while a clone is alive"); + let message = socket + .next() + .await + .expect("close frame after last drop") + .expect("frame"); + assert!(matches!(message, Message::Close(_))); + }); + let connection = ResponsesWebSocketConnection::connect_url( + &format!("ws://{address}"), + &HashMap::new(), + None, + ) + .await + .expect("connect"); + let clone = connection.clone(); + drop(connection); + tokio::time::sleep(Duration::from_millis(250)).await; + drop(clone); + tokio::time::timeout(Duration::from_secs(2), server) + .await + .expect("server completes") + .expect("server task"); + } + + #[tokio::test] + async fn send_text_is_not_blocked_by_a_pending_recv_text() { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind"); + let address = listener.local_addr().expect("address"); + let server = tokio::spawn(async move { + let (stream, _) = listener.accept().await.expect("accept"); + let mut socket = accept_async(stream).await.expect("handshake"); + let message = socket.next().await.expect("send frame").expect("frame"); + assert_eq!(message, Message::Text("ping".into())); + socket + .send(Message::Text("pong".into())) + .await + .expect("reply"); + assert!(matches!(socket.next().await, Some(Ok(Message::Close(_))))); + }); + let connection = ResponsesWebSocketConnection::connect_url( + &format!("ws://{address}"), + &HashMap::new(), + None, + ) + .await + .expect("connect"); + let pending = tokio::spawn({ + let connection = connection.clone(); + async move { connection.recv_text().await } + }); + tokio::time::sleep(Duration::from_millis(100)).await; + tokio::time::timeout(Duration::from_secs(1), connection.send_text("ping".into())) + .await + .expect("send completes while recv is pending") + .expect("send succeeds"); + let received = tokio::time::timeout(Duration::from_secs(1), pending) + .await + .expect("recv completes") + .expect("recv task"); + assert_eq!(received.expect("recv succeeds"), Some("pong".to_string())); + connection.close().await.expect("close"); + tokio::time::timeout(Duration::from_secs(2), server) + .await + .expect("server completes") + .expect("server task"); + } +} diff --git a/litellm-rust/crates/core/src/responses/mod.rs b/litellm-rust/crates/core/src/responses/mod.rs index 5ec5a2caef8..8713932a954 100644 --- a/litellm-rust/crates/core/src/responses/mod.rs +++ b/litellm-rust/crates/core/src/responses/mod.rs @@ -1,3 +1,4 @@ +pub mod connection; pub mod instrumentation; pub mod types; pub mod websocket; diff --git a/litellm-rust/crates/core/tests/workspace_crate_allowlist.rs b/litellm-rust/crates/core/tests/workspace_crate_allowlist.rs index fc0ab2b62a3..b9cd2cb935f 100644 --- a/litellm-rust/crates/core/tests/workspace_crate_allowlist.rs +++ b/litellm-rust/crates/core/tests/workspace_crate_allowlist.rs @@ -110,3 +110,73 @@ fn crates_directory_matches_allowlist() { let expected: BTreeSet = EXPECTED_CRATE_DIRS.iter().map(|s| s.to_string()).collect(); assert_eq!(actual, expected, "{MISMATCH}"); } + +/// Parse the dependency names out of a crate manifest's `[dependencies]` +/// table. +/// +/// Same hand-rolled approach as [`parse_members`]: take the lines after the +/// `[dependencies]` header up to the next table header (a line starting with +/// `[`), then keep the text before ` = ` on each non-comment line, trimmed of +/// the `.workspace`-style shorthand suffix. +fn parse_dependencies(manifest: &str) -> BTreeSet { + let Some((_, after_table)) = manifest.split_once("[dependencies]") else { + return BTreeSet::new(); + }; + + after_table + .lines() + .map(str::trim) + .take_while(|line| !line.starts_with('[')) + .filter_map(|line| { + let (name, _) = line.split_once(" = ")?; + let name = name.split('.').next().unwrap_or(name); + (!name.is_empty() && !name.starts_with('#')).then(|| name.to_string()) + }) + .collect() +} + +fn crate_manifest(root: &Path, crate_dir: &str) -> String { + fs::read_to_string(root.join("crates").join(crate_dir).join("Cargo.toml")) + .unwrap_or_else(|error| panic!("{crate_dir}/Cargo.toml should be readable: {error}")) +} + +/// The python bridge depends on the domain layers, never on the gateway: the +/// bridge and the axum server are alternative hosts over `litellm-core`, so a +/// bridge -> gateway edge would drag the server crate into every cdylib build +/// and let provider I/O creep back out of core. +#[test] +fn python_bridge_dependencies_stay_on_core_and_interop() { + let manifest = crate_manifest(&workspace_root(), "python-bridge"); + let dependencies = parse_dependencies(&manifest); + + assert!( + dependencies.contains("litellm-core"), + "python-bridge must depend on litellm-core, got {dependencies:?}" + ); + assert!( + dependencies.contains("litellm-python-interop"), + "python-bridge must depend on litellm-python-interop, got {dependencies:?}" + ); + assert!( + !dependencies.contains("litellm-ai-gateway"), + "python-bridge must not depend on litellm-ai-gateway (it belongs behind the \ + gateway's own host surface, not the cdylib), got {dependencies:?}" + ); +} + +/// `litellm-core` stays a pure Rust SDK: Python bindings (pyo3, pythonize, +/// pyo3-async-runtimes) live in `litellm-python-interop` / +/// `litellm-python-bridge`, never in core. +#[test] +fn core_dependencies_stay_python_free() { + let manifest = crate_manifest(&workspace_root(), "core"); + let dependencies = parse_dependencies(&manifest); + + for banned in ["pyo3", "pyo3-async-runtimes", "pythonize"] { + assert!( + !dependencies.contains(banned), + "litellm-core must not depend on {banned}; Python binding crates own it, \ + got {dependencies:?}" + ); + } +} diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index 014346ac5a9..bee11ee1e4d 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -14,17 +14,12 @@ default = ["abi3"] abi3 = ["pyo3/abi3-py310"] extension-module = ["pyo3/extension-module"] panic-test = [] -trace-parity = [ - "dep:tracing", - "litellm-core/observability", - "litellm-ai-gateway/trace-parity", -] +trace-parity = ["litellm-core/observability"] [dependencies] futures-util.workspace = true -tracing = { workspace = true, optional = true } +tracing.workspace = true litellm-core = { workspace = true, features = ["bedrock-auth"] } -litellm-ai-gateway = { workspace = true, default-features = false } litellm-python-interop.workspace = true pyo3.workspace = true pyo3-async-runtimes.workspace = true diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index f123c322f8d..8b291fd7a7f 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -9,14 +9,14 @@ mod execution; #[cfg(feature = "trace-parity")] mod function_trace; mod marshal; +mod ocr_callbacks; 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 litellm_core::responses::connection::ResponsesWebSocketConnection as RustResponsesWebSocketConnection; use pyo3::prelude::*; use pyo3::types::PyAny; @@ -65,7 +65,15 @@ impl ResponsesWebSocketConnection { let mut observer = callback_adapter .map(|adapter| crate::callback_bindings::python_async_session(adapter, py)) .transpose()?; - let request = ResponsesWebSocketRequest { url: request.url }; + let headers = litellm_core::http_utils::string_headers( + "Responses WebSocket", + options.extra_headers.clone(), + ) + .map_err(core_error_to_pyerr)? + .into_iter() + .collect(); + let url = request.url; + let timeout = options.timeout; pyo3_async_runtimes::tokio::future_into_py(py, async move { if let Some(observer) = observer.as_mut() { let decision = observer @@ -83,9 +91,7 @@ impl ResponsesWebSocketConnection { } } } - let inner = match RustResponsesWebSocketConnection::connect(request, &options, &context) - .await - { + let inner = match RustResponsesWebSocketConnection::connect_url(&url, &headers, timeout).await { Ok(inner) => inner, Err(error) => { if let Some(observer) = observer.as_mut() { @@ -153,6 +159,7 @@ mod _native { litellm_python_interop::callback_runtime::register(module)?; super::callback_bindings::register(module)?; + super::ocr_callbacks::register(module)?; super::errors::register(module)?; let ready_endpoints = PyDict::new(module.py()); module.add("ready_endpoints", ready_endpoints)?; diff --git a/litellm-rust/crates/python-bridge/src/ocr_callbacks.rs b/litellm-rust/crates/python-bridge/src/ocr_callbacks.rs new file mode 100644 index 00000000000..74dfc0f1c72 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/ocr_callbacks.rs @@ -0,0 +1,82 @@ +use std::num::NonZeroUsize; + +use litellm_core::ocr::observers::{OcrObserver, OcrPostCall, OcrPreCall}; +use litellm_python_interop::callback_runtime::{AsyncContext, CallbackRuntime, SyncContext}; +use pyo3::prelude::*; + +use crate::constants::OCR_CALLBACK_CAPACITY; +use crate::execution::PythonCallContext; + +litellm_core::ocr_observer_catalog!(crate::bind_python_hooks, + pub(crate) struct PythonOcrSession; + trait OcrObserver; +); + +#[pyclass(frozen)] +struct OcrRuntime(CallbackRuntime); + +pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { + let capacity = NonZeroUsize::new(OCR_CALLBACK_CAPACITY) + .expect("OCR callback capacity is a positive constant"); + module.add( + "__ocr_callback_runtime__", + OcrRuntime(CallbackRuntime::new(module, capacity)?), + ) +} + +pub(crate) enum PythonOcrObserver { + Disabled, + Sync(PythonOcrSession), + Async(PythonOcrSession), +} + +impl PythonOcrObserver { + 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("__ocr_callback_runtime__")? + .extract::>()? + .0 + .clone(); + if context.asynchronous { + Ok(Self::Async(PythonOcrSession::new( + adapter.bind(py), + runtime.async_context(py)?, + )?)) + } else { + Ok(Self::Sync(PythonOcrSession::new( + adapter.bind(py), + runtime.sync_context(py)?, + )?)) + } + } +} + +impl OcrObserver for PythonOcrObserver { + type Error = PyErr; + + #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] + async fn pre_call(&mut self, input: &OcrPreCall) -> PyResult<()> { + match self { + Self::Disabled => Ok(()), + Self::Sync(session) => session.pre_call(input).await, + Self::Async(session) => session.pre_call(input).await, + } + } + + #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] + async fn post_call(&mut self, input: &OcrPostCall) -> PyResult<()> { + match self { + Self::Disabled => Ok(()), + Self::Sync(session) => session.post_call(input).await, + Self::Async(session) => session.post_call(input).await, + } + } +} diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr.rs b/litellm-rust/crates/python-bridge/src/routes/ocr.rs index 5236d872346..ae04a501101 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr.rs @@ -1,11 +1,10 @@ -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_with_observer as run_route; +use crate::marshal::{NativeRequestContext, NativeRequestOptions, required_value}; +use crate::ocr_callbacks::PythonOcrObserver; use litellm_core::Error; +use litellm_core::ocr::{OcrRequest, ocr_with_observer}; use litellm_core::request_context::LiteLlmRequestContext; +use litellm_core::request_options::RequestOptions; use pyo3::prelude::*; use serde_json::{Map, Value}; use std::future::Future; @@ -26,7 +25,7 @@ fn prepare_ocr( python_context: crate::execution::PythonCallContext<'_>, ) -> PyResult> + Send + 'static> { let context: LiteLlmRequestContext = context.into(); - let provider_admitted = litellm_ai_gateway::io::ocr::ocr_admitted( + let provider_admitted = litellm_core::ocr::ocr_admitted( &input.model, options.provider("mistral"), context.capabilities.request_format.as_deref(), @@ -52,19 +51,22 @@ fn prepare_ocr( input.document }; let document: Value = litellm_python_interop::from_py(document.bind(py))?; - let mut observer = PythonProviderObserver::new(callback_adapter, python_context)?; + let document = required_value("document", document, Value::is_object, "dict")?; + let call_id = context.litellm_call_id.clone(); + let options: RequestOptions = options.into(); + let mut observer = PythonOcrObserver::new(callback_adapter, python_context)?; Ok(async move { - run_route( + ocr_with_observer( OcrRequest { model: &input.model, document, + api_key: options.api_key.as_deref(), + api_base: options.api_base.as_deref(), + custom_llm_provider: options.custom_llm_provider.as_deref(), + extra_headers: options.extra_headers, optional_params: input.optional_params, - }, - &options.into(), - &context, - RequestHooks { - callbacks: Vec::new(), - guardrails: Vec::new(), + timeout: options.timeout, + litellm_call_id: call_id.as_deref(), }, &mut observer, )