use std::{
collections::HashMap,
sync::{Arc, Mutex},
};
use crate::id::ConversationId;
use crate::{
completion::Message,
wasm_compat::{WasmBoxedFuture, WasmCompatSend, WasmCompatSync},
};
#[cfg(not(target_family = "wasm"))]
pub type MemoryBackendError = Box<dyn std::error::Error + Send + Sync + 'static>;
#[cfg(target_family = "wasm")]
pub type MemoryBackendError = Box<dyn std::error::Error + 'static>;
#[derive(Debug, thiserror::Error)]
pub enum MemoryError {
#[error("Memory backend error: {0}")]
Backend(#[source] MemoryBackendError),
#[error("Memory policy error: {0}")]
Policy(String),
#[error("Memory internal error: {0}")]
Internal(String),
}
impl MemoryError {
pub fn backend<E>(source: E) -> Self
where
E: Into<MemoryBackendError>,
{
Self::Backend(source.into())
}
}
pub trait ConversationMemory: WasmCompatSend + WasmCompatSync {
fn load<'a>(
&'a self,
conversation_id: &'a ConversationId,
) -> WasmBoxedFuture<'a, Result<Vec<Message>, MemoryError>>;
fn append<'a>(
&'a self,
conversation_id: &'a ConversationId,
messages: Vec<Message>,
) -> WasmBoxedFuture<'a, Result<(), MemoryError>>;
fn clear<'a>(
&'a self,
conversation_id: &'a ConversationId,
) -> WasmBoxedFuture<'a, Result<(), MemoryError>>;
}
macro_rules! forward_memory_trait {
(ConversationMemory: $($ptr:ident)+) => {$(
impl<M> ConversationMemory for $ptr<M>
where
M: ConversationMemory + ?Sized,
{
fn load<'a>(
&'a self,
conversation_id: &'a ConversationId,
) -> WasmBoxedFuture<'a, Result<Vec<Message>, MemoryError>> {
(**self).load(conversation_id)
}
fn append<'a>(
&'a self,
conversation_id: &'a ConversationId,
messages: Vec<Message>,
) -> WasmBoxedFuture<'a, Result<(), MemoryError>> {
(**self).append(conversation_id, messages)
}
fn clear<'a>(
&'a self,
conversation_id: &'a ConversationId,
) -> WasmBoxedFuture<'a, Result<(), MemoryError>> {
(**self).clear(conversation_id)
}
}
)+};
(DemotionHook: $($ptr:ident)+) => {$(
impl<H> DemotionHook for $ptr<H>
where
H: DemotionHook + ?Sized,
{
fn on_demote<'a>(
&'a self,
conversation_id: &'a ConversationId,
messages: Vec<Message>,
) -> WasmBoxedFuture<'a, Result<(), MemoryError>> {
(**self).on_demote(conversation_id, messages)
}
}
)+};
(Compactor: $($ptr:ident)+) => {$(
impl<C> Compactor for $ptr<C>
where
C: Compactor + ?Sized,
{
type Artifact = C::Artifact;
fn compact<'a>(
&'a self,
conversation_id: &'a ConversationId,
evicted: &'a [Message],
carry_over: Option<&'a Self::Artifact>,
) -> WasmBoxedFuture<'a, Result<Self::Artifact, MemoryError>> {
(**self).compact(conversation_id, evicted, carry_over)
}
}
)+};
}
forward_memory_trait!(ConversationMemory: Arc Box);
pub trait MessageFilter:
Fn(Vec<Message>) -> Vec<Message> + WasmCompatSend + WasmCompatSync
{
}
impl<F> MessageFilter for F where
F: Fn(Vec<Message>) -> Vec<Message> + WasmCompatSend + WasmCompatSync
{
}
pub trait DemotionHook: WasmCompatSend + WasmCompatSync {
fn on_demote<'a>(
&'a self,
conversation_id: &'a ConversationId,
messages: Vec<Message>,
) -> WasmBoxedFuture<'a, Result<(), MemoryError>>;
}
#[derive(Debug, Default, Clone, Copy)]
pub struct NoopDemotionHook;
impl DemotionHook for NoopDemotionHook {
fn on_demote<'a>(
&'a self,
_conversation_id: &'a ConversationId,
_messages: Vec<Message>,
) -> WasmBoxedFuture<'a, Result<(), MemoryError>> {
Box::pin(async move { Ok(()) })
}
}
forward_memory_trait!(DemotionHook: Arc);
pub trait Compactor: WasmCompatSend + WasmCompatSync {
type Artifact: Into<Message> + Clone + WasmCompatSend + WasmCompatSync + 'static;
fn compact<'a>(
&'a self,
conversation_id: &'a ConversationId,
evicted: &'a [Message],
carry_over: Option<&'a Self::Artifact>,
) -> WasmBoxedFuture<'a, Result<Self::Artifact, MemoryError>>;
}
forward_memory_trait!(Compactor: Arc);
#[derive(Clone, Default)]
pub struct InMemoryConversationMemory {
inner: Arc<Mutex<HashMap<ConversationId, Vec<Message>>>>,
filter: Option<Arc<dyn MessageFilter>>,
}
impl InMemoryConversationMemory {
pub fn new() -> Self {
Self::default()
}
pub fn with_filter<F>(mut self, filter: F) -> Self
where
F: MessageFilter + 'static,
{
self.filter = Some(Arc::new(filter));
self
}
fn lock(
&self,
) -> Result<std::sync::MutexGuard<'_, HashMap<ConversationId, Vec<Message>>>, MemoryError> {
self.inner
.lock()
.map_err(|e| MemoryError::Internal(e.to_string()))
}
}
impl std::fmt::Debug for InMemoryConversationMemory {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("InMemoryConversationMemory")
.field("filter", &self.filter.as_ref().map(|_| "<filter>"))
.finish()
}
}
impl ConversationMemory for InMemoryConversationMemory {
fn load<'a>(
&'a self,
conversation_id: &'a ConversationId,
) -> WasmBoxedFuture<'a, Result<Vec<Message>, MemoryError>> {
Box::pin(async move {
let messages = {
let guard = self.lock()?;
guard.get(conversation_id).cloned().unwrap_or_default()
};
match &self.filter {
Some(filter) => Ok(filter(messages)),
None => Ok(messages),
}
})
}
fn append<'a>(
&'a self,
conversation_id: &'a ConversationId,
messages: Vec<Message>,
) -> WasmBoxedFuture<'a, Result<(), MemoryError>> {
Box::pin(async move {
let mut guard = self.lock()?;
guard
.entry(conversation_id.clone())
.or_default()
.extend(messages);
Ok(())
})
}
fn clear<'a>(
&'a self,
conversation_id: &'a ConversationId,
) -> WasmBoxedFuture<'a, Result<(), MemoryError>> {
Box::pin(async move {
let mut guard = self.lock()?;
guard.remove(conversation_id);
Ok(())
})
}
}
#[cfg(test)]
mod tests;