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(vars: &HashMap<String, serde_json::Value>) -> Vec<Message> {
194 let mut messages = Vec::new();
195 for value in vars.values() {
196 match value {
197 serde_json::Value::Array(items) => {
198 for item in items {
199 if let Ok(msg) = serde_json::from_value::<Message>(item.clone()) {
200 messages.push(msg);
201 } else if let Some(s) = item.as_str() {
202 messages.push(Message::system(s));
203 }
204 }
205 }
206 serde_json::Value::String(s) => messages.push(Message::system(s)),
207 _ => {}
208 }
209 }
210 messages
211}
212
213#[cfg(test)]
214mod tests {
215 use super::*;
216
217 #[test]
218 fn test_chat_message_history() {
219 let mut history = ChatMessageHistory::new();
220
221 history.add_user_message("hello");
222 history.add_ai_message("Hello! How can I help you?");
223 history.add_user_message("introduce yourself");
224
225 assert_eq!(history.len(), 3);
226 assert!(!history.is_empty());
227 }
228
229 #[test]
230 fn test_chat_message_history_to_string() {
231 let mut history = ChatMessageHistory::new();
232
233 history.add_user_message("hello");
234 history.add_ai_message("Hello!");
235
236 let str = history.to_string();
237 assert!(str.contains("Human: hello"));
238 assert!(str.contains("AI: Hello!"));
239 }
240
241 #[test]
242 fn test_chat_message_history_clear() {
243 let mut history = ChatMessageHistory::new();
244
245 history.add_user_message("test");
246 assert_eq!(history.len(), 1);
247
248 history.clear();
249 assert_eq!(history.len(), 0);
250 assert!(history.is_empty());
251 }
252
253 #[test]
255 fn test_memory_variables_to_messages() {
256 let msg = Message::ai("你好");
258 let mut vars = HashMap::new();
259 vars.insert(
260 "history".to_string(),
261 serde_json::json!([serde_json::to_value(&msg).unwrap()]),
262 );
263 let messages = memory_variables_to_messages(&vars);
264 assert_eq!(messages.len(), 1);
265 assert_eq!(messages[0].content, "你好");
266
267 let mut vars = HashMap::new();
269 vars.insert(
270 "history".to_string(),
271 serde_json::Value::String("Human: 在吗\nAI: 在".to_string()),
272 );
273 let messages = memory_variables_to_messages(&vars);
274 assert_eq!(messages.len(), 1);
275 assert_eq!(messages[0].message_type, lc_schema::MessageType::System);
276 }
277}