use std::collections::VecDeque;
use thiserror::Error;
use crate::{
Message, ToolSpec,
history::{char_tail, truncate_chars},
};
#[derive(Debug, Clone)]
pub struct ContextConfig {
pub max_tokens: usize,
pub reserve_output_tokens: usize,
pub safety_margin_tokens: usize,
pub bytes_per_token: usize,
pub summary_max_chars: usize,
}
impl Default for ContextConfig {
fn default() -> Self {
Self {
max_tokens: 128_000,
reserve_output_tokens: 8_192,
safety_margin_tokens: 2_048,
bytes_per_token: 3,
summary_max_chars: 6_000,
}
}
}
#[derive(Debug, Clone)]
pub struct ContextSelection {
pub messages: Vec<Message>,
pub before_tokens: usize,
pub after_tokens: usize,
pub removed_messages: usize,
}
#[derive(Debug, Error)]
#[error("{0}")]
pub struct ContextError(pub String);
pub trait ContextPolicy: Send + Sync {
fn select(
&self,
history: &[Message],
system_prompt: &str,
tools: &[ToolSpec],
) -> Result<ContextSelection, ContextError>;
}
pub struct BudgetContextPolicy {
config: ContextConfig,
}
impl BudgetContextPolicy {
pub fn new(config: ContextConfig) -> Result<Self, ContextError> {
if config.bytes_per_token == 0 {
return Err(ContextError(
"context.bytes_per_token must be positive".into(),
));
}
if config
.reserve_output_tokens
.saturating_add(config.safety_margin_tokens)
>= config.max_tokens
{
return Err(ContextError(
"context reserve and safety margin consume the model window".into(),
));
}
Ok(Self { config })
}
fn string_tokens(&self, value: &str) -> usize {
value.len().div_ceil(self.config.bytes_per_token)
}
fn cost(&self, messages: &[Message]) -> usize {
messages
.iter()
.map(|message| message.estimated_tokens(self.config.bytes_per_token))
.sum()
}
fn group_messages(history: &[Message]) -> Vec<&[Message]> {
let mut groups = Vec::new();
let mut start = 0;
for (index, message) in history.iter().enumerate() {
if index > start && matches!(message, Message::User { .. }) {
groups.push(&history[start..index]);
start = index;
}
}
if start < history.len() {
groups.push(&history[start..]);
}
groups
}
fn summarize(&self, messages: &[Message]) -> String {
let mut output = format!(
"[SCV compacted {} earlier messages. Bounded extracts follow.]\n",
messages.len()
);
for message in messages {
let (label, content) = match message {
Message::User { content, .. } => ("user", content.as_str()),
Message::Assistant { content, .. } => ("assistant", content.as_str()),
Message::Tool {
name,
content,
is_error,
..
} => {
let status = if *is_error { "failed" } else { "ok" };
output.push_str(&format!("tool {name} ({status}): "));
("", content.as_str())
}
Message::HistoryNote { content } => ("earlier", content.as_str()),
};
if !label.is_empty() {
output.push_str(label);
output.push_str(": ");
}
let tail = char_tail(content, 240);
output.push_str(&tail.replace('\n', " "));
output.push('\n');
if output.chars().count() >= self.config.summary_max_chars {
break;
}
}
truncate_chars(&output, self.config.summary_max_chars)
}
}
impl ContextPolicy for BudgetContextPolicy {
fn select(
&self,
history: &[Message],
system_prompt: &str,
tools: &[ToolSpec],
) -> Result<ContextSelection, ContextError> {
if history.is_empty() {
return Ok(ContextSelection {
messages: Vec::new(),
before_tokens: 0,
after_tokens: 0,
removed_messages: 0,
});
}
let tools_bytes = serde_json::to_vec(tools).map_or(0, |value| value.len());
let static_tokens = self
.string_tokens(system_prompt)
.saturating_add(tools_bytes.div_ceil(self.config.bytes_per_token))
.saturating_add(self.config.reserve_output_tokens)
.saturating_add(self.config.safety_margin_tokens);
if static_tokens >= self.config.max_tokens {
return Err(ContextError(
"system prompt and tool schemas exceed context budget".into(),
));
}
let budget = self.config.max_tokens - static_tokens;
let before_tokens = static_tokens.saturating_add(self.cost(history));
let groups = Self::group_messages(history);
let (newest, older) = groups
.split_last()
.expect("a non-empty history has at least one group");
let mut selected_cost = self.cost(newest);
if selected_cost > budget {
return Err(ContextError("newest turn exceeds context budget".into()));
}
let mut selected: VecDeque<&[Message]> = VecDeque::from([*newest]);
for group in older.iter().rev() {
let cost = self.cost(group);
if selected_cost.saturating_add(cost) > budget {
break;
}
selected.push_front(group);
selected_cost += cost;
}
let kept_messages: usize = selected.iter().map(|group| group.len()).sum();
let mut removed_messages = history.len() - kept_messages;
let selection = |note: Option<Message>,
selected: VecDeque<&[Message]>,
cost: usize,
removed_messages: usize| ContextSelection {
messages: note
.into_iter()
.chain(selected.into_iter().flatten().cloned())
.collect(),
before_tokens,
after_tokens: static_tokens.saturating_add(cost),
removed_messages,
};
if removed_messages == 0 {
return Ok(selection(None, selected, selected_cost, 0));
}
loop {
let summary = self.summarize(&history[..removed_messages]);
let note = Message::HistoryNote {
content: summary.clone(),
};
let note_cost = note.estimated_tokens(self.config.bytes_per_token);
if selected_cost.saturating_add(note_cost) <= budget {
return Ok(selection(
Some(note),
selected,
selected_cost + note_cost,
removed_messages,
));
}
if selected.len() == 1 {
let available_tokens = budget.saturating_sub(selected_cost);
let note =
fit_history_note(&summary, available_tokens, self.config.bytes_per_token)
.ok_or_else(|| {
ContextError("compaction note cannot fit context budget".into())
})?;
let note_cost = note.estimated_tokens(self.config.bytes_per_token);
return Ok(selection(
Some(note),
selected,
selected_cost + note_cost,
removed_messages,
));
}
let dropped = selected
.pop_front()
.expect("more than one group is selected");
selected_cost = selected_cost.saturating_sub(self.cost(dropped));
removed_messages += dropped.len();
}
}
}
fn fit_history_note(
content: &str,
available_tokens: usize,
bytes_per_token: usize,
) -> Option<Message> {
let chars: Vec<char> = content.chars().collect();
let mut low = 0usize;
let mut high = chars.len();
let mut best = None;
while low <= high {
let middle = low + (high - low) / 2;
let candidate = Message::HistoryNote {
content: chars[..middle].iter().collect(),
};
if candidate.estimated_tokens(bytes_per_token) <= available_tokens {
best = Some(candidate);
low = middle.saturating_add(1);
} else if middle == 0 {
break;
} else {
high = middle - 1;
}
}
best
}
#[cfg(test)]
mod tests;