Skip to main content

codei_session/
compact.rs

1use codei_config::AgentConfig;
2use codei_llm::Message;
3
4use crate::model::{to_llm_messages, Role, Session, StoredMessage};
5
6/// Token budget for context window management.
7#[derive(Debug, Clone, Copy)]
8pub struct TokenBudget {
9    pub max_tokens: u32,
10    pub compaction_threshold: f32,
11    pub keep_messages: usize,
12}
13
14impl TokenBudget {
15    pub fn limit(&self) -> u32 {
16        ((self.max_tokens as f32) * self.compaction_threshold) as u32
17    }
18
19    pub fn from_agent(agent: &AgentConfig) -> Self {
20        Self {
21            max_tokens: agent.context_window_tokens,
22            compaction_threshold: agent.compaction_threshold,
23            keep_messages: agent.compaction_keep_messages.max(2) as usize,
24        }
25    }
26}
27
28/// Rough token estimate (chars / 4).
29pub fn estimate_tokens(messages: &[Message]) -> u32 {
30    messages
31        .iter()
32        .map(|m| {
33            let content_len = m.content.as_ref().map(|c| c.len()).unwrap_or(0);
34            let tools_len = m
35                .tool_calls
36                .as_ref()
37                .map(|calls| calls.iter().map(|c| c.arguments.len()).sum::<usize>())
38                .unwrap_or(0);
39            (content_len + tools_len) as u32 / 4 + 8
40        })
41        .sum()
42}
43
44/// Rough token estimate for tool definitions attached to a chat request.
45pub fn estimate_tool_defs_tokens(tools: &[codei_llm::ToolDefinition]) -> u32 {
46    tools
47        .iter()
48        .map(|d| {
49            let len = d.name.len() + d.description.len() + d.parameters.to_string().len();
50            len as u32 / 4 + 16
51        })
52        .sum()
53}
54
55/// Cap configured max output tokens so estimated input + output fits in the context window.
56pub fn cap_output_tokens(
57    messages: &[Message],
58    tools: Option<&[codei_llm::ToolDefinition]>,
59    configured_max: u32,
60    context_window: u32,
61) -> u32 {
62    const MIN_SLACK: u32 = 1024;
63
64    let mut input_est = estimate_tokens(messages);
65    if let Some(defs) = tools {
66        input_est += estimate_tool_defs_tokens(defs);
67    }
68
69    // Slack covers tokenizer variance and underestimate from the chars/4 heuristic.
70    let slack = (input_est / 6).max(MIN_SLACK);
71    let available = context_window
72        .saturating_sub(input_est)
73        .saturating_sub(slack);
74
75    if available == 0 {
76        return 1;
77    }
78    configured_max.min(available)
79}
80
81/// Whether the session should be compacted based on estimated token usage.
82pub fn should_compact_session(session: &Session, system_prompt: &str, agent: &AgentConfig) -> bool {
83    let budget = TokenBudget::from_agent(agent);
84    let tokens = estimate_tokens(&to_llm_messages(session, system_prompt));
85    tokens > budget.limit() && session.messages.len() > budget.keep_messages
86}
87
88/// Format stored messages into a transcript for LLM summarization.
89pub fn format_transcript(messages: &[StoredMessage]) -> String {
90    let mut out = String::new();
91    for msg in messages {
92        let role = match msg.role {
93            Role::User => "User",
94            Role::Assistant => "Assistant",
95            Role::Tool => "Tool",
96            Role::System => continue,
97        };
98        let Some(text) = msg.text() else { continue };
99        let trimmed = text.trim();
100        if trimmed.is_empty() {
101            continue;
102        }
103        out.push_str(role);
104        out.push_str(": ");
105        out.push_str(&truncate_chars(trimmed, 6_000));
106        out.push_str("\n\n");
107    }
108    out
109}
110
111fn truncate_chars(text: &str, max_chars: usize) -> String {
112    if text.chars().count() <= max_chars {
113        return text.to_string();
114    }
115    let truncated: String = text.chars().take(max_chars).collect();
116    format!("{truncated}…")
117}
118
119/// Truncate older messages when over budget. Keeps system prompt and recent turns.
120/// Used as a last-resort guard when building the LLM request.
121pub fn compact_messages(messages: Vec<Message>, budget: TokenBudget) -> Vec<Message> {
122    if messages.is_empty() || estimate_tokens(&messages) <= budget.limit() {
123        return messages;
124    }
125
126    const MIN_KEEP: usize = 2;
127    let keep_recent = budget.keep_messages.max(MIN_KEEP);
128    if messages.len() <= keep_recent + 1 {
129        return messages;
130    }
131
132    let system = messages.first().cloned();
133    let mut rest: Vec<Message> = messages.into_iter().skip(1).collect();
134    if rest.len() <= keep_recent {
135        return reconstruct(system, rest, 0);
136    }
137
138    let omitted = rest.len() - keep_recent;
139    rest.drain(0..omitted);
140    reconstruct(system, rest, omitted)
141}
142
143fn reconstruct(system: Option<Message>, kept: Vec<Message>, omitted: usize) -> Vec<Message> {
144    let mut out = Vec::new();
145    if let Some(sys) = system {
146        out.push(sys);
147    }
148    if omitted > 0 {
149        out.push(Message::user(format!(
150            "[Context compacted: {omitted} earlier messages omitted to fit the context window]"
151        )));
152    }
153    out.extend(kept);
154    out
155}
156
157#[cfg(test)]
158mod tests {
159    use super::*;
160    use codei_llm::Message;
161
162    #[test]
163    fn compacts_when_over_budget() {
164        let budget = TokenBudget {
165            max_tokens: 100,
166            compaction_threshold: 0.5,
167            keep_messages: 4,
168        };
169        let messages: Vec<Message> = std::iter::once(Message::system("sys"))
170            .chain((0..20).map(|i| Message::user(format!("message {i} {}", "x".repeat(50)))))
171            .collect();
172        let compacted = compact_messages(messages, budget);
173        assert!(compacted.len() < 21);
174        assert!(compacted[0].content.as_deref() == Some("sys"));
175    }
176
177    #[test]
178    fn should_compact_when_over_threshold() {
179        use codei_config::AgentConfig;
180
181        let mut session = Session::new(std::path::PathBuf::from("/tmp"));
182        for i in 0..30 {
183            session.push_user(format!("message {i} {}", "x".repeat(80)));
184        }
185        let agent = AgentConfig {
186            context_window_tokens: 200,
187            compaction_threshold: 0.5,
188            compaction_keep_messages: 6,
189            ..Default::default()
190        };
191        assert!(should_compact_session(&session, "system prompt", &agent));
192    }
193
194    #[test]
195    fn compact_with_summary_inserts_summary_message() {
196        let mut session = Session::new(std::path::PathBuf::from("/tmp"));
197        for i in 0..5 {
198            session.push_user(format!("msg {i}"));
199        }
200        session.compact_with_summary(2, "summary text".into());
201        assert_eq!(session.messages.len(), 3);
202        assert!(session.messages[0].text().unwrap().contains("summary text"));
203        assert!(session.messages[1].text().unwrap().contains("msg 3"));
204    }
205
206    #[test]
207    fn cap_output_tokens_fits_context_window() {
208        let content = "x".repeat(21505 * 4);
209        let messages = vec![Message::user(content)];
210        let context_window = 87040;
211        let configured_max = 65536;
212        let capped = cap_output_tokens(&messages, None, configured_max, context_window);
213        assert!(capped < configured_max);
214        let input_est = estimate_tokens(&messages);
215        let slack = (input_est / 6).max(1024);
216        assert!(input_est + slack + capped <= context_window);
217    }
218}