fix(rust): harden provider debug logging

Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-07-29 21:52:56 +00:00
parent ce2891602d
commit 95d594a0fe
14 changed files with 499 additions and 135 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

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