use std::{
mem,
sync::{
atomic::{AtomicIsize, Ordering},
Arc,
},
thread::{self, JoinHandle},
};
pub struct Bursty<Global, Local> {
global: Arc<Global>,
threads: Vec<JoinHandle<Local>>,
results: Vec<Local>,
}
impl<Global, Local> Bursty<Global, Local> {
pub fn new(global: Arc<Global>, threads: Vec<JoinHandle<Local>>) -> Self {
assert!(!threads.is_empty());
let results = vec![];
Self {
global,
threads,
results,
}
}
pub fn join(&mut self) {
if self.threads.is_empty() {
return;
}
self.results
.extend(self.threads.drain(..).map(|handle| handle.join().unwrap()));
}
pub fn global(&self) -> Arc<Global> {
self.global.clone()
}
pub fn into_locals(mut self) -> Vec<Local> {
self.join();
mem::take(&mut self.results)
}
}
impl<Global, Local> Drop for Bursty<Global, Local> {
fn drop(&mut self) {
self.join();
}
}
pub struct BurstyBuilder<Global, Local> {
global: Arc<Global>,
locals: Vec<Local>,
steps: Vec<Vec<Step<Global, Local>>>,
rendez_vous: Vec<RendezVous>,
}
impl<Global, Local> BurstyBuilder<Global, Local>
where
Global: Send + Sync + 'static,
Local: Send + 'static,
{
pub fn new(global: Global, locals: Vec<Local>) -> Self {
assert!(!locals.is_empty(), "no local element, no thread will run");
let global = Arc::new(global);
let steps = {
let mut steps = vec![];
steps.resize_with(locals.len(), Vec::new);
steps
};
let rendez_vous = vec![RendezVous::new(locals.len())];
Self {
global,
locals,
steps,
rendez_vous,
}
}
pub fn add_minimal_step<Factory, Step>(&mut self, mut factory: Factory)
where
Factory: FnMut() -> Step,
Step: FnMut() + Send + 'static,
{
self.add_simple_step(move || {
let mut step = factory();
move |_: &Global, _: &mut Local| step()
})
}
pub fn add_simple_step<Factory, Step>(&mut self, mut factory: Factory)
where
Factory: FnMut() -> Step,
Step: FnMut(&Global, &mut Local) + Send + 'static,
{
self.add_complex_step(move || {
let mut step = factory();
let prep = |_: &Global, _: &mut Local| ();
let step = move |global: &Global, local: &mut Local, _: ()| step(global, local);
(prep, step)
});
}
pub fn add_complex_step<Factory, Prep, R, Step>(&mut self, mut factory: Factory)
where
Factory: FnMut() -> (Prep, Step),
Prep: FnMut(&Global, &mut Local) -> R + Send + 'static,
Step: FnMut(&Global, &mut Local, R) + Send + 'static,
{
let rendez_vous = RendezVous::new(self.locals.len());
for serie in &mut self.steps {
let rendez_vous = rendez_vous.clone();
let (mut prep, mut step) = factory();
let prev = self.rendez_vous.last().unwrap().clone();
serie.push(Box::new(move |global: &Global, local: &mut Local| {
let prepared = prep(global, local);
rendez_vous.wait_until_all_ready();
step(global, local, prepared);
prev.reset();
}));
}
self.rendez_vous.push(rendez_vous);
}
pub fn launch(mut self, iterations: usize) -> Bursty<Global, Local> {
assert!(
!self.steps.is_empty(),
"Cannot launch a burst test without a single thread"
);
assert!(
!self.steps[0].is_empty(),
"Cannot launch a burst test without a single step"
);
if self.steps[0].len() < 2 {
self.add_minimal_step(|| || ());
}
for serie in &mut self.steps {
let last = self.rendez_vous.first().unwrap().clone();
let prev = self.rendez_vous.last().unwrap().clone();
serie.push(Box::new(move |_: &Global, _: &mut Local| {
last.wait_until_all_ready();
prev.reset();
}));
}
assert!(self.steps[0].len() >= 3);
let mut threads = vec![];
let rendez_vous = Arc::new(self.rendez_vous);
for (mut local, mut serie) in self.locals.into_iter().zip(self.steps) {
let global = self.global.clone();
let rendez_vous = rendez_vous.clone();
threads.push(thread::spawn(move || {
let mut guard = PoisonGuard(rendez_vous);
let global = &*global;
for _ in 0..iterations {
for step in &mut serie {
step(global, &mut local);
}
}
guard.dismiss();
local
}));
}
let global = self.global;
Bursty::new(global, threads)
}
}
type Step<Global, Local> = Box<dyn FnMut(&Global, &mut Local) + Send + 'static>;
struct PoisonGuard(Arc<Vec<RendezVous>>);
impl PoisonGuard {
fn dismiss(&mut self) {
self.0 = Arc::default()
}
}
impl Drop for PoisonGuard {
fn drop(&mut self) {
for rendez_vous in &*self.0 {
rendez_vous.poison();
}
}
}
#[derive(Clone, Debug)]
struct RendezVous(Arc<(AtomicIsize, isize)>);
impl RendezVous {
fn new(count: usize) -> Self {
assert!(count <= (isize::MAX as usize));
Self(Arc::new((AtomicIsize::new(count as isize), count as isize)))
}
fn poison(&self) {
self.0 .0.store(-1, Ordering::Relaxed);
}
fn wait_until_all_ready(&self) {
self.0 .0.fetch_sub(1, Ordering::Relaxed);
while !self.is_ready() {}
}
fn reset(&self) {
let mut count = self.load();
while self
.0
.0
.compare_exchange(
count as isize,
self.0 .1,
Ordering::Relaxed,
Ordering::Relaxed,
)
.is_err()
{
count = self.load();
}
}
fn is_ready(&self) -> bool {
self.load() == 0
}
fn load(&self) -> usize {
let count = self.0 .0.load(Ordering::Relaxed);
if count < 0 {
self.abandon_ship()
}
count as usize
}
#[cold]
#[inline(never)]
fn abandon_ship(&self) {
panic!("Someone poisoned the well!");
}
}
#[cfg(test)]
mod tests {
use std::sync::Mutex;
use super::*;
#[derive(Clone, Debug, Default, Eq, PartialEq)]
struct LocalEvent {
iteration: usize,
step: usize,
}
#[derive(Clone, Debug, Default, Eq, PartialEq, Ord, PartialOrd)]
struct GlobalEvent {
iteration: usize,
step: usize,
thread: usize,
}
#[derive(Clone, Debug, Default, Eq, PartialEq)]
struct LocalTrace {
thread: usize,
events: Vec<LocalEvent>,
}
impl LocalTrace {
fn create(threads: usize) -> Vec<LocalTrace> {
assert!(threads > 0);
let mut result = vec![];
for t in 0..threads {
result.push(LocalTrace {
thread: t,
events: vec![],
});
}
result
}
fn expected(threads: usize, iterations: usize, steps: usize) -> Vec<LocalTrace> {
assert!(iterations > 0);
assert!(steps > 0);
let mut result = Self::create(threads);
for local in &mut result {
for i in 0..iterations {
for s in 0..steps {
local.add(i, s);
}
}
}
result
}
fn add(&mut self, iteration: usize, step: usize) {
self.events.push(LocalEvent { step, iteration });
}
}
#[derive(Debug, Default)]
struct GlobalTrace {
events: Mutex<Vec<GlobalEvent>>,
}
impl GlobalTrace {
fn create_builder(threads: usize) -> BurstyBuilder<GlobalTrace, LocalTrace> {
assert!(threads > 0);
BurstyBuilder::new(GlobalTrace::default(), LocalTrace::create(threads))
}
fn expected(threads: usize, iterations: usize, steps: usize) -> Vec<GlobalEvent> {
assert!(threads > 0);
assert!(iterations > 0);
assert!(steps > 0);
let mut result = vec![];
for i in 0..iterations {
for s in 0..steps {
for t in 0..threads {
result.push(GlobalEvent {
iteration: i,
step: s,
thread: t,
});
}
}
}
result
}
fn create_step(
step: usize,
) -> (
impl FnMut(&GlobalTrace, &mut LocalTrace) -> usize,
impl FnMut(&GlobalTrace, &mut LocalTrace, usize),
) {
let mut iteration = 0;
let prep = move |_: &GlobalTrace, _: &mut LocalTrace| {
let tmp = iteration;
iteration += 1;
tmp
};
let step = move |global: &GlobalTrace, local: &mut LocalTrace, iteration: usize| {
global.add(local.thread, iteration, step);
local.add(iteration, step);
};
(prep, step)
}
fn events(&self) -> Vec<GlobalEvent> {
fn split(slice: &mut [GlobalEvent]) -> (&mut [GlobalEvent], &mut [GlobalEvent]) {
let front = slice.first().expect("Not empty").clone();
for (i, e) in slice.iter().enumerate() {
if e.iteration != front.iteration || e.step != front.step {
return slice.split_at_mut(i);
}
}
(slice, &mut [])
}
let mut events = self.events.lock().unwrap().clone();
let mut slice = &mut events[..];
while !slice.is_empty() {
let (head, tail) = split(slice);
slice = tail;
head.sort()
}
events
}
fn add(&self, thread: usize, iteration: usize, step: usize) {
let mut events = self.events.lock().unwrap();
events.push(GlobalEvent {
iteration,
step,
thread,
});
}
}
#[test]
fn single_thread_single_step_single_iteration() {
let mut builder = GlobalTrace::create_builder(1);
builder.add_complex_step(|| GlobalTrace::create_step(0));
let mut bursty = builder.launch(1);
bursty.join();
assert_eq!(GlobalTrace::expected(1, 1, 1), bursty.global().events());
assert_eq!(LocalTrace::expected(1, 1, 1), bursty.into_locals());
}
#[test]
fn single_thread_single_step_n_iterations() {
let mut builder = GlobalTrace::create_builder(1);
builder.add_complex_step(|| GlobalTrace::create_step(0));
let mut bursty = builder.launch(3);
bursty.join();
assert_eq!(GlobalTrace::expected(1, 3, 1), bursty.global().events());
assert_eq!(LocalTrace::expected(1, 3, 1), bursty.into_locals());
}
#[test]
fn single_thread_n_steps_single_iteration() {
let mut builder = GlobalTrace::create_builder(1);
builder.add_complex_step(|| GlobalTrace::create_step(0));
builder.add_complex_step(|| GlobalTrace::create_step(1));
builder.add_complex_step(|| GlobalTrace::create_step(2));
builder.add_complex_step(|| GlobalTrace::create_step(3));
builder.add_complex_step(|| GlobalTrace::create_step(4));
let mut bursty = builder.launch(1);
bursty.join();
assert_eq!(GlobalTrace::expected(1, 1, 5), bursty.global().events());
assert_eq!(LocalTrace::expected(1, 1, 5), bursty.into_locals());
}
#[test]
fn single_thread_n_steps_n_iterations() {
let mut builder = GlobalTrace::create_builder(1);
builder.add_complex_step(|| GlobalTrace::create_step(0));
builder.add_complex_step(|| GlobalTrace::create_step(1));
builder.add_complex_step(|| GlobalTrace::create_step(2));
builder.add_complex_step(|| GlobalTrace::create_step(3));
builder.add_complex_step(|| GlobalTrace::create_step(4));
let mut bursty = builder.launch(3);
bursty.join();
assert_eq!(GlobalTrace::expected(1, 3, 5), bursty.global().events());
assert_eq!(LocalTrace::expected(1, 3, 5), bursty.into_locals());
}
#[test]
fn n_threads_single_step_single_iteration() {
let mut builder = GlobalTrace::create_builder(3);
builder.add_complex_step(|| GlobalTrace::create_step(0));
let mut bursty = builder.launch(1);
bursty.join();
assert_eq!(GlobalTrace::expected(3, 1, 1), bursty.global().events());
assert_eq!(LocalTrace::expected(3, 1, 1), bursty.into_locals());
}
#[test]
fn n_threads_single_step_n_iterations() {
let mut builder = GlobalTrace::create_builder(3);
builder.add_complex_step(|| GlobalTrace::create_step(0));
let mut bursty = builder.launch(5);
bursty.join();
assert_eq!(GlobalTrace::expected(3, 5, 1), bursty.global().events());
assert_eq!(LocalTrace::expected(3, 5, 1), bursty.into_locals());
}
#[test]
fn n_threads_n_steps_single_iteration() {
let mut builder = GlobalTrace::create_builder(3);
builder.add_complex_step(|| GlobalTrace::create_step(0));
builder.add_complex_step(|| GlobalTrace::create_step(1));
builder.add_complex_step(|| GlobalTrace::create_step(2));
builder.add_complex_step(|| GlobalTrace::create_step(3));
builder.add_complex_step(|| GlobalTrace::create_step(4));
let mut bursty = builder.launch(1);
bursty.join();
assert_eq!(GlobalTrace::expected(3, 1, 5), bursty.global().events());
assert_eq!(LocalTrace::expected(3, 1, 5), bursty.into_locals());
}
#[test]
fn n_threads_n_steps_n_iterations() {
let mut builder = GlobalTrace::create_builder(3);
builder.add_complex_step(|| GlobalTrace::create_step(0));
builder.add_complex_step(|| GlobalTrace::create_step(1));
builder.add_complex_step(|| GlobalTrace::create_step(2));
builder.add_complex_step(|| GlobalTrace::create_step(3));
builder.add_complex_step(|| GlobalTrace::create_step(4));
builder.add_complex_step(|| GlobalTrace::create_step(5));
builder.add_complex_step(|| GlobalTrace::create_step(6));
let mut bursty = builder.launch(5);
bursty.join();
assert_eq!(GlobalTrace::expected(3, 5, 7), bursty.global().events());
assert_eq!(LocalTrace::expected(3, 5, 7), bursty.into_locals());
}
}