This commit is contained in:
Yujong Lee 2026-09-07 22:02:46 -07:00
parent 0a53a00edd
commit 5b7ba86c30
68 changed files with 3588 additions and 4749 deletions

View file

@ -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",
]

View file

@ -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

View file

@ -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

View file

@ -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<Box<dyn Future<Output = ActionResult<T, Error>> + 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<PreparedAudioTranscriptionRequest, Error> {
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<ProviderAudioTranscriptionRequest, Error> {
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<ProviderAudioTranscriptionRequest, Error> {
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<PreparedAudioTranscriptionRequest, ProviderAudioTranscriptionRequest>
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))
}

View file

@ -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<Value, Error> {
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)]

View file

@ -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}")
}

View file

@ -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<String>,
pub(crate) api_base: Option<String>,
pub(crate) extra_headers: Option<Map<String, Value>>,
pub(crate) optional_params: Map<String, Value>,
pub(crate) timeout: Option<Duration>,
}
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(),
)
}
}

View file

@ -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";

View file

@ -1,5 +1,3 @@
pub mod audio_transcription;
pub mod realtime;
pub mod realtime_pool;
pub mod responses_ws;
pub(crate) mod tls;

View file

@ -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<MaybeTlsStream<TcpStream>>;
pub(crate) type UpstreamTx = SplitSink<UpstreamWs, Message>;
pub(crate) type UpstreamRx = SplitStream<UpstreamWs>;
/// 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<String, Error> {
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<UpstreamWs, Error> {
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<RealtimeEvent, Error> {
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<In, Out>(
model: &str,
mut upstream_tx: UpstreamTx,
mut upstream_rx: UpstreamRx,
prelude: Option<RealtimeEvent>,
idle_timeout: Option<Duration>,
mut observe: impl FnMut(&RealtimeEvent) + Send,
mut client_in: In,
mut client_out: Out,
) -> Result<(), Error>
where
In: Stream<Item = RealtimeEvent> + Unpin + Send,
Out: Sink<RealtimeEvent> + Unpin + Send,
<Out as Sink<RealtimeEvent>>::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<In, Out>(
model: &str,
api_key: Option<&str>,
api_base: Option<&str>,
idle_timeout: Option<Duration>,
observe: impl FnMut(&RealtimeEvent) + Send,
client_in: In,
client_out: Out,
) -> Result<(), Error>
where
In: Stream<Item = RealtimeEvent> + Unpin + Send,
Out: Sink<RealtimeEvent> + Unpin + Send,
<Out as Sink<RealtimeEvent>>::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<In, Out>(
model: &str,
handoff: crate::io::realtime_pool::WarmHandoff,
idle_timeout: Option<Duration>,
observe: impl FnMut(&RealtimeEvent) + Send,
client_in: In,
client_out: Out,
) -> Result<(), Error>
where
In: Stream<Item = RealtimeEvent> + Unpin + Send,
Out: Sink<RealtimeEvent> + Unpin + Send,
<Out as Sink<RealtimeEvent>>::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::<RealtimeEvent>();
// provider -> client (we hold `backend_rx` to read backend events)
let (client_out, mut backend_rx) = mpsc::unbounded::<RealtimeEvent>();
// 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;
}
}

View file

@ -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<String>,
}
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<WarmHandoff> {
pub fn take(&self, key: &UpstreamKey) -> Option<litellm_core::realtime::WarmConnection> {
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<WarmConnection, Error> {
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<WarmConnection, litellm_core::Error> {
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<UpstreamKey> {
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)]

View file

@ -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<MaybeTlsStream<TcpStream>>;
type UpstreamTx = SplitSink<ResponsesUpstreamWs, Message>;
type UpstreamRx = SplitStream<ResponsesUpstreamWs>;
#[derive(Clone)]
pub struct ResponsesWebSocketConnection {
socket: Arc<Mutex<Option<ResponsesUpstreamWs>>>,
}
impl ResponsesWebSocketConnection {
pub async fn connect_url(
url: &str,
headers: &HashMap<String, String>,
timeout: Option<Duration>,
) -> Result<Self, Error> {
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::<HeaderName>()
.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<Option<String>, 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<String, Error> {
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<ResponsesUpstreamWs, Error> {
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<In, Out>(
model: &str,
upstream_tx: UpstreamTx,
upstream_rx: UpstreamRx,
idle_timeout: Option<Duration>,
observe: impl FnMut(&ResponsesWsEvent) + Send,
client_in: In,
client_out: Out,
) -> Result<(), Error>
where
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
Out: Sink<ResponsesWsEvent> + 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<In, Out>(
model: &str,
mut upstream_tx: UpstreamTx,
mut upstream_rx: UpstreamRx,
idle_timeout: Option<Duration>,
mut observe: impl FnMut(&ResponsesWsEvent) + Send,
mut client_in: In,
mut client_out: Out,
) -> Result<(), Error>
where
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
Out: Sink<ResponsesWsEvent> + 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::<ResponsesWsEvent>(&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<In, Out>(
model: &str,
api_key: Option<&str>,
api_base: Option<&str>,
first_frame: Option<ResponsesWsEvent>,
idle_timeout: Option<Duration>,
mut observe: impl FnMut(&ResponsesWsEvent) + Send,
client_in: In,
client_out: Out,
) -> Result<(), Error>
where
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
Out: Sink<ResponsesWsEvent> + 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<In, Out>(
model: &str,
api_key: Option<&str>,
api_base: Option<&str>,
first_frame: Option<ResponsesWsEvent>,
idle_timeout: Option<Duration>,
observe: impl FnMut(&ResponsesWsEvent) + Send,
client_in: In,
client_out: Out,
) -> Result<(), Error>
where
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
Out: Sink<ResponsesWsEvent> + 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::<ResponsesWsEvent>();
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::<ResponsesWsEvent>();
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::<ResponsesWsEvent>();
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");
}
}

View file

@ -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;

View file

@ -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;

View file

@ -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<Arc<dyn CustomLogger>>,
/// 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<Arc<dyn CustomLogger>>,
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<Option<String>>,
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<Arc<dyn CustomLogger>> = 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<Arc<dyn CustomLogger>> = 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);
}
}

View file

@ -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<Response, MessagesRouteError> {
let content_type = upstream
.headers()
.get(CONTENT_TYPE)
.cloned()
fn stream_response(call: StreamingCall) -> Result<Response, MessagesRouteError> {
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")

View file

@ -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)

View file

@ -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<ModelRouter>,
@ -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::<RealtimeEvent>(&text).ok(),
_ => None,
}
});
// Plain forwarding sink — no observe here anymore.
let client_out = ws_sink.with(|event: RealtimeEvent| async move {
Ok::<Message, axum::Error>(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;
}

View file

@ -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<In, Out>(
pool: &RealtimePool,
model: &str,
idle_timeout: Option<Duration>,
observe: impl FnMut(&RealtimeEvent) + Send,
loggers: Arc<Vec<Arc<dyn CustomLogger>>>,
call_id: String,
metadata: RequestMetadata,
client_in: In,
client_out: Out,
) -> Result<(), Error>
) -> Result<ExecutedCall<(), Error>, Error>
where
In: Stream<Item = RealtimeEvent> + Unpin + Send,
Out: Sink<RealtimeEvent> + Unpin + Send,
@ -44,34 +51,29 @@ where
.strip_prefix("openai/")
.unwrap_or(&params.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)
}

View file

@ -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;

View file

@ -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<Vec<Arc<dyn CustomLogger>>>) -> 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<In, Out>(
@ -26,7 +49,7 @@ pub async fn run<In, Out>(
metadata: RequestMetadata,
client_in: In,
client_out: Out,
) -> Result<(), Error>
) -> Result<ExecutedCall<(), Error>, Error>
where
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
Out: Sink<ResponsesWsEvent> + 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<Vec<Arc<dyn CustomLogger>>>,
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<LoggingError>,
) -> (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
}

View file

@ -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"
);
}

View file

@ -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:

View file

@ -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

View file

@ -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<Arc<dyn CustomLogger>>,
guardrails: Vec<Arc<dyn CustomGuardrail>>,
) -> 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<S: AudioServices>(
services: &S,
request: AudioRouteRequest<'_>,
) -> ExecutedCall<Value, Error> {
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<String>,
api_base: Option<String>,
extra_headers: Option<Map<String, Value>>,
optional_params: Map<String, Value>,
timeout: Option<std::time::Duration>,
}
struct AudioRequestPolicy {
guardrail_runner: CustomGuardrailRunner,
request_metadata: RequestMetadata,
}
type AudioFuture<'a, T> = Pin<Box<dyn Future<Output = ActionResult<T, Error>> + Send + 'a>>;
impl AudioRequestPolicy {
async fn run_pre_call_guardrails(
&self,
request: PreparedAudioTranscriptionRequest,
) -> Result<PreparedAudioTranscriptionRequest, Error> {
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<ProviderAudioTranscriptionRequest, Error> {
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<ProviderAudioTranscriptionRequest, Error> {
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<PreparedAudioTranscriptionRequest, ProviderAudioTranscriptionRequest>
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}")
}

View file

@ -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<Value, Error> {
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)]

View file

@ -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<Vec<&'static str>>,
}
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 { .. })
));
}

View file

@ -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<Duration>,
}
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<Map<String, Value>>,
pub optional_params: Map<String, Value>,
pub timeout: Option<Duration>,
pub request_metadata: RequestMetadata,
pub litellm_call_id: Option<&'a str>,
}
#[derive(Clone)]
pub struct ProviderAudioTranscriptionRequest {
pub(super) model: String,

View file

@ -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<String, Value>,
pub custom_llm_provider: Option<String>,
}
#[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<Result<Self::State, Decline>, 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<Transition, Error> {
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<Result<Lifecycle<ChatCompletionsRoute>, 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));
}
}
}

View file

@ -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<ChatCompletionsResponse, Error> {
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.
///

View file

@ -41,6 +41,7 @@ pub trait CustomGuardrail: Send + Sync {
}
}
#[derive(Clone)]
pub struct CustomGuardrailRunner {
guardrails: Vec<Arc<dyn CustomGuardrail>>,
}

View file

@ -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<InitialReq, ProviderReq, Services, ProviderCall, ProviderFuture>(
&self,
context: CallLifecycleContext,
request: InitialReq,
services: Arc<Services>,
observer: Box<dyn StreamingObserver>,
provider_call: ProviderCall,
) -> Result<StreamingCall, Error>
where
Services: RequestPolicy<InitialReq, ProviderReq> + TerminalDispatcher + Clock + 'static,
ProviderCall: FnOnce(ProviderReq) -> ProviderFuture,
ProviderFuture: Future<Output = Result<StreamingSource, Error>>,
{
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,

View file

@ -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};

View file

@ -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

View file

@ -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<Box<dyn Stream<Item = Result<Bytes, Error>> + Send>>;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct StreamingMetadata {
pub status: u16,
pub content_type: Option<String>,
pub cache_control: Option<String>,
}
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<S>(
source: StreamingSource,
observer: Box<dyn StreamingObserver>,
context: CallLifecycleContext,
start_time: f64,
services: Arc<S>,
) -> 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<oneshot::Receiver<StreamTerminal>>,
context: Option<CallLifecycleContext>,
start_time: f64,
services: Arc<dyn CompletionServices>,
}
impl StreamingCompletion {
pub fn register(mut self) -> JoinHandle<TerminalRecord> {
self.spawn()
}
fn spawn(&mut self) -> JoinHandle<TerminalRecord> {
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<T> CompletionServices for T where T: Clock + TerminalDispatcher {}
struct StreamTerminal {
usage: Usage,
projection: Value,
classification: TerminalClassification,
}
struct ObservedStream {
inner: BytesStream,
observer: Box<dyn StreamingObserver>,
sender: Option<oneshot::Sender<StreamTerminal>>,
}
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<Bytes, Error>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
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(),
});
}
}

View file

@ -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(),

View file

@ -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<reqwest::Response, Error> {
) -> Result<StreamingSource, Error> {
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()))
})),
})
}

View file

@ -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<S: MessagesServices>(
.await
}
pub async fn messages_stream<S: MessagesServices + 'static>(
services: Arc<S>,
request: MessagesRequest,
_options: Options,
context: CallLifecycleContext,
) -> Result<StreamingCall, Error> {
CallLifecycle
.run_streaming(
context,
request,
services,
Box::<AnthropicUsageObserver>::default(),
|request| async move { execute_messages_provider_stream(request).await },
)
.await
}
#[derive(Default)]
struct AnthropicUsageObserver {
pending: Vec<u8>,
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::<serde_json::Value>(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::<Vec<_>>();
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<Lifecycle<MessagesRoute>, Error> {
Lifecycle::admit(&(), options).map(|result| match result {
Ok(machine) => machine,

View file

@ -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<AnthropicMessagesResponse, Error> {
@ -42,8 +43,25 @@ pub async fn messages(request: MessagesRequest) -> Result<AnthropicMessagesRespo
.into_result()
}
pub async fn messages_stream(request: MessagesRequest) -> Result<reqwest::Response, Error> {
execute_messages_provider_stream(request).await
pub async fn messages_stream(request: MessagesRequest) -> Result<StreamingCall, Error> {
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::<u128>()),
);
lifecycle::messages_stream(
Arc::new(lifecycle::NoopServices),
request,
lifecycle::Options::default(),
context,
)
.await
}
#[cfg(test)]

View file

@ -1,13 +0,0 @@
use std::sync::OnceLock;
use std::time::Duration;
pub(super) fn http_client() -> &'static reqwest::Client {
static CLIENT: OnceLock<reqwest::Client> = 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())
})
}

View file

@ -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<Map<String, Value>>,
) -> Result<Vec<(String, String)>, 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<Option<(&str, &str)>, 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::<f64>().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::<IpAddr>() {
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<Url, Error> {
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(&current_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, &current_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<Vec<u8>, 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<Value, Error> {
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::<u64>().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<Duration>,
) -> Result<Value, Error> {
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()
)
);
}
}

View file

@ -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<Value, Error> {
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())
}

View file

@ -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<Box<dyn Future<Output = ActionResult<T, Error>> + 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<PreparedOcrRequest, Error> {
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<ProviderOcrRequest, Error> {
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<Value, Error> {
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<std::time::Duration>,
upstream_headers: &[(String, String)],
) -> Result<Value, Error> {
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<PreparedOcrRequest, PreparedOcrRequest> 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<String, Value>), 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<Value, Error> {
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))
}

View file

@ -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<T> 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<S: OcrServices>(
services: &S,
request: OcrRouteRequest<'_>,
request: SettledOcrRequest,
_options: crate::lifecycle::ocr::Options,
context: CallLifecycleContext,
) -> ExecutedCall<Value, Error> {
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<PreparedOcrCall, PreparedOcrCall> for PreparedOcrPolicy {
impl crate::lifecycle::RequestPolicy<SettledOcrRequest, SettledOcrRequest> for SettledOcrPolicy {
type PreCallFuture<'a>
= std::future::Ready<crate::lifecycle::ActionResult<PreparedOcrCall, Error>>
= std::future::Ready<crate::lifecycle::ActionResult<SettledOcrRequest, Error>>
where
Self: 'a;
type DuringCallFuture<'a>
= std::future::Ready<crate::lifecycle::ActionResult<PreparedOcrCall, Error>>
= std::future::Ready<crate::lifecycle::ActionResult<SettledOcrRequest, Error>>
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<PreparedOcrCall, PreparedOcrCall> 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<OcrResponseData, Error> {
let config = prepare::provider_config(&prepared.custom_llm_provider, &prepared.model)?;
pub(crate) async fn send(request: SettledOcrRequest) -> Result<OcrResponseData, Error> {
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)
}

View file

@ -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<PreparedOcr, Error> {
pub fn prepare(request: OcrAdmissionRequest) -> Result<OcrDraft, Error> {
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<PreparedOcr, Error> {
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,
})
}

View file

@ -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<Map<String, Value>>,
pub optional_params: Map<String, Value>,
pub timeout: Option<Duration>,
pub callbacks: Vec<Arc<dyn CustomLogger>>,
pub guardrails: Vec<Arc<dyn CustomGuardrail>>,
pub request_metadata: RequestMetadata,
pub litellm_call_id: Option<&'a str>,
}
pub enum OcrRouteRequest<'a> {
Native(OcrRequest<'a>),
Prepared(PreparedOcrCall),
}
impl<'a> From<OcrRequest<'a>> for OcrRouteRequest<'a> {
fn from(request: OcrRequest<'a>) -> Self {
Self::Native(request)
}
}
impl From<PreparedOcrCall> 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<String>,
pub(crate) api_base: Option<String>,
pub(crate) extra_headers: Option<Map<String, Value>>,
pub(crate) optional_params: Map<String, Value>,
pub(crate) timeout: Option<Duration>,
}
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<String, Value>,
pub(crate) upstream_headers: Vec<(String, String)>,
pub(crate) timeout: Option<Duration>,
}

View file

@ -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<String, Value>,
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)]

View file

@ -1,2 +1,8 @@
mod streaming;
pub mod transformation;
pub mod types;
pub use streaming::{
RealtimeConnectionSpec, RealtimeRequest, WarmConnection, realtime, warmup,
};

View file

@ -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<MaybeTlsStream<TcpStream>>;
static TLS_CONFIG: OnceLock<Arc<ClientConfig>> = OnceLock::new();
#[derive(Clone, Eq)]
pub struct RealtimeConnectionSpec {
model: String,
api_key: String,
api_base: Option<String>,
}
impl RealtimeConnectionSpec {
pub fn new(
model: impl Into<String>,
api_key: Option<&str>,
api_base: Option<&str>,
) -> Result<Self, Error> {
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<H: Hasher>(&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<WarmConnection>,
pub idle_timeout: Option<Duration>,
}
pub async fn warmup(connection: &RealtimeConnectionSpec) -> Result<WarmConnection, Error> {
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<S, In, Out>(
services: &S,
request: RealtimeRequest,
context: CallLifecycleContext,
client_in: In,
client_out: Out,
) -> ExecutedCall<(), Error>
where
S: TerminalDispatcher + Clock,
In: Stream<Item = RealtimeEvent> + Unpin + Send,
Out: Sink<RealtimeEvent> + 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<In, Out>(
connection: WarmConnection,
model: &str,
idle_timeout: Duration,
observation: &mut RealtimeObservation,
mut client_in: In,
mut client_out: Out,
) -> Result<(), Error>
where
In: Stream<Item = RealtimeEvent> + Unpin + Send,
Out: Sink<RealtimeEvent> + 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::<RealtimeEvent>(&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<Out>(
client_out: &mut Out,
event: &RealtimeEvent,
model: &str,
) -> Result<(), Error>
where
Out: Sink<RealtimeEvent> + 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<RealtimeEvent, Error> {
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<Upstream, Error> {
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<Arc<ClientConfig>, 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<String, Error> {
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<Vec<TerminalRecord>>,
}
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<RealtimeEvent>, 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());
}
}

View file

@ -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<String>,
pub user_api_key_user_id: Option<String>,
pub user_api_key_team_id: Option<String>,
}
#[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<ResponsesWsLogOutcome>,
pub struct ResponsesWsObservation {
pub(crate) model: String,
pub(crate) usage: Usage,
}
#[derive(Default)]
pub struct ResponsesWsInstrumentation {
state: Mutex<InstrumentationState>,
state: Mutex<ResponsesWsObservation>,
}
impl ResponsesWsInstrumentation {
pub fn new(
litellm_call_id: impl Into<String>,
model: impl Into<String>,
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<ResponsesWsLogOutcome> {
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<Box<dyn Future<Output = ActionResult<T, Error>> + 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);
}
}

View file

@ -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<MaybeTlsStream<TcpStream>>;
static TLS_CONFIG: OnceLock<Arc<ClientConfig>> = OnceLock::new();
pub trait ResponsesWebSocketProviderConfig: Sync {
fn supports_native_websocket(&self) -> bool {
false
@ -28,36 +58,255 @@ pub trait ResponsesWebSocketProviderConfig: Sync {
) -> Result<ResponsesWsTransformResult, Error>;
}
pub fn complete_websocket_url(
api_base: Option<&str>,
pub struct ResponsesWebSocketRequest {
pub model: String,
pub api_key: Option<String>,
pub api_base: Option<String>,
pub first_frame: Option<ResponsesWsEvent>,
pub idle_timeout: Option<Duration>,
}
pub async fn responses_websocket<S, In, Out>(
services: &S,
request: ResponsesWebSocketRequest,
context: crate::lifecycle::CallLifecycleContext,
client_in: In,
client_out: Out,
) -> Result<ExecutedCall<(), Error>, Error>
where
S: TerminalDispatcher + crate::lifecycle::Clock,
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
Out: Sink<ResponsesWsEvent> + 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<In, Out>(
upstream: Upstream,
model: &str,
model_in_websocket_url: bool,
) -> String {
first_frame: Option<ResponsesWsEvent>,
idle_timeout: Duration,
instrumentation: &ResponsesWsInstrumentation,
mut client_in: In,
mut client_out: Out,
) -> Result<(), Error>
where
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
Out: Sink<ResponsesWsEvent> + 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::<ResponsesWsEvent>(&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<Upstream, Message>,
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<Upstream, Error> {
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<Arc<ClientConfig>, 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<String, Error> {
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<Vec<TerminalRecord>>,
}
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<ResponsesWsEvent>,
) = 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");
}
}

View file

@ -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::<Vec<_>>()
.await
.into_iter()
.collect::<Result<Vec<_>, _>>()
.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::<Vec<_>>().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::<Vec<_>>().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();
}

View file

@ -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<litellm_core::ocr::OcrResponseData, Error> {
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!(),
}

View file

@ -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<Value, Error> {
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::<usize>().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<String>,
response_object: Option<String>,
error_kind: Option<String>,
}
#[derive(Default)]
struct RecordingOcrLogger {
events: Mutex<Vec<RecordedLogEvent>>,
}
impl RecordingOcrLogger {
fn events(&self) -> Vec<RecordedLogEvent> {
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<GuardrailEventHook>,
events: Mutex<Vec<&'static str>>,
block_pre_call: bool,
block_during_call: bool,
}
impl RecordingOcrGuardrail {
fn new(hooks: Vec<GuardrailEventHook>) -> 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<_>>(),
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<_>>(),
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}"
);
}

View file

@ -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]]

View file

@ -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<Value>,
timeout_seconds: Option<f64>,
) -> PyResult<Bound<'py, PyAny>> {
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<Bound<'py, PyAny>> {
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<Bound<'py, PyAny>> {
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<Bound<'py, PyAny>> {
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::ResponsesWebSocketConnection>()?;
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<String> = 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");
}
}

View file

@ -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<Value> {
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<Value>,
@ -88,25 +73,6 @@ pub(crate) fn optional_timeout(timeout_seconds: Option<f64>) -> PyResult<Option<
.transpose()
}
pub(crate) fn marshal_headers(headers: Option<Value>) -> PyResult<HashMap<String, String>> {
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::*;

View file

@ -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<impl Future<Output = Result<ChatCompletionsResponse, Error>> + 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<Py<PyDict>>,
model: Option<String>,
messages: Option<Value>,
optional_params: Option<Map<String, Value>>,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
extra_headers: Option<Map<String, Value>>,
timeout: Option<std::time::Duration>,
terminal: Option<TerminalRecord>,
}
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<Option<String>> {
arguments
.get_item(name)?
.filter(|value| !value.is_none())
.map(|value| value.extract::<String>())
.transpose()
}
fn value(arguments: &Bound<'_, PyDict>, name: &str) -> PyResult<Value> {
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<Option<Map<String, Value>>> {
arguments
.get_item(name)?
.filter(|value| !value.is_none())
.map(|value| from_py(&value))
.transpose()
}
fn admission(arguments: &Bound<'_, PyDict>) -> PyResult<Admission> {
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<ChatCompletionsRoute>,
asynchronous: bool,
}
#[pymethods]
impl ChatCompletionsLifecycle {
#[new]
fn new(
arguments: &Bound<'_, PyDict>,
asynchronous: bool,
internal_call: bool,
) -> PyResult<Self> {
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<bool> {
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<bool> {
match self.machine.operation() {
Operation::Complete(outcome) => Some(outcome == Outcome::Success),
_ => None,
}
}
}
#[pyfunction]
fn invoke(
py: Python<'_>,
machine: Py<ChatCompletionsLifecycle>,
host: Py<PyAny>,
) -> PyResult<(bool, Py<PyAny>)> {
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<PyDict>) -> PyResult<Py<ChatCompletionsState>> {
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::<f64>())
.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<String, Value>,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
extra_headers: Option<Map<String, Value>>,
timeout: Option<std::time::Duration>,
call_id: String,
}
fn take_request(py: Python<'_>, state: &Py<ChatCompletionsState>) -> PyResult<OwnedRequest> {
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<litellm_core::chat_completions::types::ChatCompletionsResponse, Error> {
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<ChatCompletionsState>,
executed: ExecutedCall<litellm_core::chat_completions::types::ChatCompletionsResponse, Error>,
) -> PyResult<Py<PyAny>> {
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<ChatCompletionsState>) -> PyResult<Bound<'_, PyAny>> {
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<ChatCompletionsState>) -> PyResult<Py<PyAny>> {
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<ChatCompletionsState>) -> PyResult<Py<PyAny>> {
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::<pyo3::types::PyList>() {
return Err(PyTypeError::new_err("messages must be a list"));
}
Ok(())
}
#[pyfunction]
fn chat_completions(py: Python<'_>, arguments: Py<PyDict>) -> PyResult<Bound<'_, PyAny>> {
validate_arguments(arguments.bind(py))?;
driver(py)?.getattr("drive_sync")?.call1((arguments,))
}
#[pyfunction]
fn achat_completions(py: Python<'_>, arguments: Py<PyDict>) -> PyResult<Bound<'_, PyAny>> {
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<Value>,
#[pyo3(from_py_with = litellm_python_interop::from_py)] optional_params: Option<
Map<String, Value>,
>,
custom_llm_provider: Option<String>,
) -> PyResult<Option<String>> {
let optional_params = object_or_empty("optional_params", optional_params)?;
Ok(chat_completions_decline_reason(
) -> Option<String> {
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<serde_json::Value>,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
extra_headers: Option<serde_json::Value>,
timeout_seconds: Option<f64>,
},
prepare = prepare_chat_completions,
errors = chat_completions_error_to_pyerr,
extra = [chat_completions_decline],
fn driver(py: Python<'_>) -> PyResult<&Bound<'_, PyModule>> {
static DRIVER: PyOnceLock<Py<PyModule>> = 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::<ChatCompletionsLifecycle>())?;
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)
}

View file

@ -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")?;

View file

@ -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<Py<PyDict>>,
headers: Option<Py<PyDict>>,
logging: Option<Py<PyAny>>,
prepared: Option<PreparedOcr>,
pre_call: Option<Py<PyDict>>,
endpoint: Option<OcrEndpoint>,
asynchronous: bool,
terminal: Option<TerminalRecord>,
}
@ -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<PyDict>, 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<PyDict>, asynchronous: bool) -> PyResul
.get_item("document")?
.ok_or_else(|| PyValueError::new_err("OCR requires document"))?
.cast_into::<PyDict>()?;
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::<PyDict>()?;
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<PyDict>, 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<PyDict>, 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<PyDict>, 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<OcrState>) -> 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<OcrState>) -> PyResult<OcrWireRequest> {
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<OcrState>) -> PyResult<OcrWireRequest> {
.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<OcrState>) -> PyResult<Bound<'_, PyAny>> {
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<OcrState>) -> PyResult<Bound<'_, PyAny>> {
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<PyDict>) -> PyResult<Py<PyAny>> {
#[pyfunction]
fn send_sync(py: Python<'_>, state: Py<OcrState>) -> PyResult<Py<PyAny>> {
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<OcrState>) -> PyResult<Py<PyAny>> {
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::<OcrLifecycle>())?;
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<OcrState>) -> PyResult<Py<PyAny>> {
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)

View file

@ -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)?)
}

View file

@ -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,

View file

@ -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,

View file

@ -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,

View file

@ -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 "",
}

View file

@ -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")

View file

@ -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()

View file

@ -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