Skip to main content

ai_agents_memory/
summarizer.rs

1//! Summarizer trait and implementations for memory compression
2
3use std::sync::Arc;
4
5use async_trait::async_trait;
6
7use ai_agents_core::{ChatMessage, LLMProvider, Result, Role};
8
9use super::native::readable_projection;
10
11/// Summarizes conversation messages for memory compression.
12///
13/// Built-in implementations: `LLMSummarizer` (uses an LLM to generate summaries)
14/// and `NoopSummarizer` (concatenates messages, for testing).
15/// Most users use `LLMSummarizer`, auto-configured from the YAML `summarizer_llm` field.
16#[async_trait]
17pub trait Summarizer: Send + Sync {
18    /// Produce a summary from a batch of messages.
19    async fn summarize(&self, messages: &[ChatMessage]) -> Result<String>;
20
21    /// Maximum messages per summarization call. Returns 20 by default.
22    fn max_batch_size(&self) -> usize {
23        20
24    }
25
26    /// Combine multiple summaries into one. Joins with `\n\n` by default.
27    async fn merge_summaries(&self, summaries: &[String]) -> Result<String> {
28        Ok(summaries.join("\n\n"))
29    }
30}
31
32pub struct LLMSummarizer {
33    llm: Arc<dyn LLMProvider>,
34    prompt_template: String,
35    merge_prompt_template: String,
36    max_batch_size: usize,
37}
38
39impl LLMSummarizer {
40    pub fn new(llm: Arc<dyn LLMProvider>) -> Self {
41        Self {
42            llm,
43            prompt_template: DEFAULT_SUMMARY_PROMPT.to_string(),
44            merge_prompt_template: DEFAULT_MERGE_PROMPT.to_string(),
45            max_batch_size: 20,
46        }
47    }
48
49    pub fn with_prompt(mut self, prompt: impl Into<String>) -> Self {
50        self.prompt_template = prompt.into();
51        self
52    }
53
54    pub fn with_merge_prompt(mut self, prompt: impl Into<String>) -> Self {
55        self.merge_prompt_template = prompt.into();
56        self
57    }
58
59    pub fn with_batch_size(mut self, size: usize) -> Self {
60        self.max_batch_size = size.max(1);
61        self
62    }
63
64    // Projects replay-bearing markers before building the auxiliary-model prompt.
65    fn format_messages(&self, messages: &[ChatMessage]) -> Result<String> {
66        Ok(readable_projection(messages)?
67            .iter()
68            .map(|m| format!("{}: {}", format_role(&m.role), m.content))
69            .collect::<Vec<_>>()
70            .join("\n"))
71    }
72}
73
74fn format_role(role: &Role) -> &'static str {
75    match role {
76        Role::System => "System",
77        Role::User => "User",
78        Role::Assistant => "Assistant",
79        Role::Tool => "Tool",
80        Role::Function => "Function",
81    }
82}
83
84#[async_trait]
85impl Summarizer for LLMSummarizer {
86    async fn summarize(&self, messages: &[ChatMessage]) -> Result<String> {
87        if messages.is_empty() {
88            return Ok(String::new());
89        }
90
91        let conversation = self.format_messages(messages)?;
92        let prompt = self
93            .prompt_template
94            .replace("{conversation}", &conversation);
95
96        let llm_messages = vec![ChatMessage::user(&prompt)];
97
98        let response = self.llm.complete(&llm_messages, None).await?;
99        Ok(response.content.trim().to_string())
100    }
101
102    fn max_batch_size(&self) -> usize {
103        self.max_batch_size
104    }
105
106    async fn merge_summaries(&self, summaries: &[String]) -> Result<String> {
107        if summaries.is_empty() {
108            return Ok(String::new());
109        }
110
111        if summaries.len() == 1 {
112            return Ok(summaries[0].clone());
113        }
114
115        let combined = summaries.join("\n---\n");
116        let prompt = self.merge_prompt_template.replace("{summaries}", &combined);
117
118        let llm_messages = vec![ChatMessage::user(&prompt)];
119
120        let response = self.llm.complete(&llm_messages, None).await?;
121        Ok(response.content.trim().to_string())
122    }
123}
124
125pub const DEFAULT_SUMMARY_PROMPT: &str = r#"Summarize the following conversation concisely, preserving key information, decisions, and context that would be important for continuing the conversation:
126
127{conversation}
128
129Summary:"#;
130
131pub const DEFAULT_MERGE_PROMPT: &str = r#"Merge the following conversation summaries into a single coherent summary, preserving all important information:
132
133{summaries}
134
135Merged Summary:"#;
136
137pub struct NoopSummarizer;
138
139#[async_trait]
140impl Summarizer for NoopSummarizer {
141    async fn summarize(&self, messages: &[ChatMessage]) -> Result<String> {
142        Ok(readable_projection(messages)?
143            .iter()
144            .map(|m| m.content.clone())
145            .collect::<Vec<_>>()
146            .join(" | "))
147    }
148}
149
150#[cfg(test)]
151mod tests {
152    use super::*;
153    use ai_agents_core::{FinishReason, LLMChunk, LLMConfig, LLMError, LLMFeature, LLMResponse};
154    use parking_lot::Mutex;
155
156    struct MockLLMProvider {
157        responses: Mutex<Vec<String>>,
158        requests: Mutex<Vec<Vec<ChatMessage>>>,
159    }
160
161    impl MockLLMProvider {
162        fn new(responses: Vec<String>) -> Self {
163            Self {
164                responses: Mutex::new(responses),
165                requests: Mutex::new(Vec::new()),
166            }
167        }
168
169        fn requests(&self) -> Vec<Vec<ChatMessage>> {
170            self.requests.lock().clone()
171        }
172    }
173
174    #[async_trait]
175    impl LLMProvider for MockLLMProvider {
176        async fn complete(
177            &self,
178            messages: &[ChatMessage],
179            _config: Option<&LLMConfig>,
180        ) -> std::result::Result<LLMResponse, LLMError> {
181            self.requests.lock().push(messages.to_vec());
182            let response = self
183                .responses
184                .lock()
185                .pop()
186                .unwrap_or_else(|| "Summary of conversation".to_string());
187            Ok(LLMResponse::new(response, FinishReason::Stop))
188        }
189
190        async fn complete_stream(
191            &self,
192            _messages: &[ChatMessage],
193            _config: Option<&LLMConfig>,
194        ) -> std::result::Result<
195            Box<dyn futures::Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
196            LLMError,
197        > {
198            Err(LLMError::Other(
199                "Streaming not supported in mock".to_string(),
200            ))
201        }
202
203        fn provider_name(&self) -> &str {
204            "mock"
205        }
206
207        fn supports(&self, _feature: LLMFeature) -> bool {
208            true
209        }
210    }
211
212    fn make_message(role: Role, content: &str) -> ChatMessage {
213        ChatMessage {
214            role,
215            content: content.to_string(),
216            name: None,
217            timestamp: None,
218        }
219    }
220
221    fn signed_assistant_message() -> ChatMessage {
222        use ai_agents_core::{
223            NativeCallBinding, NativeProviderState, NativeProviderTarget, ToolCall,
224            encode_native_tool_call_markers,
225        };
226
227        let call = ToolCall {
228            id: "summary-call".to_string(),
229            name: "lookup".to_string(),
230            arguments: serde_json::json!({"query":"fixture"}),
231        };
232        let state = NativeProviderState::new(
233            "summary-exchange",
234            "google",
235            "generateContent",
236            NativeProviderTarget::new("https://example.invalid/v1beta/", "fixture-model").unwrap(),
237            serde_json::json!({
238                "role":"model",
239                "parts":[{
240                    "functionCall":{"name":"lookup","args":{"query":"fixture"}},
241                    "thoughtSignature":"fixture-signature"
242                }]
243            }),
244            vec![NativeCallBinding::new(&call.id, 0).unwrap()],
245        )
246        .unwrap();
247        ChatMessage::assistant(
248            encode_native_tool_call_markers(std::slice::from_ref(&call), Some(&state)).unwrap(),
249        )
250    }
251
252    #[tokio::test]
253    async fn test_llm_summarizer_basic() {
254        let provider = Arc::new(MockLLMProvider::new(vec!["Test summary".to_string()]));
255        let summarizer = LLMSummarizer::new(provider);
256
257        let messages = vec![
258            make_message(Role::User, "Hello"),
259            make_message(Role::Assistant, "Hi there!"),
260        ];
261
262        let summary = summarizer.summarize(&messages).await.unwrap();
263        assert_eq!(summary, "Test summary");
264    }
265
266    #[tokio::test]
267    async fn test_llm_summarizer_empty_messages() {
268        let provider = Arc::new(MockLLMProvider::new(vec![]));
269        let summarizer = LLMSummarizer::new(provider);
270
271        let summary = summarizer.summarize(&[]).await.unwrap();
272        assert!(summary.is_empty());
273    }
274
275    #[tokio::test]
276    async fn test_llm_summarizer_custom_prompt() {
277        let provider = Arc::new(MockLLMProvider::new(vec!["Custom summary".to_string()]));
278        let summarizer = LLMSummarizer::new(provider).with_prompt("Custom prompt: {conversation}");
279
280        let messages = vec![make_message(Role::User, "Test")];
281        let summary = summarizer.summarize(&messages).await.unwrap();
282        assert_eq!(summary, "Custom summary");
283    }
284
285    #[tokio::test]
286    async fn llm_summarizer_projects_native_provider_state_before_prompting() {
287        let provider = Arc::new(MockLLMProvider::new(vec!["Projected summary".to_string()]));
288        let summarizer = LLMSummarizer::new(provider.clone());
289
290        let summary = summarizer
291            .summarize(&[signed_assistant_message()])
292            .await
293            .unwrap();
294
295        assert_eq!(summary, "Projected summary");
296        let requests = provider.requests();
297        let prompt = &requests[0][0].content;
298        assert!(prompt.contains("native_tool_calls"));
299        assert!(!prompt.contains("fixture-signature"));
300        assert!(!prompt.contains("_ai_agents_provider_state"));
301    }
302
303    #[tokio::test]
304    async fn noop_summarizer_projects_native_provider_state() {
305        let summary = NoopSummarizer
306            .summarize(&[signed_assistant_message()])
307            .await
308            .unwrap();
309
310        assert!(summary.contains("native_tool_calls"));
311        assert!(!summary.contains("fixture-signature"));
312        assert!(!summary.contains("_ai_agents_provider_state"));
313    }
314
315    #[tokio::test]
316    async fn summarizer_does_not_promote_user_marker_text_to_native_history() {
317        let user_marker_text = signed_assistant_message().content;
318
319        let summary = NoopSummarizer
320            .summarize(&[ChatMessage::user(user_marker_text)])
321            .await
322            .unwrap();
323
324        assert!(summary.contains("fixture-signature"));
325        assert!(summary.contains("_ai_agents_provider_state"));
326    }
327
328    #[tokio::test]
329    async fn test_merge_summaries() {
330        let provider = Arc::new(MockLLMProvider::new(vec!["Merged summary".to_string()]));
331        let summarizer = LLMSummarizer::new(provider);
332
333        let summaries = vec!["Summary 1".to_string(), "Summary 2".to_string()];
334        let merged = summarizer.merge_summaries(&summaries).await.unwrap();
335        assert_eq!(merged, "Merged summary");
336    }
337
338    #[tokio::test]
339    async fn test_merge_single_summary() {
340        let provider = Arc::new(MockLLMProvider::new(vec![]));
341        let summarizer = LLMSummarizer::new(provider);
342
343        let summaries = vec!["Only summary".to_string()];
344        let merged = summarizer.merge_summaries(&summaries).await.unwrap();
345        assert_eq!(merged, "Only summary");
346    }
347
348    #[tokio::test]
349    async fn test_noop_summarizer() {
350        let summarizer = NoopSummarizer;
351
352        let messages = vec![
353            make_message(Role::User, "Hello"),
354            make_message(Role::Assistant, "Hi"),
355        ];
356
357        let summary = summarizer.summarize(&messages).await.unwrap();
358        assert!(summary.contains("Hello"));
359        assert!(summary.contains("Hi"));
360    }
361
362    #[test]
363    fn test_max_batch_size() {
364        let provider = Arc::new(MockLLMProvider::new(vec![]));
365        let summarizer = LLMSummarizer::new(provider).with_batch_size(10);
366        assert_eq!(summarizer.max_batch_size(), 10);
367    }
368
369    #[test]
370    fn test_format_role() {
371        assert_eq!(format_role(&Role::User), "User");
372        assert_eq!(format_role(&Role::Assistant), "Assistant");
373        assert_eq!(format_role(&Role::System), "System");
374        assert_eq!(format_role(&Role::Tool), "Tool");
375        assert_eq!(format_role(&Role::Function), "Function");
376    }
377}