use crate::error::RragResult;
use crate::storage::{Memory, MemoryValue};
use std::sync::Arc;
pub struct WorkingMemory {
storage: Arc<dyn Memory>,
namespace: String,
auto_clear: bool,
}
impl WorkingMemory {
pub fn new(storage: Arc<dyn Memory>, session_id: String) -> Self {
let namespace = format!("session::{}::working", session_id);
Self {
storage,
namespace,
auto_clear: true,
}
}
pub fn new_persistent(storage: Arc<dyn Memory>, session_id: String) -> Self {
let namespace = format!("session::{}::working", session_id);
Self {
storage,
namespace,
auto_clear: false,
}
}
pub async fn set(&self, key: &str, value: impl Into<MemoryValue>) -> RragResult<()> {
let full_key = self.make_key(key);
self.storage.set(&full_key, value.into()).await
}
pub async fn get(&self, key: &str) -> RragResult<Option<MemoryValue>> {
let full_key = self.make_key(key);
self.storage.get(&full_key).await
}
pub async fn delete(&self, key: &str) -> RragResult<bool> {
let full_key = self.make_key(key);
self.storage.delete(&full_key).await
}
pub async fn exists(&self, key: &str) -> RragResult<bool> {
let full_key = self.make_key(key);
self.storage.exists(&full_key).await
}
pub async fn clear(&self) -> RragResult<()> {
self.storage.clear(Some(&self.namespace)).await
}
pub async fn keys(&self) -> RragResult<Vec<String>> {
use crate::storage::MemoryQuery;
let query = MemoryQuery::new().with_namespace(self.namespace.clone());
let all_keys = self.storage.keys(&query).await?;
let prefix = format!("{}::", self.namespace);
let keys = all_keys
.into_iter()
.filter_map(|k| k.strip_prefix(&prefix).map(String::from))
.collect();
Ok(keys)
}
pub async fn set_many(&self, pairs: &[(&str, MemoryValue)]) -> RragResult<()> {
let full_pairs: Vec<(String, MemoryValue)> = pairs
.iter()
.map(|(k, v)| (self.make_key(k), v.clone()))
.collect();
self.storage.mset(&full_pairs).await
}
pub async fn get_many(&self, keys: &[&str]) -> RragResult<Vec<Option<MemoryValue>>> {
let full_keys: Vec<String> = keys.iter().map(|k| self.make_key(k)).collect();
self.storage.mget(&full_keys).await
}
pub async fn count(&self) -> RragResult<usize> {
self.storage.count(Some(&self.namespace)).await
}
fn make_key(&self, key: &str) -> String {
format!("{}::{}", self.namespace, key)
}
pub fn disable_auto_clear(&mut self) {
self.auto_clear = false;
}
pub fn enable_auto_clear(&mut self) {
self.auto_clear = true;
}
}
impl Drop for WorkingMemory {
fn drop(&mut self) {
if self.auto_clear {
tracing::debug!(
namespace = %self.namespace,
"WorkingMemory dropped with auto_clear enabled - cleanup deferred"
);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::storage::InMemoryStorage;
#[tokio::test]
async fn test_working_memory_basic_operations() {
let storage = Arc::new(InMemoryStorage::new());
let working = WorkingMemory::new(storage, "test-session".to_string());
working
.set("temp_result", MemoryValue::from(42i64))
.await
.unwrap();
let value = working.get("temp_result").await.unwrap();
assert_eq!(value.unwrap().as_integer(), Some(42));
assert!(working.exists("temp_result").await.unwrap());
assert!(!working.exists("nonexistent").await.unwrap());
assert!(working.delete("temp_result").await.unwrap());
assert!(!working.exists("temp_result").await.unwrap());
}
#[tokio::test]
async fn test_working_memory_multiple_operations() {
let storage = Arc::new(InMemoryStorage::new());
let working = WorkingMemory::new(storage, "test-session".to_string());
let pairs = [
("key1", MemoryValue::from("value1")),
("key2", MemoryValue::from(100i64)),
("key3", MemoryValue::from(true)),
];
working.set_many(&pairs).await.unwrap();
let keys = ["key1", "key2", "key3"];
let values = working.get_many(&keys).await.unwrap();
assert_eq!(values[0].as_ref().unwrap().as_string(), Some("value1"));
assert_eq!(values[1].as_ref().unwrap().as_integer(), Some(100));
assert_eq!(values[2].as_ref().unwrap().as_boolean(), Some(true));
assert_eq!(working.count().await.unwrap(), 3);
let all_keys = working.keys().await.unwrap();
assert_eq!(all_keys.len(), 3);
}
#[tokio::test]
async fn test_working_memory_clear() {
let storage = Arc::new(InMemoryStorage::new());
let working = WorkingMemory::new(storage, "test-session".to_string());
working
.set("key1", MemoryValue::from("value1"))
.await
.unwrap();
working
.set("key2", MemoryValue::from("value2"))
.await
.unwrap();
assert_eq!(working.count().await.unwrap(), 2);
working.clear().await.unwrap();
assert_eq!(working.count().await.unwrap(), 0);
}
#[tokio::test]
async fn test_working_memory_namespace_isolation() {
let storage = Arc::new(InMemoryStorage::new());
let working1 = WorkingMemory::new(storage.clone(), "session1".to_string());
let working2 = WorkingMemory::new(storage.clone(), "session2".to_string());
working1
.set("data", MemoryValue::from("session1-data"))
.await
.unwrap();
working2
.set("data", MemoryValue::from("session2-data"))
.await
.unwrap();
let value1 = working1.get("data").await.unwrap();
let value2 = working2.get("data").await.unwrap();
assert_eq!(value1.unwrap().as_string(), Some("session1-data"));
assert_eq!(value2.unwrap().as_string(), Some("session2-data"));
}
}