Skip to main content

behest_runtime/compaction/
select.rs

1//! Message selection for compaction.
2//!
3//! Implements turn-based selection: conversation messages are grouped into
4//! "turns" (bounded by non-compaction user messages), and the most recent
5//! `tail_turns` turns are preserved. Within the preserved range, a token
6//! budget (`keep_tokens`) is enforced from-back, so the oldest retained
7//! messages are dropped if the budget is exceeded.
8//!
9//! Ported from OpenCode V1: `packages/opencode/src/session/compaction.ts` select().
10
11use crate::token::estimate_record_tokens;
12use behest_store::MessageRecord;
13
14/// Result of the selection algorithm.
15#[derive(Debug, Clone)]
16pub struct SelectionResult {
17    /// Messages to be compacted (summarised by the compaction LLM).
18    pub head: Vec<MessageRecord>,
19    /// Messages to retain as recent context.
20    pub tail: Vec<MessageRecord>,
21    /// The message ID of the first retained message (tail boundary).
22    pub tail_start_id: Option<uuid::Uuid>,
23}
24
25/// A conversation turn — bounded by non-compaction user messages.
26#[derive(Debug)]
27struct Turn {
28    /// Indices into the source message slice.
29    indices: Vec<usize>,
30}
31
32/// Groups messages into turns, where each turn starts at a non-compaction
33/// user message and extends until the next non-compaction user message
34/// (or end of slice).
35fn 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/// Selects which messages to compact and which to retain.
57///
58/// # Arguments
59/// * `messages` - The full session message history, in chronological order.
60/// * `tail_turns` - Number of recent turns to preserve intact (default 2).
61/// * `keep_tokens` - Token budget for the retained tail.
62///
63/// # Returns
64/// A [`SelectionResult`] with `head` (to compact), `tail` (to retain),
65/// and `tail_start_id` (first retained message ID).
66#[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    // Take the most recent N turns
91    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    // Collect head indices
101    let head_indices: Vec<usize> = head_turns
102        .iter()
103        .flat_map(|t| t.indices.iter().copied())
104        .collect();
105
106    // Walk backward through tail turns, accumulating tokens
107    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                // Budget exceeded — stop adding more from this turn
119                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            // This partial turn pushes us over budget; stop here.
128            // The remaining items in this turn and earlier turns become head.
129            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        // Reverse to get chronological order inside this turn
155        turn_indices.reverse();
156        let mut combined = turn_indices;
157        combined.append(&mut tail_indices);
158        tail_indices = combined;
159    }
160
161    // If we got here, all preserved turns fit within budget.
162    // Any overflow from an earlier incomplete accumulation goes to head.
163    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        // Budget is large enough
214        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        // Keep only 1 turn
231        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        // Tail should start with Turn 3
235        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        // All messages should be in one turn (one user message)
248        let result = select(&messages, 2, 8_000);
249        // head should be empty since it's just one turn
250        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        // Extremely tight budget
266        let result = select(&messages, 2, 10);
267        // Head should contain most messages
268        assert!(!result.head.is_empty());
269        // Tail should still have at least the last user message
270        assert!(!result.tail.is_empty());
271    }
272}