use std::collections::HashMap;
use std::sync::{Arc, RwLock};
use std::time::Duration as StdDuration;
use crate::sparkv::{Config as SparKVConfig, Error as SparKVError, HashMapSparKV};
use chrono::{Duration as ChronoDuration, Utc};
use serde_json::Value;
use crate::authz::metrics::MetricsCollector;
use super::config::{ConfigValidationError, DataStoreConfig};
use super::entry::DataEntry;
use super::error::DataError;
const RWLOCK_EXPECT_MESSAGE: &str = "DataStore storage lock should not be poisoned";
const INFINITE_TTL_SECS: i64 = 315_360_000;
const MAX_SAFE_DURATION_SECS: u64 = (i64::MAX / 1000) as u64;
pub(crate) struct DataStore {
storage: RwLock<HashMapSparKV<DataEntry>>,
config: DataStoreConfig,
metrics: Arc<MetricsCollector>,
}
impl DataStore {
pub(crate) fn new(
config: DataStoreConfig,
metrics: Arc<MetricsCollector>,
) -> Result<Self, ConfigValidationError> {
config.validate()?;
let sparkv_config = SparKVConfig {
max_items: config.max_entries,
max_item_size: config.max_entry_size,
max_ttl: config.max_ttl.map_or_else(
|| ChronoDuration::seconds(INFINITE_TTL_SECS),
std_duration_to_chrono_duration,
),
default_ttl: config.default_ttl.map_or_else(
|| ChronoDuration::seconds(INFINITE_TTL_SECS),
std_duration_to_chrono_duration,
),
auto_clear_expired: true,
earliest_expiration_eviction: false,
};
let size_calculator: Option<fn(&DataEntry) -> usize> =
Some(|entry| serde_json::to_string(entry).map_or(0, |s| s.len()));
Ok(Self {
storage: RwLock::new(HashMapSparKV::with_config_and_sizer(
sparkv_config,
size_calculator,
)),
config,
metrics,
})
}
pub(crate) fn push(
&self,
key: &str,
value: Value,
ttl: Option<StdDuration>,
) -> Result<(), DataError> {
if key.is_empty() {
let err = DataError::InvalidKey;
self.metrics.record_error(&err);
return Err(err);
}
if let Some(explicit_ttl) = ttl
&& let Some(max_ttl) = self.config.max_ttl
&& explicit_ttl > max_ttl
{
let err = DataError::TTLExceeded {
requested: explicit_ttl,
max: max_ttl,
};
self.metrics.record_error(&err);
return Err(err);
}
let effective_ttl_chrono =
get_effective_ttl(ttl, self.config.default_ttl, self.config.max_ttl);
let effective_ttl_std = effective_ttl_chrono.to_std().unwrap_or({
StdDuration::ZERO
});
let entry = DataEntry::new(key.to_string(), value, Some(effective_ttl_std));
let entry_size = serde_json::to_string(&entry)
.map_err(|e| {
let err = DataError::from(e);
self.metrics.record_error(&err);
err
})?
.len();
if self.config.max_entry_size > 0 && entry_size > self.config.max_entry_size {
let err = DataError::ValueTooLarge {
size: entry_size,
max: self.config.max_entry_size,
};
self.metrics.record_error(&err);
return Err(err);
}
let chrono_ttl = effective_ttl_chrono;
let mut storage = self.storage.write().expect(RWLOCK_EXPECT_MESSAGE);
storage
.set_with_ttl(key, entry, chrono_ttl, &[])
.map_err(|e| {
let err = match e {
SparKVError::CapacityExceeded => DataError::StorageLimitExceeded {
max: self.config.max_entries,
},
SparKVError::ItemSizeExceeded => DataError::ValueTooLarge {
size: entry_size,
max: self.config.max_entry_size,
},
SparKVError::TTLTooLong => DataError::TTLExceeded {
requested: ttl.unwrap_or_default(),
max: self
.config
.max_ttl
.unwrap_or(StdDuration::from_secs(INFINITE_TTL_SECS as u64)),
},
};
self.metrics.record_error(&err);
err
})?;
self.metrics.record_data_push();
Ok(())
}
pub(crate) fn get(&self, key: &str) -> Option<Value> {
self.get_entry(key).map(|entry| entry.value)
}
pub(crate) fn get_entry(&self, key: &str) -> Option<DataEntry> {
if self.config.enable_metrics {
let mut storage = self.storage.write().expect(RWLOCK_EXPECT_MESSAGE);
let mut entry = storage.get(key)?.clone();
let now = chrono::Utc::now();
if let Some(expires_at) = entry.expires_at
&& now > expires_at
{
storage.pop(key);
return None;
}
entry.increment_access();
let remaining_ttl = if let Some(expires_at) = entry.expires_at {
expires_at
.signed_duration_since(now)
.to_std()
.ok()
.map_or_else(
|| {
std_duration_to_chrono_duration(StdDuration::ZERO)
},
std_duration_to_chrono_duration,
)
} else {
get_effective_ttl(None, self.config.default_ttl, self.config.max_ttl)
};
let _ = storage.set_with_ttl(key, entry.clone(), remaining_ttl, &[]);
self.metrics.record_data_get();
Some(entry)
} else {
let storage = self.storage.read().expect(RWLOCK_EXPECT_MESSAGE);
let entry = storage.get(key)?.clone();
if let Some(expires_at) = entry.expires_at
&& chrono::Utc::now() > expires_at
{
return None;
}
self.metrics.record_data_get();
Some(entry)
}
}
pub(crate) fn remove(&self, key: &str) -> bool {
let mut storage = self.storage.write().expect(RWLOCK_EXPECT_MESSAGE);
let removed = storage.pop(key).is_some();
if removed {
self.metrics.record_data_remove();
}
removed
}
pub(crate) fn clear(&self) {
let mut storage = self.storage.write().expect(RWLOCK_EXPECT_MESSAGE);
storage.clear();
}
pub(crate) fn count(&self) -> usize {
let storage = self.storage.read().expect(RWLOCK_EXPECT_MESSAGE);
let now = chrono::Utc::now();
storage
.iter()
.filter(|(_, entry)| !entry.is_expired(now))
.count()
}
pub(crate) fn get_all(&self) -> HashMap<String, Value> {
let storage = self.storage.read().expect(RWLOCK_EXPECT_MESSAGE);
if storage.is_empty() {
return HashMap::new();
}
let now = chrono::Utc::now();
storage
.iter()
.filter(|(_, entry)| !entry.is_expired(now))
.map(|(k, entry)| (k.clone(), entry.value.clone()))
.collect()
}
pub(crate) fn list_entries(&self) -> Vec<DataEntry> {
let storage = self.storage.read().expect(RWLOCK_EXPECT_MESSAGE);
if storage.is_empty() {
return Vec::new();
}
let now = Utc::now();
storage
.iter()
.filter(|(_, entry)| !entry.is_expired(now))
.map(|(_, entry)| entry.clone())
.collect()
}
pub(crate) fn config(&self) -> &DataStoreConfig {
&self.config
}
pub(crate) fn total_size(&self) -> usize {
let storage = self.storage.read().expect(RWLOCK_EXPECT_MESSAGE);
let now = chrono::Utc::now();
storage
.iter()
.filter(|(_, entry)| !entry.is_expired(now))
.map(|(_, entry)| serde_json::to_string(entry).map_or(0, |s| s.len()))
.sum()
}
}
pub(super) fn std_duration_to_chrono_duration(d: StdDuration) -> ChronoDuration {
let secs = d.as_secs();
let nanos = d.subsec_nanos();
let secs_capped = secs.min(MAX_SAFE_DURATION_SECS);
#[allow(clippy::cast_possible_wrap)]
let secs_i64 = secs_capped as i64;
ChronoDuration::seconds(secs_i64) + ChronoDuration::nanoseconds(i64::from(nanos))
}
fn get_effective_ttl(
ttl: Option<StdDuration>,
default_ttl: Option<StdDuration>,
max_ttl: Option<StdDuration>,
) -> ChronoDuration {
let requested_ttl = ttl.or(default_ttl);
let effective = requested_ttl.unwrap_or(StdDuration::from_secs(INFINITE_TTL_SECS as u64));
let capped = if let Some(max) = max_ttl {
effective.min(max)
} else {
effective
};
std_duration_to_chrono_duration(capped)
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[cfg(not(target_arch = "wasm32"))]
use std::thread;
use std::time::Duration as StdDuration;
use test_utils::assert_eq;
fn create_test_store() -> DataStore {
let metrics = Arc::new(MetricsCollector::new(0));
DataStore::new(DataStoreConfig::default(), metrics).expect("should create store")
}
#[test]
fn test_push_and_get() {
let store = create_test_store();
store
.push("key1", json!("value1"), None)
.expect("failed to push simple value");
assert_eq!(store.get("key1"), Some(json!("value1")));
let complex_value = json!({
"name": "test",
"count": 42,
"active": true
});
store
.push("key2", complex_value.clone(), None)
.expect("failed to push complex value");
assert_eq!(store.get("key2"), Some(complex_value));
}
#[test]
fn test_push_replace_existing_key() {
let store = create_test_store();
store
.push("key1", json!("value1"), None)
.expect("failed to push initial value");
assert_eq!(store.get("key1"), Some(json!("value1")));
store
.push("key1", json!("value2"), None)
.expect("failed to replace value");
assert_eq!(store.get("key1"), Some(json!("value2")));
}
#[test]
fn test_push_empty_key() {
let store = create_test_store();
let result = store.push("", json!("value"), None);
assert!(
matches!(result, Err(DataError::InvalidKey)),
"push with empty key should return DataError::InvalidKey"
);
assert!(matches!(result, Err(DataError::InvalidKey)));
}
#[test]
fn test_get_nonexistent_key() {
let store = create_test_store();
assert_eq!(store.get("nonexistent"), None);
}
#[test]
fn test_remove() {
let store = create_test_store();
store
.push("key1", json!("value1"), None)
.expect("failed to push value for remove test");
assert_eq!(store.get("key1"), Some(json!("value1")));
assert!(store.remove("key1"));
assert_eq!(store.get("key1"), None);
assert!(!store.remove("nonexistent"));
}
#[test]
fn test_clear() {
let store = create_test_store();
store
.push("key1", json!("value1"), None)
.expect("failed to push key1 for clear test");
store
.push("key2", json!("value2"), None)
.expect("failed to push key2 for clear test");
assert_eq!(store.count(), 2);
store.clear();
assert_eq!(store.count(), 0);
assert_eq!(store.get("key1"), None);
assert_eq!(store.get("key2"), None);
}
#[test]
fn test_count() {
let store = create_test_store();
assert_eq!(store.count(), 0);
store
.push("key1", json!("value1"), None)
.expect("failed to push key1 for count test");
assert_eq!(store.count(), 1);
store
.push("key2", json!("value2"), None)
.expect("failed to push key2 for count test");
assert_eq!(store.count(), 2);
store.remove("key1");
assert_eq!(store.count(), 1);
}
#[test]
fn test_get_all() {
let store = create_test_store();
store
.push("key1", json!("value1"), None)
.expect("failed to push key1 for get_all test");
store
.push("key2", json!("value2"), None)
.expect("failed to push key2 for get_all test");
let all = store.get_all();
assert_eq!(all.len(), 2);
assert_eq!(all.get("key1"), Some(&json!("value1")));
assert_eq!(all.get("key2"), Some(&json!("value2")));
}
#[test]
#[cfg(not(target_arch = "wasm32"))]
fn test_ttl_expiration() {
let store = create_test_store();
store
.push("key1", json!("value1"), Some(StdDuration::from_millis(100)))
.expect("failed to push value with TTL");
assert_eq!(store.get("key1"), Some(json!("value1")));
thread::sleep(StdDuration::from_millis(200));
assert_eq!(store.get("key1"), None);
}
#[test]
fn test_max_entries() {
let config = DataStoreConfig {
max_entries: 2,
..Default::default()
};
let store = DataStore::new(config, Arc::new(MetricsCollector::new(0)))
.expect("should create store");
store
.push("key1", json!("value1"), None)
.expect("failed to push key1 for max_entries test");
store
.push("key2", json!("value2"), None)
.expect("failed to push key2 for max_entries test");
let result = store.push("key3", json!("value3"), None);
assert!(
matches!(result, Err(DataError::StorageLimitExceeded { max: 2 })),
"push with max_entries=2 should fail with StorageLimitExceeded when adding third entry"
);
}
#[test]
fn test_max_entry_size() {
let config = DataStoreConfig {
max_entry_size: 200,
..Default::default()
};
let store = DataStore::new(config, Arc::new(MetricsCollector::new(0)))
.expect("should create store");
store
.push("key1", json!("x"), None)
.expect("failed to push small value for max_entry_size test");
let large_value = json!(
"this is a very long string that exceeds the limit and it needs to be even longer to exceed 200 bytes including metadata"
);
let result = store.push("key2", large_value, None);
assert!(
matches!(result, Err(DataError::ValueTooLarge { .. })),
"push with value exceeding max_entry_size should fail with ValueTooLarge"
);
}
#[test]
fn test_max_ttl() {
let config = DataStoreConfig {
max_ttl: Some(StdDuration::from_secs(60)),
..Default::default()
};
let store = DataStore::new(config, Arc::new(MetricsCollector::new(0)))
.expect("should create store");
store
.push("key1", json!("value1"), Some(StdDuration::from_secs(30)))
.expect("failed to push value with valid TTL");
let result = store.push("key2", json!("value2"), Some(StdDuration::from_secs(120)));
assert!(
matches!(result, Err(DataError::TTLExceeded { .. })),
"push with TTL exceeding max_ttl should fail with TTLExceeded"
);
}
#[test]
#[cfg(not(target_arch = "wasm32"))]
fn test_default_ttl() {
let config = DataStoreConfig {
default_ttl: Some(StdDuration::from_millis(100)),
..Default::default()
};
let store = DataStore::new(config, Arc::new(MetricsCollector::new(0)))
.expect("should create store");
store
.push("key1", json!("value1"), None)
.expect("failed to push value with default TTL");
assert_eq!(store.get("key1"), Some(json!("value1")));
thread::sleep(StdDuration::from_millis(200));
assert_eq!(store.get("key1"), None);
}
#[test]
fn test_various_json_types() {
let store = create_test_store();
store
.push("str", json!("test"), None)
.expect("failed to push string value");
assert_eq!(store.get("str"), Some(json!("test")));
store
.push("num", json!(42), None)
.expect("failed to push number value");
assert_eq!(store.get("num"), Some(json!(42)));
store
.push("bool", json!(true), None)
.expect("failed to push boolean value");
assert_eq!(store.get("bool"), Some(json!(true)));
store
.push("arr", json!([1, 2, 3]), None)
.expect("failed to push array value");
assert_eq!(store.get("arr"), Some(json!([1, 2, 3])));
let obj = json!({
"a": 1,
"b": "test",
"c": [1, 2, 3]
});
store
.push("obj", obj.clone(), None)
.expect("failed to push object value");
assert_eq!(store.get("obj"), Some(obj));
}
#[test]
#[cfg(not(target_arch = "wasm32"))]
fn test_thread_safety() {
let store = create_test_store();
let store = std::sync::Arc::new(store);
let mut handles = vec![];
for i in 0..10 {
let store_clone = store.clone();
let handle = thread::spawn(move || {
for j in 0..10 {
let key = format!("key_{i}_{j}");
store_clone
.push(&key, json!(format!("value_{}_{}", i, j)), None)
.expect("failed to push value in thread");
}
});
handles.push(handle);
}
for handle in handles {
handle.join().expect("thread panicked");
}
assert_eq!(store.count(), 100);
let mut read_handles = vec![];
for i in 0..10 {
let store_clone = store.clone();
let handle = thread::spawn(move || {
for j in 0..10 {
let key = format!("key_{i}_{j}");
let expected = json!(format!("value_{}_{}", i, j));
assert_eq!(store_clone.get(&key), Some(expected));
}
});
read_handles.push(handle);
}
for handle in read_handles {
handle.join().expect("read thread panicked");
}
}
#[test]
#[cfg(not(target_arch = "wasm32"))]
fn test_concurrent_remove() {
let store = create_test_store();
let store = std::sync::Arc::new(store);
for i in 0..20 {
let key = format!("key_{i}");
store
.push(&key, json!(i), None)
.expect("failed to populate store for concurrent remove test");
}
let mut handles = vec![];
for i in 0..10 {
let store_clone = store.clone();
let handle = thread::spawn(move || {
store_clone.remove(&format!("key_{i}"));
});
handles.push(handle);
}
for handle in handles {
handle.join().expect("remove thread panicked");
}
assert!(store.count() < 20);
}
#[test]
fn test_get_entry_with_metadata() {
let store = create_test_store();
store
.push("key1", json!("value1"), Some(StdDuration::from_secs(60)))
.expect("failed to push value");
let entry = store.get_entry("key1").expect("entry should exist");
assert_eq!(entry.key, "key1");
assert_eq!(entry.value, json!("value1"));
assert_eq!(entry.data_type, crate::CedarType::String);
assert_eq!(entry.access_count, 1); assert!(entry.expires_at.is_some());
}
#[test]
fn test_metrics_tracking() {
let config = DataStoreConfig {
enable_metrics: true,
..Default::default()
};
let store = DataStore::new(config, Arc::new(MetricsCollector::new(0)))
.expect("should create store");
store
.push("key1", json!("value1"), None)
.expect("failed to push value");
let entry1 = store.get_entry("key1").expect("entry should exist");
assert_eq!(entry1.access_count, 1);
let entry2 = store.get_entry("key1").expect("entry should exist");
assert_eq!(entry2.access_count, 2);
let entry3 = store.get_entry("key1").expect("entry should exist");
assert_eq!(entry3.access_count, 3);
}
#[test]
fn test_metrics_disabled() {
let config = DataStoreConfig {
enable_metrics: false,
..Default::default()
};
let store = DataStore::new(config, Arc::new(MetricsCollector::new(0)))
.expect("should create store");
store
.push("key1", json!("value1"), None)
.expect("failed to push value");
let entry1 = store.get_entry("key1").expect("entry should exist");
assert_eq!(entry1.access_count, 0);
let entry2 = store.get_entry("key1").expect("entry should exist");
assert_eq!(entry2.access_count, 0); }
#[test]
fn test_cedar_type_inference() {
use crate::CedarType;
let store = create_test_store();
store
.push("string", json!("test"), None)
.expect("failed to push string");
store
.push("number", json!(42), None)
.expect("failed to push number");
store
.push("bool", json!(true), None)
.expect("failed to push bool");
store
.push("array", json!([1, 2, 3]), None)
.expect("failed to push array");
store
.push("object", json!({"key": "value"}), None)
.expect("failed to push object");
store
.push("entity", json!({"type": "User", "id": "123"}), None)
.expect("failed to push entity");
assert_eq!(
store.get_entry("string").unwrap().data_type,
CedarType::String
);
assert_eq!(
store.get_entry("number").unwrap().data_type,
CedarType::Long
);
assert_eq!(store.get_entry("bool").unwrap().data_type, CedarType::Bool);
assert_eq!(store.get_entry("array").unwrap().data_type, CedarType::Set);
assert_eq!(
store.get_entry("object").unwrap().data_type,
CedarType::Record
);
assert_eq!(
store.get_entry("entity").unwrap().data_type,
CedarType::Entity
);
}
#[test]
fn test_config_validation() {
let valid_config = DataStoreConfig {
default_ttl: Some(StdDuration::from_secs(300)),
max_ttl: Some(StdDuration::from_secs(3600)),
..Default::default()
};
assert!(
DataStore::new(valid_config, Arc::new(MetricsCollector::new(0))).is_ok(),
"expected DataStore::new() to succeed with valid DataStoreConfig"
);
let invalid_config = DataStoreConfig {
default_ttl: Some(StdDuration::from_secs(7200)),
max_ttl: Some(StdDuration::from_secs(3600)),
..Default::default()
};
assert!(
matches!(
DataStore::new(invalid_config, Arc::new(MetricsCollector::new(0))),
Err(ConfigValidationError::DefaultTtlExceedsMax { .. })
),
"expected DataStore::new() to return ConfigValidationError when default_ttl exceeds max_ttl"
);
}
#[test]
fn test_list_entries() {
let store = create_test_store();
store
.push("alpha", json!("a"), None)
.expect("push should succeed");
store
.push("beta", json!("b"), None)
.expect("push should succeed");
let entries = store.list_entries();
assert_eq!(entries.len(), 2);
let keys: Vec<&str> = entries.iter().map(|e| e.key.as_str()).collect();
assert!(keys.contains(&"alpha"));
assert!(keys.contains(&"beta"));
}
#[test]
fn test_config_accessor() {
let config = DataStoreConfig {
max_entries: 100,
max_entry_size: 512,
enable_metrics: true,
..Default::default()
};
let store = DataStore::new(config, Arc::new(MetricsCollector::new(0)))
.expect("should create store");
let retrieved_config = store.config();
assert_eq!(retrieved_config.max_entries, 100);
assert_eq!(retrieved_config.max_entry_size, 512);
assert!(retrieved_config.enable_metrics);
}
#[test]
fn test_get_all_returns_all_values() {
let store = create_test_store();
store
.push("user_role", json!("admin"), None)
.expect("push should succeed");
store
.push(
"feature_flags",
json!({"dark_mode": true, "beta": false}),
None,
)
.expect("push should succeed");
store
.push("rate_limit", json!(100), None)
.expect("push should succeed");
let all_data = store.get_all();
assert_eq!(all_data.len(), 3);
assert_eq!(all_data.get("user_role"), Some(&json!("admin")));
assert_eq!(
all_data.get("feature_flags"),
Some(&json!({"dark_mode": true, "beta": false}))
);
assert_eq!(all_data.get("rate_limit"), Some(&json!(100)));
}
#[test]
fn test_get_all_empty_store() {
let store = create_test_store();
let all_data = store.get_all();
assert!(all_data.is_empty());
}
#[test]
fn test_get_all_returns_values_not_metadata() {
let store = create_test_store();
store
.push("key", json!({"nested": {"value": 42}}), None)
.expect("push should succeed");
let all_data = store.get_all();
let value = all_data.get("key").expect("key should exist");
assert_eq!(value, &json!({"nested": {"value": 42}}));
}
#[test]
fn test_get_all_suitable_for_context_injection() {
let store = create_test_store();
store
.push("device_type", json!("mobile"), None)
.expect("push should succeed");
store
.push("geo", json!({"country": "US", "region": "CA"}), None)
.expect("push should succeed");
store
.push("permissions", json!(["read", "write"]), None)
.expect("push should succeed");
let all_data = store.get_all();
let data_value: Value = Value::Object(all_data.into_iter().collect());
assert!(data_value.is_object());
let obj = data_value.as_object().unwrap();
assert_eq!(obj.get("device_type"), Some(&json!("mobile")));
assert_eq!(
obj.get("geo"),
Some(&json!({"country": "US", "region": "CA"}))
);
assert_eq!(obj.get("permissions"), Some(&json!(["read", "write"])));
}
#[test]
fn test_total_size_calculation() {
let store = create_test_store();
assert_eq!(store.total_size(), 0);
store
.push("key1", json!("short"), None)
.expect("push should succeed");
let size_after_one = store.total_size();
assert!(
size_after_one > 0,
"size should be positive after adding entry"
);
store
.push("key2", json!({"nested": {"data": "value"}}), None)
.expect("push should succeed");
let size_after_two = store.total_size();
assert!(
size_after_two > size_after_one,
"size should increase after adding more entries"
);
store.remove("key1");
let size_after_remove = store.total_size();
assert!(
size_after_remove < size_after_two,
"size should decrease after removing entry"
);
}
#[test]
fn test_memory_alert_threshold_validation() {
let config = DataStoreConfig {
memory_alert_threshold: 80.0,
..Default::default()
};
assert!(config.validate().is_ok());
let config_zero = DataStoreConfig {
memory_alert_threshold: 0.0,
..Default::default()
};
assert!(config_zero.validate().is_ok());
let config_hundred = DataStoreConfig {
memory_alert_threshold: 100.0,
..Default::default()
};
assert!(config_hundred.validate().is_ok());
let config_negative = DataStoreConfig {
memory_alert_threshold: -1.0,
..Default::default()
};
assert!(
config_negative.validate().is_err(),
"memory_alert_threshold = -1.0 should fail validation"
);
let config_over = DataStoreConfig {
memory_alert_threshold: 101.0,
..Default::default()
};
assert!(
config_over.validate().is_err(),
"memory_alert_threshold = 101.0 should fail validation"
);
}
#[test]
#[cfg(not(target_arch = "wasm32"))]
fn test_stress_concurrent_read_write() {
use std::sync::Arc;
use std::thread;
let store = Arc::new(create_test_store());
let mut handles = vec![];
for writer_id in 0..5 {
let store_clone = store.clone();
let handle = thread::spawn(move || {
for i in 0..100 {
let key = format!("writer_{writer_id}_key_{i}");
store_clone
.push(&key, json!({"writer": writer_id, "value": i}), None)
.expect("concurrent push should succeed");
}
});
handles.push(handle);
}
for reader_id in 0..10 {
let store_clone = store.clone();
let handle = thread::spawn(move || {
for i in 0..200 {
let key = format!("writer_{}_key_{}", reader_id % 5, i % 100);
let _ = store_clone.get(&key);
let _ = store_clone.count();
}
});
handles.push(handle);
}
for handle in handles {
handle.join().expect("thread should not panic");
}
let count = store.count();
assert_eq!(
count, 500,
"should have 5 writers * 100 entries = 500 entries"
);
}
#[test]
#[cfg(not(target_arch = "wasm32"))]
fn test_stress_concurrent_mixed_operations() {
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::thread;
let store = Arc::new(create_test_store());
let successful_pushes = Arc::new(AtomicUsize::new(0));
let successful_removes = Arc::new(AtomicUsize::new(0));
for i in 0..50 {
store
.push(&format!("pre_{i}"), json!(i), None)
.expect("pre-populate should succeed");
}
let mut handles = vec![];
for writer_id in 0..3 {
let store_clone = store.clone();
let pushes = successful_pushes.clone();
let handle = thread::spawn(move || {
for i in 0..50 {
let key = format!("writer_{writer_id}_item_{i}");
if store_clone.push(&key, json!({"id": i}), None).is_ok() {
pushes.fetch_add(1, Ordering::SeqCst);
}
}
});
handles.push(handle);
}
for remover_id in 0..2 {
let store_clone = store.clone();
let removes = successful_removes.clone();
let handle = thread::spawn(move || {
for i in 0..25 {
let key = format!("pre_{}", remover_id * 25 + i);
if store_clone.remove(&key) {
removes.fetch_add(1, Ordering::SeqCst);
}
}
});
handles.push(handle);
}
for _ in 0..5 {
let store_clone = store.clone();
let handle = thread::spawn(move || {
for i in 0..100 {
let _ = store_clone.get(&format!("pre_{}", i % 50));
let _ = store_clone.get(&format!("writer_0_item_{}", i % 50));
let _ = store_clone.list_entries();
}
});
handles.push(handle);
}
for handle in handles {
handle.join().expect("thread should not panic");
}
let final_pushes = successful_pushes.load(Ordering::SeqCst);
let final_removes = successful_removes.load(Ordering::SeqCst);
let final_count = store.count();
assert!(final_pushes > 0, "should have completed some pushes");
assert!(final_removes > 0, "should have completed some removes");
assert!(final_count > 0, "store should not be empty");
}
#[test]
#[cfg(not(target_arch = "wasm32"))]
fn test_stress_rapid_clear_while_writing() {
use std::sync::Arc;
use std::thread;
let store = Arc::new(create_test_store());
let mut handles = vec![];
for writer_id in 0..3 {
let store_clone = store.clone();
let handle = thread::spawn(move || {
for i in 0..200 {
let key = format!("key_{writer_id}_{i}");
let _ = store_clone.push(&key, json!(i), None);
thread::yield_now();
}
});
handles.push(handle);
}
let store_clone = store.clone();
let clear_handle = thread::spawn(move || {
for _ in 0..10 {
thread::sleep(StdDuration::from_micros(100));
store_clone.clear();
}
});
handles.push(clear_handle);
for handle in handles {
handle.join().expect("thread should not panic");
}
}
#[test]
#[cfg(not(target_arch = "wasm32"))]
fn test_concurrent_get_all_for_context() {
use std::sync::Arc;
use std::thread;
let store = Arc::new(create_test_store());
for i in 0..10 {
store
.push(&format!("data_{i}"), json!({"index": i}), None)
.expect("pre-populate should succeed");
}
let mut handles = vec![];
for _ in 0..10 {
let store_clone = store.clone();
let handle = thread::spawn(move || {
for _ in 0..100 {
let data = store_clone.get_all();
assert!(
!data.is_empty() || store_clone.count() == 0,
"get_all should return data or store is empty"
);
}
});
handles.push(handle);
}
let store_clone = store.clone();
let write_handle = thread::spawn(move || {
for i in 0..50 {
let key = format!("new_data_{i}");
let _ = store_clone.push(&key, json!({"new_index": i}), None);
}
});
handles.push(write_handle);
for handle in handles {
handle.join().expect("thread should not panic");
}
}
}