diff --git a/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs b/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs index 35243a6dcf0..4336cd6e3a6 100644 --- a/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs @@ -2,8 +2,10 @@ use litellm_core::audio_transcription::{ AudioTranscriptionRequest as CoreAudioTranscriptionRequest, ProviderAudioTranscriptionRequest, prepare_audio_transcription_provider_call, }; -use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming}; 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; @@ -12,12 +14,8 @@ use super::types::PreparedAudioTranscriptionRequest; use litellm_core::integrations::custom_guardrail::{ CustomGuardrailRunner, GuardrailContext, GuardrailError, GuardrailRequest, }; -use litellm_core::integrations::custom_logger::{ - CallType, CallbackTiming, CallbackValue, CustomLoggerRunner, LoggingError, ModelCallDetails, -}; -use litellm_core::integrations::types::{ - RequestMetadata, StandardLoggingMetadata, StandardLoggingPayload, -}; +use litellm_core::integrations::custom_logger::{CallType, CustomLoggerRunner, LogFuture}; +use litellm_core::integrations::types::RequestMetadata; pub(crate) struct AudioTranscriptionLifecycleHooks { logger_runner: CustomLoggerRunner, @@ -25,8 +23,7 @@ pub(crate) struct AudioTranscriptionLifecycleHooks { request_metadata: RequestMetadata, } -type AudioFuture<'a, T> = Pin> + Send + 'a>>; -type AudioLogFuture<'a> = Pin + Send + 'a>>; +type AudioFuture<'a, T> = Pin> + Send + 'a>>; impl AudioTranscriptionLifecycleHooks { pub(crate) fn new( @@ -144,50 +141,25 @@ impl AudioTranscriptionLifecycleHooks { })?; Ok(request.with_body(body)) } - - fn logging_payload( - &self, - context: &CallLifecycleContext, - timing: &CallLifecycleTiming, - ) -> StandardLoggingPayload { - StandardLoggingPayload { - id: context.litellm_call_id.clone(), - litellm_call_id: context.litellm_call_id.clone(), - call_type: context.call_type.clone(), - model: context.model.clone(), - custom_llm_provider: context.custom_llm_provider.clone(), - response_cost: 0.0, - prompt_tokens: 0, - completion_tokens: 0, - total_tokens: 0, - start_time: timing.start_time, - end_time: timing.end_time, - stream: false, - metadata: StandardLoggingMetadata { - user_api_key_hash: self.request_metadata.user_api_key_hash.clone(), - user_api_key_user_id: self.request_metadata.user_api_key_user_id.clone(), - user_api_key_team_id: self.request_metadata.user_api_key_team_id.clone(), - ..Default::default() - }, - messages: None, - } - } } -impl CallLifecycleHooks +impl RequestPolicy for AudioTranscriptionLifecycleHooks { type PreCallFuture<'a> = AudioFuture<'a, PreparedAudioTranscriptionRequest>; type DuringCallFuture<'a> = AudioFuture<'a, ProviderAudioTranscriptionRequest>; - type SuccessFuture<'a> = AudioLogFuture<'a>; - type FailureFuture<'a> = AudioLogFuture<'a>; fn async_pre_call_hook<'a>( &'a self, _context: &'a CallLifecycleContext, request: PreparedAudioTranscriptionRequest, ) -> Self::PreCallFuture<'a> { - Box::pin(async move { self.run_pre_call_guardrails(request).await }) + 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>( @@ -195,59 +167,31 @@ impl CallLifecycleHooks Self::DuringCallFuture<'a> { - Box::pin(async move { self.prepare_provider_request(request).await }) - } - - fn async_log_success_event<'a>( - &'a self, - context: &'a CallLifecycleContext, - response: &'a Value, - timing: &'a CallLifecycleTiming, - ) -> Self::SuccessFuture<'a> { Box::pin(async move { - if self.logger_runner.is_empty() { - return; + match self.prepare_provider_request(request).await { + Ok(request) => ActionResult::Replace(request), + Err(error) => ActionResult::Reject(error), } - self.logger_runner - .async_log_success_event( - &ModelCallDetails::from_standard_logging_payload( - self.logging_payload(context, timing), - ), - &CallbackValue::new("audio_transcription", response.clone()), - CallbackTiming::new(timing.start_time, timing.end_time), - ) - .await; }) } +} - fn async_log_failure_event<'a>( - &'a self, - context: &'a CallLifecycleContext, - error: &'a Error, - timing: &'a CallLifecycleTiming, - ) -> Self::FailureFuture<'a> { - Box::pin(async move { - if self.logger_runner.is_empty() { - return; - } - let logging_error = LoggingError { - message: error.to_string(), - kind: core_error_kind(error).to_string(), - }; - self.logger_runner - .async_log_failure_event( - &ModelCallDetails::from_standard_logging_payload( - self.logging_payload(context, timing), - ) - .with_failure_error(logging_error.clone()), - Some(&CallbackValue::new( - "error", - json!({"message": logging_error.message, "kind": logging_error.kind}), - )), - CallbackTiming::new(timing.start_time, timing.end_time), - ) - .await; - }) +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() } } @@ -266,19 +210,3 @@ fn guardrail_context(metadata: &RequestMetadata) -> GuardrailContext { fn guardrail_error_to_core_error(error: GuardrailError) -> Error { Error::InvalidRequest(format!("{}: {}", error.kind, error.message)) } - -fn core_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", - } -} diff --git a/litellm-rust/crates/ai-gateway/src/audio_transcription/mod.rs b/litellm-rust/crates/ai-gateway/src/audio_transcription/mod.rs index 03d621b8414..b39a4b5600a 100644 --- a/litellm-rust/crates/ai-gateway/src/audio_transcription/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/mod.rs @@ -1,6 +1,6 @@ use litellm_core::Error; use litellm_core::audio_transcription::execute_audio_transcription_provider_call; -use litellm_core::call_lifecycle::CallLifecycle; +use litellm_core::lifecycle::CallLifecycle; use serde_json::Value; mod hooks; @@ -15,7 +15,7 @@ pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> Resu let PreparedAudioTranscriptionCall { request, hooks } = prepare_audio_transcription_call(request); CallLifecycle::default() - .run_request(request, &hooks, execute_audio_transcription_provider_call) + .run_request_result(request, &hooks, execute_audio_transcription_provider_call) .await } diff --git a/litellm-rust/crates/ai-gateway/src/audio_transcription/types.rs b/litellm-rust/crates/ai-gateway/src/audio_transcription/types.rs index 72f9ef7a47e..ee715e60e01 100644 --- a/litellm-rust/crates/ai-gateway/src/audio_transcription/types.rs +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/types.rs @@ -1,7 +1,7 @@ use std::sync::Arc; use std::time::Duration; -use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleRequest}; +use litellm_core::lifecycle::{CallLifecycleContext, CallLifecycleRequest}; use serde_json::{Map, Value}; use litellm_core::integrations::custom_guardrail::CustomGuardrail; diff --git a/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs index bb9f3851a77..45d9cafffd1 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs @@ -34,7 +34,7 @@ async fn handle( Json(body): Json, ) -> Result { let extra_headers = forwarded_headers(&headers)?; - match service::run(&state.router, body, extra_headers) + match service::run(&state.router, state.loggers, body, extra_headers) .await .map_err(MessagesRouteError::from)? { diff --git a/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs b/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs index 5434719987b..78d471f47f9 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs @@ -2,10 +2,64 @@ use std::sync::Arc; 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, +}; +use litellm_core::messages::lifecycle::{self, Options}; +use litellm_core::messages::messages_stream; use litellm_core::messages::types::MessagesRequest; -use litellm_core::messages::{messages, messages_stream}; use litellm_core::router::Router; use serde_json::{Map, Value}; +use std::future::{Ready, ready}; + +pub(crate) struct GatewayTerminalDispatcher { + runner: CustomLoggerRunner, +} + +impl GatewayTerminalDispatcher { + pub(crate) fn new(loggers: Arc>>) -> Self { + Self { + runner: CustomLoggerRunner::new(loggers.as_ref().clone()), + } + } +} + +impl Clock for GatewayTerminalDispatcher { + 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 RequestPolicy for GatewayTerminalDispatcher { + type PreCallFuture<'a> = Ready>; + type DuringCallFuture<'a> = Ready>; + + fn async_pre_call_hook<'a>( + &'a self, + _: &'a CallLifecycleContext, + request: MessagesRequest, + ) -> Self::PreCallFuture<'a> { + ready(ActionResult::Continue(request)) + } + + fn async_during_call_hook<'a>( + &'a self, + _: &'a CallLifecycleContext, + request: MessagesRequest, + ) -> Self::DuringCallFuture<'a> { + ready(ActionResult::Continue(request)) + } +} + +impl TerminalDispatcher for GatewayTerminalDispatcher { + fn dispatch<'a>(&'a self, terminal: &'a TerminalRecord) -> LogFuture<'a> { + self.runner.dispatch(terminal) + } +} pub(crate) enum MessagesResponse { Json(Value), @@ -20,6 +74,7 @@ pub(crate) enum MessagesResponse { )] pub async fn run( router: &Arc, + loggers: Arc>>, body: Value, extra_headers: Option>, ) -> Result { @@ -50,11 +105,11 @@ pub async fn run( ); let request = MessagesRequest { - model: provider_model, + model: provider_model.to_string(), body, - api_key: deployment.litellm_params.api_key.as_deref(), - api_base: deployment.litellm_params.api_base.as_deref(), - custom_llm_provider, + api_key: deployment.litellm_params.api_key.clone(), + api_base: deployment.litellm_params.api_base.clone(), + custom_llm_provider: custom_llm_provider.map(str::to_string), extra_headers, timeout: None, }; @@ -62,7 +117,22 @@ pub async fn run( return messages_stream(request).await.map(MessagesResponse::Stream); } - let response = messages(request).await?; + let services = GatewayTerminalDispatcher::new(loggers); + let context = CallLifecycleContext::new( + "messages", + provider_model, + custom_llm_provider.unwrap_or(ANTHROPIC_MESSAGES_PROVIDER), + format!( + "messages-{}", + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .map(|duration| duration.as_nanos()) + .unwrap_or(0) + ), + ); + let response = lifecycle::messages(&services, request, Options::default(), context) + .await + .into_result()?; serde_json::to_value(response) .map(MessagesResponse::Json) .map_err(|err| { diff --git a/litellm-rust/crates/ai-gateway/src/routes/responses/service.rs b/litellm-rust/crates/ai-gateway/src/routes/responses/service.rs index e525e80cbf8..e85fa16135b 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/responses/service.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/responses/service.rs @@ -3,7 +3,7 @@ use std::time::Duration; use futures_util::{Sink, Stream}; use litellm_core::Error; -use litellm_core::call_lifecycle::{CallLifecycle, CallLifecycleContext}; +use litellm_core::lifecycle::{CallLifecycle, CallLifecycleContext}; use litellm_core::responses::instrumentation::{ ResponsesWsCallbackPayload, ResponsesWsInstrumentation, ResponsesWsLogOutcome, ResponsesWsMetadata, @@ -57,7 +57,7 @@ where let observer_instrumentation = Arc::clone(&instrumentation); let context = CallLifecycleContext::new("responses_websocket", model, "openai", call_id); let result = CallLifecycle::default() - .run(context, (), instrumentation.as_ref(), |_| async move { + .run_result(context, (), instrumentation.as_ref(), |_| async move { crate::io::responses_ws::async_responses_websocket( provider_model, params.api_key.as_deref(), diff --git a/litellm-rust/crates/core/src/call_lifecycle/mod.rs b/litellm-rust/crates/core/src/call_lifecycle/mod.rs deleted file mode 100644 index cd84c05a510..00000000000 --- a/litellm-rust/crates/core/src/call_lifecycle/mod.rs +++ /dev/null @@ -1,463 +0,0 @@ -use std::future::Future; -use std::time::{Instant, SystemTime, UNIX_EPOCH}; - -use crate::Error; -use crate::lifecycle::ActionResult; - -pub mod types; - -pub use types::{ - CallLifecycleContext, CallLifecyclePhase, CallLifecyclePhaseTiming, CallLifecycleRequest, - CallLifecycleTiming, -}; - -pub trait RequestPolicy: Send + Sync { - type PreCallFuture<'a>: Future> + Send + 'a - where - Self: 'a, - InitialReq: 'a, - ProviderReq: 'a; - - type DuringCallFuture<'a>: Future> + Send + 'a - where - Self: 'a, - InitialReq: 'a, - ProviderReq: 'a; - - fn async_pre_call_hook<'a>( - &'a self, - context: &'a CallLifecycleContext, - request: InitialReq, - ) -> Self::PreCallFuture<'a>; - - fn async_during_call_hook<'a>( - &'a self, - context: &'a CallLifecycleContext, - request: InitialReq, - ) -> Self::DuringCallFuture<'a>; -} - -pub trait TerminalDispatcher: Send + Sync { - fn dispatch<'a>( - &'a self, - terminal: &'a crate::lifecycle::TerminalRecord, - ) -> crate::integrations::custom_logger::LogFuture<'a>; -} - -pub trait CallLifecycleHooks: Send + Sync { - type PreCallFuture<'a>: Future> + Send + 'a - where - Self: 'a, - InitialReq: 'a, - ProviderReq: 'a, - Resp: 'a; - - type DuringCallFuture<'a>: Future> + Send + 'a - where - Self: 'a, - InitialReq: 'a, - ProviderReq: 'a, - Resp: 'a; - - type SuccessFuture<'a>: Future + Send + 'a - where - Self: 'a, - Resp: 'a; - - type FailureFuture<'a>: Future + Send + 'a - where - Self: 'a; - - fn async_pre_call_hook<'a>( - &'a self, - context: &'a CallLifecycleContext, - request: InitialReq, - ) -> Self::PreCallFuture<'a>; - - fn async_during_call_hook<'a>( - &'a self, - context: &'a CallLifecycleContext, - request: InitialReq, - ) -> Self::DuringCallFuture<'a>; - - fn async_log_success_event<'a>( - &'a self, - context: &'a CallLifecycleContext, - response: &'a Resp, - timing: &'a CallLifecycleTiming, - ) -> Self::SuccessFuture<'a>; - - fn async_log_failure_event<'a>( - &'a self, - context: &'a CallLifecycleContext, - error: &'a Error, - timing: &'a CallLifecycleTiming, - ) -> Self::FailureFuture<'a>; -} - -pub trait CallLifecycleObserver: Send + Sync { - fn on_phase_start(&self, _context: &CallLifecycleContext, _phase: CallLifecyclePhase) {} - - fn on_phase_end(&self, _context: &CallLifecycleContext, _timing: &CallLifecyclePhaseTiming) {} -} - -#[derive(Default)] -pub struct NoopCallLifecycleObserver; - -impl CallLifecycleObserver for NoopCallLifecycleObserver {} - -pub struct CallLifecycle<'a> { - observer: &'a dyn CallLifecycleObserver, -} - -impl<'a> CallLifecycle<'a> { - pub fn new(observer: &'a dyn CallLifecycleObserver) -> Self { - Self { observer } - } - - pub async fn run_request( - &self, - request: InitialReq, - hooks: &Hooks, - provider_call: ProviderCall, - ) -> Result - where - InitialReq: CallLifecycleRequest, - Hooks: CallLifecycleHooks, - ProviderCall: FnOnce(ProviderReq) -> ProviderFuture, - ProviderFuture: Future>, - { - let context = request.lifecycle_context(); - self.run(context, request, hooks, provider_call).await - } - - pub async fn run_result( - &self, - context: CallLifecycleContext, - request: InitialReq, - hooks: &Hooks, - provider_call: ProviderCall, - ) -> Result - where - Hooks: CallLifecycleHooks, - ProviderCall: FnOnce(ProviderReq) -> ProviderFuture, - ProviderFuture: Future>, - { - self.run(context, request, hooks, provider_call).await - } - - pub async fn run( - &self, - context: CallLifecycleContext, - request: InitialReq, - hooks: &Hooks, - provider_call: ProviderCall, - ) -> Result - where - Hooks: CallLifecycleHooks, - ProviderCall: FnOnce(ProviderReq) -> ProviderFuture, - ProviderFuture: Future>, - { - let call_start = epoch_seconds(); - let mut phases = Vec::new(); - - let pre_call = self.start_phase(&context, CallLifecyclePhase::PreCall); - let request = match hooks.async_pre_call_hook(&context, request).await { - Ok(request) => { - phases.push(self.finish_phase(&context, pre_call)); - request - } - Err(error) => { - phases.push(self.finish_phase(&context, pre_call)); - self.log_failure(&context, hooks, &error, call_start, &mut phases) - .await; - return Err(error); - } - }; - - let during_call = self.start_phase(&context, CallLifecyclePhase::DuringCall); - let provider_request = match hooks.async_during_call_hook(&context, request).await { - Ok(request) => { - phases.push(self.finish_phase(&context, during_call)); - request - } - Err(error) => { - phases.push(self.finish_phase(&context, during_call)); - self.log_failure(&context, hooks, &error, call_start, &mut phases) - .await; - return Err(error); - } - }; - - let provider_phase = self.start_phase(&context, CallLifecyclePhase::ProviderCall); - let result = provider_call(provider_request).await; - phases.push(self.finish_phase(&context, provider_phase)); - - match &result { - Ok(response) => { - let success_phase = self.start_phase(&context, CallLifecyclePhase::SuccessCallback); - let timing = CallLifecycleTiming::new(call_start, epoch_seconds(), phases.clone()); - hooks - .async_log_success_event(&context, response, &timing) - .await; - phases.push(self.finish_phase(&context, success_phase)); - } - Err(error) => { - self.log_failure(&context, hooks, error, call_start, &mut phases) - .await; - } - } - - result - } - - async fn log_failure( - &self, - context: &CallLifecycleContext, - hooks: &Hooks, - error: &Error, - call_start: f64, - phases: &mut Vec, - ) where - Hooks: CallLifecycleHooks, - { - let failure_phase = self.start_phase(context, CallLifecyclePhase::FailureCallback); - let timing = CallLifecycleTiming::new(call_start, epoch_seconds(), phases.clone()); - hooks.async_log_failure_event(context, error, &timing).await; - phases.push(self.finish_phase(context, failure_phase)); - } - - fn start_phase(&self, context: &CallLifecycleContext, phase: CallLifecyclePhase) -> PhaseStart { - self.observer.on_phase_start(context, phase); - PhaseStart { - phase, - start_time: epoch_seconds(), - started_at: Instant::now(), - } - } - - fn finish_phase( - &self, - context: &CallLifecycleContext, - phase_start: PhaseStart, - ) -> CallLifecyclePhaseTiming { - let timing = CallLifecyclePhaseTiming { - phase: phase_start.phase, - start_time: phase_start.start_time, - end_time: epoch_seconds(), - duration: phase_start.started_at.elapsed(), - }; - self.observer.on_phase_end(context, &timing); - timing - } -} - -impl Default for CallLifecycle<'static> { - fn default() -> Self { - static OBSERVER: NoopCallLifecycleObserver = NoopCallLifecycleObserver; - Self::new(&OBSERVER) - } -} - -struct PhaseStart { - phase: CallLifecyclePhase, - start_time: f64, - started_at: Instant, -} - -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 std::pin::Pin; - use std::sync::Mutex; - - type BoxFuture<'a, T> = Pin + Send + 'a>>; - - #[derive(Default)] - struct RecordingHooks { - events: Mutex>, - } - - struct RecordingRequest(String); - - impl CallLifecycleRequest for RecordingRequest { - fn lifecycle_context(&self) -> CallLifecycleContext { - CallLifecycleContext::new("ocr", "mistral-ocr-latest", "mistral", "call_1") - } - } - - impl RecordingHooks { - fn events(&self) -> Vec<&'static str> { - self.events.lock().unwrap().clone() - } - } - - impl CallLifecycleHooks for RecordingHooks { - type PreCallFuture<'a> = BoxFuture<'a, Result>; - type DuringCallFuture<'a> = BoxFuture<'a, Result>; - type SuccessFuture<'a> = BoxFuture<'a, ()>; - type FailureFuture<'a> = BoxFuture<'a, ()>; - - fn async_pre_call_hook<'a>( - &'a self, - _context: &'a CallLifecycleContext, - request: String, - ) -> Self::PreCallFuture<'a> { - Box::pin(async move { - self.events.lock().unwrap().push("pre_call"); - Ok(format!("{request}:pre")) - }) - } - - fn async_during_call_hook<'a>( - &'a self, - _context: &'a CallLifecycleContext, - request: String, - ) -> Self::DuringCallFuture<'a> { - Box::pin(async move { - self.events.lock().unwrap().push("during_call"); - Ok(format!("{request}:during")) - }) - } - - fn async_log_success_event<'a>( - &'a self, - _context: &'a CallLifecycleContext, - _response: &'a String, - timing: &'a CallLifecycleTiming, - ) -> Self::SuccessFuture<'a> { - Box::pin(async move { - assert!(timing.end_time >= timing.start_time); - assert_eq!(timing.phases.len(), 3); - self.events.lock().unwrap().push("success"); - }) - } - - fn async_log_failure_event<'a>( - &'a self, - _context: &'a CallLifecycleContext, - _error: &'a Error, - _timing: &'a CallLifecycleTiming, - ) -> Self::FailureFuture<'a> { - Box::pin(async move { - self.events.lock().unwrap().push("failure"); - }) - } - } - - impl CallLifecycleHooks for RecordingHooks { - type PreCallFuture<'a> = BoxFuture<'a, Result>; - type DuringCallFuture<'a> = BoxFuture<'a, Result>; - type SuccessFuture<'a> = BoxFuture<'a, ()>; - type FailureFuture<'a> = BoxFuture<'a, ()>; - - fn async_pre_call_hook<'a>( - &'a self, - _context: &'a CallLifecycleContext, - request: RecordingRequest, - ) -> Self::PreCallFuture<'a> { - Box::pin(async move { - self.events.lock().unwrap().push("pre_call"); - Ok(RecordingRequest(format!("{}:pre", request.0))) - }) - } - - fn async_during_call_hook<'a>( - &'a self, - _context: &'a CallLifecycleContext, - request: RecordingRequest, - ) -> Self::DuringCallFuture<'a> { - Box::pin(async move { - self.events.lock().unwrap().push("during_call"); - Ok(format!("{}:during", request.0)) - }) - } - - fn async_log_success_event<'a>( - &'a self, - _context: &'a CallLifecycleContext, - _response: &'a String, - _timing: &'a CallLifecycleTiming, - ) -> Self::SuccessFuture<'a> { - Box::pin(async move { - self.events.lock().unwrap().push("success"); - }) - } - - fn async_log_failure_event<'a>( - &'a self, - _context: &'a CallLifecycleContext, - _error: &'a Error, - _timing: &'a CallLifecycleTiming, - ) -> Self::FailureFuture<'a> { - Box::pin(async move { - self.events.lock().unwrap().push("failure"); - }) - } - } - - #[tokio::test] - async fn lifecycle_runs_hooks_around_provider_call() { - let hooks = RecordingHooks::default(); - let response = CallLifecycle::default() - .run( - CallLifecycleContext::new("ocr", "mistral-ocr-latest", "mistral", "call_1"), - "request".to_string(), - &hooks, - |request| async move { - assert_eq!(request, "request:pre:during"); - Ok("response".to_string()) - }, - ) - .await - .expect("call succeeds"); - - assert_eq!(response, "response"); - assert_eq!(hooks.events(), vec!["pre_call", "during_call", "success"]); - } - - #[tokio::test] - async fn lifecycle_logs_failure_when_provider_fails() { - let hooks = RecordingHooks::default(); - let error = CallLifecycle::default() - .run( - CallLifecycleContext::new("ocr", "mistral-ocr-latest", "mistral", "call_1"), - "request".to_string(), - &hooks, - |_request| async move { - Err::(Error::Network("provider down".to_string())) - }, - ) - .await - .expect_err("call fails"); - - assert_eq!(error, Error::Network("provider down".to_string())); - assert_eq!(hooks.events(), vec!["pre_call", "during_call", "failure"]); - } - - #[tokio::test] - async fn lifecycle_can_run_any_request_with_embedded_context() { - let hooks = RecordingHooks::default(); - let response = CallLifecycle::default() - .run_request( - RecordingRequest("request".to_string()), - &hooks, - |request| async move { - assert_eq!(request, "request:pre:during"); - Ok("response".to_string()) - }, - ) - .await - .expect("call succeeds"); - - assert_eq!(response, "response"); - assert_eq!(hooks.events(), vec!["pre_call", "during_call", "success"]); - } -} diff --git a/litellm-rust/crates/core/src/call_lifecycle/types.rs b/litellm-rust/crates/core/src/call_lifecycle/types.rs deleted file mode 100644 index 8819c8830d2..00000000000 --- a/litellm-rust/crates/core/src/call_lifecycle/types.rs +++ /dev/null @@ -1,75 +0,0 @@ -use std::time::Duration; - -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct CallLifecycleContext { - pub call_type: String, - pub model: String, - pub custom_llm_provider: String, - pub litellm_call_id: String, -} - -impl CallLifecycleContext { - pub fn new( - call_type: impl Into, - model: impl Into, - custom_llm_provider: impl Into, - litellm_call_id: impl Into, - ) -> Self { - Self { - call_type: call_type.into(), - model: model.into(), - custom_llm_provider: custom_llm_provider.into(), - litellm_call_id: litellm_call_id.into(), - } - } -} - -pub trait CallLifecycleRequest { - fn lifecycle_context(&self) -> CallLifecycleContext; -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum CallLifecyclePhase { - PreCall, - DuringCall, - ProviderCall, - SuccessCallback, - FailureCallback, -} - -impl CallLifecyclePhase { - pub fn as_str(self) -> &'static str { - match self { - Self::PreCall => "pre_call", - Self::DuringCall => "during_call", - Self::ProviderCall => "provider_call", - Self::SuccessCallback => "success_callback", - Self::FailureCallback => "failure_callback", - } - } -} - -#[derive(Clone, Copy, Debug, PartialEq)] -pub struct CallLifecyclePhaseTiming { - pub phase: CallLifecyclePhase, - pub start_time: f64, - pub end_time: f64, - pub duration: Duration, -} - -#[derive(Clone, Debug, PartialEq)] -pub struct CallLifecycleTiming { - pub start_time: f64, - pub end_time: f64, - pub phases: Vec, -} - -impl CallLifecycleTiming { - pub fn new(start_time: f64, end_time: f64, phases: Vec) -> Self { - Self { - start_time, - end_time, - phases, - } - } -} diff --git a/litellm-rust/crates/core/src/integrations/custom_logger/mod.rs b/litellm-rust/crates/core/src/integrations/custom_logger/mod.rs index 6430c7cbbff..9d38703dd57 100644 --- a/litellm-rust/crates/core/src/integrations/custom_logger/mod.rs +++ b/litellm-rust/crates/core/src/integrations/custom_logger/mod.rs @@ -96,18 +96,23 @@ impl CustomLoggerRunner { } } -impl crate::call_lifecycle::TerminalDispatcher for CustomLoggerRunner { +impl crate::lifecycle::TerminalDispatcher for CustomLoggerRunner { fn dispatch<'a>(&'a self, terminal: &'a crate::lifecycle::TerminalRecord) -> LogFuture<'a> { Box::pin(async move { let details = ModelCallDetails::from(terminal); let response = CallbackValue::new( - match &terminal.projection { - crate::lifecycle::RouteProjection::Ocr { .. } => "ocr", - crate::lifecycle::RouteProjection::Messages { .. } => "messages", - crate::lifecycle::RouteProjection::ChatCompletions { .. } => "chat_completion", - crate::lifecycle::RouteProjection::Audio { .. } => "audio", - crate::lifecycle::RouteProjection::Realtime { .. } => "realtime", - crate::lifecycle::RouteProjection::ResponsesWs { .. } => "responses_websocket", + match (&terminal.classification, &terminal.projection) { + (crate::lifecycle::TerminalClassification::Failure { .. }, _) => "error", + (_, crate::lifecycle::RouteProjection::Ocr { .. }) => "ocr", + (_, crate::lifecycle::RouteProjection::Messages { .. }) => "messages", + (_, crate::lifecycle::RouteProjection::ChatCompletions { .. }) => { + "chat_completion" + } + (_, crate::lifecycle::RouteProjection::Audio { .. }) => "audio_transcription", + (_, crate::lifecycle::RouteProjection::Realtime { .. }) => "realtime", + (_, crate::lifecycle::RouteProjection::ResponsesWs { .. }) => { + "responses_websocket" + } }, terminal.projection.value().clone(), ); @@ -335,8 +340,8 @@ mod tests { #[tokio::test] async fn terminal_dispatcher_fans_out_shared_record() { - use crate::call_lifecycle::TerminalDispatcher; use crate::integrations::types::Usage; + use crate::lifecycle::TerminalDispatcher; use crate::lifecycle::terminal::CostInputs; use crate::lifecycle::{RouteProjection, TerminalClassification, TerminalRecord}; diff --git a/litellm-rust/crates/core/src/lib.rs b/litellm-rust/crates/core/src/lib.rs index 2d9ebb836b2..89c3c05dfbb 100644 --- a/litellm-rust/crates/core/src/lib.rs +++ b/litellm-rust/crates/core/src/lib.rs @@ -1,6 +1,5 @@ pub mod audio_transcription; pub mod caching; -pub mod call_lifecycle; pub mod chat_completions; pub mod constants; pub mod error; diff --git a/litellm-rust/crates/core/src/lifecycle/execution.rs b/litellm-rust/crates/core/src/lifecycle/execution.rs new file mode 100644 index 00000000000..b4d58e8697e --- /dev/null +++ b/litellm-rust/crates/core/src/lifecycle/execution.rs @@ -0,0 +1,394 @@ +use std::future::Future; +use std::time::{SystemTime, UNIX_EPOCH}; + +use serde::Serialize; +use serde_json::{Value, json}; + +use crate::Error; +use crate::integrations::custom_logger::{CallbackTiming, LogFuture}; +use crate::integrations::types::{StandardLoggingMetadata, Usage}; + +use super::terminal::CostInputs; +use super::{ActionResult, ExecutedCall, RouteProjection, TerminalClassification, TerminalRecord}; + +#[derive(Clone, Debug, PartialEq)] +pub struct CallLifecycleContext { + pub call_type: String, + pub model: String, + pub custom_llm_provider: String, + pub litellm_call_id: String, + pub trace_id: Option, + pub attempt: u32, + pub usage: Usage, + pub response_cost: f64, + pub metadata: StandardLoggingMetadata, +} + +impl CallLifecycleContext { + pub fn new( + call_type: impl Into, + model: impl Into, + custom_llm_provider: impl Into, + litellm_call_id: impl Into, + ) -> Self { + Self { + call_type: call_type.into(), + model: model.into(), + custom_llm_provider: custom_llm_provider.into(), + litellm_call_id: litellm_call_id.into(), + trace_id: None, + attempt: 1, + usage: Usage::default(), + response_cost: 0.0, + metadata: StandardLoggingMetadata::default(), + } + } + + pub fn with_metadata(mut self, metadata: StandardLoggingMetadata) -> Self { + self.metadata = metadata; + self + } + + fn terminal( + &self, + timing: CallbackTiming, + classification: TerminalClassification, + value: Value, + ) -> TerminalRecord { + TerminalRecord { + call_id: self.litellm_call_id.clone(), + trace_id: self.trace_id.clone(), + attempt: self.attempt, + call_type: self.call_type.clone(), + model: self.model.clone(), + provider: self.custom_llm_provider.clone(), + timing, + usage: self.usage, + cost_inputs: CostInputs { + response_cost: self.response_cost, + metadata: self.metadata.clone(), + }, + classification, + projection: projection(&self.call_type, value), + } + } +} + +pub trait CallLifecycleRequest { + fn lifecycle_context(&self) -> CallLifecycleContext; +} + +pub trait RequestPolicy: Send + Sync { + type PreCallFuture<'a>: Future> + Send + where + Self: 'a; + + type DuringCallFuture<'a>: Future> + Send + where + Self: 'a; + + fn async_pre_call_hook<'a>( + &'a self, + context: &'a CallLifecycleContext, + request: InitialReq, + ) -> Self::PreCallFuture<'a>; + + fn async_during_call_hook<'a>( + &'a self, + context: &'a CallLifecycleContext, + request: InitialReq, + ) -> Self::DuringCallFuture<'a>; +} + +pub trait TerminalDispatcher: Send + Sync { + fn dispatch<'a>(&'a self, terminal: &'a TerminalRecord) -> LogFuture<'a>; +} + +pub trait Clock: Send + Sync { + fn now(&self) -> f64; +} + +#[derive(Default)] +pub struct SystemClock; + +impl Clock for SystemClock { + fn now(&self) -> f64 { + epoch_seconds() + } +} + +#[derive(Default)] +pub struct CallLifecycle; + +impl CallLifecycle { + pub async fn run( + &self, + context: CallLifecycleContext, + request: InitialReq, + policy: &Policy, + provider_call: ProviderCall, + ) -> ExecutedCall + where + Resp: Serialize, + Policy: RequestPolicy + TerminalDispatcher, + ProviderCall: FnOnce(ProviderReq) -> ProviderFuture, + ProviderFuture: Future>, + { + self.run_with_clock(context, request, policy, &SystemClock, provider_call) + .await + } + + pub async fn run_with_clock< + InitialReq, + ProviderReq, + Resp, + Policy, + ClockImpl, + ProviderCall, + ProviderFuture, + >( + &self, + context: CallLifecycleContext, + request: InitialReq, + policy: &Policy, + clock: &ClockImpl, + provider_call: ProviderCall, + ) -> ExecutedCall + where + Resp: Serialize, + Policy: RequestPolicy + TerminalDispatcher, + ClockImpl: Clock, + ProviderCall: FnOnce(ProviderReq) -> ProviderFuture, + ProviderFuture: Future>, + { + let start_time = clock.now(); + let request = match policy.async_pre_call_hook(&context, request).await { + ActionResult::Continue(request) | ActionResult::Replace(request) => request, + ActionResult::Reject(error) => { + return failure(policy, clock, &context, error, start_time).await; + } + }; + let provider_request = match policy.async_during_call_hook(&context, request).await { + ActionResult::Continue(request) | ActionResult::Replace(request) => request, + ActionResult::Reject(error) => { + return failure(policy, clock, &context, error, start_time).await; + } + }; + match provider_call(provider_request).await { + Ok(response) => { + let terminal = context.terminal( + CallbackTiming::new(start_time, clock.now()), + TerminalClassification::Success, + serde_json::to_value(&response).unwrap_or(Value::Null), + ); + let _ = policy.dispatch(&terminal).await; + ExecutedCall::Success { response, terminal } + } + Err(error) => failure(policy, clock, &context, error, start_time).await, + } + } + + pub async fn run_result( + &self, + context: CallLifecycleContext, + request: InitialReq, + policy: &Policy, + provider_call: ProviderCall, + ) -> Result + where + Resp: Serialize, + Policy: RequestPolicy + TerminalDispatcher, + ProviderCall: FnOnce(ProviderReq) -> ProviderFuture, + ProviderFuture: Future>, + { + self.run(context, request, policy, provider_call) + .await + .into_result() + } + + pub async fn run_request_result< + InitialReq, + ProviderReq, + Resp, + Policy, + ProviderCall, + ProviderFuture, + >( + &self, + request: InitialReq, + policy: &Policy, + provider_call: ProviderCall, + ) -> Result + where + InitialReq: CallLifecycleRequest, + Resp: Serialize, + Policy: RequestPolicy + TerminalDispatcher, + ProviderCall: FnOnce(ProviderReq) -> ProviderFuture, + ProviderFuture: Future>, + { + let context = request.lifecycle_context(); + self.run_result(context, request, policy, provider_call) + .await + } +} + +async fn failure( + policy: &Policy, + clock: &ClockImpl, + context: &CallLifecycleContext, + error: Error, + start_time: f64, +) -> ExecutedCall +where + Policy: TerminalDispatcher, + ClockImpl: Clock, +{ + let kind = error_kind(&error).to_string(); + let message = error.to_string(); + let terminal = context.terminal( + CallbackTiming::new(start_time, clock.now()), + TerminalClassification::Failure { + kind: kind.clone(), + message: message.clone(), + }, + json!({"message": message, "kind": kind}), + ); + let _ = policy.dispatch(&terminal).await; + ExecutedCall::Failure { error, terminal } +} + +fn projection(call_type: &str, value: Value) -> RouteProjection { + match call_type { + "ocr" => RouteProjection::Ocr { value }, + "messages" => RouteProjection::Messages { value }, + "chat_completion" | "acompletion" => RouteProjection::ChatCompletions { value }, + "audio_transcription" => RouteProjection::Audio { value }, + "realtime" => RouteProjection::Realtime { value }, + _ => RouteProjection::ResponsesWs { value }, + } +} + +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 epoch_seconds() -> f64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|duration| duration.as_secs_f64()) + .unwrap_or(0.0) +} + +#[cfg(test)] +mod tests { + use std::pin::Pin; + use std::sync::Mutex; + + use super::*; + + type PolicyFuture<'a, T> = Pin> + Send + 'a>>; + + #[derive(Default)] + struct RecordingPolicy { + terminals: Mutex>, + reject: bool, + } + + impl RequestPolicy for RecordingPolicy { + type PreCallFuture<'a> = PolicyFuture<'a, String>; + type DuringCallFuture<'a> = PolicyFuture<'a, String>; + + fn async_pre_call_hook<'a>( + &'a self, + _context: &'a CallLifecycleContext, + request: String, + ) -> Self::PreCallFuture<'a> { + Box::pin(async move { + if self.reject { + ActionResult::Reject(Error::InvalidRequest("blocked".to_string())) + } else { + ActionResult::Replace(format!("{request}:pre")) + } + }) + } + + fn async_during_call_hook<'a>( + &'a self, + _context: &'a CallLifecycleContext, + request: String, + ) -> Self::DuringCallFuture<'a> { + Box::pin(async move { ActionResult::Replace(format!("{request}:during")) }) + } + } + + impl TerminalDispatcher for RecordingPolicy { + fn dispatch<'a>(&'a self, terminal: &'a TerminalRecord) -> LogFuture<'a> { + Box::pin(async move { + self.terminals.lock().unwrap().push(terminal.clone()); + Ok(()) + }) + } + } + + #[tokio::test] + async fn run_applies_replacements_and_returns_terminal_record() { + let policy = RecordingPolicy::default(); + let executed = CallLifecycle + .run( + CallLifecycleContext::new("ocr", "model", "provider", "call-1"), + "request".to_string(), + &policy, + |request| async move { + assert_eq!(request, "request:pre:during"); + Ok(request) + }, + ) + .await; + + assert!(matches!( + executed, + ExecutedCall::Success { ref terminal, .. } if terminal.call_id == "call-1" + )); + assert_eq!(policy.terminals.lock().unwrap().len(), 1); + } + + #[tokio::test] + async fn rejection_dispatches_and_retains_failure_record() { + let policy = RecordingPolicy { + reject: true, + ..Default::default() + }; + let executed = CallLifecycle + .run( + CallLifecycleContext::new("ocr", "model", "provider", "call-2"), + "request".to_string(), + &policy, + |request| async move { Ok(request) }, + ) + .await; + + assert!(matches!( + executed, + ExecutedCall::Failure { + error: Error::InvalidRequest(_), + ref terminal, + } if terminal.call_id == "call-2" + )); + assert!(matches!( + policy.terminals.lock().unwrap()[0].classification, + TerminalClassification::Failure { .. } + )); + } +} diff --git a/litellm-rust/crates/core/src/lifecycle/mod.rs b/litellm-rust/crates/core/src/lifecycle/mod.rs index 9292999759f..581822c4e6c 100644 --- a/litellm-rust/crates/core/src/lifecycle/mod.rs +++ b/litellm-rust/crates/core/src/lifecycle/mod.rs @@ -1,5 +1,6 @@ pub mod action; pub mod executed; +pub mod execution; pub mod machine; pub mod ocr; pub mod terminal; @@ -7,6 +8,10 @@ pub mod types; pub use action::{ActionBinding, Owner, ResultPolicy}; pub use executed::ExecutedCall; +pub use execution::{ + CallLifecycle, CallLifecycleContext, CallLifecycleRequest, Clock, RequestPolicy, SystemClock, + TerminalDispatcher, +}; pub use machine::{Lifecycle, LifecycleRoute}; pub use terminal::{RouteProjection, TerminalClassification, TerminalRecord}; pub use types::{ActionKind, ActionResult, Delivery, ErrorDisposition, FailurePolicy, Outcome}; diff --git a/litellm-rust/crates/core/src/lifecycle/terminal.rs b/litellm-rust/crates/core/src/lifecycle/terminal.rs index 63c2869254c..fb6d9c689a1 100644 --- a/litellm-rust/crates/core/src/lifecycle/terminal.rs +++ b/litellm-rust/crates/core/src/lifecycle/terminal.rs @@ -39,6 +39,16 @@ impl RouteProjection { | Self::ResponsesWs { value } => value, } } + + fn logging_input(&self) -> Option { + match self { + Self::Messages { value } | Self::ChatCompletions { value } => Some(value.clone()), + Self::Ocr { .. } + | Self::Audio { .. } + | Self::Realtime { .. } + | Self::ResponsesWs { .. } => None, + } + } } #[derive(Clone, Debug, PartialEq, Serialize)] @@ -59,10 +69,7 @@ pub struct TerminalRecord { impl From<&TerminalRecord> for StandardLoggingPayload { fn from(record: &TerminalRecord) -> Self { Self { - id: record - .trace_id - .clone() - .unwrap_or_else(|| record.call_id.clone()), + id: record.call_id.clone(), litellm_call_id: record.call_id.clone(), call_type: record.call_type.clone(), model: record.model.clone(), @@ -78,7 +85,7 @@ impl From<&TerminalRecord> for StandardLoggingPayload { RouteProjection::Realtime { .. } | RouteProjection::ResponsesWs { .. } ), metadata: record.cost_inputs.metadata.clone(), - messages: Some(record.projection.value().clone()), + messages: record.projection.logging_input(), } } } @@ -143,11 +150,11 @@ mod tests { let payload = StandardLoggingPayload::from(&record); let details = ModelCallDetails::from(&record); - assert_eq!(payload.id, "trace-1"); + assert_eq!(payload.id, "call-1"); assert_eq!(payload.litellm_call_id, "call-1"); assert_eq!(payload.prompt_tokens, 3); assert_eq!(payload.response_cost, 0.25); - assert_eq!(payload.messages, Some(record.projection.value().clone())); + assert_eq!(payload.messages, None); assert_eq!( details.standard_logging_payload.unwrap().litellm_call_id, "call-1" diff --git a/litellm-rust/crates/core/src/messages/handler.rs b/litellm-rust/crates/core/src/messages/handler.rs index 61ff81bcdc8..fe73c49670c 100644 --- a/litellm-rust/crates/core/src/messages/handler.rs +++ b/litellm-rust/crates/core/src/messages/handler.rs @@ -9,7 +9,7 @@ use super::types::{AnthropicMessagesResponse, MessagesRequest}; #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] pub(super) async fn execute_messages_provider_call( - request: MessagesRequest<'_>, + request: MessagesRequest, ) -> Result { let request = prepare_provider_request(request)?; let mut request_builder = http_client().post(&request.url).json(&request.body); @@ -43,7 +43,7 @@ pub(super) async fn execute_messages_provider_call( } pub(super) async fn execute_messages_provider_stream( - request: MessagesRequest<'_>, + request: MessagesRequest, ) -> Result { let request = prepare_provider_request(request)?; if request.provider != ANTHROPIC_MESSAGES_PROVIDER { diff --git a/litellm-rust/crates/core/src/messages/lifecycle.rs b/litellm-rust/crates/core/src/messages/lifecycle.rs new file mode 100644 index 00000000000..63210bf3f24 --- /dev/null +++ b/litellm-rust/crates/core/src/messages/lifecycle.rs @@ -0,0 +1,252 @@ +use std::future::{Ready, ready}; + +use crate::Error; +use crate::integrations::custom_logger::{LogError, LogFuture}; +use crate::lifecycle::{ + ActionBinding, ActionKind, ActionResult, CallLifecycle, CallLifecycleContext, Clock, Delivery, + ErrorDisposition, ExecutedCall, FailurePolicy, Lifecycle, LifecycleRoute, Outcome, Owner, + RequestPolicy, ResultPolicy, TerminalDispatcher, TerminalRecord, +}; + +use super::handler::execute_messages_provider_call; +use super::types::{AnthropicMessagesResponse, MessagesRequest}; + +#[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, Default)] +pub struct Options { + pub asynchronous: bool, + pub internal_call: bool, + pub call_id: Option, + pub trace_id: Option, +} + +#[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 operation: Operation, + pub error: ErrorDisposition, +} + +#[derive(Debug)] +pub struct MessagesState { + operation: Operation, + outcome: Outcome, + asynchronous: bool, + internal_call: bool, +} + +#[derive(Debug)] +pub struct MessagesRoute; + +impl LifecycleRoute for MessagesRoute { + type Admission = (); + type Options = Options; + type Context = Observations; + type Operation = Operation; + type Observation = Observations; + type Outcome = Outcome; + type Transition = Transition; + type Error = Error; + type Decline = std::convert::Infallible; + type State = MessagesState; + + fn admit(_: &(), options: Options) -> Result, Error> { + Ok(Ok(MessagesState { + operation: Operation::Setup, + outcome: Outcome::Success, + asynchronous: options.asynchronous, + internal_call: options.internal_call, + })) + } + + fn operation(state: &Self::State) -> Operation { + state.operation + } + + fn advance( + state: &mut Self::State, + outcome: Outcome, + observations: Observations, + ) -> Result { + use Operation::*; + + if matches!(state.operation, Complete(_)) { + return Err(Error::InvalidRequest( + "messages 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 { + operation: state.operation, + 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 trait MessagesServices: + RequestPolicy + TerminalDispatcher + Clock +{ +} + +impl MessagesServices for T where + T: RequestPolicy + TerminalDispatcher + Clock +{ +} + +#[derive(Default)] +pub struct NoopServices; + +impl Clock for NoopServices { + 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 RequestPolicy for NoopServices { + type PreCallFuture<'a> = Ready>; + type DuringCallFuture<'a> = Ready>; + + fn async_pre_call_hook<'a>( + &'a self, + _: &'a CallLifecycleContext, + request: MessagesRequest, + ) -> Self::PreCallFuture<'a> { + ready(ActionResult::Continue(request)) + } + + fn async_during_call_hook<'a>( + &'a self, + _: &'a CallLifecycleContext, + request: MessagesRequest, + ) -> Self::DuringCallFuture<'a> { + ready(ActionResult::Continue(request)) + } +} + +impl TerminalDispatcher for NoopServices { + fn dispatch<'a>(&'a self, _: &'a TerminalRecord) -> LogFuture<'a> { + Box::pin(async { Ok::<(), LogError>(()) }) + } +} + +pub async fn messages( + services: &S, + request: MessagesRequest, + _options: Options, + context: CallLifecycleContext, +) -> ExecutedCall { + CallLifecycle + .run_with_clock(context, request, services, services, |request| async move { + execute_messages_provider_call(request).await + }) + .await +} + +pub fn machine(options: Options) -> Result, Error> { + Lifecycle::admit(&(), options).map(|result| match result { + Ok(machine) => machine, + Err(never) => match never {}, + }) +} diff --git a/litellm-rust/crates/core/src/messages/mod.rs b/litellm-rust/crates/core/src/messages/mod.rs index cfa8bda1104..a703e7efe9c 100644 --- a/litellm-rust/crates/core/src/messages/mod.rs +++ b/litellm-rust/crates/core/src/messages/mod.rs @@ -8,22 +8,41 @@ //! can splice the event stream to its own caller. use crate::Error; +use crate::constants::ANTHROPIC_MESSAGES_PROVIDER; mod client; mod common_utils; mod handler; +pub mod lifecycle; mod prepare; pub mod transformation; pub mod types; -use handler::{execute_messages_provider_call, execute_messages_provider_stream}; +use handler::execute_messages_provider_stream; use types::{AnthropicMessagesResponse, MessagesRequest}; -#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] -pub async fn messages(request: MessagesRequest<'_>) -> Result { - execute_messages_provider_call(request).await +pub async fn messages(request: MessagesRequest) -> Result { + let provider = request + .custom_llm_provider + .as_deref() + .or_else(|| request.model.split_once('/').map(|(provider, _)| provider)) + .unwrap_or(ANTHROPIC_MESSAGES_PROVIDER); + let context = crate::lifecycle::CallLifecycleContext::new( + "messages", + &request.model, + provider, + format!("{:032x}", rand::random::()), + ); + lifecycle::messages( + &lifecycle::NoopServices, + request, + lifecycle::Options::default(), + context, + ) + .await + .into_result() } -pub async fn messages_stream(request: MessagesRequest<'_>) -> Result { +pub async fn messages_stream(request: MessagesRequest) -> Result { execute_messages_provider_stream(request).await } diff --git a/litellm-rust/crates/core/src/messages/prepare.rs b/litellm-rust/crates/core/src/messages/prepare.rs index ec83d03f535..45569857882 100644 --- a/litellm-rust/crates/core/src/messages/prepare.rs +++ b/litellm-rust/crates/core/src/messages/prepare.rs @@ -7,22 +7,24 @@ use super::types::{MessagesRequest, ProviderMessagesRequest}; use serde_json::{Map, Value}; pub(super) fn prepare_provider_request( - request: MessagesRequest<'_>, + request: MessagesRequest, ) -> Result { - let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider) - .or_else(|| { - request - .custom_llm_provider - .map(|provider| CustomLlmProvider { - model: request.model, - custom_llm_provider: provider, - }) - }) - .ok_or_else(|| { - Error::InvalidProvider( - "unable to resolve custom_llm_provider for messages request".to_string(), - ) - })?; + let provider_info = + get_custom_llm_provider(&request.model, request.custom_llm_provider.as_deref()) + .or_else(|| { + request + .custom_llm_provider + .as_deref() + .map(|provider| CustomLlmProvider { + model: &request.model, + custom_llm_provider: provider, + }) + }) + .ok_or_else(|| { + Error::InvalidProvider( + "unable to resolve custom_llm_provider for messages request".to_string(), + ) + })?; let model = provider_info.model.to_string(); let provider = provider_info.custom_llm_provider; @@ -30,8 +32,12 @@ pub(super) fn prepare_provider_request( .ok_or_else(|| Error::InvalidProvider(provider.to_string()))?; let env_lookup = |key: &str| std::env::var(key).ok(); - let headers = - validate_environment(config, request.extra_headers, request.api_key, &env_lookup)?; + let headers = validate_environment( + config, + request.extra_headers, + request.api_key.as_deref(), + &env_lookup, + )?; let typed_request = serde_json::from_value(request.body).map_err(|err| { Error::InvalidRequest(format!("invalid Anthropic messages request: {err}")) @@ -43,7 +49,7 @@ pub(super) fn prepare_provider_request( )) })?; - let url = config.complete_url(request.api_base, &model, &env_lookup)?; + let url = config.complete_url(request.api_base.as_deref(), &model, &env_lookup)?; Ok(ProviderMessagesRequest { provider: provider.to_string(), diff --git a/litellm-rust/crates/core/src/messages/tests.rs b/litellm-rust/crates/core/src/messages/tests.rs index df9f7051011..38f0e20f7d3 100644 --- a/litellm-rust/crates/core/src/messages/tests.rs +++ b/litellm-rust/crates/core/src/messages/tests.rs @@ -132,7 +132,7 @@ async fn messages_round_trip_builds_azure_request_and_passes_response_through() }); let response = messages(MessagesRequest { - model: "claude-sonnet-4-5", + model: "claude-sonnet-4-5".into(), body: json!({ "model": "claude-sonnet-4-5", "max_tokens": 1024, @@ -145,9 +145,9 @@ async fn messages_round_trip_builds_azure_request_and_passes_response_through() }] }] }), - api_key: Some("sk-azure"), - api_base: Some(&format!("http://{addr}")), - custom_llm_provider: Some("azure_ai"), + api_key: Some("sk-azure".into()), + api_base: Some(format!("http://{addr}")), + custom_llm_provider: Some("azure_ai".into()), extra_headers: None, timeout: Some(Duration::from_secs(5)), }) @@ -195,15 +195,15 @@ async fn messages_round_trip_builds_native_anthropic_request() { }); let response = messages(MessagesRequest { - model: "claude-sonnet-4-5", + model: "claude-sonnet-4-5".into(), body: json!({ "model": "claude-sonnet-4-5", "max_tokens": 1024, "messages": [{"role": "user", "content": "hi"}] }), - api_key: Some("sk-ant"), - api_base: Some(&format!("http://{addr}")), - custom_llm_provider: Some("anthropic"), + api_key: Some("sk-ant".into()), + api_base: Some(format!("http://{addr}")), + custom_llm_provider: Some("anthropic".into()), extra_headers: None, timeout: Some(Duration::from_secs(5)), }) @@ -252,11 +252,11 @@ async fn messages_does_not_duplicate_auth_when_x_api_key_supplied() { ); messages(MessagesRequest { - model: "claude-sonnet-4-5", + model: "claude-sonnet-4-5".into(), body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}), - api_key: Some("rust-fallback-key"), - api_base: Some(&format!("http://{addr}")), - custom_llm_provider: Some("azure_ai"), + api_key: Some("rust-fallback-key".into()), + api_base: Some(format!("http://{addr}")), + custom_llm_provider: Some("azure_ai".into()), extra_headers: Some(headers), timeout: Some(Duration::from_secs(5)), }) @@ -306,11 +306,11 @@ async fn messages_forwards_entra_id_bearer_without_requiring_api_key() { ); messages(MessagesRequest { - model: "claude-sonnet-4-5", + model: "claude-sonnet-4-5".into(), body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}), api_key: None, - api_base: Some(&format!("http://{addr}")), - custom_llm_provider: Some("azure_ai"), + api_base: Some(format!("http://{addr}")), + custom_llm_provider: Some("azure_ai".into()), extra_headers: Some(headers), timeout: Some(Duration::from_secs(5)), }) @@ -330,11 +330,11 @@ async fn messages_forwards_entra_id_bearer_without_requiring_api_key() { #[tokio::test] async fn messages_requires_auth_when_no_key_and_no_header() { let err = messages(MessagesRequest { - model: "claude-sonnet-4-5", + model: "claude-sonnet-4-5".into(), body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}), api_key: None, - api_base: Some("http://127.0.0.1:1"), - custom_llm_provider: Some("azure_ai"), + api_base: Some("http://127.0.0.1:1".into()), + custom_llm_provider: Some("azure_ai".into()), extra_headers: None, timeout: Some(Duration::from_millis(50)), }) @@ -368,11 +368,11 @@ async fn messages_ignores_malformed_authorization_and_uses_api_key() { ); messages(MessagesRequest { - model: "claude-sonnet-4-5", + model: "claude-sonnet-4-5".into(), body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}), - api_key: Some("sk-azure"), - api_base: Some(&format!("http://{addr}")), - custom_llm_provider: Some("azure_ai"), + api_key: Some("sk-azure".into()), + api_base: Some(format!("http://{addr}")), + custom_llm_provider: Some("azure_ai".into()), extra_headers: Some(headers), timeout: Some(Duration::from_secs(5)), }) @@ -409,11 +409,11 @@ async fn messages_maps_provider_error_status_to_http_error() { }); let err = messages(MessagesRequest { - model: "claude-sonnet-4-5", + model: "claude-sonnet-4-5".into(), body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}), - api_key: Some("sk-azure"), - api_base: Some(&format!("http://{addr}")), - custom_llm_provider: Some("azure_ai"), + api_key: Some("sk-azure".into()), + api_base: Some(format!("http://{addr}")), + custom_llm_provider: Some("azure_ai".into()), extra_headers: None, timeout: Some(Duration::from_secs(5)), }) @@ -426,11 +426,11 @@ async fn messages_maps_provider_error_status_to_http_error() { #[tokio::test] async fn messages_rejects_unsupported_provider() { let err = messages(MessagesRequest { - model: "claude-3-5-sonnet", + model: "claude-3-5-sonnet".into(), body: json!({"model": "claude-3-5-sonnet", "max_tokens": 8, "messages": []}), - api_key: Some("sk"), - api_base: Some("http://127.0.0.1:1"), - custom_llm_provider: Some("openai"), + api_key: Some("sk".into()), + api_base: Some("http://127.0.0.1:1".into()), + custom_llm_provider: Some("openai".into()), extra_headers: None, timeout: Some(Duration::from_millis(50)), }) diff --git a/litellm-rust/crates/core/src/messages/types.rs b/litellm-rust/crates/core/src/messages/types.rs index b9f807c29fd..f55a9c9dfa7 100644 --- a/litellm-rust/crates/core/src/messages/types.rs +++ b/litellm-rust/crates/core/src/messages/types.rs @@ -5,12 +5,12 @@ use serde_json::{Map, Value}; use super::transformation::AnthropicMessagesProviderConfig; -pub struct MessagesRequest<'a> { - pub model: &'a str, +pub struct MessagesRequest { + pub model: String, pub body: Value, - pub api_key: Option<&'a str>, - pub api_base: Option<&'a str>, - pub custom_llm_provider: Option<&'a str>, + pub api_key: Option, + pub api_base: Option, + pub custom_llm_provider: Option, pub extra_headers: Option>, pub timeout: Option, } diff --git a/litellm-rust/crates/core/src/ocr/hooks.rs b/litellm-rust/crates/core/src/ocr/hooks.rs index e385a0de775..99e6b7ea8a2 100644 --- a/litellm-rust/crates/core/src/ocr/hooks.rs +++ b/litellm-rust/crates/core/src/ocr/hooks.rs @@ -1,5 +1,7 @@ -use crate::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming}; use crate::error::Error; +use crate::lifecycle::{ + ActionResult, CallLifecycleContext, RequestPolicy, TerminalDispatcher, TerminalRecord, +}; use crate::providers::reducto::ocr::transformation::{ build_upload_request, extract_document_source, extract_upload_file_id, }; @@ -13,12 +15,8 @@ use super::runtime_types::{PreparedOcrRequest, ProviderOcrRequest}; use crate::integrations::custom_guardrail::{ CustomGuardrailRunner, GuardrailContext, GuardrailError, GuardrailRequest, }; -use crate::integrations::custom_logger::{ - CallType, CallbackTiming, CallbackValue, CustomLoggerRunner, LoggingError, ModelCallDetails, -}; -use crate::integrations::types::{ - RequestMetadata, StandardLoggingMetadata, StandardLoggingPayload, -}; +use crate::integrations::custom_logger::{CallType, CustomLoggerRunner, LogFuture}; +use crate::integrations::types::RequestMetadata; pub(crate) struct OcrLifecycleHooks { logger_runner: CustomLoggerRunner, @@ -26,8 +24,7 @@ pub(crate) struct OcrLifecycleHooks { request_metadata: RequestMetadata, } -type OcrFuture<'a, T> = Pin> + Send + 'a>>; -type OcrLogFuture<'a> = Pin + Send + 'a>>; +type OcrFuture<'a, T> = Pin> + Send + 'a>>; impl OcrLifecycleHooks { pub(crate) fn new( @@ -157,33 +154,6 @@ impl OcrLifecycleHooks { parse_ocr_during_call_guardrail_request(guardrail_request) } - fn standard_logging_payload( - &self, - context: &CallLifecycleContext, - timing: &CallLifecycleTiming, - ) -> StandardLoggingPayload { - StandardLoggingPayload { - id: context.litellm_call_id.clone(), - litellm_call_id: context.litellm_call_id.clone(), - call_type: context.call_type.clone(), - model: context.model.clone(), - custom_llm_provider: context.custom_llm_provider.clone(), - response_cost: 0.0, - prompt_tokens: 0, - completion_tokens: 0, - total_tokens: 0, - start_time: timing.start_time, - end_time: timing.end_time, - stream: false, - metadata: StandardLoggingMetadata { - user_api_key_hash: self.request_metadata.user_api_key_hash.clone(), - user_api_key_user_id: self.request_metadata.user_api_key_user_id.clone(), - user_api_key_team_id: self.request_metadata.user_api_key_team_id.clone(), - ..Default::default() - }, - messages: None, - } - } } async fn upload_reducto_document( @@ -243,18 +213,21 @@ async fn upload_reducto_document( Ok(json!({"type": "document_url", "document_url": file_id})) } -impl CallLifecycleHooks for OcrLifecycleHooks { +impl RequestPolicy for OcrLifecycleHooks { type PreCallFuture<'a> = OcrFuture<'a, PreparedOcrRequest>; type DuringCallFuture<'a> = OcrFuture<'a, PreparedOcrRequest>; - type SuccessFuture<'a> = OcrLogFuture<'a>; - type FailureFuture<'a> = OcrLogFuture<'a>; fn async_pre_call_hook<'a>( &'a self, _context: &'a CallLifecycleContext, request: PreparedOcrRequest, ) -> Self::PreCallFuture<'a> { - Box::pin(async move { self.run_pre_call_guardrails(request).await }) + 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>( @@ -262,76 +235,24 @@ impl CallLifecycleHooks for OcrLi _context: &'a CallLifecycleContext, request: PreparedOcrRequest, ) -> Self::DuringCallFuture<'a> { - Box::pin(async move { Ok(request) }) + Box::pin(async move { ActionResult::Continue(request) }) } +} - #[tracing::instrument( - name = "success_callback", - target = "litellm::function_trace", - level = "trace", - skip_all - )] - fn async_log_success_event<'a>( - &'a self, - context: &'a CallLifecycleContext, - response: &'a Value, - timing: &'a CallLifecycleTiming, - ) -> Self::SuccessFuture<'a> { - Box::pin(async move { - if self.logger_runner.is_empty() { - return; - } - let response_obj = CallbackValue::new("ocr", response.clone()); - self.logger_runner - .async_log_success_event( - &ModelCallDetails::from_standard_logging_payload( - self.standard_logging_payload(context, timing), - ), - &response_obj, - CallbackTiming::new(timing.start_time, timing.end_time), - ) - .await; - }) +impl TerminalDispatcher for OcrLifecycleHooks { + 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 }) } +} - #[tracing::instrument( - name = "failure_callback", - target = "litellm::function_trace", - level = "trace", - skip_all - )] - fn async_log_failure_event<'a>( - &'a self, - context: &'a CallLifecycleContext, - error: &'a Error, - timing: &'a CallLifecycleTiming, - ) -> Self::FailureFuture<'a> { - Box::pin(async move { - if self.logger_runner.is_empty() { - return; - } - let logging_error = LoggingError { - message: error.to_string(), - kind: core_error_kind(error).to_string(), - }; - let response_obj = CallbackValue::new( - "error", - json!({ - "message": logging_error.message, - "kind": logging_error.kind, - }), - ); - self.logger_runner - .async_log_failure_event( - &ModelCallDetails::from_standard_logging_payload( - self.standard_logging_payload(context, timing), - ) - .with_failure_error(logging_error), - Some(&response_obj), - CallbackTiming::new(timing.start_time, timing.end_time), - ) - .await; - }) +fn request_metadata(metadata: &RequestMetadata) -> crate::integrations::types::StandardLoggingMetadata { + crate::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() } } diff --git a/litellm-rust/crates/core/src/ocr/runtime_types.rs b/litellm-rust/crates/core/src/ocr/runtime_types.rs index 5e9e531a0c3..d27f6a92188 100644 --- a/litellm-rust/crates/core/src/ocr/runtime_types.rs +++ b/litellm-rust/crates/core/src/ocr/runtime_types.rs @@ -1,7 +1,7 @@ use std::sync::Arc; use std::time::Duration; -use crate::call_lifecycle::{CallLifecycleContext, CallLifecycleRequest}; +use crate::lifecycle::{CallLifecycleContext, CallLifecycleRequest}; use crate::ocr::transformation::OcrProviderConfig; use serde_json::{Map, Value}; diff --git a/litellm-rust/crates/core/src/responses/instrumentation.rs b/litellm-rust/crates/core/src/responses/instrumentation.rs index b1098f4d386..5f29dbb143e 100644 --- a/litellm-rust/crates/core/src/responses/instrumentation.rs +++ b/litellm-rust/crates/core/src/responses/instrumentation.rs @@ -6,7 +6,10 @@ use std::time::{SystemTime, UNIX_EPOCH}; use serde_json::Value; use crate::Error; -use crate::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming}; +use crate::integrations::custom_logger::{LogError, LogFuture}; +use crate::lifecycle::{ + ActionResult, CallLifecycleContext, RequestPolicy, TerminalDispatcher, TerminalRecord, +}; use crate::responses::types::{ResponsesWsEvent, ResponsesWsEventType}; #[derive(Clone, Debug, Default, PartialEq, Eq)] @@ -205,20 +208,18 @@ impl ResponsesWsInstrumentation { } } -type LifecycleFuture<'a, T> = Pin> + Send + 'a>>; +type LifecycleFuture<'a, T> = Pin> + Send + 'a>>; -impl CallLifecycleHooks<(), (), ()> for ResponsesWsInstrumentation { +impl RequestPolicy<(), ()> for ResponsesWsInstrumentation { type PreCallFuture<'a> = LifecycleFuture<'a, ()>; type DuringCallFuture<'a> = LifecycleFuture<'a, ()>; - type SuccessFuture<'a> = Pin + Send + 'a>>; - type FailureFuture<'a> = Pin + Send + 'a>>; fn async_pre_call_hook<'a>( &'a self, _context: &'a CallLifecycleContext, request: (), ) -> Self::PreCallFuture<'a> { - Box::pin(async move { Ok(request) }) + Box::pin(async move { ActionResult::Continue(request) }) } fn async_during_call_hook<'a>( @@ -226,34 +227,21 @@ impl CallLifecycleHooks<(), (), ()> for ResponsesWsInstrumentation { _context: &'a CallLifecycleContext, request: (), ) -> Self::DuringCallFuture<'a> { - Box::pin(async move { Ok(request) }) + Box::pin(async move { ActionResult::Continue(request) }) } +} - fn async_log_success_event<'a>( - &'a self, - _context: &'a CallLifecycleContext, - _response: &'a (), - _timing: &'a CallLifecycleTiming, - ) -> Self::SuccessFuture<'a> { +impl TerminalDispatcher for ResponsesWsInstrumentation { + fn dispatch<'a>(&'a self, terminal: &'a TerminalRecord) -> LogFuture<'a> { Box::pin(async move { - let outcome = self.success_outcome(); - if let Ok(mut state) = self.state.lock() { - state.outcome = Some(outcome); - } - }) - } - - fn async_log_failure_event<'a>( - &'a self, - _context: &'a CallLifecycleContext, - _error: &'a Error, - _timing: &'a CallLifecycleTiming, - ) -> Self::FailureFuture<'a> { - Box::pin(async move { - let outcome = self.failure_outcome(); + 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>(()) }) } } @@ -332,9 +320,9 @@ mod tests { async fn lifecycle_records_success_outcome_for_provider_completion() { let instrumentation = ResponsesWsInstrumentation::new("call-1", "gpt-5", ResponsesWsMetadata::default()); - let result = crate::call_lifecycle::CallLifecycle::default() + let result = crate::lifecycle::CallLifecycle .run( - crate::call_lifecycle::CallLifecycleContext::new( + crate::lifecycle::CallLifecycleContext::new( "responses_websocket", "gpt-5", "openai", @@ -346,7 +334,10 @@ mod tests { ) .await; - assert!(result.is_ok()); + assert!(matches!( + result, + crate::lifecycle::ExecutedCall::Success { .. } + )); assert!(matches!( instrumentation.take_outcome(), Some(ResponsesWsLogOutcome::Success { .. }) diff --git a/litellm-rust/crates/core/tests/messages_lifecycle.rs b/litellm-rust/crates/core/tests/messages_lifecycle.rs new file mode 100644 index 00000000000..9fbdabd71c8 --- /dev/null +++ b/litellm-rust/crates/core/tests/messages_lifecycle.rs @@ -0,0 +1,160 @@ +use std::future::Future; +use std::pin::Pin; +use std::sync::Mutex; + +use litellm_core::Error; +use litellm_core::integrations::custom_logger::{LogError, LogFuture}; +use litellm_core::lifecycle::{ + ActionResult, CallLifecycleContext, Clock, RequestPolicy, TerminalClassification, + TerminalDispatcher, TerminalRecord, +}; +use litellm_core::messages::lifecycle::{Options, messages}; +use litellm_core::messages::types::MessagesRequest; +use serde_json::json; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; +use tokio::net::TcpListener; + +type PolicyFuture<'a, T> = Pin> + Send + 'a>>; + +#[derive(Default)] +struct Services { + reject: bool, + terminals: Mutex>, +} + +impl Clock for Services { + fn now(&self) -> f64 { + 10.0 + } +} + +impl RequestPolicy for Services { + type PreCallFuture<'a> = PolicyFuture<'a, MessagesRequest>; + type DuringCallFuture<'a> = PolicyFuture<'a, MessagesRequest>; + + fn async_pre_call_hook<'a>( + &'a self, + _: &'a CallLifecycleContext, + request: MessagesRequest, + ) -> Self::PreCallFuture<'a> { + Box::pin(async move { + if self.reject { + ActionResult::Reject(Error::InvalidRequest("blocked".into())) + } else { + ActionResult::Continue(request) + } + }) + } + + fn async_during_call_hook<'a>( + &'a self, + _: &'a CallLifecycleContext, + request: MessagesRequest, + ) -> Self::DuringCallFuture<'a> { + Box::pin(async move { ActionResult::Continue(request) }) + } +} + +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 request(api_base: String) -> MessagesRequest { + MessagesRequest { + model: "claude-test".into(), + body: json!({ + "model": "claude-test", + "max_tokens": 8, + "messages": [{"role": "user", "content": "hello"}] + }), + api_key: Some("test-key".into()), + api_base: Some(api_base), + custom_llm_provider: Some("anthropic".into()), + extra_headers: None, + timeout: None, + } +} + +fn context() -> CallLifecycleContext { + CallLifecycleContext::new("messages", "claude-test", "anthropic", "call-1") +} + +async fn upstream(status: u16) -> (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 body = if status == 200 { + r#"{"id":"msg_1","type":"message","role":"assistant","model":"claude-test","content":[],"stop_reason":"end_turn","stop_sequence":null}"# + } else { + r#"{"error":"failed"}"# + }; + let response = format!( + "HTTP/1.1 {status} Test\r\ncontent-type: application/json\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) +} + +#[tokio::test] +async fn success_dispatches_exactly_one_terminal() { + let (api_base, server) = upstream(200).await; + let services = Services::default(); + let result = messages(&services, request(api_base), Options::default(), context()).await; + + result.into_result().expect("messages succeeds"); + server.await.unwrap(); + let terminals = services.terminals.lock().unwrap(); + assert_eq!(terminals.len(), 1); + assert_eq!(terminals[0].classification, TerminalClassification::Success); +} + +#[tokio::test] +async fn provider_failure_dispatches_exactly_one_terminal() { + let (api_base, server) = upstream(500).await; + let services = Services::default(); + let result = messages(&services, request(api_base), Options::default(), context()).await; + + assert!(matches!( + result.into_result(), + Err(Error::Http { status: 500, .. }) + )); + server.await.unwrap(); + let terminals = services.terminals.lock().unwrap(); + assert_eq!(terminals.len(), 1); + assert!(matches!( + terminals[0].classification, + TerminalClassification::Failure { .. } + )); +} + +#[tokio::test] +async fn pre_call_rejection_never_touches_socket() { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let api_base = format!("http://{}", listener.local_addr().unwrap()); + let services = Services { + reject: true, + ..Default::default() + }; + let result = messages(&services, request(api_base), Options::default(), context()).await; + + assert!(matches!( + result.into_result(), + Err(Error::InvalidRequest(_)) + )); + assert_eq!(services.terminals.lock().unwrap().len(), 1); + assert!( + tokio::time::timeout(std::time::Duration::from_millis(50), listener.accept()) + .await + .is_err() + ); +} diff --git a/litellm-rust/crates/python-bridge/src/driver.rs b/litellm-rust/crates/python-bridge/src/driver.rs new file mode 100644 index 00000000000..331eecff6a5 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/driver.rs @@ -0,0 +1,48 @@ +use std::ffi::CString; + +use pyo3::prelude::*; + +const DRIVE: &str = r#" +def drive_sync(arguments): + host = Host(arguments, False) + while host.machine.complete() is None: + try: + _invoke(host.machine, host) + except Exception as error: + host.advance(1, error) + except BaseException as error: + host.advance(2, error) + else: + host.advance(0) + return host.result() + +async def drive_async(arguments): + host = Host(arguments, True) + while host.machine.complete() is None: + try: + awaiting, value = _invoke(host.machine, host) + if awaiting: + await value + except Exception as error: + host.advance(1, error) + except BaseException as error: + host.advance(2, error) + else: + host.advance(0) + return host.result() +"#; + +pub(crate) fn compile<'py>( + py: Python<'py>, + route: &str, + host: &str, +) -> PyResult> { + let source = CString::new(format!("{host}\n{DRIVE}")).map_err(|_| { + pyo3::exceptions::PyValueError::new_err("driver source contains a null byte") + })?; + let filename = CString::new(format!("{route}_driver.py")) + .map_err(|_| pyo3::exceptions::PyValueError::new_err("invalid driver route name"))?; + let module_name = CString::new(format!("_{route}_driver")) + .map_err(|_| pyo3::exceptions::PyValueError::new_err("invalid driver route name"))?; + PyModule::from_code(py, &source, &filename, &module_name) +} diff --git a/litellm-rust/crates/python-bridge/src/errors.rs b/litellm-rust/crates/python-bridge/src/errors.rs index 914e2e1e033..ec85252ffb4 100644 --- a/litellm-rust/crates/python-bridge/src/errors.rs +++ b/litellm-rust/crates/python-bridge/src/errors.rs @@ -18,12 +18,21 @@ pyo3::create_exception!( pub(crate) fn core_error_to_pyerr(err: Error) -> PyErr { match err { - Error::Auth(message) => PyValueError::new_err(message), - Error::InvalidProvider(_) - | Error::InvalidRequest(_) - | Error::InvalidType { .. } - | Error::MissingField(_) => PyValueError::new_err(err.to_string()), - other => PyRuntimeError::new_err(other.to_string()), + Error::InvalidProvider(_) => PyValueError::new_err("Invalid provider configuration"), + Error::InvalidRequest(_) => PyValueError::new_err("Invalid provider request"), + Error::InvalidType { .. } | Error::MissingField(_) => { + PyValueError::new_err(err.to_string()) + } + Error::Auth(_) => PyValueError::new_err("Provider authentication failed"), + Error::Http { status, .. } => { + PyRuntimeError::new_err(format!("Provider request failed (HTTP {status})")) + } + Error::Network(_) | Error::Connect(_) => { + PyRuntimeError::new_err("Provider transport failed") + } + Error::InvalidResponse(_) => PyRuntimeError::new_err("Invalid provider response"), + Error::Routing(_) => PyRuntimeError::new_err("Provider routing failed"), + Error::Unsupported(_) => PyRuntimeError::new_err("Operation is not supported"), } } @@ -35,22 +44,65 @@ pub(crate) fn core_error_to_pyerr(err: Error) -> PyErr { /// done the work and billed for it. pub(crate) fn chat_completions_error_to_pyerr(err: Error) -> PyErr { match err { - Error::Unsupported(_) - | Error::Auth(_) - | Error::InvalidProvider(_) - | Error::InvalidRequest(_) - | Error::InvalidType { .. } - | Error::MissingField(_) - | Error::Routing(_) - // Nothing reached the provider, so serving it on Python cannot double - // bill and is the only way the caller gets an answer at all. - | Error::Connect(_) => RustBridgeDeclined::new_err(err.to_string()), - Error::Http { status, body } => { - RustUpstreamError::new_err((status, format!("{status}: {body}"))) + Error::Unsupported(_) => RustBridgeDeclined::new_err("Operation is not supported"), + Error::Auth(_) => RustBridgeDeclined::new_err("Provider authentication failed"), + Error::InvalidProvider(_) => RustBridgeDeclined::new_err("Invalid provider configuration"), + Error::InvalidRequest(_) => RustBridgeDeclined::new_err("Invalid provider request"), + Error::InvalidType { .. } | Error::MissingField(_) => { + RustBridgeDeclined::new_err(err.to_string()) } - Error::Network(message) | Error::InvalidResponse(message) => { - RustUpstreamError::new_err((0u16, message)) + Error::Routing(_) => RustBridgeDeclined::new_err("Provider routing failed"), + Error::Connect(_) => RustBridgeDeclined::new_err("Could not reach provider"), + Error::Http { status, .. } => { + RustUpstreamError::new_err((status, format!("Provider request failed (HTTP {status})"))) } + Error::Network(_) => RustUpstreamError::new_err((0u16, "Provider transport failed")), + Error::InvalidResponse(_) => { + RustUpstreamError::new_err((0u16, "Invalid provider response")) + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn public_error_mappers_do_not_expose_private_details() { + Python::initialize(); + Python::attach(|_| { + for error in [ + Error::Auth("credential secret".into()), + Error::Http { + status: 500, + body: "response secret".into(), + }, + Error::Network("network secret".into()), + Error::Connect("connection secret".into()), + Error::InvalidResponse("parse secret".into()), + Error::Routing("routing secret".into()), + ] { + assert!(!core_error_to_pyerr(error).to_string().contains("secret")); + } + + for error in [ + Error::Auth("credential secret".into()), + Error::Http { + status: 500, + body: "response secret".into(), + }, + Error::Network("network secret".into()), + Error::Connect("connection secret".into()), + Error::InvalidResponse("parse secret".into()), + Error::Routing("routing secret".into()), + ] { + assert!( + !chat_completions_error_to_pyerr(error) + .to_string() + .contains("secret") + ); + } + }); } } diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 64a5c5bda76..f4a77fbef2e 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -1,4 +1,5 @@ mod diagnostics; +mod driver; mod errors; #[cfg(feature = "trace-parity")] mod function_trace; @@ -30,8 +31,8 @@ impl ResponsesWebSocketConnection { timeout_seconds: Option, ) -> PyResult> { let headers = marshal_headers(headers)?; - let timeout = optional_timeout(timeout_seconds); - pyo3_async_runtimes::tokio::future_into_py(py, async move { + 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)?; @@ -41,27 +42,27 @@ impl ResponsesWebSocketConnection { fn send_text<'py>(&self, py: Python<'py>, text: String) -> PyResult> { let inner = self.inner.clone(); - pyo3_async_runtimes::tokio::future_into_py(py, async move { + litellm_python_interop::run_async_py(py, async move { inner.send_text(text).await.map_err(core_error_to_pyerr) }) } fn recv_text<'py>(&self, py: Python<'py>) -> PyResult> { let inner = self.inner.clone(); - pyo3_async_runtimes::tokio::future_into_py(py, async move { + litellm_python_interop::run_async_py(py, async move { inner.recv_text().await.map_err(core_error_to_pyerr) }) } fn close<'py>(&self, py: Python<'py>) -> PyResult> { let inner = self.inner.clone(); - pyo3_async_runtimes::tokio::future_into_py(py, async move { + litellm_python_interop::run_async_py(py, async move { inner.close().await.map_err(core_error_to_pyerr) }) } } -#[pymodule(gil_used = false)] +#[pymodule(gil_used = true)] mod _native { use pyo3::prelude::*; diff --git a/litellm-rust/crates/python-bridge/src/marshal.rs b/litellm-rust/crates/python-bridge/src/marshal.rs index a14e4b55d82..cdb4319b146 100644 --- a/litellm-rust/crates/python-bridge/src/marshal.rs +++ b/litellm-rust/crates/python-bridge/src/marshal.rs @@ -1,7 +1,7 @@ use std::collections::HashMap; use std::time::Duration; -use pyo3::exceptions::PyValueError; +use pyo3::exceptions::{PyTypeError, PyValueError}; use pyo3::prelude::*; use serde_json::{Map, Value}; @@ -31,7 +31,7 @@ impl RouteOptions { api_base: inputs.api_base, custom_llm_provider: inputs.custom_llm_provider, extra_headers: optional_object("extra_headers", inputs.extra_headers)?, - timeout: optional_timeout(inputs.timeout_seconds), + timeout: optional_timeout(inputs.timeout_seconds)?, }) } } @@ -45,7 +45,7 @@ pub(crate) fn required_value( if expected(&value) { return Ok(value); } - Err(PyValueError::new_err(format!( + Err(PyTypeError::new_err(format!( "{name} must be a {expected_name}" ))) } @@ -70,18 +70,22 @@ fn optional_object( fn object(name: &'static str, value: Value) -> PyResult> { match value { Value::Object(map) => Ok(map), - _ => Err(PyValueError::new_err(format!("{name} must be a dict"))), + _ => Err(PyTypeError::new_err(format!("{name} must be a dict"))), } } -pub(crate) fn optional_timeout(timeout_seconds: Option) -> Option { - timeout_seconds.and_then(|secs| { - if secs.is_finite() && secs > 0.0 { - Some(Duration::from_secs_f64(secs)) - } else { - None - } - }) +pub(crate) fn optional_timeout(timeout_seconds: Option) -> PyResult> { + timeout_seconds + .map(|seconds| { + if seconds <= 0.0 { + return Err(PyValueError::new_err( + "timeout_seconds must be greater than zero", + )); + } + Duration::try_from_secs_f64(seconds) + .map_err(|_| PyValueError::new_err("timeout_seconds must be a finite duration")) + }) + .transpose() } pub(crate) fn marshal_headers(headers: Option) -> PyResult> { @@ -90,7 +94,7 @@ pub(crate) fn marshal_headers(headers: Option) -> PyResult Value::Object(Map::new()), }; let Value::Object(headers) = value else { - return Err(PyValueError::new_err("headers must be a dict")); + return Err(PyTypeError::new_err("headers must be a dict")); }; headers .into_iter() @@ -102,3 +106,28 @@ pub(crate) fn marshal_headers(headers: Option) -> PyResult(py)); + } + }); + } + + #[test] + fn timeout_accepts_none_and_positive_finite_values() { + assert_eq!(optional_timeout(None).unwrap(), None); + assert_eq!( + optional_timeout(Some(1.5)).unwrap(), + Some(Duration::from_millis(1500)) + ); + } +} diff --git a/litellm-rust/crates/python-bridge/src/routes/definition.rs b/litellm-rust/crates/python-bridge/src/routes/definition.rs index 45de81d8d80..e7fe7214e9e 100644 --- a/litellm-rust/crates/python-bridge/src/routes/definition.rs +++ b/litellm-rust/crates/python-bridge/src/routes/definition.rs @@ -228,11 +228,7 @@ mod tests { "atranscription", "(model, audio, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None)", ), - ( - "messages", - "amessages", - "(model, body, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None)", - ), + ("messages", "amessages", "(arguments)"), ( "chat_completions", "achat_completions", @@ -277,23 +273,30 @@ mod tests { assert_eq!( sync_chat_error.to_string(), - "ValueError: messages must be a list" + "TypeError: messages must be a list" ); assert_eq!(async_chat_error.to_string(), sync_chat_error.to_string()); let invalid_body = PyList::empty(py); + let invalid_arguments = PyDict::new(py); + invalid_arguments + .set_item("model", "model") + .expect("arguments should accept model"); + invalid_arguments + .set_item("body", &invalid_body) + .expect("arguments should accept body"); let sync_messages_error = module .getattr("messages") - .and_then(|function| function.call1(("model", &invalid_body))) + .and_then(|function| function.call1((&invalid_arguments,))) .expect_err("sync Messages should reject a non-dict body"); let async_messages_error = module .getattr("amessages") - .and_then(|function| function.call1(("model", &invalid_body))) + .and_then(|function| function.call1((&invalid_arguments,))) .expect_err("async Messages should reject a non-dict body"); assert_eq!( sync_messages_error.to_string(), - "ValueError: body must be a dict" + "TypeError: body must be a dict" ); assert_eq!( async_messages_error.to_string(), @@ -320,7 +323,7 @@ mod tests { assert_eq!( sync_error.to_string(), - "ValueError: extra_headers must be a dict" + "TypeError: extra_headers must be a dict" ); assert_eq!(async_error.to_string(), sync_error.to_string()); } @@ -349,7 +352,7 @@ mod tests { function.call(("model", &invalid_messages), Some(&chat_kwargs)) }) .expect_err("messages should be validated first"); - assert_eq!(error.to_string(), "ValueError: messages must be a list"); + assert_eq!(error.to_string(), "TypeError: messages must be a list"); let valid_messages = PyList::empty(py); let error = module @@ -358,7 +361,7 @@ mod tests { .expect_err("optional_params should be validated before headers"); assert_eq!( error.to_string(), - "ValueError: optional_params must be a dict" + "TypeError: optional_params must be a dict" ); let headers_kwargs = PyDict::new(py); @@ -366,11 +369,21 @@ mod tests { .set_item("extra_headers", &invalid) .expect("kwargs should accept extra_headers"); let invalid_body = PyList::empty(py); + let invalid_arguments = PyDict::new(py); + invalid_arguments + .set_item("model", "model") + .expect("arguments should accept model"); + invalid_arguments + .set_item("body", &invalid_body) + .expect("arguments should accept body"); + invalid_arguments + .set_item("extra_headers", &invalid) + .expect("arguments should accept extra_headers"); let error = module .getattr("messages") - .and_then(|function| function.call(("model", &invalid_body), Some(&headers_kwargs))) + .and_then(|function| function.call1((&invalid_arguments,))) .expect_err("body should be validated before headers"); - assert_eq!(error.to_string(), "ValueError: body must be a dict"); + assert_eq!(error.to_string(), "TypeError: body must be a dict"); let invalid_payload = PyModule::new(py, "invalid_payload").expect("invalid payload should be created"); diff --git a/litellm-rust/crates/python-bridge/src/routes/messages.rs b/litellm-rust/crates/python-bridge/src/routes/messages.rs index f69b5e9251d..dedfae4108f 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages.rs @@ -1,65 +1,372 @@ -use litellm_core::Error; -use litellm_core::messages::messages as run_messages; -use litellm_core::messages::types::{AnthropicMessagesResponse, MessagesRequest}; +use litellm_core::lifecycle::{ErrorDisposition, Lifecycle, Outcome}; +use litellm_core::messages::lifecycle::{MessagesRoute, Observations, Operation, Options, machine}; +use litellm_core::messages::types::MessagesRequest; +use litellm_python_interop::{Pythonized, from_py, run_async_value, run_sync_value}; +use pyo3::exceptions::{PyRuntimeError, PyTypeError, PyValueError}; use pyo3::prelude::*; -use serde_json::Value; -use std::future::Future; +use pyo3::pyclass::{PyTraverseError, PyVisit}; +use pyo3::sync::PyOnceLock; +use pyo3::types::PyDict; +use serde_json::{Map, Value}; use crate::errors::core_error_to_pyerr; -use crate::marshal::{RouteOptions, RouteOptionsInputs, required_value}; +use crate::marshal::optional_timeout; -fn prepare_messages( - inputs: MessagesInputs, -) -> PyResult> + Send + 'static> { - let body = required_value("body", inputs.body, Value::is_object, "dict")?; - 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 MessagesState { + arguments: Option>, + request: Option, +} - Ok(async move { - let RouteOptions { - model, - api_key, - api_base, - custom_llm_provider, - extra_headers, +#[pymethods] +impl MessagesState { + 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.request.take()) + }; + drop(roots); + } +} + +fn scalar(arguments: &Bound<'_, PyDict>, name: &str) -> PyResult> { + arguments + .get_item(name)? + .filter(|value| !value.is_none()) + .map(|value| value.extract::()) + .transpose() +} + +fn decode_state(py: Python<'_>, arguments: Py) -> PyResult { + let bag = arguments.bind(py); + let body = bag + .get_item(pyo3::intern!(py, "body"))? + .ok_or_else(|| PyValueError::new_err("messages requires body"))?; + let timeout = optional_timeout( + bag.get_item(pyo3::intern!(py, "timeout_seconds"))? + .filter(|value| !value.is_none()) + .map(|value| value.extract::()) + .transpose()?, + )?; + Ok(MessagesState { + request: Some(MessagesRequest { + body: from_py(&body)?, + model: scalar(bag, "model")? + .ok_or_else(|| PyValueError::new_err("messages requires model"))?, + api_key: scalar(bag, "api_key")?, + api_base: scalar(bag, "api_base")?, + custom_llm_provider: scalar(bag, "custom_llm_provider")?, + extra_headers: bag + .get_item("extra_headers")? + .filter(|value| !value.is_none()) + .map(|value| from_py::>(&value)) + .transpose()?, timeout, - } = options; - run_messages(MessagesRequest { - model: &model, - body, - api_key: api_key.as_deref(), - api_base: api_base.as_deref(), - custom_llm_provider: custom_llm_provider.as_deref(), - extra_headers, - timeout, - }) - .await + }), + arguments: Some(arguments), }) } -bridge_route! { - sync = messages, - asynchronous = amessages, - inputs = MessagesInputs, - required = { - model: String, - #[pyo3(from_py_with = litellm_python_interop::from_py)] - body: serde_json::Value, - }, - optional = { - api_key: Option, - api_base: Option, - custom_llm_provider: Option, - #[pyo3(from_py_with = litellm_python_interop::from_py)] - extra_headers: Option, - timeout_seconds: Option, - }, - prepare = prepare_messages, - errors = core_error_to_pyerr, +#[pyclass] +struct MessagesLifecycle { + machine: Lifecycle, + asynchronous: bool, +} + +#[pymethods] +impl MessagesLifecycle { + #[new] + fn new(asynchronous: bool, internal_call: bool) -> PyResult { + Ok(Self { + machine: machine(Options { + asynchronous, + internal_call, + ..Options::default() + }) + .map_err(core_error_to_pyerr)?, + asynchronous, + }) + } + + fn advance( + &mut self, + outcome: u8, + logger_available: bool, + has_fallbacks: bool, + ) -> PyResult { + let outcome = match outcome { + 0 => Outcome::Success, + 1 => Outcome::Failure, + _ => Outcome::Abort, + }; + self.machine + .advance( + outcome, + Observations { + logger_available, + has_fallbacks, + }, + ) + .map(|transition| transition.error == ErrorDisposition::Replace) + .map_err(core_error_to_pyerr) + } + + fn complete(&self) -> Option { + match self.machine.operation() { + Operation::Complete(outcome) => Some(outcome == Outcome::Success), + _ => None, + } + } +} + +#[pyfunction] +fn invoke( + py: Python<'_>, + machine: Py, + host: Py, +) -> PyResult<(bool, Py)> { + let (operation, asynchronous) = { + let machine = machine.borrow(py); + (machine.machine.operation(), machine.asynchronous) + }; + let (method, awaiting) = match operation { + Operation::Setup => ("setup", false), + Operation::DeploymentPre => ("deployment_pre", true), + Operation::Prepare => ("prepare", false), + Operation::Send if asynchronous => ("send", true), + Operation::Send => ("send_sync", false), + Operation::DeploymentSuccess => ("deployment_success", true), + Operation::DeploymentFailure => ("deployment_failure", true), + Operation::SyncSuccess => ("sync_success", false), + Operation::AsyncSuccess => ("async_success", false), + Operation::SyncSuccessIfNeeded => ("sync_success_if_needed", false), + Operation::SyncFailure => ("sync_failure", false), + Operation::AsyncFailure => ("async_failure", true), + Operation::Restore => ("restore", false), + Operation::Complete(_) => { + return Err(PyRuntimeError::new_err("messages lifecycle is complete")); + } + }; + Ok((awaiting, host.getattr(py, method)?.call0(py)?)) +} + +#[pyfunction] +fn prepare(py: Python<'_>, arguments: Py) -> PyResult> { + let state = decode_state(py, arguments)?; + let bag = state.arguments.as_ref().unwrap().bind(py); + let logging = bag + .get_item("litellm_logging_obj")? + .filter(|value| !value.is_none()) + .map(|value| value.unbind()) + .ok_or_else(|| PyRuntimeError::new_err("messages logging was not initialized"))?; + let additional = PyDict::new(py); + additional.set_item( + pyo3::intern!(py, "complete_input_dict"), + bag.get_item(pyo3::intern!(py, "body"))?, + )?; + additional.set_item( + pyo3::intern!(py, "api_base"), + bag.get_item(pyo3::intern!(py, "api_base"))?, + )?; + additional.set_item( + pyo3::intern!(py, "headers"), + bag.get_item(pyo3::intern!(py, "extra_headers"))?, + )?; + let kwargs = PyDict::new(py); + kwargs.set_item( + "input", + bag.get_item("messages")? + .unwrap_or_else(|| py.None().into_bound(py)), + )?; + kwargs.set_item("api_key", "")?; + kwargs.set_item("additional_args", additional)?; + logging + .bind(py) + .call_method(pyo3::intern!(py, "pre_call"), (), Some(&kwargs))?; + Py::new(py, state) +} + +fn take_request(py: Python<'_>, state: &Py) -> PyResult { + state + .borrow_mut(py) + .request + .take() + .ok_or_else(|| PyRuntimeError::new_err("messages request was already sent or cleared")) +} + +fn validate_arguments(arguments: &Bound<'_, PyDict>) -> PyResult<()> { + let body = arguments + .get_item("body")? + .ok_or_else(|| PyValueError::new_err("messages requires body"))?; + if !body.is_instance_of::() { + return Err(PyTypeError::new_err("body must be a dict")); + } + if let Some(headers) = arguments + .get_item("extra_headers")? + .filter(|value| !value.is_none()) + && !headers.is_instance_of::() + { + return Err(PyTypeError::new_err("extra_headers must be a dict")); + } + Ok(()) +} + +#[pyfunction] +fn send(py: Python<'_>, state: Py) -> PyResult> { + let request = take_request(py, &state)?; + litellm_python_interop::run_async_py(py, async move { + let _state = state; + let response = run_async_value( + litellm_core::messages::messages(request), + core_error_to_pyerr, + ) + .await?; + Ok(Pythonized(response)) + }) +} + +#[pyfunction] +fn send_sync(py: Python<'_>, state: Py) -> PyResult> { + let request = take_request(py, &state)?; + let response = run_sync_value( + py, + litellm_core::messages::messages(request), + core_error_to_pyerr, + )?; + Ok(Pythonized(response).into_pyobject(py)?.unbind().into_any()) +} + +#[pyfunction] +fn messages(py: Python<'_>, arguments: Py) -> PyResult> { + validate_arguments(arguments.bind(py))?; + driver(py)?.getattr("drive_sync")?.call1((arguments,)) +} + +#[pyfunction] +fn amessages(py: Python<'_>, arguments: Py) -> PyResult> { + validate_arguments(arguments.bind(py))?; + driver(py)?.getattr("drive_async")?.call1((arguments,)) +} + +fn driver(py: Python<'_>) -> PyResult<&Bound<'_, PyModule>> { + static DRIVER: PyOnceLock> = PyOnceLock::new(); + if let Some(module) = DRIVER.get(py) { + return Ok(module.bind(py)); + } + let module = crate::driver::compile(py, "messages", HOST)?; + module.add("_Lifecycle", py.get_type::())?; + module.add("_invoke", wrap_pyfunction!(invoke, &module)?)?; + module.add("_prepare", wrap_pyfunction!(prepare, &module)?)?; + module.add("_send", wrap_pyfunction!(send, &module)?)?; + module.add("_send_sync", wrap_pyfunction!(send_sync, &module)?)?; + 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.messages import initialize_logging, invoke_terminal + +class Host: + def __init__(self, arguments, asynchronous): + self.machine = _Lifecycle(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, 'amessages') + 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 = _send_sync(self.state) + self.end = datetime.now() + + async def send(self): + self.response = await _send(self.state) + self.end = datetime.now() + + async def deployment_success(self): + self.response = await utils.async_post_call_success_deployment_hook(self.current, self.response, CallTypes.aanthropic_messages) + + async def deployment_failure(self): + await utils.async_post_call_failure_deployment_hook(self.current, self.error, 'amessages') + + def terminal(self, action, value): + return invoke_terminal(action, (self.arguments, self.current, self.state), self.logger, 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!(messages, module)?)?; + crate::routes::definition::add_function(module, wrap_pyfunction!(amessages, module)?)?; + Ok(()) +} + +#[cfg(feature = "trace-parity")] +pub(super) fn register_trace(module: &Bound<'_, PyModule>) -> PyResult<()> { + register(module) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn decode_state_rejects_invalid_timeouts_without_panicking() { + Python::initialize(); + Python::attach(|py| { + for timeout in [-1.0, 0.0, f64::NAN, f64::INFINITY, f64::MAX] { + let arguments = PyDict::new(py); + arguments.set_item("model", "model").unwrap(); + arguments.set_item("body", PyDict::new(py)).unwrap(); + arguments.set_item("timeout_seconds", timeout).unwrap(); + let error = match decode_state(py, arguments.unbind()) { + Ok(_) => panic!("invalid timeout should fail normally"), + Err(error) => error, + }; + assert!(error.is_instance_of::(py)); + } + }); + } } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr.rs b/litellm-rust/crates/python-bridge/src/routes/ocr.rs index b83e48e438a..2fa9599c4c9 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr.rs @@ -330,9 +330,11 @@ fn prepare(py: Python<'_>, arguments: Py, asynchronous: bool) -> PyResul let request = decode_request(py, bag)?; let model = request.model.clone(); let custom_llm_provider = request.custom_llm_provider.clone(); - let prepared = litellm_core::ocr::prepare::prepare(request).map_err(|error| { - request_error_to_pyerr(py, error, &model, custom_llm_provider.as_deref()) - })?; + let prepared = py + .detach(|| litellm_core::ocr::prepare::prepare(request)) + .map_err(|error| { + request_error_to_pyerr(py, error, &model, custom_llm_provider.as_deref()) + })?; let document = bag .get_item("document")? .ok_or_else(|| PyValueError::new_err("OCR requires document"))? @@ -341,9 +343,11 @@ fn prepare(py: Python<'_>, arguments: Py, asynchronous: bool) -> PyResul .into_bound(py) .cast_into::()?; match prepared.document_projection { - OcrDocumentProjection::RetainedDocument => body.set_item("document", &document)?, + OcrDocumentProjection::RetainedDocument => { + body.set_item(pyo3::intern!(py, "document"), &document)? + } OcrDocumentProjection::ShallowCopyDocument => { - body.set_item("document", document.copy()?)? + body.set_item(pyo3::intern!(py, "document"), document.copy()?)? } OcrDocumentProjection::Transformed => {} } @@ -364,7 +368,10 @@ fn prepare(py: Python<'_>, arguments: Py, asynchronous: bool) -> PyResul .call1((bag, asynchronous))?; let litellm_params = PyDict::new(py); litellm_params.set_item("litellm_call_id", bag.get_item("litellm_call_id")?)?; - litellm_params.set_item("api_base", bag.get_item("api_base")?)?; + litellm_params.set_item( + pyo3::intern!(py, "api_base"), + bag.get_item(pyo3::intern!(py, "api_base"))?, + )?; let update = PyDict::new(py); update.set_item("kwargs", bag)?; update.set_item("model", &prepared.model)?; @@ -375,13 +382,13 @@ fn prepare(py: Python<'_>, arguments: Py, asynchronous: bool) -> PyResul let additional_args = PyDict::new(py); additional_args.set_item("complete_input_dict", &body)?; - additional_args.set_item("api_base", &prepared.url)?; - additional_args.set_item("headers", &headers)?; + additional_args.set_item(pyo3::intern!(py, "api_base"), &prepared.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("pre_call", (), Some(&pre_call))?; + logging.call_method(pyo3::intern!(py, "pre_call"), (), Some(&pre_call))?; let logging = logging.unbind(); Py::new( @@ -427,7 +434,7 @@ fn request(py: Python<'_>, state: &Py) -> PyResult { #[pyfunction] fn send(py: Python<'_>, state: Py) -> PyResult> { let (prepared, headers, body) = request(py, &state)?; - pyo3_async_runtimes::tokio::future_into_py(py, async move { + litellm_python_interop::run_async_py(py, async move { let _state = state; let model = prepared.model.clone(); let provider = prepared.custom_llm_provider.clone(); @@ -443,17 +450,20 @@ fn send(py: Python<'_>, state: Py) -> PyResult> { #[pyfunction] fn finish(py: Python<'_>, response: Py) -> PyResult> { - let fields = response.bind(py); - let native_response = fields.get_item("provider_native_response")?; + let fields = response.bind(py).copy()?; + let native_response = fields.get_item(pyo3::intern!(py, "provider_native_response"))?; if native_response.is_some() { - fields.del_item("provider_native_response")?; + fields.del_item(pyo3::intern!(py, "provider_native_response"))?; } let response = py .import("litellm.llms.base_llm.ocr.transformation")? .getattr("OCRResponse")? - .call((), Some(fields))?; + .call((), Some(&fields))?; if let Some(native_response) = native_response.filter(|value| !value.is_none()) { - response.call_method1("set_provider_native_response", (native_response,))?; + response.call_method1( + pyo3::intern!(py, "set_provider_native_response"), + (native_response,), + )?; } Ok(response.unbind()) } @@ -484,7 +494,7 @@ fn ocr(py: Python<'_>, arguments: Py) -> PyResult> { #[pyfunction] fn aocr(py: Python<'_>, arguments: Py) -> PyResult> { - driver(py)?.getattr("drive")?.call1((arguments,)) + driver(py)?.getattr("drive_async")?.call1((arguments,)) } // Compilation can re-enter through audit hooks; publish only a finished module. @@ -493,9 +503,10 @@ fn driver(py: Python<'_>) -> PyResult<&Bound<'_, PyModule>> { if let Some(module) = DRIVER.get(py) { return Ok(module.bind(py)); } - let module = PyModule::from_code( + let module = crate::driver::compile( py, - c"from datetime import datetime + "ocr", + "from datetime import datetime from litellm import utils from litellm.types.utils import CallTypes from litellm.rust_bridge.ocr import initialize_logging, invoke_terminal @@ -581,36 +592,7 @@ class Host: return self.response raise self.error -def drive_sync(arguments): - host = Host(arguments, False) - while host.machine.complete() is None: - try: - _invoke(host.machine, host) - except Exception as error: - host.advance(1, error) - except BaseException as error: - host.advance(2, error) - else: - host.advance(0) - return host.result() - -async def drive(arguments): - host = Host(arguments, True) - while host.machine.complete() is None: - try: - awaiting, value = _invoke(host.machine, host) - if awaiting: - await value - except Exception as error: - host.advance(1, error) - except BaseException as error: - host.advance(2, error) - else: - host.advance(0) - return host.result() ", - c"ocr_driver.py", - c"_ocr_driver", )?; module.add("_Lifecycle", py.get_type::())?; module.add("_invoke", wrap_pyfunction!(invoke, &module)?)?; @@ -724,6 +706,30 @@ mod tests { }); } + #[test] + fn finish_does_not_modify_the_input_dictionary() { + Python::initialize(); + Python::attach(|py| { + let fields = PyDict::new(py); + fields + .set_item("provider_native_response", "native") + .unwrap(); + fields.set_item("model", "model").unwrap(); + + let _ = finish(py, fields.clone().unbind()); + + assert_eq!( + fields + .get_item("provider_native_response") + .unwrap() + .unwrap() + .extract::() + .unwrap(), + "native" + ); + }); + } + #[test] fn callback_decline_is_terminal_and_identity_is_reused() { Python::initialize(); diff --git a/litellm-rust/crates/python-interop/src/execution.rs b/litellm-rust/crates/python-interop/src/execution.rs index 5928bc5355b..6b0b254a8be 100644 --- a/litellm-rust/crates/python-interop/src/execution.rs +++ b/litellm-rust/crates/python-interop/src/execution.rs @@ -78,6 +78,19 @@ where }) } +pub fn run_async_py(py: Python<'_>, future: F) -> PyResult> +where + T: for<'py> IntoPyObject<'py> + Send + 'static, + F: Future> + Send + 'static, +{ + pyo3_async_runtimes::tokio::future_into_py(py, async move { + AssertUnwindSafe(future) + .catch_unwind() + .await + .map_err(panic_to_pyerr)? + }) +} + pub async fn run_async_value(future: F, map_error: fn(E) -> PyErr) -> PyResult where T: Send + 'static, @@ -192,6 +205,15 @@ mod tests { ) } + #[pyfunction] + fn async_py_panic(py: Python<'_>) -> PyResult> { + run_async_py(py, async move { + panic!("ordinary async future panicked"); + #[allow(unreachable_code)] + Ok(true) + }) + } + #[pyfunction] fn runtime_worker_count() -> usize { pyo3_async_runtimes::tokio::get_runtime() @@ -455,4 +477,38 @@ asyncio.run(exercise()) .expect("result delivery should leave Tokio workers responsive"); }); } + + #[test] + fn ordinary_async_runner_maps_panics() { + Python::initialize(); + Python::attach(|py| { + let module = PyModule::new(py, "runtime").expect("module should be created"); + module + .add_function( + wrap_pyfunction!(async_py_panic, &module).expect("function should wrap"), + ) + .expect("function should register"); + let locals = PyDict::new(py); + locals.set_item("runtime", &module).unwrap(); + py.run( + c" +import asyncio + +async def exercise(): + try: + await runtime.async_py_panic() + except BaseException as error: + assert type(error).__name__ == 'PanicException' + assert str(error) == 'ordinary async future panicked' + else: + raise AssertionError('panic was not raised') + +asyncio.run(exercise()) +", + Some(&locals), + Some(&locals), + ) + .unwrap(); + }); + } } diff --git a/litellm-rust/crates/python-interop/src/lib.rs b/litellm-rust/crates/python-interop/src/lib.rs index bd638ef8a73..5601c6bdad3 100644 --- a/litellm-rust/crates/python-interop/src/lib.rs +++ b/litellm-rust/crates/python-interop/src/lib.rs @@ -1,5 +1,5 @@ mod execution; mod marshal; -pub use execution::{run_async, run_async_value, run_sync, run_sync_value}; +pub use execution::{run_async, run_async_py, run_async_value, run_sync, run_sync_value}; pub use marshal::{Pythonized, from_py, panic_to_pyerr, to_py}; diff --git a/litellm-rust/crates/python-interop/src/marshal.rs b/litellm-rust/crates/python-interop/src/marshal.rs index 46b7d04b54b..04b49341101 100644 --- a/litellm-rust/crates/python-interop/src/marshal.rs +++ b/litellm-rust/crates/python-interop/src/marshal.rs @@ -1,7 +1,7 @@ use std::any::Any; use std::panic::{AssertUnwindSafe, catch_unwind}; -use pyo3::exceptions::PyValueError; +use pyo3::exceptions::{PyTypeError, PyValueError}; use pyo3::panic::PanicException; use pyo3::prelude::*; use serde::Serialize; @@ -11,7 +11,7 @@ pub fn from_py(value: &Bound<'_, PyAny>) -> PyResult where T: DeserializeOwned, { - pythonize::depythonize(value).map_err(|error| PyValueError::new_err(error.to_string())) + pythonize::depythonize(value).map_err(|error| PyTypeError::new_err(error.to_string())) } pub fn to_py(py: Python<'_>, value: &T) -> PyResult> diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 2f561809940..9c4dcf7afdf 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -2243,16 +2243,6 @@ class BaseLLMHTTPHandler: # internally -- this only deduplicates the success path. request_body_json: Final = json.dumps(request_body) - logging_obj.pre_call( - input=[{"role": "user", "content": request_body_json}], - api_key="", - additional_args={ - "complete_input_dict": request_body, - "api_base": str(request_url), - "headers": headers, - }, - ) - rust_messages_response: Final = await self._maybe_rust_anthropic_messages( custom_llm_provider=custom_llm_provider, litellm_params=litellm_params, @@ -2267,6 +2257,12 @@ class BaseLLMHTTPHandler: stream=stream or False, custom_llm_provider=custom_llm_provider, ), + arguments={ + **kwargs, + "messages": messages, + "litellm_logging_obj": logging_obj, + "litellm_params": litellm_params, + }, ) if rust_messages_response is not None: if stream: @@ -2283,6 +2279,16 @@ class BaseLLMHTTPHandler: kwargs=kwargs, ) + logging_obj.pre_call( + input=[{"role": "user", "content": request_body_json}], + api_key="", + additional_args={ + "complete_input_dict": request_body, + "api_base": str(request_url), + "headers": headers, + }, + ) + response: Final = await self._async_post_anthropic_messages_with_http_error_retry( async_httpx_client=async_httpx_client, request_url=request_url, @@ -2423,6 +2429,7 @@ class BaseLLMHTTPHandler: headers: dict, request_body: dict, timeout: float | httpx.Timeout | None, + arguments: dict[str, object] | None = None, ) -> AnthropicMessagesResponse | None: if custom_llm_provider not in ("azure_ai", "anthropic"): return None @@ -2438,6 +2445,7 @@ class BaseLLMHTTPHandler: upstream_body: Final = {key: value for key, value in request_body.items() if key != "stream"} try: rust_response: Final = await rust_messages_bridge.amessages( + arguments=arguments or {}, model=model, body=upstream_body, api_key=api_key, diff --git a/litellm/rust_bridge/_lifecycle.py b/litellm/rust_bridge/_lifecycle.py new file mode 100644 index 00000000000..75aa86a24d4 --- /dev/null +++ b/litellm/rust_bridge/_lifecycle.py @@ -0,0 +1,22 @@ +from __future__ import annotations + +from datetime import datetime + + +def initialize_logging(arguments: dict[str, object], asynchronous: bool, route: str) -> object: + from litellm.rust_bridge.ocr import initialize_logging as initialize_ocr_logging + + return initialize_ocr_logging(arguments, asynchronous, route) + + +def invoke_terminal( + action: str, + roots: object, + logger: object, + value: object, + start_time: datetime, + end_time: datetime, +) -> object: + from litellm.rust_bridge.ocr import invoke_terminal as invoke_ocr_terminal + + return invoke_ocr_terminal(action, roots, logger, value, start_time, end_time) diff --git a/litellm/rust_bridge/messages.py b/litellm/rust_bridge/messages.py index 40d0ddf622b..80409688789 100644 --- a/litellm/rust_bridge/messages.py +++ b/litellm/rust_bridge/messages.py @@ -1,42 +1,22 @@ -"""Thin Python wrapper for the native Rust Anthropic Messages bridge.""" - from __future__ import annotations from collections.abc import Awaitable from dataclasses import dataclass from typing import Final, Protocol, cast -import httpx - +from litellm.rust_bridge._lifecycle import ( + initialize_logging as initialize_lifecycle_logging, +) +from litellm.rust_bridge._lifecycle import invoke_terminal from litellm.rust_bridge.timeouts import timeout_to_seconds class RustMessages(Protocol): - def __call__( - self, - model: str, - body: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - timeout_seconds: float | None, - ) -> dict[str, object]: - raise NotImplementedError + def __call__(self, arguments: dict[str, object]) -> dict[str, object]: ... class RustAmessages(Protocol): - def __call__( - self, - model: str, - body: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - timeout_seconds: float | None, - ) -> Awaitable[dict[str, object]]: - raise NotImplementedError + def __call__(self, arguments: dict[str, object]) -> Awaitable[dict[str, object]]: ... class _Unset: @@ -52,7 +32,7 @@ class _RustMessagesState: amessages: RustAmessages | None = None -_STATE: Final[_RustMessagesState] = _RustMessagesState() +_STATE: Final = _RustMessagesState() def set_rust_messages( @@ -71,10 +51,8 @@ def load_rust_messages() -> RustMessages | None: return _STATE.messages from litellm.rust_bridge import get_native_bridge - native_bridge: Final = get_native_bridge() - if native_bridge is None: - return None - return cast(RustMessages, getattr(native_bridge, "messages", None)) + bridge: Final = get_native_bridge() + return cast(RustMessages, getattr(bridge, "messages", None)) if bridge is not None else None def load_rust_amessages() -> RustAmessages | None: @@ -82,10 +60,34 @@ def load_rust_amessages() -> RustAmessages | None: return _STATE.amessages from litellm.rust_bridge import get_native_bridge - native_bridge: Final = get_native_bridge() - if native_bridge is None: - return None - return cast(RustAmessages, getattr(native_bridge, "amessages", None)) + bridge: Final = get_native_bridge() + return cast(RustAmessages, getattr(bridge, "amessages", None)) if bridge is not None else None + + +def initialize_logging(arguments: dict[str, object], asynchronous: bool) -> object: + return initialize_lifecycle_logging(arguments, asynchronous, "messages") + + +def _arguments( + arguments: dict[str, object], + model: str, + body: dict[str, object], + api_key: str | None, + api_base: str | None, + custom_llm_provider: str | None, + extra_headers: dict[str, object] | None, + timeout: object, +) -> dict[str, object]: + return { + **arguments, + "model": model, + "body": body, + "api_key": api_key, + "api_base": api_base, + "custom_llm_provider": custom_llm_provider, + "extra_headers": extra_headers, + "timeout_seconds": timeout_to_seconds(timeout), + } def messages( @@ -96,19 +98,14 @@ def messages( api_base: str | None, custom_llm_provider: str | None, extra_headers: dict[str, object] | None, - timeout: float | httpx.Timeout | None, + timeout: object, + arguments: dict[str, object] | None = None, ) -> dict[str, object] | None: - rust_messages: Final = load_rust_messages() - if rust_messages is None: + implementation: Final = load_rust_messages() + if implementation is None: return None - return rust_messages( - model=model, - body=body, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - timeout_seconds=timeout_to_seconds(timeout), + return implementation( + arguments=_arguments(arguments or {}, model, body, api_key, api_base, custom_llm_provider, extra_headers, timeout) ) @@ -120,17 +117,15 @@ async def amessages( api_base: str | None, custom_llm_provider: str | None, extra_headers: dict[str, object] | None, - timeout: float | httpx.Timeout | None, + timeout: object, + arguments: dict[str, object] | None = None, ) -> dict[str, object] | None: - rust_amessages: Final = load_rust_amessages() - if rust_amessages is None: + implementation: Final = load_rust_amessages() + if implementation is None: return None - return await rust_amessages( - model=model, - body=body, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - timeout_seconds=timeout_to_seconds(timeout), + return await implementation( + arguments=_arguments(arguments or {}, model, body, api_key, api_base, custom_llm_provider, extra_headers, timeout) ) + + +__all__ = ["amessages", "initialize_logging", "invoke_terminal", "load_rust_amessages", "load_rust_messages", "messages", "set_rust_messages"] diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index 1825e71346e..2d9e8aec9f3 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -55,7 +55,7 @@ async def aocr(arguments: dict[str, object]) -> OCRResponse: return await implementation(arguments) -def initialize_logging(arguments: dict[str, object], asynchronous: bool) -> object: +def initialize_logging(arguments: dict[str, object], asynchronous: bool, route: str = "ocr") -> object: import litellm from litellm import utils from litellm.integrations.custom_logger import CustomLogger @@ -155,7 +155,7 @@ def initialize_logging(arguments: dict[str, object], asynchronous: bool) -> obje model=str(arguments["model"]), messages="default-message-value", stream=False, - call_type="aocr" if asynchronous else "ocr", + call_type=f"a{route}" if asynchronous else route, start_time=datetime.now(), # noqa: DTZ005 # Logging preserves the legacy naive timestamp contract litellm_call_id=call_id, function_id=str(arguments.get("id") or ""), diff --git a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py index a30474245c6..6f3839750ed 100644 --- a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py +++ b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py @@ -40,25 +40,9 @@ class RecordingMessages: def __call__( self, - model: str, - body: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - timeout_seconds: float | None, + arguments: dict[str, object], ) -> dict[str, object]: - self.calls.append( - { - "model": model, - "body": body, - "api_key": api_key, - "api_base": api_base, - "custom_llm_provider": custom_llm_provider, - "extra_headers": extra_headers, - "timeout_seconds": timeout_seconds, - } - ) + self.calls.append(arguments) return dict(FAKE_MESSAGES_RESPONSE) @@ -68,25 +52,9 @@ class RecordingAsyncMessages: async def __call__( self, - model: str, - body: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - timeout_seconds: float | None, + arguments: dict[str, object], ) -> dict[str, object]: - self.calls.append( - { - "model": model, - "body": body, - "api_key": api_key, - "api_base": api_base, - "custom_llm_provider": custom_llm_provider, - "extra_headers": extra_headers, - "timeout_seconds": timeout_seconds, - } - ) + self.calls.append(arguments) return dict(FAKE_MESSAGES_RESPONSE)