From 95d594a0fece445ccf7452b3a1dc1974358b08f0 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 29 Jul 2026 21:52:56 +0000 Subject: [PATCH] fix(rust): harden provider debug logging Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> --- .../src/integrations/logging/console.rs | 34 ++- .../src/integrations/logging/mod.rs | 7 +- .../crates/ai-gateway/src/ocr/common_utils.rs | 13 +- .../crates/ai-gateway/src/ocr/handler.rs | 4 +- .../crates/ai-gateway/src/ocr/tests.rs | 3 +- litellm-rust/crates/core/src/lib.rs | 1 + .../crates/core/src/logging/events.rs | 26 +- litellm-rust/crates/core/src/logging/http.rs | 89 +++--- litellm-rust/crates/core/src/logging/mod.rs | 58 ++-- .../crates/core/src/logging/redaction.rs | 48 ++++ .../crates/core/src/logging/stream.rs | 80 +++++- .../crates/core/src/messages/common_utils.rs | 9 - .../crates/core/src/messages/tests.rs | 253 +++++++++++++++++- litellm-rust/crates/core/src/utils.rs | 9 + 14 files changed, 499 insertions(+), 135 deletions(-) create mode 100644 litellm-rust/crates/core/src/utils.rs diff --git a/litellm-rust/crates/ai-gateway/src/integrations/logging/console.rs b/litellm-rust/crates/ai-gateway/src/integrations/logging/console.rs index 8bbd9f79e02..89884829acd 100644 --- a/litellm-rust/crates/ai-gateway/src/integrations/logging/console.rs +++ b/litellm-rust/crates/ai-gateway/src/integrations/logging/console.rs @@ -4,7 +4,7 @@ use std::sync::{Arc, Mutex}; use colored_json::{ColorMode, ColoredFormatter, Output, PrettyFormatter}; -use super::{LogSink, ProviderDebugEvent}; +use super::{LogEvent, LogSink}; #[derive(Clone, Copy)] enum RenderMode { @@ -66,24 +66,24 @@ fn render_mode() -> &'static RenderMode { }) } -fn header(event: &ProviderDebugEvent) -> String { +fn header(event: &LogEvent) -> String { match event { - ProviderDebugEvent::Request(value) => { + LogEvent::Request(value) => { format!("provider.request {} {}", value.call_id, value.provider) } - ProviderDebugEvent::Response(value) => format!( + LogEvent::Response(value) => format!( "provider.response {} {} status={} duration_ms={}", value.call_id, value.provider, value.status, value.duration_ms ), - ProviderDebugEvent::StreamStarted(value) => format!( + LogEvent::StreamStarted(value) => format!( "provider.stream.started {} {} status={}", value.call_id, value.provider, value.status ), - ProviderDebugEvent::StreamCompleted(value) => format!( + LogEvent::StreamCompleted(value) => format!( "provider.stream.completed {} {} duration_ms={}", value.call_id, value.provider, value.duration_ms ), - ProviderDebugEvent::Error(value) => format!( + LogEvent::Error(value) => format!( "provider.error {} {}{} duration_ms={}", value.call_id, value.provider, @@ -104,7 +104,7 @@ fn decorate(value: &str, color_mode: ColorMode, code: &str) -> String { } impl LogSink for ConsoleDebugHook { - fn emit(&self, event: &ProviderDebugEvent) { + fn emit(&self, event: &LogEvent) { let Ok(mut output) = self.output.lock() else { return; }; @@ -151,7 +151,7 @@ mod tests { use serde_json::json; use super::*; - use litellm_core::logging::{RequestEventInput, request_event}; + use litellm_core::logging::{LogEvent, ProviderRequestEvent}; struct Buffer(Arc>>); @@ -170,14 +170,18 @@ mod tests { fn compact_output_is_canonical_json() { let buffer = Arc::new(Mutex::new(Vec::new())); let hook = ConsoleDebugHook::with_writer_and_mode(Box::new(Buffer(buffer.clone())), false); - let event = request_event(RequestEventInput { + let event = LogEvent::Request(ProviderRequestEvent { + source: "litellm-rust", call_id: "call_01".to_string(), provider: "anthropic".to_string(), model: "claude".to_string(), stream: false, + method: "POST", url: "https://example.test".to_string(), - headers: Vec::new(), + headers: Default::default(), body: json!({"prompt": "visible"}), + body_truncated: None, + body_original_bytes: None, }); let expected = serde_json::to_value(&event).expect("event serializes"); hook.emit(&event); @@ -193,14 +197,18 @@ mod tests { fn pretty_output_has_header_separator_and_indented_payload() { let buffer = Arc::new(Mutex::new(Vec::new())); let hook = ConsoleDebugHook::with_writer_and_mode(Box::new(Buffer(buffer.clone())), true); - let event = request_event(RequestEventInput { + let event = LogEvent::Request(ProviderRequestEvent { + source: "litellm-rust", call_id: "call_01".to_string(), provider: "anthropic".to_string(), model: "claude".to_string(), stream: false, + method: "POST", url: "https://example.test".to_string(), - headers: Vec::new(), + headers: Default::default(), body: json!({"prompt": "visible"}), + body_truncated: None, + body_original_bytes: None, }); hook.emit(&event); let output = diff --git a/litellm-rust/crates/ai-gateway/src/integrations/logging/mod.rs b/litellm-rust/crates/ai-gateway/src/integrations/logging/mod.rs index 93daa8c1e26..cb7fcb93592 100644 --- a/litellm-rust/crates/ai-gateway/src/integrations/logging/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/integrations/logging/mod.rs @@ -1,8 +1,7 @@ pub mod console; pub use litellm_core::logging::{ - BodySnapshot, ErrorEventInput, LogSink, ProviderDebugEvent, ProviderErrorEvent, - ProviderRequestEvent, ProviderResponseEvent, ProviderStreamCompletedEvent, - ProviderStreamStartedEvent, RequestEventInput, ResponseBody, ResponseEventInput, error_event, - request_event, response_event, stream_completed, stream_started, + BodySnapshot, ErrorEventInput, LogEvent, LogSink, ProviderErrorEvent, ProviderRequestEvent, + ProviderResponseEvent, ProviderStreamCompletedEvent, ProviderStreamStartedEvent, + RequestEventInput, ResponseBody, ResponseEventInput, }; diff --git a/litellm-rust/crates/ai-gateway/src/ocr/common_utils.rs b/litellm-rust/crates/ai-gateway/src/ocr/common_utils.rs index 9bc2818b6e7..c146ef26228 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/common_utils.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/common_utils.rs @@ -20,19 +20,10 @@ use litellm_core::providers::vertex_ai::ocr::transformation::{ use crate::client::http_client; -const ERROR_BODY_MAX_CHARS: usize = 256; const AZURE_DOCUMENT_INTELLIGENCE_POLL_TIMEOUT_SECS: u64 = 120; const DEFAULT_MAX_IMAGE_URL_DOWNLOAD_SIZE_MB: f64 = 50.0; const MAX_SAFE_FETCH_REDIRECTS: usize = 10; -pub(super) fn truncate_error_body(body: &str) -> String { - if body.chars().count() <= ERROR_BODY_MAX_CHARS { - return body.to_string(); - } - let truncated: String = body.chars().take(ERROR_BODY_MAX_CHARS).collect(); - format!("{truncated}... (truncated)") -} - pub(super) fn ocr_provider_config( provider: &str, model: &str, @@ -269,7 +260,7 @@ pub(super) async fn convert_document_url_to_data_uri(document: Value) -> CoreRes let body = response.text().await.unwrap_or_default(); return Err(CoreError::Http { status: status.as_u16(), - body: truncate_error_body(&body), + body: litellm_core::utils::truncate_error_body(&body), }); } let content_type = response @@ -383,7 +374,7 @@ pub(super) async fn poll_document_intelligence( if !status.is_success() { return Err(CoreError::Http { status: status.as_u16(), - body: truncate_error_body(&text), + body: litellm_core::utils::truncate_error_body(&text), }); } let response_json: Value = serde_json::from_str(&text).map_err(|err| { diff --git a/litellm-rust/crates/ai-gateway/src/ocr/handler.rs b/litellm-rust/crates/ai-gateway/src/ocr/handler.rs index e5a3c79a29d..52b6d27a944 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/handler.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/handler.rs @@ -4,7 +4,7 @@ use litellm_core::logging::http::{JsonRequest, execute_json}; use litellm_core::ocr::transformation::OcrResponseHandling; use serde_json::Value; -use super::common_utils::{poll_document_intelligence, truncate_error_body}; +use super::common_utils::poll_document_intelligence; use super::types::ProviderOcrRequest; use crate::client::http_client; @@ -78,7 +78,7 @@ pub(crate) async fn execute_ocr_provider_call(request: ProviderOcrRequest) -> Co if !status.is_success() { return Err(CoreError::Http { status: status.as_u16(), - body: truncate_error_body(&text), + body: litellm_core::utils::truncate_error_body(&text), }); } diff --git a/litellm-rust/crates/ai-gateway/src/ocr/tests.rs b/litellm-rust/crates/ai-gateway/src/ocr/tests.rs index 7b4576516c5..68f1a195310 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/tests.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/tests.rs @@ -7,7 +7,7 @@ use serde_json::{Map, Value, json}; use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::net::{TcpListener, TcpStream}; -use super::common_utils::{has_header, ocr_provider_config, string_headers, truncate_error_body}; +use super::common_utils::{has_header, ocr_provider_config, string_headers}; use super::{OcrRequest, ocr}; use crate::integrations::custom_guardrail::{ CustomGuardrail, GuardrailContext, GuardrailDecision, GuardrailError, GuardrailEventHook, @@ -17,6 +17,7 @@ use crate::integrations::custom_logger::{ CallbackTiming, CallbackValue, CustomLogger, LogFuture, ModelCallDetails, }; use crate::integrations::types::RequestMetadata; +use litellm_core::utils::truncate_error_body; async fn read_http_headers(socket: &mut TcpStream) -> String { let mut request = Vec::new(); diff --git a/litellm-rust/crates/core/src/lib.rs b/litellm-rust/crates/core/src/lib.rs index a03bcba67c2..e80ea45fe1e 100644 --- a/litellm-rust/crates/core/src/lib.rs +++ b/litellm-rust/crates/core/src/lib.rs @@ -11,5 +11,6 @@ pub mod realtime; pub mod responses; pub mod router; pub mod routing_utils; +pub mod utils; pub use error::{CoreError, CoreResult}; diff --git a/litellm-rust/crates/core/src/logging/events.rs b/litellm-rust/crates/core/src/logging/events.rs index 7db390daa2c..0830cd7077b 100644 --- a/litellm-rust/crates/core/src/logging/events.rs +++ b/litellm-rust/crates/core/src/logging/events.rs @@ -7,7 +7,7 @@ use super::redaction::{redact_headers, redact_url, snapshot_json}; #[derive(Clone, Debug, Serialize)] #[serde(tag = "event")] -pub enum ProviderDebugEvent { +pub enum LogEvent { #[serde(rename = "provider.request")] Request(ProviderRequestEvent), #[serde(rename = "provider.response")] @@ -114,9 +114,9 @@ pub struct ProviderErrorEvent { pub body: Option, } -pub fn request_event(input: RequestEventInput) -> ProviderDebugEvent { +pub(crate) fn request_event(input: RequestEventInput) -> LogEvent { let snapshot = snapshot_json(input.body); - ProviderDebugEvent::Request(ProviderRequestEvent { + LogEvent::Request(ProviderRequestEvent { source: "litellm-rust", call_id: input.call_id, provider: input.provider, @@ -131,9 +131,9 @@ pub fn request_event(input: RequestEventInput) -> ProviderDebugEvent { }) } -pub fn response_event(input: ResponseEventInput) -> ProviderDebugEvent { +pub(crate) fn response_event(input: ResponseEventInput) -> LogEvent { let snapshot = input.body.snapshot(); - ProviderDebugEvent::Response(ProviderResponseEvent { + LogEvent::Response(ProviderResponseEvent { source: "litellm-rust", call_id: input.call_id, provider: input.provider, @@ -146,8 +146,8 @@ pub fn response_event(input: ResponseEventInput) -> ProviderDebugEvent { }) } -pub fn error_event(input: ErrorEventInput) -> ProviderDebugEvent { - ProviderDebugEvent::Error(ProviderErrorEvent { +pub(crate) fn error_event(input: ErrorEventInput) -> LogEvent { + LogEvent::Error(ProviderErrorEvent { source: "litellm-rust", call_id: input.call_id, provider: input.provider, @@ -159,13 +159,13 @@ pub fn error_event(input: ErrorEventInput) -> ProviderDebugEvent { }) } -pub fn stream_started( +pub(crate) fn stream_started( call_id: String, provider: String, status: u16, content_type: Option, -) -> ProviderDebugEvent { - ProviderDebugEvent::StreamStarted(ProviderStreamStartedEvent { +) -> LogEvent { + LogEvent::StreamStarted(ProviderStreamStartedEvent { source: "litellm-rust", call_id, provider, @@ -174,15 +174,15 @@ pub fn stream_started( }) } -pub fn stream_completed( +pub(crate) fn stream_completed( call_id: String, provider: String, duration_ms: u128, bytes_received: usize, frames_received: usize, events_decoded: usize, -) -> ProviderDebugEvent { - ProviderDebugEvent::StreamCompleted(ProviderStreamCompletedEvent { +) -> LogEvent { + LogEvent::StreamCompleted(ProviderStreamCompletedEvent { source: "litellm-rust", call_id, provider, diff --git a/litellm-rust/crates/core/src/logging/http.rs b/litellm-rust/crates/core/src/logging/http.rs index 06cf7671708..196e40b8e2f 100644 --- a/litellm-rust/crates/core/src/logging/http.rs +++ b/litellm-rust/crates/core/src/logging/http.rs @@ -1,3 +1,4 @@ +use std::sync::Arc; use std::time::Duration; use serde::de::DeserializeOwned; @@ -9,7 +10,7 @@ use crate::error::CoreError; use super::{CallLogger, ResponseBody}; pub struct JsonRequest { - pub logger: std::sync::Arc, + pub logger: Arc, pub model: String, pub stream: bool, pub url: String, @@ -24,13 +25,17 @@ pub async fn execute_json( ) -> CoreResult { let body_bytes = serde_json::to_vec(&request.body) .map_err(|error| CoreError::InvalidRequest(error.to_string()))?; - request.logger.request_about_to_be_sent( - request.model, - request.stream, - request.url.clone(), - request.headers.clone(), - request.body, - ); + request + .logger + .request_about_to_be_sent(super::RequestEventInput { + call_id: request.logger.call_id().to_string(), + provider: request.logger.provider().to_string(), + model: request.model, + stream: request.stream, + url: request.url.clone(), + headers: request.headers.clone(), + body: request.body, + }); let builder = request.headers.iter().fold( client.post(&request.url).body(body_bytes), |builder, (name, value)| builder.header(name, value), @@ -73,7 +78,7 @@ pub async fn execute_json( let body = serde_json::from_str(&text) .map(ResponseBody::Json) .unwrap_or(ResponseBody::Binary { - media_type, + media_type: media_type.clone(), bytes: text.len(), }); request.logger.failure( @@ -84,43 +89,45 @@ pub async fn execute_json( ); return Err(CoreError::Http { status: status.as_u16(), - body: crate::messages::common_utils::truncate_error_body(&text), + body: crate::utils::truncate_error_body(&text), }); } - if !media_type - .as_deref() - .is_some_and(|value| value.to_ascii_lowercase().contains("json")) - { - request.logger.response_received( - status.as_u16(), - headers, - ResponseBody::Binary { - media_type, - bytes: text.len(), - }, - ); - return Err(CoreError::InvalidResponse( - "provider response was not JSON".to_string(), - )); - } let value = serde_json::from_str::(&text).map_err(|error| { request.logger.failure( Some(status.as_u16()), "invalid_json", error.to_string(), Some(ResponseBody::Binary { - media_type, + media_type: media_type.clone(), bytes: text.len(), }), ); CoreError::InvalidResponse(format!("invalid provider response JSON: {error}")) })?; - let typed = T::deserialize(&value).map_err(|error| { - CoreError::InvalidResponse(format!("invalid provider response: {error}")) - })?; + let typed = T::deserialize(&value); + let response_body = if media_type + .as_deref() + .is_some_and(|value| value.to_ascii_lowercase().contains("json")) + { + ResponseBody::Json(value) + } else { + ResponseBody::Binary { + media_type, + bytes: text.len(), + } + }; request .logger - .response_received(status.as_u16(), headers, ResponseBody::Json(value)); + .response_received(status.as_u16(), headers, response_body); + let typed = typed.map_err(|error| { + request.logger.failure( + Some(status.as_u16()), + "invalid_json", + error.to_string(), + None, + ); + CoreError::InvalidResponse(format!("invalid provider response: {error}")) + })?; Ok(typed) } @@ -130,13 +137,17 @@ pub async fn execute_stream( ) -> CoreResult { let body_bytes = serde_json::to_vec(&request.body) .map_err(|error| CoreError::InvalidRequest(error.to_string()))?; - request.logger.request_about_to_be_sent( - request.model, - true, - request.url.clone(), - request.headers.clone(), - request.body, - ); + request + .logger + .request_about_to_be_sent(super::RequestEventInput { + call_id: request.logger.call_id().to_string(), + provider: request.logger.provider().to_string(), + model: request.model, + stream: true, + url: request.url.clone(), + headers: request.headers.clone(), + body: request.body, + }); let builder = request.headers.iter().fold( client.post(&request.url).body(body_bytes), |builder, (name, value)| builder.header(name, value), @@ -189,7 +200,7 @@ pub async fn execute_stream( ); return Err(CoreError::Http { status: status.as_u16(), - body: crate::messages::common_utils::truncate_error_body(&text), + body: crate::utils::truncate_error_body(&text), }); } let content_type = response diff --git a/litellm-rust/crates/core/src/logging/mod.rs b/litellm-rust/crates/core/src/logging/mod.rs index 3ac980c8036..5e4b3f3e2e0 100644 --- a/litellm-rust/crates/core/src/logging/mod.rs +++ b/litellm-rust/crates/core/src/logging/mod.rs @@ -11,7 +11,7 @@ use std::time::Instant; use crate::call_lifecycle::CallLifecycleContext; pub trait LogSink: Send + Sync { - fn emit(&self, event: &ProviderDebugEvent); + fn emit(&self, event: &LogEvent); } pub struct CallLogger { @@ -35,23 +35,18 @@ impl CallLogger { } } - pub fn request_about_to_be_sent( - &self, - model: String, - stream: bool, - url: String, - headers: Vec<(String, String)>, - body: serde_json::Value, - ) { - self.emit(events::request_event(events::RequestEventInput { - call_id: self.context.litellm_call_id.clone(), - provider: self.context.custom_llm_provider.clone(), - model, - stream, - url, - headers, - body, - })); + pub(crate) fn call_id(&self) -> &str { + &self.context.litellm_call_id + } + + pub(crate) fn provider(&self) -> &str { + &self.context.custom_llm_provider + } + + pub fn request_about_to_be_sent(&self, mut input: events::RequestEventInput) { + input.call_id = self.context.litellm_call_id.clone(); + input.provider = self.context.custom_llm_provider.clone(); + self.emit(events::request_event(input)); } pub fn response_received( @@ -114,7 +109,7 @@ impl CallLogger { self.events_decoded.fetch_add(events, Ordering::Relaxed); } - fn emit(&self, event: ProviderDebugEvent) { + fn emit(&self, event: LogEvent) { if let Some(sink) = &self.sink { sink.emit(&event); } @@ -122,10 +117,9 @@ impl CallLogger { } pub use events::{ - BodySnapshot, ErrorEventInput, ProviderDebugEvent, ProviderErrorEvent, ProviderRequestEvent, + BodySnapshot, ErrorEventInput, LogEvent, ProviderErrorEvent, ProviderRequestEvent, ProviderResponseEvent, ProviderStreamCompletedEvent, ProviderStreamStartedEvent, - RequestEventInput, ResponseBody, ResponseEventInput, error_event, request_event, - response_event, stream_completed, stream_started, + RequestEventInput, ResponseBody, ResponseEventInput, }; #[cfg(test)] @@ -135,10 +129,10 @@ mod tests { use super::*; #[derive(Clone, Default)] - struct RecordingSink(Arc>>); + struct RecordingSink(Arc>>); impl LogSink for RecordingSink { - fn emit(&self, event: &ProviderDebugEvent) { + fn emit(&self, event: &LogEvent) { self.0.lock().expect("recording lock").push(event.clone()); } } @@ -148,13 +142,15 @@ mod tests { let sink = RecordingSink::default(); let context = CallLifecycleContext::new("messages", "claude", "anthropic", "req_123"); let logger = CallLogger::new(&context, Some(Arc::new(sink.clone()))); - logger.request_about_to_be_sent( - "claude".to_string(), - false, - "https://example.test/v1/messages".to_string(), - vec![("authorization".to_string(), "Bearer secret".to_string())], - serde_json::json!({"token": "secret", "prompt": "visible"}), - ); + logger.request_about_to_be_sent(events::RequestEventInput { + call_id: String::new(), + provider: String::new(), + model: "claude".to_string(), + stream: false, + url: "https://example.test/v1/messages".to_string(), + headers: vec![("authorization".to_string(), "Bearer secret".to_string())], + body: serde_json::json!({"token": "secret", "prompt": "visible"}), + }); let events = sink.0.lock().expect("recording lock"); let serialized = serde_json::to_string(&events[0]).expect("event serializes"); diff --git a/litellm-rust/crates/core/src/logging/redaction.rs b/litellm-rust/crates/core/src/logging/redaction.rs index a97a3ecdf5a..22e07ffbc7e 100644 --- a/litellm-rust/crates/core/src/logging/redaction.rs +++ b/litellm-rust/crates/core/src/logging/redaction.rs @@ -117,3 +117,51 @@ fn is_secret_key(key: &str) -> bool { | "x-amz-security-token" ) } + +#[cfg(test)] +mod tests { + use serde_json::json; + + use super::{PROVIDER_DEBUG_BODY_MAX_BYTES, redact_headers, redact_url, snapshot_json}; + + #[test] + fn redacts_sensitive_headers() { + let headers = redact_headers(&[ + ("authorization".to_string(), "Bearer secret".to_string()), + ("content-type".to_string(), "application/json".to_string()), + ]); + assert_eq!(headers["authorization"], "[REDACTED]"); + assert_eq!(headers["content-type"], "application/json"); + } + + #[test] + fn redacts_sensitive_query_parameters() { + let url = redact_url( + "https://example.test/v1%3A0/invoke?X-Amz-Signature=secret&keep=value&key=hidden", + ); + assert!(url.contains("X-Amz-Signature=%5BREDACTED%5D")); + assert!(url.contains("key=%5BREDACTED%5D")); + assert!(url.contains("keep=value")); + assert!(url.contains("/v1%3A0/invoke")); + } + + #[test] + fn redacts_nested_body_keys() { + let snapshot = snapshot_json(json!({ + "outer": {"token": "secret", "visible": "keep"}, + "items": [{"password": "hidden"}] + })); + assert_eq!(snapshot.body["outer"]["token"], "[REDACTED]"); + assert_eq!(snapshot.body["outer"]["visible"], "keep"); + assert_eq!(snapshot.body["items"][0]["password"], "[REDACTED]"); + } + + #[test] + fn records_body_truncation_metadata() { + let snapshot = snapshot_json(json!({"content": "x".repeat(PROVIDER_DEBUG_BODY_MAX_BYTES)})); + assert_eq!(snapshot.body_truncated, Some(true)); + assert!( + snapshot.body_original_bytes.expect("original bytes") > PROVIDER_DEBUG_BODY_MAX_BYTES + ); + } +} diff --git a/litellm-rust/crates/core/src/logging/stream.rs b/litellm-rust/crates/core/src/logging/stream.rs index 358f1926242..0ba09c86845 100644 --- a/litellm-rust/crates/core/src/logging/stream.rs +++ b/litellm-rust/crates/core/src/logging/stream.rs @@ -1,3 +1,4 @@ +use std::fmt::Display; use std::sync::Arc; use bytes::Bytes; @@ -11,7 +12,8 @@ pub fn count_forwarded_stream( logger: Arc, ) -> impl Stream> where - S: Stream>, + S: Stream> + Send + 'static, + E: Display + Send + 'static, { futures_util::stream::unfold( (Box::pin(stream), logger, false, false), @@ -29,7 +31,8 @@ where .enumerate() .filter(|(index, byte)| { **byte == b'\n' - && ((*index > 0 && bytes[*index - 1] == b'\n') || trailing_newline) + && ((*index > 0 && bytes[*index - 1] == b'\n') + || (*index == 0 && trailing_newline)) }) .count(); let trailing_newline = bytes.last().copied() == Some(b'\n'); @@ -37,15 +40,76 @@ where Some((Ok(bytes), (stream, logger, trailing_newline, failed))) } Some(Err(error)) => { - logger.failure( - None, - "stream_error", - "provider stream failed".to_string(), - None, - ); + logger.failure(None, "stream_error", error.to_string(), None); Some((Err(error), (stream, logger, trailing_newline, true))) } } }, ) } + +#[cfg(test)] +mod tests { + use std::sync::{Arc, Mutex}; + + use bytes::Bytes; + use futures_util::StreamExt; + + use crate::call_lifecycle::CallLifecycleContext; + + use super::super::{CallLogger, LogEvent, LogSink}; + use super::count_forwarded_stream; + + #[derive(Clone, Default)] + struct Sink(Arc>>); + + impl LogSink for Sink { + fn emit(&self, event: &LogEvent) { + self.0.lock().expect("lock").push(event.clone()); + } + } + + fn logger(sink: Sink) -> Arc { + Arc::new(CallLogger::new( + &CallLifecycleContext::new("messages", "model", "anthropic", "call"), + Some(Arc::new(sink)), + )) + } + + #[tokio::test] + async fn counts_only_delimiters_and_handles_split_boundaries() { + let sink = Sink::default(); + let logger = logger(sink.clone()); + let stream = futures_util::stream::iter([ + Ok::<_, std::io::Error>(Bytes::from_static(b"data: a\n")), + Ok(Bytes::from_static(b"\ndata: b\nx\n")), + ]); + let chunks = count_forwarded_stream(stream, logger) + .collect::>() + .await; + assert_eq!(chunks.len(), 2); + let events = sink.0.lock().expect("lock"); + let LogEvent::StreamCompleted(completed) = events.last().expect("completion") else { + panic!("expected completion"); + }; + assert_eq!(completed.events_decoded, 1); + } + + #[tokio::test] + async fn newline_after_non_delimiter_does_not_count_from_carry() { + let sink = Sink::default(); + let logger = logger(sink.clone()); + let stream = futures_util::stream::iter([ + Ok::<_, std::io::Error>(Bytes::from_static(b"data: a\n")), + Ok(Bytes::from_static(b"x\n")), + ]); + let _ = count_forwarded_stream(stream, logger) + .collect::>() + .await; + let events = sink.0.lock().expect("lock"); + let LogEvent::StreamCompleted(completed) = events.last().expect("completion") else { + panic!("expected completion"); + }; + assert_eq!(completed.events_decoded, 0); + } +} diff --git a/litellm-rust/crates/core/src/messages/common_utils.rs b/litellm-rust/crates/core/src/messages/common_utils.rs index f6b4eea0dfe..daa53bcb527 100644 --- a/litellm-rust/crates/core/src/messages/common_utils.rs +++ b/litellm-rust/crates/core/src/messages/common_utils.rs @@ -1,20 +1,11 @@ use serde_json::{Map, Value}; -use crate::constants::MESSAGES_ERROR_BODY_MAX_CHARS; use crate::error::{CoreError, CoreResult, json_type_name}; use crate::providers::anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG; use crate::providers::azure_ai::messages::transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG; use super::transformation::AnthropicMessagesProviderConfig; -pub(crate) fn truncate_error_body(body: &str) -> String { - if body.chars().count() <= MESSAGES_ERROR_BODY_MAX_CHARS { - return body.to_string(); - } - let truncated: String = body.chars().take(MESSAGES_ERROR_BODY_MAX_CHARS).collect(); - format!("{truncated}... (truncated)") -} - pub(super) fn messages_provider_config( provider: &str, ) -> Option<&'static dyn AnthropicMessagesProviderConfig> { diff --git a/litellm-rust/crates/core/src/messages/tests.rs b/litellm-rust/crates/core/src/messages/tests.rs index 6ce34fb2ffb..69cfd8a7cbe 100644 --- a/litellm-rust/crates/core/src/messages/tests.rs +++ b/litellm-rust/crates/core/src/messages/tests.rs @@ -1,16 +1,18 @@ +use std::sync::{Arc, Mutex}; use std::time::Duration; +use crate::logging::{LogEvent, LogSink}; +use futures_util::StreamExt; use serde_json::{Map, Value, json}; use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::net::{TcpListener, TcpStream}; use crate::error::CoreError; +use crate::utils::truncate_error_body; -use super::common_utils::{ - has_bearer_auth, has_header, messages_provider_config, string_headers, truncate_error_body, -}; -use super::messages; +use super::common_utils::{has_bearer_auth, has_header, messages_provider_config, string_headers}; use super::types::MessagesRequest; +use super::{messages, messages_stream}; async fn read_http_request(socket: &mut TcpStream) -> String { let mut request = Vec::new(); @@ -455,3 +457,246 @@ async fn messages_rejects_unsupported_provider() { assert!(matches!(err, CoreError::InvalidProvider(provider) if provider == "openai")); } + +#[derive(Clone, Default)] +struct RecordingSink(Arc>>); + +impl LogSink for RecordingSink { + fn emit(&self, event: &LogEvent) { + self.0.lock().expect("recording lock").push(event.clone()); + } +} + +#[tokio::test] +async fn messages_debug_logs_transformed_request_and_redacted_response() { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); + let addr = listener.local_addr().expect("addr"); + let server = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.expect("accepts request"); + let request = read_http_request(&mut socket).await; + let response = r#"{"id":"msg_1","type":"message","role":"assistant","content":[{"type":"text","text":"visible"}],"model":"claude","stop_reason":"end_turn","usage":{"input_tokens":1,"output_tokens":1},"token":"response-secret"}"#; + socket + .write_all(write_response(response).as_bytes()) + .await + .expect("writes response"); + request + }); + let sink = RecordingSink::default(); + let response = messages(MessagesRequest { + model: "claude", + body: json!({ + "model": "claude", + "max_tokens": 8, + "messages": [{"role": "user", "content": "prompt visible", "token": "request-secret"}] + }), + api_key: Some("request-secret"), + api_base: Some(&format!("http://{addr}")), + custom_llm_provider: Some("anthropic"), + extra_headers: Some( + json!({ + "Authorization": "Bearer bearer-secret", + "X-Amz-Security-Token": "session-secret" + }) + .as_object() + .expect("headers") + .clone(), + ), + timeout: Some(Duration::from_secs(5)), + litellm_call_id: Some("debug-success"), + logging_sink: Some(Arc::new(sink.clone())), + }) + .await + .expect("messages succeeds"); + assert_eq!(response.content[0]["text"], "visible"); + let request = server.await.expect("server task"); + let events = sink.0.lock().expect("recording lock"); + let serialized = serde_json::to_string(&*events).expect("events serialize"); + assert!(serialized.contains("prompt visible")); + assert!(!serialized.contains("request-secret")); + assert!(!serialized.contains("bearer-secret")); + assert!(!serialized.contains("session-secret")); + assert!(!serialized.contains("response-secret")); + let LogEvent::Request(request_event) = &events[0] else { + panic!("request event first"); + }; + assert!(request_event.url.ends_with("/v1/messages")); + assert_eq!(request_event.headers["content-type"], "application/json"); + let sent_body = request.split_once("\r\n\r\n").expect("request body").1; + assert!(sent_body.contains("prompt visible")); + let response_event = events + .iter() + .find_map(|event| match event { + LogEvent::Response(value) => Some(value), + _ => None, + }) + .expect("response event"); + assert_eq!(response_event.status, 200); + assert!(response_event.duration_ms < 10_000); +} + +#[tokio::test] +async fn messages_debug_http_failure_emits_one_error_and_none_emits_nothing() { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); + let addr = listener.local_addr().expect("addr"); + let server = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.expect("accepts request"); + let _ = read_http_request(&mut socket).await; + let body = r#"{"token":"provider-secret","error":"no"}"#; + let response = format!( + "HTTP/1.1 401 Unauthorized\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", + body.len(), + body + ); + socket + .write_all(response.as_bytes()) + .await + .expect("writes response"); + }); + let sink = RecordingSink::default(); + let err = messages(MessagesRequest { + model: "claude", + body: json!({"model": "claude", "max_tokens": 8, "messages": []}), + api_key: Some("secret"), + api_base: Some(&format!("http://{addr}")), + custom_llm_provider: Some("anthropic"), + extra_headers: None, + timeout: Some(Duration::from_secs(5)), + litellm_call_id: Some("debug-error"), + logging_sink: Some(Arc::new(sink.clone())), + }) + .await + .expect_err("request fails"); + assert!(matches!(err, CoreError::Http { status: 401, .. })); + server.await.expect("server task"); + { + let events = sink.0.lock().expect("recording lock"); + assert_eq!( + events + .iter() + .filter(|event| matches!(event, LogEvent::Error(_))) + .count(), + 1 + ); + assert!( + !events + .iter() + .any(|event| matches!(event, LogEvent::Response(_))) + ); + } + + let none_sink = messages(MessagesRequest { + model: "claude", + body: json!({"model": "claude", "max_tokens": 8, "messages": []}), + api_key: Some("secret"), + api_base: Some("http://127.0.0.1:1"), + custom_llm_provider: Some("anthropic"), + extra_headers: None, + timeout: Some(Duration::from_millis(20)), + litellm_call_id: Some("debug-disabled"), + logging_sink: None, + }) + .await + .expect_err("disabled test endpoint fails"); + assert!(matches!(none_sink, CoreError::Network(_))); +} + +#[tokio::test] +async fn messages_debug_stream_counts_clean_and_interrupted_streams() { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); + let addr = listener.local_addr().expect("addr"); + let server = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.expect("accepts request"); + let _ = read_http_request(&mut socket).await; + let body = b"data: one\n\ndata: two\n\n"; + let response = format!( + "HTTP/1.1 200 OK\r\ncontent-type: text/event-stream\r\ncontent-length: {}\r\nconnection: close\r\n\r\n", + body.len() + ); + socket + .write_all(response.as_bytes()) + .await + .expect("headers"); + socket.write_all(body).await.expect("body"); + }); + let sink = RecordingSink::default(); + let upstream = messages_stream(MessagesRequest { + model: "claude", + body: json!({"model": "claude", "max_tokens": 8, "stream": true, "messages": []}), + api_key: Some("secret"), + api_base: Some(&format!("http://{addr}")), + custom_llm_provider: Some("anthropic"), + extra_headers: None, + timeout: Some(Duration::from_secs(5)), + litellm_call_id: Some("debug-stream"), + logging_sink: Some(Arc::new(sink.clone())), + }) + .await + .expect("stream starts"); + let chunks = crate::logging::stream::count_forwarded_stream( + upstream.response.bytes_stream(), + upstream.logger, + ); + let forwarded = futures_util::StreamExt::collect::>(chunks).await; + assert!(!forwarded.is_empty()); + assert!(forwarded.iter().all(Result::is_ok)); + server.await.expect("server task"); + let events = sink.0.lock().expect("recording lock"); + let completed = events + .iter() + .find_map(|event| match event { + LogEvent::StreamCompleted(value) => Some(value), + _ => None, + }) + .expect("completion"); + assert!(completed.bytes_received > 0); + assert!(completed.frames_received > 0); + assert!(completed.events_decoded > 0); +} + +#[tokio::test] +async fn messages_debug_interrupted_stream_emits_error_without_completion() { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); + let addr = listener.local_addr().expect("addr"); + let server = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.expect("accepts request"); + let _ = read_http_request(&mut socket).await; + socket + .write_all( + b"HTTP/1.1 200 OK\r\ncontent-type: text/event-stream\r\ncontent-length: 100\r\nconnection: close\r\n\r\nshort", + ) + .await + .expect("writes partial response"); + }); + let sink = RecordingSink::default(); + let upstream = messages_stream(MessagesRequest { + model: "claude", + body: json!({"model": "claude", "max_tokens": 8, "stream": true, "messages": []}), + api_key: Some("secret"), + api_base: Some(&format!("http://{addr}")), + custom_llm_provider: Some("anthropic"), + extra_headers: None, + timeout: Some(Duration::from_secs(5)), + litellm_call_id: Some("debug-interrupted"), + logging_sink: Some(Arc::new(sink.clone())), + }) + .await + .expect("stream starts"); + let results = crate::logging::stream::count_forwarded_stream( + upstream.response.bytes_stream(), + upstream.logger, + ) + .collect::>() + .await; + assert!(results.iter().any(Result::is_err)); + server.await.expect("server task"); + let events = sink.0.lock().expect("recording lock"); + assert!(events.iter().any(|event| matches!( + event, + LogEvent::Error(value) if value.kind == "stream_error" + ))); + assert!( + !events + .iter() + .any(|event| matches!(event, LogEvent::StreamCompleted(_))) + ); +} diff --git a/litellm-rust/crates/core/src/utils.rs b/litellm-rust/crates/core/src/utils.rs new file mode 100644 index 00000000000..c5beac28823 --- /dev/null +++ b/litellm-rust/crates/core/src/utils.rs @@ -0,0 +1,9 @@ +use crate::constants::MESSAGES_ERROR_BODY_MAX_CHARS; + +pub fn truncate_error_body(body: &str) -> String { + if body.chars().count() <= MESSAGES_ERROR_BODY_MAX_CHARS { + return body.to_string(); + } + let truncated: String = body.chars().take(MESSAGES_ERROR_BODY_MAX_CHARS).collect(); + format!("{truncated}... (truncated)") +}