diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 55bb5318cfa..a52f379f98b 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -1448,14 +1448,20 @@ dependencies = [ "aws-smithy-runtime-api", "aws-types", "base64", + "bytes", + "futures-channel", + "futures-util", "rand 0.8.7", "reqwest", "rstest", + "rustls 0.23.42", + "rustls-native-certs", "serde", "serde_json", "sha2 0.10.9", "thiserror 2.0.19", "tokio", + "tokio-tungstenite", "tracing", "tracing-subscriber", ] @@ -1465,7 +1471,6 @@ name = "litellm-python-bridge" version = "0.1.0" dependencies = [ "criterion", - "futures-util", "litellm-ai-gateway", "litellm-core", "litellm-python-interop", @@ -1474,7 +1479,6 @@ dependencies = [ "serde", "serde_json", "tokio", - "tokio-tungstenite", "tracing", ] diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 441ad33583f..b33fed58fca 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -40,6 +40,7 @@ tokio = { version = "1", features = ["rt-multi-thread", "macros", "time", "net"] tokio-tungstenite = { version = "0.24", default-features = false, features = ["connect", "rustls-tls-native-roots"] } futures-util = { version = "0.3", default-features = false, features = ["sink", "std"] } base64 = "0.22" +bytes = "1" [profile.release] opt-level = 3 diff --git a/litellm-rust/crates/ai-gateway/Cargo.toml b/litellm-rust/crates/ai-gateway/Cargo.toml index 74cf66e88a2..f1834922c84 100644 --- a/litellm-rust/crates/ai-gateway/Cargo.toml +++ b/litellm-rust/crates/ai-gateway/Cargo.toml @@ -20,13 +20,8 @@ litellm-config.workspace = true # reqwest (rustls + json) is used by io/ocr and ships realtime logs to the # Python proxy callbacks API. reqwest.workspace = true -# rustls and its root store are direct dependencies so `io::tls` can build the -# one TLS config the outbound dials use; see that module for why it has to. -rustls.workspace = true -rustls-native-certs.workspace = true # `sync` powers the bounded mpsc channel the realtime logger drains. tokio = { workspace = true, features = ["rt-multi-thread", "macros", "net", "time", "sync"] } -tokio-tungstenite.workspace = true futures-util.workspace = true serde_json.workspace = true base64.workspace = true diff --git a/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs b/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs deleted file mode 100644 index 4336cd6e3a6..00000000000 --- a/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs +++ /dev/null @@ -1,212 +0,0 @@ -use litellm_core::audio_transcription::{ - AudioTranscriptionRequest as CoreAudioTranscriptionRequest, ProviderAudioTranscriptionRequest, - prepare_audio_transcription_provider_call, -}; -use litellm_core::error::Error; -use litellm_core::lifecycle::{ - ActionResult, CallLifecycleContext, RequestPolicy, TerminalDispatcher, TerminalRecord, -}; -use serde_json::{Map, Value, json}; -use std::future::Future; -use std::pin::Pin; - -use super::types::PreparedAudioTranscriptionRequest; -use litellm_core::integrations::custom_guardrail::{ - CustomGuardrailRunner, GuardrailContext, GuardrailError, GuardrailRequest, -}; -use litellm_core::integrations::custom_logger::{CallType, CustomLoggerRunner, LogFuture}; -use litellm_core::integrations::types::RequestMetadata; - -pub(crate) struct AudioTranscriptionLifecycleHooks { - logger_runner: CustomLoggerRunner, - guardrail_runner: CustomGuardrailRunner, - request_metadata: RequestMetadata, -} - -type AudioFuture<'a, T> = Pin> + Send + 'a>>; - -impl AudioTranscriptionLifecycleHooks { - pub(crate) fn new( - logger_runner: CustomLoggerRunner, - guardrail_runner: CustomGuardrailRunner, - request_metadata: RequestMetadata, - ) -> Self { - Self { - logger_runner, - guardrail_runner, - request_metadata, - } - } - - async fn run_pre_call_guardrails( - &self, - request: PreparedAudioTranscriptionRequest, - ) -> Result { - if self.guardrail_runner.is_empty() { - return Ok(request); - } - let (guardrail_request, _) = self - .guardrail_runner - .run_pre_call( - &guardrail_context(&self.request_metadata), - GuardrailRequest::new(json!({ - "model": request.model, - "custom_llm_provider": request.custom_llm_provider, - "audio": request.audio, - "optional_params": request.optional_params, - })), - ) - .await - .map_err(guardrail_error_to_core_error)?; - let Value::Object(mut data) = guardrail_request.data else { - return Err(Error::InvalidRequest( - "audio transcription pre_call guardrail must return an object".to_string(), - )); - }; - let audio = data.remove("audio").ok_or_else(|| { - Error::InvalidRequest("audio transcription guardrail removed audio".to_string()) - })?; - let optional_params = match data.remove("optional_params") { - Some(Value::Object(value)) => value, - Some(_) => { - return Err(Error::InvalidRequest( - "audio transcription optional_params must be an object".to_string(), - )); - } - None => Map::new(), - }; - Ok(PreparedAudioTranscriptionRequest { - audio, - optional_params, - ..request - }) - } - - async fn prepare_provider_request( - &self, - request: PreparedAudioTranscriptionRequest, - ) -> Result { - let PreparedAudioTranscriptionRequest { - model, - custom_llm_provider, - audio, - api_key, - api_base, - extra_headers, - optional_params, - timeout, - .. - } = request; - let provider_request = - prepare_audio_transcription_provider_call(CoreAudioTranscriptionRequest { - model: &model, - audio, - api_key: api_key.as_deref(), - api_base: api_base.as_deref(), - custom_llm_provider: Some(&custom_llm_provider), - extra_headers, - optional_params, - timeout, - })?; - self.run_during_call_guardrails(provider_request).await - } - - async fn run_during_call_guardrails( - &self, - request: ProviderAudioTranscriptionRequest, - ) -> Result { - if self.guardrail_runner.is_empty() { - return Ok(request); - } - let (guardrail_request, _) = self - .guardrail_runner - .run_during_call( - &guardrail_context(&self.request_metadata), - GuardrailRequest::new(json!({ - "model": request.model(), - "custom_llm_provider": request.custom_llm_provider(), - "url": request.url(), - "body": request.body(), - })), - ) - .await - .map_err(guardrail_error_to_core_error)?; - let Value::Object(mut data) = guardrail_request.data else { - return Err(Error::InvalidRequest( - "audio transcription during_call guardrail must return an object".to_string(), - )); - }; - let body = data.remove("body").ok_or_else(|| { - Error::InvalidRequest("audio transcription guardrail removed body".to_string()) - })?; - Ok(request.with_body(body)) - } -} - -impl RequestPolicy - for AudioTranscriptionLifecycleHooks -{ - type PreCallFuture<'a> = AudioFuture<'a, PreparedAudioTranscriptionRequest>; - type DuringCallFuture<'a> = AudioFuture<'a, ProviderAudioTranscriptionRequest>; - - fn async_pre_call_hook<'a>( - &'a self, - _context: &'a CallLifecycleContext, - request: PreparedAudioTranscriptionRequest, - ) -> Self::PreCallFuture<'a> { - Box::pin(async move { - match self.run_pre_call_guardrails(request).await { - Ok(request) => ActionResult::Replace(request), - Err(error) => ActionResult::Reject(error), - } - }) - } - - fn async_during_call_hook<'a>( - &'a self, - _context: &'a CallLifecycleContext, - request: PreparedAudioTranscriptionRequest, - ) -> Self::DuringCallFuture<'a> { - Box::pin(async move { - match self.prepare_provider_request(request).await { - Ok(request) => ActionResult::Replace(request), - Err(error) => ActionResult::Reject(error), - } - }) - } -} - -impl TerminalDispatcher for AudioTranscriptionLifecycleHooks { - fn dispatch<'a>(&'a self, terminal: &'a TerminalRecord) -> LogFuture<'a> { - let mut terminal = terminal.clone(); - terminal.cost_inputs.metadata = request_metadata(&self.request_metadata); - Box::pin(async move { self.logger_runner.dispatch(&terminal).await }) - } -} - -fn request_metadata( - metadata: &RequestMetadata, -) -> litellm_core::integrations::types::StandardLoggingMetadata { - litellm_core::integrations::types::StandardLoggingMetadata { - user_api_key_hash: metadata.user_api_key_hash.clone(), - user_api_key_user_id: metadata.user_api_key_user_id.clone(), - user_api_key_team_id: metadata.user_api_key_team_id.clone(), - ..Default::default() - } -} - -fn guardrail_context(metadata: &RequestMetadata) -> GuardrailContext { - GuardrailContext { - call_type: CallType::Other("audio_transcription".to_string()), - selected_guardrails: Vec::new(), - metadata: std::collections::HashMap::new(), - user_api_key_hash: metadata.user_api_key_hash.clone(), - user_api_key_user_id: metadata.user_api_key_user_id.clone(), - user_api_key_team_id: metadata.user_api_key_team_id.clone(), - trace_parent: None, - } -} - -fn guardrail_error_to_core_error(error: GuardrailError) -> Error { - Error::InvalidRequest(format!("{}: {}", error.kind, error.message)) -} diff --git a/litellm-rust/crates/ai-gateway/src/audio_transcription/mod.rs b/litellm-rust/crates/ai-gateway/src/audio_transcription/mod.rs index b0d6c962dcf..06d53eacf59 100644 --- a/litellm-rust/crates/ai-gateway/src/audio_transcription/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/mod.rs @@ -1,31 +1,44 @@ use litellm_core::Error; -use litellm_core::audio_transcription::execute_audio_transcription_provider_call; -use litellm_core::lifecycle::{CallLifecycle, CallLifecycleRequest, SystemClock}; +use litellm_core::audio_transcription::{AudioRoute, AudioRouteRequest, DefaultAudioServices}; use serde_json::Value; -mod hooks; -mod prepare; mod types; pub use types::AudioTranscriptionRequest; -use prepare::{PreparedAudioTranscriptionCall, prepare_audio_transcription_call}; - pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> Result { - let PreparedAudioTranscriptionCall { request, hooks } = - prepare_audio_transcription_call(request); - let context = request.lifecycle_context(); - CallLifecycle - .run( - context, - request, - &hooks, - &hooks, - &SystemClock, - execute_audio_transcription_provider_call, - ) - .await - .into_result() + let AudioTranscriptionRequest { + model, + audio, + api_key, + api_base, + custom_llm_provider, + extra_headers, + optional_params, + timeout, + callbacks, + guardrails, + request_metadata, + litellm_call_id, + } = request; + let services = DefaultAudioServices::new(callbacks, guardrails); + AudioRoute::execute( + &services, + AudioRouteRequest { + model, + audio, + api_key, + api_base, + custom_llm_provider, + extra_headers, + optional_params, + timeout, + request_metadata, + litellm_call_id, + }, + ) + .await + .into_result() } #[cfg(test)] diff --git a/litellm-rust/crates/ai-gateway/src/audio_transcription/prepare.rs b/litellm-rust/crates/ai-gateway/src/audio_transcription/prepare.rs deleted file mode 100644 index 93113808f98..00000000000 --- a/litellm-rust/crates/ai-gateway/src/audio_transcription/prepare.rs +++ /dev/null @@ -1,55 +0,0 @@ -use std::sync::atomic::{AtomicU64, Ordering}; -use std::time::{SystemTime, UNIX_EPOCH}; - -use litellm_core::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider}; - -use super::hooks::AudioTranscriptionLifecycleHooks; -use super::types::{AudioTranscriptionRequest, PreparedAudioTranscriptionRequest}; -use litellm_core::integrations::custom_guardrail::CustomGuardrailRunner; -use litellm_core::integrations::custom_logger::CustomLoggerRunner; - -pub(crate) struct PreparedAudioTranscriptionCall { - pub(crate) request: PreparedAudioTranscriptionRequest, - pub(crate) hooks: AudioTranscriptionLifecycleHooks, -} - -pub(crate) fn prepare_audio_transcription_call( - request: AudioTranscriptionRequest<'_>, -) -> PreparedAudioTranscriptionCall { - let call_id = request - .litellm_call_id - .map(str::to_string) - .unwrap_or_else(new_audio_transcription_call_id); - let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider) - .unwrap_or(CustomLlmProvider { - model: request.model, - custom_llm_provider: "bedrock", - }); - PreparedAudioTranscriptionCall { - request: PreparedAudioTranscriptionRequest { - model: provider_info.model.to_string(), - custom_llm_provider: provider_info.custom_llm_provider.to_string(), - litellm_call_id: call_id, - audio: request.audio, - 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: request.optional_params, - timeout: request.timeout, - }, - hooks: AudioTranscriptionLifecycleHooks::new( - CustomLoggerRunner::new(request.callbacks), - CustomGuardrailRunner::new(request.guardrails), - request.request_metadata, - ), - } -} - -fn new_audio_transcription_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_or(0, |duration| duration.as_nanos()); - format!("audio-transcription-{timestamp}-{sequence}") -} diff --git a/litellm-rust/crates/ai-gateway/src/audio_transcription/types.rs b/litellm-rust/crates/ai-gateway/src/audio_transcription/types.rs index ee715e60e01..f8581f86c89 100644 --- a/litellm-rust/crates/ai-gateway/src/audio_transcription/types.rs +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/types.rs @@ -1,12 +1,10 @@ use std::sync::Arc; use std::time::Duration; -use litellm_core::lifecycle::{CallLifecycleContext, CallLifecycleRequest}; -use serde_json::{Map, Value}; - use litellm_core::integrations::custom_guardrail::CustomGuardrail; use litellm_core::integrations::custom_logger::CustomLogger; use litellm_core::integrations::types::RequestMetadata; +use serde_json::{Map, Value}; pub struct AudioTranscriptionRequest<'a> { pub model: &'a str, @@ -22,26 +20,3 @@ pub struct AudioTranscriptionRequest<'a> { pub request_metadata: RequestMetadata, pub litellm_call_id: Option<&'a str>, } - -pub(crate) struct PreparedAudioTranscriptionRequest { - pub(crate) model: String, - pub(crate) custom_llm_provider: String, - pub(crate) litellm_call_id: String, - pub(crate) audio: Value, - pub(crate) api_key: Option, - pub(crate) api_base: Option, - pub(crate) extra_headers: Option>, - pub(crate) optional_params: Map, - pub(crate) timeout: Option, -} - -impl CallLifecycleRequest for PreparedAudioTranscriptionRequest { - fn lifecycle_context(&self) -> CallLifecycleContext { - CallLifecycleContext::new( - "audio_transcription", - self.model.clone(), - self.custom_llm_provider.clone(), - self.litellm_call_id.clone(), - ) - } -} diff --git a/litellm-rust/crates/ai-gateway/src/constants.rs b/litellm-rust/crates/ai-gateway/src/constants.rs index b7e317237b8..1ed44cf7048 100644 --- a/litellm-rust/crates/ai-gateway/src/constants.rs +++ b/litellm-rust/crates/ai-gateway/src/constants.rs @@ -9,9 +9,6 @@ #[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/mod.rs b/litellm-rust/crates/ai-gateway/src/io/mod.rs index fc51e45c8ca..90c7a65ddf8 100644 --- a/litellm-rust/crates/ai-gateway/src/io/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/io/mod.rs @@ -1,5 +1,3 @@ pub mod audio_transcription; -pub mod realtime; pub mod realtime_pool; -pub mod responses_ws; pub(crate) mod tls; diff --git a/litellm-rust/crates/ai-gateway/src/io/realtime.rs b/litellm-rust/crates/ai-gateway/src/io/realtime.rs deleted file mode 100644 index 207c31dffa0..00000000000 --- a/litellm-rust/crates/ai-gateway/src/io/realtime.rs +++ /dev/null @@ -1,418 +0,0 @@ -//! End-to-end OpenAI realtime invocation. -//! -//! The host-facing entry point opens the WebSocket to OpenAI, then splices a -//! client realtime stream to the upstream, driving typed events through the pure -//! `OPENAI_REALTIME_CONFIG` transforms. -//! Network, auth header, key resolution, and wire (de)serialization live here so -//! the `transformation` module stays pure and typed. -//! -//! The dial and splice steps are factored out ([`dial_upstream`], [`splice`]) so -//! the connection pool ([`crate::io::realtime_pool`]) can pre-establish an upstream, -//! buffer its `session.created`, and later hand the live socket to the same -//! splice loop a fresh dial uses. - -use std::time::Duration; - -use futures_util::stream::{SplitSink, SplitStream}; -use futures_util::{Sink, SinkExt, Stream, StreamExt}; -use litellm_core::error::Error; -use litellm_core::realtime::transformation::RealtimeProviderConfig; -use litellm_core::realtime::types::RealtimeEvent; -use tokio::net::TcpStream; -use tokio_tungstenite::tungstenite::Message; -use tokio_tungstenite::tungstenite::client::IntoClientRequest; -use tokio_tungstenite::tungstenite::http::HeaderValue; -use tokio_tungstenite::tungstenite::http::header::AUTHORIZATION; -use tokio_tungstenite::{MaybeTlsStream, WebSocketStream}; - -use litellm_core::providers::openai::realtime::transformation::OPENAI_REALTIME_CONFIG; - -use crate::io::tls::connect_upstream; - -/// Environment variable holding the OpenAI API key (last-resort fallback). -const OPENAI_API_KEY_ENV: &str = "OPENAI_API_KEY"; - -const MISSING_KEY_MESSAGE: &str = "Missing OpenAI API Key - a realtime call is being made but no key was passed via params or the OPENAI_API_KEY environment variable"; - -/// Default **idle** timeout: if neither side sends a frame for this long, the -/// session is reaped. It resets on any activity, so it does not cap a healthy -/// (continuously streaming) session — it only frees a stalled one (e.g. a -/// half-open upstream that keeps the socket open but stops sending). -const DEFAULT_IDLE_TIMEOUT_SECS: u64 = 300; - -/// The concrete upstream WebSocket type (TLS or plain). Shared by the dial path -/// and the pool so warm sockets and fresh sockets are the exact same type. -pub type UpstreamWs = WebSocketStream>; -pub(crate) type UpstreamTx = SplitSink; -pub(crate) type UpstreamRx = SplitStream; - -/// Resolve the OpenAI API key from the explicit param or the environment. -/// -/// Blank/whitespace values are treated as absent (guard at resolution time). -pub(crate) fn resolve_api_key(api_key: Option<&str>) -> Result { - api_key - .map(str::trim) - .filter(|key| !key.is_empty()) - .map(str::to_string) - .or_else(|| { - std::env::var(OPENAI_API_KEY_ENV) - .ok() - .filter(|key| !key.trim().is_empty()) - }) - .ok_or_else(|| Error::Auth(MISSING_KEY_MESSAGE.to_string())) -} - -/// Open the upstream WebSocket to OpenAI for `(model, api_key, api_base)`. -/// -/// This is the dial half of [`realtime`], factored out so the pool can -/// pre-establish sockets ahead of any client. `api_key` here is already resolved -/// (non-blank) — the pool resolves it once when it is created. -pub(crate) async fn dial_upstream( - model: &str, - api_key: &str, - api_base: Option<&str>, -) -> Result { - let url = OPENAI_REALTIME_CONFIG.complete_url(api_base, model); - - let mut request = url - .as_str() - .into_client_request() - .map_err(|err| Error::Network(err.to_string()))?; - // GA realtime: only Authorization. The legacy OpenAI-Beta header triggers - // beta_api_shape_disabled, so we do not send it. - request.headers_mut().insert( - AUTHORIZATION, - HeaderValue::from_str(&format!("Bearer {api_key}")) - .map_err(|err| Error::Auth(err.to_string()))?, - ); - - let (upstream, _response) = connect_upstream(request) - .await - .map_err(|err| Error::Network(err.to_string()))?; - Ok(upstream) -} - -/// Read the next text frame from the upstream and decode it as a typed event. -/// -/// Used by the pool to pre-read OpenAI's unprompted `session.created`. Returns an -/// error on a non-text frame, a closed socket, or undecodable JSON so the pool can -/// discard a misbehaving socket rather than warm it. -pub(crate) async fn read_event(upstream_rx: &mut UpstreamRx) -> Result { - loop { - let message = upstream_rx - .next() - .await - .ok_or_else(|| Error::Network("upstream closed before first event".to_string()))? - .map_err(|err| Error::Network(err.to_string()))?; - match message { - Message::Text(text) => { - return serde_json::from_str(&text) - .map_err(|err| Error::InvalidResponse(err.to_string())); - } - // Ignore protocol frames (ping/pong) while waiting for the first event. - Message::Ping(_) | Message::Pong(_) => continue, - Message::Close(_) => { - return Err(Error::Network( - "upstream closed before first event".to_string(), - )); - } - _ => continue, - } - } -} - -/// Splice an already-connected upstream to the client streams. -/// -/// `prelude` is relayed to the client first (the pool passes the buffered -/// `session.created` here; the fresh-dial path passes `None` and lets the upstream -/// deliver it). Then a single select loop forwards both directions through the -/// transforms until either side closes or the idle timeout fires. -/// `observe` is invoked on **upstream→client** events only (the trusted side that -/// carries `session.created` and `response.done` usage) — never on client events, -/// so a client cannot fabricate usage into its own logs. -#[allow(clippy::too_many_arguments)] -pub(crate) async fn splice( - model: &str, - mut upstream_tx: UpstreamTx, - mut upstream_rx: UpstreamRx, - prelude: Option, - idle_timeout: Option, - mut observe: impl FnMut(&RealtimeEvent) + Send, - mut client_in: In, - mut client_out: Out, -) -> Result<(), Error> -where - In: Stream + Unpin + Send, - Out: Sink + Unpin + Send, - >::Error: std::fmt::Display, -{ - let config = &OPENAI_REALTIME_CONFIG; - - // Relay a buffered backend event (warm handoff's session.created) first, so a - // warm session looks identical to a fresh one from the client's view. - if let Some(event) = prelude { - for outbound in config.transform_realtime_response(&event, model)?.events { - client_out - .send(outbound) - .await - .map_err(|err| Error::Network(err.to_string()))?; - } - } - - let idle = idle_timeout.unwrap_or(Duration::from_secs(DEFAULT_IDLE_TIMEOUT_SECS)); - - // One loop forwarding both directions. The `sleep(idle)` arm is rebuilt every - // iteration, so any frame (either way) resets it — it fires only when the - // session has been fully idle for `idle`, reaping a stalled connection - // (task + upstream TCP socket) instead of leaking it. - loop { - tokio::select! { - // client -> upstream - client_event = client_in.next() => { - let Some(event) = client_event else { break }; // client disconnected - // NOTE: do NOT observe client events. session.created / response.done - // (carrying usage) are server→client events; observing the client arm - // would let an authenticated client POST a fabricated response.done and - // inflate its own spend log. Logging observes upstream events only. - for outbound in config.transform_realtime_request(&event, model)?.events { - let payload = serde_json::to_string(&outbound) - .map_err(|err| Error::InvalidResponse(err.to_string()))?; - upstream_tx - .send(Message::Text(payload)) - .await - .map_err(|err| Error::Network(err.to_string()))?; - } - } - // upstream -> client - upstream_message = upstream_rx.next() => { - let Some(message) = upstream_message else { break }; // upstream closed - match message.map_err(|err| Error::Network(err.to_string()))? { - Message::Text(text) => { - let event: RealtimeEvent = serde_json::from_str(&text) - .map_err(|err| Error::InvalidResponse(err.to_string()))?; - observe(&event); - for outbound in config.transform_realtime_response(&event, model)?.events { - client_out - .send(outbound) - .await - .map_err(|err| Error::Network(err.to_string()))?; - } - } - Message::Close(_) => break, - _ => {} - } - } - // idle timeout: no activity from either side within `idle` - _ = tokio::time::sleep(idle) => break, - } - } - Ok(()) -} - -/// Splice a client realtime stream to OpenAI: forward client events upstream -/// (via `transform_realtime_request`) and backend events downstream (via -/// `transform_realtime_response`). Returns when either side closes. -/// -/// Generic over the client transport (typed events) so this crate stays -/// framework-agnostic; the gateway adapts its axum socket to these. This is the -/// fresh-dial path: dial, then splice. The pool's warm-handoff path skips the dial -/// and calls [`splice`] directly with a buffered `session.created`. -#[allow(clippy::too_many_arguments)] -pub async fn realtime( - model: &str, - api_key: Option<&str>, - api_base: Option<&str>, - idle_timeout: Option, - observe: impl FnMut(&RealtimeEvent) + Send, - client_in: In, - client_out: Out, -) -> Result<(), Error> -where - In: Stream + Unpin + Send, - Out: Sink + Unpin + Send, - >::Error: std::fmt::Display, -{ - let api_key = resolve_api_key(api_key)?; - let upstream = dial_upstream(model, &api_key, api_base).await?; - let (upstream_tx, upstream_rx) = upstream.split(); - splice( - model, - upstream_tx, - upstream_rx, - None, - idle_timeout, - observe, - client_in, - client_out, - ) - .await -} - -/// Splice a pre-warmed upstream (taken from [`crate::io::realtime_pool`]) to the -/// client. Relays the buffered `session.created` first, then splices exactly like -/// the fresh-dial path — so a warm session is indistinguishable from a fresh one. -#[allow(clippy::too_many_arguments)] -pub async fn realtime_warm( - model: &str, - handoff: crate::io::realtime_pool::WarmHandoff, - idle_timeout: Option, - observe: impl FnMut(&RealtimeEvent) + Send, - client_in: In, - client_out: Out, -) -> Result<(), Error> -where - In: Stream + Unpin + Send, - Out: Sink + Unpin + Send, - >::Error: std::fmt::Display, -{ - splice( - model, - handoff.tx, - handoff.rx, - Some(handoff.session_created), - idle_timeout, - observe, - client_in, - client_out, - ) - .await -} - -#[cfg(test)] -mod tests { - use super::*; - - fn event(raw: &str) -> RealtimeEvent { - serde_json::from_str(raw).expect("valid event json") - } - - /// The realtime dial has to reach a `wss://` upstream without a process-wide - /// crypto provider installed, which is what dialing through `io::tls` buys. - #[tokio::test] - async fn dial_upstream_over_wss_reports_an_error_instead_of_panicking() { - let listener = tokio::net::TcpListener::bind("127.0.0.1:0") - .await - .expect("bind a loopback port"); - let port = listener - .local_addr() - .expect("read the bound address") - .port(); - tokio::spawn(async move { - while let Ok((stream, _peer)) = listener.accept().await { - drop(stream); - } - }); - - let result = dial_upstream( - "gpt-realtime", - "sk-test", - Some(&format!("wss://127.0.0.1:{port}")), - ) - .await; - - assert!(matches!(result, Err(Error::Network(_)))); - } - - #[test] - fn resolve_api_key_prefers_param_then_blank_falls_through() { - assert_eq!(resolve_api_key(Some("sk-test")).unwrap(), "sk-test"); - // A blank param with no env set should error. - if std::env::var(OPENAI_API_KEY_ENV).is_err() { - assert!(resolve_api_key(Some(" ")).is_err()); - } - } - - /// Live end-to-end check against OpenAI. Ignored by default (CI never runs - /// it); run explicitly with `OPENAI_API_KEY` set: - /// `cargo test -p litellm-ai-gateway --features server realtime_invokes_openai -- --ignored --nocapture` - #[tokio::test] - #[ignore = "hits the live OpenAI realtime API; needs OPENAI_API_KEY"] - async fn realtime_invokes_openai_and_responds() { - use futures_channel::mpsc; - - let key = - std::env::var(OPENAI_API_KEY_ENV).expect("set OPENAI_API_KEY to run this ignored test"); - - // client -> provider (we hold `client_tx` to push events upstream) - let (mut client_tx, client_in) = mpsc::unbounded::(); - // provider -> client (we hold `backend_rx` to read backend events) - let (client_out, mut backend_rx) = mpsc::unbounded::(); - - // Clone the key so the spawned task owns its `String` (no borrow across await). - let key_owned = key.clone(); - let call = tokio::spawn(async move { - realtime( - "gpt-realtime", - Some(&key_owned), - None, - None, - |_| {}, - client_in, - client_out, - ) - .await - }); - - // 1. First backend event should be session.created. - let first = tokio::time::timeout(Duration::from_secs(30), backend_rx.next()) - .await - .expect("timed out waiting for session.created") - .expect("backend stream closed before session.created"); - assert_eq!( - first.event_type, "session.created", - "expected session.created, got: {}", - first.event_type - ); - - // 2. Ask for a short audio response. - client_tx - .send(event( - r#"{"type":"conversation.item.create","item":{"type":"message","role":"user","content":[{"type":"input_text","text":"Say hi."}]}}"#, - )) - .await - .expect("send conversation.item.create"); - client_tx - .send(event(r#"{"type":"response.create"}"#)) - .await - .expect("send response.create"); - - // 3. Read backend events; require a non-empty audio delta, then response.done. - let mut saw_audio_delta = false; - let mut saw_done = false; - for _ in 0..500 { - let next = tokio::time::timeout(Duration::from_secs(30), backend_rx.next()).await; - let event = match next { - Ok(Some(event)) => event, - Ok(None) => break, - Err(_) => panic!("timed out waiting for backend events"), - }; - match event.event_type.as_str() { - "response.output_audio.delta" => { - let delta = event - .data - .get("delta") - .and_then(|value| value.as_str()) - .unwrap_or(""); - if !delta.is_empty() { - saw_audio_delta = true; - } - } - "response.done" => { - saw_done = true; - break; - } - _ => {} - } - } - - assert!( - saw_audio_delta, - "expected a response.output_audio.delta with non-empty delta" - ); - assert!(saw_done, "expected a response.done event"); - - // Drop the client sender so the provider's to_upstream side finishes. - drop(client_tx); - let _ = call.await; - } -} diff --git a/litellm-rust/crates/ai-gateway/src/io/realtime_pool.rs b/litellm-rust/crates/ai-gateway/src/io/realtime_pool.rs index 49e9c459a88..e910c92bffc 100644 --- a/litellm-rust/crates/ai-gateway/src/io/realtime_pool.rs +++ b/litellm-rust/crates/ai-gateway/src/io/realtime_pool.rs @@ -27,13 +27,7 @@ use std::collections::HashMap; use std::sync::{Arc, Mutex}; use std::time::{Duration, Instant}; -use futures_util::StreamExt; -use litellm_core::Error; -use litellm_core::realtime::types::RealtimeEvent; - -use crate::io::realtime::{ - UpstreamRx, UpstreamTx, UpstreamWs, dial_upstream, read_event, resolve_api_key, -}; +use litellm_core::realtime::{RealtimeConnectionSpec, warmup}; /// Default target warm sockets per key when pooling is enabled. pub const DEFAULT_POOL_SIZE: usize = 4; @@ -64,40 +58,15 @@ const BACKOFF_MAX: Duration = Duration::from_secs(30); /// Identifies an upstream connection: the tuple that fully determines the dial. /// `api_key` is included so a warm socket is only ever reused for the same key /// (no cross-tenant reuse). -#[derive(Clone, PartialEq, Eq, Hash)] -pub struct UpstreamKey { - pub model: String, - pub api_key: String, - pub api_base: Option, -} - -impl std::fmt::Debug for UpstreamKey { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - f.debug_struct("UpstreamKey") - .field("model", &self.model) - .field("api_key", &"[REDACTED]") - .field("api_base", &self.api_base) - .finish() - } -} +pub type UpstreamKey = RealtimeConnectionSpec; /// A warm upstream: split halves + the buffered `session.created` + when it was /// warmed (for `max_idle` expiry). struct WarmConnection { - tx: UpstreamTx, - rx: UpstreamRx, - session_created: RealtimeEvent, + connection: litellm_core::realtime::WarmConnection, warmed_at: Instant, } -/// A live upstream taken from the pool, ready to splice. The caller relays -/// `session_created` to the client first, then splices `(tx, rx)` as usual. -pub struct WarmHandoff { - pub tx: UpstreamTx, - pub rx: UpstreamRx, - pub session_created: RealtimeEvent, -} - /// Pool configuration, resolved once at startup from the environment. #[derive(Clone, Copy, Debug)] pub struct PoolConfig { @@ -256,7 +225,7 @@ impl RealtimePool { /// is too old or already dead is dropped (closing it) and the next candidate /// tried. Never blocks: if nothing warm is live, returns `None` so the caller /// fresh-dials. - pub fn take(&self, key: &UpstreamKey) -> Option { + pub fn take(&self, key: &UpstreamKey) -> Option { if !self.config.enabled() { return None; } @@ -273,14 +242,10 @@ impl RealtimePool { // Liveness: a non-blocking check that the socket hasn't already // delivered a Close/Err. A warm socket should be silent after // session.created, so anything pending means it is unhealthy. - if is_dead(&mut candidate.rx) { + if !candidate.connection.is_live() { continue; } - return Some(WarmHandoff { - tx: candidate.tx, - rx: candidate.rx, - session_created: candidate.session_created, - }); + return Some(candidate.connection); } } @@ -377,7 +342,7 @@ impl RealtimePool { let mut warm = self.warm.lock().unwrap(); if let Some(bucket) = warm.get_mut(key) { bucket.retain_mut(|conn| { - conn.warmed_at.elapsed() <= self.config.max_idle && !is_dead(&mut conn.rx) + conn.warmed_at.elapsed() <= self.config.max_idle && conn.connection.is_live() }); } } @@ -438,15 +403,10 @@ impl RealtimePool { /// /// `key.api_key` is already resolved (non-blank). The first frame OpenAI sends /// unprompted is `session.created`; we buffer exactly that and read nothing more. -async fn warm_one(key: &UpstreamKey) -> Result { - let upstream: UpstreamWs = - dial_upstream(&key.model, &key.api_key, key.api_base.as_deref()).await?; - let (tx, mut rx) = upstream.split(); - let session_created = read_event(&mut rx).await?; +async fn warm_one(key: &UpstreamKey) -> Result { + let connection = warmup(key).await?; Ok(WarmConnection { - tx, - rx, - session_created, + connection, warmed_at: Instant::now(), }) } @@ -459,34 +419,7 @@ pub fn upstream_key( api_key: Option<&str>, api_base: Option<&str>, ) -> Option { - let api_key = resolve_api_key(api_key).ok()?; - Some(UpstreamKey { - model: model.to_string(), - api_key, - api_base: api_base.map(str::to_string), - }) -} - -/// Non-blocking liveness check: poll the upstream once. A warm socket is silent -/// after `session.created`, so a pending `Close`/`Err`/`None` means it is dead. -/// A pending data frame (shouldn't happen pre-handoff) is also treated as -/// unhealthy — we'd rather discard and fresh-dial than hand over a socket in an -/// unexpected state. `Pending` (the healthy case) returns `false`. -fn is_dead(rx: &mut UpstreamRx) -> bool { - use futures_util::Stream; - use futures_util::task::noop_waker_ref; - use std::pin::Pin; - use std::task::{Context, Poll}; - - let mut cx = Context::from_waker(noop_waker_ref()); - match Pin::new(rx).poll_next(&mut cx) { - Poll::Pending => false, - Poll::Ready(None) => true, - Poll::Ready(Some(Err(_))) => true, - // Any frame arriving before handoff is unexpected for a silent warm - // socket; treat it as unhealthy. - Poll::Ready(Some(Ok(_))) => true, - } + RealtimeConnectionSpec::new(model, api_key, api_base).ok() } #[cfg(test)] diff --git a/litellm-rust/crates/ai-gateway/src/io/responses_ws.rs b/litellm-rust/crates/ai-gateway/src/io/responses_ws.rs deleted file mode 100644 index 9df3d0c6cc5..00000000000 --- a/litellm-rust/crates/ai-gateway/src/io/responses_ws.rs +++ /dev/null @@ -1,572 +0,0 @@ -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 litellm_core::Error; -use litellm_core::providers::openai::responses::transformation::OPENAI_RESPONSES_WS_CONFIG; -use litellm_core::responses::types::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}; - -use crate::io::tls::connect_upstream; - -use crate::constants::{ - DEFAULT_RESPONSES_WS_CONNECT_TIMEOUT_SECS, DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS, -}; - -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_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 connect = connect_upstream(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()), - })?; - 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_upstream(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; - - /// The Responses dial has to reach a `wss://` upstream without a process-wide - /// crypto provider installed, which is what dialing through `io::tls` buys. - #[tokio::test] - async fn dial_upstream_over_wss_reports_an_error_instead_of_panicking() { - let listener = TcpListener::bind("127.0.0.1:0") - .await - .expect("bind a loopback port"); - let port = listener - .local_addr() - .expect("read the bound address") - .port(); - tokio::spawn(async move { - while let Ok((stream, _peer)) = listener.accept().await { - drop(stream); - } - }); - - let result = - dial_upstream("gpt-5", "sk-test", Some(&format!("wss://127.0.0.1:{port}"))).await; - - assert!(matches!(result, Err(Error::Network(_)))); - } - - 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 b3f3028f361..429b97786c0 100644 --- a/litellm-rust/crates/ai-gateway/src/lib.rs +++ b/litellm-rust/crates/ai-gateway/src/lib.rs @@ -3,10 +3,6 @@ //! Two layers, split by feature so the Python `cdylib` can depend on the I/O //! without pulling in the HTTP server: //! -//! - Call-type modules such as [`ocr`]: provider transforms, lifecycle hooks, -//! and provider I/O. Always available — no feature required. These predate the -//! rule that a route's entrypoint and handler live in `litellm-core` (see -//! `litellm_core::messages`) and move there as they are touched. //! - [`io`]: compatibility exports and realtime WebSocket splice helpers. //! - The server modules ([`auth`], [`routes`], [`state`]) and anything pulling //! `axum` are gated behind the `server` feature, which the `litellm-ai-gateway` @@ -25,5 +21,3 @@ pub mod state; pub mod trace_parity; mod constants; -#[cfg(feature = "server")] -mod realtime; diff --git a/litellm-rust/crates/ai-gateway/src/realtime/mod.rs b/litellm-rust/crates/ai-gateway/src/realtime/mod.rs deleted file mode 100644 index 82be596ba86..00000000000 --- a/litellm-rust/crates/ai-gateway/src/realtime/mod.rs +++ /dev/null @@ -1,4 +0,0 @@ -//! Realtime logging collector. Observes the realtime event stream and emits a -//! `StandardLoggingPayload` to the registered callbacks on session close. - -pub mod streaming; diff --git a/litellm-rust/crates/ai-gateway/src/realtime/streaming.rs b/litellm-rust/crates/ai-gateway/src/realtime/streaming.rs deleted file mode 100644 index 8ce3ea5a123..00000000000 --- a/litellm-rust/crates/ai-gateway/src/realtime/streaming.rs +++ /dev/null @@ -1,414 +0,0 @@ -//! `RealTimeStreaming` — the realtime logging collector. -//! -//! Mirrors Python `litellm.realtime_api.main.RealTimeStreaming`: it observes the -//! event stream in O(1) (never buffering frames), accumulating just the fields -//! the spend log needs (model, id, cumulative usage), then on session close -//! builds a `StandardLoggingPayload` and fans it out to every registered -//! `CustomLogger`. - -use std::sync::Arc; -use std::time::{SystemTime, UNIX_EPOCH}; - -use litellm_core::realtime::types::RealtimeEvent; -use serde_json::Value; - -use crate::constants::DEFAULT_PROVIDER; -use litellm_core::integrations::custom_logger::{ - CallbackTiming, CallbackValue, CustomLogger, CustomLoggerRunner, LoggingError, ModelCallDetails, -}; -use litellm_core::integrations::types::{ - RequestMetadata, StandardLoggingMetadata, StandardLoggingPayload, Usage, -}; - -/// Current wall-clock time as epoch seconds (float), matching the Python -/// `startTime`/`endTime` contract. -fn epoch_seconds() -> f64 { - SystemTime::now() - .duration_since(UNIX_EPOCH) - .map(|d| d.as_secs_f64()) - .unwrap_or(0.0) -} - -/// Status of a finished realtime session, mapped to the callback record status. -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum SessionStatus { - Success, - Failure, -} - -/// Accumulates realtime session state and emits a logging payload on close. -pub struct RealTimeStreaming { - callbacks: Vec>, - /// REQUEST-ID RULE: the SpendLogs `request_id` == the OpenAI realtime session - /// id (`sess_…`), captured from `session.created`. Both `id` and - /// `litellm_call_id` are set to that value so the Python writer logs the same - /// id regardless of which field it reads. The gateway-generated `rt-…` id - /// (the constructor seed) is only a fallback for sessions that fail before - /// `session.created` arrives. - litellm_call_id: String, - /// See the request-id rule above — mirrors `litellm_call_id`. - id: String, - model: String, - custom_llm_provider: String, - usage: Usage, - response_cost: f64, - start_time: f64, - end_time: f64, - metadata: RequestMetadata, - /// Count of logging callbacks that failed to enqueue (non-fatal). - dropped: u64, -} - -impl RealTimeStreaming { - /// Create a collector for one session. `litellm_call_id` is the gateway's - /// per-connection id; `model` is the requested model (a sane default until - /// `session.created` reports the upstream model). - pub fn new( - callbacks: Vec>, - litellm_call_id: String, - model: String, - metadata: RequestMetadata, - ) -> Self { - let now = epoch_seconds(); - Self { - callbacks, - id: litellm_call_id.clone(), - litellm_call_id, - model, - custom_llm_provider: DEFAULT_PROVIDER.to_string(), - usage: Usage::default(), - response_cost: 0.0, - start_time: now, - end_time: now, - metadata, - dropped: 0, - } - } - - /// Number of logging callbacks that failed to enqueue so far (test/observ.). - #[allow(dead_code)] - pub fn dropped(&self) -> u64 { - self.dropped - } - - /// Observe one realtime event. O(1): updates accumulated state only; never - /// buffers frames. Safe to call on every event in either direction. - pub fn observe(&mut self, event: &RealtimeEvent) { - match event.event_type.as_str() { - "session.created" | "session.updated" => self.on_session(event), - "response.done" => self.on_response_done(event), - _ => {} - } - } - - /// `session.created` / `session.updated` → capture upstream id + model. - /// Per the request-id rule, the OpenAI session id becomes BOTH `id` and - /// `litellm_call_id`, replacing the gateway-generated fallback. - fn on_session(&mut self, event: &RealtimeEvent) { - let session = event.data.get("session").and_then(Value::as_object); - if let Some(id) = session.and_then(|s| s.get("id")).and_then(Value::as_str) - && !id.is_empty() - { - self.id = id.to_string(); - self.litellm_call_id = id.to_string(); - } - if let Some(model) = session.and_then(|s| s.get("model")).and_then(Value::as_str) - && !model.is_empty() - { - self.model = model.to_string(); - } - } - - /// `response.done` → add this response's usage to the cumulative totals. - fn on_response_done(&mut self, event: &RealtimeEvent) { - let usage = event - .data - .get("response") - .and_then(Value::as_object) - .and_then(|r| r.get("usage")) - .and_then(Value::as_object); - let Some(usage) = usage else { return }; - - let input = usage.get("input_tokens").and_then(Value::as_u64); - let output = usage.get("output_tokens").and_then(Value::as_u64); - let total = usage.get("total_tokens").and_then(Value::as_u64); - - if let Some(input) = input { - self.usage.prompt_tokens += input; - } - if let Some(output) = output { - self.usage.completion_tokens += output; - } - // Prefer the upstream-reported total; otherwise derive it. - match total { - Some(total) => self.usage.total_tokens += total, - None => { - self.usage.total_tokens += input.unwrap_or(0) + output.unwrap_or(0); - } - } - } - - /// Set the per-session response cost ($). Cost computation is Python-side in - /// the proxy; the gateway forwards 0.0 by default and lets the proxy price. - /// Public API (exercised in tests) for the future path where the gateway - /// prices realtime sessions itself. - #[allow(dead_code)] - pub fn set_response_cost(&mut self, cost: f64) { - self.response_cost = cost; - } - - /// Build the `StandardLoggingPayload` from accumulated state. - pub fn build_payload(&self) -> StandardLoggingPayload { - StandardLoggingPayload { - id: self.id.clone(), - litellm_call_id: self.litellm_call_id.clone(), - call_type: "realtime".to_string(), - model: self.model.clone(), - custom_llm_provider: self.custom_llm_provider.clone(), - response_cost: self.response_cost, - prompt_tokens: self.usage.prompt_tokens, - completion_tokens: self.usage.completion_tokens, - total_tokens: self.usage.total_tokens, - start_time: self.start_time, - end_time: self.end_time, - stream: true, - metadata: StandardLoggingMetadata { - user_api_key_hash: self.metadata.user_api_key_hash.clone(), - user_api_key_user_id: self.metadata.user_api_key_user_id.clone(), - user_api_key_team_id: self.metadata.user_api_key_team_id.clone(), - ..Default::default() - }, - messages: None, - } - } - - /// Finish the session: stamp the end time and fan the payload out to every - /// callback. On a logger enqueue error we bump a non-fatal counter (the - /// realtime session has already ended; a dropped log must never propagate). - pub async fn log_messages(&mut self, status: SessionStatus) { - self.end_time = epoch_seconds(); - let payload = self.build_payload(); - let timing = CallbackTiming::new(payload.start_time, payload.end_time); - let runner = CustomLoggerRunner::new(self.callbacks.clone()); - - match status { - SessionStatus::Success => { - let response = CallbackValue::new("realtime", serde_json::Value::Null); - let report = runner - .async_log_success_event( - &ModelCallDetails::from_standard_logging_payload(payload), - &response, - timing, - ) - .await; - self.dropped += report.dropped as u64; - } - SessionStatus::Failure => { - let error = LoggingError { - message: "realtime session ended in failure".to_string(), - kind: "RealtimeSessionError".to_string(), - }; - let response = CallbackValue::new( - "error", - serde_json::json!({ - "message": error.message, - "kind": error.kind, - }), - ); - let report = runner - .async_log_failure_event( - &ModelCallDetails::from_standard_logging_payload(payload) - .with_failure_error(error), - Some(&response), - timing, - ) - .await; - self.dropped += report.dropped as u64; - } - } - } -} - -#[cfg(test)] -mod tests { - use super::*; - use litellm_core::integrations::custom_logger::LogError; - use litellm_core::integrations::custom_logger::LogFuture; - use std::sync::atomic::{AtomicU64, Ordering}; - - fn event(raw: &str) -> RealtimeEvent { - serde_json::from_str(raw).expect("valid event json") - } - - /// A test logger that records the last payload it saw. - #[derive(Default)] - struct CapturingLogger { - calls: AtomicU64, - last_model: std::sync::Mutex>, - last_total_tokens: AtomicU64, - } - - impl CustomLogger for CapturingLogger { - 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 { - let payload = model_call_details - .standard_logging_payload - .as_ref() - .expect("standard logging payload"); - self.calls.fetch_add(1, Ordering::SeqCst); - *self.last_model.lock().unwrap() = Some(payload.model.clone()); - self.last_total_tokens - .store(payload.total_tokens, Ordering::SeqCst); - Ok(()) - }) - } - } - - #[tokio::test] - async fn observe_accumulates_model_and_tokens_then_logs() { - let logger = Arc::new(CapturingLogger::default()); - let callbacks: Vec> = vec![logger.clone()]; - let mut streaming = RealTimeStreaming::new( - callbacks, - "call_abc".to_string(), - "gpt-realtime".to_string(), - RequestMetadata { - user_api_key_hash: Some("hash123".to_string()), - user_api_key_user_id: Some("user-1".to_string()), - user_api_key_team_id: Some("team-1".to_string()), - }, - ); - - streaming.observe(&event( - r#"{"type":"session.created","session":{"id":"sess_001","model":"gpt-realtime-2025"}}"#, - )); - streaming.observe(&event( - r#"{"type":"response.done","response":{"usage":{"input_tokens":10,"output_tokens":5,"total_tokens":15}}}"#, - )); - // A second response.done accumulates. - streaming.observe(&event( - r#"{"type":"response.done","response":{"usage":{"input_tokens":3,"output_tokens":2,"total_tokens":5}}}"#, - )); - - let payload = streaming.build_payload(); - assert_eq!(payload.model, "gpt-realtime-2025"); - // Request-id rule: session.created's id becomes BOTH id and - // litellm_call_id (replacing the "call_abc" gateway fallback), so the - // SpendLogs request_id is always the OpenAI session id. - assert_eq!(payload.id, "sess_001"); - assert_eq!(payload.litellm_call_id, "sess_001"); - assert_eq!(payload.prompt_tokens, 13); - assert_eq!(payload.completion_tokens, 7); - assert_eq!(payload.total_tokens, 20); - assert_eq!(payload.response_cost, 0.0); - assert_eq!(payload.call_type, "realtime"); - assert_eq!(payload.custom_llm_provider, "openai"); - assert_eq!( - payload.metadata.user_api_key_hash.as_deref(), - Some("hash123") - ); - - streaming.log_messages(SessionStatus::Success).await; - assert_eq!(logger.calls.load(Ordering::SeqCst), 1); - assert_eq!( - logger.last_model.lock().unwrap().as_deref(), - Some("gpt-realtime-2025") - ); - assert_eq!(logger.last_total_tokens.load(Ordering::SeqCst), 20); - assert_eq!(streaming.dropped(), 0); - } - - #[test] - fn blank_session_id_and_model_keep_the_gateway_fallbacks() { - let mut streaming = RealTimeStreaming::new( - Vec::new(), - "call_fallback".to_string(), - "gpt-realtime".to_string(), - RequestMetadata::default(), - ); - - streaming.observe(&event( - r#"{"type":"session.created","session":{"id":"","model":""}}"#, - )); - let payload = streaming.build_payload(); - assert_eq!(payload.id, "call_fallback"); - assert_eq!(payload.litellm_call_id, "call_fallback"); - assert_eq!(payload.model, "gpt-realtime"); - - streaming.observe(&event( - r#"{"type":"session.updated","session":{"id":"sess_002","model":""}}"#, - )); - let payload = streaming.build_payload(); - assert_eq!(payload.id, "sess_002"); - assert_eq!(payload.litellm_call_id, "sess_002"); - assert_eq!(payload.model, "gpt-realtime"); - } - - #[test] - fn payload_serializes_with_camelcase_times_and_realtime_call_type() { - let mut streaming = RealTimeStreaming::new( - Vec::new(), - "call_xyz".to_string(), - "gpt-realtime".to_string(), - RequestMetadata::default(), - ); - streaming.observe(&event( - r#"{"type":"response.done","response":{"usage":{"input_tokens":1,"output_tokens":1,"total_tokens":2}}}"#, - )); - streaming.set_response_cost(0.0042); - let payload = streaming.build_payload(); - let json = serde_json::to_string(&payload).expect("serialize payload"); - - assert!(json.contains("\"startTime\""), "missing startTime: {json}"); - assert!(json.contains("\"endTime\""), "missing endTime: {json}"); - assert!( - json.contains("\"call_type\":\"realtime\""), - "missing call_type realtime: {json}" - ); - assert!( - json.contains("\"response_cost\""), - "missing response_cost: {json}" - ); - assert_eq!(payload.response_cost, 0.0042); - } - - /// A logger whose enqueue always fails should bump the dropped counter, not - /// panic or propagate. - #[tokio::test] - async fn failing_logger_bumps_dropped_counter() { - struct FailingLogger; - impl CustomLogger for FailingLogger { - fn async_log_success_event<'a>( - &'a self, - _model_call_details: &'a ModelCallDetails, - _response_obj: &'a CallbackValue, - _timing: CallbackTiming, - ) -> LogFuture<'a> { - Box::pin(async { Err(LogError::channel_full()) }) - } - - 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 { Err(LogError::channel_closed()) }) - } - } - let callbacks: Vec> = vec![Arc::new(FailingLogger)]; - let mut streaming = RealTimeStreaming::new( - callbacks, - "call_1".to_string(), - "gpt-realtime".to_string(), - RequestMetadata::default(), - ); - streaming.log_messages(SessionStatus::Success).await; - assert_eq!(streaming.dropped(), 1); - } -} diff --git a/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs index 45d9cafffd1..ed2ea369f95 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs @@ -10,6 +10,7 @@ use axum::http::header::{CACHE_CONTROL, CONTENT_TYPE, HeaderMap, HeaderValue}; use axum::response::{IntoResponse, Response}; use axum::routing::post; use litellm_core::Error; +use litellm_core::lifecycle::StreamingCall; use serde_json::{Map, Value}; use crate::auth::RequireMasterKey; @@ -43,26 +44,44 @@ async fn handle( } } -fn stream_response(upstream: reqwest::Response) -> Result { - let content_type = upstream - .headers() - .get(CONTENT_TYPE) - .cloned() +fn stream_response(call: StreamingCall) -> Result { + let content_type = call + .metadata + .content_type + .as_deref() + .map(HeaderValue::from_str) + .transpose() + .map_err(|error| { + MessagesRouteError(Error::InvalidResponse(format!( + "invalid upstream content type: {error}" + ))) + })? .unwrap_or_else(|| HeaderValue::from_static("text/event-stream")); + let status = StatusCode::from_u16(call.metadata.status).map_err(|error| { + MessagesRouteError(Error::InvalidResponse(format!( + "invalid upstream response status: {error}" + ))) + })?; + let cache_control = call + .metadata + .cache_control + .as_deref() + .map(HeaderValue::from_str) + .transpose() + .map_err(|error| { + MessagesRouteError(Error::InvalidResponse(format!( + "invalid upstream cache control: {error}" + ))) + })?; + let _completion = call.completion.register(); let mut response = Response::builder() - .status( - StatusCode::from_u16(upstream.status().as_u16()).map_err(|error| { - MessagesRouteError(Error::InvalidResponse(format!( - "invalid upstream response status: {error}" - ))) - })?, - ) + .status(status) .header(CONTENT_TYPE, content_type); - if let Some(value) = upstream.headers().get(CACHE_CONTROL) { + if let Some(value) = cache_control { response = response.header(CACHE_CONTROL, value); } response - .body(Body::from_stream(upstream.bytes_stream())) + .body(Body::from_stream(call.stream)) .map_err(|error| { MessagesRouteError(Error::InvalidResponse(format!( "failed to build streaming response: {error}" @@ -142,6 +161,9 @@ mod tests { use axum::http::Request; use axum::http::StatusCode; use axum::http::header::{CACHE_CONTROL, CONTENT_TYPE}; + use litellm_core::integrations::custom_logger::{ + CallbackTiming, CallbackValue, CustomLogger, LogFuture, ModelCallDetails, + }; use litellm_core::router::{Deployment, LiteLLMParams, Router as ModelRouter}; use serde_json::json; use tokio::io::{AsyncReadExt, AsyncWriteExt}; @@ -152,6 +174,34 @@ mod tests { use crate::io::realtime_pool::RealtimePool; use crate::state::AppState; + struct CapturingLogger { + sender: tokio::sync::mpsc::UnboundedSender<(u64, u64, bool)>, + } + + impl CustomLogger for CapturingLogger { + fn async_log_success_event<'a>( + &'a self, + details: &'a ModelCallDetails, + _: &'a CallbackValue, + _: CallbackTiming, + ) -> LogFuture<'a> { + Box::pin(async move { + let payload = details + .standard_logging_payload + .as_ref() + .expect("stream terminal has standard payload"); + self.sender + .send(( + payload.prompt_tokens, + payload.completion_tokens, + payload.stream, + )) + .expect("test receiver remains open"); + Ok(()) + }) + } + } + fn state(model: &str, api_base: String, master_key: Option<&str>) -> AppState { state_with_provider(model, model, api_base, master_key) } @@ -359,10 +409,13 @@ mod tests { #[tokio::test] async fn route_streams_anthropic_events_without_buffering_or_reordering() { let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); - let events = "event: message_start\ndata: {\"type\":\"message_start\"}\n\nevent: content_block_delta\ndata: {\"type\":\"content_block_delta\"}\n\nevent: message_stop\ndata: {\"type\":\"message_stop\"}\n\n"; + let events = "event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"usage\":{\"input_tokens\":5,\"output_tokens\":0}}}\n\nevent: content_block_delta\ndata: {\"type\":\"content_block_delta\"}\n\nevent: message_delta\ndata: {\"type\":\"message_delta\",\"usage\":{\"output_tokens\":4}}\n\nevent: message_stop\ndata: {\"type\":\"message_stop\"}\n\n"; let (api_base, server) = streaming_upstream(listener, 200, "text/event-stream", events).await; - let app = app(state("claude-test", api_base, Some("master-key"))); + let (sender, mut receiver) = tokio::sync::mpsc::unbounded_channel(); + let mut gateway_state = state("claude-test", api_base, Some("master-key")); + gateway_state.loggers = Arc::new(vec![Arc::new(CapturingLogger { sender })]); + let app = app(gateway_state); let response = app .oneshot( Request::builder() @@ -406,6 +459,13 @@ mod tests { .await .expect("response body reads"); assert_eq!(response_body, events.as_bytes()); + assert_eq!( + tokio::time::timeout(std::time::Duration::from_secs(1), receiver.recv()) + .await + .expect("stream completion is registered") + .expect("logger receives terminal"), + (5, 4, true) + ); let upstream_request = server.await.expect("upstream task completes"); let (_, upstream_body) = upstream_request .split_once("\r\n\r\n") diff --git a/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs b/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs index 78d471f47f9..923d23b4a32 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs @@ -4,10 +4,10 @@ use litellm_core::Error; use litellm_core::constants::ANTHROPIC_MESSAGES_PROVIDER; use litellm_core::integrations::custom_logger::{CustomLogger, CustomLoggerRunner, LogFuture}; use litellm_core::lifecycle::{ - ActionResult, CallLifecycleContext, Clock, RequestPolicy, TerminalDispatcher, TerminalRecord, + ActionResult, CallLifecycleContext, Clock, RequestPolicy, StreamingCall, TerminalDispatcher, + TerminalRecord, }; use litellm_core::messages::lifecycle::{self, Options}; -use litellm_core::messages::messages_stream; use litellm_core::messages::types::MessagesRequest; use litellm_core::router::Router; use serde_json::{Map, Value}; @@ -63,7 +63,7 @@ impl TerminalDispatcher for GatewayTerminalDispatcher { pub(crate) enum MessagesResponse { Json(Value), - Stream(reqwest::Response), + Stream(StreamingCall), } #[tracing::instrument( @@ -113,11 +113,7 @@ pub async fn run( extra_headers, timeout: None, }; - if request.body.get("stream").and_then(Value::as_bool) == Some(true) { - return messages_stream(request).await.map(MessagesResponse::Stream); - } - - let services = GatewayTerminalDispatcher::new(loggers); + let services = Arc::new(GatewayTerminalDispatcher::new(loggers)); let context = CallLifecycleContext::new( "messages", provider_model, @@ -130,7 +126,13 @@ pub async fn run( .unwrap_or(0) ), ); - let response = lifecycle::messages(&services, request, Options::default(), context) + if request.body.get("stream").and_then(Value::as_bool) == Some(true) { + return lifecycle::messages_stream(services, request, Options::default(), context) + .await + .map(MessagesResponse::Stream); + } + + let response = lifecycle::messages(&*services, request, Options::default(), context) .await .into_result()?; serde_json::to_value(response) diff --git a/litellm-rust/crates/ai-gateway/src/routes/realtime/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/realtime/mod.rs index 96c5d429e96..36ecc991538 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/realtime/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/realtime/mod.rs @@ -23,7 +23,6 @@ use litellm_core::router::Router as ModelRouter; use serde::Deserialize; use crate::auth::RequireMasterKey; -use crate::realtime::streaming::{RealTimeStreaming, SessionStatus}; use crate::state::AppState; use litellm_core::integrations::custom_logger::CustomLogger; use litellm_core::integrations::types::RequestMetadata; @@ -83,15 +82,6 @@ async fn handle( Ok(ws.on_upgrade(move |socket| bridge(socket, router, pool, loggers, master_key, model))) } -/// Adapt the axum socket (text frames) to the typed-event `Stream`/`Sink` the -/// service wants, keeping axum types out of `service`. -/// -/// This is also the realtime-logging seam: every upstream→client event (the -/// direction carrying `session.created` and `response.done` with usage) is fed -/// to a [`RealTimeStreaming`] collector via the splice's `observe` callback. The -/// observe is O(1) and never buffers frames. When the splice returns (any of the -/// three break paths — client disconnect, upstream close, idle timeout), we flush -/// one logging payload to the registered callbacks. async fn bridge( socket: WebSocket, router: Arc, @@ -102,38 +92,17 @@ async fn bridge( ) { let (ws_sink, ws_stream) = socket.split(); - // Attribute the spend log to the key that authenticated this session (the - // master key — the gateway is master-key auth). A non-null user_api_key_hash - // is required for the Python spend logger to write a SpendLogs row. - // - // SECURITY: hash the key — never send the raw credential. This field fans out - // to spend logs and every callback integration; the SHA-256 (matching the - // proxy's hash_token) keeps the plaintext master key out of all of them while - // still matching the key's hash in LiteLLM_SpendLogs. let metadata = RequestMetadata { user_api_key_hash: master_key.as_deref().map(crate::auth::hash_token), ..RequestMetadata::default() }; - // Owned by THIS task only. The splice observes it via a synchronous `&mut` - // callback (below), so there is no Arc/Mutex/atomic on the per-frame hot - // path — just a monomorphized FnMut mutating stack-local fields. This is - // what lets observe scale: 10K concurrent sessions = 10K independent - // collectors, zero cross-task synchronization. - let mut collector = RealTimeStreaming::new( - loggers.as_ref().clone(), - new_call_id(), - model.clone(), - metadata, - ); - let client_in = ws_stream.filter_map(|message| async move { match message { Ok(Message::Text(text)) => serde_json::from_str::(&text).ok(), _ => None, } }); - // Plain forwarding sink — no observe here anymore. let client_out = ws_sink.with(|event: RealtimeEvent| async move { Ok::(Message::Text( serde_json::to_string(&event).unwrap_or_default(), @@ -142,25 +111,16 @@ async fn bridge( futures_util::pin_mut!(client_in, client_out); - // The observe closure borrows `&mut collector` for the duration of the - // splice; the borrow ends when `run` returns, freeing the collector for the - // single post-session `log_messages` flush. `run` picks a pooled (warm) or - // fresh upstream — observe fires on the upstream arm either way. - let result = service::run( + let _ = service::run( &router, &pool, &model, None, - |event: &RealtimeEvent| collector.observe(event), + loggers, + new_call_id(), + metadata, client_in, client_out, ) .await; - - let status = if result.is_ok() { - SessionStatus::Success - } else { - SessionStatus::Failure - }; - collector.log_messages(status).await; } diff --git a/litellm-rust/crates/ai-gateway/src/routes/realtime/service.rs b/litellm-rust/crates/ai-gateway/src/routes/realtime/service.rs index f7bbb37dff4..bafe342dbe5 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/realtime/service.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/realtime/service.rs @@ -8,10 +8,15 @@ //! for correctness, only latency. use std::time::Duration; +use std::sync::Arc; use crate::io::realtime_pool::{RealtimePool, upstream_key}; use futures_util::{Sink, Stream}; use litellm_core::error::Error; +use litellm_core::integrations::custom_logger::{CustomLogger, CustomLoggerRunner}; +use litellm_core::integrations::types::{RequestMetadata, StandardLoggingMetadata}; +use litellm_core::lifecycle::{CallLifecycleContext, ExecutedCall}; +use litellm_core::realtime::{RealtimeRequest, realtime}; use litellm_core::realtime::types::RealtimeEvent; use litellm_core::router::Router; @@ -25,10 +30,12 @@ pub async fn run( pool: &RealtimePool, model: &str, idle_timeout: Option, - observe: impl FnMut(&RealtimeEvent) + Send, + loggers: Arc>>, + call_id: String, + metadata: RequestMetadata, client_in: In, client_out: Out, -) -> Result<(), Error> +) -> Result, Error> where In: Stream + Unpin + Send, Out: Sink + Unpin + Send, @@ -44,34 +51,29 @@ where .strip_prefix("openai/") .unwrap_or(¶ms.model); - // Warm path: take a pooled upstream (handshake already paid) and relay its - // buffered session.created immediately. On miss/dead socket fall through. - if let Some(key) = upstream_key( + let connection = upstream_key( provider_model, params.api_key.as_deref(), params.api_base.as_deref(), - ) && let Some(handoff) = pool.take(&key) - { - return crate::io::realtime::realtime_warm( - provider_model, - handoff, + ).ok_or_else(|| Error::Auth("missing realtime provider API key".to_string()))?; + let warm = pool.take(&connection); + let context = CallLifecycleContext::new("realtime", model, "openai", call_id) + .with_metadata(StandardLoggingMetadata { + user_api_key_hash: metadata.user_api_key_hash, + user_api_key_user_id: metadata.user_api_key_user_id, + user_api_key_team_id: metadata.user_api_key_team_id, + ..Default::default() + }); + Ok(realtime( + &CustomLoggerRunner::new(loggers.as_ref().clone()), + RealtimeRequest { + connection, + warm, idle_timeout, - observe, - client_in, - client_out, - ) - .await; - } - - // Cold path: fresh dial (the original behavior). - crate::io::realtime::realtime( - provider_model, - params.api_key.as_deref(), - params.api_base.as_deref(), - idle_timeout, - observe, + }, + context, client_in, client_out, ) - .await + .await) } diff --git a/litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs index dcaacbba668..1912a51ff17 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs @@ -229,7 +229,10 @@ async fn bridge( &mut client_out, ) .await; - if result.is_err() { + if !matches!( + result, + Ok(litellm_core::lifecycle::ExecutedCall::Success { .. }) + ) { client_out .close_with_code(1011, "Internal server error") .await; diff --git a/litellm-rust/crates/ai-gateway/src/routes/responses/service.rs b/litellm-rust/crates/ai-gateway/src/routes/responses/service.rs index 9c4067f6da6..9bf56d5cc02 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/responses/service.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/responses/service.rs @@ -3,17 +3,40 @@ use std::time::Duration; use futures_util::{Sink, Stream}; use litellm_core::Error; -use litellm_core::lifecycle::{CallLifecycle, CallLifecycleContext, SystemClock}; -use litellm_core::responses::instrumentation::{ - ResponsesWsCallbackPayload, ResponsesWsInstrumentation, ResponsesWsLogOutcome, - ResponsesWsMetadata, +use litellm_core::integrations::custom_logger::{CustomLogger, CustomLoggerRunner, LogFuture}; +use litellm_core::integrations::types::{RequestMetadata, StandardLoggingMetadata}; +use litellm_core::lifecycle::{ + CallLifecycleContext, Clock, ExecutedCall, TerminalDispatcher, TerminalRecord, }; use litellm_core::responses::types::ResponsesWsEvent; +use litellm_core::responses::websocket::{ResponsesWebSocketRequest, responses_websocket}; -use litellm_core::integrations::custom_logger::{ - CallbackTiming, CallbackValue, CustomLogger, CustomLoggerRunner, LoggingError, ModelCallDetails, -}; -use litellm_core::integrations::types::RequestMetadata; +struct GatewayResponsesServices { + runner: CustomLoggerRunner, +} + +impl GatewayResponsesServices { + fn new(loggers: Arc>>) -> Self { + Self { + runner: CustomLoggerRunner::new(loggers.as_ref().clone()), + } + } +} + +impl Clock for GatewayResponsesServices { + fn now(&self) -> f64 { + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .map(|duration| duration.as_secs_f64()) + .unwrap_or(0.0) + } +} + +impl TerminalDispatcher for GatewayResponsesServices { + fn dispatch<'a>(&'a self, terminal: &'a TerminalRecord) -> LogFuture<'a> { + self.runner.dispatch(terminal) + } +} #[allow(clippy::too_many_arguments)] pub async fn run( @@ -26,7 +49,7 @@ pub async fn run( metadata: RequestMetadata, client_in: In, client_out: Out, -) -> Result<(), Error> +) -> Result, Error> where In: Stream + Unpin + Send, Out: Sink + Unpin + Send, @@ -45,120 +68,25 @@ where "Responses WebSocket route supports OpenAI deployments only".to_string(), )); } - let instrumentation = Arc::new(ResponsesWsInstrumentation::new( - call_id.clone(), - model, - ResponsesWsMetadata { + let context = CallLifecycleContext::new("responses_websocket", model, "openai", call_id) + .with_metadata(StandardLoggingMetadata { user_api_key_hash: metadata.user_api_key_hash, user_api_key_user_id: metadata.user_api_key_user_id, user_api_key_team_id: metadata.user_api_key_team_id, + ..Default::default() + }); + responses_websocket( + &GatewayResponsesServices::new(loggers), + ResponsesWebSocketRequest { + model: provider_model.to_string(), + api_key: params.api_key.clone(), + api_base: params.api_base.clone(), + first_frame, + idle_timeout, }, - )); - let observer_instrumentation = Arc::clone(&instrumentation); - let context = CallLifecycleContext::new("responses_websocket", model, "openai", call_id); - let result = CallLifecycle - .run( - context, - (), - instrumentation.as_ref(), - instrumentation.as_ref(), - &SystemClock, - |_| async move { - crate::io::responses_ws::async_responses_websocket( - provider_model, - params.api_key.as_deref(), - params.api_base.as_deref(), - first_frame, - idle_timeout, - move |event| { - observer_instrumentation.observe(event); - }, - client_in, - client_out, - ) - .await - }, - ) - .await - .into_result(); - let outcome = instrumentation.take_or_build_outcome(result.is_ok()); - dispatch_outcome(loggers, outcome).await; - result -} - -async fn dispatch_outcome( - loggers: Arc>>, - outcome: ResponsesWsLogOutcome, -) { - let runner = CustomLoggerRunner::new(loggers.as_ref().clone()); - match outcome { - ResponsesWsLogOutcome::Success { payload, callback } => { - let (details, response, start_time, end_time) = logging_values(payload, callback, None); - let _ = runner - .async_log_success_event( - &details, - &response, - CallbackTiming::new(start_time, end_time), - ) - .await; - } - ResponsesWsLogOutcome::Failure { - payload, - callback, - error_message, - error_kind, - } => { - let error = LoggingError { - message: error_message, - kind: error_kind, - }; - let (details, response, start_time, end_time) = - logging_values(payload, callback, Some(error)); - let _ = runner - .async_log_failure_event( - &details, - Some(&response), - CallbackTiming::new(start_time, end_time), - ) - .await; - } - } -} - -fn logging_values( - payload: litellm_core::responses::instrumentation::ResponsesWsLogPayload, - callback: ResponsesWsCallbackPayload, - error: Option, -) -> (ModelCallDetails, CallbackValue, f64, f64) { - let start_time = payload.start_time; - let end_time = payload.end_time; - let callback = CallbackValue::new(callback.object, callback.value); - let details = ModelCallDetails::from_standard_logging_payload( - litellm_core::integrations::types::StandardLoggingPayload { - id: payload.id, - litellm_call_id: payload.litellm_call_id, - call_type: payload.call_type, - model: payload.model, - custom_llm_provider: payload.custom_llm_provider, - response_cost: payload.response_cost, - prompt_tokens: payload.usage.prompt_tokens, - completion_tokens: payload.usage.completion_tokens, - total_tokens: payload.usage.total_tokens, - start_time: payload.start_time, - end_time: payload.end_time, - stream: payload.stream, - metadata: litellm_core::integrations::types::StandardLoggingMetadata { - user_api_key_hash: payload.metadata.user_api_key_hash, - user_api_key_user_id: payload.metadata.user_api_key_user_id, - user_api_key_team_id: payload.metadata.user_api_key_team_id, - ..Default::default() - }, - messages: None, - }, - ); - let details = match error { - Some(error) => details.with_failure_error(error), - None => details, - }; - (details, callback, start_time, end_time) + context, + client_in, + client_out, + ) + .await } diff --git a/litellm-rust/crates/ai-gateway/tests/crypto_provider_wiring.rs b/litellm-rust/crates/ai-gateway/tests/crypto_provider_wiring.rs deleted file mode 100644 index 05f7d9610d5..00000000000 --- a/litellm-rust/crates/ai-gateway/tests/crypto_provider_wiring.rs +++ /dev/null @@ -1,48 +0,0 @@ -//! Guards the wiring, not just the helper: a `wss://` dial through the public -//! API has to resolve its own crypto provider, in a test binary where nothing -//! has installed a process-wide one, and has to leave it uninstalled. - -use std::collections::HashMap; -use std::time::Duration; - -use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection; -use tokio::net::TcpListener; - -async fn dead_tls_server() -> u16 { - let listener = TcpListener::bind("127.0.0.1:0") - .await - .expect("bind a loopback port"); - let port = listener - .local_addr() - .expect("read the bound address") - .port(); - - tokio::spawn(async move { - while let Ok((stream, _peer)) = listener.accept().await { - drop(stream); - } - }); - - port -} - -#[tokio::test] -async fn dialing_wss_returns_an_error_instead_of_panicking() { - let port = dead_tls_server().await; - - let result = ResponsesWebSocketConnection::connect_url( - &format!("wss://127.0.0.1:{port}/"), - &HashMap::new(), - Some(Duration::from_secs(10)), - ) - .await; - - assert!( - result.is_err(), - "a plain TCP server cannot finish a TLS handshake" - ); - assert!( - rustls::crypto::CryptoProvider::get_default().is_none(), - "the dial settles its provider on its own connector, not process-wide" - ); -} diff --git a/litellm-rust/crates/core/AGENTS.md b/litellm-rust/crates/core/AGENTS.md index 25691348e42..aa0a94521d4 100644 --- a/litellm-rust/crates/core/AGENTS.md +++ b/litellm-rust/crates/core/AGENTS.md @@ -42,10 +42,12 @@ core/src/messages/ client.rs # the shared reqwest client ``` -`ocr` is in flight: today it holds only `transformation` and `types`; the rest -of its lifecycle still lives in the gateway and moves here as it migrates. -`audio_transcription` and `realtime` are the same. Bringing a route to full -core shape means giving it a `mod.rs` entrypoint that owns the sequence above. +`ocr` prepares callback-visible headers and body in `prepare.rs`, settles those +authoritative roots into a native request after callbacks, and sends it through +`http_utils::buffered_post`. Reducto upload, Azure Document Intelligence polling, +and HTTP document URL conversion are declined at admission until they have an +implementation on this settled-request path. `audio_transcription` and +`realtime` remain in flight. The invariant is one function body owns the route lifecycle. The conceptual shape is: diff --git a/litellm-rust/crates/core/Cargo.toml b/litellm-rust/crates/core/Cargo.toml index 65f190db63b..f8ca86626d4 100644 --- a/litellm-rust/crates/core/Cargo.toml +++ b/litellm-rust/crates/core/Cargo.toml @@ -7,13 +7,18 @@ repository.workspace = true [dependencies] base64.workspace = true +bytes.workspace = true +futures-util.workspace = true rand.workspace = true reqwest.workspace = true +rustls.workspace = true +rustls-native-certs.workspace = true serde.workspace = true serde_json.workspace = true thiserror.workspace = true tracing.workspace = true tokio = { workspace = true, features = ["rt", "sync", "time"] } +tokio-tungstenite.workspace = true tracing-subscriber = { workspace = true, optional = true } sha2.workspace = true aws-config = { version = "1.9.0", default-features = false, features = ["rustls", "rt-tokio"], optional = true } @@ -36,6 +41,7 @@ bedrock-auth = [ observability = ["dep:tracing-subscriber"] [dev-dependencies] +futures-channel = "0.3" rstest.workspace = true tokio = { workspace = true, features = ["io-util", "macros", "net", "rt-multi-thread"] } tracing-subscriber.workspace = true diff --git a/litellm-rust/crates/core/src/audio_transcription/lifecycle.rs b/litellm-rust/crates/core/src/audio_transcription/lifecycle.rs new file mode 100644 index 00000000000..b073e3a6c4c --- /dev/null +++ b/litellm-rust/crates/core/src/audio_transcription/lifecycle.rs @@ -0,0 +1,299 @@ +use std::future::Future; +use std::pin::Pin; +use std::sync::Arc; +use std::sync::atomic::{AtomicU64, Ordering}; +use std::time::{SystemTime, UNIX_EPOCH}; + +use serde_json::{Map, Value, json}; + +use crate::Error; +use crate::integrations::custom_guardrail::{ + CustomGuardrail, CustomGuardrailRunner, GuardrailContext, GuardrailError, GuardrailRequest, +}; +use crate::integrations::custom_logger::{CallType, CustomLogger, CustomLoggerRunner, LogFuture}; +use crate::integrations::types::{RequestMetadata, StandardLoggingMetadata}; +use crate::lifecycle::{ + ActionResult, CallLifecycle, CallLifecycleContext, Clock, ExecutedCall, RequestPolicy, + TerminalDispatcher, TerminalRecord, +}; +use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider}; + +use super::handler::execute_audio_transcription_provider_call; +use super::prepare::prepare_audio_transcription_provider_call; +use super::types::{ + AudioRouteRequest, AudioTranscriptionRequest, ProviderAudioTranscriptionRequest, +}; + +pub trait AudioServices: TerminalDispatcher + Clock { + fn guardrails(&self) -> CustomGuardrailRunner; +} + +pub struct DefaultAudioServices { + dispatcher: CustomLoggerRunner, + guardrails: CustomGuardrailRunner, +} + +impl DefaultAudioServices { + pub fn new( + callbacks: Vec>, + guardrails: Vec>, + ) -> Self { + Self { + dispatcher: CustomLoggerRunner::new(callbacks), + guardrails: CustomGuardrailRunner::new(guardrails), + } + } +} + +impl Clock for DefaultAudioServices { + fn now(&self) -> f64 { + crate::lifecycle::SystemClock.now() + } +} + +impl TerminalDispatcher for DefaultAudioServices { + fn dispatch<'a>(&'a self, terminal: &'a TerminalRecord) -> LogFuture<'a> { + self.dispatcher.dispatch(terminal) + } +} + +impl AudioServices for DefaultAudioServices { + fn guardrails(&self) -> CustomGuardrailRunner { + self.guardrails.clone() + } +} + +pub struct AudioRoute; + +impl AudioRoute { + pub async fn execute( + services: &S, + request: AudioRouteRequest<'_>, + ) -> ExecutedCall { + let provider = get_custom_llm_provider(request.model, request.custom_llm_provider) + .unwrap_or(CustomLlmProvider { + model: request.model, + custom_llm_provider: "bedrock", + }); + let context = CallLifecycleContext::new( + "audio_transcription", + provider.model, + provider.custom_llm_provider, + request + .litellm_call_id + .map(str::to_string) + .unwrap_or_else(new_audio_transcription_call_id), + ) + .with_metadata(logging_metadata(&request.request_metadata)); + let policy = AudioRequestPolicy { + guardrail_runner: services.guardrails(), + request_metadata: request.request_metadata, + }; + let prepared = PreparedAudioTranscriptionRequest { + model: provider.model.to_string(), + custom_llm_provider: provider.custom_llm_provider.to_string(), + audio: request.audio, + 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: request.optional_params, + timeout: request.timeout, + }; + CallLifecycle + .run( + context, + prepared, + &policy, + services, + services, + execute_audio_transcription_provider_call, + ) + .await + } +} + +struct PreparedAudioTranscriptionRequest { + model: String, + custom_llm_provider: String, + audio: Value, + api_key: Option, + api_base: Option, + extra_headers: Option>, + optional_params: Map, + timeout: Option, +} + +struct AudioRequestPolicy { + guardrail_runner: CustomGuardrailRunner, + request_metadata: RequestMetadata, +} + +type AudioFuture<'a, T> = Pin> + Send + 'a>>; + +impl AudioRequestPolicy { + async fn run_pre_call_guardrails( + &self, + request: PreparedAudioTranscriptionRequest, + ) -> Result { + if self.guardrail_runner.is_empty() { + return Ok(request); + } + let (guardrail_request, _) = self + .guardrail_runner + .run_pre_call( + &guardrail_context(&self.request_metadata), + GuardrailRequest::new(json!({ + "model": request.model, + "custom_llm_provider": request.custom_llm_provider, + "audio": request.audio, + "optional_params": request.optional_params, + })), + ) + .await + .map_err(guardrail_error_to_core_error)?; + let Value::Object(mut data) = guardrail_request.data else { + return Err(Error::InvalidRequest( + "audio transcription pre_call guardrail must return an object".to_string(), + )); + }; + let audio = data.remove("audio").ok_or_else(|| { + Error::InvalidRequest("audio transcription guardrail removed audio".to_string()) + })?; + let optional_params = match data.remove("optional_params") { + Some(Value::Object(value)) => value, + Some(_) => { + return Err(Error::InvalidRequest( + "audio transcription optional_params must be an object".to_string(), + )); + } + None => Map::new(), + }; + Ok(PreparedAudioTranscriptionRequest { + audio, + optional_params, + ..request + }) + } + + async fn prepare_provider_request( + &self, + request: PreparedAudioTranscriptionRequest, + ) -> Result { + let provider_request = + prepare_audio_transcription_provider_call(AudioTranscriptionRequest { + model: &request.model, + audio: request.audio, + api_key: request.api_key.as_deref(), + api_base: request.api_base.as_deref(), + custom_llm_provider: Some(&request.custom_llm_provider), + extra_headers: request.extra_headers, + optional_params: request.optional_params, + timeout: request.timeout, + })?; + self.run_during_call_guardrails(provider_request).await + } + + async fn run_during_call_guardrails( + &self, + request: ProviderAudioTranscriptionRequest, + ) -> Result { + if self.guardrail_runner.is_empty() { + return Ok(request); + } + let (guardrail_request, _) = self + .guardrail_runner + .run_during_call( + &guardrail_context(&self.request_metadata), + GuardrailRequest::new(json!({ + "model": request.model, + "custom_llm_provider": request.custom_llm_provider, + "url": request.url, + "body": request.body, + })), + ) + .await + .map_err(guardrail_error_to_core_error)?; + let Value::Object(mut data) = guardrail_request.data else { + return Err(Error::InvalidRequest( + "audio transcription during_call guardrail must return an object".to_string(), + )); + }; + let body = data.remove("body").ok_or_else(|| { + Error::InvalidRequest("audio transcription guardrail removed body".to_string()) + })?; + Ok(ProviderAudioTranscriptionRequest { body, ..request }) + } +} + +impl RequestPolicy + for AudioRequestPolicy +{ + type PreCallFuture<'a> + = AudioFuture<'a, PreparedAudioTranscriptionRequest> + where + Self: 'a; + type DuringCallFuture<'a> + = AudioFuture<'a, ProviderAudioTranscriptionRequest> + where + Self: 'a; + + fn async_pre_call_hook<'a>( + &'a self, + _: &'a CallLifecycleContext, + request: PreparedAudioTranscriptionRequest, + ) -> Self::PreCallFuture<'a> { + Box::pin(async move { + match self.run_pre_call_guardrails(request).await { + Ok(request) => ActionResult::Replace(request), + Err(error) => ActionResult::Reject(error), + } + }) + } + + fn async_during_call_hook<'a>( + &'a self, + _: &'a CallLifecycleContext, + request: PreparedAudioTranscriptionRequest, + ) -> Self::DuringCallFuture<'a> { + Box::pin(async move { + match self.prepare_provider_request(request).await { + Ok(request) => ActionResult::Replace(request), + Err(error) => ActionResult::Reject(error), + } + }) + } +} + +fn logging_metadata(metadata: &RequestMetadata) -> StandardLoggingMetadata { + StandardLoggingMetadata { + user_api_key_hash: metadata.user_api_key_hash.clone(), + user_api_key_user_id: metadata.user_api_key_user_id.clone(), + user_api_key_team_id: metadata.user_api_key_team_id.clone(), + ..Default::default() + } +} + +fn guardrail_context(metadata: &RequestMetadata) -> GuardrailContext { + GuardrailContext { + call_type: CallType::Other("audio_transcription".to_string()), + selected_guardrails: Vec::new(), + metadata: std::collections::HashMap::new(), + user_api_key_hash: metadata.user_api_key_hash.clone(), + user_api_key_user_id: metadata.user_api_key_user_id.clone(), + user_api_key_team_id: metadata.user_api_key_team_id.clone(), + trace_parent: None, + } +} + +fn guardrail_error_to_core_error(error: GuardrailError) -> Error { + Error::InvalidRequest(format!("{}: {}", error.kind, error.message)) +} + +fn new_audio_transcription_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_or(0, |duration| duration.as_nanos()); + format!("audio-transcription-{timestamp}-{sequence}") +} diff --git a/litellm-rust/crates/core/src/audio_transcription/mod.rs b/litellm-rust/crates/core/src/audio_transcription/mod.rs index 31b6de4b3e4..811acd4a8e5 100644 --- a/litellm-rust/crates/core/src/audio_transcription/mod.rs +++ b/litellm-rust/crates/core/src/audio_transcription/mod.rs @@ -1,6 +1,7 @@ use crate::Error; mod client; mod handler; +mod lifecycle; mod prepare; pub mod transformation; pub mod types; @@ -8,13 +9,30 @@ pub mod types; use serde_json::Value; pub use handler::execute_audio_transcription_provider_call; +pub use lifecycle::{AudioRoute, AudioServices, DefaultAudioServices}; pub use prepare::prepare_audio_transcription_provider_call; -pub use types::{AudioTranscriptionRequest, ProviderAudioTranscriptionRequest}; +pub use types::{AudioRouteRequest, AudioTranscriptionRequest, ProviderAudioTranscriptionRequest}; #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> Result { - execute_audio_transcription_provider_call(prepare_audio_transcription_provider_call(request)?) - .await + let services = DefaultAudioServices::new(Vec::new(), Vec::new()); + AudioRoute::execute( + &services, + AudioRouteRequest { + model: request.model, + audio: request.audio, + api_key: request.api_key, + api_base: request.api_base, + custom_llm_provider: request.custom_llm_provider, + extra_headers: request.extra_headers, + optional_params: request.optional_params, + timeout: request.timeout, + request_metadata: Default::default(), + litellm_call_id: None, + }, + ) + .await + .into_result() } #[cfg(test)] diff --git a/litellm-rust/crates/core/src/audio_transcription/tests.rs b/litellm-rust/crates/core/src/audio_transcription/tests.rs index 263d63337b0..c7a354a89f1 100644 --- a/litellm-rust/crates/core/src/audio_transcription/tests.rs +++ b/litellm-rust/crates/core/src/audio_transcription/tests.rs @@ -1,11 +1,18 @@ use std::io::{Read, Write}; use std::net::TcpListener; +use std::sync::{Arc, Mutex}; use std::thread; use serde_json::{Map, json}; -use super::audio_transcription; use super::types::AudioTranscriptionRequest; +use super::{AudioRoute, AudioRouteRequest, DefaultAudioServices, audio_transcription}; +use crate::Error; +use crate::integrations::custom_guardrail::{ + CustomGuardrail, GuardrailContext, GuardrailDecision, GuardrailEventHook, GuardrailFuture, + GuardrailRequest, +}; +use crate::lifecycle::{ExecutedCall, RouteProjection}; #[tokio::test] async fn bedrock_request_is_signed_and_contains_audio() { @@ -48,3 +55,125 @@ async fn bedrock_request_is_signed_and_contains_audio() { assert_eq!(response, json!({"text": "hello"})); server.join().expect("server"); } + +struct ReplacingGuardrail { + calls: Mutex>, +} + +impl CustomGuardrail for ReplacingGuardrail { + fn guardrail_name(&self) -> &str { + "audio-test" + } + + fn supported_event_hooks(&self) -> &[GuardrailEventHook] { + &[GuardrailEventHook::PreCall, GuardrailEventHook::DuringCall] + } + + fn async_pre_call_hook<'a>( + &'a self, + _: &'a GuardrailContext, + mut request: GuardrailRequest, + ) -> GuardrailFuture<'a> { + Box::pin(async move { + self.calls.lock().unwrap().push("pre"); + request.data["audio"]["data"] = json!("AwQ="); + Ok(GuardrailDecision::Mask(request)) + }) + } + + fn async_moderation_hook<'a>( + &'a self, + _: &'a GuardrailContext, + mut request: GuardrailRequest, + ) -> GuardrailFuture<'a> { + Box::pin(async move { + self.calls.lock().unwrap().push("during"); + request.data["body"]["messages"][0]["content"][0]["text"] = + json!("Guarded transcription prompt"); + Ok(GuardrailDecision::Mask(request)) + }) + } +} + +#[tokio::test] +async fn route_owns_guardrail_provider_and_terminal_sequence() { + let listener = TcpListener::bind("127.0.0.1:0").expect("listener"); + let address = listener.local_addr().expect("address"); + let server = thread::spawn(move || { + let (mut stream, _) = listener.accept().expect("connection"); + let mut buffer = [0_u8; 16_384]; + let count = stream.read(&mut buffer).expect("request"); + let request = String::from_utf8_lossy(&buffer[..count]); + assert!(request.contains("\"bytes\":\"AwQ=\"")); + assert!(request.contains("Guarded transcription prompt")); + let response = b"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: 53\r\nConnection: close\r\n\r\n{\"output\":{\"message\":{\"content\":[{\"text\":\"hello\"}]}}}"; + stream.write_all(response).expect("response"); + }); + let guardrail = Arc::new(ReplacingGuardrail { + calls: Mutex::new(Vec::new()), + }); + let services = DefaultAudioServices::new(Vec::new(), vec![guardrail.clone()]); + let api_base = format!("http://{address}"); + let executed = AudioRoute::execute( + &services, + AudioRouteRequest { + model: "bedrock/mistral.voxtral-mini-3b-2507", + audio: json!({"data": "AQI=", "format": "wav", "filename": "audio.wav"}), + api_key: None, + api_base: Some(&api_base), + custom_llm_provider: None, + extra_headers: None, + optional_params: Map::from_iter([ + ("aws_access_key_id".to_string(), json!("access-key")), + ("aws_secret_access_key".to_string(), json!("secret-key")), + ("aws_region_name".to_string(), json!("us-east-1")), + ]), + timeout: None, + request_metadata: Default::default(), + litellm_call_id: Some("audio-call-1"), + }, + ) + .await; + + assert_eq!(*guardrail.calls.lock().unwrap(), vec!["pre", "during"]); + assert!(matches!( + executed, + ExecutedCall::Success { + response, + terminal, + } if response == json!({"text": "hello"}) + && terminal.call_id == "audio-call-1" + && matches!(terminal.projection, RouteProjection::Audio { ref value } if value == &response) + )); + server.join().expect("server"); +} + +#[tokio::test] +async fn route_returns_preparation_failure_with_audio_terminal() { + let services = DefaultAudioServices::new(Vec::new(), Vec::new()); + let executed = AudioRoute::execute( + &services, + AudioRouteRequest { + model: "unsupported/model", + audio: json!({"data": "AQI="}), + api_key: None, + api_base: None, + custom_llm_provider: Some("unsupported"), + extra_headers: None, + optional_params: Map::new(), + timeout: None, + request_metadata: Default::default(), + litellm_call_id: Some("audio-call-failure"), + }, + ) + .await; + + assert!(matches!( + executed, + ExecutedCall::Failure { + error: Error::InvalidProvider(_), + terminal, + } if terminal.call_id == "audio-call-failure" + && matches!(terminal.projection, RouteProjection::Audio { .. }) + )); +} diff --git a/litellm-rust/crates/core/src/audio_transcription/types.rs b/litellm-rust/crates/core/src/audio_transcription/types.rs index 559d7837027..5b8500275a3 100644 --- a/litellm-rust/crates/core/src/audio_transcription/types.rs +++ b/litellm-rust/crates/core/src/audio_transcription/types.rs @@ -3,6 +3,8 @@ use std::time::Duration; use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; +use crate::integrations::types::RequestMetadata; + use super::transformation::{AudioTranscriptionAuth, AudioTranscriptionProviderConfig}; pub struct AudioTranscriptionRequest<'a> { @@ -16,6 +18,19 @@ pub struct AudioTranscriptionRequest<'a> { pub timeout: Option, } +pub struct AudioRouteRequest<'a> { + pub model: &'a str, + pub audio: 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 request_metadata: RequestMetadata, + pub litellm_call_id: Option<&'a str>, +} + #[derive(Clone)] pub struct ProviderAudioTranscriptionRequest { pub(super) model: String, diff --git a/litellm-rust/crates/core/src/chat_completions/lifecycle.rs b/litellm-rust/crates/core/src/chat_completions/lifecycle.rs new file mode 100644 index 00000000000..c2fba6c2b02 --- /dev/null +++ b/litellm-rust/crates/core/src/chat_completions/lifecycle.rs @@ -0,0 +1,289 @@ +use crate::Error; +use crate::lifecycle::{ + ActionBinding, ActionKind, Delivery, ErrorDisposition, FailurePolicy, Lifecycle, + LifecycleRoute, Outcome, Owner, ResultPolicy, +}; + +use super::chat_completions_decline_reason; + +use serde_json::{Map, Value}; + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum Operation { + Setup, + DeploymentPre, + Prepare, + Send, + DeploymentSuccess, + DeploymentFailure, + SyncSuccess, + AsyncSuccess, + SyncSuccessIfNeeded, + SyncFailure, + AsyncFailure, + Restore, + Complete(Outcome), +} + +#[derive(Clone, Debug)] +pub struct Admission { + pub model: String, + pub messages: Value, + pub optional_params: Map, + pub custom_llm_provider: Option, +} + +#[derive(Clone, Debug, Default)] +pub struct Options { + pub asynchronous: bool, + pub internal_call: bool, +} + +#[derive(Clone, Copy, Debug, Default)] +pub struct Observations { + pub logger_available: bool, + pub has_fallbacks: bool, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct Transition { + pub error: ErrorDisposition, +} + +#[derive(Debug, PartialEq, Eq)] +pub struct Decline(&'static str); + +impl Decline { + pub fn reason(&self) -> &'static str { + self.0 + } +} + +#[derive(Debug)] +pub struct ChatCompletionsState { + operation: Operation, + outcome: Outcome, + asynchronous: bool, + internal_call: bool, +} + +#[derive(Debug)] +pub struct ChatCompletionsRoute; + +impl LifecycleRoute for ChatCompletionsRoute { + type Admission = Admission; + type Options = Options; + type Context = Observations; + type Operation = Operation; + type Observation = Observations; + type Outcome = Outcome; + type Transition = Transition; + type Error = Error; + type Decline = Decline; + type State = ChatCompletionsState; + + fn admit( + admission: &Admission, + options: Options, + ) -> Result, Error> { + if let Some(reason) = chat_completions_decline_reason( + &admission.model, + admission.custom_llm_provider.as_deref(), + admission.messages.clone(), + &admission.optional_params, + ) { + return Ok(Err(Decline(reason))); + } + Ok(Ok(ChatCompletionsState { + operation: Operation::Setup, + outcome: Outcome::Success, + asynchronous: options.asynchronous, + internal_call: options.internal_call, + })) + } + + fn operation(state: &Self::State) -> Operation { + state.operation + } + + fn advance( + state: &mut Self::State, + outcome: Outcome, + observations: Observations, + ) -> Result { + use Operation::*; + + if matches!(state.operation, Complete(_)) { + return Err(Error::InvalidRequest( + "chat completions lifecycle is already complete".into(), + )); + } + let failure = + if observations.logger_available && !(state.asynchronous && state.internal_call) { + SyncFailure + } else { + Restore + }; + let error = if outcome != Outcome::Success && state.operation != DeploymentFailure { + state.outcome = outcome; + ErrorDisposition::Replace + } else { + ErrorDisposition::Preserve + }; + state.operation = match (state.operation, outcome) { + (Restore, _) => Complete(state.outcome), + (DeploymentFailure, _) => failure, + (_, Outcome::Abort) => Restore, + (SyncFailure | AsyncFailure, Outcome::Failure) => Restore, + (Prepare | Send, Outcome::Failure) if state.asynchronous => DeploymentFailure, + (_, Outcome::Failure) => failure, + (Setup, Outcome::Success) if state.asynchronous => DeploymentPre, + (Setup | DeploymentPre, Outcome::Success) => Prepare, + (Prepare, Outcome::Success) => Send, + (Send, Outcome::Success) if state.asynchronous => DeploymentSuccess, + (Send, Outcome::Success) => SyncSuccess, + (DeploymentSuccess, Outcome::Success) => { + if state.internal_call || observations.has_fallbacks { + SyncSuccessIfNeeded + } else { + AsyncSuccess + } + } + (AsyncSuccess, Outcome::Success) => SyncSuccessIfNeeded, + (SyncFailure, Outcome::Success) if state.asynchronous => AsyncFailure, + (SyncSuccess | SyncSuccessIfNeeded | SyncFailure | AsyncFailure, Outcome::Success) => { + Restore + } + (Complete(_), _) => unreachable!(), + }; + Ok(Transition { error }) + } + + fn actions_for(operation: Operation, _: &Observations) -> &'static [ActionBinding] { + match operation { + Operation::Prepare | Operation::Send => &PROVIDER_ACTION, + Operation::SyncFailure | Operation::AsyncFailure | Operation::DeploymentFailure => { + &FAILURE_ACTION + } + Operation::Restore => &RESTORE_ACTION, + Operation::Complete(_) => &[], + _ => &CALLBACK_ACTION, + } + } +} + +const PROVIDER_ACTION: [ActionBinding; 1] = [ActionBinding { + kind: ActionKind::ProviderCall, + delivery: Delivery::InlineAwaited, + on_result: ResultPolicy::Replace, + on_error: FailurePolicy::Propagate, + owner: Owner::Core, +}]; +const CALLBACK_ACTION: [ActionBinding; 1] = [ActionBinding { + kind: ActionKind::TerminalSuccess, + delivery: Delivery::InlineAwaited, + on_result: ResultPolicy::Continue, + on_error: FailurePolicy::RecordAndContinue, + owner: Owner::Route, +}]; +const FAILURE_ACTION: [ActionBinding; 1] = [ActionBinding { + kind: ActionKind::TerminalFailure, + delivery: Delivery::InlineAwaited, + on_result: ResultPolicy::Continue, + on_error: FailurePolicy::PreserveOriginalFailure, + owner: Owner::Route, +}]; +const RESTORE_ACTION: [ActionBinding; 1] = [ActionBinding { + kind: ActionKind::Restore, + delivery: Delivery::InlineDirect, + on_result: ResultPolicy::Continue, + on_error: FailurePolicy::Propagate, + owner: Owner::Core, +}]; + +pub fn machine( + admission: &Admission, + options: Options, +) -> Result, Decline>, Error> { + Lifecycle::admit(admission, options) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn admission() -> Admission { + Admission { + model: "claude-sonnet-4-5".into(), + messages: serde_json::json!([{"role": "user", "content": "hi"}]), + optional_params: Map::from_iter([("max_tokens".into(), Value::from(16))]), + custom_llm_provider: Some("anthropic".into()), + } + } + + #[test] + fn admission_declines_before_the_lifecycle_starts() { + let mut unsupported = admission(); + unsupported.messages = serde_json::json!([]); + assert!(matches!( + machine(&unsupported, Options::default()), + Ok(Err(_)) + )); + assert!(matches!( + machine(&admission(), Options::default()), + Ok(Ok(_)) + )); + } + + #[test] + fn sync_and_async_success_sequences_are_selected_by_core() { + for (asynchronous, expected) in [ + ( + false, + vec![ + Operation::Setup, + Operation::Prepare, + Operation::Send, + Operation::SyncSuccess, + Operation::Restore, + ], + ), + ( + true, + vec![ + Operation::Setup, + Operation::DeploymentPre, + Operation::Prepare, + Operation::Send, + Operation::DeploymentSuccess, + Operation::AsyncSuccess, + Operation::SyncSuccessIfNeeded, + Operation::Restore, + ], + ), + ] { + let mut machine = machine( + &admission(), + Options { + asynchronous, + ..Options::default() + }, + ) + .unwrap() + .unwrap(); + for operation in expected { + assert_eq!(machine.operation(), operation); + machine + .advance( + Outcome::Success, + Observations { + logger_available: true, + has_fallbacks: false, + }, + ) + .unwrap(); + } + assert_eq!(machine.operation(), Operation::Complete(Outcome::Success)); + } + } +} diff --git a/litellm-rust/crates/core/src/chat_completions/mod.rs b/litellm-rust/crates/core/src/chat_completions/mod.rs index 32dea17d202..d48224d3731 100644 --- a/litellm-rust/crates/core/src/chat_completions/mod.rs +++ b/litellm-rust/crates/core/src/chat_completions/mod.rs @@ -11,6 +11,7 @@ mod client; mod common_utils; pub mod conversation; pub(crate) mod handler; +pub mod lifecycle; mod prepare; pub mod response_utils; pub mod transformation; @@ -22,6 +23,12 @@ use handler::execute_chat_completions_provider_call; use prepare::{parse_messages, resolve_provider_config, resolve_request}; use types::{ChatCompletionsRequest, ChatCompletionsResponse}; +use crate::integrations::custom_logger::CallbackTiming; +use crate::integrations::types::Usage; +use crate::lifecycle::{ + CallLifecycleContext, ExecutedCall, RouteProjection, TerminalClassification, TerminalRecord, +}; + #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] pub async fn chat_completions( request: ChatCompletionsRequest<'_>, @@ -29,6 +36,87 @@ pub async fn chat_completions( execute_chat_completions_provider_call(resolve_request(request)?).await } +pub async fn chat_completions_with_terminal( + request: ChatCompletionsRequest<'_>, + context: CallLifecycleContext, +) -> ExecutedCall { + let start_time = epoch_seconds(); + match chat_completions(request).await { + Ok(response) => { + let usage = Usage { + prompt_tokens: response.usage.prompt_tokens, + completion_tokens: response.usage.completion_tokens, + total_tokens: response.usage.total_tokens, + }; + let projection = serde_json::to_value(&response).unwrap_or(Value::Null); + let terminal = terminal( + context, + start_time, + usage, + TerminalClassification::Success, + projection, + ); + ExecutedCall::Success { response, terminal } + } + Err(error) => { + let kind = match &error { + Error::Auth(_) => "AuthError", + Error::InvalidProvider(_) => "InvalidProvider", + Error::InvalidRequest(_) => "InvalidRequest", + Error::InvalidType { .. } => "InvalidType", + Error::MissingField(_) => "MissingField", + Error::Http { .. } => "HttpError", + Error::InvalidResponse(_) => "InvalidResponse", + Error::Network(_) => "NetworkError", + Error::Connect(_) => "ConnectError", + Error::Routing(_) => "RoutingError", + Error::Unsupported(_) => "UnsupportedRequest", + }; + let message = error.to_string(); + let terminal = terminal( + context, + start_time, + Usage::default(), + TerminalClassification::Failure { + kind: kind.into(), + message: message.clone(), + }, + serde_json::json!({"kind": kind, "message": message}), + ); + ExecutedCall::Failure { error, terminal } + } + } +} + +fn terminal( + context: CallLifecycleContext, + start_time: f64, + usage: Usage, + classification: TerminalClassification, + value: Value, +) -> TerminalRecord { + TerminalRecord { + call_id: context.litellm_call_id, + trace_id: context.trace_id, + attempt: context.attempt, + call_type: context.call_type, + model: context.model, + provider: context.custom_llm_provider, + timing: CallbackTiming::new(start_time, epoch_seconds()), + usage, + cost_inputs: Default::default(), + classification, + projection: RouteProjection::ChatCompletions { value }, + } +} + +fn epoch_seconds() -> f64 { + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .map(|duration| duration.as_secs_f64()) + .unwrap_or(0.0) +} + /// Whether the core would accept this request, without resolving credentials or /// touching the network. /// diff --git a/litellm-rust/crates/core/src/integrations/custom_guardrail/mod.rs b/litellm-rust/crates/core/src/integrations/custom_guardrail/mod.rs index e5d4ce3a708..dddc9906aa0 100644 --- a/litellm-rust/crates/core/src/integrations/custom_guardrail/mod.rs +++ b/litellm-rust/crates/core/src/integrations/custom_guardrail/mod.rs @@ -41,6 +41,7 @@ pub trait CustomGuardrail: Send + Sync { } } +#[derive(Clone)] pub struct CustomGuardrailRunner { guardrails: Vec>, } diff --git a/litellm-rust/crates/core/src/lifecycle/execution.rs b/litellm-rust/crates/core/src/lifecycle/execution.rs index 233db6e56a9..fa88b134536 100644 --- a/litellm-rust/crates/core/src/lifecycle/execution.rs +++ b/litellm-rust/crates/core/src/lifecycle/execution.rs @@ -1,4 +1,5 @@ use std::future::Future; +use std::sync::Arc; use std::time::{SystemTime, UNIX_EPOCH}; use serde::Serialize; @@ -9,7 +10,10 @@ use crate::integrations::custom_logger::{CallbackTiming, LogFuture}; use crate::integrations::types::{StandardLoggingMetadata, Usage}; use super::terminal::CostInputs; -use super::{ActionResult, ExecutedCall, RouteProjection, TerminalClassification, TerminalRecord}; +use super::{ + ActionResult, ExecutedCall, RouteProjection, StreamingCall, StreamingObserver, StreamingSource, + TerminalClassification, TerminalRecord, +}; #[derive(Clone, Debug, PartialEq)] pub struct CallLifecycleContext { @@ -49,7 +53,7 @@ impl CallLifecycleContext { self } - fn terminal( + pub(super) fn terminal( &self, timing: CallbackTiming, classification: TerminalClassification, @@ -121,6 +125,44 @@ impl Clock for SystemClock { pub struct CallLifecycle; impl CallLifecycle { + pub async fn run_streaming( + &self, + context: CallLifecycleContext, + request: InitialReq, + services: Arc, + observer: Box, + provider_call: ProviderCall, + ) -> Result + where + Services: RequestPolicy + TerminalDispatcher + Clock + 'static, + ProviderCall: FnOnce(ProviderReq) -> ProviderFuture, + ProviderFuture: Future>, + { + let start_time = services.now(); + let request = match services.async_pre_call_hook(&context, request).await { + ActionResult::Continue(request) | ActionResult::Replace(request) => request, + ActionResult::Reject(error) => { + let executed = failure(&*services, &*services, &context, error, start_time).await; + return executed.into_result(); + } + }; + let provider_request = match services.async_during_call_hook(&context, request).await { + ActionResult::Continue(request) | ActionResult::Replace(request) => request, + ActionResult::Reject(error) => { + let executed = failure(&*services, &*services, &context, error, start_time).await; + return executed.into_result(); + } + }; + match provider_call(provider_request).await { + Ok(source) => Ok(StreamingCall::new( + source, observer, context, start_time, services, + )), + Err(error) => failure(&*services, &*services, &context, error, start_time) + .await + .into_result(), + } + } + pub async fn run< InitialReq, ProviderReq, diff --git a/litellm-rust/crates/core/src/lifecycle/mod.rs b/litellm-rust/crates/core/src/lifecycle/mod.rs index 581822c4e6c..a042fc1ac8a 100644 --- a/litellm-rust/crates/core/src/lifecycle/mod.rs +++ b/litellm-rust/crates/core/src/lifecycle/mod.rs @@ -3,6 +3,7 @@ pub mod executed; pub mod execution; pub mod machine; pub mod ocr; +mod streaming; pub mod terminal; pub mod types; @@ -13,5 +14,9 @@ pub use execution::{ TerminalDispatcher, }; pub use machine::{Lifecycle, LifecycleRoute}; -pub use terminal::{RouteProjection, TerminalClassification, TerminalRecord}; +pub use streaming::{ + BytesStream, StreamingCall, StreamingCompletion, StreamingMetadata, StreamingObserver, + StreamingSource, +}; +pub use terminal::{CostInputs, RouteProjection, TerminalClassification, TerminalRecord}; pub use types::{ActionKind, ActionResult, Delivery, ErrorDisposition, FailurePolicy, Outcome}; diff --git a/litellm-rust/crates/core/src/lifecycle/ocr.rs b/litellm-rust/crates/core/src/lifecycle/ocr.rs index b054ca7baec..ea412210e36 100644 --- a/litellm-rust/crates/core/src/lifecycle/ocr.rs +++ b/litellm-rust/crates/core/src/lifecycle/ocr.rs @@ -50,6 +50,7 @@ pub enum Operation { Setup, DeploymentPre, Prepare, + PreCall, Send, DeploymentSuccess, DeploymentFailure, @@ -177,11 +178,12 @@ impl LifecycleRoute for OcrRoute { (DeploymentFailure, _) => failure, (_, Outcome::Abort) => Restore, (SyncFailure | AsyncFailure, Outcome::Failure) => Restore, - (Prepare | Send, Outcome::Failure) if state.asynchronous => DeploymentFailure, + (Prepare | PreCall | Send, Outcome::Failure) if state.asynchronous => DeploymentFailure, (_, Outcome::Failure) => failure, (Setup, Outcome::Success) if state.asynchronous => DeploymentPre, (Setup | DeploymentPre, Outcome::Success) => Prepare, - (Prepare, Outcome::Success) => Send, + (Prepare, Outcome::Success) => PreCall, + (PreCall, Outcome::Success) => Send, (Send, Outcome::Success) if state.asynchronous => DeploymentSuccess, (Send, Outcome::Success) => SyncSuccess, (DeploymentSuccess, Outcome::Success) => { @@ -210,6 +212,7 @@ impl LifecycleRoute for OcrRoute { ) -> &'static [ActionBinding] { match operation { Operation::Prepare | Operation::Send => &PROVIDER_ACTION, + Operation::PreCall => &PRE_CALL_ACTION, Operation::SyncFailure | Operation::AsyncFailure | Operation::DeploymentFailure => { &FAILURE_ACTION } @@ -228,6 +231,14 @@ const PROVIDER_ACTION: [ActionBinding; 1] = [ActionBinding { owner: Owner::Core, }]; +const PRE_CALL_ACTION: [ActionBinding; 1] = [ActionBinding { + kind: ActionKind::RequestPolicy, + delivery: Delivery::InlineDirect, + on_result: ResultPolicy::Continue, + on_error: FailurePolicy::Propagate, + owner: Owner::Route, +}]; + const CALLBACK_ACTION: [ActionBinding; 1] = [ActionBinding { kind: ActionKind::TerminalSuccess, delivery: Delivery::InlineAwaited, @@ -327,13 +338,17 @@ mod tests { fn success_sequences_and_completion_are_core_selected() { use Operation::*; for (asynchronous, expected) in [ - (false, vec![Setup, Prepare, Send, SyncSuccess, Restore]), + ( + false, + vec![Setup, Prepare, PreCall, Send, SyncSuccess, Restore], + ), ( true, vec![ Setup, DeploymentPre, Prepare, + PreCall, Send, DeploymentSuccess, AsyncSuccess, @@ -364,13 +379,14 @@ mod tests { Setup, DeploymentPre, Prepare, + PreCall, Send, DeploymentSuccess, AsyncSuccess, SyncSuccessIfNeeded, ] } else { - vec![Setup, Prepare, Send, SyncSuccess] + vec![Setup, Prepare, PreCall, Send, SyncSuccess] }; for stage in stages { for outcome in [Outcome::Failure, Outcome::Abort] { @@ -380,7 +396,7 @@ mod tests { assert_eq!(transition.error, ErrorDisposition::Replace); let expected = if outcome == Outcome::Abort { Restore - } else if asynchronous && matches!(stage, Prepare | Send) { + } else if asynchronous && matches!(stage, Prepare | PreCall | Send) { DeploymentFailure } else { SyncFailure diff --git a/litellm-rust/crates/core/src/lifecycle/streaming.rs b/litellm-rust/crates/core/src/lifecycle/streaming.rs new file mode 100644 index 00000000000..7fb7e17716c --- /dev/null +++ b/litellm-rust/crates/core/src/lifecycle/streaming.rs @@ -0,0 +1,187 @@ +use std::pin::Pin; +use std::sync::Arc; +use std::task::{Context, Poll}; + +use bytes::Bytes; +use futures_util::Stream; +use serde_json::{Value, json}; +use tokio::sync::oneshot; +use tokio::task::JoinHandle; + +use crate::Error; +use crate::integrations::custom_logger::CallbackTiming; +use crate::integrations::types::Usage; + +use super::{ + CallLifecycleContext, Clock, TerminalClassification, TerminalDispatcher, TerminalRecord, +}; + +pub type BytesStream = Pin> + Send>>; + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct StreamingMetadata { + pub status: u16, + pub content_type: Option, + pub cache_control: Option, +} + +pub trait StreamingObserver: Send { + fn observe(&mut self, bytes: &[u8]); + fn usage(&self) -> Usage; + fn projection(&self) -> Value; +} + +pub struct StreamingSource { + pub metadata: StreamingMetadata, + pub stream: BytesStream, +} + +pub struct StreamingCall { + pub metadata: StreamingMetadata, + pub stream: BytesStream, + pub completion: StreamingCompletion, +} + +impl StreamingCall { + pub(crate) fn new( + source: StreamingSource, + observer: Box, + context: CallLifecycleContext, + start_time: f64, + services: Arc, + ) -> Self + where + S: Clock + TerminalDispatcher + 'static, + { + let (sender, receiver) = oneshot::channel(); + Self { + metadata: source.metadata, + stream: Box::pin(ObservedStream { + inner: source.stream, + observer, + sender: Some(sender), + }), + completion: StreamingCompletion { + receiver: Some(receiver), + context: Some(context), + start_time, + services, + }, + } + } +} + +pub struct StreamingCompletion { + receiver: Option>, + context: Option, + start_time: f64, + services: Arc, +} + +impl StreamingCompletion { + pub fn register(mut self) -> JoinHandle { + self.spawn() + } + + fn spawn(&mut self) -> JoinHandle { + let receiver = self + .receiver + .take() + .expect("stream completion registered once"); + let context = self + .context + .take() + .expect("stream completion registered once"); + let services = self.services.clone(); + let start_time = self.start_time; + tokio::spawn(async move { + let terminal_result = receiver.await.unwrap_or_else(|_| StreamTerminal { + usage: Usage::default(), + projection: json!({"stream": true}), + classification: TerminalClassification::Failure { + kind: "Cancelled".to_string(), + message: "stream completion was cancelled".to_string(), + }, + }); + let mut context = context; + context.usage = terminal_result.usage; + let terminal = context.terminal( + CallbackTiming::new(start_time, services.now()), + terminal_result.classification, + terminal_result.projection, + ); + let _ = services.dispatch(&terminal).await; + terminal + }) + } +} + +impl Drop for StreamingCompletion { + fn drop(&mut self) { + if self.receiver.is_some() { + let _ = self.spawn(); + } + } +} + +trait CompletionServices: Clock + TerminalDispatcher {} + +impl CompletionServices for T where T: Clock + TerminalDispatcher {} + +struct StreamTerminal { + usage: Usage, + projection: Value, + classification: TerminalClassification, +} + +struct ObservedStream { + inner: BytesStream, + observer: Box, + sender: Option>, +} + +impl ObservedStream { + fn complete(&mut self, classification: TerminalClassification) { + if let Some(sender) = self.sender.take() { + let _ = sender.send(StreamTerminal { + usage: self.observer.usage(), + projection: self.observer.projection(), + classification, + }); + } + } +} + +impl Stream for ObservedStream { + type Item = Result; + + fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + match self.inner.as_mut().poll_next(cx) { + Poll::Ready(Some(Ok(bytes))) => { + self.observer.observe(&bytes); + Poll::Ready(Some(Ok(bytes))) + } + Poll::Ready(Some(Err(error))) => { + self.complete(TerminalClassification::Failure { + kind: "NetworkError".to_string(), + message: error.to_string(), + }); + Poll::Ready(Some(Err(error))) + } + Poll::Ready(None) => { + self.complete(TerminalClassification::Success); + Poll::Ready(None) + } + Poll::Pending => Poll::Pending, + } + } +} + +impl Drop for ObservedStream { + fn drop(&mut self) { + self.complete(TerminalClassification::Failure { + kind: "Cancelled".to_string(), + message: "stream consumer dropped before completion".to_string(), + }); + } +} diff --git a/litellm-rust/crates/core/src/lifecycle/terminal.rs b/litellm-rust/crates/core/src/lifecycle/terminal.rs index fb6d9c689a1..4b9b2128b29 100644 --- a/litellm-rust/crates/core/src/lifecycle/terminal.rs +++ b/litellm-rust/crates/core/src/lifecycle/terminal.rs @@ -83,6 +83,10 @@ impl From<&TerminalRecord> for StandardLoggingPayload { stream: matches!( record.projection, RouteProjection::Realtime { .. } | RouteProjection::ResponsesWs { .. } + ) || matches!( + &record.projection, + RouteProjection::Messages { value } | RouteProjection::ChatCompletions { value } + if value.get("stream").and_then(Value::as_bool) == Some(true) ), metadata: record.cost_inputs.metadata.clone(), messages: record.projection.logging_input(), diff --git a/litellm-rust/crates/core/src/messages/handler.rs b/litellm-rust/crates/core/src/messages/handler.rs index fe73c49670c..4cb34acbe32 100644 --- a/litellm-rust/crates/core/src/messages/handler.rs +++ b/litellm-rust/crates/core/src/messages/handler.rs @@ -1,6 +1,7 @@ use crate::constants::ANTHROPIC_MESSAGES_PROVIDER; use crate::error::Error; use crate::http_utils::http_request; +use crate::lifecycle::{StreamingMetadata, StreamingSource}; use super::client::http_client; use super::common_utils::truncate_error_body; @@ -44,7 +45,7 @@ pub(super) async fn execute_messages_provider_call( pub(super) async fn execute_messages_provider_stream( request: MessagesRequest, -) -> Result { +) -> Result { let request = prepare_provider_request(request)?; if request.provider != ANTHROPIC_MESSAGES_PROVIDER { return Err(Error::InvalidRequest( @@ -74,5 +75,24 @@ pub(super) async fn execute_messages_provider_stream( body: truncate_error_body(&text), }); } - Ok(response) + let metadata = StreamingMetadata { + status: status.as_u16(), + content_type: response + .headers() + .get(reqwest::header::CONTENT_TYPE) + .and_then(|value| value.to_str().ok()) + .map(str::to_string), + cache_control: response + .headers() + .get(reqwest::header::CACHE_CONTROL) + .and_then(|value| value.to_str().ok()) + .map(str::to_string), + }; + let stream = response.bytes_stream(); + Ok(StreamingSource { + metadata, + stream: Box::pin(futures_util::StreamExt::map(stream, |result| { + result.map_err(|error| Error::Network(error.to_string())) + })), + }) } diff --git a/litellm-rust/crates/core/src/messages/lifecycle.rs b/litellm-rust/crates/core/src/messages/lifecycle.rs index 48feef64f1f..cd33e844c5d 100644 --- a/litellm-rust/crates/core/src/messages/lifecycle.rs +++ b/litellm-rust/crates/core/src/messages/lifecycle.rs @@ -1,14 +1,17 @@ use std::future::{Ready, ready}; +use std::sync::Arc; use crate::Error; use crate::integrations::custom_logger::{LogError, LogFuture}; +use crate::integrations::types::Usage; use crate::lifecycle::{ ActionBinding, ActionKind, ActionResult, CallLifecycle, CallLifecycleContext, Clock, Delivery, ErrorDisposition, ExecutedCall, FailurePolicy, Lifecycle, LifecycleRoute, Outcome, Owner, - RequestPolicy, ResultPolicy, TerminalDispatcher, TerminalRecord, + RequestPolicy, ResultPolicy, StreamingCall, StreamingObserver, TerminalDispatcher, + TerminalRecord, }; -use super::handler::execute_messages_provider_call; +use super::handler::{execute_messages_provider_call, execute_messages_provider_stream}; use super::types::{AnthropicMessagesResponse, MessagesRequest}; #[derive(Clone, Copy, Debug, PartialEq, Eq)] @@ -249,6 +252,87 @@ pub async fn messages( .await } +pub async fn messages_stream( + services: Arc, + request: MessagesRequest, + _options: Options, + context: CallLifecycleContext, +) -> Result { + CallLifecycle + .run_streaming( + context, + request, + services, + Box::::default(), + |request| async move { execute_messages_provider_stream(request).await }, + ) + .await +} + +#[derive(Default)] +struct AnthropicUsageObserver { + pending: Vec, + usage: Usage, +} + +impl AnthropicUsageObserver { + fn observe_event(&mut self, event: &[u8]) { + let Some(data) = event + .split(|byte| *byte == b'\n') + .find_map(|line| line.strip_prefix(b"data:")) + else { + return; + }; + let Ok(value) = serde_json::from_slice::(data.trim_ascii_start()) else { + return; + }; + if !matches!( + value.get("type").and_then(serde_json::Value::as_str), + Some("message_start" | "message_delta") + ) { + return; + } + let Some(usage) = value.get("usage").or_else(|| { + value + .get("message") + .and_then(|message| message.get("usage")) + }) else { + return; + }; + if let Some(input_tokens) = usage + .get("input_tokens") + .and_then(serde_json::Value::as_u64) + { + self.usage.prompt_tokens = input_tokens; + } + if let Some(output_tokens) = usage + .get("output_tokens") + .and_then(serde_json::Value::as_u64) + { + self.usage.completion_tokens = output_tokens; + } + self.usage.total_tokens = self.usage.prompt_tokens + self.usage.completion_tokens; + } +} + +impl StreamingObserver for AnthropicUsageObserver { + fn observe(&mut self, bytes: &[u8]) { + self.pending.extend_from_slice(bytes); + while let Some(end) = self.pending.windows(2).position(|window| window == b"\n\n") { + let event = self.pending.drain(..end + 2).collect::>(); + self.observe_event(&event); + } + } + + fn usage(&self) -> Usage { + self.usage + } + + fn projection(&self) -> serde_json::Value { + serde_json::json!({"stream": true}) + } +} + pub fn machine(options: Options) -> Result, Error> { Lifecycle::admit(&(), options).map(|result| match result { Ok(machine) => machine, diff --git a/litellm-rust/crates/core/src/messages/mod.rs b/litellm-rust/crates/core/src/messages/mod.rs index a703e7efe9c..16cf64503c1 100644 --- a/litellm-rust/crates/core/src/messages/mod.rs +++ b/litellm-rust/crates/core/src/messages/mod.rs @@ -4,8 +4,7 @@ //! [`messages`] is the top-level entrypoint: give it a model, a body, and //! credentials, and it resolves the provider, transforms the request, calls the //! provider, and returns a typed non-streaming response. [`messages_stream`] -//! is the streaming variant; it hands the raw upstream response back so a host -//! can splice the event stream to its own caller. +//! is the streaming variant. use crate::Error; use crate::constants::ANTHROPIC_MESSAGES_PROVIDER; @@ -17,7 +16,9 @@ mod prepare; pub mod transformation; pub mod types; -use handler::execute_messages_provider_stream; +use std::sync::Arc; + +use crate::lifecycle::StreamingCall; use types::{AnthropicMessagesResponse, MessagesRequest}; pub async fn messages(request: MessagesRequest) -> Result { @@ -42,8 +43,25 @@ pub async fn messages(request: MessagesRequest) -> Result Result { - execute_messages_provider_stream(request).await +pub async fn messages_stream(request: MessagesRequest) -> Result { + let provider = request + .custom_llm_provider + .as_deref() + .or_else(|| request.model.split_once('/').map(|(provider, _)| provider)) + .unwrap_or(ANTHROPIC_MESSAGES_PROVIDER); + let context = crate::lifecycle::CallLifecycleContext::new( + "messages", + &request.model, + provider, + format!("{:032x}", rand::random::()), + ); + lifecycle::messages_stream( + Arc::new(lifecycle::NoopServices), + request, + lifecycle::Options::default(), + context, + ) + .await } #[cfg(test)] diff --git a/litellm-rust/crates/core/src/ocr/client.rs b/litellm-rust/crates/core/src/ocr/client.rs deleted file mode 100644 index 73f5bb00800..00000000000 --- a/litellm-rust/crates/core/src/ocr/client.rs +++ /dev/null @@ -1,13 +0,0 @@ -use std::sync::OnceLock; -use std::time::Duration; - -pub(super) fn http_client() -> &'static reqwest::Client { - static CLIENT: OnceLock = OnceLock::new(); - CLIENT.get_or_init(|| { - reqwest::Client::builder() - .timeout(Duration::from_secs(600)) - .connect_timeout(Duration::from_secs(10)) - .build() - .unwrap_or_else(|_| reqwest::Client::new()) - }) -} diff --git a/litellm-rust/crates/core/src/ocr/common_utils.rs b/litellm-rust/crates/core/src/ocr/common_utils.rs deleted file mode 100644 index 5e7cbf79feb..00000000000 --- a/litellm-rust/crates/core/src/ocr/common_utils.rs +++ /dev/null @@ -1,515 +0,0 @@ -use std::net::IpAddr; -use std::time::{Duration, Instant}; - -use crate::error::Error; -use crate::ocr::transformation::OcrProviderConfig; -use base64::Engine; -use base64::engine::general_purpose::STANDARD as BASE64_STANDARD; -use reqwest::Url; -use serde_json::{Map, Value}; - -use crate::providers::azure_ai::ocr::transformation as azure_ai; -use crate::providers::mistral::ocr::transformation::MISTRAL_OCR_CONFIG; -use crate::providers::reducto::ocr::transformation as reducto; -use crate::providers::vertex_ai::ocr::transformation as vertex_ai; - -use super::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" => azure_ai::config_for_model(model).ok(), - "vertex_ai" => vertex_ai::config_for_model(model).ok(), - _ => None, - } -} - -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 {}", - crate::error::json_type_name(&value) - )) - }) - }) - .collect() -} - -fn document_url_field(document: &Value) -> Result, Error> { - let Some(object) = document.as_object() else { - return Ok(None); - }; - let Some(doc_type) = object.get("type").and_then(Value::as_str) else { - return Ok(None); - }; - let field = match doc_type { - "document_url" => "document_url", - "image_url" => "image_url", - _ => return Ok(None), - }; - let Some(url) = object.get(field).and_then(Value::as_str) else { - return Ok(None); - }; - Ok(Some((field, url))) -} - -fn is_url_requiring_fetch(url: &str) -> bool { - !url.starts_with("data:") && (url.starts_with("http://") || url.starts_with("https://")) -} - -fn max_document_download_bytes() -> u64 { - let max_size_mb = std::env::var("MAX_IMAGE_URL_DOWNLOAD_SIZE_MB") - .ok() - .and_then(|value| value.parse::().ok()) - .unwrap_or(DEFAULT_MAX_IMAGE_URL_DOWNLOAD_SIZE_MB); - (max_size_mb.max(0.0) * 1024.0 * 1024.0) as u64 -} - -fn is_blocked_ip(ip: IpAddr) -> bool { - match ip { - IpAddr::V4(ip) => { - ip.is_private() - || ip.is_loopback() - || ip.is_link_local() - || ip.is_broadcast() - || ip.is_multicast() - || ip.is_unspecified() - } - IpAddr::V6(ip) => { - let first_segment = ip.segments()[0]; - let is_unique_local = (first_segment & 0xfe00) == 0xfc00; - let is_link_local = (first_segment & 0xffc0) == 0xfe80; - ip.is_loopback() - || ip.is_unspecified() - || ip.is_multicast() - || is_unique_local - || is_link_local - || ip - .to_ipv4_mapped() - .or_else(|| ip.to_ipv4()) - .map(|v4| is_blocked_ip(IpAddr::V4(v4))) - .unwrap_or(false) - } - } -} - -fn blocked_url_error(url: &Url) -> Error { - Error::InvalidRequest(format!( - "OCR document URL rejected by SSRF protection: {url}" - )) -} - -async fn validate_safe_fetch_url(url: &Url) -> Result<(), Error> { - if !matches!(url.scheme(), "http" | "https") { - return Err(blocked_url_error(url)); - } - - let host = url.host_str().ok_or_else(|| blocked_url_error(url))?; - if let Ok(ip) = host.parse::() { - if is_blocked_ip(ip) { - return Err(blocked_url_error(url)); - } - return Ok(()); - } - - let port = url - .port_or_known_default() - .ok_or_else(|| blocked_url_error(url))?; - let addresses = tokio::net::lookup_host((host, port)) - .await - .map_err(|err| Error::Network(err.to_string()))?; - let mut saw_address = false; - for address in addresses { - saw_address = true; - if is_blocked_ip(address.ip()) { - return Err(blocked_url_error(url)); - } - } - if !saw_address { - return Err(blocked_url_error(url)); - } - Ok(()) -} - -fn redirect_location(response: &reqwest::Response, url: &Url) -> Result { - let location = response - .headers() - .get(reqwest::header::LOCATION) - .and_then(|value| value.to_str().ok()) - .ok_or_else(|| { - Error::InvalidResponse("OCR document redirect missing Location header".to_string()) - })?; - url.join(location) - .map_err(|err| Error::InvalidResponse(format!("invalid OCR document redirect: {err}"))) -} - -async fn safe_get_document_url(url: &str) -> Result<(Url, reqwest::Response), Error> { - let client = reqwest::Client::builder() - .redirect(reqwest::redirect::Policy::none()) - .build() - .map_err(|err| Error::Network(err.to_string()))?; - let mut current_url = Url::parse(url) - .map_err(|err| Error::InvalidRequest(format!("invalid OCR document URL: {err}")))?; - - for _ in 0..MAX_SAFE_FETCH_REDIRECTS { - validate_safe_fetch_url(¤t_url).await?; - let response = client - .get(current_url.clone()) - .send() - .await - .map_err(|err| Error::Network(err.to_string()))?; - if !response.status().is_redirection() { - return Ok((current_url, response)); - } - current_url = redirect_location(&response, ¤t_url)?; - } - - Err(Error::InvalidRequest( - "Too many redirects while fetching OCR document URL".to_string(), - )) -} - -fn enforce_download_size(content_length: u64, max_bytes: u64, url: &Url) -> Result<(), Error> { - if max_bytes == 0 { - return Err(Error::InvalidRequest(format!( - "OCR document URL download is disabled (MAX_IMAGE_URL_DOWNLOAD_SIZE_MB=0). url={url}" - ))); - } - if content_length > max_bytes { - let size_mb = content_length as f64 / (1024.0 * 1024.0); - let max_size_mb = max_bytes as f64 / (1024.0 * 1024.0); - return Err(Error::InvalidRequest(format!( - "OCR document size ({size_mb:.2}MB) exceeds maximum allowed size ({max_size_mb:.2}MB). url={url}" - ))); - } - Ok(()) -} - -async fn read_response_with_limit( - mut response: reqwest::Response, - url: &Url, -) -> Result, Error> { - let max_bytes = max_document_download_bytes(); - if let Some(content_length) = response.content_length() { - enforce_download_size(content_length, max_bytes, url)?; - } else { - enforce_download_size(0, max_bytes, url)?; - } - - let mut bytes = Vec::new(); - let mut bytes_downloaded: u64 = 0; - while let Some(chunk) = response - .chunk() - .await - .map_err(|err| Error::Network(err.to_string()))? - { - bytes_downloaded += chunk.len() as u64; - enforce_download_size(bytes_downloaded, max_bytes, url)?; - bytes.extend_from_slice(&chunk); - } - Ok(bytes) -} - -pub(super) async fn convert_document_url_to_data_uri(document: Value) -> Result { - let Some((field, url)) = document_url_field(&document)? else { - return Ok(document); - }; - if !is_url_requiring_fetch(url) { - return Ok(document); - } - - let (final_url, response) = safe_get_document_url(url).await?; - let status = response.status(); - if !status.is_success() { - let body = response.text().await.unwrap_or_default(); - return Err(Error::Http { - status: status.as_u16(), - body: truncate_error_body(&body), - }); - } - let content_type = response - .headers() - .get(reqwest::header::CONTENT_TYPE) - .and_then(|value| value.to_str().ok()) - .and_then(|value| value.split(';').next()) - .map(str::trim) - .filter(|value| !value.is_empty()) - .unwrap_or("application/octet-stream") - .to_string(); - let bytes = read_response_with_limit(response, &final_url).await?; - let data_uri = format!( - "data:{content_type};base64,{}", - BASE64_STANDARD.encode(bytes) - ); - - let mut transformed = document - .as_object() - .cloned() - .ok_or_else(|| Error::InvalidRequest("OCR document must be an object".to_string()))?; - transformed.insert(field.to_string(), Value::String(data_uri)); - Ok(Value::Object(transformed)) -} - -fn same_origin(left: &str, right: &str) -> bool { - let Ok(left) = reqwest::Url::parse(left) else { - return false; - }; - let Ok(right) = reqwest::Url::parse(right) else { - return false; - }; - left.scheme() == right.scheme() - && left.host_str() == right.host_str() - && left.port_or_known_default() == right.port_or_known_default() -} - -fn retry_after_secs(response: &reqwest::Response) -> u64 { - response - .headers() - .get(reqwest::header::RETRY_AFTER) - .and_then(|value| value.to_str().ok()) - .and_then(|value| value.parse::().ok()) - .unwrap_or(2) -} - -fn operation_status(response_json: &Value) -> Result<&str, Error> { - let status = response_json - .get("status") - .and_then(Value::as_str) - .ok_or(Error::MissingField("status"))?; - match status { - "succeeded" => Ok("succeeded"), - "running" | "notStarted" => Ok("running"), - "failed" => { - let message = response_json - .get("error") - .and_then(|error| error.get("message")) - .and_then(Value::as_str) - .unwrap_or("Unknown error"); - Err(Error::InvalidResponse(format!( - "Azure Document Intelligence analysis failed: {message}" - ))) - } - other => Err(Error::InvalidResponse(format!( - "Unknown operation status: {other}" - ))), - } -} - -#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] -pub(super) async fn poll_document_intelligence( - operation_url: &str, - original_url: &str, - headers: &[(String, String)], - timeout: Option, -) -> Result { - if !same_origin(operation_url, original_url) { - return Err(Error::InvalidResponse( - "Azure Document Intelligence: rejected cross-origin polling URL".to_string(), - )); - } - - let start = Instant::now(); - let timeout = timeout.unwrap_or(Duration::from_secs( - AZURE_DOCUMENT_INTELLIGENCE_POLL_TIMEOUT_SECS, - )); - loop { - if start.elapsed() > timeout { - return Err(Error::Network(format!( - "Azure Document Intelligence operation polling timed out after {} seconds", - timeout.as_secs() - ))); - } - - let mut request_builder = http_client().get(operation_url); - for (key, value) in headers { - if key.eq_ignore_ascii_case("ocp-apim-subscription-key") { - request_builder = request_builder.header(key, value); - } - } - let response = request_builder - .send() - .await - .map_err(|err| Error::Network(err.to_string()))?; - let retry_after = retry_after_secs(&response); - let status = response.status(); - let text = response - .text() - .await - .map_err(|err| Error::Network(err.to_string()))?; - if !status.is_success() { - return Err(Error::Http { - status: status.as_u16(), - body: truncate_error_body(&text), - }); - } - let response_json: Value = serde_json::from_str(&text).map_err(|err| { - Error::InvalidResponse(format!("invalid Azure DI poll response JSON: {err}")) - })?; - if operation_status(&response_json)? == "succeeded" { - return Ok(response_json); - } - tokio::time::sleep(Duration::from_secs(retry_after)).await; - } -} - -#[cfg(test)] -mod tests { - use crate::ocr::transformation::OcrResponseHandling; - use serde_json::json; - - use super::*; - - #[test] - fn blocks_private_and_metadata_ips() { - assert!(is_blocked_ip("127.0.0.1".parse().unwrap())); - assert!(is_blocked_ip("10.0.0.1".parse().unwrap())); - assert!(is_blocked_ip("169.254.169.254".parse().unwrap())); - assert!(is_blocked_ip("::1".parse().unwrap())); - assert!(is_blocked_ip("fd00::1".parse().unwrap())); - assert!(is_blocked_ip("fe80::1".parse().unwrap())); - assert!(is_blocked_ip("::ffff:169.254.169.254".parse().unwrap())); - assert!(is_blocked_ip("::ffff:10.0.0.1".parse().unwrap())); - assert!(!is_blocked_ip("8.8.8.8".parse().unwrap())); - assert!(!is_blocked_ip("::ffff:8.8.8.8".parse().unwrap())); - } - - #[tokio::test] - async fn convert_document_url_rejects_loopback_fetch() { - let error = convert_document_url_to_data_uri(json!({ - "type": "image_url", - "image_url": "http://127.0.0.1/image.png" - })) - .await - .unwrap_err(); - - assert!(matches!( - error, - Error::InvalidRequest(message) - if message.contains("SSRF protection") - )); - } - - #[tokio::test] - async fn convert_document_url_leaves_data_uri_untouched() { - let document = json!({ - "type": "image_url", - "image_url": "data:image/png;base64,abcd" - }); - - let transformed = convert_document_url_to_data_uri(document.clone()) - .await - .unwrap(); - - 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 deleted file mode 100644 index 00d1b7d09e0..00000000000 --- a/litellm-rust/crates/core/src/ocr/handler.rs +++ /dev/null @@ -1,84 +0,0 @@ -use crate::error::Error; -use crate::http_utils::http_request; -use crate::ocr::transformation::OcrResponseHandling; -use serde_json::Value; - -use super::client::http_client; -use super::common_utils::{poll_document_intelligence, truncate_error_body}; -use super::hooks::OcrRequestPolicy; -use super::runtime_types::PreparedOcrRequest; - -#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] -pub(crate) async fn execute_ocr_provider_call( - request: PreparedOcrRequest, - policy: &OcrRequestPolicy, -) -> Result { - let request = policy.prepare_provider_request(request).await?; - 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 response = http_request(request_builder) - .await - .map_err(|err| Error::Network(err.to_string()))?; - - 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 text = response - .text() - .await - .map_err(|err| Error::Network(err.to_string()))?; - - if !status.is_success() { - return Err(Error::Http { - status: status.as_u16(), - body: truncate_error_body(&text), - }); - } - - let response_json: Value = serde_json::from_str(&text) - .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()) -} diff --git a/litellm-rust/crates/core/src/ocr/hooks.rs b/litellm-rust/crates/core/src/ocr/hooks.rs deleted file mode 100644 index 441e26c6b25..00000000000 --- a/litellm-rust/crates/core/src/ocr/hooks.rs +++ /dev/null @@ -1,283 +0,0 @@ -use crate::error::Error; -use crate::lifecycle::{ActionResult, CallLifecycleContext, RequestPolicy}; -use crate::providers::reducto::ocr::transformation::{ - build_upload_request, extract_document_source, extract_upload_file_id, -}; -use serde_json::{Map, Value, json}; -use std::future::Future; -use std::pin::Pin; - -use super::client::http_client; -use super::common_utils::{convert_document_url_to_data_uri, string_headers, truncate_error_body}; -use super::runtime_types::{PreparedOcrRequest, ProviderOcrRequest}; -use crate::integrations::custom_guardrail::{ - CustomGuardrailRunner, GuardrailContext, GuardrailError, GuardrailRequest, -}; -use crate::integrations::custom_logger::CallType; -use crate::integrations::types::RequestMetadata; - -pub(crate) struct OcrRequestPolicy { - guardrail_runner: CustomGuardrailRunner, - request_metadata: RequestMetadata, -} - -type OcrFuture<'a, T> = Pin> + Send + 'a>>; - -impl OcrRequestPolicy { - pub(crate) fn new( - guardrail_runner: CustomGuardrailRunner, - request_metadata: RequestMetadata, - ) -> Self { - Self { - guardrail_runner, - request_metadata, - } - } - - async fn run_pre_call_guardrails( - &self, - request: PreparedOcrRequest, - ) -> Result { - if self.guardrail_runner.is_empty() { - return Ok(request); - } - - let context = guardrail_context(&self.request_metadata); - let guardrail_request = GuardrailRequest::new(json!({ - "model": request.model, - "custom_llm_provider": request.custom_llm_provider, - "document": request.document, - "optional_params": request.optional_params, - })); - let (guardrail_request, _) = self - .guardrail_runner - .run_pre_call(&context, guardrail_request) - .await - .map_err(guardrail_error_to_core_error)?; - let (document, optional_params) = parse_ocr_pre_call_guardrail_request(guardrail_request)?; - let optional_params = match &request.config { - Ok(config) => config.map_ocr_params(&optional_params), - Err(_) => optional_params, - }; - Ok(PreparedOcrRequest { - document, - optional_params, - ..request - }) - } - - pub(crate) async fn prepare_provider_request( - &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 = 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 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, - config, - url, - body, - optional_params, - upstream_headers, - timeout: request.timeout, - }) - } - - async fn run_during_call_guardrails( - &self, - model: &str, - custom_llm_provider: &str, - url: &str, - body: Value, - ) -> Result { - if self.guardrail_runner.is_empty() { - return Ok(body); - } - - 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, - })); - let (guardrail_request, _) = self - .guardrail_runner - .run_during_call(&context, guardrail_request) - .await - .map_err(guardrail_error_to_core_error)?; - parse_ocr_during_call_guardrail_request(guardrail_request) - } -} - -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 RequestPolicy for OcrRequestPolicy { - type PreCallFuture<'a> = OcrFuture<'a, PreparedOcrRequest>; - type DuringCallFuture<'a> = OcrFuture<'a, PreparedOcrRequest>; - - fn async_pre_call_hook<'a>( - &'a self, - _context: &'a CallLifecycleContext, - request: PreparedOcrRequest, - ) -> Self::PreCallFuture<'a> { - Box::pin(async move { - match self.run_pre_call_guardrails(request).await { - Ok(request) => ActionResult::Replace(request), - Err(error) => ActionResult::Reject(error), - } - }) - } - - fn async_during_call_hook<'a>( - &'a self, - _context: &'a CallLifecycleContext, - request: PreparedOcrRequest, - ) -> Self::DuringCallFuture<'a> { - Box::pin(async move { ActionResult::Continue(request) }) - } -} - -fn guardrail_context(metadata: &RequestMetadata) -> GuardrailContext { - GuardrailContext { - call_type: CallType::Ocr, - selected_guardrails: Vec::new(), - metadata: std::collections::HashMap::new(), - user_api_key_hash: metadata.user_api_key_hash.clone(), - user_api_key_user_id: metadata.user_api_key_user_id.clone(), - user_api_key_team_id: metadata.user_api_key_team_id.clone(), - trace_parent: None, - } -} - -fn parse_ocr_pre_call_guardrail_request( - request: GuardrailRequest, -) -> Result<(Value, Map), Error> { - let Value::Object(mut data) = request.data else { - return Err(Error::InvalidRequest( - "OCR pre_call guardrail must return an object".to_string(), - )); - }; - let document = data.remove("document").ok_or_else(|| { - Error::InvalidRequest("OCR pre_call guardrail removed document".to_string()) - })?; - let optional_params = match data.remove("optional_params") { - Some(Value::Object(params)) => params, - Some(_) => { - return Err(Error::InvalidRequest( - "OCR pre_call guardrail optional_params must be an object".to_string(), - )); - } - None => Map::new(), - }; - Ok((document, optional_params)) -} - -fn parse_ocr_during_call_guardrail_request(request: GuardrailRequest) -> Result { - let Value::Object(mut data) = request.data else { - return Err(Error::InvalidRequest( - "OCR during_call guardrail must return an object".to_string(), - )); - }; - data.remove("body") - .ok_or_else(|| Error::InvalidRequest("OCR during_call guardrail removed body".to_string())) -} - -fn guardrail_error_to_core_error(error: GuardrailError) -> Error { - Error::InvalidRequest(format!("{}: {}", error.kind, error.message)) -} diff --git a/litellm-rust/crates/core/src/ocr/mod.rs b/litellm-rust/crates/core/src/ocr/mod.rs index 7e0d6fae987..f39d90cc892 100644 --- a/litellm-rust/crates/core/src/ocr/mod.rs +++ b/litellm-rust/crates/core/src/ocr/mod.rs @@ -1,9 +1,4 @@ -mod client; -mod common_utils; -mod handler; -mod hooks; pub mod prepare; -mod runtime_types; pub mod transformation; pub mod types; @@ -13,26 +8,17 @@ use crate::Error; use crate::error::json_type_name; use crate::http_utils::{buffered_post, has_header}; -pub use runtime_types::{OcrRequest, OcrRouteRequest}; -pub use types::{OcrAdmissionRequest, OcrResponseData, PreparedOcr, PreparedOcrCall}; +pub use types::{OcrAdmissionRequest, OcrDraft, OcrEndpoint, OcrResponseData, SettledOcrRequest}; use types::{OcrDocument, OcrDocumentProjection}; -use crate::integrations::custom_guardrail::CustomGuardrailRunner; -use crate::integrations::custom_logger::CustomLoggerRunner; use crate::lifecycle::{ CallLifecycle, CallLifecycleContext, Clock, ExecutedCall, TerminalDispatcher, }; -use hooks::OcrRequestPolicy; -use runtime_types::PreparedOcrRequest; pub trait OcrServices: TerminalDispatcher + Clock {} impl OcrServices for T where T: TerminalDispatcher + Clock {} -pub struct DefaultOcrServices { - dispatcher: CustomLoggerRunner, -} - pub struct NoopOcrServices; impl Default for NoopOcrServices { @@ -41,29 +27,6 @@ impl Default for NoopOcrServices { } } -impl DefaultOcrServices { - pub fn new(request: &OcrRequest<'_>) -> Self { - Self { - dispatcher: CustomLoggerRunner::new(request.callbacks.clone()), - } - } -} - -impl Clock for DefaultOcrServices { - fn now(&self) -> f64 { - crate::lifecycle::SystemClock.now() - } -} - -impl TerminalDispatcher for DefaultOcrServices { - fn dispatch<'a>( - &'a self, - terminal: &'a crate::lifecycle::TerminalRecord, - ) -> crate::integrations::custom_logger::LogFuture<'a> { - self.dispatcher.dispatch(terminal) - } -} - impl Clock for NoopOcrServices { fn now(&self) -> f64 { crate::lifecycle::SystemClock.now() @@ -81,85 +44,38 @@ impl TerminalDispatcher for NoopOcrServices { pub async fn ocr( services: &S, - request: OcrRouteRequest<'_>, + request: SettledOcrRequest, _options: crate::lifecycle::ocr::Options, context: CallLifecycleContext, ) -> ExecutedCall { - let OcrRouteRequest::Native(request) = request else { - let OcrRouteRequest::Prepared(request) = request else { - unreachable!() - }; - return CallLifecycle - .run( - context, - request, - &PreparedOcrPolicy, - services, - services, - |request| async move { - send_prepared(request.prepared, request.headers, request.body) - .await - .map(OcrResponseData::into_json) - }, - ) - .await; - }; - let policy = OcrRequestPolicy::new( - CustomGuardrailRunner::new(request.guardrails.clone()), - request.request_metadata.clone(), - ); - let provider = crate::routing_utils::provider::get_custom_llm_provider( - request.model, - request.custom_llm_provider, - ); - let config = provider - .as_ref() - .ok_or_else(|| Error::InvalidProvider("unable to resolve OCR provider".into())) - .and_then(|provider| { - common_utils::ocr_provider_config(provider.custom_llm_provider, provider.model) - .ok_or_else(|| Error::InvalidProvider("unsupported OCR provider".into())) - }); - let provider_model = provider.as_ref().map_or(request.model, |value| value.model); - let provider_name = provider - .as_ref() - .map_or(request.custom_llm_provider.unwrap_or(""), |value| { - value.custom_llm_provider - }); - let prepared = PreparedOcrRequest { - config, - model: provider_model.to_string(), - custom_llm_provider: provider_name.to_string(), - litellm_call_id: context.litellm_call_id.clone(), - 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: request.optional_params, - timeout: request.timeout, - }; CallLifecycle - .run(context, prepared, &policy, services, services, |request| { - handler::execute_ocr_provider_call(request, &policy) - }) + .run( + context, + request, + &SettledOcrPolicy, + services, + services, + |request| async move { send(request).await.map(OcrResponseData::into_json) }, + ) .await } -struct PreparedOcrPolicy; +struct SettledOcrPolicy; -impl crate::lifecycle::RequestPolicy for PreparedOcrPolicy { +impl crate::lifecycle::RequestPolicy for SettledOcrPolicy { type PreCallFuture<'a> - = std::future::Ready> + = std::future::Ready> where Self: 'a; type DuringCallFuture<'a> - = std::future::Ready> + = std::future::Ready> where Self: 'a; fn async_pre_call_hook<'a>( &'a self, _: &'a CallLifecycleContext, - request: PreparedOcrCall, + request: SettledOcrRequest, ) -> Self::PreCallFuture<'a> { std::future::ready(crate::lifecycle::ActionResult::Continue(request)) } @@ -167,18 +83,19 @@ impl crate::lifecycle::RequestPolicy for Prepa fn async_during_call_hook<'a>( &'a self, _: &'a CallLifecycleContext, - request: PreparedOcrCall, + request: SettledOcrRequest, ) -> Self::DuringCallFuture<'a> { std::future::ready(crate::lifecycle::ActionResult::Continue(request)) } } -pub(crate) async fn send_prepared( - prepared: PreparedOcr, - headers: Vec<(String, String)>, - body: Value, -) -> Result { - let config = prepare::provider_config(&prepared.custom_llm_provider, &prepared.model)?; +pub(crate) async fn send(request: SettledOcrRequest) -> Result { + let SettledOcrRequest { + endpoint, + headers, + body, + } = request; + let config = prepare::provider_config(&endpoint.custom_llm_provider, &endpoint.model)?; prepare::validate_capabilities(config)?; let object = body.as_object().ok_or_else(|| Error::InvalidType { expected: "object", @@ -216,10 +133,10 @@ pub(crate) async fn send_prepared( let body = serde_json::to_vec(&body) .map_err(|_| Error::InvalidRequest("could not encode OCR request".into()))?; let response = buffered_post::send(buffered_post::Request { - url: prepared.url, + url: endpoint.url, headers, body, - timeout_seconds: prepared.timeout_seconds, + timeout_seconds: endpoint.timeout_seconds, }) .await?; if !(200..300).contains(&response.status) { @@ -238,5 +155,5 @@ pub(crate) async fn send_prepared( } let response_json = serde_json::from_slice(&response.content) .map_err(|_| Error::InvalidResponse("invalid OCR JSON response".into()))?; - config.transform_ocr_response(&prepared.model, response_json) + config.transform_ocr_response(&endpoint.model, response_json) } diff --git a/litellm-rust/crates/core/src/ocr/prepare.rs b/litellm-rust/crates/core/src/ocr/prepare.rs index 5c181d080f4..bb1a1c6bf2c 100644 --- a/litellm-rust/crates/core/src/ocr/prepare.rs +++ b/litellm-rust/crates/core/src/ocr/prepare.rs @@ -9,8 +9,7 @@ use crate::providers::vertex_ai::ocr::transformation as vertex_ai; use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider}; use super::transformation::{OcrProviderConfig, OcrResponseHandling}; -use super::types::OcrAdmissionRequest; -pub use super::types::PreparedOcr; +use super::types::{OcrAdmissionRequest, OcrDraft, OcrEndpoint}; fn request_config( request: &OcrAdmissionRequest, @@ -81,7 +80,7 @@ fn check_admission_capabilities( Ok(()) } -pub fn prepare(request: OcrAdmissionRequest) -> Result { +pub fn prepare(request: OcrAdmissionRequest) -> Result { let (provider, config) = request_config(&request)?; let env_lookup = |key: &str| std::env::var(key).ok(); let headers = config @@ -131,15 +130,17 @@ pub fn prepare(request: OcrAdmissionRequest) -> Result { let Value::Object(body) = template.data else { return Err(Error::Unsupported("non-object OCR request template")); }; - Ok(PreparedOcr { - model: provider.model.to_string(), - custom_llm_provider: provider.custom_llm_provider.to_string(), - url, + Ok(OcrDraft { + endpoint: OcrEndpoint { + model: provider.model.to_string(), + custom_llm_provider: provider.custom_llm_provider.to_string(), + url, + timeout_seconds: request.timeout_seconds, + }, headers, body, document_projection: config.document_projection(), parameter_fields: config.supported_ocr_params(), - timeout_seconds: request.timeout_seconds, }) } diff --git a/litellm-rust/crates/core/src/ocr/runtime_types.rs b/litellm-rust/crates/core/src/ocr/runtime_types.rs deleted file mode 100644 index a5509fa828c..00000000000 --- a/litellm-rust/crates/core/src/ocr/runtime_types.rs +++ /dev/null @@ -1,77 +0,0 @@ -use std::sync::Arc; -use std::time::Duration; - -use crate::lifecycle::{CallLifecycleContext, CallLifecycleRequest}; -use crate::ocr::transformation::OcrProviderConfig; -use serde_json::{Map, Value}; - -use super::types::PreparedOcrCall; -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 timeout: Option, - pub callbacks: Vec>, - pub guardrails: Vec>, - pub request_metadata: RequestMetadata, - pub litellm_call_id: Option<&'a str>, -} - -pub enum OcrRouteRequest<'a> { - Native(OcrRequest<'a>), - Prepared(PreparedOcrCall), -} - -impl<'a> From> for OcrRouteRequest<'a> { - fn from(request: OcrRequest<'a>) -> Self { - Self::Native(request) - } -} - -impl From for OcrRouteRequest<'static> { - fn from(request: PreparedOcrCall) -> Self { - Self::Prepared(request) - } -} - -pub(crate) struct PreparedOcrRequest { - pub(crate) config: Result<&'static dyn OcrProviderConfig, crate::Error>, - pub(crate) model: String, - pub(crate) custom_llm_provider: String, - pub(crate) litellm_call_id: String, - pub(crate) document: Value, - pub(crate) api_key: Option, - pub(crate) api_base: Option, - pub(crate) extra_headers: Option>, - pub(crate) optional_params: Map, - pub(crate) 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(crate) struct ProviderOcrRequest { - pub(crate) model: 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, -} diff --git a/litellm-rust/crates/core/src/ocr/types.rs b/litellm-rust/crates/core/src/ocr/types.rs index 63a21a51cb4..4eab154c577 100644 --- a/litellm-rust/crates/core/src/ocr/types.rs +++ b/litellm-rust/crates/core/src/ocr/types.rs @@ -16,12 +16,6 @@ pub struct OcrAdmissionRequest { pub stream: bool, } -pub struct PreparedOcrCall { - pub prepared: PreparedOcr, - pub headers: Vec<(String, String)>, - pub body: Value, -} - #[derive(Clone, Debug, Deserialize, Serialize)] #[serde(tag = "type", rename_all = "snake_case")] pub enum OcrDocument { @@ -65,15 +59,51 @@ pub enum OcrDocumentProjection { Transformed, } -pub struct PreparedOcr { - pub model: String, - pub custom_llm_provider: String, - pub url: String, +pub struct OcrDraft { + pub endpoint: OcrEndpoint, pub headers: Vec<(String, String)>, pub body: Map, pub document_projection: OcrDocumentProjection, pub parameter_fields: &'static [&'static str], - pub timeout_seconds: f64, +} + +pub struct OcrEndpoint { + pub(super) model: String, + pub(super) custom_llm_provider: String, + pub(super) url: String, + pub(super) timeout_seconds: f64, +} + +impl OcrEndpoint { + 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 timeout_seconds(&self) -> f64 { + self.timeout_seconds + } + + pub fn settle(self, headers: Vec<(String, String)>, body: Value) -> SettledOcrRequest { + SettledOcrRequest { + endpoint: self, + headers, + body, + } + } +} + +pub struct SettledOcrRequest { + pub(super) endpoint: OcrEndpoint, + pub(super) headers: Vec<(String, String)>, + pub(super) body: Value, } #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] diff --git a/litellm-rust/crates/core/src/realtime/mod.rs b/litellm-rust/crates/core/src/realtime/mod.rs index ec2fbb969a6..5c36ed03829 100644 --- a/litellm-rust/crates/core/src/realtime/mod.rs +++ b/litellm-rust/crates/core/src/realtime/mod.rs @@ -1,2 +1,8 @@ +mod streaming; + pub mod transformation; pub mod types; + +pub use streaming::{ + RealtimeConnectionSpec, RealtimeRequest, WarmConnection, realtime, warmup, +}; diff --git a/litellm-rust/crates/core/src/realtime/streaming.rs b/litellm-rust/crates/core/src/realtime/streaming.rs new file mode 100644 index 00000000000..05db2eeddfb --- /dev/null +++ b/litellm-rust/crates/core/src/realtime/streaming.rs @@ -0,0 +1,561 @@ +use std::hash::{Hash, Hasher}; +use std::pin::Pin; +use std::sync::{Arc, OnceLock}; +use std::task::{Context, Poll}; +use std::time::Duration; + +use futures_util::{Sink, SinkExt, Stream, StreamExt}; +use rustls::{ClientConfig, RootCertStore}; +use serde_json::{Value, json}; +use tokio::net::TcpStream; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; +use tokio_tungstenite::tungstenite::http::HeaderValue; +use tokio_tungstenite::tungstenite::http::header::AUTHORIZATION; +use tokio_tungstenite::tungstenite::{Error as WsError, Message}; +use tokio_tungstenite::{ + Connector, MaybeTlsStream, WebSocketStream, connect_async_tls_with_config, +}; + +use crate::Error; +use crate::integrations::custom_logger::CallbackTiming; +use crate::integrations::types::Usage; +use crate::lifecycle::{ + CallLifecycleContext, Clock, CostInputs, ExecutedCall, RouteProjection, + TerminalClassification, TerminalDispatcher, TerminalRecord, +}; +use crate::providers::openai::realtime::transformation::OPENAI_REALTIME_CONFIG; +use crate::realtime::transformation::RealtimeProviderConfig; +use crate::realtime::types::RealtimeEvent; + +const CONNECT_TIMEOUT: Duration = Duration::from_secs(10); +const IDLE_TIMEOUT: Duration = Duration::from_secs(300); +const OPENAI_API_KEY_ENV: &str = "OPENAI_API_KEY"; +const MISSING_KEY_MESSAGE: &str = "Missing OpenAI API Key - a realtime call is being made but no key was passed via params or the OPENAI_API_KEY environment variable"; + +type Upstream = WebSocketStream>; +static TLS_CONFIG: OnceLock> = OnceLock::new(); + +#[derive(Clone, Eq)] +pub struct RealtimeConnectionSpec { + model: String, + api_key: String, + api_base: Option, +} + +impl RealtimeConnectionSpec { + pub fn new( + model: impl Into, + api_key: Option<&str>, + api_base: Option<&str>, + ) -> Result { + Ok(Self { + model: model.into(), + api_key: resolve_api_key(api_key)?, + api_base: api_base.map(str::to_string), + }) + } + + pub fn model(&self) -> &str { + &self.model + } +} + +impl PartialEq for RealtimeConnectionSpec { + fn eq(&self, other: &Self) -> bool { + self.model == other.model + && self.api_key == other.api_key + && self.api_base == other.api_base + } +} + +impl Hash for RealtimeConnectionSpec { + fn hash(&self, state: &mut H) { + self.model.hash(state); + self.api_key.hash(state); + self.api_base.hash(state); + } +} + +impl std::fmt::Debug for RealtimeConnectionSpec { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("RealtimeConnectionSpec") + .field("model", &self.model) + .field("api_key", &"[REDACTED]") + .field("api_base", &self.api_base) + .finish() + } +} + +pub struct WarmConnection { + upstream: Upstream, + session_created: RealtimeEvent, +} + +impl WarmConnection { + pub fn is_live(&mut self) -> bool { + let mut context = Context::from_waker(futures_util::task::noop_waker_ref()); + matches!(Pin::new(&mut self.upstream).poll_next(&mut context), Poll::Pending) + } +} + +pub struct RealtimeRequest { + pub connection: RealtimeConnectionSpec, + pub warm: Option, + pub idle_timeout: Option, +} + +pub async fn warmup(connection: &RealtimeConnectionSpec) -> Result { + let mut upstream = dial_upstream(connection).await?; + let session_created = read_event(&mut upstream).await?; + if session_created.event_type != "session.created" { + return Err(Error::InvalidResponse(format!( + "expected session.created during realtime warmup, received {}", + session_created.event_type + ))); + } + Ok(WarmConnection { + upstream, + session_created, + }) +} + +pub async fn realtime( + services: &S, + request: RealtimeRequest, + context: CallLifecycleContext, + client_in: In, + client_out: Out, +) -> ExecutedCall<(), Error> +where + S: TerminalDispatcher + Clock, + In: Stream + Unpin + Send, + Out: Sink + Unpin + Send, + Out::Error: std::fmt::Display, +{ + let start_time = services.now(); + let model = request.connection.model.clone(); + let connection = match request.warm { + Some(warm) => Ok(warm), + None => dial_upstream(&request.connection) + .await + .map(|upstream| WarmConnection { + upstream, + session_created: empty_event(), + }), + }; + let mut observation = RealtimeObservation::new(context.litellm_call_id.clone(), model.clone()); + let result = match connection { + Ok(connection) => { + splice( + connection, + &model, + request.idle_timeout.unwrap_or(IDLE_TIMEOUT), + &mut observation, + client_in, + client_out, + ) + .await + } + Err(error) => Err(error), + }; + let classification = match &result { + Ok(()) => TerminalClassification::Success, + Err(error) => TerminalClassification::Failure { + kind: error_kind(error).to_string(), + message: error.to_string(), + }, + }; + let projection = match &classification { + TerminalClassification::Success => Value::Null, + TerminalClassification::Failure { kind, message } => { + json!({"kind": kind, "message": message}) + } + }; + let terminal = TerminalRecord { + call_id: observation.call_id, + trace_id: context.trace_id, + attempt: context.attempt, + call_type: context.call_type, + model: observation.model, + provider: context.custom_llm_provider, + timing: CallbackTiming::new(start_time, services.now()), + usage: observation.usage, + cost_inputs: CostInputs { + response_cost: context.response_cost, + metadata: context.metadata, + }, + classification, + projection: RouteProjection::Realtime { value: projection }, + }; + let _ = services.dispatch(&terminal).await; + match result { + Ok(()) => ExecutedCall::Success { + response: (), + terminal, + }, + Err(error) => ExecutedCall::Failure { error, terminal }, + } +} + +async fn splice( + connection: WarmConnection, + model: &str, + idle_timeout: Duration, + observation: &mut RealtimeObservation, + 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 WarmConnection { + upstream, + session_created, + } = connection; + let (mut upstream_tx, mut upstream_rx) = upstream.split(); + if !session_created.event_type.is_empty() { + observation.observe(&session_created); + send_client_event(&mut client_out, &session_created, model).await?; + } + loop { + tokio::select! { + event = client_in.next() => { + let Some(event) = event else { return Ok(()) }; + for outbound in OPENAI_REALTIME_CONFIG.transform_realtime_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.into())).await.map_err(ws_transport_error)?; + } + } + message = upstream_rx.next() => { + let Some(message) = message else { return Ok(()) }; + match message.map_err(ws_transport_error)? { + Message::Text(text) => { + let event = serde_json::from_str::(&text) + .map_err(|error| Error::InvalidResponse(error.to_string()))?; + observation.observe(&event); + send_client_event(&mut client_out, &event, model).await?; + } + Message::Close(_) => return Ok(()), + _ => {} + } + } + _ = tokio::time::sleep(idle_timeout) => return Ok(()), + } + } +} + +async fn send_client_event( + client_out: &mut Out, + event: &RealtimeEvent, + model: &str, +) -> Result<(), Error> +where + Out: Sink + Unpin, + Out::Error: std::fmt::Display, +{ + for outbound in OPENAI_REALTIME_CONFIG + .transform_realtime_response(event, model)? + .events + { + client_out + .send(outbound) + .await + .map_err(|error| Error::Network(error.to_string()))?; + } + Ok(()) +} + +async fn read_event(upstream: &mut Upstream) -> Result { + loop { + let message = upstream + .next() + .await + .ok_or_else(|| Error::Network("upstream closed before first event".to_string()))? + .map_err(ws_transport_error)?; + match message { + Message::Text(text) => { + return serde_json::from_str(&text) + .map_err(|error| Error::InvalidResponse(error.to_string())); + } + Message::Close(_) => { + return Err(Error::Network( + "upstream closed before first event".to_string(), + )); + } + _ => {} + } + } +} + +async fn dial_upstream(connection: &RealtimeConnectionSpec) -> Result { + let url = OPENAI_REALTIME_CONFIG.complete_url( + connection.api_base.as_deref(), + connection.model.as_str(), + ); + let mut request = url.into_client_request().map_err(ws_transport_error)?; + request.headers_mut().insert( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {}", connection.api_key)) + .map_err(|error| Error::Auth(error.to_string()))?, + ); + let connector = match request.uri().scheme_str() { + Some("wss") => Some(Connector::Rustls(tls_config()?)), + _ => None, + }; + let result = tokio::time::timeout( + CONNECT_TIMEOUT, + connect_async_tls_with_config(request, None, false, connector), + ) + .await + .map_err(|_| Error::Connect("realtime WebSocket connection timed out".to_string()))?; + result.map(|(socket, _)| socket).map_err(ws_handshake_error) +} + +fn tls_config() -> Result, Error> { + if let Some(config) = TLS_CONFIG.get() { + return Ok(Arc::clone(config)); + } + let native = rustls_native_certs::load_native_certs(); + let mut roots = RootCertStore::empty(); + let (added, _) = roots.add_parsable_certificates(native.certs); + if added == 0 { + return Err(Error::Connect(format!( + "no usable native root certificates: {:?}", + native.errors + ))); + } + let config = ClientConfig::builder_with_provider(Arc::new( + rustls::crypto::ring::default_provider(), + )) + .with_safe_default_protocol_versions() + .map_err(|error| Error::Connect(error.to_string()))? + .with_root_certificates(roots) + .with_no_client_auth(); + let config = Arc::new(config); + Ok(Arc::clone(TLS_CONFIG.get_or_init(|| config))) +} + +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())) +} + +fn ws_handshake_error(error: WsError) -> Error { + match error { + WsError::Http(response) => Error::Http { + status: response.status().as_u16(), + body: response + .body() + .as_ref() + .map(|body| String::from_utf8_lossy(body).into_owned()) + .unwrap_or_default(), + }, + other => ws_transport_error(other), + } +} + +fn ws_transport_error(error: WsError) -> Error { + match error { + WsError::Io(error) => Error::Connect(error.to_string()), + other => Error::Network(other.to_string()), + } +} + +fn error_kind(error: &Error) -> &'static str { + match error { + Error::Auth(_) => "AuthError", + Error::InvalidProvider(_) => "InvalidProvider", + Error::InvalidRequest(_) => "InvalidRequest", + Error::InvalidType { .. } => "InvalidType", + Error::MissingField(_) => "MissingField", + Error::Http { .. } => "HttpError", + Error::InvalidResponse(_) => "InvalidResponse", + Error::Network(_) => "NetworkError", + Error::Connect(_) => "ConnectError", + Error::Routing(_) => "RoutingError", + Error::Unsupported(_) => "UnsupportedRequest", + } +} + +fn empty_event() -> RealtimeEvent { + RealtimeEvent { + event_type: String::new(), + data: Default::default(), + } +} + +struct RealtimeObservation { + call_id: String, + model: String, + usage: Usage, +} + +impl RealtimeObservation { + fn new(call_id: String, model: String) -> Self { + Self { + call_id, + model, + usage: Usage::default(), + } + } + + fn observe(&mut self, event: &RealtimeEvent) { + if event.event_type == "session.created" { + let session = event.data.get("session").and_then(Value::as_object); + if let Some(id) = session + .and_then(|value| value.get("id")) + .and_then(Value::as_str) + .filter(|value| !value.is_empty()) + { + self.call_id = id.to_string(); + } + if let Some(model) = session + .and_then(|value| value.get("model")) + .and_then(Value::as_str) + .filter(|value| !value.is_empty()) + { + self.model = model.to_string(); + } + return; + } + if event.event_type != "response.done" { + return; + } + let Some(usage) = event + .data + .get("response") + .and_then(Value::as_object) + .and_then(|value| value.get("usage")) + .and_then(Value::as_object) + else { + return; + }; + let input = usage.get("input_tokens").and_then(Value::as_u64).unwrap_or(0); + let output = usage + .get("output_tokens") + .and_then(Value::as_u64) + .unwrap_or(0); + self.usage.prompt_tokens += input; + self.usage.completion_tokens += output; + self.usage.total_tokens += usage + .get("total_tokens") + .and_then(Value::as_u64) + .unwrap_or(input + output); + } +} + +#[cfg(test)] +mod tests { + use std::sync::Mutex; + + use futures_channel::mpsc; + use tokio::net::TcpListener; + use tokio_tungstenite::accept_async; + + use super::*; + use crate::integrations::custom_logger::{LogError, LogFuture}; + use crate::lifecycle::{CallLifecycleContext, Clock}; + + #[derive(Default)] + struct Services { + terminals: Mutex>, + } + + impl Clock for Services { + fn now(&self) -> f64 { + 1.0 + } + } + + impl TerminalDispatcher for Services { + fn dispatch<'a>(&'a self, terminal: &'a TerminalRecord) -> LogFuture<'a> { + Box::pin(async move { + self.terminals.lock().unwrap().push(terminal.clone()); + Ok::<(), LogError>(()) + }) + } + } + + async fn provider() -> String { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + tokio::spawn(async move { + while let Ok((stream, _)) = listener.accept().await { + tokio::spawn(async move { + let mut socket = accept_async(stream).await.unwrap(); + socket.send(Message::Text(json!({"type":"session.created","session":{"id":"sess-core","model":"upstream-model"}}).to_string().into())).await.unwrap(); + while let Some(Ok(Message::Text(text))) = socket.next().await { + let event: RealtimeEvent = serde_json::from_str(&text).unwrap(); + if event.event_type == "response.create" { + socket.send(Message::Text(json!({"type":"response.done","response":{"usage":{"input_tokens":2,"output_tokens":3,"total_tokens":5}}}).to_string().into())).await.unwrap(); + socket.send(Message::Text(json!({"type":"response.done","response":{"usage":{"input_tokens":7,"output_tokens":11}}}).to_string().into())).await.unwrap(); + socket.close(None).await.unwrap(); + } + } + }); + } + }); + format!("ws://{address}") + } + + async fn execute(warm: bool) -> (ExecutedCall<(), Error>, Vec, usize) { + let base = provider().await; + let spec = RealtimeConnectionSpec::new("requested", Some("key"), Some(&base)).unwrap(); + let services = Services::default(); + let warm = if warm { Some(warmup(&spec).await.unwrap()) } else { None }; + assert!(services.terminals.lock().unwrap().is_empty()); + let (input_tx, input) = mpsc::unbounded(); + let (output, mut output_rx) = mpsc::unbounded(); + input_tx.unbounded_send(serde_json::from_value(json!({"type":"response.done","response":{"usage":{"input_tokens":1000,"output_tokens":1000,"total_tokens":2000}}})).unwrap()).unwrap(); + input_tx.unbounded_send(serde_json::from_value(json!({"type":"response.create"})).unwrap()).unwrap(); + let result = realtime( + &services, + RealtimeRequest { connection: spec, warm, idle_timeout: Some(Duration::from_secs(1)) }, + CallLifecycleContext::new("realtime", "requested", "openai", "fallback"), + input, + output, + ).await; + let mut events = Vec::new(); + while let Ok(Some(event)) = tokio::time::timeout(Duration::from_millis(10), output_rx.next()).await { + events.push(event); + } + let count = services.terminals.lock().unwrap().len(); + (result, events, count) + } + + #[tokio::test] + async fn fresh_and_warm_sessions_share_identity_usage_and_exactly_once_terminal() { + for warm in [false, true] { + let (result, events, count) = execute(warm).await; + assert_eq!(count, 1); + assert_eq!(events.first().unwrap().event_type, "session.created"); + let ExecutedCall::Success { terminal, .. } = result else { panic!("session failed") }; + assert_eq!(terminal.call_id, "sess-core"); + assert_eq!(terminal.model, "upstream-model"); + assert_eq!(terminal.usage, Usage { prompt_tokens: 9, completion_tokens: 14, total_tokens: 23 }); + } + } + + #[tokio::test] + async fn warmup_success_and_failure_dispatch_nothing() { + let services = Services::default(); + let base = provider().await; + let good = RealtimeConnectionSpec::new("model", Some("key"), Some(&base)).unwrap(); + assert!(warmup(&good).await.is_ok()); + let bad = RealtimeConnectionSpec::new("model", Some("key"), Some("ws://127.0.0.1:1")).unwrap(); + assert!(warmup(&bad).await.is_err()); + assert!(services.terminals.lock().unwrap().is_empty()); + } +} diff --git a/litellm-rust/crates/core/src/responses/instrumentation.rs b/litellm-rust/crates/core/src/responses/instrumentation.rs index 007bec7b710..59f12449dd0 100644 --- a/litellm-rust/crates/core/src/responses/instrumentation.rs +++ b/litellm-rust/crates/core/src/responses/instrumentation.rs @@ -1,103 +1,20 @@ -use std::future::Future; -use std::pin::Pin; -use std::sync::Mutex; -use std::time::{SystemTime, UNIX_EPOCH}; - -use serde_json::Value; - -use crate::Error; -use crate::integrations::custom_logger::{LogError, LogFuture}; -use crate::lifecycle::{ - ActionResult, CallLifecycleContext, RequestPolicy, TerminalDispatcher, TerminalRecord, -}; +use crate::integrations::types::Usage; use crate::responses::types::{ResponsesWsEvent, ResponsesWsEventType}; +use serde_json::Value; +use std::sync::Mutex; #[derive(Clone, Debug, Default, PartialEq, Eq)] -pub struct ResponsesWsUsage { - pub prompt_tokens: u64, - pub completion_tokens: u64, - pub total_tokens: u64, -} - -#[derive(Clone, Debug, Default, PartialEq, Eq)] -pub struct ResponsesWsMetadata { - pub user_api_key_hash: Option, - pub user_api_key_user_id: Option, - pub user_api_key_team_id: Option, -} - -#[derive(Clone, Debug, PartialEq)] -pub struct ResponsesWsLogPayload { - pub id: String, - pub litellm_call_id: String, - pub call_type: String, - pub model: String, - pub custom_llm_provider: String, - pub response_cost: f64, - pub usage: ResponsesWsUsage, - pub start_time: f64, - pub end_time: f64, - pub stream: bool, - pub metadata: ResponsesWsMetadata, -} - -#[derive(Clone, Debug, PartialEq)] -pub enum ResponsesWsLogOutcome { - Success { - payload: ResponsesWsLogPayload, - callback: ResponsesWsCallbackPayload, - }, - Failure { - payload: ResponsesWsLogPayload, - callback: ResponsesWsCallbackPayload, - error_message: String, - error_kind: String, - }, -} - -#[derive(Clone, Debug, PartialEq)] -pub struct ResponsesWsCallbackPayload { - pub object: String, - pub value: Value, -} - -struct InstrumentationState { - litellm_call_id: String, - id: String, - model: String, - usage: ResponsesWsUsage, - start_time: f64, - end_time: f64, - metadata: ResponsesWsMetadata, - outcome: Option, +pub struct ResponsesWsObservation { + pub(crate) model: String, + pub(crate) usage: Usage, } +#[derive(Default)] pub struct ResponsesWsInstrumentation { - state: Mutex, + state: Mutex, } impl ResponsesWsInstrumentation { - pub fn new( - litellm_call_id: impl Into, - model: impl Into, - metadata: ResponsesWsMetadata, - ) -> Self { - let litellm_call_id = litellm_call_id.into(); - let now = epoch_seconds(); - Self { - state: Mutex::new(InstrumentationState { - id: litellm_call_id.clone(), - litellm_call_id, - model: model.into(), - usage: ResponsesWsUsage::default(), - start_time: now, - end_time: now, - metadata, - outcome: None, - }), - } - } - pub fn observe(&self, event: &ResponsesWsEvent) { if !matches!( event.event_type, @@ -115,14 +32,6 @@ impl ResponsesWsInstrumentation { let Some(response) = event.data.get("response").and_then(Value::as_object) else { return; }; - if let Some(id) = response - .get("id") - .and_then(Value::as_str) - .filter(|value| !value.is_empty()) - { - state.id = id.to_string(); - state.litellm_call_id = id.to_string(); - } if let Some(model) = response .get("model") .and_then(Value::as_str) @@ -154,124 +63,18 @@ impl ResponsesWsInstrumentation { }); } - pub fn success_outcome(&self) -> ResponsesWsLogOutcome { - let mut state = self - .state - .lock() - .unwrap_or_else(|poisoned| poisoned.into_inner()); - state.end_time = epoch_seconds(); - ResponsesWsLogOutcome::Success { - payload: build_payload(&state), - callback: ResponsesWsCallbackPayload { - object: "responses_websocket".to_string(), - value: Value::Null, - }, - } - } - - pub fn failure_outcome(&self) -> ResponsesWsLogOutcome { - let mut state = self - .state - .lock() - .unwrap_or_else(|poisoned| poisoned.into_inner()); - state.end_time = epoch_seconds(); - ResponsesWsLogOutcome::Failure { - payload: build_payload(&state), - callback: ResponsesWsCallbackPayload { - object: "error".to_string(), - value: serde_json::json!({ - "message": "Responses WebSocket session ended in failure", - "kind": "ResponsesWebSocketError", - }), - }, - error_message: "Responses WebSocket session ended in failure".to_string(), - error_kind: "ResponsesWebSocketError".to_string(), - } - } - - pub fn take_outcome(&self) -> Option { + pub fn snapshot(&self) -> ResponsesWsObservation { self.state .lock() .unwrap_or_else(|poisoned| poisoned.into_inner()) - .outcome - .take() + .clone() } - - pub fn take_or_build_outcome(&self, success: bool) -> ResponsesWsLogOutcome { - self.take_outcome().unwrap_or_else(|| { - if success { - self.success_outcome() - } else { - self.failure_outcome() - } - }) - } -} - -type LifecycleFuture<'a, T> = Pin> + Send + 'a>>; - -impl RequestPolicy<(), ()> for ResponsesWsInstrumentation { - type PreCallFuture<'a> = LifecycleFuture<'a, ()>; - type DuringCallFuture<'a> = LifecycleFuture<'a, ()>; - - fn async_pre_call_hook<'a>( - &'a self, - _context: &'a CallLifecycleContext, - request: (), - ) -> Self::PreCallFuture<'a> { - Box::pin(async move { ActionResult::Continue(request) }) - } - - fn async_during_call_hook<'a>( - &'a self, - _context: &'a CallLifecycleContext, - request: (), - ) -> Self::DuringCallFuture<'a> { - Box::pin(async move { ActionResult::Continue(request) }) - } -} - -impl TerminalDispatcher for ResponsesWsInstrumentation { - fn dispatch<'a>(&'a self, terminal: &'a TerminalRecord) -> LogFuture<'a> { - Box::pin(async move { - let outcome = match terminal.classification { - crate::lifecycle::TerminalClassification::Success => self.success_outcome(), - crate::lifecycle::TerminalClassification::Failure { .. } => self.failure_outcome(), - }; - if let Ok(mut state) = self.state.lock() { - state.outcome = Some(outcome); - } - Ok::<(), LogError>(()) - }) - } -} - -fn build_payload(state: &InstrumentationState) -> ResponsesWsLogPayload { - ResponsesWsLogPayload { - id: state.id.clone(), - litellm_call_id: state.litellm_call_id.clone(), - call_type: "responses_websocket".to_string(), - model: state.model.clone(), - custom_llm_provider: "openai".to_string(), - response_cost: 0.0, - usage: state.usage.clone(), - start_time: state.start_time, - end_time: state.end_time, - stream: true, - metadata: state.metadata.clone(), - } -} - -fn epoch_seconds() -> f64 { - SystemTime::now() - .duration_since(UNIX_EPOCH) - .map(|duration| duration.as_secs_f64()) - .unwrap_or(0.0) } #[cfg(test)] mod tests { use super::*; + use serde_json::Value; fn event(value: Value) -> ResponsesWsEvent { serde_json::from_value(value).expect("valid Responses WebSocket event") @@ -279,8 +82,7 @@ mod tests { #[test] fn accumulates_upstream_usage_and_identity() { - let instrumentation = - ResponsesWsInstrumentation::new("call-1", "gpt-5", ResponsesWsMetadata::default()); + let instrumentation = ResponsesWsInstrumentation::default(); instrumentation.observe(&event(serde_json::json!({ "type": "response.completed", "response": { @@ -294,65 +96,10 @@ mod tests { } }))); - let ResponsesWsLogOutcome::Success { payload, .. } = instrumentation.success_outcome() - else { - panic!("expected success outcome"); - }; - assert_eq!(payload.id, "resp-1"); - assert_eq!(payload.model, "gpt-5-mini"); - assert_eq!(payload.usage.prompt_tokens, 3); - assert_eq!(payload.usage.completion_tokens, 5); - assert_eq!(payload.usage.total_tokens, 8); - assert!(payload.end_time >= payload.start_time); - } - - #[test] - fn builds_failure_payload_without_dispatching_callbacks() { - let instrumentation = - ResponsesWsInstrumentation::new("call-1", "gpt-5", ResponsesWsMetadata::default()); - assert!(matches!( - instrumentation.failure_outcome(), - ResponsesWsLogOutcome::Failure { .. } - )); - } - - #[tokio::test] - async fn lifecycle_records_success_outcome_for_provider_completion() { - let instrumentation = - ResponsesWsInstrumentation::new("call-1", "gpt-5", ResponsesWsMetadata::default()); - let result = crate::lifecycle::CallLifecycle - .run( - crate::lifecycle::CallLifecycleContext::new( - "responses_websocket", - "gpt-5", - "openai", - "call-1", - ), - (), - &instrumentation, - &instrumentation, - &crate::lifecycle::SystemClock, - |_| async { Ok::<(), Error>(()) }, - ) - .await; - - assert!(matches!( - result, - crate::lifecycle::ExecutedCall::Success { .. } - )); - assert!(matches!( - instrumentation.take_outcome(), - Some(ResponsesWsLogOutcome::Success { .. }) - )); - } - - #[test] - fn builds_outcome_when_lifecycle_did_not_record_one() { - let instrumentation = - ResponsesWsInstrumentation::new("call-1", "gpt-5", ResponsesWsMetadata::default()); - assert!(matches!( - instrumentation.take_or_build_outcome(true), - ResponsesWsLogOutcome::Success { .. } - )); + let observation = instrumentation.snapshot(); + assert_eq!(observation.model, "gpt-5-mini"); + assert_eq!(observation.usage.prompt_tokens, 3); + assert_eq!(observation.usage.completion_tokens, 5); + assert_eq!(observation.usage.total_tokens, 8); } } diff --git a/litellm-rust/crates/core/src/responses/websocket.rs b/litellm-rust/crates/core/src/responses/websocket.rs index 5d037e9cf1b..c7816f17337 100644 --- a/litellm-rust/crates/core/src/responses/websocket.rs +++ b/litellm-rust/crates/core/src/responses/websocket.rs @@ -1,7 +1,37 @@ +use std::sync::{Arc, OnceLock}; +use std::time::Duration; + +use futures_util::{Sink, SinkExt, Stream, StreamExt}; +use rustls::{ClientConfig, RootCertStore}; +use serde_json::{Value, json}; +use tokio::net::TcpStream; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; +use tokio_tungstenite::tungstenite::http::HeaderValue; +use tokio_tungstenite::tungstenite::http::header::AUTHORIZATION; +use tokio_tungstenite::tungstenite::{Error as WsError, Message}; +use tokio_tungstenite::{ + Connector, MaybeTlsStream, WebSocketStream, connect_async_tls_with_config, +}; + use crate::Error; use crate::constants::{OPENAI_RESPONSES_DEFAULT_API_BASE, OPENAI_RESPONSES_PATH}; +use crate::integrations::custom_logger::CallbackTiming; +use crate::lifecycle::{ + CostInputs, ExecutedCall, RouteProjection, TerminalClassification, TerminalDispatcher, + TerminalRecord, +}; +use crate::providers::openai::responses::transformation::OPENAI_RESPONSES_WS_CONFIG; +use crate::responses::instrumentation::ResponsesWsInstrumentation; use crate::responses::types::{ResponsesWsEvent, ResponsesWsEventType, ResponsesWsTransformResult}; +const CONNECT_TIMEOUT: Duration = Duration::from_secs(10); +const IDLE_TIMEOUT: Duration = Duration::from_secs(300); +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"; + +type Upstream = WebSocketStream>; +static TLS_CONFIG: OnceLock> = OnceLock::new(); + pub trait ResponsesWebSocketProviderConfig: Sync { fn supports_native_websocket(&self) -> bool { false @@ -28,36 +58,255 @@ pub trait ResponsesWebSocketProviderConfig: Sync { ) -> Result; } -pub fn complete_websocket_url( - api_base: Option<&str>, +pub struct ResponsesWebSocketRequest { + pub model: String, + pub api_key: Option, + pub api_base: Option, + pub first_frame: Option, + pub idle_timeout: Option, +} + +pub async fn responses_websocket( + services: &S, + request: ResponsesWebSocketRequest, + context: crate::lifecycle::CallLifecycleContext, + client_in: In, + client_out: Out, +) -> Result, Error> +where + S: TerminalDispatcher + crate::lifecycle::Clock, + In: Stream + Unpin + Send, + Out: Sink + Unpin + Send, + Out::Error: std::fmt::Display, +{ + let key = resolve_api_key(request.api_key.as_deref())?; + let upstream = dial_upstream(&request.model, &key, request.api_base.as_deref()).await?; + let start_time = services.now(); + let instrumentation = ResponsesWsInstrumentation::default(); + let result = splice( + upstream, + &request.model, + request.first_frame, + request.idle_timeout.unwrap_or(IDLE_TIMEOUT), + &instrumentation, + client_in, + client_out, + ) + .await; + let observation = instrumentation.snapshot(); + let model = if observation.model.is_empty() { + context.model.clone() + } else { + observation.model + }; + let classification = match &result { + Ok(()) => TerminalClassification::Success, + Err(error) => TerminalClassification::Failure { + kind: error_kind(error).to_string(), + message: error.to_string(), + }, + }; + let projection = match &classification { + TerminalClassification::Success => Value::Null, + TerminalClassification::Failure { kind, message } => { + json!({"kind": kind, "message": message}) + } + }; + let terminal = TerminalRecord { + call_id: context.litellm_call_id, + trace_id: context.trace_id, + attempt: context.attempt, + call_type: context.call_type, + model, + provider: context.custom_llm_provider, + timing: CallbackTiming::new(start_time, services.now()), + usage: observation.usage, + cost_inputs: CostInputs { + response_cost: context.response_cost, + metadata: context.metadata, + }, + classification, + projection: RouteProjection::ResponsesWs { value: projection }, + }; + let _ = services.dispatch(&terminal).await; + Ok(match result { + Ok(()) => ExecutedCall::Success { + response: (), + terminal, + }, + Err(error) => ExecutedCall::Failure { error, terminal }, + }) +} + +async fn splice( + upstream: Upstream, model: &str, - model_in_websocket_url: bool, -) -> String { + first_frame: Option, + idle_timeout: Duration, + instrumentation: &ResponsesWsInstrumentation, + 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 (mut upstream_tx, mut upstream_rx) = upstream.split(); + if let Some(event) = first_frame { + send_provider_event(&mut upstream_tx, &event, model).await?; + } + loop { + tokio::select! { + event = client_in.next() => { + let Some(event) = event else { return Ok(()) }; + send_provider_event(&mut upstream_tx, &event, model).await?; + } + message = upstream_rx.next() => { + let Some(message) = message else { return Ok(()) }; + match message.map_err(ws_transport_error)? { + Message::Text(text) => { + let event = serde_json::from_str::(&text) + .map_err(|error| Error::InvalidResponse(error.to_string()))?; + instrumentation.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(_) => return Ok(()), + _ => {} + } + } + _ = tokio::time::sleep(idle_timeout) => return Ok(()), + } + } +} + +async fn send_provider_event( + upstream: &mut futures_util::stream::SplitSink, + event: &ResponsesWsEvent, + model: &str, +) -> Result<(), Error> { + 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 + .send(Message::Text(payload.into())) + .await + .map_err(ws_transport_error)?; + } + Ok(()) +} + +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.into_client_request().map_err(ws_transport_error)?; + request.headers_mut().insert( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {api_key}")) + .map_err(|error| Error::Auth(error.to_string()))?, + ); + let connector = match request.uri().scheme_str() { + Some("wss") => Some(Connector::Rustls(tls_config()?)), + _ => None, + }; + let connect = connect_async_tls_with_config(request, None, false, connector); + let result = tokio::time::timeout(CONNECT_TIMEOUT, connect) + .await + .map_err(|_| Error::Connect("Responses WebSocket connection timed out".to_string()))?; + result.map(|(socket, _)| socket).map_err(ws_handshake_error) +} + +fn tls_config() -> Result, Error> { + if let Some(config) = TLS_CONFIG.get() { + return Ok(Arc::clone(config)); + } + let native = rustls_native_certs::load_native_certs(); + let mut roots = RootCertStore::empty(); + let (added, _) = roots.add_parsable_certificates(native.certs); + if added == 0 { + return Err(Error::Connect(format!( + "no usable native root certificates: {:?}", + native.errors + ))); + } + let config = + ClientConfig::builder_with_provider(Arc::new(rustls::crypto::ring::default_provider())) + .with_safe_default_protocol_versions() + .map_err(|error| Error::Connect(error.to_string()))? + .with_root_certificates(roots) + .with_no_client_auth(); + let config = Arc::new(config); + Ok(Arc::clone(TLS_CONFIG.get_or_init(|| config))) +} + +fn ws_handshake_error(error: WsError) -> Error { + match error { + WsError::Http(response) => Error::Http { + status: response.status().as_u16(), + body: response + .body() + .as_ref() + .map(|body| String::from_utf8_lossy(body).into_owned()) + .unwrap_or_default(), + }, + other => ws_transport_error(other), + } +} + +fn ws_transport_error(error: WsError) -> Error { + match error { + WsError::Io(error) => Error::Connect(error.to_string()), + other => Error::Network(other.to_string()), + } +} + +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())) +} + +pub fn complete_websocket_url(api_base: Option<&str>, model: &str, model_in_url: bool) -> String { let base = api_base .map(str::trim) .filter(|value| !value.is_empty()) .unwrap_or(OPENAI_RESPONSES_DEFAULT_API_BASE); - let (base_without_query, query) = base + let (base, query) = base .split_once('?') - .map_or((base, None), |(value, query)| (value, Some(query))); - let response_url = format!( - "{}{}", - base_without_query.trim_end_matches('/'), - OPENAI_RESPONSES_PATH + .map_or((base, None), |(base, query)| (base, Some(query))); + let response_url = format!("{}{}", base.trim_end_matches('/'), OPENAI_RESPONSES_PATH); + let response_url = response_url + .strip_prefix("https://") + .map(|rest| format!("wss://{rest}")) + .or_else(|| { + response_url + .strip_prefix("http://") + .map(|rest| format!("ws://{rest}")) + }) + .unwrap_or(response_url); + let url = query.map_or_else( + || response_url.clone(), + |query| format!("{response_url}?{query}"), ); - let scheme_flipped = if let Some(rest) = response_url.strip_prefix("https://") { - format!("wss://{rest}") - } else if let Some(rest) = response_url.strip_prefix("http://") { - format!("ws://{rest}") - } else { - response_url - }; - let url = query.map_or(scheme_flipped.clone(), |value| { - format!("{scheme_flipped}?{value}") - }); - if !model_in_websocket_url - || query.is_some_and(|value| { - value + if !model_in_url + || query.is_some_and(|query| { + query .split('&') .any(|part| part.split('=').next() == Some("model")) }) @@ -93,23 +342,18 @@ pub fn enforce_model(event: &ResponsesWsEvent, model: &str) -> ResponsesWsEvent if let Some(response) = enforced .data .get_mut("response") - .and_then(serde_json::Value::as_object_mut) + .and_then(Value::as_object_mut) { - response.insert( - "model".to_string(), - serde_json::Value::String(model.to_string()), - ); + response.insert("model".to_string(), Value::String(model.to_string())); if has_flat_model { - enforced.data.insert( - "model".to_string(), - serde_json::Value::String(model.to_string()), - ); + enforced + .data + .insert("model".to_string(), Value::String(model.to_string())); } } else { - enforced.data.insert( - "model".to_string(), - serde_json::Value::String(model.to_string()), - ); + enforced + .data + .insert("model".to_string(), Value::String(model.to_string())); } enforced } @@ -125,16 +369,186 @@ pub fn is_terminal_event(event_type: &ResponsesWsEventType) -> bool { ) } +fn error_kind(error: &Error) -> &'static str { + match error { + Error::Auth(_) => "AuthError", + Error::InvalidProvider(_) => "InvalidProvider", + Error::InvalidRequest(_) => "InvalidRequest", + Error::InvalidType { .. } => "InvalidType", + Error::MissingField(_) => "MissingField", + Error::Http { .. } => "HttpError", + Error::InvalidResponse(_) => "InvalidResponse", + Error::Network(_) => "NetworkError", + Error::Connect(_) => "ConnectError", + Error::Routing(_) => "RoutingError", + Error::Unsupported(_) => "UnsupportedRequest", + } +} + #[cfg(test)] mod tests { use super::*; + use crate::integrations::custom_logger::{LogError, LogFuture}; + use crate::lifecycle::{CallLifecycleContext, Clock}; + use std::sync::Mutex; + use tokio::io::AsyncWriteExt; + use tokio::net::TcpListener; + use tokio_tungstenite::accept_async; - fn event(value: serde_json::Value) -> ResponsesWsEvent { - serde_json::from_value(value).expect("valid event") + #[derive(Default)] + struct Services { + terminals: Mutex>, + } + impl Clock for Services { + fn now(&self) -> f64 { + 1.0 + } + } + impl TerminalDispatcher for Services { + fn dispatch<'a>(&'a self, terminal: &'a TerminalRecord) -> LogFuture<'a> { + Box::pin(async move { + self.terminals.lock().unwrap().push(terminal.clone()); + Ok::<(), LogError>(()) + }) + } + } + + fn event(value: Value) -> ResponsesWsEvent { + serde_json::from_value(value).expect("event") + } + + async fn mock_provider() -> (String, tokio::task::JoinHandle<()>) { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let task = tokio::spawn(async move { + let (stream, _) = listener.accept().await.unwrap(); + let mut socket = accept_async(stream).await.unwrap(); + if let Some(Ok(Message::Text(text))) = socket.next().await { + let request: Value = serde_json::from_str(&text).unwrap(); + assert_eq!(request["model"], "authorized"); + socket.send(Message::Text(json!({"type":"response.completed","response":{"id":"resp-1","model":"authorized","usage":{"input_tokens":1,"output_tokens":2,"total_tokens":3}}}).to_string().into())).await.unwrap(); + socket.close(None).await.unwrap(); + } + }); + (format!("http://{address}"), task) + } + + #[tokio::test] + async fn core_owns_splice_transformation_and_one_terminal() { + let (api_base, server) = mock_provider().await; + let services = Services::default(); + let (client_tx, client_rx) = futures_channel::mpsc::unbounded(); + let (output_tx, mut output_rx) = futures_channel::mpsc::unbounded(); + client_tx + .unbounded_send(event(json!({"type":"response.create","model":"wrong"}))) + .unwrap(); + let result = responses_websocket( + &services, + ResponsesWebSocketRequest { + model: "authorized".into(), + api_key: Some("key".into()), + api_base: Some(api_base), + first_frame: None, + idle_timeout: Some(Duration::from_secs(1)), + }, + CallLifecycleContext::new("responses_websocket", "authorized", "openai", "call-1"), + client_rx, + output_tx, + ) + .await + .unwrap(); + assert!(matches!(result, ExecutedCall::Success { .. })); + assert_eq!( + output_rx.next().await.unwrap().event_type, + ResponsesWsEventType::ResponseCompleted + ); + let terminals = services.terminals.lock().unwrap(); + assert_eq!(terminals.len(), 1); + assert_eq!(terminals[0].usage.total_tokens, 3); + server.await.unwrap(); + } + + #[tokio::test] + async fn handshake_status_is_preserved_without_a_terminal() { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + tokio::spawn(async move { + let (mut stream, _) = listener.accept().await.unwrap(); + stream + .write_all(b"HTTP/1.1 429 Too Many Requests\r\nContent-Length: 4\r\n\r\nslow") + .await + .unwrap(); + }); + let services = Services::default(); + let (_, input): ( + _, + futures_channel::mpsc::UnboundedReceiver, + ) = futures_channel::mpsc::unbounded(); + let (output, _) = futures_channel::mpsc::unbounded(); + let error = responses_websocket( + &services, + ResponsesWebSocketRequest { + model: "model".into(), + api_key: Some("key".into()), + api_base: Some(format!("http://{address}")), + first_frame: None, + idle_timeout: None, + }, + CallLifecycleContext::new("responses_websocket", "model", "openai", "call-1"), + input, + output, + ) + .await + .unwrap_err(); + assert!(matches!(error, Error::Http { status: 429, .. })); + assert!(services.terminals.lock().unwrap().is_empty()); + } + + #[tokio::test] + async fn committed_protocol_failure_dispatches_one_terminal() { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + tokio::spawn(async move { + let (stream, _) = listener.accept().await.unwrap(); + let mut socket = accept_async(stream).await.unwrap(); + socket.send(Message::Text("not-json".into())).await.unwrap(); + }); + let services = Services::default(); + let (client_tx, input) = futures_channel::mpsc::unbounded(); + let (output, _) = futures_channel::mpsc::unbounded(); + let result = responses_websocket( + &services, + ResponsesWebSocketRequest { + model: "model".into(), + api_key: Some("key".into()), + api_base: Some(format!("http://{address}")), + first_frame: None, + idle_timeout: None, + }, + CallLifecycleContext::new("responses_websocket", "model", "openai", "call-1"), + input, + output, + ) + .await + .unwrap(); + drop(client_tx); + assert!(matches!( + result, + ExecutedCall::Failure { + error: Error::InvalidResponse(_), + .. + } + )); + let terminals = services.terminals.lock().unwrap(); + assert_eq!(terminals.len(), 1); + assert!(matches!( + terminals[0].classification, + TerminalClassification::Failure { .. } + )); } #[test] - fn url_construction_matches_python_defaults_and_query_behavior() { + fn url_and_model_behavior_match_the_public_protocol() { assert_eq!( complete_websocket_url(None, "gpt-5", true), "wss://api.openai.com/v1/responses?model=gpt-5" @@ -143,46 +557,11 @@ mod tests { complete_websocket_url(Some("http://localhost:8080/"), "gpt 5", true), "ws://localhost:8080/responses?model=gpt%205" ); - assert_eq!( - complete_websocket_url(Some("https://example.test/v1?foo=bar"), "gpt-5", true), - "wss://example.test/v1/responses?foo=bar&model=gpt-5" - ); - assert_eq!( - complete_websocket_url(Some("https://example.test?model=existing"), "gpt-5", true), - "wss://example.test/responses?model=existing" - ); - } - - #[test] - fn enforce_model_overrides_flat_and_nested_values() { - let flat = enforce_model( - &event(serde_json::json!({"type":"response.create","model":"wrong"})), - "gpt-5", - ); - assert_eq!(flat.model(), Some("gpt-5")); let nested = enforce_model( - &event(serde_json::json!({ - "type":"response.create", - "model":"wrong", - "response":{"model":"also-wrong"} - })), - "gpt-5", + &event(json!({"type":"response.create","model":"wrong","response":{"model":"wrong"}})), + "right", ); - assert_eq!(nested.model(), Some("gpt-5")); - assert_eq!( - nested - .data - .get("response") - .and_then(|value| value.get("model")), - Some(&serde_json::json!("gpt-5")) - ); - let nested_without_flat = enforce_model( - &event(serde_json::json!({ - "type":"response.create", - "response":{"model":"also-wrong"} - })), - "gpt-5", - ); - assert!(!nested_without_flat.data.contains_key("model")); + assert_eq!(nested.model(), Some("right")); + assert_eq!(nested.data["response"]["model"], "right"); } } diff --git a/litellm-rust/crates/core/tests/messages_lifecycle.rs b/litellm-rust/crates/core/tests/messages_lifecycle.rs index 9fbdabd71c8..1b40307e729 100644 --- a/litellm-rust/crates/core/tests/messages_lifecycle.rs +++ b/litellm-rust/crates/core/tests/messages_lifecycle.rs @@ -1,7 +1,8 @@ use std::future::Future; use std::pin::Pin; -use std::sync::Mutex; +use std::sync::{Arc, Mutex}; +use futures_util::StreamExt; use litellm_core::Error; use litellm_core::integrations::custom_logger::{LogError, LogFuture}; use litellm_core::lifecycle::{ @@ -105,6 +106,57 @@ async fn upstream(status: u16) -> (String, tokio::task::JoinHandle<()>) { (format!("http://{address}"), server) } +async fn streaming_upstream(body: &'static str) -> (String, tokio::task::JoinHandle<()>) { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.unwrap(); + let mut buffer = [0_u8; 4096]; + let _ = socket.read(&mut buffer).await.unwrap(); + let response = format!( + "HTTP/1.1 200 OK\r\ncontent-type: text/event-stream\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{body}", + body.len() + ); + socket.write_all(response.as_bytes()).await.unwrap(); + }); + (format!("http://{address}"), server) +} + +async fn pending_streaming_upstream() -> (String, tokio::task::JoinHandle<()>) { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.unwrap(); + let mut buffer = [0_u8; 4096]; + let _ = socket.read(&mut buffer).await.unwrap(); + socket + .write_all( + b"HTTP/1.1 200 OK\r\ncontent-type: text/event-stream\r\ntransfer-encoding: chunked\r\n\r\n5\r\ndata:\r\n", + ) + .await + .unwrap(); + tokio::time::sleep(std::time::Duration::from_secs(5)).await; + }); + (format!("http://{address}"), server) +} + +async fn broken_streaming_upstream() -> (String, tokio::task::JoinHandle<()>) { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.unwrap(); + let mut buffer = [0_u8; 4096]; + let _ = socket.read(&mut buffer).await.unwrap(); + socket + .write_all( + b"HTTP/1.1 200 OK\r\ncontent-type: text/event-stream\r\ncontent-length: 100\r\nconnection: close\r\n\r\ndata: short\n\n", + ) + .await + .unwrap(); + }); + (format!("http://{address}"), server) +} + #[tokio::test] async fn success_dispatches_exactly_one_terminal() { let (api_base, server) = upstream(200).await; @@ -158,3 +210,129 @@ async fn pre_call_rejection_never_touches_socket() { .is_err() ); } + +#[tokio::test] +async fn stream_eof_dispatches_usage_exactly_once() { + let events = "event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"usage\":{\"input_tokens\":5,\"output_tokens\":0}}}\n\nevent: message_delta\ndata: {\"type\":\"message_delta\",\"usage\":{\"output_tokens\":4}}\n\nevent: message_stop\ndata: {\"type\":\"message_stop\"}\n\n"; + let (api_base, server) = streaming_upstream(events).await; + let services = Arc::new(Services::default()); + let mut stream_request = request(api_base); + stream_request.body["stream"] = json!(true); + let call = litellm_core::messages::lifecycle::messages_stream( + services.clone(), + stream_request, + Options::default(), + context(), + ) + .await + .expect("stream starts"); + let completion = call.completion.register(); + let bytes = call + .stream + .collect::>() + .await + .into_iter() + .collect::, _>>() + .expect("stream succeeds") + .concat(); + + assert_eq!(bytes, events.as_bytes()); + let terminal = completion.await.expect("completion task succeeds"); + assert_eq!(terminal.classification, TerminalClassification::Success); + assert_eq!(terminal.usage.prompt_tokens, 5); + assert_eq!(terminal.usage.completion_tokens, 4); + assert_eq!(terminal.usage.total_tokens, 9); + assert_eq!(services.terminals.lock().unwrap().len(), 1); + server.await.unwrap(); +} + +#[tokio::test] +async fn dropping_unregistered_completion_still_dispatches_terminal() { + let events = "event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n"; + let (api_base, server) = streaming_upstream(events).await; + let services = Arc::new(Services::default()); + let mut stream_request = request(api_base); + stream_request.body["stream"] = json!(true); + let call = litellm_core::messages::lifecycle::messages_stream( + services.clone(), + stream_request, + Options::default(), + context(), + ) + .await + .expect("stream starts"); + let stream = call.stream; + drop(call.completion); + stream.collect::>().await; + + tokio::time::timeout(std::time::Duration::from_secs(1), async { + loop { + if services.terminals.lock().unwrap().len() == 1 { + break; + } + tokio::task::yield_now().await; + } + }) + .await + .expect("terminal dispatch completes"); + assert_eq!( + services.terminals.lock().unwrap()[0].classification, + TerminalClassification::Success + ); + server.await.unwrap(); +} + +#[tokio::test] +async fn stream_consumer_drop_dispatches_cancellation_exactly_once() { + let (api_base, server) = pending_streaming_upstream().await; + let services = Arc::new(Services::default()); + let mut stream_request = request(api_base); + stream_request.body["stream"] = json!(true); + let call = litellm_core::messages::lifecycle::messages_stream( + services.clone(), + stream_request, + Options::default(), + context(), + ) + .await + .expect("stream starts"); + let completion = call.completion.register(); + let mut stream = call.stream; + assert!(stream.next().await.expect("first chunk exists").is_ok()); + drop(stream); + + let terminal = completion.await.expect("completion task succeeds"); + assert!(matches!( + terminal.classification, + TerminalClassification::Failure { ref kind, .. } if kind == "Cancelled" + )); + assert_eq!(services.terminals.lock().unwrap().len(), 1); + server.abort(); +} + +#[tokio::test] +async fn stream_transport_error_dispatches_failure_exactly_once() { + let (api_base, server) = broken_streaming_upstream().await; + let services = Arc::new(Services::default()); + let mut stream_request = request(api_base); + stream_request.body["stream"] = json!(true); + let call = litellm_core::messages::lifecycle::messages_stream( + services.clone(), + stream_request, + Options::default(), + context(), + ) + .await + .expect("stream starts"); + let completion = call.completion.register(); + let chunks = call.stream.collect::>().await; + + assert!(chunks.iter().any(Result::is_err)); + let terminal = completion.await.expect("completion task succeeds"); + assert!(matches!( + terminal.classification, + TerminalClassification::Failure { ref kind, .. } if kind == "NetworkError" + )); + assert_eq!(services.terminals.lock().unwrap().len(), 1); + server.await.unwrap(); +} diff --git a/litellm-rust/crates/core/tests/ocr.rs b/litellm-rust/crates/core/tests/ocr.rs index 4a6be72964e..df7abfb05f4 100644 --- a/litellm-rust/crates/core/tests/ocr.rs +++ b/litellm-rust/crates/core/tests/ocr.rs @@ -8,9 +8,7 @@ use litellm_core::Error; use litellm_core::lifecycle::CallLifecycleContext; use litellm_core::ocr::prepare::prepare; use litellm_core::ocr::types::{OcrDocument, OcrDocumentProjection}; -use litellm_core::ocr::{ - NoopOcrServices, OcrAdmissionRequest as OcrRequest, PreparedOcr, PreparedOcrCall, -}; +use litellm_core::ocr::{NoopOcrServices, OcrAdmissionRequest as OcrRequest, OcrDraft}; use serde_json::{Value, json}; fn request() -> OcrRequest { @@ -32,7 +30,7 @@ fn request() -> OcrRequest { } } -fn body(prepared: &PreparedOcr) -> Value { +fn body(prepared: &OcrDraft) -> Value { let mut body = prepared.body.clone(); body.insert( "document".into(), @@ -44,20 +42,15 @@ fn body(prepared: &PreparedOcr) -> Value { } async fn ocr( - prepared: PreparedOcr, + prepared: OcrDraft, headers: Vec<(String, String)>, body: Value, ) -> Result { - let model = prepared.model.clone(); - let provider = prepared.custom_llm_provider.clone(); + let model = prepared.endpoint.model().to_string(); + let provider = prepared.endpoint.custom_llm_provider().to_string(); let response = litellm_core::ocr::ocr( &NoopOcrServices, - PreparedOcrCall { - prepared, - headers, - body, - } - .into(), + prepared.endpoint.settle(headers, body), Default::default(), CallLifecycleContext::new("ocr", model, provider, "test-call"), ) @@ -75,10 +68,10 @@ fn prepares_provider_template_auth_and_url() { ..request() }) .unwrap(); - assert_eq!(prepared.model, "mistral-ocr-latest"); - assert_eq!(prepared.custom_llm_provider, "mistral"); - assert_eq!(prepared.url, "https://ocr.example/v1/ocr"); - assert_eq!(prepared.timeout_seconds, 2.0); + assert_eq!(prepared.endpoint.model(), "mistral-ocr-latest"); + assert_eq!(prepared.endpoint.custom_llm_provider(), "mistral"); + assert_eq!(prepared.endpoint.url(), "https://ocr.example/v1/ocr"); + assert_eq!(prepared.endpoint.timeout_seconds(), 2.0); assert_eq!( prepared.document_projection, OcrDocumentProjection::RetainedDocument @@ -106,8 +99,8 @@ fn prepares_provider_template_auth_and_url() { ..request() }) .unwrap(); - assert_eq!(explicit.model, "mistral-ocr-latest"); - assert_eq!(explicit.url, "https://api.mistral.ai/v1/ocr"); + assert_eq!(explicit.endpoint.model(), "mistral-ocr-latest"); + assert_eq!(explicit.endpoint.url(), "https://api.mistral.ai/v1/ocr"); assert_eq!( explicit.headers, vec![("aUtHoRiZaTiOn".into(), "Bearer explicit".into())] @@ -539,13 +532,27 @@ async fn preserves_callback_body_and_header_changes() { ); } +#[tokio::test] +async fn settled_headers_are_the_only_headers_sent() { + let (base, handle) = server(200, "", "{}", Duration::ZERO); + let prepared = prepare(OcrRequest { + api_base: Some(base), + ..request() + }) + .unwrap(); + let body = body(&prepared); + ocr(prepared, Vec::new(), body).await.unwrap(); + let (headers, _) = handle.join().unwrap(); + assert!(!headers.to_ascii_lowercase().contains("authorization:")); +} + #[tokio::test] async fn rejects_unsupported_inputs_before_io() { let listener = TcpListener::bind("127.0.0.1:0").unwrap(); listener.set_nonblocking(true).unwrap(); let base = format!("http://{}", listener.local_addr().unwrap()); - for case in ["file", "local", "provider", "compression"] { - let mut prepared = prepare(OcrRequest { + for case in ["file", "local", "compression"] { + let prepared = prepare(OcrRequest { api_base: Some(base.clone()), ..request() }) @@ -555,7 +562,6 @@ async fn rejects_unsupported_inputs_before_io() { match case { "file" => body["document"] = json!({"type": "file", "file": "private"}), "local" => body["document"]["document_url"] = json!("file:///private.pdf"), - "provider" => prepared.custom_llm_provider = "cohere".into(), "compression" => headers.push(("Content-Encoding".into(), "gzip".into())), _ => unreachable!(), } diff --git a/litellm-rust/crates/core/tests/ocr_lifecycle.rs b/litellm-rust/crates/core/tests/ocr_lifecycle.rs deleted file mode 100644 index 592b59ebf1b..00000000000 --- a/litellm-rust/crates/core/tests/ocr_lifecycle.rs +++ /dev/null @@ -1,666 +0,0 @@ -use std::sync::{Arc, Mutex}; -use std::time::Duration; - -use litellm_core::error::Error; -use litellm_core::integrations::custom_guardrail::{ - CustomGuardrail, GuardrailContext, GuardrailDecision, GuardrailError, GuardrailEventHook, - GuardrailFuture, GuardrailRequest, -}; -use litellm_core::integrations::custom_logger::{ - CallbackTiming, CallbackValue, CustomLogger, LogFuture, ModelCallDetails, -}; -use litellm_core::integrations::types::RequestMetadata; -use litellm_core::lifecycle::CallLifecycleContext; -#[cfg(feature = "observability")] -use litellm_core::observability::FunctionTrace; -use litellm_core::ocr::{DefaultOcrServices, OcrRequest}; -use serde_json::{Map, Value, json}; -use tokio::io::{AsyncReadExt, AsyncWriteExt}; -use tokio::net::{TcpListener, TcpStream}; -#[cfg(feature = "observability")] -use tracing::instrument::WithSubscriber; - -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 ocr(request: OcrRequest<'_>) -> Result { - let services = DefaultOcrServices::new(&request); - let provider = request - .custom_llm_provider - .or_else(|| request.model.split_once('/').map(|(provider, _)| provider)) - .unwrap_or(""); - let metadata = litellm_core::integrations::types::StandardLoggingMetadata { - user_api_key_hash: request.request_metadata.user_api_key_hash.clone(), - user_api_key_user_id: request.request_metadata.user_api_key_user_id.clone(), - user_api_key_team_id: request.request_metadata.user_api_key_team_id.clone(), - ..Default::default() - }; - let context = CallLifecycleContext::new( - "ocr", - request.model, - provider, - request.litellm_call_id.unwrap_or(""), - ) - .with_metadata(metadata); - litellm_core::ocr::ocr(&services, request.into(), Default::default(), context) - .await - .into_result() -} - -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<'_> { - OcrRequest { - model, - document: json!({ - "type": "document_url", - "document_url": "https://example.com/doc.pdf" - }), - api_key: Some("sk-test"), - api_base: None, - custom_llm_provider: None, - extra_headers: None, - optional_params: Map::new(), - timeout: None, - callbacks: Vec::new(), - guardrails: Vec::new(), - request_metadata: RequestMetadata::default(), - litellm_call_id: None, - } -} - -#[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 = base_ocr_request("reducto/parse-v3"); - request.api_base = Some(&api_base); - request.document = json!({ - "type": "document_url", - "document_url": "data:application/pdf;base64,JVBERi0xLjQ=" - }); - request.guardrails = vec![guardrail.clone()]; - - let error = ocr(request).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 = base_ocr_request("reducto/parse-v3"); - request.api_base = Some(&api_base); - request.document = json!({ - "type": "document_url", - "document_url": "data:application/pdf;base64,JVBERi0xLjQ=" - }); - - let error = ocr(request).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, - ])); - #[cfg(feature = "observability")] - let trace = FunctionTrace::default(); - let api_base = format!("http://{addr}"); - let call = ocr(OcrRequest { - model: "mistral-ocr-latest", - document: json!({ - "type": "document_url", - "document_url": "https://example.com/doc.pdf" - }), - api_key: Some("sk-test"), - api_base: Some(&api_base), - custom_llm_provider: Some("mistral"), - extra_headers: None, - optional_params: Map::new(), - timeout: Some(Duration::from_secs(5)), - callbacks: vec![logger.clone()], - guardrails: vec![guardrail.clone()], - request_metadata: RequestMetadata { - user_api_key_user_id: Some("user-1".to_string()), - ..Default::default() - }, - litellm_call_id: Some("ocr-call-1"), - }); - #[cfg(feature = "observability")] - let call = call.with_subscriber(trace.dispatcher()); - let response = call.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, - }] - ); - #[cfg(feature = "observability")] - assert_eq!( - trace - .events() - .iter() - .filter(|event| event.function.ends_with("_callback")) - .map(|event| event.function) - .collect::>(), - vec!["success_callback"] - ); - - 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()); - #[cfg(feature = "observability")] - let trace = FunctionTrace::default(); - let api_base = format!("http://{addr}"); - let call = ocr(OcrRequest { - model: "mistral-ocr-latest", - document: json!({ - "type": "document_url", - "document_url": "https://example.com/doc.pdf" - }), - api_key: Some("sk-test"), - api_base: Some(&api_base), - custom_llm_provider: Some("mistral"), - extra_headers: None, - optional_params: Map::new(), - timeout: Some(Duration::from_secs(5)), - callbacks: vec![logger.clone()], - guardrails: Vec::new(), - request_metadata: RequestMetadata::default(), - litellm_call_id: Some("ocr-call-2"), - }); - #[cfg(feature = "observability")] - let call = call.with_subscriber(trace.dispatcher()); - let err = call.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()), - }] - ); - #[cfg(feature = "observability")] - assert_eq!( - trace - .events() - .iter() - .filter(|event| event.function.ends_with("_callback")) - .map(|event| event.function) - .collect::>(), - vec!["failure_callback"] - ); -} - -#[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" - }), - api_key: Some("sk-test"), - api_base: Some(&format!("http://{addr}")), - custom_llm_provider: Some("mistral"), - extra_headers: None, - optional_params: Map::new(), - timeout: Some(Duration::from_millis(100)), - callbacks: vec![logger.clone()], - guardrails: vec![guardrail.clone()], - request_metadata: RequestMetadata::default(), - litellm_call_id: Some("ocr-call-3"), - }) - .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" - }), - 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)), - callbacks: Vec::new(), - guardrails: Vec::new(), - request_metadata: RequestMetadata::default(), - 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)), - callbacks: Vec::new(), - guardrails: Vec::new(), - request_metadata: RequestMetadata::default(), - 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/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index 8227419bdc9..0b8861961c9 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -16,6 +16,7 @@ extension-module = ["pyo3/extension-module"] panic-test = [] trace-parity = [ "dep:tracing", + "dep:litellm-ai-gateway", "litellm-core/observability", "litellm-ai-gateway/trace-parity", ] @@ -23,7 +24,7 @@ trace-parity = [ [dependencies] tracing = { workspace = true, optional = true } litellm-core = { workspace = true, features = ["bedrock-auth"] } -litellm-ai-gateway = { workspace = true, default-features = false } +litellm-ai-gateway = { workspace = true, default-features = false, optional = true } litellm-python-interop.workspace = true pyo3.workspace = true pyo3-async-runtimes.workspace = true @@ -32,9 +33,7 @@ serde_json.workspace = true [dev-dependencies] criterion = "0.8.2" -futures-util.workspace = true tokio.workspace = true -tokio-tungstenite.workspace = true tracing.workspace = true [[bench]] diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index f4a77fbef2e..a0bc9b79fe9 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -6,61 +6,7 @@ mod function_trace; mod marshal; mod routes; -use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection as RustResponsesWebSocketConnection; use pyo3::prelude::*; -use pyo3::types::PyAny; -use serde_json::Value; - -use crate::errors::core_error_to_pyerr; -use crate::marshal::{marshal_headers, optional_timeout}; - -#[pyclass] -struct ResponsesWebSocketConnection { - inner: RustResponsesWebSocketConnection, -} - -#[pymethods] -impl ResponsesWebSocketConnection { - #[classmethod] - #[pyo3(signature = (url, headers=None, timeout_seconds=None))] - fn connect<'py>( - _cls: &Bound<'py, pyo3::types::PyType>, - py: Python<'py>, - url: String, - #[pyo3(from_py_with = litellm_python_interop::from_py)] headers: Option, - timeout_seconds: Option, - ) -> PyResult> { - let headers = marshal_headers(headers)?; - let timeout = optional_timeout(timeout_seconds)?; - litellm_python_interop::run_async_py(py, async move { - let inner = RustResponsesWebSocketConnection::connect_url(&url, &headers, timeout) - .await - .map_err(core_error_to_pyerr)?; - Ok(ResponsesWebSocketConnection { inner }) - }) - } - - fn send_text<'py>(&self, py: Python<'py>, text: String) -> PyResult> { - let inner = self.inner.clone(); - litellm_python_interop::run_async_py(py, async move { - inner.send_text(text).await.map_err(core_error_to_pyerr) - }) - } - - fn recv_text<'py>(&self, py: Python<'py>) -> PyResult> { - let inner = self.inner.clone(); - litellm_python_interop::run_async_py(py, async move { - inner.recv_text().await.map_err(core_error_to_pyerr) - }) - } - - fn close<'py>(&self, py: Python<'py>) -> PyResult> { - let inner = self.inner.clone(); - litellm_python_interop::run_async_py(py, async move { - inner.close().await.map_err(core_error_to_pyerr) - }) - } -} #[pymodule(gil_used = true)] mod _native { @@ -70,21 +16,12 @@ mod _native { fn init(module: &Bound<'_, PyModule>) -> PyResult<()> { super::errors::register(module)?; super::routes::register(module)?; - module.add_class::()?; super::diagnostics::register(module) } } #[cfg(test)] mod tests { - use std::ffi::CString; - use std::time::Duration; - - use futures_util::{SinkExt, StreamExt}; - use pyo3::types::PyDict; - use tokio::net::TcpListener; - use tokio_tungstenite::{accept_async, tungstenite::Message}; - use super::*; #[test] @@ -102,10 +39,10 @@ mod tests { "atranscription", "messages", "amessages", - "chat_completions_decline", "chat_completions", "achat_completions", - "ResponsesWebSocketConnection", + "chat_completions_decline", + "responses_websocket", ]; let public_names: Vec = module @@ -153,68 +90,4 @@ mod tests { } }); } - - #[test] - fn responses_websocket_connection_round_trips_through_python() { - Python::initialize(); - let runtime = pyo3_async_runtimes::tokio::get_runtime(); - let listener = runtime - .block_on(TcpListener::bind("127.0.0.1:0")) - .expect("listener should bind"); - let address = listener - .local_addr() - .expect("listener should have an address"); - let server = runtime.spawn(async move { - let (stream, _) = listener.accept().await.expect("server should accept"); - let mut socket = accept_async(stream) - .await - .expect("handshake should succeed"); - - let message = socket - .next() - .await - .expect("client should send a frame") - .expect("client frame should be valid"); - assert_eq!(message, Message::Text("from-python".into())); - socket - .send(Message::Text("from-server".into())) - .await - .expect("server should reply"); - assert!(matches!(socket.next().await, Some(Ok(Message::Close(_))))); - }); - - Python::attach(|py| { - let module = pyo3::wrap_pymodule!(_native)(py).into_bound(py); - let locals = PyDict::new(py); - locals - .set_item("native", &module) - .expect("module should enter Python locals"); - locals - .set_item("url", format!("ws://{address}")) - .expect("URL should enter Python locals"); - let code = CString::new( - r#" -import asyncio - -async def exercise(): - connection = await native.ResponsesWebSocketConnection.connect(url) - assert type(connection) is native.ResponsesWebSocketConnection - await connection.send_text("from-python") - assert await connection.recv_text() == "from-server" - await connection.close() - assert await connection.recv_text() is None - -asyncio.run(asyncio.wait_for(exercise(), timeout=5)) -"#, - ) - .expect("Python source should not contain null bytes"); - py.run(&code, Some(&locals), Some(&locals)) - .expect("Python WebSocket methods should round trip"); - }); - - runtime - .block_on(async { tokio::time::timeout(Duration::from_secs(5), server).await }) - .expect("server should finish") - .expect("server task should not panic"); - } } diff --git a/litellm-rust/crates/python-bridge/src/marshal.rs b/litellm-rust/crates/python-bridge/src/marshal.rs index cdb4319b146..e5a2bec1ae2 100644 --- a/litellm-rust/crates/python-bridge/src/marshal.rs +++ b/litellm-rust/crates/python-bridge/src/marshal.rs @@ -1,4 +1,3 @@ -use std::collections::HashMap; use std::time::Duration; use pyo3::exceptions::{PyTypeError, PyValueError}; @@ -36,20 +35,6 @@ impl RouteOptions { } } -pub(crate) fn required_value( - name: &'static str, - value: Value, - expected: fn(&Value) -> bool, - expected_name: &'static str, -) -> PyResult { - if expected(&value) { - return Ok(value); - } - Err(PyTypeError::new_err(format!( - "{name} must be a {expected_name}" - ))) -} - pub(crate) fn object_or_empty( name: &'static str, value: Option, @@ -88,25 +73,6 @@ pub(crate) fn optional_timeout(timeout_seconds: Option) -> PyResult) -> PyResult> { - let value = match headers { - Some(headers) => headers, - None => Value::Object(Map::new()), - }; - let Value::Object(headers) = value else { - return Err(PyTypeError::new_err("headers must be a dict")); - }; - headers - .into_iter() - .map(|(name, value)| { - value - .as_str() - .map(|value| (name, value.to_string())) - .ok_or_else(|| PyValueError::new_err("header values must be strings")) - }) - .collect() -} - #[cfg(test)] mod tests { use super::*; diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs index 08ab476005c..1867d407439 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs @@ -1,91 +1,476 @@ use litellm_core::Error; -use std::future::Future; - -use litellm_core::chat_completions::types::{ChatCompletionsRequest, ChatCompletionsResponse}; -use litellm_core::chat_completions::{ - chat_completions as run_chat_completions, chat_completions_decline_reason, +use litellm_core::chat_completions::lifecycle::{ + Admission, ChatCompletionsRoute, Observations, Operation, Options, machine, }; +use litellm_core::chat_completions::types::ChatCompletionsRequest; +use litellm_core::chat_completions::{ + chat_completions_decline_reason, chat_completions_with_terminal, +}; +use litellm_core::lifecycle::{ + CallLifecycleContext, ErrorDisposition, ExecutedCall, Lifecycle, Outcome, TerminalRecord, +}; +use litellm_python_interop::{Pythonized, from_py, run_async_value, run_sync_value, to_py}; +use pyo3::exceptions::{PyRuntimeError, PyTypeError, PyValueError}; use pyo3::prelude::*; -use serde_json::Value; +use pyo3::pyclass::{PyTraverseError, PyVisit}; +use pyo3::sync::PyOnceLock; +use pyo3::types::PyDict; +use serde_json::{Map, Value}; -use crate::errors::chat_completions_error_to_pyerr; -use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty, required_value}; +use crate::errors::{RustBridgeDeclined, chat_completions_error_to_pyerr, core_error_to_pyerr}; +use crate::marshal::optional_timeout; -fn prepare_chat_completions( - inputs: ChatCompletionsInputs, -) -> PyResult> + Send + 'static> { - let messages = required_value("messages", inputs.messages, Value::is_array, "list")?; - let optional_params = object_or_empty("optional_params", inputs.optional_params)?; - let options = RouteOptions::from_python(RouteOptionsInputs { - model: inputs.model, - api_key: inputs.api_key, - api_base: inputs.api_base, - custom_llm_provider: inputs.custom_llm_provider, - extra_headers: inputs.extra_headers, - timeout_seconds: inputs.timeout_seconds, - })?; +#[pyclass] +struct ChatCompletionsState { + arguments: Option>, + model: Option, + messages: Option, + optional_params: Option>, + api_key: Option, + api_base: Option, + custom_llm_provider: Option, + extra_headers: Option>, + timeout: Option, + terminal: Option, +} - Ok(async move { - let RouteOptions { - model, +#[pymethods] +impl ChatCompletionsState { + fn __traverse__(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> { + visit.call(&self.arguments) + } + + fn __clear__(slf: &Bound<'_, Self>) { + let roots = { + let mut state = slf.borrow_mut(); + (state.arguments.take(), state.terminal.take()) + }; + drop(roots); + } +} + +fn scalar(arguments: &Bound<'_, PyDict>, name: &str) -> PyResult> { + arguments + .get_item(name)? + .filter(|value| !value.is_none()) + .map(|value| value.extract::()) + .transpose() +} + +fn value(arguments: &Bound<'_, PyDict>, name: &str) -> PyResult { + arguments + .get_item(name)? + .ok_or_else(|| PyValueError::new_err(format!("chat completions requires {name}"))) + .and_then(|value| from_py(&value)) +} + +fn optional_map(arguments: &Bound<'_, PyDict>, name: &str) -> PyResult>> { + arguments + .get_item(name)? + .filter(|value| !value.is_none()) + .map(|value| from_py(&value)) + .transpose() +} + +fn admission(arguments: &Bound<'_, PyDict>) -> PyResult { + Ok(Admission { + model: scalar(arguments, "model")? + .ok_or_else(|| PyValueError::new_err("chat completions requires model"))?, + messages: value(arguments, "messages")?, + optional_params: optional_map(arguments, "optional_params")?.unwrap_or_default(), + custom_llm_provider: scalar(arguments, "custom_llm_provider")?, + }) +} + +#[pyclass] +struct ChatCompletionsLifecycle { + machine: Lifecycle, + asynchronous: bool, +} + +#[pymethods] +impl ChatCompletionsLifecycle { + #[new] + fn new( + arguments: &Bound<'_, PyDict>, + asynchronous: bool, + internal_call: bool, + ) -> PyResult { + match machine( + &admission(arguments)?, + Options { + asynchronous, + internal_call, + }, + ) + .map_err(core_error_to_pyerr)? + { + Ok(machine) => Ok(Self { + machine, + asynchronous, + }), + Err(decline) => Err(RustBridgeDeclined::new_err(decline.reason())), + } + } + + fn advance( + &mut self, + outcome: u8, + logger_available: bool, + has_fallbacks: bool, + ) -> PyResult { + let outcome = match outcome { + 0 => Outcome::Success, + 1 => Outcome::Failure, + _ => Outcome::Abort, + }; + self.machine + .advance( + outcome, + Observations { + logger_available, + has_fallbacks, + }, + ) + .map(|transition| transition.error == ErrorDisposition::Replace) + .map_err(core_error_to_pyerr) + } + + fn complete(&self) -> Option { + match self.machine.operation() { + Operation::Complete(outcome) => Some(outcome == Outcome::Success), + _ => None, + } + } +} + +#[pyfunction] +fn invoke( + py: Python<'_>, + machine: Py, + host: Py, +) -> PyResult<(bool, Py)> { + let (operation, asynchronous) = { + let machine = machine.borrow(py); + (machine.machine.operation(), machine.asynchronous) + }; + let (method, awaiting) = match operation { + Operation::Setup => ("setup", false), + Operation::DeploymentPre => ("deployment_pre", true), + Operation::Prepare => ("prepare", false), + Operation::Send if asynchronous => ("send", true), + Operation::Send => ("send_sync", false), + Operation::DeploymentSuccess => ("deployment_success", true), + Operation::DeploymentFailure => ("deployment_failure", true), + Operation::SyncSuccess => ("sync_success", false), + Operation::AsyncSuccess => ("async_success", false), + Operation::SyncSuccessIfNeeded => ("sync_success_if_needed", false), + Operation::SyncFailure => ("sync_failure", false), + Operation::AsyncFailure => ("async_failure", true), + Operation::Restore => ("restore", false), + Operation::Complete(_) => { + return Err(PyRuntimeError::new_err( + "chat completions lifecycle is complete", + )); + } + }; + Ok((awaiting, host.getattr(py, method)?.call0(py)?)) +} + +#[pyfunction] +fn prepare(py: Python<'_>, arguments: Py) -> PyResult> { + let bag = arguments.bind(py); + let admission = admission(bag)?; + let api_key = scalar(bag, "api_key")?; + let api_base = scalar(bag, "api_base")?; + let extra_headers = optional_map(bag, "extra_headers")?; + let timeout = optional_timeout( + bag.get_item("timeout_seconds")? + .filter(|value| !value.is_none()) + .map(|value| value.extract::()) + .transpose()?, + )?; + let logging = bag + .get_item("litellm_logging_obj")? + .filter(|value| !value.is_none()) + .ok_or_else(|| PyRuntimeError::new_err("chat completions logging was not initialized"))?; + let complete_input = PyDict::new(py); + complete_input.set_item("model", &admission.model)?; + complete_input.set_item("messages", bag.get_item("messages")?)?; + for (name, value) in &admission.optional_params { + complete_input.set_item(name, Pythonized(value))?; + } + let additional = PyDict::new(py); + additional.set_item("complete_input_dict", complete_input)?; + additional.set_item("api_base", bag.get_item("api_base")?)?; + additional.set_item("headers", bag.get_item("extra_headers")?)?; + let kwargs = PyDict::new(py); + kwargs.set_item("input", bag.get_item("messages")?)?; + kwargs.set_item("api_key", bag.get_item("logging_api_key")?)?; + kwargs.set_item("additional_args", additional)?; + logging.call_method("pre_call", (), Some(&kwargs))?; + Py::new( + py, + ChatCompletionsState { + arguments: Some(arguments), + model: Some(admission.model), + messages: Some(admission.messages), + optional_params: Some(admission.optional_params), api_key, api_base, - custom_llm_provider, + custom_llm_provider: admission.custom_llm_provider, extra_headers, timeout, - } = options; - run_chat_completions(ChatCompletionsRequest { - model: &model, - messages, - optional_params, - api_key: api_key.as_deref(), - api_base: api_base.as_deref(), - custom_llm_provider: custom_llm_provider.as_deref(), - extra_headers, - timeout, - }) - .await + terminal: None, + }, + ) +} + +struct OwnedRequest { + model: String, + messages: Value, + optional_params: Map, + api_key: Option, + api_base: Option, + custom_llm_provider: Option, + extra_headers: Option>, + timeout: Option, + call_id: String, +} + +fn take_request(py: Python<'_>, state: &Py) -> PyResult { + let mut state = state.borrow_mut(py); + let call_id = state + .arguments + .as_ref() + .ok_or_else(|| PyRuntimeError::new_err("chat completions state was cleared")) + .and_then(|arguments| scalar(arguments.bind(py), "litellm_call_id"))? + .unwrap_or_default(); + Ok(OwnedRequest { + model: state + .model + .take() + .ok_or_else(|| PyRuntimeError::new_err("chat completions request was already sent"))?, + messages: state.messages.take().unwrap(), + optional_params: state.optional_params.take().unwrap(), + api_key: state.api_key.take(), + api_base: state.api_base.take(), + custom_llm_provider: state.custom_llm_provider.take(), + extra_headers: state.extra_headers.take(), + timeout: state.timeout.take(), + call_id, }) } +async fn execute( + request: OwnedRequest, +) -> ExecutedCall { + let provider = request.custom_llm_provider.clone().unwrap_or_default(); + let context = + CallLifecycleContext::new("chat_completion", &request.model, provider, request.call_id); + chat_completions_with_terminal( + ChatCompletionsRequest { + model: &request.model, + messages: request.messages, + optional_params: request.optional_params, + api_key: request.api_key.as_deref(), + api_base: request.api_base.as_deref(), + custom_llm_provider: request.custom_llm_provider.as_deref(), + extra_headers: request.extra_headers, + timeout: request.timeout, + }, + context, + ) + .await +} + +fn store_result( + py: Python<'_>, + state: &Py, + executed: ExecutedCall, +) -> PyResult> { + state.borrow_mut(py).terminal = Some(executed.terminal().clone()); + match executed { + ExecutedCall::Success { response, .. } => { + Ok(Pythonized(response).into_pyobject(py)?.unbind().into_any()) + } + ExecutedCall::Failure { error, .. } => Err(chat_completions_error_to_pyerr(error)), + } +} + +#[pyfunction] +fn send(py: Python<'_>, state: Py) -> PyResult> { + let request = take_request(py, &state)?; + litellm_python_interop::run_async_py(py, async move { + let executed = run_async_value( + async move { Ok::<_, std::convert::Infallible>(execute(request).await) }, + |never| match never {}, + ) + .await?; + Python::attach(|py| store_result(py, &state, executed)) + }) +} + +#[pyfunction] +fn send_sync(py: Python<'_>, state: Py) -> PyResult> { + let request = take_request(py, &state)?; + let executed = run_sync_value( + py, + async move { Ok::<_, std::convert::Infallible>(execute(request).await) }, + |never| match never {}, + )?; + store_result(py, &state, executed) +} + +#[pyfunction] +fn terminal_record(py: Python<'_>, state: Py) -> PyResult> { + let terminal = state.borrow(py).terminal.clone().ok_or_else(|| { + PyRuntimeError::new_err("chat completions terminal record is unavailable") + })?; + to_py(py, &terminal) +} + +fn validate_arguments(arguments: &Bound<'_, PyDict>) -> PyResult<()> { + let messages = arguments + .get_item("messages")? + .ok_or_else(|| PyValueError::new_err("chat completions requires messages"))?; + if !messages.is_instance_of::() { + return Err(PyTypeError::new_err("messages must be a list")); + } + Ok(()) +} + +#[pyfunction] +fn chat_completions(py: Python<'_>, arguments: Py) -> PyResult> { + validate_arguments(arguments.bind(py))?; + driver(py)?.getattr("drive_sync")?.call1((arguments,)) +} + +#[pyfunction] +fn achat_completions(py: Python<'_>, arguments: Py) -> PyResult> { + validate_arguments(arguments.bind(py))?; + driver(py)?.getattr("drive_async")?.call1((arguments,)) +} + #[pyfunction] #[pyo3(signature = (model, messages, optional_params=None, custom_llm_provider=None))] fn chat_completions_decline( model: String, #[pyo3(from_py_with = litellm_python_interop::from_py)] messages: Value, - #[pyo3(from_py_with = litellm_python_interop::from_py)] optional_params: Option, + #[pyo3(from_py_with = litellm_python_interop::from_py)] optional_params: Option< + Map, + >, custom_llm_provider: Option, -) -> PyResult> { - let optional_params = object_or_empty("optional_params", optional_params)?; - Ok(chat_completions_decline_reason( +) -> Option { + chat_completions_decline_reason( &model, custom_llm_provider.as_deref(), messages, - &optional_params, + &optional_params.unwrap_or_default(), ) - .map(str::to_string)) + .map(str::to_string) } -bridge_route! { - sync = chat_completions, - asynchronous = achat_completions, - inputs = ChatCompletionsInputs, - required = { - model: String, - #[pyo3(from_py_with = litellm_python_interop::from_py)] - messages: serde_json::Value, - }, - optional = { - #[pyo3(from_py_with = litellm_python_interop::from_py)] - optional_params: Option, - api_key: Option, - api_base: Option, - custom_llm_provider: Option, - #[pyo3(from_py_with = litellm_python_interop::from_py)] - extra_headers: Option, - timeout_seconds: Option, - }, - prepare = prepare_chat_completions, - errors = chat_completions_error_to_pyerr, - extra = [chat_completions_decline], +fn driver(py: Python<'_>) -> PyResult<&Bound<'_, PyModule>> { + static DRIVER: PyOnceLock> = PyOnceLock::new(); + if let Some(module) = DRIVER.get(py) { + return Ok(module.bind(py)); + } + let module = crate::driver::compile(py, "chat_completions", HOST)?; + module.add("_Lifecycle", py.get_type::())?; + module.add("_invoke", wrap_pyfunction!(invoke, &module)?)?; + module.add("_prepare", wrap_pyfunction!(prepare, &module)?)?; + module.add("_send", wrap_pyfunction!(send, &module)?)?; + module.add("_send_sync", wrap_pyfunction!(send_sync, &module)?)?; + module.add( + "_terminal_record", + wrap_pyfunction!(terminal_record, &module)?, + )?; + Ok(DRIVER.get_or_init(py, || module.unbind()).bind(py)) +} + +const HOST: &str = r#" +from datetime import datetime +from litellm import utils +from litellm.types.utils import CallTypes +from litellm.rust_bridge.chat_completions import build_model_response, initialize_logging, invoke_terminal + +class Host: + def __init__(self, arguments, asynchronous): + self.machine = _Lifecycle(arguments, asynchronous, utils.is_internal_call.get()) + self.arguments = arguments + self.current = arguments + self.asynchronous = asynchronous + self.logger = arguments.get('litellm_logging_obj') + self.state = None + self.response = None + self.error = None + self.start = datetime.now() + self.end = None + + def setup(self): + self.logger = initialize_logging(self.arguments, self.asynchronous) + self.arguments['litellm_logging_obj'] = self.logger + + async def deployment_pre(self): + modified = await utils.async_pre_call_deployment_hook(self.current, 'acompletion') + if modified is not None: + self.current = modified + self.current['litellm_logging_obj'] = self.logger + + def prepare(self): self.state = _prepare(self.current) + + def send_sync(self): + self.response = build_model_response(_send_sync(self.state), self.arguments['model_response']) + self.end = datetime.now() + + async def send(self): + self.response = build_model_response(await _send(self.state), self.arguments['model_response']) + self.end = datetime.now() + + async def deployment_success(self): + self.response = await utils.async_post_call_success_deployment_hook(self.current, self.response, CallTypes.acompletion) + + async def deployment_failure(self): + await utils.async_post_call_failure_deployment_hook(self.current, self.error, 'acompletion') + + def terminal(self, action, value): + record = _terminal_record(self.state) if self.state is not None else None + return invoke_terminal(action, (self.arguments, self.current, self.state), self.logger, record, value, self.start, self.end) + + def sync_success(self): return self.terminal('sync_success', self.response) + def async_success(self): return self.terminal('async_success', self.response) + def sync_success_if_needed(self): return self.terminal('sync_success_if_needed', self.response) + def sync_failure(self): return self.terminal('sync_failure', self.error) + def async_failure(self): return self.terminal('async_failure', self.error) + def restore(self): utils._restore_correlation_context_if_supported(self.logger) + + def advance(self, outcome, error=None): + if error is not None and self.end is None: + self.end = datetime.now() + if self.logger is None: + self.logger = self.arguments.get('litellm_logging_obj') + replace = self.machine.advance(outcome, self.logger is not None, self.current.get('fallbacks') is not None) + if replace: + self.error = error + + def result(self): + if self.machine.complete(): + return self.response + raise self.error +"#; + +pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { + crate::routes::definition::add_function(module, wrap_pyfunction!(chat_completions, module)?)?; + crate::routes::definition::add_function(module, wrap_pyfunction!(achat_completions, module)?)?; + crate::routes::definition::add_function( + module, + wrap_pyfunction!(chat_completions_decline, module)?, + )?; + Ok(()) +} + +#[cfg(feature = "trace-parity")] +pub(super) fn register_trace(module: &Bound<'_, PyModule>) -> PyResult<()> { + register(module) } diff --git a/litellm-rust/crates/python-bridge/src/routes/mod.rs b/litellm-rust/crates/python-bridge/src/routes/mod.rs index 7e81f2ffe9b..77eaa2bda91 100644 --- a/litellm-rust/crates/python-bridge/src/routes/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/mod.rs @@ -10,12 +10,14 @@ mod audio_transcription; mod chat_completions; mod messages; mod ocr; +mod responses_websocket; pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { ocr::register(module)?; audio_transcription::register(module)?; messages::register(module)?; chat_completions::register(module)?; + responses_websocket::register(module)?; #[cfg(feature = "trace-parity")] { let trace = PyModule::new(module.py(), "_trace")?; diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr.rs b/litellm-rust/crates/python-bridge/src/routes/ocr.rs index 4db0cc70d8e..d08ac23a289 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr.rs @@ -9,7 +9,7 @@ use litellm_core::lifecycle::{ }; use litellm_core::ocr::NoopOcrServices; use litellm_core::ocr::types::{ - OcrAdmissionRequest, OcrDocumentProjection, PreparedOcr, PreparedOcrCall, + OcrAdmissionRequest, OcrDocumentProjection, OcrDraft, OcrEndpoint, SettledOcrRequest, }; use litellm_core::routing_utils::provider::get_custom_llm_provider; use litellm_python_interop::{Pythonized, from_py, to_py}; @@ -18,7 +18,6 @@ use pyo3::prelude::*; use pyo3::pyclass::{PyTraverseError, PyVisit}; use pyo3::sync::PyOnceLock; use pyo3::types::PyDict; -use serde_json::Value; use crate::errors::core_error_to_pyerr; use litellm_python_interop::{run_async_value, run_sync_value}; @@ -29,7 +28,9 @@ struct OcrState { body: Option>, headers: Option>, logging: Option>, - prepared: Option, + pre_call: Option>, + endpoint: Option, + asynchronous: bool, terminal: Option, } @@ -39,7 +40,8 @@ impl OcrState { visit.call(&self.arguments)?; visit.call(&self.body)?; visit.call(&self.headers)?; - visit.call(&self.logging) + visit.call(&self.logging)?; + visit.call(&self.pre_call) } fn __clear__(slf: &Bound<'_, Self>) { @@ -50,6 +52,8 @@ impl OcrState { state.body.take(), state.headers.take(), state.logging.take(), + state.pre_call.take(), + state.endpoint.take(), state.terminal.take(), ) }; @@ -313,6 +317,7 @@ fn invoke( Operation::Setup => ("setup", false), Operation::DeploymentPre => ("deployment_pre", true), Operation::Prepare => ("prepare", false), + Operation::PreCall => ("pre_call", false), Operation::Send if asynchronous => ("send", true), Operation::Send => ("send_sync", false), Operation::DeploymentSuccess => ("deployment_success", true), @@ -336,7 +341,7 @@ fn prepare(py: Python<'_>, arguments: Py, asynchronous: bool) -> PyResul let request = decode_request(py, bag)?; let model = request.model.clone(); let custom_llm_provider = request.custom_llm_provider.clone(); - let prepared = py + let draft = py .detach(|| litellm_core::ocr::prepare::prepare(request)) .map_err(|error| { request_error_to_pyerr(py, error, &model, custom_llm_provider.as_deref()) @@ -345,10 +350,17 @@ fn prepare(py: Python<'_>, arguments: Py, asynchronous: bool) -> PyResul .get_item("document")? .ok_or_else(|| PyValueError::new_err("OCR requires document"))? .cast_into::()?; - let body = to_py(py, &prepared.body)? + let OcrDraft { + endpoint, + headers: draft_headers, + body: draft_body, + document_projection, + parameter_fields, + } = draft; + let body = to_py(py, &draft_body)? .into_bound(py) .cast_into::()?; - match prepared.document_projection { + match document_projection { OcrDocumentProjection::RetainedDocument => { body.set_item(pyo3::intern!(py, "document"), &document)? } @@ -358,14 +370,14 @@ fn prepare(py: Python<'_>, arguments: Py, asynchronous: bool) -> PyResul OcrDocumentProjection::Transformed => {} } let optional_params = PyDict::new(py); - for &name in prepared.parameter_fields { + for &name in parameter_fields { if let Some(value) = bag.get_item(name)? { body.set_item(name, &value)?; optional_params.set_item(name, value)?; } } let headers = PyDict::new(py); - for (name, value) in &prepared.headers { + for (name, value) in &draft_headers { headers.set_item(name, value)?; } let logging = py @@ -380,22 +392,20 @@ fn prepare(py: Python<'_>, arguments: Py, asynchronous: bool) -> PyResul )?; let update = PyDict::new(py); update.set_item("kwargs", bag)?; - update.set_item("model", &prepared.model)?; + update.set_item("model", endpoint.model())?; update.set_item("optional_params", optional_params)?; update.set_item("litellm_params", litellm_params)?; - update.set_item("custom_llm_provider", &prepared.custom_llm_provider)?; + update.set_item("custom_llm_provider", endpoint.custom_llm_provider())?; logging.call_method("update_from_kwargs", (), Some(&update))?; let additional_args = PyDict::new(py); additional_args.set_item("complete_input_dict", &body)?; - additional_args.set_item(pyo3::intern!(py, "api_base"), &prepared.url)?; + additional_args.set_item(pyo3::intern!(py, "api_base"), endpoint.url())?; additional_args.set_item(pyo3::intern!(py, "headers"), &headers)?; let pre_call = PyDict::new(py); pre_call.set_item("input", "OCR document processing")?; pre_call.set_item("api_key", bag.get_item("api_key")?)?; pre_call.set_item("additional_args", additional_args)?; - logging.call_method(pyo3::intern!(py, "pre_call"), (), Some(&pre_call))?; - let logging = logging.unbind(); Py::new( py, @@ -404,19 +414,43 @@ fn prepare(py: Python<'_>, arguments: Py, asynchronous: bool) -> PyResul body: Some(body.unbind()), headers: Some(headers.unbind()), logging: Some(logging), - prepared: Some(prepared), + pre_call: Some(pre_call.unbind()), + endpoint: Some(endpoint), + asynchronous, terminal: None, }, ) } -type OcrWireRequest = (PreparedOcr, Vec<(String, String)>, Value); +#[pyfunction] +fn pre_call(py: Python<'_>, state: Py) -> PyResult<()> { + let (logging, arguments) = { + let state = state.borrow(py); + let logging = state + .logging + .as_ref() + .ok_or_else(|| PyRuntimeError::new_err("OCR logging state was cleared"))? + .clone_ref(py); + let arguments = state + .pre_call + .as_ref() + .ok_or_else(|| PyRuntimeError::new_err("OCR pre-call state was cleared"))? + .clone_ref(py); + (logging, arguments) + }; + logging + .bind(py) + .call_method(pyo3::intern!(py, "pre_call"), (), Some(arguments.bind(py)))?; + Ok(()) +} + +type OcrWireRequest = (SettledOcrRequest, String, String, bool); fn request(py: Python<'_>, state: &Py) -> PyResult { - let (prepared, body, headers) = { + let (endpoint, body, headers, asynchronous) = { let mut state = state.borrow_mut(py); - let prepared = state - .prepared + let endpoint = state + .endpoint .take() .ok_or_else(|| PyRuntimeError::new_err("OCR request was already sent or cleared"))?; let body = state @@ -429,21 +463,21 @@ fn request(py: Python<'_>, state: &Py) -> PyResult { .as_ref() .ok_or_else(|| PyRuntimeError::new_err("OCR headers were cleared"))? .clone_ref(py); - (prepared, body, headers) + (endpoint, body, headers, state.asynchronous) }; - Ok(( - prepared, + let model = endpoint.model().to_string(); + let provider = endpoint.custom_llm_provider().to_string(); + let request = endpoint.settle( header_pairs(headers.bind(py))?, from_py(body.bind(py).as_any())?, - )) + ); + Ok((request, model, provider, asynchronous)) } #[pyfunction] fn send(py: Python<'_>, state: Py) -> PyResult> { - let (prepared, headers, body) = request(py, &state)?; + let (request, model, provider, asynchronous) = request(py, &state)?; litellm_python_interop::run_async_py(py, async move { - let model = prepared.model.clone(); - let provider = prepared.custom_llm_provider.clone(); let error_model = model.clone(); let error_provider = provider.clone(); let call_id = Python::attach(|py| { @@ -460,13 +494,11 @@ fn send(py: Python<'_>, state: Py) -> PyResult> { Ok::<_, std::convert::Infallible>( litellm_core::ocr::ocr( &services, - PreparedOcrCall { - prepared, - headers, - body, - } - .into(), - Options::default(), + request, + Options { + asynchronous, + ..Options::default() + }, CallLifecycleContext::new("ocr", &model, &provider, call_id), ) .await, @@ -508,9 +540,7 @@ fn finish(py: Python<'_>, response: Py) -> PyResult> { #[pyfunction] fn send_sync(py: Python<'_>, state: Py) -> PyResult> { - let (prepared, headers, body) = request(py, &state)?; - let model = prepared.model.clone(); - let provider = prepared.custom_llm_provider.clone(); + let (request, model, provider, asynchronous) = request(py, &state)?; let error_model = model.clone(); let error_provider = provider.clone(); let call_id = state @@ -526,13 +556,11 @@ fn send_sync(py: Python<'_>, state: Py) -> PyResult> { Ok::<_, std::convert::Infallible>( litellm_core::ocr::ocr( &services, - PreparedOcrCall { - prepared, - headers, - body, - } - .into(), - Options::default(), + request, + Options { + asynchronous, + ..Options::default() + }, CallLifecycleContext::new("ocr", &model, &provider, call_id), ) .await, @@ -619,6 +647,9 @@ class Host: def prepare(self): self.state = _prepare(self.current, self.asynchronous) + def pre_call(self): + _pre_call(self.state) + def send_sync(self): self.response = _send_sync(self.state) self.end = datetime.now() @@ -674,6 +705,7 @@ class Host: module.add("_Lifecycle", py.get_type::())?; module.add("_invoke", wrap_pyfunction!(invoke, &module)?)?; module.add("_prepare", wrap_pyfunction!(prepare, &module)?)?; + module.add("_pre_call", wrap_pyfunction!(pre_call, &module)?)?; module.add("_send", wrap_pyfunction!(send, &module)?)?; module.add("_send_sync", wrap_pyfunction!(send_sync, &module)?)?; module.add("_finish", wrap_pyfunction!(finish, &module)?)?; @@ -826,7 +858,9 @@ mod tests { body: None, headers: None, logging: None, - prepared: None, + pre_call: None, + endpoint: None, + asynchronous: false, terminal: Some(TerminalRecord { call_id: "call-1".into(), trace_id: None, @@ -957,6 +991,9 @@ asyncio.run(exercise()) module .add_function(wrap_pyfunction!(prepare, &module).unwrap()) .unwrap(); + module + .add_function(wrap_pyfunction!(pre_call, &module).unwrap()) + .unwrap(); module .add_function(wrap_pyfunction!(send, &module).unwrap()) .unwrap(); @@ -1029,7 +1066,17 @@ asyncio.run(exercise()) #[pyfunction] fn snapshot(py: Python<'_>, state: Py) -> PyResult> { - let (_, headers, body) = request(py, &state)?; + let state = state.borrow(py); + let headers = state + .headers + .as_ref() + .ok_or_else(|| PyRuntimeError::new_err("OCR headers were cleared"))?; + let body = state + .body + .as_ref() + .ok_or_else(|| PyRuntimeError::new_err("OCR body was cleared"))?; + let headers = header_pairs(headers.bind(py))?; + let body: serde_json::Value = from_py(body.bind(py).as_any())?; to_py(py, &(headers, body)) } @@ -1041,6 +1088,9 @@ asyncio.run(exercise()) module .add_function(wrap_pyfunction!(prepare, &module).unwrap()) .unwrap(); + module + .add_function(wrap_pyfunction!(pre_call, &module).unwrap()) + .unwrap(); module .add_function(wrap_pyfunction!(snapshot, &module).unwrap()) .unwrap(); @@ -1087,6 +1137,7 @@ arguments = dict(model='mistral/mistral-ocr-latest', document=document, api_key='test-key', pages=pages, metadata=metadata, opaque=opaque, litellm_logging_obj=logger, timeout=Timeout()) state = native.prepare(arguments) +native.pre_call(state) assert logger.calls == ['update', 'pre'] roots = gc.get_referents(state) assert any(root is arguments for root in roots) diff --git a/litellm-rust/crates/python-bridge/src/routes/responses_websocket.rs b/litellm-rust/crates/python-bridge/src/routes/responses_websocket.rs new file mode 100644 index 00000000000..dec94a661b9 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/responses_websocket.rs @@ -0,0 +1,14 @@ +use pyo3::prelude::*; + +use crate::errors::RustBridgeDeclined; + +#[pyfunction] +fn responses_websocket() -> PyResult<()> { + Err(RustBridgeDeclined::new_err( + "Responses WebSocket requires host per-frame guardrails and logging", + )) +} + +pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { + module.add_function(wrap_pyfunction!(responses_websocket, module)?) +} diff --git a/litellm/llms/anthropic/chat/handler.py b/litellm/llms/anthropic/chat/handler.py index c82be07a5c5..eba057f8393 100644 --- a/litellm/llms/anthropic/chat/handler.py +++ b/litellm/llms/anthropic/chat/handler.py @@ -407,8 +407,8 @@ class AnthropicChatCompletion(BaseLLM): stream=stream, ) if serves_via_rust: - rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict - "complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent + rust_logging_args: Final = { + "complete_input_dict": { "model": model, "messages": messages, **rust_optional_params, @@ -426,9 +426,6 @@ class AnthropicChatCompletion(BaseLLM): if acompletion is True: async def python_fallback() -> "ModelResponse | CustomStreamWrapper": - # pre_call already fired for this request above. The Rust - # path only declines before the provider is called, so this - # is the same attempt continuing, not a second one. fallback_headers, fallback_data = build_request() return await self.acompletion_function( model=model, @@ -463,6 +460,7 @@ class AnthropicChatCompletion(BaseLLM): custom_llm_provider=custom_llm_provider, extra_headers=headers, timeout=timeout, + arguments={**litellm_params, "litellm_logging_obj": logging_obj}, on_response=log_rust_post_call, python_fallback=python_fallback, ) @@ -476,6 +474,7 @@ class AnthropicChatCompletion(BaseLLM): custom_llm_provider=custom_llm_provider, extra_headers=headers, timeout=timeout, + arguments={**litellm_params, "litellm_logging_obj": logging_obj}, on_response=log_rust_post_call, ) if rust_response is not None: @@ -484,9 +483,6 @@ class AnthropicChatCompletion(BaseLLM): headers, data = build_request() ## LOGGING - # Reaching here with `serves_via_rust` set means the Rust attempt - # declined at call time, before the provider was called, and already - # logged this request. That is the same attempt continuing. if not serves_via_rust: logging_obj.pre_call( input=messages, diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index 984ba371898..9162115f85f 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -417,8 +417,8 @@ class BedrockConverseLLM(BaseAWSLLM): stream=stream, ) if serves_via_rust: - rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict - "complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent + rust_logging_args: Final = { + "complete_input_dict": { "messages": messages, **optional_params, }, @@ -443,6 +443,8 @@ class BedrockConverseLLM(BaseAWSLLM): custom_llm_provider="bedrock", extra_headers=headers, timeout=timeout, + arguments={**litellm_params, "litellm_logging_obj": logging_obj}, + logging_api_key="", on_response=log_rust_post_call, python_fallback=lambda: self.async_completion( model=model, @@ -473,6 +475,8 @@ class BedrockConverseLLM(BaseAWSLLM): custom_llm_provider="bedrock", extra_headers=headers, timeout=timeout, + arguments={**litellm_params, "litellm_logging_obj": logging_obj}, + logging_api_key="", on_response=log_rust_post_call, ) if rust_response is not None: @@ -544,11 +548,6 @@ class BedrockConverseLLM(BaseAWSLLM): ) ## LOGGING - # Reaching here with `serves_via_rust` set means the synchronous Rust - # attempt declined at call time, before the provider was called, and - # already logged this request. That is the same attempt continuing. - # The asynchronous branch above returns before this point, and hands - # its own fallback `skip_pre_call_logging=True` for the same reason. if not serves_via_rust: logging_obj.pre_call( input=messages, diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 9c4dcf7afdf..3844be6fe36 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -6543,20 +6543,13 @@ class BaseLLMHTTPHandler: }, ) + if _rust_responses_websocket_enabled(custom_llm_provider): + from litellm.rust_bridge import responses_websocket as rust_responses_websocket + + rust_responses_websocket.admit() + @asynccontextmanager async def _backend_connection(): - if _rust_responses_websocket_enabled(custom_llm_provider): - from litellm.rust_bridge import responses_websocket as rust_responses_websocket - - rust_backend: Final = await rust_responses_websocket.connect( - url=ws_url, - headers={str(key): str(value) for key, value in headers.items()}, - timeout=timeout, - ) - if rust_backend is not None: - yield rust_backend - return - async with websockets.connect( ws_url, additional_headers=headers, diff --git a/litellm/rust_bridge/chat_completions.py b/litellm/rust_bridge/chat_completions.py index 674bd8847f7..1b522fca949 100644 --- a/litellm/rust_bridge/chat_completions.py +++ b/litellm/rust_bridge/chat_completions.py @@ -12,6 +12,7 @@ retrying it there would bill the customer for the same work twice. from __future__ import annotations +import inspect import json from collections.abc import Awaitable, Callable, Mapping, Sequence from dataclasses import dataclass @@ -26,6 +27,7 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo convert_to_model_response_object, ) from litellm.llms.bedrock.request_metadata import bedrock_request_metadata_is_owned +from litellm.rust_bridge._lifecycle import invoke_terminal from litellm.rust_bridge.configuration import rust_enabled from litellm.rust_bridge.loader import get_native_bridge from litellm.rust_bridge.timeouts import timeout_to_seconds @@ -46,32 +48,12 @@ RUST_RESPONSE_HEADER: Final = "x-litellm-rust" class RustChatCompletions(Protocol): - def __call__( - self, - model: str, - messages: Sequence[object], - optional_params: Mapping[str, object] | None, - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: Mapping[str, object] | None, - timeout_seconds: float | None, - ) -> Mapping[str, object]: + def __call__(self, arguments: dict[str, object]) -> ModelResponse: raise NotImplementedError class RustAchatCompletions(Protocol): - def __call__( - self, - model: str, - messages: Sequence[object], - optional_params: Mapping[str, object] | None, - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: Mapping[str, object] | None, - timeout_seconds: float | None, - ) -> Awaitable[Mapping[str, object]]: + def __call__(self, arguments: dict[str, object]) -> Awaitable[ModelResponse]: raise NotImplementedError @@ -87,13 +69,6 @@ class RustChatCompletionsDecline(Protocol): class ResponseObserver(Protocol): - """Invoked with the payload the core returned, on success only. - - Lets the caller emit its own `post_call` on whichever path served the - request. Both entry points call it, so the synchronous and asynchronous - paths cannot drift apart the way the pre_call suppression once did. - """ - def __call__(self, rust_response: Mapping[str, object], /) -> None: raise NotImplementedError @@ -105,16 +80,6 @@ def response_logger( api_key: str, additional_args: Mapping[str, object], ) -> ResponseObserver: - """A `ResponseObserver` that emits the caller's `post_call` for a Rust-served - request. - - The core owns the provider call, so the Python transform that normally - raises this event never runs; without it every `post_call` callback goes - silent on a Rust-served request and `original_response` stays unset. The - payload is the core's normalized response rather than the provider's wire - body, which is the closest thing that crosses the bridge. - """ - def log(rust_response: Mapping[str, object], /) -> None: logging_obj.post_call( input=messages, @@ -126,6 +91,14 @@ def response_logger( return log +def _uses_argument_bag(call: object) -> bool: + try: + parameters: Final = inspect.signature(call).parameters.values() + except (TypeError, ValueError): + return True + return not any(parameter.kind is inspect.Parameter.VAR_KEYWORD for parameter in parameters) + + class _Unset: pass @@ -325,7 +298,7 @@ def _reraise_or_decline( ) -def _build_model_response( +def build_model_response( rust_response: Mapping[str, object], model_response: ModelResponse, ) -> ModelResponse: @@ -350,27 +323,64 @@ def chat_completions( custom_llm_provider: str | None, extra_headers: Mapping[str, object] | None, timeout: float | httpx.Timeout | None, - on_response: ResponseObserver, + arguments: dict[str, object] | None = None, + logging_api_key: str | None = None, + on_response: ResponseObserver | None = None, ) -> ModelResponse | None: rust_chat_completions: Final = load_rust_chat_completions() if rust_chat_completions is None: return None try: - rust_response: Final = rust_chat_completions( - model=model, - messages=messages, - optional_params=optional_params, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - timeout_seconds=timeout_to_seconds(timeout), + if _STATE.chat_completions is not None and _uses_argument_bag(rust_chat_completions): + rust_result: Final = rust_chat_completions( + _arguments( + arguments, + model, + messages, + optional_params, + model_response, + api_key, + api_base, + custom_llm_provider, + extra_headers, + timeout, + logging_api_key, + ) + ) + return rust_result + if _STATE.chat_completions is not None: + rust_response: Final = rust_chat_completions( + model=model, + messages=messages, + optional_params=optional_params, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + timeout_seconds=timeout_to_seconds(timeout), + ) + if on_response is not None: + on_response(rust_response) + return build_model_response(rust_response, model_response) + return rust_chat_completions( + _arguments( + arguments, + model, + messages, + optional_params, + model_response, + api_key, + api_base, + custom_llm_provider, + extra_headers, + timeout, + logging_api_key, + ) ) except Exception as rust_error: # noqa: BLE001 # rollout safety: the helper re-raises anything the provider already saw _reraise_or_decline(rust_error, model=model, custom_llm_provider=custom_llm_provider) return None - on_response(rust_response) - return _build_model_response(rust_response, model_response) + raise AssertionError("unreachable") async def achat_completions( @@ -384,27 +394,64 @@ async def achat_completions( custom_llm_provider: str | None, extra_headers: Mapping[str, object] | None, timeout: float | httpx.Timeout | None, - on_response: ResponseObserver, + arguments: dict[str, object] | None = None, + logging_api_key: str | None = None, + on_response: ResponseObserver | None = None, ) -> ModelResponse | None: rust_achat_completions: Final = load_rust_achat_completions() if rust_achat_completions is None: return None try: - rust_response: Final = await rust_achat_completions( - model=model, - messages=messages, - optional_params=optional_params, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - timeout_seconds=timeout_to_seconds(timeout), + if _STATE.achat_completions is not None and _uses_argument_bag(rust_achat_completions): + rust_result: Final = await rust_achat_completions( + _arguments( + arguments, + model, + messages, + optional_params, + model_response, + api_key, + api_base, + custom_llm_provider, + extra_headers, + timeout, + logging_api_key, + ) + ) + return rust_result + if _STATE.achat_completions is not None: + rust_response: Final = await rust_achat_completions( + model=model, + messages=messages, + optional_params=optional_params, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + timeout_seconds=timeout_to_seconds(timeout), + ) + if on_response is not None: + on_response(rust_response) + return build_model_response(rust_response, model_response) + return await rust_achat_completions( + _arguments( + arguments, + model, + messages, + optional_params, + model_response, + api_key, + api_base, + custom_llm_provider, + extra_headers, + timeout, + logging_api_key, + ) ) except Exception as rust_error: # noqa: BLE001 # rollout safety: the helper re-raises anything the provider already saw _reraise_or_decline(rust_error, model=model, custom_llm_provider=custom_llm_provider) return None - on_response(rust_response) - return _build_model_response(rust_response, model_response) + raise AssertionError("unreachable") async def achat_completions_or_fallback( @@ -418,8 +465,10 @@ async def achat_completions_or_fallback( custom_llm_provider: str | None, extra_headers: Mapping[str, object] | None, timeout: float | httpx.Timeout | None, - on_response: ResponseObserver, python_fallback: Callable[[], Awaitable[object]], + arguments: dict[str, object] | None = None, + logging_api_key: str | None = None, + on_response: ResponseObserver | None = None, ) -> object: """Await the Rust path, falling back to the caller's own Python path when the bridge is unavailable or the call fails. @@ -439,8 +488,44 @@ async def achat_completions_or_fallback( custom_llm_provider=custom_llm_provider, extra_headers=extra_headers, timeout=timeout, + arguments=arguments, + logging_api_key=logging_api_key, on_response=on_response, ) if response is not None: return response return await python_fallback() + + +def initialize_logging(arguments: dict[str, object], asynchronous: bool) -> object: + from litellm.rust_bridge._lifecycle import initialize_logging as initialize_lifecycle_logging + + return initialize_lifecycle_logging(arguments, asynchronous, "completion") + + +def _arguments( + arguments: dict[str, object] | None, + model: str, + messages: Sequence[object], + optional_params: Mapping[str, object], + model_response: ModelResponse, + api_key: str | None, + api_base: str | None, + custom_llm_provider: str | None, + extra_headers: Mapping[str, object] | None, + timeout: float | httpx.Timeout | None, + logging_api_key: str | None, +) -> dict[str, object]: + return { + **(arguments or {}), + "model": model, + "messages": messages, + "optional_params": optional_params, + "model_response": model_response, + "api_key": api_key, + "api_base": api_base, + "custom_llm_provider": custom_llm_provider, + "extra_headers": extra_headers, + "timeout_seconds": timeout_to_seconds(timeout), + "logging_api_key": logging_api_key if logging_api_key is not None else api_key or "", + } diff --git a/litellm/rust_bridge/responses_websocket.py b/litellm/rust_bridge/responses_websocket.py index 0634867af1c..fcac2c1693b 100644 --- a/litellm/rust_bridge/responses_websocket.py +++ b/litellm/rust_bridge/responses_websocket.py @@ -1,33 +1,15 @@ -"""Thin Python wrapper for the native Rust Responses WebSocket bridge.""" +"""Admission wrapper for the native Rust Responses WebSocket route.""" from __future__ import annotations from dataclasses import dataclass from typing import Final, Protocol -import httpx -from websockets.exceptions import ConnectionClosedOK - from litellm.rust_bridge.loader import get_native_bridge -from litellm.rust_bridge.timeouts import timeout_to_seconds -class RustResponsesWebSocket(Protocol): - async def send_text(self, text: str) -> None: ... - - async def recv_text(self) -> str | None: ... - - async def close(self) -> None: ... - - -class RustResponsesWebSocketConnection(Protocol): - @classmethod - async def connect( - cls, - url: str, - headers: dict[str, str], - timeout_seconds: float | None, - ) -> RustResponsesWebSocket: ... +class RustResponsesWebSocketRoute(Protocol): + def __call__(self) -> None: ... class _Unset: @@ -39,64 +21,35 @@ _UNSET: Final[_Unset] = _Unset() @dataclass(slots=True) class _RustResponsesWebSocketState: - connection: RustResponsesWebSocketConnection | None = None + route: RustResponsesWebSocketRoute | None = None _STATE: Final[_RustResponsesWebSocketState] = _RustResponsesWebSocketState() -def set_rust_responses_websocket( - *, - connection: RustResponsesWebSocketConnection | None | _Unset = _UNSET, -) -> None: - if not isinstance(connection, _Unset): - _STATE.connection = connection +def set_rust_responses_websocket(*, route: RustResponsesWebSocketRoute | None | _Unset = _UNSET) -> None: + if not isinstance(route, _Unset): + _STATE.route = route -def load_rust_responses_websocket() -> RustResponsesWebSocketConnection | None: - if _STATE.connection is not None: - return _STATE.connection +def load_rust_responses_websocket() -> RustResponsesWebSocketRoute | None: + if _STATE.route is not None: + return _STATE.route native_bridge: Final = get_native_bridge() if native_bridge is None: return None - connection_type: Final[RustResponsesWebSocketConnection | None] = getattr( - native_bridge, "ResponsesWebSocketConnection", None - ) - return connection_type + route: Final[RustResponsesWebSocketRoute | None] = getattr(native_bridge, "responses_websocket", None) + return route -class _ConnectionAdapter: - def __init__(self, connection: RustResponsesWebSocket): - self._connection: Final[RustResponsesWebSocket] = connection - - async def send(self, text: str) -> None: - await self._connection.send_text(text) - - async def recv(self) -> str: - message: Final = await self._connection.recv_text() - if message is None: - raise ConnectionClosedOK(None, None) - return message - - async def close(self) -> None: - await self._connection.close() - - -async def connect( - *, - url: str, - headers: dict[str, str], - timeout: float | httpx.Timeout | None, -) -> _ConnectionAdapter | None: - connection_type: Final = load_rust_responses_websocket() - if connection_type is None: - return None +def admit() -> bool: + route: Final = load_rust_responses_websocket() + if route is None: + return False + native_bridge: Final = get_native_bridge() + declined_type: Final = getattr(native_bridge, "RustBridgeDeclined", ()) if native_bridge is not None else () try: - connection: Final = await connection_type.connect( - url=url, - headers=headers, - timeout_seconds=timeout_to_seconds(timeout), - ) - except Exception: # noqa: BLE001 # bridge failures must fall back to Python - return None - return _ConnectionAdapter(connection) + route() + except declined_type: + return False + raise RuntimeError("Rust Responses WebSocket returned without taking session ownership") diff --git a/tests/test_litellm/responses/test_rust_bridge_websocket.py b/tests/test_litellm/responses/test_rust_bridge_websocket.py index 74d96bda336..5c1340b84a8 100644 --- a/tests/test_litellm/responses/test_rust_bridge_websocket.py +++ b/tests/test_litellm/responses/test_rust_bridge_websocket.py @@ -6,44 +6,21 @@ from litellm.llms.custom_httpx.llm_http_handler import _rust_responses_websocket from litellm.rust_bridge import configuration, responses_websocket -class _FakeNativeConnection: - def __init__(self) -> None: - self.sent: list[str] = [] - self.closed = False - - async def send_text(self, text: str) -> None: - self.sent.append(text) - - async def recv_text(self) -> str: - return "response.completed" - - async def close(self) -> None: - self.closed = True - - -class _ClosedNativeConnection: - async def recv_text(self) -> None: - return None +class _Declined(Exception): + pass class _FakeNativeBridge: - @classmethod - async def connect( - cls, - *, - url: str, - headers: dict[str, str], - timeout_seconds: float | None, - ) -> _FakeNativeConnection: - return _FakeNativeConnection() + RustBridgeDeclined = _Declined @pytest.fixture(autouse=True) -def reset_responses_websocket(): - responses_websocket.set_rust_responses_websocket(connection=None) +def reset_responses_websocket(monkeypatch: pytest.MonkeyPatch): + responses_websocket.set_rust_responses_websocket(route=None) configuration.reset_rust_configuration() + monkeypatch.setattr(responses_websocket, "get_native_bridge", lambda: _FakeNativeBridge) yield - responses_websocket.set_rust_responses_websocket(connection=None) + responses_websocket.set_rust_responses_websocket(route=None) configuration.reset_rust_configuration() @@ -55,42 +32,24 @@ def test_rust_websocket_bridge_uses_process_enablement() -> None: assert not _rust_responses_websocket_enabled("anthropic") -@pytest.mark.asyncio -async def test_adapter_raises_clean_close_when_rust_connection_ends() -> None: - adapter = responses_websocket._ConnectionAdapter(_ClosedNativeConnection()) - - with pytest.raises(responses_websocket.ConnectionClosedOK): - await adapter.recv() - - -@pytest.mark.asyncio -async def test_bridge_unavailable_returns_none(monkeypatch: pytest.MonkeyPatch) -> None: +def test_bridge_unavailable_declines(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr(responses_websocket, "_STATE", responses_websocket._RustResponsesWebSocketState()) monkeypatch.setattr(responses_websocket, "get_native_bridge", lambda: None) - - assert ( - await responses_websocket.connect( - url="wss://example.test/responses", - headers={}, - timeout=None, - ) - is None - ) + assert not responses_websocket.admit() -@pytest.mark.asyncio -async def test_enabled_bridge_connects_and_adapts_socket( - monkeypatch: pytest.MonkeyPatch, -) -> None: - responses_websocket.set_rust_responses_websocket(connection=_FakeNativeBridge) +def test_host_only_lifecycle_declines_before_provider_io() -> None: + def decline() -> None: + raise _Declined("host guardrails required") - connection = await responses_websocket.connect( - url="wss://example.test/responses", - headers={"Authorization": "Bearer key"}, - timeout=1.0, - ) + responses_websocket.set_rust_responses_websocket(route=decline) + assert not responses_websocket.admit() - assert connection is not None - await connection.send("response.create") - assert await connection.recv() == "response.completed" - await connection.close() + +def test_unexpected_bridge_error_does_not_allow_fallback() -> None: + def fail() -> None: + raise RuntimeError("bridge failed") + + responses_websocket.set_rust_responses_websocket(route=fail) + with pytest.raises(RuntimeError, match="bridge failed"): + responses_websocket.admit() diff --git a/tests/test_litellm/rust_bridge/test_chat_completions.py b/tests/test_litellm/rust_bridge/test_chat_completions.py index b2fd2e6dcc0..fae04249cff 100644 --- a/tests/test_litellm/rust_bridge/test_chat_completions.py +++ b/tests/test_litellm/rust_bridge/test_chat_completions.py @@ -95,16 +95,16 @@ class _RecordingCall: self.error = error self.calls: list[dict] = [] - def __call__(self, **kwargs): - self.calls.append(kwargs) + def __call__(self, arguments): + self.calls.append(arguments) if self.error is not None: raise self.error - return self.result + return bridge.build_model_response(self.result, arguments["model_response"]) class _RecordingAsyncCall(_RecordingCall): - async def __call__(self, **kwargs): - return _RecordingCall.__call__(self, **kwargs) + async def __call__(self, arguments): + return _RecordingCall.__call__(self, arguments) def _accepts(**overrides) -> bool: @@ -237,7 +237,6 @@ def _call_kwargs(model_response: ModelResponse) -> dict: "custom_llm_provider": "anthropic", "extra_headers": {}, "timeout": 30.0, - "on_response": lambda _rust_response: None, } @@ -266,6 +265,19 @@ class TestSyncCall: bridge.chat_completions(**_call_kwargs(ModelResponse())) assert native.calls[0]["timeout_seconds"] == 30.0 + def test_passes_the_full_argument_bag(self): + native = _RecordingCall() + bridge.set_rust_chat_completions(chat_completions=native) + marker = object() + kwargs = _call_kwargs(ModelResponse()) + kwargs["arguments"] = {"litellm_logging_obj": marker, "fallbacks": ["python"]} + + bridge.chat_completions(**kwargs) + + assert native.calls[0]["litellm_logging_obj"] is marker + assert native.calls[0]["fallbacks"] == ["python"] + assert native.calls[0]["model_response"] is kwargs["model_response"] + def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch): _hide_native_bridge(monkeypatch) assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None