use oxicode_ai::{ContentBlock, Message, MessageContent, TextContent};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ShakeConfig {
pub protect_window_tokens: usize,
pub min_elidable_tokens: usize,
pub min_savings_tokens: usize,
}
impl Default for ShakeConfig {
fn default() -> Self {
Self {
protect_window_tokens: 16_384,
min_elidable_tokens: 400,
min_savings_tokens: 4_096,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ShakeOutcome {
Shaken {
regions_elided: usize,
tokens_saved: usize,
},
NoChange,
}
#[inline]
fn estimate_tokens(text: &str) -> usize {
text.chars().count() / 4
}
#[inline]
fn estimate_block_tokens(block: &ContentBlock) -> usize {
match block {
ContentBlock::Text(t) => estimate_tokens(&t.text),
ContentBlock::Thinking(t) => estimate_tokens(&t.thinking),
ContentBlock::Image(_) => 8,
ContentBlock::ToolCall(tc) => (tc.name.chars().count() / 4) + 12,
ContentBlock::Unknown(_) => 10,
}
}
fn estimate_message_content_tokens(content: &MessageContent) -> usize {
match content {
MessageContent::Text(s) => estimate_tokens(s),
MessageContent::Blocks(blocks) => blocks.iter().map(estimate_block_tokens).sum(),
}
}
fn estimate_message_tokens(message: &Message) -> usize {
match message {
Message::User(m) => estimate_message_content_tokens(&m.content),
Message::Assistant(m) => m.content.iter().map(estimate_block_tokens).sum(),
Message::ToolResult(m) => match m.text_content() {
Ok(text) => estimate_tokens(&text),
Err(_) => m.content.iter().map(estimate_block_tokens).sum(),
},
}
}
fn find_protect_boundary(messages: &[Message], protect_window_tokens: usize) -> usize {
if messages.is_empty() || protect_window_tokens == 0 {
return 0;
}
let mut accumulated_after: usize = 0;
for (idx, message) in messages.iter().enumerate().rev() {
if accumulated_after >= protect_window_tokens {
return idx + 1;
}
accumulated_after = accumulated_after.saturating_add(estimate_message_tokens(message));
}
messages.len()
}
#[derive(Debug, Clone)]
enum Candidate {
ToolResult {
index: usize,
tokens_saved: usize,
},
CodeBlock {
index: usize,
block_start: usize,
block_end: usize,
placeholder: String,
tokens_saved: usize,
},
}
fn collect_candidates(
messages: &[Message],
boundary: usize,
min_elidable_tokens: usize,
) -> Vec<Candidate> {
let mut candidates: Vec<Candidate> = Vec::new();
for (index, message) in messages[..boundary].iter().enumerate() {
match message {
Message::ToolResult(m) => {
let original_tokens = match m.text_content() {
Ok(text) => estimate_tokens(&text),
Err(_) => m.content.iter().map(estimate_block_tokens).sum(),
};
if original_tokens >= min_elidable_tokens {
let placeholder = format!("[tool result elided (~{original_tokens} tokens)]");
let placeholder_tokens = estimate_tokens(&placeholder);
let tokens_saved = original_tokens.saturating_sub(placeholder_tokens);
candidates.push(Candidate::ToolResult {
index,
tokens_saved,
});
}
}
Message::User(m) => {
collect_text_candidates(&m.content, index, min_elidable_tokens, &mut candidates);
}
Message::Assistant(m) => {
for block in &m.content {
if let ContentBlock::Text(t) = block {
let content = MessageContent::Text(t.text.clone());
collect_text_candidates(
&content,
index,
min_elidable_tokens,
&mut candidates,
);
}
}
}
}
}
candidates
}
fn collect_text_candidates(
content: &MessageContent,
index: usize,
min_elidable_tokens: usize,
out: &mut Vec<Candidate>,
) {
let Some(text) = content.as_str() else {
return;
};
for_each_elidable_code_block(
text,
min_elidable_tokens,
|block_start, block_end, body, lines| {
let body_tokens = estimate_tokens(body);
let placeholder = format!("\n```\n...code block elided ({lines} lines)...\n```\n");
let placeholder_tokens = estimate_tokens(&placeholder);
let tokens_saved = body_tokens.saturating_sub(placeholder_tokens);
if tokens_saved == 0 {
return;
}
out.push(Candidate::CodeBlock {
index,
block_start,
block_end,
placeholder,
tokens_saved,
});
},
);
}
fn for_each_elidable_code_block(
text: &str,
min_elidable_tokens: usize,
mut f: impl FnMut(usize, usize, &str, usize),
) {
let bytes = text.as_bytes();
let mut search_from = 0usize;
while let Some(open_rel) = find_fence_open(bytes, search_from) {
let open_start = search_from + open_rel;
let Some(close_rel) = find_fence_close(bytes, open_start + 3) else {
return;
};
let close_end = open_start + 3 + close_rel + 3;
let body_start = open_start + 3;
let body = &text[body_start..close_end - 3];
let tokens = estimate_tokens(body);
if tokens >= min_elidable_tokens {
let lines = body.lines().count();
f(open_start, close_end, body, lines);
}
search_from = close_end;
}
}
fn find_fence_open(bytes: &[u8], from: usize) -> Option<usize> {
if from + 2 >= bytes.len() {
return None;
}
let mut idx = from;
while idx + 2 < bytes.len() {
if bytes[idx] == b'`' && bytes[idx + 1] == b'`' && bytes[idx + 2] == b'`' {
return Some(idx - from);
}
idx += 1;
}
None
}
fn find_fence_close(bytes: &[u8], search_from: usize) -> Option<usize> {
if search_from + 2 >= bytes.len() {
return None;
}
let mut idx = search_from;
while idx + 2 < bytes.len() {
if bytes[idx] == b'`' && bytes[idx + 1] == b'`' && bytes[idx + 2] == b'`' {
return Some(idx - search_from);
}
idx += 1;
}
None
}
fn replace_code_block(
text: &str,
block_start: usize,
block_end: usize,
placeholder: &str,
) -> String {
let mut out = String::with_capacity(text.len());
out.push_str(&text[..block_start]);
out.push_str(placeholder);
out.push_str(&text[block_end..]);
out
}
fn apply_candidates(messages: &mut [Message], mut candidates: Vec<Candidate>) {
candidates.sort_by_key(|c| std::cmp::Reverse(c.index()));
for candidate in candidates {
apply_one(messages, candidate);
}
}
impl Candidate {
fn index(&self) -> usize {
match self {
Candidate::ToolResult { index, .. } => *index,
Candidate::CodeBlock { index, .. } => *index,
}
}
}
fn apply_one(messages: &mut [Message], candidate: Candidate) {
match candidate {
Candidate::ToolResult { index, .. } => {
let Some(Message::ToolResult(tr)) = messages.get_mut(index) else {
return;
};
let original_tokens = match tr.text_content() {
Ok(text) => estimate_tokens(&text),
Err(_) => tr.content.iter().map(estimate_block_tokens).sum(),
};
let placeholder = format!("[tool result elided (~{original_tokens} tokens)]");
tr.content = vec![ContentBlock::Text(TextContent::new(placeholder))];
}
Candidate::CodeBlock {
index,
block_start,
block_end,
placeholder,
..
} => {
let Some(message) = messages.get_mut(index) else {
return;
};
match message {
Message::User(m) => {
rewrite_message_content(&mut m.content, block_start, block_end, &placeholder)
}
Message::Assistant(m) => {
for block in &mut m.content {
if let ContentBlock::Text(t) = block
&& t.text.len() >= block_end
{
t.text =
replace_code_block(&t.text, block_start, block_end, &placeholder);
return;
}
}
}
Message::ToolResult(_) => {
}
}
}
}
}
fn rewrite_message_content(
content: &mut MessageContent,
block_start: usize,
block_end: usize,
placeholder: &str,
) {
match content {
MessageContent::Text(s) => {
if s.len() >= block_end {
*s = replace_code_block(s, block_start, block_end, placeholder);
}
}
MessageContent::Blocks(blocks) => {
for block in blocks {
if let ContentBlock::Text(t) = block
&& t.text.len() >= block_end
{
t.text = replace_code_block(&t.text, block_start, block_end, placeholder);
return;
}
}
}
}
}
#[allow(clippy::ptr_arg)]
pub fn shake(messages: &mut Vec<Message>, config: &ShakeConfig) -> ShakeOutcome {
let boundary = find_protect_boundary(messages, config.protect_window_tokens);
if boundary == 0 {
return ShakeOutcome::NoChange;
}
let candidates = collect_candidates(messages, boundary, config.min_elidable_tokens);
if candidates.is_empty() {
return ShakeOutcome::NoChange;
}
let total_savings: usize = candidates.iter().map(Candidate::tokens_saved).sum();
if total_savings < config.min_savings_tokens {
return ShakeOutcome::NoChange;
}
let regions_elided = candidates.len();
apply_candidates(messages, candidates);
ShakeOutcome::Shaken {
regions_elided,
tokens_saved: total_savings,
}
}
impl Candidate {
fn tokens_saved(&self) -> usize {
match self {
Candidate::ToolResult { tokens_saved, .. } => *tokens_saved,
Candidate::CodeBlock { tokens_saved, .. } => *tokens_saved,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use oxicode_ai::{Api, AssistantMessage, ToolResultMessage, UserMessage};
fn user_msg(text: &str) -> Message {
Message::User(UserMessage::new(text.to_string()))
}
#[allow(dead_code)]
fn assistant_msg(text: &str) -> Message {
let mut msg = AssistantMessage::new(Api::AnthropicMessages, "mock", "test-model");
msg.content.push(ContentBlock::Text(TextContent::new(text)));
Message::Assistant(msg)
}
fn tool_result_msg(tool_call_id: &str, tool_name: &str, text: &str) -> Message {
Message::ToolResult(ToolResultMessage::new(
tool_call_id.to_string(),
tool_name.to_string(),
vec![ContentBlock::Text(TextContent::new(text.to_string()))],
))
}
fn chars(n: usize) -> String {
"a".repeat(n)
}
const CFG: ShakeConfig = ShakeConfig {
protect_window_tokens: 100,
min_elidable_tokens: 50,
min_savings_tokens: 200,
};
#[test]
fn test_shake_elides_large_tool_result() {
let mut messages = vec![
tool_result_msg("call-1", "search", &chars(8_000)),
user_msg(&chars(500)),
];
let outcome = shake(&mut messages, &CFG);
match outcome {
ShakeOutcome::Shaken {
regions_elided,
tokens_saved,
} => {
assert_eq!(regions_elided, 1);
assert!(tokens_saved >= 200);
}
other => panic!("expected Shaken, got {other:?}"),
}
match &messages[0] {
Message::ToolResult(tr) => {
assert_eq!(tr.content.len(), 1);
let rendered = tr.text_content().expect("renderable");
assert!(
rendered.contains("tool result elided"),
"unexpected tool result text: {rendered:?}"
);
}
other => panic!("expected ToolResult variant, got {other:?}"),
}
}
#[test]
fn test_shake_preserves_protect_window() {
let mut messages = vec![
tool_result_msg("call-1", "search", &chars(8_000)),
tool_result_msg("call-2", "search", &chars(8_000)),
tool_result_msg("call-3", "search", &chars(8_000)),
tool_result_msg("call-4", "search", &chars(8_000)),
];
let snapshot_before: Vec<String> = messages
.iter()
.map(|m| match m {
Message::ToolResult(tr) => tr.text_content().unwrap_or_default(),
_ => String::new(),
})
.collect();
let cfg = ShakeConfig {
protect_window_tokens: 2_500,
..CFG
};
let outcome = shake(&mut messages, &cfg);
assert!(matches!(outcome, ShakeOutcome::Shaken { .. }));
let len = messages.len();
for (idx, original) in snapshot_before.iter().enumerate().rev().take(2) {
let preserved = match &messages[idx] {
Message::ToolResult(tr) => tr.text_content().unwrap_or_default(),
_ => panic!("expected ToolResult at index {idx}"),
};
assert_eq!(
&preserved, original,
"tool result at index {idx} was mutated but should be inside the protect window"
);
}
for (msg, idx) in messages.iter().take(len - 2).zip(0..) {
let rendered = match msg {
Message::ToolResult(tr) => tr.text_content().unwrap_or_default(),
_ => panic!("expected ToolResult at index {idx}"),
};
assert!(
rendered.contains("tool result elided"),
"tool result at index {idx} was not elided: {rendered:?}"
);
}
}
#[test]
fn test_shake_no_change_when_insufficient_savings() {
let mut messages = vec![
tool_result_msg("call-1", "echo", &chars(200)),
tool_result_msg("call-2", "echo", &chars(200)),
tool_result_msg("call-3", "echo", &chars(200)),
];
let snapshot_before: Vec<String> = messages
.iter()
.map(|m| match m {
Message::ToolResult(tr) => tr.text_content().unwrap_or_default(),
_ => String::new(),
})
.collect();
let cfg = ShakeConfig {
protect_window_tokens: 50,
min_elidable_tokens: 40,
min_savings_tokens: 5_000,
};
let outcome = shake(&mut messages, &cfg);
assert_eq!(outcome, ShakeOutcome::NoChange);
let snapshot_after: Vec<String> = messages
.iter()
.map(|m| match m {
Message::ToolResult(tr) => tr.text_content().unwrap_or_default(),
_ => String::new(),
})
.collect();
assert_eq!(snapshot_before, snapshot_after);
}
#[test]
fn test_shake_elides_large_code_block() {
let code_body = chars(8_000);
let user_text = format!("x\n```rust\n{code_body}\n```\n");
let mut messages = vec![user_msg(&user_text), user_msg(&chars(500))];
let outcome = shake(&mut messages, &CFG);
match outcome {
ShakeOutcome::Shaken {
regions_elided,
tokens_saved,
} => {
assert_eq!(regions_elided, 1);
assert!(tokens_saved >= 200);
}
other => panic!("expected Shaken, got {other:?}"),
}
let rendered = messages[0].text_content().expect("renderable");
assert!(
rendered.contains("code block elided"),
"code block was not replaced; got: {rendered:?}"
);
assert!(
!rendered.contains(&code_body),
"original code body should have been replaced"
);
}
}