mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
wip
This commit is contained in:
parent
0a53a00edd
commit
5b7ba86c30
68 changed files with 3588 additions and 4749 deletions
8
litellm-rust/Cargo.lock
generated
8
litellm-rust/Cargo.lock
generated
|
|
@ -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",
|
||||
]
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
}
|
||||
|
|
@ -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)]
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
}
|
||||
|
|
@ -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(),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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";
|
||||
|
|
|
|||
|
|
@ -1,5 +1,3 @@
|
|||
pub mod audio_transcription;
|
||||
pub mod realtime;
|
||||
pub mod realtime_pool;
|
||||
pub mod responses_ws;
|
||||
pub(crate) mod tls;
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
}
|
||||
}
|
||||
|
|
@ -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)]
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
}
|
||||
}
|
||||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
@ -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);
|
||||
}
|
||||
}
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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(¶ms.model);
|
||||
|
||||
// Warm path: take a pooled upstream (handshake already paid) and relay its
|
||||
// buffered session.created immediately. On miss/dead socket fall through.
|
||||
if let Some(key) = upstream_key(
|
||||
let connection = upstream_key(
|
||||
provider_model,
|
||||
params.api_key.as_deref(),
|
||||
params.api_base.as_deref(),
|
||||
) && let Some(handoff) = pool.take(&key)
|
||||
{
|
||||
return crate::io::realtime::realtime_warm(
|
||||
provider_model,
|
||||
handoff,
|
||||
).ok_or_else(|| Error::Auth("missing realtime provider API key".to_string()))?;
|
||||
let warm = pool.take(&connection);
|
||||
let context = CallLifecycleContext::new("realtime", model, "openai", call_id)
|
||||
.with_metadata(StandardLoggingMetadata {
|
||||
user_api_key_hash: metadata.user_api_key_hash,
|
||||
user_api_key_user_id: metadata.user_api_key_user_id,
|
||||
user_api_key_team_id: metadata.user_api_key_team_id,
|
||||
..Default::default()
|
||||
});
|
||||
Ok(realtime(
|
||||
&CustomLoggerRunner::new(loggers.as_ref().clone()),
|
||||
RealtimeRequest {
|
||||
connection,
|
||||
warm,
|
||||
idle_timeout,
|
||||
observe,
|
||||
client_in,
|
||||
client_out,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
// Cold path: fresh dial (the original behavior).
|
||||
crate::io::realtime::realtime(
|
||||
provider_model,
|
||||
params.api_key.as_deref(),
|
||||
params.api_base.as_deref(),
|
||||
idle_timeout,
|
||||
observe,
|
||||
},
|
||||
context,
|
||||
client_in,
|
||||
client_out,
|
||||
)
|
||||
.await
|
||||
.await)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
);
|
||||
}
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
299
litellm-rust/crates/core/src/audio_transcription/lifecycle.rs
Normal file
299
litellm-rust/crates/core/src/audio_transcription/lifecycle.rs
Normal 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}")
|
||||
}
|
||||
|
|
@ -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)]
|
||||
|
|
|
|||
|
|
@ -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 { .. })
|
||||
));
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
289
litellm-rust/crates/core/src/chat_completions/lifecycle.rs
Normal file
289
litellm-rust/crates/core/src/chat_completions/lifecycle.rs
Normal 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));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -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.
|
||||
///
|
||||
|
|
|
|||
|
|
@ -41,6 +41,7 @@ pub trait CustomGuardrail: Send + Sync {
|
|||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct CustomGuardrailRunner {
|
||||
guardrails: Vec<Arc<dyn CustomGuardrail>>,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
187
litellm-rust/crates/core/src/lifecycle/streaming.rs
Normal file
187
litellm-rust/crates/core/src/lifecycle/streaming.rs
Normal 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(),
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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()))
|
||||
})),
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
})
|
||||
}
|
||||
|
|
@ -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(¤t_url).await?;
|
||||
let response = client
|
||||
.get(current_url.clone())
|
||||
.send()
|
||||
.await
|
||||
.map_err(|err| Error::Network(err.to_string()))?;
|
||||
if !response.status().is_redirection() {
|
||||
return Ok((current_url, response));
|
||||
}
|
||||
current_url = redirect_location(&response, ¤t_url)?;
|
||||
}
|
||||
|
||||
Err(Error::InvalidRequest(
|
||||
"Too many redirects while fetching OCR document URL".to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
fn enforce_download_size(content_length: u64, max_bytes: u64, url: &Url) -> Result<(), Error> {
|
||||
if max_bytes == 0 {
|
||||
return Err(Error::InvalidRequest(format!(
|
||||
"OCR document URL download is disabled (MAX_IMAGE_URL_DOWNLOAD_SIZE_MB=0). url={url}"
|
||||
)));
|
||||
}
|
||||
if content_length > max_bytes {
|
||||
let size_mb = content_length as f64 / (1024.0 * 1024.0);
|
||||
let max_size_mb = max_bytes as f64 / (1024.0 * 1024.0);
|
||||
return Err(Error::InvalidRequest(format!(
|
||||
"OCR document size ({size_mb:.2}MB) exceeds maximum allowed size ({max_size_mb:.2}MB). url={url}"
|
||||
)));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn read_response_with_limit(
|
||||
mut response: reqwest::Response,
|
||||
url: &Url,
|
||||
) -> Result<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()
|
||||
)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -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())
|
||||
}
|
||||
|
|
@ -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))
|
||||
}
|
||||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
})
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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>,
|
||||
}
|
||||
|
|
@ -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)]
|
||||
|
|
|
|||
|
|
@ -1,2 +1,8 @@
|
|||
mod streaming;
|
||||
|
||||
pub mod transformation;
|
||||
pub mod types;
|
||||
|
||||
pub use streaming::{
|
||||
RealtimeConnectionSpec, RealtimeRequest, WarmConnection, realtime, warmup,
|
||||
};
|
||||
|
|
|
|||
561
litellm-rust/crates/core/src/realtime/streaming.rs
Normal file
561
litellm-rust/crates/core/src/realtime/streaming.rs
Normal 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());
|
||||
}
|
||||
}
|
||||
|
|
@ -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);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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!(),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
);
|
||||
}
|
||||
|
|
@ -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]]
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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::*;
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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")?;
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)?)
|
||||
}
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 "",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue