lc_memory/
summary_buffer.rs1use 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
28pub struct ConversationSummaryBufferMemory<M: BaseChatModel> {
47 llm: M,
48
49 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 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}