use foundation_db::traits::KeyValueStore;
use foundation_db::{StorageError, StorageResult};
use serde::{Deserialize, Serialize};
use crate::types::{SessionId, SessionRecord};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum MemoryTier {
Working,
Observation,
Reflection,
}
impl MemoryTier {
#[must_use]
pub fn of(record: &SessionRecord) -> Option<Self> {
match record {
SessionRecord::WorkingMemory { .. } => Some(MemoryTier::Working),
SessionRecord::Observation { .. } => Some(MemoryTier::Observation),
SessionRecord::Reflection { .. } => Some(MemoryTier::Reflection),
_ => None, }
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, Default)]
pub struct SessionMemory {
#[serde(skip_serializing_if = "Option::is_none")]
pub working: Option<SessionRecord>,
#[serde(skip_serializing_if = "Option::is_none")]
pub observation: Option<SessionRecord>,
#[serde(skip_serializing_if = "Option::is_none")]
pub reflection: Option<SessionRecord>,
}
impl SessionMemory {
fn get(&self, tier: MemoryTier) -> Option<&SessionRecord> {
match tier {
MemoryTier::Working => self.working.as_ref(),
MemoryTier::Observation => self.observation.as_ref(),
MemoryTier::Reflection => self.reflection.as_ref(),
}
}
#[must_use]
pub fn set(mut self, tier: MemoryTier, record: SessionRecord) -> Self {
match tier {
MemoryTier::Working => self.working = Some(record),
MemoryTier::Observation => self.observation = Some(record),
MemoryTier::Reflection => self.reflection = Some(record),
}
self
}
#[must_use]
pub fn key(session: &SessionId) -> String {
format!("memory:{session}")
}
}
#[async_trait::async_trait]
pub trait MemoryStore: Send + Sync {
async fn get_async(
&self,
session: &SessionId,
tier: MemoryTier,
) -> StorageResult<Option<SessionRecord>>;
async fn set_async(&self, session: &SessionId, record: &SessionRecord) -> StorageResult<()>;
async fn hydrate_async(&self, session: &SessionId) -> StorageResult<SessionMemory>;
async fn clear_async(&self, session: &SessionId) -> StorageResult<()>;
fn hydrate_sync(&self, session: &SessionId) -> StorageResult<SessionMemory> {
let _ = session;
Ok(SessionMemory::default())
}
}
pub struct KvMemoryStore<K> {
kv: K,
}
impl<K> KvMemoryStore<K> {
pub fn new(kv: K) -> Self {
Self { kv }
}
}
impl<K: Default> Default for KvMemoryStore<K> {
fn default() -> Self {
Self { kv: K::default() }
}
}
#[async_trait::async_trait]
impl<K: KeyValueStore> MemoryStore for KvMemoryStore<K> {
async fn get_async(
&self,
session: &SessionId,
tier: MemoryTier,
) -> StorageResult<Option<SessionRecord>> {
let key = SessionMemory::key(session);
let mem: Option<SessionMemory> = self.kv.get(&key)?;
Ok(mem.and_then(|m| m.get(tier).cloned()))
}
async fn set_async(&self, session: &SessionId, record: &SessionRecord) -> StorageResult<()> {
let tier = MemoryTier::of(record).ok_or_else(|| {
StorageError::Backend(
"set_async: record is not a memory variant (Working/Observation/Reflection)".into(),
)
})?;
let key = SessionMemory::key(session);
let mem: SessionMemory = self.kv.get(&key)?.unwrap_or_default();
self.kv.set(&key, mem.set(tier, record.clone()))
}
async fn hydrate_async(&self, session: &SessionId) -> StorageResult<SessionMemory> {
let key = SessionMemory::key(session);
self.kv
.get(&key)
.map(std::option::Option::unwrap_or_default)
}
async fn clear_async(&self, session: &SessionId) -> StorageResult<()> {
self.kv.delete(&SessionMemory::key(session))
}
fn hydrate_sync(&self, session: &SessionId) -> StorageResult<SessionMemory> {
let key = SessionMemory::key(session);
self.kv
.get(&key)
.map(std::option::Option::unwrap_or_default)
}
}