use std::sync::{Arc, Barrier};
use std::thread;
use std::time::Instant;
use market_square::area::area;
use crossbeam_channel::{TryRecvError, TrySendError, bounded};
use market_square::arithmetics::NumericType;
const N_WRITERS: usize = 5; type WriterIsExclusive = (); const M_READERS: usize = 5; const MESSAGES_PER_WRITER: usize = 1_000_000;
const BUFFER_CAPACITY: usize = 16384; const BATCH_SIZE: usize = 4;
fn main() {
println!("Benchmarking with:");
println!(" Writers: {}", N_WRITERS);
println!(" Readers: {}", M_READERS);
println!(" Messages per writer: {}", MESSAGES_PER_WRITER);
println!(" Buffer capacity: {}", BUFFER_CAPACITY);
println!(" Batch size for market square: {}", BATCH_SIZE);
println!(" Total messages sent: {}", N_WRITERS * MESSAGES_PER_WRITER);
println!();
benchmark_market_square();
benchmark_crossbeam();
}
fn benchmark_market_square() {
println!("--- Market Square (Broadcast) ---");
let reader_cap = (M_READERS + 1).next_power_of_two().max(8);
let (writer, mut reader) = area(BUFFER_CAPACITY, reader_cap);
reader.suspend().expect("Couldn't suspend first reader!");
let start = Instant::now();
let mut reader_handles = vec![];
let mut writer_handles = vec![];
for i in 0..M_READERS {
let mut reader = reader.create_reader_with_seed(100 + i as NumericType).unwrap();
reader_handles.push(thread::spawn(move || {
let mut count = 0;
let total_expected = (N_WRITERS * MESSAGES_PER_WRITER) as u64;
while let Ok(slice) = reader.read_with_check() {
let _ = slice.try_cleanup_old_slots::<()>();
count += slice.len() as u64;
}
if count != total_expected {
panic!("Read a different number of messages than expected! expected {}, read {}", total_expected, count);
}
count
}));
}
assert!(MESSAGES_PER_WRITER % BATCH_SIZE == 0, "MESSAGES_PER_WRITER must be a multiple of BATCH_SIZE");
let barrier = Arc::new(Barrier::new(N_WRITERS));
for _ in 0..N_WRITERS {
let writer = writer.create_writer();
let barrier = barrier.clone();
writer_handles.push(thread::spawn(move || {
barrier.wait();
for msg_idx in 0..(MESSAGES_PER_WRITER / BATCH_SIZE) {
loop {
match writer.reserve::<WriterIsExclusive>(BATCH_SIZE) {
Ok(mut reservation) => {
for i in 0..BATCH_SIZE {
reservation.get_mut(i).unwrap().write((msg_idx * BATCH_SIZE + i) as u64); }
if N_WRITERS == 1 {
match unsafe { reservation.publish::<market_square::area::Exclusive>() } {
Ok(_) => (),
Err(_) => panic!("Failed to publish reservation! This should never happen with a single writer."),
}
} else {
unsafe { reservation.publish_spin() };
}
break;
}
Err(_) => {
thread::yield_now();
}
}
}
}
}));
}
drop(reader); drop(writer);
for handle in writer_handles {
handle.join().unwrap();
}
for handle in reader_handles {
handle.join().unwrap();
}
let duration = start.elapsed();
println!("Time: {:?}", duration);
println!("Total reads: {}", M_READERS * N_WRITERS * MESSAGES_PER_WRITER);
println!("Throughput: {:.2} million reads/sec",
(M_READERS * N_WRITERS * MESSAGES_PER_WRITER) as f64 / duration.as_secs_f64() / 1_000_000.0);
println!();
}
fn benchmark_crossbeam() {
println!("--- Crossbeam MPMC (Queue) ---");
let (tx, rx) = bounded::<u64>(BUFFER_CAPACITY);
let start = Instant::now();
let mut reader_handles = vec![];
let mut writer_handles = vec![];
for _ in 0..M_READERS {
let rx = rx.clone();
reader_handles.push(thread::spawn(move || {
let mut count = 0;
loop {
match rx.try_recv() {
Err(TryRecvError::Empty) => thread::yield_now(),
Err(TryRecvError::Disconnected) => break,
Ok(_) => count += 1,
}
}
count
}));
}
let barrier = Arc::new(Barrier::new(N_WRITERS));
for _ in 0..N_WRITERS {
let tx = tx.clone();
let barrier = barrier.clone();
writer_handles.push(thread::spawn(move || {
barrier.wait();
for i in 0..MESSAGES_PER_WRITER {
loop {
match tx.try_send(i as u64) {
Ok(_) => break,
Err(TrySendError::Full(_)) => {
thread::yield_now();
},
Err(TrySendError::Disconnected(_)) => {
panic!("Channel disconnected unexpectedly!");
}
}
}
}
}));
}
drop(tx);
for handle in writer_handles {
handle.join().unwrap();
}
let mut total_reads = 0;
for handle in reader_handles {
total_reads += handle.join().unwrap();
}
let duration = start.elapsed();
println!("Time: {:?}", duration);
println!("Total reads: {}", total_reads);
println!("Throughput: {:.2} million reads/sec",
total_reads as f64 / duration.as_secs_f64() / 1_000_000.0);
println!();
}