Skip to main content

aether_core/context/
compaction.rs

1use std::sync::Arc;
2
3use tokio_stream::StreamExt;
4
5use llm::types::IsoString;
6use llm::{ChatMessage, Context, LlmResponse, MessageId, StreamingModelProvider, TokenUsage};
7
8const SUMMARIZATION_PROMPT: &str = include_str!("prompts/summarization.md");
9
10/// Result of a compaction operation
11#[derive(Debug, Clone)]
12pub struct CompactionResult {
13    /// The summary text that replaces the compacted messages
14    pub summary: String,
15    /// Number of messages that were removed/compacted
16    pub messages_removed: usize,
17    /// Token usage reported by the summarization LLM call, if any
18    pub usage: Option<TokenUsage>,
19}
20
21/// Errors that can occur during compaction
22#[derive(Debug, Clone, thiserror::Error)]
23pub enum CompactionError {
24    /// The LLM failed to generate a summary
25    #[error("summarization failed: {0}")]
26    SummarizationFailed(String),
27    /// No messages to compact
28    #[error("nothing to compact")]
29    NothingToCompact,
30}
31
32/// Configuration for context compaction
33#[derive(Debug, Clone)]
34pub struct CompactionConfig {
35    /// Threshold (0.0-1.0) at which to trigger compaction
36    pub threshold: f64,
37}
38
39impl Default for CompactionConfig {
40    fn default() -> Self {
41        Self { threshold: super::DEFAULT_COMPACTION_THRESHOLD }
42    }
43}
44
45impl CompactionConfig {
46    /// Create a new compaction config with the given threshold
47    pub fn with_threshold(threshold: f64) -> Self {
48        Self { threshold }
49    }
50}
51
52/// Compacts context by generating an LLM summary
53pub struct Compactor {
54    llm: Arc<dyn StreamingModelProvider>,
55}
56
57impl Compactor {
58    pub fn new(llm: Arc<dyn StreamingModelProvider>) -> Self {
59        Self { llm }
60    }
61
62    /// Generate a structured summary of the conversation.
63    ///
64    /// Takes the context snapshot by value since the caller applies the summary
65    /// to its live context (which may have changed while the summarization
66    /// request was in flight) via [`Context::with_compacted_summary`].
67    pub async fn compact(&self, mut context: Context) -> Result<CompactionResult, CompactionError> {
68        let messages_to_summarize = context.messages_for_summary();
69        if messages_to_summarize.is_empty() {
70            return Err(CompactionError::NothingToCompact);
71        }
72
73        let messages_removed = messages_to_summarize.len();
74
75        context.add_message(ChatMessage::User {
76            message_id: MessageId::new(),
77            content: vec![llm::ContentBlock::text(format!(
78                "{SUMMARIZATION_PROMPT}\n\nPlease perform a structured handoff of the conversation above."
79            ))],
80            timestamp: IsoString::now(),
81        });
82
83        let mut stream = self.llm.stream_response(&context);
84        let mut summary = String::new();
85        let mut usage = None;
86
87        while let Some(result) = stream.next().await {
88            match result {
89                Ok(LlmResponse::Text { chunk }) => {
90                    summary.push_str(&chunk);
91                }
92                Ok(LlmResponse::Usage { tokens }) => usage = Some(tokens),
93                Ok(LlmResponse::Done { .. }) => break,
94                Ok(LlmResponse::Error { message }) => {
95                    return Err(CompactionError::SummarizationFailed(message));
96                }
97                Err(e) => {
98                    return Err(CompactionError::SummarizationFailed(e.to_string()));
99                }
100                _ => {}
101            }
102        }
103
104        if summary.is_empty() {
105            return Err(CompactionError::SummarizationFailed("LLM returned empty summary".to_string()));
106        }
107
108        Ok(CompactionResult { summary, messages_removed, usage })
109    }
110}
111
112#[cfg(test)]
113mod tests {
114    use super::*;
115    use llm::types::IsoString;
116    use llm::{ChatMessage, ContentBlock, MessageId};
117
118    #[test]
119    fn test_compaction_config_default() {
120        let config = CompactionConfig::default();
121        assert!((config.threshold - 0.85).abs() < 0.001);
122    }
123
124    #[test]
125    fn test_compaction_config_with_threshold() {
126        let config = CompactionConfig::with_threshold(0.9);
127        assert!((config.threshold - 0.9).abs() < 0.001);
128    }
129
130    #[tokio::test]
131    async fn test_compactor_generates_summary() {
132        use llm::testing::FakeLlmProvider;
133
134        let summary_response = vec![
135            LlmResponse::Start,
136            LlmResponse::text(
137                "## Primary Goal\nTest the compaction feature\n\n## Completed Work\n- Wrote initial tests\n\n## File Changes\n- `src/main.rs` — added entry point\n\n## Key Decisions\n- Use structured handoff — preserves context better\n\n## Current State\nRunning compaction tests\n\n## Next Steps\n1. Verify all tests pass\n\n## Open Questions\n(none)\n\n## Constraints\n(none)",
138            ),
139            LlmResponse::done(),
140        ];
141
142        let fake_llm = Arc::new(FakeLlmProvider::with_single_response(summary_response));
143        let compactor = Compactor::new(fake_llm);
144
145        let context = Context::new(
146            vec![
147                ChatMessage::system("System"),
148                ChatMessage::User {
149                    message_id: MessageId::new(),
150                    content: vec![ContentBlock::text("Test message")],
151                    timestamp: IsoString::now(),
152                },
153            ],
154            vec![],
155        );
156
157        let result = compactor.compact(context).await;
158        assert!(result.is_ok());
159
160        let result = result.unwrap();
161        assert!(result.summary.contains("Primary Goal"));
162        assert!(result.summary.contains("File Changes"));
163        assert!(result.summary.contains("Next Steps"));
164        assert_eq!(result.messages_removed, 1);
165    }
166
167    #[tokio::test]
168    async fn test_compactor_handles_error() {
169        use llm::testing::FakeLlmProvider;
170
171        let error_response = vec![LlmResponse::Error { message: "API error".to_string() }];
172
173        let fake_llm = Arc::new(FakeLlmProvider::with_single_response(error_response));
174        let compactor = Compactor::new(fake_llm);
175
176        let context = Context::new(
177            vec![
178                ChatMessage::system("System"),
179                ChatMessage::User {
180                    message_id: MessageId::new(),
181                    content: vec![ContentBlock::text("Test")],
182                    timestamp: IsoString::now(),
183                },
184            ],
185            vec![],
186        );
187
188        let result = compactor.compact(context).await;
189        assert!(matches!(result, Err(CompactionError::SummarizationFailed(_))));
190    }
191
192    #[tokio::test]
193    async fn test_compactor_empty_context() {
194        use llm::testing::FakeLlmProvider;
195
196        let fake_llm = Arc::new(FakeLlmProvider::with_single_response(vec![]));
197        let compactor = Compactor::new(fake_llm);
198
199        let context = Context::new(vec![ChatMessage::system("System")], vec![]);
200
201        let result = compactor.compact(context).await;
202        assert!(matches!(result, Err(CompactionError::NothingToCompact)));
203    }
204}