use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::thread;
use mvcc::{Config, Database, Error, Mvcc, ReadCommitted, Result, Serializable, Snapshot};
#[derive(Mvcc, Clone, Debug)]
#[mvcc(table = "rows")]
struct Row {
#[mvcc(primary_key)]
id: u64,
value: i64,
label: String,
}
const KEYS: u64 = 8;
fn seeded(db: &Database) -> Result<()> {
db.transaction(|tx| {
for id in 0..KEYS {
tx.insert(Row {
id,
value: 0,
label: format!("row-{id}"),
})?;
}
Ok(())
})
}
#[test]
fn readers_survive_writers_aborting_underneath_them() -> Result<()> {
let db = Arc::new(Database::open(Config::in_memory())?);
db.register::<Row>()?;
seeded(&db)?;
let stop = Arc::new(AtomicBool::new(false));
let reads = Arc::new(AtomicU64::new(0));
let aborts = Arc::new(AtomicU64::new(0));
let writers: Vec<_> = (0..4)
.map(|t| {
let db = Arc::clone(&db);
let stop = Arc::clone(&stop);
let aborts = Arc::clone(&aborts);
thread::spawn(move || {
let mut seed = 0x9e37_79b9_7f4a_7c15u64 ^ (t + 1);
while !stop.load(Ordering::Relaxed) {
seed ^= seed << 13;
seed ^= seed >> 7;
seed ^= seed << 17;
let key = seed % KEYS;
let mut tx = db.begin_with::<Snapshot>();
match tx.update::<Row>(&key, |r| {
r.value += 1;
r.label = format!("pending-{key}");
}) {
Ok(_) => {
tx.abort();
aborts.fetch_add(1, Ordering::Relaxed);
}
Err(e) => assert!(e.is_retriable(), "unexpected: {e}"),
}
}
})
})
.collect();
let readers: Vec<_> = (0..4)
.map(|t| {
let db = Arc::clone(&db);
let stop = Arc::clone(&stop);
let reads = Arc::clone(&reads);
thread::spawn(move || -> Result<()> {
let mut seed = 0xdead_beef_0bad_f00du64 ^ (t + 1);
while !stop.load(Ordering::Relaxed) {
seed ^= seed << 13;
seed ^= seed >> 7;
seed ^= seed << 17;
let mut tx = db.begin_with::<Snapshot>();
let row = tx.get::<Row>(&(seed % KEYS))?.expect("seeded");
assert_eq!(row.value, 0, "observed an aborted write");
assert_eq!(row.label, format!("row-{}", row.id), "torn version");
reads.fetch_add(1, Ordering::Relaxed);
}
Ok(())
})
})
.collect();
thread::sleep(std::time::Duration::from_millis(400));
stop.store(true, Ordering::Relaxed);
for w in writers {
w.join().expect("writer panicked");
}
for r in readers {
r.join().expect("reader panicked")?;
}
assert!(
aborts.load(Ordering::Relaxed) > 100,
"too few aborts to be meaningful"
);
assert!(
reads.load(Ordering::Relaxed) > 100,
"too few reads to be meaningful"
);
Ok(())
}
#[test]
fn concurrent_inserts_survive_the_slot_map_growing_under_readers() -> Result<()> {
const N: u64 = 12_000;
const WRITERS: u64 = 4;
let db = Arc::new(Database::open(Config::in_memory())?);
db.register::<Row>()?;
let stop = Arc::new(AtomicBool::new(false));
let probes = Arc::new(AtomicU64::new(0));
let readers: Vec<_> = (0..4)
.map(|t| {
let db = Arc::clone(&db);
let stop = Arc::clone(&stop);
let probes = Arc::clone(&probes);
thread::spawn(move || -> Result<()> {
let mut seed = 0xfeed_face_cafe_d00du64 ^ (t + 1);
while !stop.load(Ordering::Relaxed) {
seed ^= seed << 13;
seed ^= seed >> 7;
seed ^= seed << 17;
let id = seed % N;
let mut tx = db.begin_with::<Snapshot>();
if let Some(row) = tx.get::<Row>(&id)? {
assert_eq!(row.id, id, "lookup returned another key's slot");
assert_eq!(row.label, format!("row-{id}"), "torn version");
}
probes.fetch_add(1, Ordering::Relaxed);
}
Ok(())
})
})
.collect();
let writers: Vec<_> = (0..WRITERS)
.map(|t| {
let db = Arc::clone(&db);
thread::spawn(move || -> Result<()> {
for id in (t % 2..N).step_by(2) {
let outcome = db.transaction(|tx| {
tx.insert(Row {
id,
value: 0,
label: format!("row-{id}"),
})
});
if let Err(e) = outcome
&& !matches!(e, Error::DuplicateKey { .. } | Error::WriteConflict { .. })
{
return Err(e);
}
}
Ok(())
})
})
.collect();
for w in writers {
w.join().expect("writer panicked")?;
}
stop.store(true, Ordering::Relaxed);
for r in readers {
r.join().expect("reader panicked")?;
}
let mut tx = db.begin();
let rows = tx.scan::<Row>()?;
assert_eq!(
rows.len(),
N as usize,
"slot count wrong after concurrent growth"
);
let mut ids: Vec<u64> = rows.iter().map(|r| r.id).collect();
ids.dedup();
assert_eq!(ids.len(), N as usize, "a key was installed in two slots");
drop(tx);
let mut tx = db.begin();
for id in 0..N {
assert!(
tx.get::<Row>(&id)?.is_some(),
"key {id} unreachable by lookup"
);
}
assert!(
probes.load(Ordering::Relaxed) > 1_000,
"too few concurrent probes to be meaningful"
);
Ok(())
}
#[test]
fn contended_serializable_transfers_conserve_and_terminate() -> Result<()> {
let db = Arc::new(Database::open(Config::in_memory())?);
db.register::<Row>()?;
db.transaction(|tx| {
for id in 0..KEYS {
tx.insert(Row {
id,
value: 1_000,
label: format!("row-{id}"),
})?;
}
Ok(())
})?;
let threads: Vec<_> = (0..4)
.map(|t| {
let db = Arc::clone(&db);
thread::spawn(move || -> Result<()> {
let mut seed = 0x51_7c_c1_b7u64 ^ (t + 1);
for _ in 0..300 {
seed ^= seed << 13;
seed ^= seed >> 7;
seed ^= seed << 17;
let (from, to) = (seed % KEYS, (seed >> 8) % KEYS);
if from == to {
continue;
}
let outcome = db.transaction_with::<Serializable, _, _>(|tx| {
let balance = tx.get::<Row>(&from)?.map(|r| r.value).unwrap_or(0);
if balance < 10 {
return Ok(());
}
tx.update::<Row>(&from, |r| r.value -= 10)?;
tx.update::<Row>(&to, |r| r.value += 10)?;
Ok(())
});
if let Err(e) = outcome {
match e {
Error::SerializationFailure | Error::WriteConflict { .. } => {}
e => return Err(e),
}
}
}
Ok(())
})
})
.collect();
for t in threads {
t.join().expect("worker panicked")?;
}
let mut tx = db.begin();
let total: i64 = tx.scan::<Row>()?.iter().map(|r| r.value).sum();
assert_eq!(total, KEYS as i64 * 1_000, "money was created or destroyed");
Ok(())
}
#[test]
fn only_one_transaction_may_take_a_unique_key() -> Result<()> {
use std::sync::Barrier;
#[derive(Mvcc, Clone, Debug)]
#[mvcc(table = "unique_rows")]
struct Member {
#[mvcc(primary_key)]
id: u64,
#[mvcc(index(unique))]
email: String,
}
const THREADS: u64 = 4;
const ROUNDS: u64 = 200;
let db = Arc::new(Database::open(Config::in_memory())?);
db.register::<Member>()?;
let barrier = Arc::new(Barrier::new(THREADS as usize));
let winners = Arc::new(AtomicU64::new(0));
let threads: Vec<_> = (0..THREADS)
.map(|t| {
let db = Arc::clone(&db);
let barrier = Arc::clone(&barrier);
let winners = Arc::clone(&winners);
thread::spawn(move || -> Result<()> {
for round in 0..ROUNDS {
barrier.wait();
let email = format!("round-{round}@example.com");
let outcome = db.transaction(|tx| {
tx.insert(Member {
id: round * THREADS + t,
email: email.clone(),
})
});
match outcome {
Ok(()) => {
winners.fetch_add(1, Ordering::Relaxed);
}
Err(Error::DuplicateKey { .. } | Error::WriteConflict { .. }) => {}
Err(e) => return Err(e),
}
}
Ok(())
})
})
.collect();
for t in threads {
t.join().expect("worker panicked")?;
}
let mut tx = db.begin();
let rows = tx.scan::<Member>()?;
assert_eq!(
rows.len(),
ROUNDS as usize,
"{} rows for {ROUNDS} unique keys",
rows.len()
);
let mut emails: Vec<String> = rows.iter().map(|m| m.email.clone()).collect();
emails.sort();
let distinct = emails.len();
emails.dedup();
assert_eq!(emails.len(), distinct, "a unique key was taken twice");
assert_eq!(
winners.load(Ordering::Relaxed),
ROUNDS,
"exactly one transaction per round may win"
);
Ok(())
}
#[test]
fn a_racing_insert_does_not_overwrite_a_committed_row() -> Result<()> {
use std::sync::Barrier;
const THREADS: u64 = 4;
const ROUNDS: u64 = 200;
let db = Arc::new(Database::open(Config::in_memory())?);
db.register::<Row>()?;
let barrier = Arc::new(Barrier::new(THREADS as usize));
let winners = Arc::new(AtomicU64::new(0));
let threads: Vec<_> = (0..THREADS)
.map(|t| {
let db = Arc::clone(&db);
let barrier = Arc::clone(&barrier);
let winners = Arc::clone(&winners);
thread::spawn(move || -> Result<()> {
for round in 0..ROUNDS {
barrier.wait();
let outcome = db.transaction_with::<ReadCommitted, _, _>(|tx| {
tx.insert(Row {
id: round,
value: t as i64,
label: format!("by-{t}"),
})
});
match outcome {
Ok(()) => {
winners.fetch_add(1, Ordering::Relaxed);
}
Err(Error::DuplicateKey { .. } | Error::WriteConflict { .. }) => {}
Err(e) => return Err(e),
}
}
Ok(())
})
})
.collect();
for t in threads {
t.join().expect("worker panicked")?;
}
assert_eq!(
winners.load(Ordering::Relaxed),
ROUNDS,
"every key must have exactly one successful inserter"
);
let mut tx = db.begin();
assert_eq!(tx.scan::<Row>()?.len(), ROUNDS as usize);
Ok(())
}