652 lines
23 KiB
Rust
652 lines
23 KiB
Rust
use redb::{Database, ReadableDatabase, TableDefinition};
|
|
use serde::{Serialize, de::DeserializeOwned};
|
|
use std::sync::{Arc, RwLock};
|
|
|
|
pub const STORE_TABLE: TableDefinition<&str, &[u8]> = TableDefinition::new("store");
|
|
|
|
/// Internal write request dispatched to the single database writer actor.
|
|
enum DbOp {
|
|
Insert(Vec<u8>),
|
|
Delete,
|
|
Batch {
|
|
inserts: Vec<(String, Vec<u8>)>,
|
|
deletes: Vec<String>,
|
|
},
|
|
}
|
|
|
|
/// Internal write request dispatched to the single database writer actor.
|
|
struct DbWriteTask {
|
|
key: String,
|
|
op: DbOp,
|
|
flushed_notifier: Arc<tokio::sync::Notify>,
|
|
oneshot_tx: Option<tokio::sync::oneshot::Sender<()>>,
|
|
}
|
|
|
|
/// Shared centralized write queue actor that handles all database writes serially with micro-batching.
|
|
#[derive(Clone)]
|
|
pub struct DbWriteQueue {
|
|
tx: tokio::sync::mpsc::Sender<DbWriteTask>,
|
|
}
|
|
|
|
static QUEUE_REGISTRY: std::sync::Mutex<Option<(Arc<Database>, DbWriteQueue)>> =
|
|
std::sync::Mutex::new(None);
|
|
|
|
fn get_or_create_queue(db: Arc<Database>) -> DbWriteQueue {
|
|
let mut reg = QUEUE_REGISTRY.lock().unwrap_or_else(|e| e.into_inner());
|
|
if let Some((ref existing_db, ref queue)) = *reg
|
|
&& Arc::ptr_eq(existing_db, &db)
|
|
&& !queue.tx.is_closed()
|
|
{
|
|
return queue.clone();
|
|
}
|
|
let new_queue = DbWriteQueue::new(db.clone());
|
|
*reg = Some((db, new_queue.clone()));
|
|
new_queue
|
|
}
|
|
|
|
impl DbWriteQueue {
|
|
pub fn new(db: Arc<Database>) -> Self {
|
|
let (tx, mut rx) = tokio::sync::mpsc::channel::<DbWriteTask>(1024);
|
|
|
|
tokio::spawn(async move {
|
|
while let Some(first_task) = rx.recv().await {
|
|
let mut batch = Vec::with_capacity(100);
|
|
batch.push(first_task);
|
|
|
|
// Gold Standard Micro-batching: Drain up to 100 accumulated tasks from queue without blocking
|
|
while batch.len() < 100 {
|
|
match rx.try_recv() {
|
|
Ok(task) => batch.push(task),
|
|
Err(_) => break,
|
|
}
|
|
}
|
|
|
|
let db_inner = db.clone();
|
|
let _ = tokio::task::spawn_blocking(move || {
|
|
match db_inner.begin_write() {
|
|
Ok(write_txn) => {
|
|
if let Ok(mut table) = write_txn.open_table(STORE_TABLE) {
|
|
for task in &batch {
|
|
match &task.op {
|
|
DbOp::Insert(data) => {
|
|
if let Err(e) =
|
|
table.insert(task.key.as_str(), data.as_slice())
|
|
{
|
|
tracing::error!(
|
|
"Failed to insert key '{}' into redb: {}",
|
|
task.key,
|
|
e
|
|
);
|
|
}
|
|
}
|
|
DbOp::Delete => {
|
|
if let Err(e) = table.remove(task.key.as_str()) {
|
|
tracing::error!(
|
|
"Failed to delete key '{}' from redb: {}",
|
|
task.key,
|
|
e
|
|
);
|
|
}
|
|
}
|
|
DbOp::Batch { inserts, deletes } => {
|
|
for del_k in deletes {
|
|
if let Err(e) = table.remove(del_k.as_str()) {
|
|
tracing::error!(
|
|
"Failed to delete batch key '{}' from redb: {}",
|
|
del_k,
|
|
e
|
|
);
|
|
}
|
|
}
|
|
for (ins_k, ins_bytes) in inserts {
|
|
if let Err(e) =
|
|
table.insert(ins_k.as_str(), ins_bytes.as_slice())
|
|
{
|
|
tracing::error!(
|
|
"Failed to insert batch key '{}' into redb: {}",
|
|
ins_k,
|
|
e
|
|
);
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
if let Err(e) = write_txn.commit() {
|
|
tracing::error!("Failed to commit batch to redb: {}", e);
|
|
}
|
|
}
|
|
Err(e) => {
|
|
tracing::error!(
|
|
"Failed to begin write transaction on redb writer actor: {}",
|
|
e
|
|
);
|
|
}
|
|
}
|
|
|
|
// Event-driven notification to all waiting listeners for this micro-batch
|
|
for task in batch {
|
|
if let Some(oneshot) = task.oneshot_tx {
|
|
let _ = oneshot.send(());
|
|
}
|
|
task.flushed_notifier.notify_waiters();
|
|
}
|
|
})
|
|
.await;
|
|
}
|
|
});
|
|
|
|
Self { tx }
|
|
}
|
|
|
|
pub fn push(
|
|
&self,
|
|
key: String,
|
|
data: Vec<u8>,
|
|
flushed_notifier: Arc<tokio::sync::Notify>,
|
|
) -> Option<tokio::sync::oneshot::Receiver<()>> {
|
|
self.push_op(key, DbOp::Insert(data), flushed_notifier)
|
|
}
|
|
|
|
pub fn push_delete(
|
|
&self,
|
|
key: String,
|
|
flushed_notifier: Arc<tokio::sync::Notify>,
|
|
) -> Option<tokio::sync::oneshot::Receiver<()>> {
|
|
self.push_op(key, DbOp::Delete, flushed_notifier)
|
|
}
|
|
|
|
pub fn push_batch(
|
|
&self,
|
|
key: String,
|
|
inserts: Vec<(String, Vec<u8>)>,
|
|
deletes: Vec<String>,
|
|
flushed_notifier: Arc<tokio::sync::Notify>,
|
|
) -> Option<tokio::sync::oneshot::Receiver<()>> {
|
|
self.push_op(key, DbOp::Batch { inserts, deletes }, flushed_notifier)
|
|
}
|
|
|
|
fn push_op(
|
|
&self,
|
|
key: String,
|
|
op: DbOp,
|
|
flushed_notifier: Arc<tokio::sync::Notify>,
|
|
) -> Option<tokio::sync::oneshot::Receiver<()>> {
|
|
let (oneshot_tx, oneshot_rx) = tokio::sync::oneshot::channel();
|
|
let task = DbWriteTask {
|
|
key,
|
|
op,
|
|
flushed_notifier,
|
|
oneshot_tx: Some(oneshot_tx),
|
|
};
|
|
if let Err(e) = self.tx.try_send(task) {
|
|
match e {
|
|
tokio::sync::mpsc::error::TrySendError::Full(task) => {
|
|
let tx = self.tx.clone();
|
|
let key = task.key.clone();
|
|
tokio::spawn(async move {
|
|
if let Err(err) = tx.send(task).await {
|
|
tracing::error!(
|
|
"DbWriteQueue fallback send failed for key '{}': {}",
|
|
key,
|
|
err
|
|
);
|
|
}
|
|
});
|
|
None
|
|
}
|
|
tokio::sync::mpsc::error::TrySendError::Closed(task) => {
|
|
tracing::error!(
|
|
"DbWriteQueue channel closed; unable to persist key '{}'",
|
|
task.key
|
|
);
|
|
None
|
|
}
|
|
}
|
|
} else {
|
|
Some(oneshot_rx)
|
|
}
|
|
}
|
|
|
|
pub async fn push_async(
|
|
&self,
|
|
key: String,
|
|
data: Vec<u8>,
|
|
flushed_notifier: Arc<tokio::sync::Notify>,
|
|
) -> Option<tokio::sync::oneshot::Receiver<()>> {
|
|
self.push_op_async(key, DbOp::Insert(data), flushed_notifier)
|
|
.await
|
|
}
|
|
|
|
pub async fn push_delete_async(
|
|
&self,
|
|
key: String,
|
|
flushed_notifier: Arc<tokio::sync::Notify>,
|
|
) -> Option<tokio::sync::oneshot::Receiver<()>> {
|
|
self.push_op_async(key, DbOp::Delete, flushed_notifier)
|
|
.await
|
|
}
|
|
|
|
pub async fn push_batch_async(
|
|
&self,
|
|
key: String,
|
|
inserts: Vec<(String, Vec<u8>)>,
|
|
deletes: Vec<String>,
|
|
flushed_notifier: Arc<tokio::sync::Notify>,
|
|
) -> Option<tokio::sync::oneshot::Receiver<()>> {
|
|
self.push_op_async(key, DbOp::Batch { inserts, deletes }, flushed_notifier)
|
|
.await
|
|
}
|
|
|
|
async fn push_op_async(
|
|
&self,
|
|
key: String,
|
|
op: DbOp,
|
|
flushed_notifier: Arc<tokio::sync::Notify>,
|
|
) -> Option<tokio::sync::oneshot::Receiver<()>> {
|
|
let (oneshot_tx, oneshot_rx) = tokio::sync::oneshot::channel();
|
|
let task = DbWriteTask {
|
|
key,
|
|
op,
|
|
flushed_notifier,
|
|
oneshot_tx: Some(oneshot_tx),
|
|
};
|
|
if let Err(e) = self.tx.send(task).await {
|
|
tracing::error!(
|
|
"DbWriteQueue channel closed; unable to persist key '{}'",
|
|
e.0.key
|
|
);
|
|
None
|
|
} else {
|
|
Some(oneshot_rx)
|
|
}
|
|
}
|
|
}
|
|
|
|
pub struct Store<T> {
|
|
pub cache: Arc<RwLock<T>>,
|
|
pub flushed: Arc<tokio::sync::Notify>,
|
|
key: String,
|
|
queue: DbWriteQueue,
|
|
is_corrupted: bool,
|
|
known_granular_keys: Arc<RwLock<std::collections::HashSet<String>>>,
|
|
}
|
|
|
|
impl<T: DeserializeOwned + Default + Serialize + Send + Sync + 'static> Store<T> {
|
|
pub fn new(key: &str, db: Arc<Database>) -> Self {
|
|
let (initial_data, is_corrupted, known_keys) = Self::load_from_db(key, &db);
|
|
let cache = Arc::new(RwLock::new(initial_data));
|
|
let flushed = Arc::new(tokio::sync::Notify::new());
|
|
let queue = get_or_create_queue(db);
|
|
let known_granular_keys = Arc::new(RwLock::new(known_keys));
|
|
|
|
Self {
|
|
cache,
|
|
flushed,
|
|
key: key.to_string(),
|
|
queue,
|
|
is_corrupted,
|
|
known_granular_keys,
|
|
}
|
|
}
|
|
|
|
fn load_from_db(key: &str, db: &Database) -> (T, bool, std::collections::HashSet<String>) {
|
|
let Ok(read_txn) = db.begin_read() else {
|
|
tracing::error!("Failed to begin read transaction for key '{}'", key);
|
|
return (T::default(), false, std::collections::HashSet::new());
|
|
};
|
|
let Ok(table) = read_txn.open_table(STORE_TABLE) else {
|
|
return (T::default(), false, std::collections::HashSet::new());
|
|
};
|
|
|
|
// 1. Check monolithic key first as the authoritative snapshot
|
|
match table.get(key) {
|
|
Ok(Some(value)) => match serde_json::from_slice::<T>(value.value()) {
|
|
Ok(parsed) => {
|
|
let mut known = std::collections::HashSet::new();
|
|
if let Ok(val) = serde_json::to_value(&parsed) {
|
|
for k in Self::extract_granular_keys(key, &val) {
|
|
known.insert(k);
|
|
}
|
|
}
|
|
return (parsed, false, known);
|
|
}
|
|
Err(e) => {
|
|
tracing::error!(
|
|
"CRITICAL: Corrupted data for key '{}' in database: {}. Quarantine mode active: state initialized to empty default without overwriting DB key.",
|
|
key,
|
|
e
|
|
);
|
|
return (T::default(), true, std::collections::HashSet::new());
|
|
}
|
|
},
|
|
Ok(None) => {}
|
|
Err(e) => {
|
|
tracing::error!("Failed to get key '{}' from store table: {}", key, e);
|
|
}
|
|
}
|
|
|
|
// 2. Granular prefix keys fallback: format!("{}:", key)
|
|
let prefix = format!("{}:", key);
|
|
let mut items_array = Vec::new();
|
|
let mut items_map = serde_json::Map::new();
|
|
let mut found_granular = false;
|
|
let mut known = std::collections::HashSet::new();
|
|
|
|
if let Ok(range) = table.range(prefix.as_str()..) {
|
|
for (k, v) in range.flatten() {
|
|
let k_str = k.value();
|
|
if !k_str.starts_with(&prefix) {
|
|
break;
|
|
}
|
|
found_granular = true;
|
|
known.insert(k_str.to_string());
|
|
if let Ok(val) = serde_json::from_slice::<serde_json::Value>(v.value()) {
|
|
let sub_key = &k_str[prefix.len()..];
|
|
items_array.push(val.clone());
|
|
items_map.insert(sub_key.to_string(), val);
|
|
}
|
|
}
|
|
}
|
|
|
|
if found_granular {
|
|
if let Ok(parsed) = serde_json::from_value::<T>(serde_json::Value::Array(items_array)) {
|
|
return (parsed, false, known);
|
|
}
|
|
if let Ok(parsed) = serde_json::from_value::<T>(serde_json::Value::Object(items_map)) {
|
|
return (parsed, false, known);
|
|
}
|
|
}
|
|
|
|
(T::default(), false, std::collections::HashSet::new())
|
|
}
|
|
|
|
fn extract_granular_keys(base_key: &str, val: &serde_json::Value) -> Vec<String> {
|
|
let mut keys = Vec::new();
|
|
match val {
|
|
serde_json::Value::Array(arr) => {
|
|
for (i, item) in arr.iter().enumerate() {
|
|
let sub_key = item
|
|
.get("id")
|
|
.or_else(|| item.get("name"))
|
|
.or_else(|| item.get("title"))
|
|
.and_then(|v| v.as_str())
|
|
.map(|s| s.to_string())
|
|
.unwrap_or_else(|| i.to_string());
|
|
keys.push(format!("{}:{}", base_key, sub_key));
|
|
}
|
|
}
|
|
serde_json::Value::Object(map) => {
|
|
for sub_key in map.keys() {
|
|
keys.push(format!("{}:{}", base_key, sub_key));
|
|
}
|
|
}
|
|
_ => {}
|
|
}
|
|
keys
|
|
}
|
|
|
|
fn extract_granular_entries(base_key: &str, val: &serde_json::Value) -> Vec<(String, Vec<u8>)> {
|
|
let mut granular = Vec::new();
|
|
match val {
|
|
serde_json::Value::Array(arr) => {
|
|
for (i, item) in arr.iter().enumerate() {
|
|
let sub_key = item
|
|
.get("id")
|
|
.or_else(|| item.get("name"))
|
|
.or_else(|| item.get("title"))
|
|
.and_then(|v| v.as_str())
|
|
.map(|s| s.to_string())
|
|
.unwrap_or_else(|| i.to_string());
|
|
if let Ok(item_bytes) = serde_json::to_vec(item) {
|
|
granular.push((format!("{}:{}", base_key, sub_key), item_bytes));
|
|
}
|
|
}
|
|
}
|
|
serde_json::Value::Object(map) => {
|
|
for (sub_key, item) in map {
|
|
if let Ok(item_bytes) = serde_json::to_vec(item) {
|
|
granular.push((format!("{}:{}", base_key, sub_key), item_bytes));
|
|
}
|
|
}
|
|
}
|
|
_ => {}
|
|
}
|
|
granular
|
|
}
|
|
|
|
pub fn read_with<F, R>(&self, f: F) -> R
|
|
where
|
|
F: FnOnce(&T) -> R,
|
|
{
|
|
let lock = self.cache.read().unwrap_or_else(|e| e.into_inner());
|
|
f(&lock)
|
|
}
|
|
|
|
pub fn modify<F: FnOnce(&mut T)>(&self, f: F)
|
|
where
|
|
T: Serialize + Clone,
|
|
{
|
|
if self.is_corrupted {
|
|
tracing::error!(
|
|
"CRITICAL: Refusing to persist changes for corrupted store key '{}' to prevent data loss.",
|
|
self.key
|
|
);
|
|
return;
|
|
}
|
|
|
|
// Fast mutation under critical lock section, only ONE clone taken, then immediately release the RwLock guard
|
|
let new_snapshot = {
|
|
let mut lock = self.cache.write().unwrap_or_else(|e| e.into_inner());
|
|
f(&mut lock);
|
|
(*lock).clone()
|
|
};
|
|
|
|
match self.prepare_batch(&new_snapshot) {
|
|
Ok((batch_inserts, removed_keys)) => {
|
|
self.queue.push_batch(
|
|
self.key.clone(),
|
|
batch_inserts,
|
|
removed_keys,
|
|
self.flushed.clone(),
|
|
);
|
|
}
|
|
Err(e) => tracing::error!(
|
|
"Failed to serialize memory store for key '{}': {}",
|
|
self.key,
|
|
e
|
|
),
|
|
}
|
|
}
|
|
|
|
#[allow(clippy::type_complexity)]
|
|
fn prepare_batch(
|
|
&self,
|
|
new_snapshot: &T,
|
|
) -> Result<(Vec<(String, Vec<u8>)>, Vec<String>), serde_json::Error>
|
|
where
|
|
T: Serialize,
|
|
{
|
|
let full_bytes = serde_json::to_vec(new_snapshot)?;
|
|
let granular_entries = match serde_json::from_slice::<serde_json::Value>(&full_bytes) {
|
|
Ok(val) => Self::extract_granular_entries(&self.key, &val),
|
|
Err(_) => Vec::new(),
|
|
};
|
|
|
|
let new_keys: std::collections::HashSet<String> =
|
|
granular_entries.iter().map(|(k, _)| k.clone()).collect();
|
|
let mut removed_keys = Vec::new();
|
|
{
|
|
let mut known = self
|
|
.known_granular_keys
|
|
.write()
|
|
.unwrap_or_else(|e| e.into_inner());
|
|
for old_k in known.iter() {
|
|
if !new_keys.contains(old_k) {
|
|
removed_keys.push(old_k.clone());
|
|
}
|
|
}
|
|
*known = new_keys;
|
|
}
|
|
|
|
let mut batch_inserts = Vec::with_capacity(granular_entries.len() + 1);
|
|
batch_inserts.extend(granular_entries);
|
|
batch_inserts.push((self.key.clone(), full_bytes));
|
|
|
|
Ok((batch_inserts, removed_keys))
|
|
}
|
|
|
|
pub async fn modify_async<F: FnOnce(&mut T)>(&self, f: F)
|
|
where
|
|
T: Serialize + Clone,
|
|
{
|
|
if self.is_corrupted {
|
|
tracing::error!(
|
|
"CRITICAL: Refusing to persist changes for corrupted store key '{}' to prevent data loss.",
|
|
self.key
|
|
);
|
|
return;
|
|
}
|
|
|
|
let new_snapshot = {
|
|
let mut lock = self.cache.write().unwrap_or_else(|e| e.into_inner());
|
|
f(&mut lock);
|
|
(*lock).clone()
|
|
};
|
|
|
|
match self.prepare_batch(&new_snapshot) {
|
|
Ok((batch_inserts, removed_keys)) => {
|
|
if let Some(rx) = self
|
|
.queue
|
|
.push_batch_async(
|
|
self.key.clone(),
|
|
batch_inserts,
|
|
removed_keys,
|
|
self.flushed.clone(),
|
|
)
|
|
.await
|
|
{
|
|
let _ = rx.await;
|
|
}
|
|
}
|
|
Err(e) => tracing::error!(
|
|
"Failed to serialize memory store for key '{}': {}",
|
|
self.key,
|
|
e
|
|
),
|
|
}
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
|
|
#[derive(Serialize, serde::Deserialize, Clone, Default, PartialEq, Debug)]
|
|
struct TestData {
|
|
name: String,
|
|
value: i32,
|
|
}
|
|
|
|
fn create_in_memory_test_db() -> Arc<Database> {
|
|
let db = Database::builder()
|
|
.create_with_backend(redb::backends::InMemoryBackend::new())
|
|
.unwrap();
|
|
let write_txn = db.begin_write().unwrap();
|
|
{
|
|
write_txn.open_table(STORE_TABLE).unwrap();
|
|
}
|
|
write_txn.commit().unwrap();
|
|
Arc::new(db)
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_store_read_write() {
|
|
let db = create_in_memory_test_db();
|
|
let store = Store::<TestData>::new("test_key", db.clone());
|
|
|
|
assert_eq!(store.read_with(|s| s.clone()), TestData::default());
|
|
|
|
store.modify(|data| {
|
|
data.name = "Hello".to_string();
|
|
data.value = 42;
|
|
});
|
|
|
|
// Event-driven wait for persistence completion
|
|
store.flushed.notified().await;
|
|
|
|
let store2 = Store::<TestData>::new("test_key", db.clone());
|
|
assert_eq!(
|
|
store2.read_with(|s| s.clone()),
|
|
TestData {
|
|
name: "Hello".to_string(),
|
|
value: 42
|
|
}
|
|
);
|
|
}
|
|
|
|
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
|
|
async fn test_store_concurrency() {
|
|
let db = create_in_memory_test_db();
|
|
let store = Arc::new(Store::<TestData>::new("concurrent_key", db.clone()));
|
|
|
|
let mut handles = vec![];
|
|
for _ in 0..50 {
|
|
let s = store.clone();
|
|
handles.push(tokio::spawn(async move {
|
|
s.modify(|data| {
|
|
data.value += 1;
|
|
});
|
|
}));
|
|
}
|
|
|
|
for h in handles {
|
|
h.await.unwrap();
|
|
}
|
|
|
|
// Event-driven wait for blocking writes to flush
|
|
store.flushed.notified().await;
|
|
|
|
assert_eq!(store.read_with(|s| s.value), 50);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_store_deletion_does_not_resurrect() {
|
|
let db = create_in_memory_test_db();
|
|
#[derive(Serialize, serde::Deserialize, Clone, Default, PartialEq, Debug)]
|
|
struct Item {
|
|
id: String,
|
|
name: String,
|
|
}
|
|
let store = Store::<Vec<Item>>::new("items", db.clone());
|
|
store.modify(|items| {
|
|
items.push(Item {
|
|
id: "item1".into(),
|
|
name: "First".into(),
|
|
});
|
|
items.push(Item {
|
|
id: "item2".into(),
|
|
name: "Second".into(),
|
|
});
|
|
});
|
|
store.flushed.notified().await;
|
|
|
|
// Verify both items loaded
|
|
let store_check = Store::<Vec<Item>>::new("items", db.clone());
|
|
assert_eq!(store_check.read_with(|items| items.len()), 2);
|
|
|
|
// Delete item1
|
|
store.modify(|items| {
|
|
items.retain(|i| i.id != "item1");
|
|
});
|
|
store.flushed.notified().await;
|
|
|
|
// Reload from DB into a brand new Store instance - item1 must NOT resurrect!
|
|
let store_reloaded = Store::<Vec<Item>>::new("items", db.clone());
|
|
let remaining = store_reloaded.read_with(|items| items.clone());
|
|
assert_eq!(remaining.len(), 1);
|
|
assert_eq!(remaining[0].id, "item2");
|
|
}
|
|
}
|