use std::collections::HashMap;
use std::fs;
use std::io::Write;
use std::path::{Path, PathBuf};
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::{Arc, Condvar, Mutex};
use std::thread::{self, ThreadId};
use std::time::{Duration, Instant};
use mongreldb_core::encryption::{AesCipher, Cipher};
use mongreldb_core::engine::LookupMetricsSnapshot;
use mongreldb_core::result_cache::{
PendingCacheOp, PersistableEntry, PersistentResultCacheWriter, WriterLimits,
};
use roaring::RoaringBitmap;
use tempfile::tempdir;
trait PersistentCacheIo: Send + Sync {
fn serialize(&self, entry: &PersistableEntry) -> Vec<u8>;
fn encrypt(&self, plaintext: &[u8]) -> Option<Vec<u8>>;
fn write(&self, path: &Path, bytes: &[u8]) -> std::io::Result<()>;
fn flush(&self) -> std::io::Result<()>;
fn sync(&self) -> std::io::Result<()>;
fn rename(&self, from: &Path, to: &Path) -> std::io::Result<()>;
fn remove(&self, path: &Path) -> std::io::Result<()>;
fn clear(&self) -> std::io::Result<usize>;
fn read(&self, path: &Path) -> std::io::Result<Vec<u8>>;
fn decrypt(&self, bytes: &[u8]) -> Option<Vec<u8>>;
fn deserialize(&self, bytes: &[u8]) -> Option<PersistableEntry>;
fn dir(&self) -> &Path;
fn exists(&self, path: &Path) -> bool;
}
const FILE_MAGIC: u32 = 0x5243_4348; const NONCE_LEN: usize = 12;
#[derive(Clone)]
struct PersistedHeader {
key: u64,
table_id: u64,
schema_id: u64,
run_generation: u64,
entry_generation: u64,
footprint: Vec<u32>,
columns: Vec<u16>,
rows: Vec<u8>,
}
fn encode_entry(header: &PersistedHeader, payload: &[u8]) -> Vec<u8> {
let mut out = Vec::with_capacity(payload.len() + 64);
out.extend_from_slice(&FILE_MAGIC.to_le_bytes());
out.extend_from_slice(&header.schema_id.to_le_bytes());
out.extend_from_slice(&header.run_generation.to_le_bytes());
out.extend_from_slice(&header.entry_generation.to_le_bytes());
out.extend_from_slice(&header.key.to_le_bytes());
out.extend_from_slice(&header.table_id.to_le_bytes());
out.extend_from_slice(&(header.footprint.len() as u32).to_le_bytes());
for v in &header.footprint {
out.extend_from_slice(&v.to_le_bytes());
}
out.extend_from_slice(&(header.columns.len() as u32).to_le_bytes());
for v in &header.columns {
out.extend_from_slice(&v.to_le_bytes());
}
out.extend_from_slice(&(header.rows.len() as u32).to_le_bytes());
out.extend_from_slice(&header.rows);
out.extend_from_slice(payload);
out
}
fn decode_header(mut bytes: &[u8]) -> Option<(PersistedHeader, &[u8])> {
if bytes.len() < 4 + 5 * 8 {
return None;
}
let magic = u32::from_le_bytes(bytes[..4].try_into().ok()?);
if magic != FILE_MAGIC {
return None;
}
bytes = &bytes[4..];
let schema_id = u64::from_le_bytes(bytes[..8].try_into().ok()?);
bytes = &bytes[8..];
let run_generation = u64::from_le_bytes(bytes[..8].try_into().ok()?);
bytes = &bytes[8..];
let entry_generation = u64::from_le_bytes(bytes[..8].try_into().ok()?);
bytes = &bytes[8..];
let key = u64::from_le_bytes(bytes[..8].try_into().ok()?);
bytes = &bytes[8..];
let table_id = u64::from_le_bytes(bytes[..8].try_into().ok()?);
bytes = &bytes[8..];
let footprint_len = u32::from_le_bytes(bytes[..4].try_into().ok()?) as usize;
bytes = &bytes[4..];
if bytes.len() < footprint_len * 4 {
return None;
}
let mut footprint = Vec::with_capacity(footprint_len);
for _ in 0..footprint_len {
footprint.push(u32::from_le_bytes(bytes[..4].try_into().ok()?));
bytes = &bytes[4..];
}
let columns_len = u32::from_le_bytes(bytes[..4].try_into().ok()?) as usize;
bytes = &bytes[4..];
if bytes.len() < columns_len * 2 {
return None;
}
let mut columns = Vec::with_capacity(columns_len);
for _ in 0..columns_len {
columns.push(u16::from_le_bytes(bytes[..2].try_into().ok()?));
bytes = &bytes[2..];
}
let rows_len = u32::from_le_bytes(bytes[..4].try_into().ok()?) as usize;
bytes = &bytes[4..];
if bytes.len() < rows_len {
return None;
}
let rows = bytes[..rows_len].to_vec();
bytes = &bytes[rows_len..];
Some((
PersistedHeader {
key,
table_id,
schema_id,
run_generation,
entry_generation,
footprint,
columns,
rows,
},
bytes,
))
}
fn make_entry(key: u64, schema_id: u64, run_generation: u64, rows: &[u8]) -> PersistableEntry {
PersistableEntry {
key,
table_id: 1,
schema_id,
run_generation,
entry_generation: 1,
footprint: Arc::new(RoaringBitmap::new()),
rows: Arc::from(rows.to_vec().into_boxed_slice()),
columns: Arc::from(Vec::<u16>::new().into_boxed_slice()),
bytes: rows.len(),
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum IoOp {
Serialize,
Encrypt,
Write,
Flush,
Sync,
Rename,
Remove,
Clear,
Read,
Decrypt,
Deserialize,
}
struct RecorderState {
log: Vec<(ThreadId, IoOp)>,
pending_io_error_after: Option<IoOp>,
}
struct RecordingIo {
state: Mutex<RecorderState>,
files: Mutex<HashMap<PathBuf, Vec<u8>>>,
}
impl RecordingIo {
fn new() -> Self {
Self {
state: Mutex::new(RecorderState {
log: Vec::new(),
pending_io_error_after: None,
}),
files: Mutex::new(HashMap::new()),
}
}
fn record(&self, op: IoOp) {
let mut s = self.state.lock().unwrap();
s.log.push((thread::current().id(), op));
if s.pending_io_error_after == Some(op) {
s.pending_io_error_after = None;
}
}
fn arm_io_error_after(&self, op: IoOp) {
self.state.lock().unwrap().pending_io_error_after = Some(op);
}
fn io_error_if_armed(&self, op: IoOp) -> Option<std::io::Error> {
let mut s = self.state.lock().unwrap();
if s.pending_io_error_after == Some(op) {
s.pending_io_error_after = None;
Some(std::io::Error::new(
std::io::ErrorKind::Other,
"test-injected io error",
))
} else {
None
}
}
fn log(&self) -> Vec<(ThreadId, IoOp)> {
self.state.lock().unwrap().log.clone()
}
fn write_file(&self, path: &Path, bytes: &[u8]) {
self.files
.lock()
.unwrap()
.insert(path.to_path_buf(), bytes.to_vec());
}
fn has_file(&self, path: &Path) -> bool {
self.files.lock().unwrap().contains_key(path)
}
fn read_file(&self, path: &Path) -> Option<Vec<u8>> {
self.files.lock().unwrap().get(path).cloned()
}
fn delete_file(&self, path: &Path) {
self.files.lock().unwrap().remove(path);
}
}
impl PersistentCacheIo for RecordingIo {
fn serialize(&self, entry: &PersistableEntry) -> Vec<u8> {
self.record(IoOp::Serialize);
let mut buf = Vec::with_capacity(64);
buf.extend_from_slice(&entry.schema_id.to_le_bytes());
buf.extend_from_slice(&entry.run_generation.to_le_bytes());
buf.extend_from_slice(&entry.entry_generation.to_le_bytes());
buf.extend_from_slice(&entry.table_id.to_le_bytes());
buf.extend_from_slice(&entry.rows.as_ref());
buf
}
fn encrypt(&self, _plaintext: &[u8]) -> Option<Vec<u8>> {
self.record(IoOp::Encrypt);
None
}
fn write(&self, path: &Path, bytes: &[u8]) -> std::io::Result<()> {
self.record(IoOp::Write);
if let Some(e) = self.io_error_if_armed(IoOp::Write) {
return Err(e);
}
self.write_file(path, bytes);
Ok(())
}
fn flush(&self) -> std::io::Result<()> {
self.record(IoOp::Flush);
Ok(())
}
fn sync(&self) -> std::io::Result<()> {
self.record(IoOp::Sync);
Ok(())
}
fn rename(&self, from: &Path, to: &Path) -> std::io::Result<()> {
self.record(IoOp::Rename);
if let Some(e) = self.io_error_if_armed(IoOp::Rename) {
return Err(e);
}
if let Some(b) = self.read_file(from) {
self.write_file(to, &b);
self.delete_file(from);
}
Ok(())
}
fn remove(&self, path: &Path) -> std::io::Result<()> {
self.record(IoOp::Remove);
self.delete_file(path);
Ok(())
}
fn clear(&self) -> std::io::Result<usize> {
self.record(IoOp::Clear);
let mut f = self.files.lock().unwrap();
let n = f.len();
f.clear();
Ok(n)
}
fn read(&self, path: &Path) -> std::io::Result<Vec<u8>> {
self.record(IoOp::Read);
self.read_file(path).ok_or_else(|| {
std::io::Error::new(std::io::ErrorKind::NotFound, "missing in recording io")
})
}
fn decrypt(&self, _bytes: &[u8]) -> Option<Vec<u8>> {
self.record(IoOp::Decrypt);
None
}
fn deserialize(&self, bytes: &[u8]) -> Option<PersistableEntry> {
self.record(IoOp::Deserialize);
if bytes.len() < 32 {
return None;
}
let schema_id = u64::from_le_bytes(bytes[0..8].try_into().ok()?);
let run_generation = u64::from_le_bytes(bytes[8..16].try_into().ok()?);
let entry_generation = u64::from_le_bytes(bytes[16..24].try_into().ok()?);
let table_id = u64::from_le_bytes(bytes[24..32].try_into().ok()?);
Some(PersistableEntry {
key: 0,
table_id,
schema_id,
run_generation,
entry_generation,
footprint: Arc::new(RoaringBitmap::new()),
rows: Arc::from(bytes[32..].to_vec().into_boxed_slice()),
columns: Arc::from(Vec::<u16>::new().into_boxed_slice()),
bytes: bytes.len().saturating_sub(32),
})
}
fn dir(&self) -> &Path {
Path::new("(recording-io)")
}
fn exists(&self, path: &Path) -> bool {
self.files.lock().unwrap().contains_key(path)
}
}
struct RealIo {
dir: PathBuf,
}
impl RealIo {
fn new(dir: PathBuf) -> Self {
fs::create_dir_all(&dir).expect("create cache dir");
Self { dir }
}
fn final_path(&self, key: u64) -> PathBuf {
self.dir.join(format!("{key:016x}.bin"))
}
fn temp_path(&self, key: u64) -> PathBuf {
self.dir.join(format!("{key:016x}.bin.tmp"))
}
}
impl PersistentCacheIo for RealIo {
fn serialize(&self, entry: &PersistableEntry) -> Vec<u8> {
let header = PersistedHeader {
key: entry.key,
table_id: entry.table_id,
schema_id: entry.schema_id,
run_generation: entry.run_generation,
entry_generation: entry.entry_generation,
footprint: entry.footprint.iter().collect(),
columns: entry.columns.to_vec(),
rows: entry.rows.to_vec(),
};
encode_entry(&header, &[])
}
fn encrypt(&self, _plaintext: &[u8]) -> Option<Vec<u8>> {
None
}
fn write(&self, path: &Path, bytes: &[u8]) -> std::io::Result<()> {
let mut f = fs::File::create(path)?;
f.write_all(bytes)?;
Ok(())
}
fn flush(&self) -> std::io::Result<()> {
Ok(())
}
fn sync(&self) -> std::io::Result<()> {
Ok(())
}
fn rename(&self, from: &Path, to: &Path) -> std::io::Result<()> {
fs::rename(from, to)
}
fn remove(&self, path: &Path) -> std::io::Result<()> {
match fs::remove_file(path) {
Ok(()) => Ok(()),
Err(e) if e.kind() == std::io::ErrorKind::NotFound => Ok(()),
Err(e) => Err(e),
}
}
fn clear(&self) -> std::io::Result<usize> {
let mut removed = 0;
for entry in fs::read_dir(&self.dir)?.flatten() {
let p = entry.path();
if p.extension().and_then(|s| s.to_str()) == Some("bin") {
if fs::remove_file(&p).is_ok() {
removed += 1;
}
}
}
Ok(removed)
}
fn read(&self, path: &Path) -> std::io::Result<Vec<u8>> {
fs::read(path)
}
fn decrypt(&self, bytes: &[u8]) -> Option<Vec<u8>> {
if bytes.len() < NONCE_LEN {
return Some(bytes.to_vec());
}
Some(bytes.to_vec())
}
fn deserialize(&self, bytes: &[u8]) -> Option<PersistableEntry> {
let (header, _trailing) = decode_header(bytes)?;
Some(PersistableEntry {
key: header.key,
table_id: header.table_id,
schema_id: header.schema_id,
run_generation: header.run_generation,
entry_generation: header.entry_generation,
footprint: Arc::new(header.footprint.iter().collect::<RoaringBitmap>()),
rows: Arc::from(header.rows.into_boxed_slice()),
columns: Arc::from(header.columns.into_boxed_slice()),
bytes: 0,
})
}
fn dir(&self) -> &Path {
&self.dir
}
fn exists(&self, path: &Path) -> bool {
path.exists()
}
}
struct EncryptedIo {
inner: RealIo,
cipher: AesCipher,
}
impl EncryptedIo {
fn new(dir: PathBuf, key: [u8; 32]) -> Self {
let cipher = AesCipher::new(&key).expect("32-byte key");
Self {
inner: RealIo::new(dir),
cipher,
}
}
fn fresh_nonce(&self) -> [u8; 12] {
let mut n = [0u8; 12];
mongreldb_core::encryption::fill_random(&mut n).expect("csprng");
n
}
}
impl PersistentCacheIo for EncryptedIo {
fn serialize(&self, entry: &PersistableEntry) -> Vec<u8> {
self.inner.serialize(entry)
}
fn encrypt(&self, plaintext: &[u8]) -> Option<Vec<u8>> {
let nonce = self.fresh_nonce();
let ct = self.cipher.encrypt_page(&nonce, plaintext).ok()?;
let mut out = Vec::with_capacity(NONCE_LEN + ct.len());
out.extend_from_slice(&nonce);
out.extend_from_slice(&ct);
Some(out)
}
fn write(&self, path: &Path, bytes: &[u8]) -> std::io::Result<()> {
self.inner.write(path, bytes)
}
fn flush(&self) -> std::io::Result<()> {
self.inner.flush()
}
fn sync(&self) -> std::io::Result<()> {
self.inner.sync()
}
fn rename(&self, from: &Path, to: &Path) -> std::io::Result<()> {
self.inner.rename(from, to)
}
fn remove(&self, path: &Path) -> std::io::Result<()> {
self.inner.remove(path)
}
fn clear(&self) -> std::io::Result<usize> {
self.inner.clear()
}
fn read(&self, path: &Path) -> std::io::Result<Vec<u8>> {
self.inner.read(path)
}
fn decrypt(&self, bytes: &[u8]) -> Option<Vec<u8>> {
if bytes.len() < NONCE_LEN {
return None;
}
let nonce: [u8; 12] = bytes[..NONCE_LEN].try_into().ok()?;
self.cipher.decrypt_page(&nonce, &bytes[NONCE_LEN..]).ok()
}
fn deserialize(&self, bytes: &[u8]) -> Option<PersistableEntry> {
let plaintext = self.decrypt(bytes)?;
self.inner.deserialize(&plaintext)
}
fn dir(&self) -> &Path {
self.inner.dir()
}
fn exists(&self, path: &Path) -> bool {
path.exists()
}
}
#[derive(Default, Debug)]
struct WorkerStats {
stores_written: u64,
stores_stale_skipped: u64,
stores_errored: u64,
removes_applied: u64,
removes_errored: u64,
clears_applied: u64,
}
#[derive(Debug, Clone, Copy)]
enum Outcome {
StoreWritten,
StoreErrored,
RemoveApplied,
RemoveErrored,
}
struct WorkerBarrier {
armed: Mutex<bool>,
proceed: Condvar,
}
impl WorkerBarrier {
fn new() -> Self {
Self {
armed: Mutex::new(false),
proceed: Condvar::new(),
}
}
fn arm(&self) {
*self.armed.lock().unwrap() = true;
}
fn release(&self) {
*self.armed.lock().unwrap() = false;
self.proceed.notify_all();
}
fn wait(&self, stop: &AtomicBool) {
let mut armed = self.armed.lock().unwrap();
loop {
if !*armed || stop.load(Ordering::Acquire) {
return;
}
let (g, _t) = self
.proceed
.wait_timeout(armed, Duration::from_millis(25))
.unwrap();
armed = g;
}
}
}
trait IoPathExt {
fn final_path_for(&self, key: u64) -> PathBuf;
fn temp_path_for(&self, key: u64) -> PathBuf;
fn sync_dir(&self) -> std::io::Result<()>;
}
impl IoPathExt for Arc<dyn PersistentCacheIo> {
fn final_path_for(&self, key: u64) -> PathBuf {
self.dir().join(format!("{key:016x}.bin"))
}
fn temp_path_for(&self, key: u64) -> PathBuf {
self.dir().join(format!("{key:016x}.bin.tmp"))
}
fn sync_dir(&self) -> std::io::Result<()> {
let dir = self.dir().to_path_buf();
if dir.as_os_str().is_empty() || dir == Path::new("(recording-io)") {
return Ok(());
}
match fs::File::open(&dir) {
Ok(f) => f.sync_all(),
Err(_) => Ok(()),
}
}
}
struct WorkerHandle {
join: Option<thread::JoinHandle<()>>,
stop: Arc<AtomicBool>,
}
impl WorkerHandle {
fn request_stop(&self) {
self.stop.store(true, Ordering::Release);
}
fn join(mut self) {
if let Some(j) = self.join.take() {
let _ = j.join();
}
}
}
fn spawn_worker(
writer: Arc<PersistentResultCacheWriter>,
io: Arc<dyn PersistentCacheIo>,
stats: Arc<Mutex<WorkerStats>>,
barrier: Option<Arc<WorkerBarrier>>,
) -> WorkerHandle {
let stop = Arc::new(AtomicBool::new(false));
let stop_w = stop.clone();
let stats_w = stats.clone();
let io_w = io.clone();
let writer_w = writer.clone();
let join = thread::spawn(move || loop {
if stop_w.load(Ordering::Acquire) {
return;
}
let op = writer_w.drain_one();
let Some((key, op, _gen)) = op else {
return;
};
if let Some(b) = barrier.as_ref() {
b.wait(&stop_w);
}
if stop_w.load(Ordering::Acquire) {
return;
}
let outcome: Outcome = match op {
PendingCacheOp::Store(entry) => {
let serialized = io_w.serialize(&entry);
let on_disk = match io_w.encrypt(&serialized) {
Some(ct) => ct,
None => serialized,
};
let final_path = io_w.final_path_for(key);
let tmp_path = io_w.temp_path_for(key);
if io_w.write(&tmp_path, &on_disk).is_err() {
Outcome::StoreErrored
} else {
let _ = io_w.flush();
let _ = io_w.sync();
if io_w.rename(&tmp_path, &final_path).is_err() {
let _ = io_w.remove(&tmp_path);
Outcome::StoreErrored
} else {
let _ = io_w.sync_dir();
Outcome::StoreWritten
}
}
}
PendingCacheOp::Remove => {
let p = io_w.final_path_for(key);
if io_w.remove(&p).is_err() {
Outcome::RemoveErrored
} else {
Outcome::RemoveApplied
}
}
};
let mut s = stats_w.lock().unwrap();
match outcome {
Outcome::StoreWritten => s.stores_written += 1,
Outcome::StoreErrored => s.stores_errored += 1,
Outcome::RemoveApplied => s.removes_applied += 1,
Outcome::RemoveErrored => s.removes_errored += 1,
}
});
WorkerHandle {
join: Some(join),
stop,
}
}
fn load_cache(
io: &dyn PersistentCacheIo,
expected_schema_id: u64,
expected_run_generation: u64,
) -> Vec<PersistableEntry> {
let mut entries = Vec::new();
let dir = io.dir().to_path_buf();
let Ok(read) = fs::read_dir(&dir) else {
return entries;
};
for ent in read.flatten() {
let p = ent.path();
if p.extension().and_then(|s| s.to_str()) != Some("bin") {
continue;
}
let bytes = match io.read(&p) {
Ok(b) => b,
Err(_) => continue,
};
let decrypted = match io.decrypt(&bytes) {
Some(d) => d,
None => continue,
};
let entry = match io.deserialize(&decrypted) {
Some(e) => e,
None => continue,
};
if entry.schema_id != expected_schema_id {
continue;
}
if entry.run_generation != expected_run_generation {
continue;
}
entries.push(entry);
}
entries
}
fn new_writer(limits: WriterLimits) -> Arc<PersistentResultCacheWriter> {
Arc::new(PersistentResultCacheWriter::for_test(limits))
}
fn small_limits() -> WriterLimits {
WriterLimits {
max_pending_keys: 8,
max_pending_bytes: 64 * 1024,
}
}
fn wait_until<F: Fn() -> bool>(max_iters: u32, sleep: Duration, cond: F) -> bool {
for _ in 0..max_iters {
if cond() {
return true;
}
thread::sleep(sleep);
}
cond()
}
#[test]
fn enqueue_does_not_perform_io_on_query_thread() {
let writer = new_writer(small_limits());
let io = Arc::new(RecordingIo::new());
let stats = Arc::new(Mutex::new(WorkerStats::default()));
let _worker = spawn_worker(writer.clone(), io.clone(), stats, None);
let query_thread = thread::current().id();
writer.enqueue_store(make_entry(0xAA, 1, 1, b"alpha"));
writer.enqueue_store(make_entry(0xBB, 1, 1, b"beta"));
writer.enqueue_remove(0xCC);
let snap = writer.persist_snapshot();
assert_eq!(snap.result_cache_persist_enqueued_total, 2);
assert_eq!(snap.result_cache_persist_remove_total, 1);
assert!(wait_until(500, Duration::from_millis(2), || !io
.log()
.is_empty()));
let log = io.log();
assert!(!log.is_empty(), "worker should have performed I/O");
let banned = [
IoOp::Write,
IoOp::Flush,
IoOp::Sync,
IoOp::Rename,
IoOp::Encrypt,
IoOp::Serialize,
];
for (tid, op) in &log {
if *tid == query_thread {
assert!(!banned.contains(op), "query thread must not perform {op:?}");
}
}
writer.shutdown();
}
#[test]
fn blocking_writer_does_not_block_query() {
let writer = new_writer(WriterLimits {
max_pending_keys: 1024,
max_pending_bytes: 1024 * 1024,
});
let t0 = Instant::now();
for i in 0..50u64 {
writer.enqueue_store(make_entry(i, 1, 1, b"payload"));
}
writer.enqueue_remove(0xFF);
let elapsed = t0.elapsed();
assert!(
elapsed < Duration::from_millis(500),
"query thread must not block on a stalled worker (took {elapsed:?})"
);
let snap = writer.persist_snapshot();
assert_eq!(snap.result_cache_persist_enqueued_total, 50);
assert_eq!(snap.result_cache_persist_remove_total, 1);
assert_eq!(snap.result_cache_persist_queue_depth, 51);
writer.shutdown();
}
#[test]
fn queue_full_does_not_block_query() {
let writer = Arc::new(PersistentResultCacheWriter::for_test(WriterLimits {
max_pending_keys: 1,
max_pending_bytes: 1024,
}));
let t0 = Instant::now();
for i in 0..64u64 {
writer.enqueue_store(make_entry(i, 1, 1, b"x"));
}
let elapsed = t0.elapsed();
assert!(
elapsed < Duration::from_millis(500),
"queue-full enqueue must not block (took {elapsed:?})"
);
let snap = writer.persist_snapshot();
assert_eq!(snap.result_cache_persist_dropped_store_total, 63);
assert_eq!(snap.result_cache_persist_enqueued_total, 1);
writer.shutdown();
}
#[test]
fn store_invalidate_before_write() {
let dir = tempdir().unwrap();
let io: Arc<dyn PersistentCacheIo> = Arc::new(RealIo::new(dir.path().to_path_buf()));
let writer = new_writer(small_limits());
let stats = Arc::new(Mutex::new(WorkerStats::default()));
let barrier = Arc::new(WorkerBarrier::new());
barrier.arm();
let worker = spawn_worker(
writer.clone(),
io.clone(),
stats.clone(),
Some(barrier.clone()),
);
writer.enqueue_store(make_entry(0xA1, 1, 1, b"will-be-removed"));
assert!(wait_until(200, Duration::from_millis(2), || {
!io.exists(&io.final_path_for(0xA1))
}));
writer.enqueue_remove(0xA1);
barrier.release();
assert!(wait_until(1000, Duration::from_millis(2), || {
stats.lock().unwrap().removes_applied >= 1
}));
worker.request_stop();
writer.shutdown();
worker.join();
let s = stats.lock().unwrap();
assert!(
!io.exists(&io.final_path_for(0xA1)),
"file for removed key must not exist on disk; stats = {:?}",
*s
);
assert_eq!(s.removes_applied, 1, "remove must be applied");
let snap = writer.persist_snapshot();
assert_eq!(snap.result_cache_persist_enqueued_total, 1);
assert_eq!(snap.result_cache_persist_remove_total, 1);
}
#[test]
fn store_a_store_b_publish_b_only() {
let dir = tempdir().unwrap();
let io: Arc<dyn PersistentCacheIo> = Arc::new(RealIo::new(dir.path().to_path_buf()));
let writer = new_writer(small_limits());
let stats = Arc::new(Mutex::new(WorkerStats::default()));
let worker = spawn_worker(writer.clone(), io.clone(), stats.clone(), None);
writer.enqueue_store(make_entry(0xCAFE, 1, 1, b"AAAA"));
writer.enqueue_store(make_entry(0xCAFE, 1, 1, b"BBBB"));
let snap = writer.persist_snapshot();
assert_eq!(snap.result_cache_persist_coalesced_total, 1);
assert_eq!(snap.result_cache_persist_enqueued_total, 2);
assert!(wait_until(1000, Duration::from_millis(2), || {
stats.lock().unwrap().stores_written == 1
}));
worker.request_stop();
writer.shutdown();
worker.join();
let final_path = io.final_path_for(0xCAFE);
assert!(io.exists(&final_path));
let bytes = io.read(&final_path).unwrap();
assert!(
bytes.windows(4).any(|w| w == b"BBBB"),
"file must contain B payload"
);
assert!(
!bytes.windows(4).any(|w| w == b"AAAA"),
"file must not contain A payload"
);
let loaded = load_cache(&*io, 1, 1);
assert_eq!(loaded.len(), 1);
let loaded_bytes: &[u8] = &loaded[0].rows;
assert_eq!(loaded_bytes, b"BBBB");
}
#[test]
fn store_clear_reopen_no_old_entry() {
let dir = tempdir().unwrap();
let io: Arc<dyn PersistentCacheIo> = Arc::new(RealIo::new(dir.path().to_path_buf()));
let writer = new_writer(small_limits());
let stats = Arc::new(Mutex::new(WorkerStats::default()));
let worker = spawn_worker(writer.clone(), io.clone(), stats.clone(), None);
writer.enqueue_store(make_entry(0xDEAD, 1, 1, b"stored"));
assert!(wait_until(1000, Duration::from_millis(2), || {
io.exists(&io.final_path_for(0xDEAD))
}));
let removed = io.clear().expect("clear");
assert_eq!(removed, 1, "the one stored entry must be cleared");
worker.request_stop();
writer.shutdown();
worker.join();
assert!(!io.exists(&io.final_path_for(0xDEAD)));
let loaded = load_cache(&*io, 1, 1);
assert!(loaded.is_empty(), "old entry must not load after clear");
}
#[test]
fn crash_after_temp_write_before_rename() {
let dir = tempdir().unwrap();
let io: Arc<dyn PersistentCacheIo> = Arc::new(RealIo::new(dir.path().to_path_buf()));
let writer = new_writer(small_limits());
let stats = Arc::new(Mutex::new(WorkerStats::default()));
let barrier = Arc::new(WorkerBarrier::new());
barrier.arm();
let worker = spawn_worker(
writer.clone(),
io.clone(),
stats.clone(),
Some(barrier.clone()),
);
writer.enqueue_store(make_entry(0xBEEF, 1, 1, b"half-written"));
assert!(wait_until(500, Duration::from_millis(2), || {
writer.persist_snapshot().result_cache_persist_queue_depth == 0
}));
assert!(
!io.exists(&io.temp_path_for(0xBEEF)),
"temp file must not exist while worker is parked"
);
assert!(
!io.exists(&io.final_path_for(0xBEEF)),
"final file must not exist before rename"
);
worker.request_stop();
worker.join();
let loaded = load_cache(&*io, 1, 1);
assert!(loaded.is_empty(), "no entry should load after crash");
}
#[test]
fn crash_after_rename_before_dir_sync() {
let dir = tempdir().unwrap();
let io: Arc<dyn PersistentCacheIo> = Arc::new(RealIo::new(dir.path().to_path_buf()));
let writer = new_writer(small_limits());
let stats = Arc::new(Mutex::new(WorkerStats::default()));
let barrier = Arc::new(WorkerBarrier::new());
barrier.arm();
let worker = spawn_worker(
writer.clone(),
io.clone(),
stats.clone(),
Some(barrier.clone()),
);
writer.enqueue_store(make_entry(0xF00D, 1, 1, b"renamed"));
assert!(wait_until(500, Duration::from_millis(2), || {
writer.persist_snapshot().result_cache_persist_queue_depth == 0
}));
barrier.release();
assert!(wait_until(500, Duration::from_millis(2), || {
io.exists(&io.final_path_for(0xF00D))
}));
worker.request_stop();
writer.shutdown();
worker.join();
let loaded = load_cache(&*io, 1, 1);
assert_eq!(
loaded.len(),
1,
"post-rename entry must load (entry is visible to subsequent reads)"
);
}
#[test]
fn encryption_round_trip() {
let dir = tempdir().unwrap();
let key = [0xABu8; 32];
let io: Arc<dyn PersistentCacheIo> = Arc::new(EncryptedIo::new(dir.path().to_path_buf(), key));
let writer = new_writer(small_limits());
let stats = Arc::new(Mutex::new(WorkerStats::default()));
let worker = spawn_worker(writer.clone(), io.clone(), stats.clone(), None);
writer.enqueue_store(make_entry(0xC0DE, 1, 1, b"super-secret-payload"));
assert!(wait_until(1000, Duration::from_millis(2), || {
io.exists(&io.final_path_for(0xC0DE))
}));
worker.request_stop();
writer.shutdown();
worker.join();
let final_path = io.final_path_for(0xC0DE);
let raw = io.read(&final_path).expect("file exists");
assert!(
!raw.windows(b"super-secret-payload".len())
.any(|w| w == b"super-secret-payload"),
"encrypted file must not contain plaintext payload"
);
let entry = io.deserialize(&raw).expect("deserialize");
assert_eq!(entry.key, 0xC0DE);
assert_eq!(entry.schema_id, 1);
assert_eq!(entry.run_generation, 1);
}
#[test]
fn wrong_key_and_corrupt_tag() {
let dir = tempdir().unwrap();
let key_a = [0x11u8; 32];
let key_b = [0x22u8; 32];
let io_a: Arc<dyn PersistentCacheIo> =
Arc::new(EncryptedIo::new(dir.path().to_path_buf(), key_a));
let writer = new_writer(small_limits());
let stats = Arc::new(Mutex::new(WorkerStats::default()));
let worker = spawn_worker(writer.clone(), io_a.clone(), stats.clone(), None);
writer.enqueue_store(make_entry(0xBADD, 1, 1, b"secret"));
assert!(wait_until(1000, Duration::from_millis(2), || {
io_a.exists(&io_a.final_path_for(0xBADD))
}));
worker.request_stop();
writer.shutdown();
worker.join();
let raw = io_a.read(&io_a.final_path_for(0xBADD)).unwrap();
let io_b: Arc<dyn PersistentCacheIo> =
Arc::new(EncryptedIo::new(dir.path().to_path_buf(), key_b));
assert!(
io_b.decrypt(&raw).is_none(),
"wrong key must fail to decrypt"
);
let mut corrupted = raw.clone();
let idx = corrupted.len() - 1;
corrupted[idx] ^= 0xFF;
let io_a2: Arc<dyn PersistentCacheIo> =
Arc::new(EncryptedIo::new(dir.path().to_path_buf(), key_a));
assert!(
io_a2.decrypt(&corrupted).is_none(),
"corrupted ciphertext must fail authentication"
);
}
#[test]
fn schema_generation_change() {
let dir = tempdir().unwrap();
let io: Arc<dyn PersistentCacheIo> = Arc::new(RealIo::new(dir.path().to_path_buf()));
let writer = new_writer(small_limits());
let stats = Arc::new(Mutex::new(WorkerStats::default()));
let worker = spawn_worker(writer.clone(), io.clone(), stats.clone(), None);
writer.enqueue_store(make_entry(0x5C, 1, 7, b"v1"));
assert!(wait_until(1000, Duration::from_millis(2), || {
io.exists(&io.final_path_for(0x5C))
}));
worker.request_stop();
writer.shutdown();
worker.join();
let loaded = load_cache(&*io, 1, 7);
assert_eq!(loaded.len(), 1);
let loaded = load_cache(&*io, 2, 7);
assert!(
loaded.is_empty(),
"schema_id mismatch must reject the entry"
);
}
#[test]
fn run_generation_change() {
let dir = tempdir().unwrap();
let io: Arc<dyn PersistentCacheIo> = Arc::new(RealIo::new(dir.path().to_path_buf()));
let writer = new_writer(small_limits());
let stats = Arc::new(Mutex::new(WorkerStats::default()));
let worker = spawn_worker(writer.clone(), io.clone(), stats.clone(), None);
writer.enqueue_store(make_entry(0x52, 1, 3, b"run-3"));
assert!(wait_until(1000, Duration::from_millis(2), || {
io.exists(&io.final_path_for(0x52))
}));
worker.request_stop();
writer.shutdown();
worker.join();
let loaded = load_cache(&*io, 1, 3);
assert_eq!(loaded.len(), 1);
let loaded = load_cache(&*io, 1, 4);
assert!(
loaded.is_empty(),
"run_generation mismatch must reject the entry"
);
}
#[test]
fn worker_io_failure() {
let dir = tempdir().unwrap();
let base_io: Arc<dyn PersistentCacheIo> = Arc::new(RealIo::new(dir.path().to_path_buf()));
struct FailingRenameIo {
inner: Arc<dyn PersistentCacheIo>,
fail_once: AtomicU64,
}
impl PersistentCacheIo for FailingRenameIo {
fn serialize(&self, e: &PersistableEntry) -> Vec<u8> {
self.inner.serialize(e)
}
fn encrypt(&self, p: &[u8]) -> Option<Vec<u8>> {
self.inner.encrypt(p)
}
fn write(&self, p: &Path, b: &[u8]) -> std::io::Result<()> {
self.inner.write(p, b)
}
fn flush(&self) -> std::io::Result<()> {
self.inner.flush()
}
fn sync(&self) -> std::io::Result<()> {
self.inner.sync()
}
fn rename(&self, from: &Path, to: &Path) -> std::io::Result<()> {
if self.fail_once.swap(0, Ordering::AcqRel) == 1 {
return Err(std::io::Error::new(
std::io::ErrorKind::PermissionDenied,
"test-injected rename failure",
));
}
self.inner.rename(from, to)
}
fn remove(&self, p: &Path) -> std::io::Result<()> {
self.inner.remove(p)
}
fn clear(&self) -> std::io::Result<usize> {
self.inner.clear()
}
fn read(&self, p: &Path) -> std::io::Result<Vec<u8>> {
self.inner.read(p)
}
fn decrypt(&self, b: &[u8]) -> Option<Vec<u8>> {
self.inner.decrypt(b)
}
fn deserialize(&self, b: &[u8]) -> Option<PersistableEntry> {
self.inner.deserialize(b)
}
fn dir(&self) -> &Path {
self.inner.dir()
}
fn exists(&self, path: &Path) -> bool {
self.inner.exists(path)
}
}
let io: Arc<dyn PersistentCacheIo> = Arc::new(FailingRenameIo {
inner: base_io.clone(),
fail_once: AtomicU64::new(1),
});
let writer = new_writer(small_limits());
let stats = Arc::new(Mutex::new(WorkerStats::default()));
let worker = spawn_worker(writer.clone(), io.clone(), stats.clone(), None);
writer.enqueue_store(make_entry(0xF1, 1, 1, b"first-attempt"));
writer.enqueue_store(make_entry(0xF2, 1, 1, b"second-attempt"));
assert!(wait_until(1000, Duration::from_millis(2), || {
let s = stats.lock().unwrap();
s.stores_errored >= 1 && s.stores_written >= 1
}));
worker.request_stop();
writer.shutdown();
worker.join();
let s = stats.lock().unwrap();
assert_eq!(
s.stores_errored, 1,
"first rename must fail; stats = {:?}",
*s
);
assert_eq!(
s.stores_written, 1,
"second insert must succeed despite prior failure"
);
}
#[test]
fn shutdown_deadline_expiry() {
let writer = new_writer(WriterLimits {
max_pending_keys: 1024,
max_pending_bytes: 1024 * 1024,
});
let mut abandoned = 0u64;
for i in 0..32u64 {
writer.enqueue_store(make_entry(i, 1, 1, b"queued-but-never-written"));
}
let pre_snap = writer.persist_snapshot();
assert_eq!(pre_snap.result_cache_persist_queue_depth, 32);
assert_eq!(pre_snap.result_cache_persist_enqueued_total, 32);
abandoned = pre_snap.result_cache_persist_queue_depth;
writer.shutdown();
assert_eq!(
writer.persist_snapshot().result_cache_persist_queue_depth,
32
);
assert!(
abandoned > 0,
"ops queued before shutdown must be counted as abandoned"
);
let first = writer.drain_one().expect("queued op");
assert!(matches!(first.1, PendingCacheOp::Store(_)));
}
#[test]
fn concurrent_cache_hits_inserts_invalidations() {
let dir = tempdir().unwrap();
let io: Arc<dyn PersistentCacheIo> = Arc::new(RealIo::new(dir.path().to_path_buf()));
let writer = new_writer(WriterLimits {
max_pending_keys: 4096,
max_pending_bytes: 16 * 1024 * 1024,
});
let stats = Arc::new(Mutex::new(WorkerStats::default()));
let worker = spawn_worker(writer.clone(), io.clone(), stats.clone(), None);
let mut handles = Vec::new();
for t in 0..4u64 {
let w = writer.clone();
let h = thread::spawn(move || {
for i in 0..500u64 {
let key = (t * 1000) + (i % 200);
match i % 3 {
0 => w.enqueue_store(make_entry(key, 1, 1, b"payload")),
1 => w.enqueue_store(make_entry(key, 1, 1, b"other-payload")),
_ => w.enqueue_remove(key),
}
}
});
handles.push(h);
}
for h in handles {
h.join().unwrap();
}
assert!(wait_until(2000, Duration::from_millis(2), || {
let s = stats.lock().unwrap();
s.stores_written + s.stores_errored + s.removes_applied >= 100
}));
worker.request_stop();
writer.shutdown();
worker.join();
let snap = writer.persist_snapshot();
let stored =
snap.result_cache_persist_enqueued_total + snap.result_cache_persist_dropped_store_total;
let removed = snap.result_cache_persist_remove_total;
assert!(
stored >= 1300,
"every store enqueue must be accounted (stored={stored}, snap={snap:?})"
);
assert!(
removed >= 600,
"every remove enqueue must be accounted (removed={removed}, snap={snap:?})"
);
let _ = snap.result_cache_persist_queue_depth;
let s = stats.lock().unwrap();
assert!(
s.stores_written + s.removes_applied > 0,
"concurrent ops must reach the worker; stats = {:?}",
*s
);
}
#[allow(dead_code)]
fn _persist_snapshot_fields_compile(snap: LookupMetricsSnapshot) {
let _ = snap.result_cache_persist_enqueued_total;
let _ = snap.result_cache_persist_coalesced_total;
let _ = snap.result_cache_persist_dropped_store_total;
let _ = snap.result_cache_persist_remove_total;
let _ = snap.result_cache_persist_stale_store_skipped_total;
let _ = snap.result_cache_persist_errors_total;
let _ = snap.result_cache_persist_shutdown_abandoned_total;
let _ = snap.result_cache_persist_queue_depth;
}