use crate::settings::dynamic::{DynamicBackend, DynamicResult};
use async_trait::async_trait;
use parking_lot::RwLock;
use std::collections::HashMap;
use std::sync::Arc;
use std::time::{Duration, Instant};
#[derive(Clone)]
struct ValueEntry {
value: serde_json::Value,
expires_at: Option<Instant>,
}
impl ValueEntry {
fn new(value: serde_json::Value, ttl: Option<u64>) -> Self {
Self {
value,
expires_at: ttl.map(|secs| Instant::now() + Duration::from_secs(secs)),
}
}
fn is_expired(&self) -> bool {
self.expires_at
.map(|expires| Instant::now() >= expires)
.unwrap_or(false)
}
}
pub struct MemoryBackend {
data: Arc<RwLock<HashMap<String, ValueEntry>>>,
}
impl MemoryBackend {
pub fn new() -> Self {
Self {
data: Arc::new(RwLock::new(HashMap::new())),
}
}
pub fn cleanup_expired(&self) {
let mut data = self.data.write();
data.retain(|_, entry| !entry.is_expired());
}
pub fn len(&self) -> usize {
self.data.read().len()
}
pub fn is_empty(&self) -> bool {
self.data.read().is_empty()
}
pub fn clear(&self) {
self.data.write().clear();
}
}
impl Default for MemoryBackend {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl DynamicBackend for MemoryBackend {
async fn get(&self, key: &str) -> DynamicResult<Option<serde_json::Value>> {
let data = self.data.read();
if let Some(entry) = data.get(key) {
if entry.is_expired() {
drop(data);
self.data.write().remove(key);
Ok(None)
} else {
Ok(Some(entry.value.clone()))
}
} else {
Ok(None)
}
}
async fn set(
&self,
key: &str,
value: &serde_json::Value,
ttl: Option<u64>,
) -> DynamicResult<()> {
let entry = ValueEntry::new(value.clone(), ttl);
self.data.write().insert(key.to_string(), entry);
Ok(())
}
async fn delete(&self, key: &str) -> DynamicResult<()> {
self.data.write().remove(key);
Ok(())
}
async fn exists(&self, key: &str) -> DynamicResult<bool> {
let data = self.data.read();
if let Some(entry) = data.get(key) {
if entry.is_expired() {
drop(data);
self.data.write().remove(key);
Ok(false)
} else {
Ok(true)
}
} else {
Ok(false)
}
}
async fn keys(&self) -> DynamicResult<Vec<String>> {
let data = self.data.read();
let valid_keys: Vec<String> = data
.iter()
.filter(|(_, entry)| !entry.is_expired())
.map(|(key, _)| key.clone())
.collect();
if valid_keys.len() < data.len() {
drop(data);
self.cleanup_expired();
}
Ok(valid_keys)
}
}
#[cfg(test)]
mod tests {
use super::*;
async fn poll_until<F, Fut>(
timeout: std::time::Duration,
interval: std::time::Duration,
mut condition: F,
) -> Result<(), String>
where
F: FnMut() -> Fut,
Fut: std::future::Future<Output = bool>,
{
let start = std::time::Instant::now();
while start.elapsed() < timeout {
if condition().await {
return Ok(());
}
tokio::time::sleep(interval).await;
}
Err(format!("Condition not met within {:?}", timeout))
}
#[tokio::test]
async fn test_basic_operations() {
let backend = MemoryBackend::new();
backend
.set("key1", &serde_json::json!("value1"), None)
.await
.unwrap();
let value = backend.get("key1").await.unwrap();
assert_eq!(value, Some(serde_json::json!("value1")));
assert!(backend.exists("key1").await.unwrap());
assert!(!backend.exists("nonexistent").await.unwrap());
backend.delete("key1").await.unwrap();
assert!(!backend.exists("key1").await.unwrap());
}
#[tokio::test]
async fn test_ttl_expiration() {
let backend = MemoryBackend::new();
backend
.set("temp_key", &serde_json::json!("temp_value"), Some(1))
.await
.unwrap();
assert!(backend.exists("temp_key").await.unwrap());
poll_until(
Duration::from_millis(1200),
Duration::from_millis(50),
|| async { !backend.exists("temp_key").await.unwrap() },
)
.await
.expect("Key should expire within 1200ms");
assert!(!backend.exists("temp_key").await.unwrap());
assert_eq!(backend.get("temp_key").await.unwrap(), None);
}
#[tokio::test]
async fn test_multiple_values() {
let backend = MemoryBackend::new();
backend
.set("string", &serde_json::json!("text"), None)
.await
.unwrap();
backend
.set("number", &serde_json::json!(42), None)
.await
.unwrap();
backend
.set("boolean", &serde_json::json!(true), None)
.await
.unwrap();
backend
.set("object", &serde_json::json!({"key": "value"}), None)
.await
.unwrap();
let keys = backend.keys().await.unwrap();
assert_eq!(keys.len(), 4);
assert!(keys.contains(&"string".to_string()));
assert!(keys.contains(&"number".to_string()));
assert!(keys.contains(&"boolean".to_string()));
assert!(keys.contains(&"object".to_string()));
assert_eq!(
backend.get("string").await.unwrap(),
Some(serde_json::json!("text"))
);
assert_eq!(
backend.get("number").await.unwrap(),
Some(serde_json::json!(42))
);
assert_eq!(
backend.get("boolean").await.unwrap(),
Some(serde_json::json!(true))
);
assert_eq!(
backend.get("object").await.unwrap(),
Some(serde_json::json!({"key": "value"}))
);
}
#[tokio::test]
async fn test_overwrite_value() {
let backend = MemoryBackend::new();
backend
.set("key", &serde_json::json!("value1"), None)
.await
.unwrap();
assert_eq!(
backend.get("key").await.unwrap(),
Some(serde_json::json!("value1"))
);
backend
.set("key", &serde_json::json!("value2"), None)
.await
.unwrap();
assert_eq!(
backend.get("key").await.unwrap(),
Some(serde_json::json!("value2"))
);
backend
.set("key", &serde_json::json!("value3"), Some(60))
.await
.unwrap();
assert_eq!(
backend.get("key").await.unwrap(),
Some(serde_json::json!("value3"))
);
}
#[tokio::test]
async fn test_cleanup_expired() {
let backend = MemoryBackend::new();
backend
.set("permanent", &serde_json::json!("forever"), None)
.await
.unwrap();
backend
.set("temp1", &serde_json::json!("expires"), Some(1))
.await
.unwrap();
backend
.set("temp2", &serde_json::json!("expires"), Some(1))
.await
.unwrap();
assert_eq!(backend.len(), 3);
tokio::time::sleep(tokio::time::Duration::from_millis(1100)).await;
backend.cleanup_expired();
assert_eq!(backend.len(), 1);
assert!(backend.exists("permanent").await.unwrap());
assert!(!backend.exists("temp1").await.unwrap());
assert!(!backend.exists("temp2").await.unwrap());
}
#[tokio::test]
async fn test_keys_filters_expired() {
let backend = MemoryBackend::new();
backend
.set("active1", &serde_json::json!("value1"), None)
.await
.unwrap();
backend
.set("active2", &serde_json::json!("value2"), None)
.await
.unwrap();
backend
.set("expired", &serde_json::json!("value"), Some(1))
.await
.unwrap();
tokio::time::sleep(Duration::from_millis(1100)).await;
let keys = backend.keys().await.unwrap();
assert_eq!(keys.len(), 2);
assert!(keys.contains(&"active1".to_string()));
assert!(keys.contains(&"active2".to_string()));
assert!(!keys.contains(&"expired".to_string()));
}
#[tokio::test]
async fn test_clear() {
let backend = MemoryBackend::new();
backend
.set("key1", &serde_json::json!("value1"), None)
.await
.unwrap();
backend
.set("key2", &serde_json::json!("value2"), None)
.await
.unwrap();
backend
.set("key3", &serde_json::json!("value3"), Some(60))
.await
.unwrap();
assert_eq!(backend.len(), 3);
assert!(!backend.is_empty());
backend.clear();
assert_eq!(backend.len(), 0);
assert!(backend.is_empty());
assert_eq!(backend.keys().await.unwrap().len(), 0);
}
#[tokio::test]
async fn test_concurrent_access() {
use std::sync::Arc;
let backend = Arc::new(MemoryBackend::new());
let mut handles = vec![];
for i in 0..10 {
let backend_clone = backend.clone();
let handle = tokio::spawn(async move {
let key = format!("key{}", i);
let value = serde_json::json!(i);
backend_clone.set(&key, &value, None).await.unwrap();
let retrieved = backend_clone.get(&key).await.unwrap();
assert_eq!(retrieved, Some(value));
});
handles.push(handle);
}
for handle in handles {
handle.await.unwrap();
}
let keys = backend.keys().await.unwrap();
assert_eq!(keys.len(), 10);
}
#[tokio::test]
async fn test_default_implementation() {
let backend = MemoryBackend::default();
assert!(backend.is_empty());
backend
.set("test", &serde_json::json!("value"), None)
.await
.unwrap();
assert_eq!(backend.len(), 1);
}
}