use std::future::Future;
use std::path::{Path, PathBuf};
use std::pin::Pin;
use std::sync::atomic::{AtomicBool, AtomicI64, Ordering};
use std::sync::{Arc, Mutex};
use corium_core::{Cardinality, EntityId, Partition, Schema, Value, ValueType};
use corium_db::{Db, attribute};
use corium_log::{TransactionLog, VersionedLog};
use corium_store::{
BlobId, BlobIdStream, BlobStore, DbRoot, MemoryStore, RootStore, StoreError, db_root_name,
};
use corium_transactor::lease::{self, Lease, LeaseError};
use corium_transactor::{EmbeddedTransactor, TransactError};
use corium_tx::{EntityRef, TxItem, TxOp};
const DB: &str = "sim";
const TTL_MS: i64 = 5_000;
type TakeoverFuture = Pin<Box<dyn Future<Output = ()> + Send>>;
type Takeover = Box<dyn FnOnce() -> TakeoverFuture + Send>;
fn schema() -> (Schema, EntityId) {
let a = EntityId::new(Partition::Db as u32, 100);
let mut schema = Schema::default();
schema.insert(attribute(100, ValueType::Long, Cardinality::One, None));
(schema, a)
}
fn long_values(db: &Db) -> Vec<i64> {
let mut values: Vec<i64> = db
.datoms()
.into_iter()
.filter_map(|datom| match datom.v {
Value::Long(v) => Some(v),
_ => None,
})
.collect();
values.sort_unstable();
values
}
#[derive(Debug)]
struct Install {
by_active: bool,
after_takeover: bool,
root: DbRoot,
}
struct Script {
ops_before_takeover: usize,
fired: bool,
takeover: Option<Takeover>,
installs: Vec<Install>,
}
struct World {
raw: Arc<MemoryStore>,
script: Mutex<Script>,
clock: AtomicI64,
}
impl World {
fn now(&self) -> i64 {
self.clock.load(Ordering::SeqCst)
}
async fn tick(&self) {
let takeover = {
let mut script = self.script.lock().expect("script lock");
if script.fired || script.ops_before_takeover > 0 {
script.ops_before_takeover = script.ops_before_takeover.saturating_sub(1);
None
} else {
script.fired = true;
script.takeover.take()
}
};
if let Some(takeover) = takeover {
takeover().await;
}
}
fn fired(&self) -> bool {
self.script.lock().expect("script lock").fired
}
async fn force_fire(&self) {
let takeover = {
let mut script = self.script.lock().expect("script lock");
script.fired = true;
script.takeover.take()
};
if let Some(takeover) = takeover {
takeover().await;
}
}
}
#[derive(Clone)]
struct ScriptedStore {
world: Arc<World>,
counted: bool,
}
#[async_trait::async_trait]
impl BlobStore for ScriptedStore {
async fn put(&self, bytes: &[u8]) -> Result<BlobId, StoreError> {
self.world.raw.put(bytes).await
}
async fn get(&self, id: &BlobId) -> Result<Option<Vec<u8>>, StoreError> {
self.world.raw.get(id).await
}
async fn delete(&self, id: &BlobId) -> Result<(), StoreError> {
self.world.raw.delete(id).await
}
async fn list(&self) -> Result<BlobIdStream, StoreError> {
self.world.raw.list().await
}
}
#[async_trait::async_trait]
impl RootStore for ScriptedStore {
async fn get_root(&self, name: &str) -> Result<Option<Vec<u8>>, StoreError> {
if self.counted {
self.world.tick().await;
}
self.world.raw.get_root(name).await
}
async fn cas_root(
&self,
name: &str,
expected: Option<&[u8]>,
new: &[u8],
) -> Result<(), StoreError> {
if self.counted {
self.world.tick().await;
}
let result = self.world.raw.cas_root(name, expected, new).await;
if result.is_ok()
&& name == db_root_name(DB)
&& let Some(root) = DbRoot::decode(new)
{
let mut script = self.world.script.lock().expect("script lock");
let after_takeover = script.fired;
script.installs.push(Install {
by_active: self.counted,
after_takeover,
root,
});
}
result
}
async fn delete_root(&self, name: &str) -> Result<(), StoreError> {
self.world.raw.delete_root(name).await
}
async fn list_roots(&self, prefix: &str) -> Result<Vec<String>, StoreError> {
self.world.raw.list_roots(prefix).await
}
}
#[derive(Debug, Eq, PartialEq)]
enum Commit {
Acked,
RefusedPre,
RefusedPost,
}
enum Step {
Commit(i64),
Publish,
Renew,
}
struct Writer {
store: ScriptedStore,
lease: Lease,
transactor: EmbeddedTransactor,
attr: EntityId,
acked: Vec<i64>,
refused_appended: Vec<i64>,
}
impl Writer {
async fn start(store: ScriptedStore, logs: &Path, owner: &str) -> Result<Self, LeaseError> {
let now = store.world.now();
let held = lease::acquire(&store, DB, owner, "", TTL_MS, now).await?;
let log: Arc<dyn TransactionLog> =
Arc::new(VersionedLog::open(logs, DB, held.version).expect("open log"));
let (schema, attr) = schema();
let transactor = EmbeddedTransactor::recover(schema, log).expect("recover");
Ok(Self {
store,
lease: held,
transactor,
attr,
acked: Vec::new(),
refused_appended: Vec::new(),
})
}
async fn commit(&mut self, value: i64) -> Commit {
if lease::verify(&self.store, DB, &self.lease).await.is_err() {
return Commit::RefusedPre;
}
if self.store.counted {
self.store.world.tick().await;
}
self.transactor
.transact([TxItem::Op(TxOp::Add(
EntityRef::Temp("e".into()),
self.attr,
Value::Long(value),
))])
.expect("append to own version file");
if lease::verify(&self.store, DB, &self.lease).await.is_err() {
self.refused_appended.push(value);
return Commit::RefusedPost;
}
self.acked.push(value);
Commit::Acked
}
async fn publish(&self) -> Result<DbRoot, TransactError> {
self.transactor
.publish_indexes(&self.store, &db_root_name(DB), self.lease.version)
.await
}
async fn renew(&mut self) -> Result<(), LeaseError> {
let now = self.store.world.now();
self.lease = lease::renew(&self.store, DB, &self.lease, TTL_MS, now).await?;
Ok(())
}
}
struct TakeoverResult {
recovered: Vec<i64>,
acked: Vec<i64>,
lease_version: u64,
}
#[derive(Default)]
struct ActiveMirror {
acked: Vec<i64>,
in_flight: Option<i64>,
}
#[derive(Default)]
struct Coverage {
refused_post: bool,
}
#[allow(clippy::too_many_lines)]
async fn run_scenario(pause_at: usize, resume: bool) -> Coverage {
let logs_dir = tempfile::tempdir().expect("tempdir");
let logs: PathBuf = logs_dir.path().to_path_buf();
let world = Arc::new(World {
raw: Arc::new(MemoryStore::default()),
script: Mutex::new(Script {
ops_before_takeover: pause_at,
fired: false,
takeover: None,
installs: Vec::new(),
}),
clock: AtomicI64::new(1_000),
});
let mirror = Arc::new(Mutex::new(ActiveMirror::default()));
let takeover_result: Arc<Mutex<Option<TakeoverResult>>> = Arc::new(Mutex::new(None));
let takeover_failed = Arc::new(AtomicBool::new(false));
{
let hook_world = Arc::clone(&world);
let mirror = Arc::clone(&mirror);
let takeover_result = Arc::clone(&takeover_result);
let takeover_failed = Arc::clone(&takeover_failed);
let logs = logs.clone();
let hook = move || {
Box::pin(async move {
let store = ScriptedStore {
world: Arc::clone(&hook_world),
counted: false,
};
let expiry = store
.world
.raw
.get_root(&db_root_name(DB))
.await
.expect("read root")
.as_deref()
.and_then(DbRoot::decode)
.map_or(0, |root| root.lease_expires_unix_ms);
hook_world.clock.fetch_max(expiry + 1, Ordering::SeqCst);
let mut b = match Writer::start(store, &logs, "owner-b").await {
Ok(b) => b,
Err(error) => {
eprintln!("takeover failed: {error}");
takeover_failed.store(true, Ordering::SeqCst);
return;
}
};
let recovered = long_values(&b.transactor.db());
{
let mirror = mirror.lock().expect("mirror lock");
for value in &mirror.acked {
assert!(
recovered.contains(value),
"takeover lost acked value {value} (recovered {recovered:?})"
);
}
for value in &recovered {
assert!(
mirror.acked.contains(value) || mirror.in_flight == Some(*value),
"takeover invented value {value}"
);
}
}
assert_eq!(b.commit(100).await, Commit::Acked, "standby serves writes");
assert_eq!(b.commit(101).await, Commit::Acked);
b.publish().await.expect("standby publishes");
*takeover_result.lock().expect("result lock") = Some(TakeoverResult {
recovered,
acked: b.acked.clone(),
lease_version: b.lease.version,
});
}) as Pin<Box<dyn Future<Output = ()> + Send>>
};
world.script.lock().expect("script lock").takeover = Some(Box::new(hook));
}
let store_a = ScriptedStore {
world: Arc::clone(&world),
counted: true,
};
let mut a = match Writer::start(store_a, &logs, "owner-a").await {
Ok(a) => a,
Err(LeaseError::Held { .. }) if world.fired() => {
return Coverage::default();
}
Err(error) => panic!("A failed to start: {error}"),
};
let steps = [
Step::Commit(1),
Step::Commit(2),
Step::Publish,
Step::Renew,
Step::Commit(3),
Step::Publish,
Step::Commit(4),
];
for step in steps {
if !resume && world.fired() {
break;
}
match step {
Step::Commit(value) => {
mirror.lock().expect("mirror lock").in_flight = Some(value);
let outcome = a.commit(value).await;
let mut mirror = mirror.lock().expect("mirror lock");
mirror.in_flight = None;
if outcome == Commit::Acked {
mirror.acked.push(value);
assert!(!world.fired(), "A acked value {value} after the takeover");
}
}
Step::Publish => {
let result = a.publish().await;
if world.fired() {
assert!(
result.is_err(),
"A published an index root after the takeover"
);
}
}
Step::Renew => {
let result = a.renew().await;
if world.fired() {
assert!(result.is_err(), "A renewed a lost lease");
}
}
}
}
if !world.fired() {
world.force_fire().await;
}
assert!(
!takeover_failed.load(Ordering::SeqCst),
"standby failed to take over"
);
if resume {
let outcome = a.commit(99).await;
assert_ne!(outcome, Commit::Acked, "deposed A acked a transaction");
assert!(
a.publish().await.is_err(),
"deposed A published an index root"
);
assert!(a.renew().await.is_err(), "deposed A renewed the lease");
}
let coverage = Coverage {
refused_post: !a.refused_appended.is_empty(),
};
drop(a);
let result = takeover_result
.lock()
.expect("result lock")
.take()
.expect("takeover ran");
{
let script = world.script.lock().expect("script lock");
for install in &script.installs {
assert!(
!(install.by_active && install.after_takeover),
"A installed a root after the takeover: {:?}",
install.root
);
}
let mut last = (0_u64, 0_u64);
for install in &script.installs {
let next = (install.root.lease_version, install.root.index_basis_t);
assert!(
next >= last,
"published root regressed from {last:?} to {next:?}"
);
last = next;
}
assert!(
script
.installs
.iter()
.any(|install| install.root.lease_version == result.lease_version),
"takeover never rewrote the root under its lease version"
);
}
let final_log = VersionedLog::open_read_only(&logs, DB).expect("open merged log");
let records = final_log.replay().expect("merged log is contiguous");
for pair in records.windows(2) {
assert_eq!(pair[1].t, pair[0].t + 1, "hole in the merged log");
}
let (schema, _) = schema();
let mut replayed = Db::new(schema);
for record in &records {
replayed = replayed.with_transaction(record.t, &record.datoms);
}
let values = long_values(&replayed);
let acked: Vec<i64> = mirror
.lock()
.expect("mirror lock")
.acked
.iter()
.copied()
.chain(result.acked.iter().copied())
.collect();
for value in &acked {
assert!(
values.contains(value),
"acked value {value} missing from the final log (log has {values:?})"
);
}
let mut unique = values.clone();
unique.dedup();
assert_eq!(unique, values, "duplicate transaction in the final log");
for value in &values {
assert!(
acked.contains(value) || a_could_have_appended(*value),
"final log contains unexplained value {value}"
);
}
for value in &result.recovered {
assert!(values.contains(value) || result.acked.contains(value));
}
coverage
}
fn a_could_have_appended(value: i64) -> bool {
[1, 2, 3, 4, 99].contains(&value)
}
const MAX_BOUNDARIES: usize = 40;
#[tokio::test]
async fn takeover_at_every_boundary_preserves_acked_and_never_double_publishes() {
let mut fence_exercised = false;
for pause_at in 0..MAX_BOUNDARIES {
for resume in [false, true] {
fence_exercised |= run_scenario(pause_at, resume).await.refused_post;
}
}
assert!(
fence_exercised,
"no timing hit the post-append fence; the sweep lost its \
append-vs-takeover race coverage"
);
}
#[tokio::test]
async fn workload_fits_within_enumerated_boundaries() {
let logs_dir = tempfile::tempdir().expect("tempdir");
let world = Arc::new(World {
raw: Arc::new(MemoryStore::default()),
script: Mutex::new(Script {
ops_before_takeover: MAX_BOUNDARIES,
fired: false,
takeover: None,
installs: Vec::new(),
}),
clock: AtomicI64::new(1_000),
});
let store = ScriptedStore {
world: Arc::clone(&world),
counted: true,
};
let mut a = Writer::start(store, logs_dir.path(), "owner-a")
.await
.expect("start");
for value in [1, 2] {
assert_eq!(a.commit(value).await, Commit::Acked);
}
a.publish().await.expect("publish");
a.renew().await.expect("renew");
for value in [3, 4] {
assert_eq!(a.commit(value).await, Commit::Acked);
}
a.publish().await.expect("publish");
assert!(
world
.script
.lock()
.expect("script lock")
.ops_before_takeover
> 0,
"workload used more boundaries than MAX_BOUNDARIES enumerates"
);
}