mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-13 23:11:40 +00:00
fix(rust): harden provider debug logging
Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>
This commit is contained in:
parent
ce2891602d
commit
95d594a0fe
14 changed files with 499 additions and 135 deletions
|
|
@ -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<Mutex<Vec<u8>>>);
|
||||
|
||||
|
|
@ -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 =
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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| {
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
});
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
|
|
|||
|
|
@ -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<Value>,
|
||||
}
|
||||
|
||||
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<String>,
|
||||
) -> 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,
|
||||
|
|
|
|||
|
|
@ -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<CallLogger>,
|
||||
pub logger: Arc<CallLogger>,
|
||||
pub model: String,
|
||||
pub stream: bool,
|
||||
pub url: String,
|
||||
|
|
@ -24,13 +25,17 @@ pub async fn execute_json<T: DeserializeOwned>(
|
|||
) -> CoreResult<T> {
|
||||
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<T: DeserializeOwned>(
|
|||
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<T: DeserializeOwned>(
|
|||
);
|
||||
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::<Value>(&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<reqwest::Response> {
|
||||
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
|
||||
|
|
|
|||
|
|
@ -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<Mutex<Vec<ProviderDebugEvent>>>);
|
||||
struct RecordingSink(Arc<Mutex<Vec<LogEvent>>>);
|
||||
|
||||
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");
|
||||
|
|
|
|||
|
|
@ -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
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
use std::fmt::Display;
|
||||
use std::sync::Arc;
|
||||
|
||||
use bytes::Bytes;
|
||||
|
|
@ -11,7 +12,8 @@ pub fn count_forwarded_stream<S, E>(
|
|||
logger: Arc<CallLogger>,
|
||||
) -> impl Stream<Item = Result<Bytes, E>>
|
||||
where
|
||||
S: Stream<Item = Result<Bytes, E>>,
|
||||
S: Stream<Item = Result<Bytes, E>> + 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<Mutex<Vec<LogEvent>>>);
|
||||
|
||||
impl LogSink for Sink {
|
||||
fn emit(&self, event: &LogEvent) {
|
||||
self.0.lock().expect("lock").push(event.clone());
|
||||
}
|
||||
}
|
||||
|
||||
fn logger(sink: Sink) -> Arc<CallLogger> {
|
||||
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::<Vec<_>>()
|
||||
.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::<Vec<_>>()
|
||||
.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);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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> {
|
||||
|
|
|
|||
|
|
@ -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<Mutex<Vec<LogEvent>>>);
|
||||
|
||||
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::<Vec<_>>(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::<Vec<_>>()
|
||||
.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(_)))
|
||||
);
|
||||
}
|
||||
|
|
|
|||
9
litellm-rust/crates/core/src/utils.rs
Normal file
9
litellm-rust/crates/core/src/utils.rs
Normal file
|
|
@ -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)")
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue