use async_trait::async_trait;
use lc_schema::Message;
use std::collections::HashMap;
#[derive(Debug, thiserror::Error)]
pub enum MemoryError {
#[error("Failed to load memory: {0}")]
LoadError(String),
#[error("Failed to save memory: {0}")]
SaveError(String),
#[error("Failed to clear memory: {0}")]
ClearError(String),
#[error("Memory error: {0}")]
Other(String),
}
#[async_trait]
pub trait BaseMemory: Send + Sync {
fn memory_variables(&self) -> Vec<&str>;
async fn load_memory_variables(
&self,
inputs: &HashMap<String, String>,
) -> Result<HashMap<String, serde_json::Value>, MemoryError>;
async fn save_context(
&mut self,
inputs: &HashMap<String, String>,
outputs: &HashMap<String, String>,
) -> Result<(), MemoryError>;
async fn clear(&mut self) -> Result<(), MemoryError>;
}
pub trait BaseChatMemory: BaseMemory {
fn messages(&self) -> &[Message];
fn add_message(&mut self, message: Message);
fn add_user_message(&mut self, content: &str) {
self.add_message(Message::human(content));
}
fn add_ai_message(&mut self, content: &str) {
self.add_message(Message::ai(content));
}
}
#[derive(Debug, Clone)]
pub struct ChatMessageHistory {
messages: Vec<Message>,
}
impl ChatMessageHistory {
pub fn new() -> Self {
Self {
messages: Vec::new(),
}
}
pub fn from_messages(messages: Vec<Message>) -> Self {
Self { messages }
}
pub fn add_message(&mut self, message: Message) {
self.messages.push(message);
}
pub fn add_user_message(&mut self, content: &str) {
self.add_message(Message::human(content));
}
pub fn add_ai_message(&mut self, content: &str) {
self.add_message(Message::ai(content));
}
pub fn add_system_message(&mut self, content: &str) {
self.add_message(Message::system(content));
}
pub fn messages(&self) -> &[Message] {
&self.messages
}
pub fn clear(&mut self) {
self.messages.clear();
}
pub fn len(&self) -> usize {
self.messages.len()
}
pub fn is_empty(&self) -> bool {
self.messages.is_empty()
}
}
impl std::fmt::Display for ChatMessageHistory {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let formatted: String = self
.messages
.iter()
.map(|msg| {
let role = match msg.message_type {
lc_schema::MessageType::Human => "Human",
lc_schema::MessageType::AI => "AI",
lc_schema::MessageType::System => "System",
lc_schema::MessageType::Tool { .. } => "Tool",
};
format!("{}: {}", role, msg.content)
})
.collect::<Vec<_>>()
.join("\n");
write!(f, "{}", formatted)
}
}
impl Default for ChatMessageHistory {
fn default() -> Self {
Self::new()
}
}
pub fn memory_variables_to_messages(
vars: &HashMap<String, serde_json::Value>,
) -> Vec<Message> {
let mut messages = Vec::new();
for value in vars.values() {
match value {
serde_json::Value::Array(items) => {
for item in items {
if let Ok(msg) = serde_json::from_value::<Message>(item.clone()) {
messages.push(msg);
} else if let Some(s) = item.as_str() {
messages.push(Message::system(s));
}
}
}
serde_json::Value::String(s) => messages.push(Message::system(s)),
_ => {}
}
}
messages
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_chat_message_history() {
let mut history = ChatMessageHistory::new();
history.add_user_message("hello");
history.add_ai_message("Hello! How can I help you?");
history.add_user_message("introduce yourself");
assert_eq!(history.len(), 3);
assert!(!history.is_empty());
}
#[test]
fn test_chat_message_history_to_string() {
let mut history = ChatMessageHistory::new();
history.add_user_message("hello");
history.add_ai_message("Hello!");
let str = history.to_string();
assert!(str.contains("Human: hello"));
assert!(str.contains("AI: Hello!"));
}
#[test]
fn test_chat_message_history_clear() {
let mut history = ChatMessageHistory::new();
history.add_user_message("test");
assert_eq!(history.len(), 1);
history.clear();
assert_eq!(history.len(), 0);
assert!(history.is_empty());
}
#[test]
fn test_memory_variables_to_messages() {
let msg = Message::ai("你好");
let mut vars = HashMap::new();
vars.insert(
"history".to_string(),
serde_json::json!([serde_json::to_value(&msg).unwrap()]),
);
let messages = memory_variables_to_messages(&vars);
assert_eq!(messages.len(), 1);
assert_eq!(messages[0].content, "你好");
let mut vars = HashMap::new();
vars.insert(
"history".to_string(),
serde_json::Value::String("Human: 在吗\nAI: 在".to_string()),
);
let messages = memory_variables_to_messages(&vars);
assert_eq!(messages.len(), 1);
assert_eq!(messages[0].message_type, lc_schema::MessageType::System);
}
}