use crate::message::Message;
use crate::telemetry::Metrics;
use std::collections::HashSet;
use std::future::Future;
use std::pin::Pin;
#[derive(Debug, Clone)]
pub struct OptimizationConfig {
pub preserve_recent: usize,
pub pinned_ids: HashSet<String>,
pub token_budget: usize,
pub frozen_prefix_len: usize,
}
impl Default for OptimizationConfig {
fn default() -> Self {
Self {
preserve_recent: 10,
pinned_ids: HashSet::new(),
token_budget: 4096,
frozen_prefix_len: 0,
}
}
}
impl OptimizationConfig {
pub fn with_budget(token_budget: usize) -> Self {
Self {
token_budget,
..Default::default()
}
}
pub fn preserve_recent(mut self, n: usize) -> Self {
self.preserve_recent = n;
self
}
pub fn pin(mut self, id: impl Into<String>) -> Self {
self.pinned_ids.insert(id.into());
self
}
pub fn pin_all(mut self, ids: impl IntoIterator<Item = impl Into<String>>) -> Self {
for id in ids {
self.pinned_ids.insert(id.into());
}
self
}
pub fn frozen_prefix(mut self, n: usize) -> Self {
self.frozen_prefix_len = n;
self
}
}
pub trait ContextOptimizer: Send + Sync {
fn optimize<'a>(
&'a self,
messages: Vec<Message>,
config: &'a OptimizationConfig,
) -> Pin<Box<dyn Future<Output = Vec<Message>> + Send + 'a>>;
fn estimate_tokens(&self, message: &Message) -> usize {
message.content_length() / 4 + 1
}
fn estimate_total_tokens(&self, messages: &[Message]) -> usize {
messages.iter().map(|m| self.estimate_tokens(m)).sum()
}
}
#[derive(Debug, Clone, Default)]
pub struct RecencyOptimizer;
impl RecencyOptimizer {
pub fn new() -> Self {
Self
}
}
impl ContextOptimizer for RecencyOptimizer {
fn optimize<'a>(
&'a self,
messages: Vec<Message>,
config: &'a OptimizationConfig,
) -> Pin<Box<dyn Future<Output = Vec<Message>> + Send + 'a>> {
Box::pin(async move {
let total = messages.len();
if total <= config.preserve_recent {
return messages;
}
let preserve_from = total.saturating_sub(config.preserve_recent);
let result: Vec<Message> = messages
.into_iter()
.enumerate()
.filter(|(i, m)| {
if let Some(id) = m.id() {
if config.pinned_ids.contains(id) {
return true;
}
}
*i >= preserve_from
})
.map(|(_, m)| m)
.collect();
let metrics = Metrics::global();
metrics.record_optimization("recency", total, result.len());
result
})
}
}
#[derive(Debug, Clone, Default)]
pub struct PriorityOptimizer {
priority_key: String,
default_priority: i32,
}
impl PriorityOptimizer {
pub fn new() -> Self {
Self {
priority_key: "priority".to_string(),
default_priority: 0,
}
}
pub fn with_priority_key(mut self, key: impl Into<String>) -> Self {
self.priority_key = key.into();
self
}
pub fn with_default_priority(mut self, priority: i32) -> Self {
self.default_priority = priority;
self
}
}
impl ContextOptimizer for PriorityOptimizer {
fn optimize<'a>(
&'a self,
messages: Vec<Message>,
config: &'a OptimizationConfig,
) -> Pin<Box<dyn Future<Output = Vec<Message>> + Send + 'a>> {
Box::pin(async move {
let total = messages.len();
if total <= config.preserve_recent {
return messages;
}
let preserve_from = total.saturating_sub(config.preserve_recent);
let mut indexed: Vec<(usize, Message, i32)> = messages
.into_iter()
.enumerate()
.map(|(i, m)| {
let priority = if m
.id()
.map(|id| config.pinned_ids.contains(id))
.unwrap_or(false)
{
i32::MAX
} else if i >= preserve_from {
i32::MAX - 1
} else {
self.default_priority
};
(i, m, priority)
})
.collect();
indexed.sort_by(|a, b| b.2.cmp(&a.2).then(a.0.cmp(&b.0)));
let current_tokens: usize = indexed
.iter()
.map(|(_, m, _)| self.estimate_tokens(m))
.sum();
let mut tokens_to_remove = current_tokens.saturating_sub(config.token_budget);
let mut keep = vec![true; indexed.len()];
for i in (0..indexed.len()).rev() {
if tokens_to_remove == 0 {
break;
}
if indexed[i].2 >= i32::MAX - 1 {
continue;
}
let msg_tokens = self.estimate_tokens(&indexed[i].1);
keep[i] = false;
tokens_to_remove = tokens_to_remove.saturating_sub(msg_tokens);
}
let result: Vec<_> = indexed
.into_iter()
.enumerate()
.filter(|(i, _)| keep[*i])
.map(|(_, (_, m, _))| m)
.collect();
let metrics = Metrics::global();
metrics.record_optimization("priority", total, result.len());
result
})
}
}
#[derive(Debug, Clone)]
pub struct TruncationOptimizer {
max_message_length: usize,
suffix: String,
}
impl Default for TruncationOptimizer {
fn default() -> Self {
Self {
max_message_length: 1000,
suffix: "... [truncated]".to_string(),
}
}
}
impl TruncationOptimizer {
pub fn new() -> Self {
Self::default()
}
pub fn with_max_length(mut self, max: usize) -> Self {
self.max_message_length = max;
self
}
pub fn with_suffix(mut self, suffix: impl Into<String>) -> Self {
self.suffix = suffix.into();
self
}
fn truncate_message(&self, mut message: Message) -> Message {
let text = message.text();
if text.len() > self.max_message_length {
let truncated = format!(
"{}{}",
&text[..self.max_message_length.saturating_sub(self.suffix.len())],
self.suffix
);
match &message {
Message::User { id, .. } => {
message = if let Some(id) = id {
Message::user_with_id(truncated, id)
} else {
Message::user(truncated)
};
}
Message::Assistant { id, .. } => {
message = if let Some(id) = id {
Message::assistant_with_id(truncated, id)
} else {
Message::assistant(truncated)
};
}
}
}
message
}
}
impl ContextOptimizer for TruncationOptimizer {
fn optimize<'a>(
&'a self,
messages: Vec<Message>,
config: &'a OptimizationConfig,
) -> Pin<Box<dyn Future<Output = Vec<Message>> + Send + 'a>> {
Box::pin(async move {
let total = messages.len();
let preserve_from = total.saturating_sub(config.preserve_recent);
let result: Vec<Message> = messages
.into_iter()
.enumerate()
.map(|(i, m)| {
let is_pinned = m
.id()
.map(|id| config.pinned_ids.contains(id))
.unwrap_or(false);
let is_recent = i >= preserve_from;
if is_pinned || is_recent {
m
} else {
self.truncate_message(m)
}
})
.collect();
let metrics = Metrics::global();
metrics.record_optimization("truncation", total, result.len());
result
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_recency_optimizer() {
let messages = vec![
Message::user("Message 1"),
Message::user("Message 2"),
Message::user("Message 3"),
Message::user("Message 4"),
Message::user("Message 5"),
];
let config = OptimizationConfig::default().preserve_recent(2);
let optimizer = RecencyOptimizer::new();
let result = optimizer.optimize(messages, &config).await;
assert_eq!(result.len(), 2);
assert_eq!(result[0].text(), "Message 4");
assert_eq!(result[1].text(), "Message 5");
}
#[tokio::test]
async fn test_recency_optimizer_with_pinned() {
let messages = vec![
Message::user_with_id("Important message", "pinned-1"),
Message::user("Message 2"),
Message::user("Message 3"),
Message::user("Message 4"),
Message::user("Message 5"),
];
let config = OptimizationConfig::default()
.preserve_recent(2)
.pin("pinned-1");
let optimizer = RecencyOptimizer::new();
let result = optimizer.optimize(messages, &config).await;
assert_eq!(result.len(), 3);
assert_eq!(result[0].text(), "Important message");
assert_eq!(result[1].text(), "Message 4");
assert_eq!(result[2].text(), "Message 5");
}
#[tokio::test]
async fn test_truncation_optimizer() {
let long_text = "a".repeat(2000);
let messages = vec![Message::user(&long_text), Message::user("Short message")];
let config = OptimizationConfig::default().preserve_recent(1);
let optimizer = TruncationOptimizer::new().with_max_length(100);
let result = optimizer.optimize(messages, &config).await;
assert_eq!(result.len(), 2);
assert!(result[0].text().len() <= 100);
assert!(result[0].text().ends_with("[truncated]"));
assert_eq!(result[1].text(), "Short message"); }
#[test]
fn test_optimization_config_builder() {
let config = OptimizationConfig::with_budget(8192)
.preserve_recent(20)
.pin("msg-1")
.pin("msg-2");
assert_eq!(config.token_budget, 8192);
assert_eq!(config.preserve_recent, 20);
assert!(config.pinned_ids.contains("msg-1"));
assert!(config.pinned_ids.contains("msg-2"));
}
}