use std::collections::BTreeSet;
use std::error::Error;
use std::sync::{Arc, Barrier};
use std::thread;
use std::time::Duration;
use beamr::module::ModuleRegistry;
use beamr::scheduler::{Scheduler, SchedulerConfig};
use super::{ShardMode, ShardRouter};
const TIMEOUT: Duration = Duration::from_secs(5);
fn scheduler() -> Result<Arc<Scheduler>, Box<dyn Error>> {
let scheduler = Scheduler::new(SchedulerConfig::default(), Arc::new(ModuleRegistry::new()))
.map_err(|message| -> Box<dyn Error> { message.into() })?;
Ok(Arc::new(scheduler))
}
fn build_router(dir: &std::path::Path, shard_count: usize) -> Result<ShardRouter, Box<dyn Error>> {
ShardRouter::new(
scheduler()?,
dir,
shard_count,
crate::tree::TreePolicy::V1_DEFAULT,
ShardMode::Create,
crate::db::root_advance::RootAdvanceSeam::new(),
)
.ok_or_else(|| "non-zero shard_count".into())
}
#[test]
fn nothing_is_materialised_until_first_touch() -> Result<(), Box<dyn Error>> {
let dir = tempfile::tempdir()?;
let router = build_router(dir.path(), 4096)?;
assert_eq!(router.materialised_shard_ids(), Vec::<usize>::new());
let first = router.handle_for_shard(7).map_err(|error| error.message)?;
assert_eq!(router.materialised_shard_ids(), vec![7]);
let again = router.handle_for_shard(7).map_err(|error| error.message)?;
assert_eq!(first.pid(), again.pid());
assert_eq!(router.materialised_shard_ids(), vec![7]);
router.shutdown_all(TIMEOUT);
Ok(())
}
#[test]
fn out_of_range_shard_id_is_rejected() -> Result<(), Box<dyn Error>> {
let dir = tempfile::tempdir()?;
let router = build_router(dir.path(), 4)?;
let error = router
.handle_for_shard(4)
.err()
.ok_or("id 4 must be rejected for a 4-shard router")?;
assert_eq!(error.shard_id, 4);
router.shutdown_all(TIMEOUT);
Ok(())
}
#[test]
fn concurrent_first_touch_of_one_cold_shard_spawns_exactly_one_actor() -> Result<(), Box<dyn Error>>
{
const THREADS: usize = 24;
let dir = tempfile::tempdir()?;
let router = Arc::new(build_router(dir.path(), 16)?);
let barrier = Arc::new(Barrier::new(THREADS));
let mut joins = Vec::with_capacity(THREADS);
for _ in 0..THREADS {
let router = Arc::clone(&router);
let barrier = Arc::clone(&barrier);
joins.push(thread::spawn(move || -> Result<u64, String> {
barrier.wait();
router
.handle_for_shard(3)
.map(|handle| handle.pid())
.map_err(|error| error.message)
}));
}
let mut pids: BTreeSet<u64> = BTreeSet::new();
for join in joins {
let pid = join.join().map_err(|_| "worker thread panicked")??;
pids.insert(pid);
}
assert_eq!(
pids.len(),
1,
"cold shard must spawn exactly one actor, saw {pids:?}"
);
assert_eq!(router.materialised_shard_ids(), vec![3]);
router.shutdown_all(TIMEOUT);
Ok(())
}
#[test]
fn concurrent_writes_to_one_cold_shard_all_persist() -> Result<(), Box<dyn Error>> {
const THREADS: usize = 16;
let dir = tempfile::tempdir()?;
let router = Arc::new(build_router(dir.path(), 16)?);
let barrier = Arc::new(Barrier::new(THREADS));
let mut joins = Vec::with_capacity(THREADS);
for worker in 0..THREADS {
let router = Arc::clone(&router);
let barrier = Arc::clone(&barrier);
joins.push(thread::spawn(move || -> Result<Vec<u8>, String> {
barrier.wait();
let handle = router.handle_for_shard(3).map_err(|error| error.message)?;
let key = format!("k{worker:02}").into_bytes();
handle
.put_with_ttl(key.clone(), worker.to_le_bytes().to_vec(), None, TIMEOUT)
.map_err(|error| error.to_string())?;
Ok(key)
}));
}
let mut keys: Vec<Vec<u8>> = Vec::with_capacity(THREADS);
for join in joins {
keys.push(join.join().map_err(|_| "worker thread panicked")??);
}
let handle = router.handle_for_shard(3).map_err(|error| error.message)?;
handle.commit(TIMEOUT)?;
for (worker, key) in keys.iter().enumerate() {
let value = handle.get(key.clone(), TIMEOUT)?;
assert_eq!(
value,
Some(worker.to_le_bytes().to_vec()),
"every concurrent write to the cold shard must persist"
);
}
router.shutdown_all(TIMEOUT);
Ok(())
}