use super::{ToolDispatchContext, ToolDispatchResult, ToolMiddleware, ToolPipeline};
use crate::message::{ToolContent, ToolContentPart};
use std::future::Future;
use std::pin::Pin;
pub struct OutputLimitMiddleware {
max_chars: usize,
}
const TRUNCATION_MARKER: &str = "\n[truncated]";
fn truncate_marked(text: &str, budget: usize) -> String {
let marker_len = TRUNCATION_MARKER.chars().count();
let kept = budget.saturating_sub(marker_len);
let truncated: String = text.chars().take(kept).collect();
format!("{truncated}{TRUNCATION_MARKER}")
}
impl OutputLimitMiddleware {
#[must_use]
pub fn new(max_chars: usize) -> Self {
Self { max_chars }
}
}
impl ToolMiddleware for OutputLimitMiddleware {
fn name(&self) -> &'static str {
"output_limit"
}
fn dispatch<'a>(
&'a self,
ctx: &'a mut ToolDispatchContext,
next: &'a ToolPipeline,
) -> Pin<Box<dyn Future<Output = ToolDispatchResult> + Send + 'a>> {
let max_chars = self.max_chars;
Box::pin(async move {
let mut result = next.dispatch(ctx).await;
if max_chars == 0 {
return result;
}
match result.output {
ToolContent::Text(ref text) => {
let char_count = text.chars().count();
if char_count > max_chars {
result.output = ToolContent::Text(truncate_marked(text, max_chars));
}
}
ToolContent::Multipart(ref mut parts) => {
let mut remaining = max_chars;
for part in parts.iter_mut() {
if let ToolContentPart::Text { text } = part {
let char_count = text.chars().count();
if char_count > remaining {
if remaining == 0 {
text.clear();
} else {
*text = truncate_marked(text, remaining);
}
remaining = 0;
} else {
remaining = remaining.saturating_sub(char_count);
}
}
}
}
}
result
})
}
}