use std::collections::HashMap;
use std::future::Future;
use std::pin::Pin;
use std::sync::Mutex;
use serde::{Deserialize, Serialize};
use crate::error::MemoryError;
use crate::tenant::TenantContext;
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct MemoryRecord {
pub key: String,
pub value: String,
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct MemoryQuery {
pub key: String,
}
pub trait MemoryProvider: Send + Sync {
fn remember<'a>(
&'a self,
tenant: &'a TenantContext,
session_id: &'a str,
record: MemoryRecord,
) -> Pin<Box<dyn Future<Output = Result<(), MemoryError>> + Send + 'a>>;
fn recall<'a>(
&'a self,
tenant: &'a TenantContext,
session_id: &'a str,
query: &'a MemoryQuery,
) -> Pin<Box<dyn Future<Output = Result<Option<MemoryRecord>, MemoryError>> + Send + 'a>>;
}
#[derive(Default)]
pub struct InMemoryMemoryProvider {
entries: Mutex<HashMap<(String, String, String), MemoryRecord>>,
}
impl InMemoryMemoryProvider {
pub fn new() -> Self {
Self {
entries: Mutex::new(HashMap::new()),
}
}
}
impl MemoryProvider for InMemoryMemoryProvider {
fn remember<'a>(
&'a self,
tenant: &'a TenantContext,
session_id: &'a str,
record: MemoryRecord,
) -> Pin<Box<dyn Future<Output = Result<(), MemoryError>> + Send + 'a>> {
let key = (
tenant.key_prefix(),
session_id.to_string(),
record.key.clone(),
);
let result = self
.entries
.lock()
.map_err(|_| MemoryError::Backend("memory mutex poisoned".to_string()))
.map(|mut entries| {
entries.insert(key, record);
});
Box::pin(async move { result })
}
fn recall<'a>(
&'a self,
tenant: &'a TenantContext,
session_id: &'a str,
query: &'a MemoryQuery,
) -> Pin<Box<dyn Future<Output = Result<Option<MemoryRecord>, MemoryError>> + Send + 'a>> {
let key = (
tenant.key_prefix(),
session_id.to_string(),
query.key.clone(),
);
let result = self
.entries
.lock()
.map_err(|_| MemoryError::Backend("memory mutex poisoned".to_string()))
.map(|entries| entries.get(&key).cloned());
Box::pin(async move { result })
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used)]
mod tests {
use super::*;
fn tenant() -> TenantContext {
TenantContext::new("acme", "prod")
}
#[tokio::test]
async fn remember_then_recall_roundtrips() {
let provider = InMemoryMemoryProvider::new();
let t = tenant();
provider
.remember(
&t,
"sess-1",
MemoryRecord {
key: "fav_color".to_string(),
value: "green".to_string(),
},
)
.await
.unwrap();
let got = provider
.recall(
&t,
"sess-1",
&MemoryQuery {
key: "fav_color".to_string(),
},
)
.await
.unwrap();
assert_eq!(
got,
Some(MemoryRecord {
key: "fav_color".to_string(),
value: "green".to_string(),
})
);
}
#[tokio::test]
async fn recall_missing_key_returns_none() {
let provider = InMemoryMemoryProvider::new();
let got = provider
.recall(
&tenant(),
"sess-1",
&MemoryQuery {
key: "nope".to_string(),
},
)
.await
.unwrap();
assert!(got.is_none());
}
#[tokio::test]
async fn recall_is_isolated_across_tenants() {
let provider = InMemoryMemoryProvider::new();
provider
.remember(
&TenantContext::new("acme", "prod"),
"sess-1",
MemoryRecord {
key: "fav_color".to_string(),
value: "green".to_string(),
},
)
.await
.unwrap();
let other = provider
.recall(
&TenantContext::new("globex", "prod"),
"sess-1",
&MemoryQuery {
key: "fav_color".to_string(),
},
)
.await
.unwrap();
assert!(other.is_none());
}
}