behest_runtime/compaction/
select.rs1use crate::token::estimate_record_tokens;
12use behest_store::MessageRecord;
13
14#[derive(Debug, Clone)]
16pub struct SelectionResult {
17 pub head: Vec<MessageRecord>,
19 pub tail: Vec<MessageRecord>,
21 pub tail_start_id: Option<uuid::Uuid>,
23}
24
25#[derive(Debug)]
27struct Turn {
28 indices: Vec<usize>,
30}
31
32fn turns(messages: &[MessageRecord]) -> Vec<Turn> {
36 let mut result = Vec::new();
37 let mut current = Vec::new();
38
39 for (i, msg) in messages.iter().enumerate() {
40 if !msg.is_compaction && msg.role == behest_store::MessageRole::User && !current.is_empty()
41 {
42 result.push(Turn {
43 indices: std::mem::take(&mut current),
44 });
45 }
46 current.push(i);
47 }
48
49 if !current.is_empty() {
50 result.push(Turn { indices: current });
51 }
52
53 result
54}
55
56#[must_use]
67pub fn select(
68 messages: &[MessageRecord],
69 tail_turns: usize,
70 keep_tokens: usize,
71) -> SelectionResult {
72 if messages.is_empty() {
73 return SelectionResult {
74 head: Vec::new(),
75 tail: Vec::new(),
76 tail_start_id: None,
77 };
78 }
79
80 let all_turns = turns(messages);
81
82 if all_turns.is_empty() {
83 return SelectionResult {
84 head: messages.to_vec(),
85 tail: Vec::new(),
86 tail_start_id: None,
87 };
88 }
89
90 let preserved_turns = if all_turns.len() <= tail_turns {
92 all_turns.len()
93 } else {
94 tail_turns
95 };
96
97 let head_turns = &all_turns[..all_turns.len() - preserved_turns];
98 let tail_turns_slice = &all_turns[all_turns.len() - preserved_turns..];
99
100 let head_indices: Vec<usize> = head_turns
102 .iter()
103 .flat_map(|t| t.indices.iter().copied())
104 .collect();
105
106 let mut tail_indices: Vec<usize> = Vec::new();
108 let mut accumulated = 0usize;
109
110 for turn in tail_turns_slice.iter().rev() {
111 let mut turn_accumulated = 0usize;
112 let mut turn_indices: Vec<usize> = Vec::new();
113
114 for &idx in turn.indices.iter().rev() {
115 let tokens = estimate_record_tokens(&messages[idx]);
116
117 if accumulated + tokens > keep_tokens && !tail_indices.is_empty() {
118 break;
120 }
121
122 turn_accumulated += tokens;
123 turn_indices.push(idx);
124 }
125
126 if accumulated + turn_accumulated > keep_tokens && !tail_indices.is_empty() {
127 let overflow: Vec<usize> = turn
130 .indices
131 .iter()
132 .copied()
133 .filter(|i| !turn_indices.contains(i))
134 .collect();
135 let mut full_head: Vec<usize> = head_indices.into_iter().chain(overflow).collect();
136 full_head.sort_unstable();
137
138 tail_indices.reverse();
139
140 let head: Vec<MessageRecord> = full_head.iter().map(|&i| messages[i].clone()).collect();
141 let tail: Vec<MessageRecord> =
142 tail_indices.iter().map(|&i| messages[i].clone()).collect();
143 let tail_start_id = tail.first().map(|m| m.id);
144
145 return SelectionResult {
146 head,
147 tail,
148 tail_start_id,
149 };
150 }
151
152 accumulated += turn_accumulated;
153
154 turn_indices.reverse();
156 let mut combined = turn_indices;
157 combined.append(&mut tail_indices);
158 tail_indices = combined;
159 }
160
161 tail_indices.sort_unstable();
164
165 let head: Vec<MessageRecord> = head_indices.iter().map(|&i| messages[i].clone()).collect();
166 let tail: Vec<MessageRecord> = tail_indices.iter().map(|&i| messages[i].clone()).collect();
167 let tail_start_id = tail.first().map(|m| m.id);
168
169 SelectionResult {
170 head,
171 tail,
172 tail_start_id,
173 }
174}
175
176#[cfg(test)]
177#[allow(clippy::unwrap_used)]
178mod tests {
179 use super::*;
180 use behest_provider::ContentPart;
181 use uuid::Uuid;
182
183 fn make_msg(role: behest_store::MessageRole, text: &str) -> MessageRecord {
184 MessageRecord::new(Uuid::now_v7(), role, vec![ContentPart::text(text)])
185 }
186
187 fn make_user(text: &str) -> MessageRecord {
188 make_msg(behest_store::MessageRole::User, text)
189 }
190
191 fn make_assistant(text: &str) -> MessageRecord {
192 make_msg(behest_store::MessageRole::Assistant, text)
193 }
194
195 fn make_tool(text: &str) -> MessageRecord {
196 make_msg(behest_store::MessageRole::Tool, text)
197 }
198
199 #[test]
200 fn empty_messages_returns_empty() {
201 let result = select(&[], 2, 8_000);
202 assert!(result.head.is_empty());
203 assert!(result.tail.is_empty());
204 assert!(result.tail_start_id.is_none());
205 }
206
207 #[test]
208 fn single_turn_within_budget() {
209 let messages = vec![make_user("Hello"), make_assistant("Hi there!")];
210
211 let total = estimate_record_tokens(&messages[0]) + estimate_record_tokens(&messages[1]);
212
213 let result = select(&messages, 2, total * 2);
215 assert!(result.head.is_empty());
216 assert_eq!(result.tail.len(), 2);
217 }
218
219 #[test]
220 fn multiple_turns_preserves_recent() {
221 let messages = vec![
222 make_user("Turn 1"),
223 make_assistant("Response 1"),
224 make_user("Turn 2"),
225 make_assistant("Response 2"),
226 make_user("Turn 3"),
227 make_assistant("Response 3"),
228 ];
229
230 let result = select(&messages, 1, 8_000);
232 assert!(!result.head.is_empty(), "head should have older turns");
233 assert!(!result.tail.is_empty(), "tail should have recent turn");
234 assert_eq!(result.tail[0].role, behest_store::MessageRole::User);
236 }
237
238 #[test]
239 fn compact_tool_messages_are_in_turn() {
240 let messages = vec![
241 make_user("Use tool"),
242 make_assistant("Calling tool..."),
243 make_tool("Tool result"),
244 make_assistant("Done with tool"),
245 ];
246
247 let result = select(&messages, 2, 8_000);
249 assert!(result.head.is_empty());
251 assert_eq!(result.tail.len(), 4);
252 }
253
254 #[test]
255 fn tiny_budget_forces_compaction() {
256 let messages: Vec<MessageRecord> = (0..5)
257 .flat_map(|i| {
258 vec![
259 make_user(&format!("Question {i}")),
260 make_assistant(&format!("Answer {i}")),
261 ]
262 })
263 .collect();
264
265 let result = select(&messages, 2, 10);
267 assert!(!result.head.is_empty());
269 assert!(!result.tail.is_empty());
271 }
272}