1use codei_config::AgentConfig;
2use codei_llm::Message;
3
4use crate::model::{to_llm_messages, Role, Session, StoredMessage};
5
6#[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
28pub 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
44pub 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
55pub 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 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
81pub 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
88pub 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
119pub 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}