use std::future::Future;
use std::pin::Pin;
use crate::error::Result;
use crate::types::Message;
pub trait Memory: Send + Sync {
fn add_message(&self, message: &Message) -> impl Future<Output = Result<()>> + Send;
fn get_messages(&self) -> impl Future<Output = Result<Vec<Message>>> + Send;
fn with_messages<R, F>(&self, f: F) -> impl Future<Output = Result<R>> + Send
where
F: FnOnce(&[Message]) -> R + Send,
R: Send,
{
async move { Ok(f(&self.get_messages().await?)) }
}
fn clear(&self) -> impl Future<Output = Result<()>> + Send;
}
pub trait ErasedMemory: Send + Sync {
fn add_message_erased<'a>(
&'a self,
message: &'a Message,
) -> Pin<Box<dyn Future<Output = Result<()>> + Send + 'a>>;
fn get_messages_erased<'a>(
&'a self,
) -> Pin<Box<dyn Future<Output = Result<Vec<Message>>> + Send + 'a>>;
fn with_messages_erased<'a>(
&'a self,
f: Box<dyn FnOnce(&[Message]) + Send + 'a>,
) -> Pin<Box<dyn Future<Output = Result<()>> + Send + 'a>> {
Box::pin(async move {
let messages = self.get_messages_erased().await?;
f(&messages);
Ok(())
})
}
fn clear_erased<'a>(&'a self) -> Pin<Box<dyn Future<Output = Result<()>> + Send + 'a>>;
}
impl<T: Memory> ErasedMemory for T {
fn add_message_erased<'a>(
&'a self,
message: &'a Message,
) -> Pin<Box<dyn Future<Output = Result<()>> + Send + 'a>> {
Box::pin(self.add_message(message))
}
fn get_messages_erased<'a>(
&'a self,
) -> Pin<Box<dyn Future<Output = Result<Vec<Message>>> + Send + 'a>> {
Box::pin(self.get_messages())
}
fn with_messages_erased<'a>(
&'a self,
f: Box<dyn FnOnce(&[Message]) + Send + 'a>,
) -> Pin<Box<dyn Future<Output = Result<()>> + Send + 'a>> {
Box::pin(self.with_messages(f))
}
fn clear_erased<'a>(&'a self) -> Pin<Box<dyn Future<Output = Result<()>> + Send + 'a>> {
Box::pin(self.clear())
}
}
pub type SharedMemory = std::sync::Arc<dyn ErasedMemory>;
#[cfg(test)]
mod tests {
use super::{ErasedMemory, Memory, SharedMemory};
use crate::{Message, Result, Role};
use std::future::Future;
use std::pin::Pin;
use std::sync::{Arc, Mutex};
struct VecMemory(Mutex<Vec<Message>>);
impl Memory for VecMemory {
async fn add_message(&self, message: &Message) -> Result<()> {
self.0.lock().unwrap().push(message.clone());
Ok(())
}
async fn get_messages(&self) -> Result<Vec<Message>> {
Ok(self.0.lock().unwrap().clone())
}
async fn clear(&self) -> Result<()> {
self.0.lock().unwrap().clear();
Ok(())
}
}
#[tokio::test]
async fn memory_is_implementable_from_core_alone() {
let mem = VecMemory(Mutex::new(Vec::new()));
mem.add_message(&Message::user("hi")).await.unwrap();
assert_eq!(mem.get_messages().await.unwrap().len(), 1);
assert_eq!(mem.get_messages().await.unwrap()[0].role, Role::User);
mem.clear().await.unwrap();
assert!(mem.get_messages().await.unwrap().is_empty());
let shared: SharedMemory = Arc::new(VecMemory(Mutex::new(Vec::new())));
shared
.add_message_erased(&Message::user("x"))
.await
.unwrap();
assert_eq!(shared.get_messages_erased().await.unwrap().len(), 1);
}
#[tokio::test]
async fn with_messages_default_matches_get_messages() {
let mem = VecMemory(Mutex::new(Vec::new()));
mem.add_message(&Message::user("one")).await.unwrap();
mem.add_message(&Message::assistant("two")).await.unwrap();
let owned = mem.get_messages().await.unwrap();
let borrowed = mem
.with_messages(|messages| {
messages
.iter()
.map(|m| (m.role.clone(), m.content.clone()))
.collect::<Vec<_>>()
})
.await
.unwrap();
assert_eq!(borrowed.len(), owned.len());
for (seen, expected) in borrowed.iter().zip(&owned) {
assert_eq!(seen.0, expected.role);
assert_eq!(seen.1, expected.content);
}
}
#[tokio::test]
async fn with_messages_erased_works_through_shared_memory() {
let shared: SharedMemory = Arc::new(VecMemory(Mutex::new(Vec::new())));
shared
.add_message_erased(&Message::user("x"))
.await
.unwrap();
shared
.add_message_erased(&Message::assistant("y"))
.await
.unwrap();
let mut seen: Vec<Option<String>> = Vec::new();
shared
.with_messages_erased(Box::new(|messages| {
seen.extend(messages.iter().map(|m| m.content.clone()));
}))
.await
.unwrap();
assert_eq!(seen, vec![Some("x".into()), Some("y".into())]);
}
struct DirectErased(Mutex<Vec<Message>>);
impl ErasedMemory for DirectErased {
fn add_message_erased<'a>(
&'a self,
message: &'a Message,
) -> Pin<Box<dyn Future<Output = Result<()>> + Send + 'a>> {
Box::pin(async move {
self.0.lock().unwrap().push(message.clone());
Ok(())
})
}
fn get_messages_erased<'a>(
&'a self,
) -> Pin<Box<dyn Future<Output = Result<Vec<Message>>> + Send + 'a>> {
Box::pin(async move { Ok(self.0.lock().unwrap().clone()) })
}
fn clear_erased<'a>(&'a self) -> Pin<Box<dyn Future<Output = Result<()>> + Send + 'a>> {
Box::pin(async move {
self.0.lock().unwrap().clear();
Ok(())
})
}
}
#[tokio::test]
async fn direct_erased_impl_gets_default_with_messages_erased() {
let shared: SharedMemory = Arc::new(DirectErased(Mutex::new(Vec::new())));
shared
.add_message_erased(&Message::user("hello"))
.await
.unwrap();
let mut count = 0usize;
shared
.with_messages_erased(Box::new(|messages| count = messages.len()))
.await
.unwrap();
assert_eq!(count, 1);
}
}