use std::collections::BTreeSet;
use std::sync::{Arc, Barrier};
use std::thread;
use haematite::db::{Database, DatabaseConfig};
type TestResult = Result<(), Box<dyn std::error::Error>>;
const SHARDS: usize = 4;
const WRITERS: usize = 6;
const READERS: usize = 4;
const KEYS_PER_WRITER: u64 = 24;
const ROUNDS: u64 = 20;
const READER_PASSES: usize = 10;
fn key_for(writer: usize, index: u64) -> Vec<u8> {
format!("w{writer:02}-k{index:04}").into_bytes()
}
fn config(dir: &std::path::Path) -> DatabaseConfig {
DatabaseConfig {
data_dir: dir.to_path_buf(),
shard_count: SHARDS,
distributed: None,
}
}
#[test]
#[ignore = "concurrency soak: tens of seconds of shard-actor round-trips; run with --ignored"]
fn many_writers_across_shards_with_concurrent_readers() -> TestResult {
let dir = tempfile::tempdir()?;
let db = Arc::new(Database::create(config(dir.path()))?);
let mut shards_hit = BTreeSet::new();
for w in 0..WRITERS {
for i in 0..KEYS_PER_WRITER {
shards_hit.insert(db.shard_for(&key_for(w, i)));
}
}
assert_eq!(
shards_hit.len(),
SHARDS,
"test keys must cover all {SHARDS} shards (covered: {shards_hit:?})"
);
let start = Arc::new(Barrier::new(WRITERS + READERS));
let mut handles: Vec<thread::JoinHandle<Result<(), String>>> = Vec::new();
for w in 0..WRITERS {
let db = Arc::clone(&db);
let start = Arc::clone(&start);
handles.push(thread::spawn(move || -> Result<(), String> {
start.wait();
for i in 0..KEYS_PER_WRITER {
let key = key_for(w, i);
db.cas(key.clone(), None, 0)
.map_err(|e| format!("create: {e}"))?;
for round in 0..ROUNDS {
db.cas(key.clone(), Some(round), round + 1)
.map_err(|e| format!("advance: {e}"))?;
}
}
Ok(())
}));
}
for _ in 0..READERS {
let db = Arc::clone(&db);
let start = Arc::clone(&start);
handles.push(thread::spawn(move || -> Result<(), String> {
start.wait();
for _pass in 0..READER_PASSES {
for w in 0..WRITERS {
for i in 0..KEYS_PER_WRITER {
let observed = db
.read_value(&key_for(w, i))
.map_err(|e| format!("read: {e}"))?;
match observed {
Some(value) if value > ROUNDS => {
return Err(format!("out-of-range value {value} (> {ROUNDS})"));
}
_ => {}
}
}
}
}
Ok(())
}));
}
for handle in handles {
handle
.join()
.map_err(|_| "worker thread panicked")?
.map_err(|e| format!("worker failed: {e}"))?;
}
for w in 0..WRITERS {
for i in 0..KEYS_PER_WRITER {
let key = key_for(w, i);
assert_eq!(
db.read_value(&key)?,
Some(ROUNDS),
"key {} did not reach final value",
String::from_utf8_lossy(&key)
);
}
}
Ok(())
}