1use async_trait::async_trait;
5use lc_schema::Message;
6use std::collections::HashMap;
7
8#[derive(Debug, thiserror::Error)]
10pub enum MemoryError {
11 #[error("Failed to load memory: {0}")]
13 LoadError(String),
14
15 #[error("Failed to save memory: {0}")]
17 SaveError(String),
18
19 #[error("Failed to clear memory: {0}")]
21 ClearError(String),
22
23 #[error("Memory error: {0}")]
25 Other(String),
26}
27
28#[async_trait]
32pub trait BaseMemory: Send + Sync {
33 fn memory_variables(&self) -> Vec<&str>;
37
38 async fn load_memory_variables(
46 &self,
47 inputs: &HashMap<String, String>,
48 ) -> Result<HashMap<String, serde_json::Value>, MemoryError>;
49
50 async fn save_context(
61 &mut self,
62 inputs: &HashMap<String, String>,
63 outputs: &HashMap<String, String>,
64 ) -> Result<(), MemoryError>;
65
66 async fn clear(&mut self) -> Result<(), MemoryError>;
68}
69
70pub trait BaseChatMemory: BaseMemory {
78 fn messages(&self) -> &[Message];
80
81 fn add_message(&mut self, message: Message);
83
84 fn add_user_message(&mut self, content: &str) {
86 self.add_message(Message::human(content));
87 }
88
89 fn add_ai_message(&mut self, content: &str) {
91 self.add_message(Message::ai(content));
92 }
93}
94
95#[derive(Debug, Clone)]
99pub struct ChatMessageHistory {
100 messages: Vec<Message>,
102}
103
104impl ChatMessageHistory {
105 pub fn new() -> Self {
107 Self {
108 messages: Vec::new(),
109 }
110 }
111
112 pub fn from_messages(messages: Vec<Message>) -> Self {
114 Self { messages }
115 }
116
117 pub fn add_message(&mut self, message: Message) {
119 self.messages.push(message);
120 }
121
122 pub fn add_user_message(&mut self, content: &str) {
124 self.add_message(Message::human(content));
125 }
126
127 pub fn add_ai_message(&mut self, content: &str) {
129 self.add_message(Message::ai(content));
130 }
131
132 pub fn add_system_message(&mut self, content: &str) {
134 self.add_message(Message::system(content));
135 }
136
137 pub fn messages(&self) -> &[Message] {
139 &self.messages
140 }
141
142 pub fn clear(&mut self) {
144 self.messages.clear();
145 }
146
147 pub fn len(&self) -> usize {
149 self.messages.len()
150 }
151
152 pub fn is_empty(&self) -> bool {
154 self.messages.is_empty()
155 }
156}
157
158impl std::fmt::Display for ChatMessageHistory {
159 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
160 let formatted: String = self
161 .messages
162 .iter()
163 .map(|msg| {
164 let role = match msg.message_type {
165 lc_schema::MessageType::Human => "Human",
166 lc_schema::MessageType::AI => "AI",
167 lc_schema::MessageType::System => "System",
168 lc_schema::MessageType::Tool { .. } => "Tool",
169 };
170 format!("{}: {}", role, msg.content)
171 })
172 .collect::<Vec<_>>()
173 .join("\n");
174 write!(f, "{}", formatted)
175 }
176}
177
178impl Default for ChatMessageHistory {
179 fn default() -> Self {
180 Self::new()
181 }
182}
183
184pub fn memory_variables_to_messages(
194 vars: &HashMap<String, serde_json::Value>,
195) -> Vec<Message> {
196 let mut messages = Vec::new();
197 for value in vars.values() {
198 match value {
199 serde_json::Value::Array(items) => {
200 for item in items {
201 if let Ok(msg) = serde_json::from_value::<Message>(item.clone()) {
202 messages.push(msg);
203 } else if let Some(s) = item.as_str() {
204 messages.push(Message::system(s));
205 }
206 }
207 }
208 serde_json::Value::String(s) => messages.push(Message::system(s)),
209 _ => {}
210 }
211 }
212 messages
213}
214
215#[cfg(test)]
216mod tests {
217 use super::*;
218
219 #[test]
220 fn test_chat_message_history() {
221 let mut history = ChatMessageHistory::new();
222
223 history.add_user_message("hello");
224 history.add_ai_message("Hello! How can I help you?");
225 history.add_user_message("introduce yourself");
226
227 assert_eq!(history.len(), 3);
228 assert!(!history.is_empty());
229 }
230
231 #[test]
232 fn test_chat_message_history_to_string() {
233 let mut history = ChatMessageHistory::new();
234
235 history.add_user_message("hello");
236 history.add_ai_message("Hello!");
237
238 let str = history.to_string();
239 assert!(str.contains("Human: hello"));
240 assert!(str.contains("AI: Hello!"));
241 }
242
243 #[test]
244 fn test_chat_message_history_clear() {
245 let mut history = ChatMessageHistory::new();
246
247 history.add_user_message("test");
248 assert_eq!(history.len(), 1);
249
250 history.clear();
251 assert_eq!(history.len(), 0);
252 assert!(history.is_empty());
253 }
254
255 #[test]
257 fn test_memory_variables_to_messages() {
258 let msg = Message::ai("你好");
260 let mut vars = HashMap::new();
261 vars.insert(
262 "history".to_string(),
263 serde_json::json!([serde_json::to_value(&msg).unwrap()]),
264 );
265 let messages = memory_variables_to_messages(&vars);
266 assert_eq!(messages.len(), 1);
267 assert_eq!(messages[0].content, "你好");
268
269 let mut vars = HashMap::new();
271 vars.insert(
272 "history".to_string(),
273 serde_json::Value::String("Human: 在吗\nAI: 在".to_string()),
274 );
275 let messages = memory_variables_to_messages(&vars);
276 assert_eq!(messages.len(), 1);
277 assert_eq!(messages[0].message_type, lc_schema::MessageType::System);
278 }
279}