Skip to main content

lc_memory/
summary_buffer.rs

1// lc-memory/src/summary_buffer.rs
2//! Conversation Summary Buffer Memory
3//!
4//! Combines summary and full conversation, balancing token consumption and conversation quality.
5
6use async_trait::async_trait;
7use serde_json::Value;
8use std::collections::HashMap;
9
10use super::base::{BaseMemory, ChatMessageHistory, MemoryError};
11use lc_core::language_models::BaseChatModel;
12use lc_core::language_models::LLMResult;
13use lc_core::runnables::Runnable;
14use lc_prompts::PromptTemplate;
15use lc_schema::Message;
16
17const DEFAULT_SUMMARY_PROMPT: &str =
18    "Progressively summarize the conversation, adding new content to the previous summary.
19
20Current summary:
21{summary}
22
23New lines of conversation:
24{new_lines}
25
26New summary:";
27
28/// Conversation Summary Buffer Memory
29///
30/// Combines summary and full conversation:
31/// - Keeps the last k rounds of full conversation (ensuring fluency)
32/// - Summarizes older conversations (saving tokens)
33///
34/// # Example
35/// ```ignore
36/// use lc_memory::ConversationSummaryBufferMemory;
37/// use lc_providers::OpenAIChat;
38///
39/// let llm = OpenAIChat::new(config);
40/// let memory = ConversationSummaryBufferMemory::new(llm, 5); // Keep last 5 rounds
41///
42/// // After 20 rounds:
43/// // - First 15 rounds -> summary
44/// // - Last 5 rounds -> full conversation
45/// ```
46pub struct ConversationSummaryBufferMemory<M: BaseChatModel> {
47    llm: M,
48
49    /// M67: Removed `Mutex<String>` - &mut self already guarantees exclusive access
50    buffer: String,
51    chat_memory: ChatMessageHistory,
52
53    max_token_limit: usize,
54
55    input_key: String,
56    output_key: String,
57    memory_key: String,
58
59    summary_prompt: String,
60    return_messages: bool,
61}
62
63impl<M: BaseChatModel> ConversationSummaryBufferMemory<M> {
64    pub fn new(llm: M, max_token_limit: usize) -> Self {
65        Self {
66            llm,
67            buffer: String::new(),
68            chat_memory: ChatMessageHistory::new(),
69            max_token_limit,
70            input_key: "input".to_string(),
71            output_key: "output".to_string(),
72            memory_key: "history".to_string(),
73            summary_prompt: DEFAULT_SUMMARY_PROMPT.to_string(),
74            return_messages: false,
75        }
76    }
77
78    pub fn with_input_key(mut self, key: impl Into<String>) -> Self {
79        self.input_key = key.into();
80        self
81    }
82
83    pub fn with_output_key(mut self, key: impl Into<String>) -> Self {
84        self.output_key = key.into();
85        self
86    }
87
88    pub fn with_memory_key(mut self, key: impl Into<String>) -> Self {
89        self.memory_key = key.into();
90        self
91    }
92
93    pub fn with_summary_prompt(mut self, prompt: impl Into<String>) -> Self {
94        self.summary_prompt = prompt.into();
95        self
96    }
97
98    pub fn with_return_messages(mut self, return_messages: bool) -> Self {
99        self.return_messages = return_messages;
100        self
101    }
102
103    pub fn chat_memory(&self) -> &ChatMessageHistory {
104        &self.chat_memory
105    }
106
107    pub fn chat_memory_mut(&mut self) -> &mut ChatMessageHistory {
108        &mut self.chat_memory
109    }
110
111    pub fn max_token_limit(&self) -> usize {
112        self.max_token_limit
113    }
114
115    pub async fn buffer(&self) -> String {
116        self.buffer.clone()
117    }
118
119    fn estimate_tokens(text: &str) -> usize {
120        text.len() / 4
121    }
122
123    fn prune_messages(&self, messages: &[Message]) -> Vec<Message> {
124        let total_tokens = messages
125            .iter()
126            .map(|m| Self::estimate_tokens(&m.content))
127            .sum::<usize>();
128
129        if total_tokens <= self.max_token_limit {
130            return messages.to_vec();
131        }
132
133        let mut kept_messages = Vec::new();
134        let mut current_tokens = 0;
135
136        for msg in messages.iter().rev() {
137            let msg_tokens = Self::estimate_tokens(&msg.content);
138            if current_tokens + msg_tokens <= self.max_token_limit {
139                kept_messages.push(msg.clone());
140                current_tokens += msg_tokens;
141            } else {
142                break;
143            }
144        }
145
146        kept_messages.reverse();
147        kept_messages
148    }
149
150    async fn predict_new_summary(&self, new_lines: &str) -> Result<String, MemoryError> {
151        let buffer = self.buffer.clone();
152
153        let prompt = {
154            let template = PromptTemplate::new(&self.summary_prompt);
155            let mut vars: std::collections::HashMap<&str, &str> = std::collections::HashMap::new();
156            vars.insert("summary", buffer.as_str());
157            vars.insert("new_lines", new_lines);
158            template
159                .format(&vars)
160                .unwrap_or_else(|_| self.summary_prompt.clone())
161        };
162
163        let messages = vec![Message::human(&prompt)];
164
165        let result =
166            self.llm.invoke(messages, None).await.map_err(|e| {
167                MemoryError::SaveError(format!("LLM summary generation failed: {}", e))
168            })?;
169
170        Ok(result.content)
171    }
172}
173
174#[async_trait]
175impl<M: BaseChatModel + Send + Sync + 'static> BaseMemory for ConversationSummaryBufferMemory<M>
176where
177    <M as Runnable<Vec<Message>, LLMResult>>::Error: std::fmt::Display,
178{
179    fn memory_variables(&self) -> Vec<&str> {
180        vec![&self.memory_key]
181    }
182
183    async fn load_memory_variables(
184        &self,
185        _inputs: &HashMap<String, String>,
186    ) -> Result<HashMap<String, Value>, MemoryError> {
187        let mut result = HashMap::new();
188
189        let buffer = self.buffer.clone();
190        let messages = self.chat_memory.messages();
191        let pruned = self.prune_messages(messages);
192
193        if self.return_messages {
194            let mut all_messages = Vec::new();
195
196            if !buffer.is_empty() {
197                all_messages.push(Message::system(&buffer));
198            }
199
200            all_messages.extend(pruned);
201
202            let messages_value: Vec<Value> = all_messages
203                .iter()
204                .map(|m| serde_json::to_value(m).unwrap_or(Value::Null))
205                .collect();
206
207            result.insert(self.memory_key.clone(), Value::Array(messages_value));
208        } else {
209            let mut history = String::new();
210
211            if !buffer.is_empty() {
212                history.push_str(&format!("Summary: {}\n\n", buffer));
213            }
214
215            for msg in &pruned {
216                let role = match msg.message_type {
217                    lc_schema::MessageType::Human => "Human",
218                    lc_schema::MessageType::AI => "AI",
219                    lc_schema::MessageType::System => "System",
220                    lc_schema::MessageType::Tool { .. } => "Tool",
221                };
222                history.push_str(&format!("{}: {}\n", role, msg.content));
223            }
224
225            result.insert(self.memory_key.clone(), Value::String(history));
226        }
227
228        Ok(result)
229    }
230
231    async fn save_context(
232        &mut self,
233        inputs: &HashMap<String, String>,
234        outputs: &HashMap<String, String>,
235    ) -> Result<(), MemoryError> {
236        let empty = String::new();
237        let input = inputs.get(&self.input_key).unwrap_or(&empty);
238        let output = outputs.get(&self.output_key).unwrap_or(&empty);
239
240        self.chat_memory.add_user_message(input);
241        self.chat_memory.add_ai_message(output);
242
243        let messages = self.chat_memory.messages();
244        let total_tokens = messages
245            .iter()
246            .map(|m| Self::estimate_tokens(&m.content))
247            .sum::<usize>();
248
249        if total_tokens > self.max_token_limit {
250            let pruned = self.prune_messages(messages);
251
252            let pruned_count = pruned.len();
253
254            if messages.len() > pruned_count {
255                let messages_to_summarize: Vec<&Message> = messages
256                    .iter()
257                    .take(messages.len() - pruned_count)
258                    .collect();
259
260                if !messages_to_summarize.is_empty() {
261                    let new_lines: String = messages_to_summarize
262                        .iter()
263                        .map(|m| {
264                            let role = match m.message_type {
265                                lc_schema::MessageType::Human => "Human",
266                                lc_schema::MessageType::AI => "AI",
267                                lc_schema::MessageType::System => "System",
268                                lc_schema::MessageType::Tool { .. } => "Tool",
269                            };
270                            format!("{}: {}", role, m.content)
271                        })
272                        .collect::<Vec<_>>()
273                        .join("\n");
274
275                    let new_summary = self.predict_new_summary(&new_lines).await?;
276
277                    self.buffer = new_summary;
278                }
279
280                self.chat_memory.clear();
281                for msg in pruned {
282                    if matches!(msg.message_type, lc_schema::MessageType::Human) {
283                        self.chat_memory.add_user_message(&msg.content);
284                    } else if matches!(msg.message_type, lc_schema::MessageType::AI) {
285                        self.chat_memory.add_ai_message(&msg.content);
286                    } else if matches!(msg.message_type, lc_schema::MessageType::System) {
287                        // H28: Preserve System messages during pruning
288                        self.chat_memory.add_system_message(&msg.content);
289                    }
290                }
291            }
292        }
293
294        Ok(())
295    }
296
297    async fn clear(&mut self) -> Result<(), MemoryError> {
298        self.buffer = String::new();
299        self.chat_memory.clear();
300        Ok(())
301    }
302}
303
304#[cfg(test)]
305mod tests {
306    use super::*;
307    use lc_providers::{OpenAIChat, OpenAIConfig};
308
309    fn create_test_config() -> OpenAIConfig {
310        OpenAIConfig::default()
311    }
312
313    #[test]
314    fn test_new() {
315        let llm = OpenAIChat::new(create_test_config());
316        let memory: ConversationSummaryBufferMemory<OpenAIChat> =
317            ConversationSummaryBufferMemory::new(llm, 1000);
318
319        assert_eq!(memory.memory_variables(), vec!["history"]);
320        assert_eq!(memory.max_token_limit(), 1000);
321    }
322
323    #[test]
324    fn test_with_options() {
325        let llm = OpenAIChat::new(create_test_config());
326        let memory: ConversationSummaryBufferMemory<OpenAIChat> =
327            ConversationSummaryBufferMemory::new(llm, 500)
328                .with_input_key("question")
329                .with_output_key("answer")
330                .with_memory_key("context")
331                .with_return_messages(true);
332
333        assert_eq!(memory.input_key, "question");
334        assert_eq!(memory.output_key, "answer");
335        assert_eq!(memory.memory_key, "context");
336        assert!(memory.return_messages);
337    }
338
339    #[test]
340    fn test_estimate_tokens() {
341        let text1 = "Hello";
342        let text2 = "Hello World";
343        let text3 = "This is some Chinese text";
344
345        assert!(ConversationSummaryBufferMemory::<OpenAIChat>::estimate_tokens(text1) > 0);
346        assert!(
347            ConversationSummaryBufferMemory::<OpenAIChat>::estimate_tokens(text2)
348                > ConversationSummaryBufferMemory::<OpenAIChat>::estimate_tokens(text1)
349        );
350        assert!(ConversationSummaryBufferMemory::<OpenAIChat>::estimate_tokens(text3) > 0);
351    }
352
353    #[test]
354    fn test_prune_messages_within_limit() {
355        let llm = OpenAIChat::new(create_test_config());
356        let memory: ConversationSummaryBufferMemory<OpenAIChat> =
357            ConversationSummaryBufferMemory::new(llm, 1000);
358
359        let messages = vec![
360            Message::human("Short message 1"),
361            Message::ai("Short reply 1"),
362        ];
363
364        let pruned = memory.prune_messages(&messages);
365
366        assert_eq!(pruned.len(), 2);
367    }
368
369    #[tokio::test]
370    async fn test_buffer_initial_empty() {
371        let llm = OpenAIChat::new(create_test_config());
372        let memory: ConversationSummaryBufferMemory<OpenAIChat> =
373            ConversationSummaryBufferMemory::new(llm, 1000);
374
375        let buffer = memory.buffer().await;
376        assert!(buffer.is_empty());
377    }
378
379    #[tokio::test]
380    async fn test_load_memory_variables_empty() {
381        let llm = OpenAIChat::new(create_test_config());
382        let memory: ConversationSummaryBufferMemory<OpenAIChat> =
383            ConversationSummaryBufferMemory::new(llm, 1000);
384
385        let vars = memory.load_memory_variables(&HashMap::new()).await.unwrap();
386        let history = vars.get("history").unwrap().as_str().unwrap();
387
388        assert!(history.is_empty());
389    }
390
391    #[tokio::test]
392    async fn test_clear() {
393        let llm = OpenAIChat::new(create_test_config());
394        let mut memory: ConversationSummaryBufferMemory<OpenAIChat> =
395            ConversationSummaryBufferMemory::new(llm, 1000);
396
397        memory.chat_memory.add_user_message("test");
398        memory.chat_memory.add_ai_message("reply");
399
400        memory.buffer = "Test summary".to_string();
401
402        memory.clear().await.unwrap();
403
404        assert!(memory.buffer().await.is_empty());
405        assert_eq!(memory.chat_memory().len(), 0);
406    }
407}