use crate::handles::{Barrier, Consumer, Cursor, HandleInner, Producer};
use crate::ringbuffer::RingBuffer;
use crate::wait::{WaitPhased, WaitSleep};
use std::collections::{HashMap, HashSet};
use std::error::Error;
use std::fmt::{Debug, Display, Formatter};
use std::hash::{BuildHasherDefault, Hasher};
use std::sync::Arc;
use std::time::Duration;
#[derive(Debug)]
pub struct DisruptorBuilder<F, E, W> {
buffer_size: usize,
event_factory: Option<F>,
provided_buffer: Option<Box<[E]>>,
wait_strategy: W,
followed_by: U64Map<Vec<u64>>,
follows: U64Map<Follows>,
follows_lead: U64Set,
handles_map: U64Map<Handle>,
overlapping_ids: U64Set,
}
pub(crate) const BACKOFF_WAIT: WaitPhased<WaitSleep> = WaitPhased::new(
Duration::from_millis(1),
Duration::from_millis(1),
WaitSleep::new(Duration::from_micros(50)),
);
impl<E> DisruptorBuilder<fn() -> E, E, WaitPhased<WaitSleep>> {
pub fn with_buffer(buffer: Box<[E]>) -> Self
where
E: Sync,
{
DisruptorBuilder {
buffer_size: buffer.len(),
event_factory: None,
provided_buffer: Some(buffer),
wait_strategy: BACKOFF_WAIT,
followed_by: U64Map::default(),
follows: U64Map::default(),
follows_lead: U64Set::default(),
handles_map: U64Map::default(),
overlapping_ids: U64Set::default(),
}
}
}
impl<F, E> DisruptorBuilder<F, E, WaitPhased<WaitSleep>> {
pub fn new(size: usize, event_factory: F) -> Self
where
E: Sync,
F: FnMut() -> E,
{
DisruptorBuilder {
buffer_size: size,
event_factory: Some(event_factory),
provided_buffer: None,
wait_strategy: BACKOFF_WAIT,
followed_by: U64Map::default(),
follows: U64Map::default(),
follows_lead: U64Set::default(),
handles_map: U64Map::default(),
overlapping_ids: U64Set::default(),
}
}
}
impl<F, E, W> DisruptorBuilder<F, E, W>
where
F: FnMut() -> E,
W: Clone,
{
pub fn add_handle(mut self, id: u64, handle: Handle, follows: Follows) -> Self {
if self.follows.contains_key(&id) {
self.overlapping_ids.insert(id);
}
self.handles_map.insert(id, handle);
self.followed_by.entry(id).or_default();
match &follows {
Follows::LeadProducer => {
self.follows_lead.insert(id);
}
Follows::Handles(ids) => ids.iter().for_each(|follow_id| {
self.followed_by
.entry(*follow_id)
.and_modify(|vec| vec.push(id))
.or_insert_with(|| vec![id]);
}),
}
self.follows.insert(id, follows);
self
}
pub fn extend_handles(self, iter: impl IntoIterator<Item = (u64, Handle, Follows)>) -> Self {
let mut this = self;
for (id, handle, follows) in iter {
this = this.add_handle(id, handle, follows);
}
this
}
pub fn wait_strategy<W2>(self, strategy: W2) -> DisruptorBuilder<F, E, W2>
where
W2: Clone,
{
DisruptorBuilder {
buffer_size: self.buffer_size,
event_factory: self.event_factory,
provided_buffer: None,
wait_strategy: strategy,
followed_by: self.followed_by,
follows: self.follows,
follows_lead: self.follows_lead,
handles_map: self.handles_map,
overlapping_ids: self.overlapping_ids,
}
}
pub fn build(mut self) -> Result<DisruptorHandles<E, W>, BuildError> {
self.validate()?;
let buffer = Arc::new(self.construct_buffer());
let lead_cursor = Arc::new(Cursor::start());
let (producers, consumers, cursor_map) = self.construct_handles(&lead_cursor, &buffer);
let barrier = self.construct_lead_barrier(cursor_map);
let lead = HandleInner::new(lead_cursor, barrier, buffer, self.wait_strategy.clone());
Ok(DisruptorHandles {
lead: Some(lead.into_producer()),
producers,
consumers,
})
}
fn construct_buffer(&mut self) -> RingBuffer<E> {
match (self.event_factory.take(), self.provided_buffer.take()) {
(Some(event_factory), None) => {
RingBuffer::from_factory(self.buffer_size, event_factory)
}
(None, Some(buffer)) => RingBuffer::from_buffer(buffer),
_ => unreachable!("guaranteed by DisruptorBuilder construction methods"),
}
}
fn construct_lead_barrier(&self, cursor_map: U64Map<Arc<Cursor>>) -> Barrier {
let mut cursors: Vec<_> = self
.followed_by
.iter()
.filter(|(_, followed_by)| followed_by.is_empty())
.map(|(id, _)| Arc::clone(cursor_map.get(id).unwrap()))
.collect();
assert!(!cursors.is_empty());
match cursors.len() {
1 => Barrier::one(cursors.pop().unwrap()),
_ => Barrier::many(cursors.into_boxed_slice()),
}
}
#[allow(clippy::type_complexity)]
fn construct_handles(
&self,
lead_cursor: &Arc<Cursor>,
buffer: &Arc<RingBuffer<E>>,
) -> (
U64Map<Producer<E, W, false>>,
U64Map<Consumer<E, W>>,
U64Map<Arc<Cursor>>,
) {
let mut producers = U64Map::default();
let mut consumers = U64Map::default();
let mut cursor_map = U64Map::default();
fn get_cursor(id: u64, map: &mut U64Map<Arc<Cursor>>) -> Arc<Cursor> {
let cursor = map.entry(id).or_insert_with(|| Arc::new(Cursor::start()));
Arc::clone(cursor)
}
for (&id, follows) in &self.follows {
let cursor = get_cursor(id, &mut cursor_map);
let buf = Arc::clone(buffer);
let barrier = match follows {
Follows::LeadProducer => Barrier::one(Arc::clone(lead_cursor)),
Follows::Handles(ids) if ids.len() == 1 => {
Barrier::one(get_cursor(ids[0], &mut cursor_map))
}
Follows::Handles(ids) => {
let follows_cursors = ids
.iter()
.map(|follow_id| get_cursor(*follow_id, &mut cursor_map))
.collect();
Barrier::many(follows_cursors)
}
};
let handle = HandleInner::new(cursor, barrier, buf, self.wait_strategy.clone());
match self.handles_map.get(&id).unwrap() {
Handle::Producer => {
producers.insert(id, handle.into_producer());
}
Handle::Consumer => {
consumers.insert(id, handle.into_consumer());
}
}
}
(producers, consumers, cursor_map)
}
fn validate(&self) -> Result<(), BuildError> {
if self.follows.is_empty() {
return Err(BuildError::EmptyGraph);
}
for id in self.followed_by.keys() {
if !self.follows.contains_key(id) {
return Err(BuildError::UnregisteredID(*id));
}
}
if self.buffer_size == 0 || !self.buffer_size.is_power_of_two() {
return Err(BuildError::BufferSize(self.buffer_size));
}
if !self.overlapping_ids.is_empty() {
return Err(BuildError::OverlappingIDs(
self.overlapping_ids.iter().copied().collect(),
));
}
let chains = validate_graph(&self.followed_by, &self.follows_lead)?;
validate_order(&self.handles_map, chains)?;
Ok(())
}
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub enum BuildError {
BufferSize(usize),
OverlappingIDs(Vec<u64>),
EmptyGraph,
UnorderedProducer(Vec<u64>, u64),
UnregisteredID(u64),
GraphCycle(u64),
DisconnectedNode(u64),
}
impl Display for BuildError {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
let str = match self {
BuildError::BufferSize(size) => {
format!("RingBuffer size must be a non-zero power of 2; given size: {size}")
}
BuildError::OverlappingIDs(ids) => format!("Found overlapping ids: {ids:?}"),
BuildError::EmptyGraph => "Graph empty, no handles added".to_owned(),
BuildError::UnorderedProducer(chain, id) => {
format!("Chain of handle ids: {chain:?} does not contain producer id: {id}")
}
BuildError::UnregisteredID(id) => format!("Unregistered id: {id} referred to in graph"),
BuildError::GraphCycle(id) => format!("Cycle in graph for id: {id}"),
BuildError::DisconnectedNode(id) => format!("id: {id} disconnected from graph"),
};
write!(f, "{str}")
}
}
impl Error for BuildError {}
fn validate_graph(graph: &U64Map<Vec<u64>>, roots: &U64Set) -> Result<Vec<Vec<u64>>, BuildError> {
let mut visiting = U64Set::default();
let mut visited = U64Set::default();
let mut chains = Vec::new();
for node in roots {
let result = visit(*node, &mut visiting, &mut visited, Vec::new(), graph)?;
chains.extend(result);
}
for node in graph.keys() {
if !visited.contains(node) {
return Err(BuildError::DisconnectedNode(*node));
}
}
Ok(chains)
}
fn visit(
node: u64,
visiting: &mut U64Set,
visited: &mut U64Set,
mut chain: Vec<u64>,
graph: &U64Map<Vec<u64>>,
) -> Result<Vec<Vec<u64>>, BuildError> {
if visiting.contains(&node) {
return Err(BuildError::GraphCycle(node));
}
visiting.insert(node);
chain.push(node);
let mut chains = Vec::new();
let children = graph.get(&node).ok_or(BuildError::UnregisteredID(node))?;
if children.is_empty() {
chains.push(chain);
} else {
for child in children {
let result = visit(*child, visiting, visited, chain.clone(), graph)?;
chains.extend(result);
}
}
visiting.remove(&node);
visited.insert(node);
Ok(chains)
}
fn validate_order(handles_map: &U64Map<Handle>, chains: Vec<Vec<u64>>) -> Result<(), BuildError> {
let producer_ids: U64Set = handles_map
.iter()
.filter(|(_, h)| matches!(h, Handle::Producer))
.map(|(id, _)| *id)
.collect();
if producer_ids.is_empty() {
return Ok(());
}
for chain in chains {
for id in &producer_ids {
if !chain.contains(id) {
return Err(BuildError::UnorderedProducer(chain, *id));
}
}
}
Ok(())
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum Follows {
LeadProducer,
Handles(Vec<u64>),
}
#[derive(Copy, Clone, Debug, Eq, PartialEq)]
pub enum Handle {
Producer,
Consumer,
}
#[derive(Debug)]
pub struct DisruptorHandles<E, W> {
lead: Option<Producer<E, W, true>>,
producers: U64Map<Producer<E, W, false>>,
consumers: U64Map<Consumer<E, W>>,
}
impl<E, W> DisruptorHandles<E, W> {
#[must_use = "Disruptor will stall if any handle is not used "]
pub fn take_lead(&mut self) -> Option<Producer<E, W, true>> {
self.lead.take()
}
#[must_use = "Disruptor will stall if any handle is not used"]
pub fn take_producer(&mut self, id: u64) -> Option<Producer<E, W, false>> {
self.producers.remove(&id)
}
#[must_use = "Disruptor will stall if any handle is not used"]
pub fn take_consumer(&mut self, id: u64) -> Option<Consumer<E, W>> {
self.consumers.remove(&id)
}
#[must_use = "Disruptor will stall if any handle is not used"]
pub fn drain_producers(&mut self) -> impl Iterator<Item = (u64, Producer<E, W, false>)> + '_ {
self.producers.drain()
}
#[must_use = "Disruptor will stall if any handle is not used"]
pub fn drain_consumers(&mut self) -> impl Iterator<Item = (u64, Consumer<E, W>)> + '_ {
self.consumers.drain()
}
pub fn is_empty(&self) -> bool {
self.lead.is_none() && self.consumers.is_empty() && self.producers.is_empty()
}
}
#[cfg(debug_assertions)]
impl<E, W> Drop for DisruptorHandles<E, W> {
fn drop(&mut self) {
if !self.is_empty() {
println!("WARNING: DisruptorHandles not empty when dropped");
}
}
}
#[derive(Copy, Clone, Debug, Default)]
struct HashU64(u64);
impl Hasher for HashU64 {
fn finish(&self) -> u64 {
self.0
}
fn write(&mut self, _: &[u8]) {
unreachable!("`write` should never be called");
}
fn write_u64(&mut self, i: u64) {
self.0 = i;
}
}
type U64Map<V> = HashMap<u64, V, BuildHasherDefault<HashU64>>;
type U64Set = HashSet<u64, BuildHasherDefault<HashU64>>;
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_find_cycle() {
let result = DisruptorBuilder::new(64, || 0)
.extend_handles([
(0, Handle::Consumer, Follows::LeadProducer),
(1, Handle::Consumer, Follows::Handles(vec![0, 3])),
(2, Handle::Consumer, Follows::Handles(vec![0])),
(3, Handle::Consumer, Follows::Handles(vec![1])),
(4, Handle::Consumer, Follows::Handles(vec![2])),
])
.build();
assert_eq!(result.err().unwrap(), BuildError::GraphCycle(1));
let result = DisruptorBuilder::new(64, || 0)
.extend_handles([
(0, Handle::Consumer, Follows::LeadProducer),
(1, Handle::Consumer, Follows::Handles(vec![0])),
(2, Handle::Consumer, Follows::Handles(vec![1, 4])),
(3, Handle::Consumer, Follows::Handles(vec![2])),
(4, Handle::Consumer, Follows::Handles(vec![2])),
])
.build();
assert_eq!(result.err().unwrap(), BuildError::GraphCycle(2));
}
#[test]
fn test_find_disconnected_branch() {
let graph = U64Map::from_iter([
(0, vec![1]),
(1, vec![2]),
(2, vec![3]),
(3, vec![]),
(4, vec![5]),
(5, vec![4]),
]);
let roots = U64Set::from_iter([0]);
assert_eq!(
validate_graph(&graph, &roots),
Err(BuildError::DisconnectedNode(4))
);
}
#[test]
fn test_find_unregistered_id() {
let graph = U64Map::from_iter([(0, vec![1])]);
let roots = U64Set::from_iter([0]);
assert_eq!(
validate_graph(&graph, &roots),
Err(BuildError::UnregisteredID(1))
);
let graph = U64Map::from_iter([(0, vec![1, 2]), (1, vec![3]), (2, vec![3])]);
assert_eq!(
validate_graph(&graph, &roots),
Err(BuildError::UnregisteredID(3))
);
}
#[test]
fn test_find_cycle_self_follow() {
let graph = U64Map::from_iter([(0, vec![1]), (1, vec![1])]);
let roots = U64Set::from_iter([0]);
assert_eq!(
validate_graph(&graph, &roots),
Err(BuildError::GraphCycle(1))
);
}
#[test]
fn test_find_chains() {
let graph = U64Map::from_iter([
(0, vec![1, 2]),
(1, vec![3]),
(2, vec![4]),
(3, vec![4]),
(4, vec![]),
]);
let roots = U64Set::from_iter([0]);
let chains = vec![vec![0, 1, 3, 4], vec![0, 2, 4]];
assert_eq!(validate_graph(&graph, &roots), Ok(chains));
}
#[test]
fn test_find_chains_multiple_roots() {
let graph = U64Map::from_iter([
(0, vec![3]),
(1, vec![4]),
(2, vec![5]),
(3, vec![6]),
(4, vec![6]),
(5, vec![6]),
(6, vec![]),
]);
let roots = U64Set::from_iter([0, 1, 2]);
let chains = vec![vec![0, 3, 6], vec![1, 4, 6], vec![2, 5, 6]];
assert_eq!(validate_graph(&graph, &roots), Ok(chains));
}
#[test]
fn test_find_chains_fan_out_in_uneven() {
let graph = U64Map::from_iter([
(0, vec![1, 4]),
(1, vec![2]),
(2, vec![3]),
(4, vec![5]),
(5, vec![6]),
(6, vec![3]),
(3, vec![7]),
(7, vec![]),
]);
let roots = U64Set::from_iter([0]);
let chains = vec![vec![0, 1, 2, 3, 7], vec![0, 4, 5, 6, 3, 7]];
assert_eq!(validate_graph(&graph, &roots), Ok(chains));
}
#[test]
fn test_producer_partial_order() {
let handles_map = U64Map::from_iter([
(0, Handle::Consumer),
(1, Handle::Consumer),
(2, Handle::Consumer),
(3, Handle::Consumer),
(4, Handle::Producer),
(5, Handle::Consumer),
(6, Handle::Consumer),
(7, Handle::Consumer),
]);
let chains = vec![vec![0, 1, 2, 3, 7], vec![0, 4, 5, 6, 3, 7]];
let result = validate_order(&handles_map, chains);
assert_eq!(
result,
Err(BuildError::UnorderedProducer(vec![0, 1, 2, 3, 7], 4))
)
}
#[test]
fn test_producer_total_order() {
let handles_map = U64Map::from_iter([
(0, Handle::Consumer),
(1, Handle::Consumer),
(2, Handle::Consumer),
(3, Handle::Producer),
(4, Handle::Consumer),
(5, Handle::Consumer),
(6, Handle::Consumer),
(7, Handle::Consumer),
]);
let chains = vec![vec![0, 1, 2, 3, 7], vec![0, 4, 5, 6, 3, 7]];
let result = validate_order(&handles_map, chains);
assert_eq!(result, Ok(()))
}
#[test]
fn test_producer_total_order2() {
let handles_map = U64Map::from_iter([
(0, Handle::Producer),
(1, Handle::Consumer),
(2, Handle::Consumer),
(3, Handle::Consumer),
(4, Handle::Consumer),
(5, Handle::Consumer),
(6, Handle::Consumer),
(7, Handle::Consumer),
]);
let chains = vec![vec![0, 1, 2, 3, 7], vec![0, 4, 5, 6, 3, 7]];
let result = validate_order(&handles_map, chains);
assert_eq!(result, Ok(()))
}
#[test]
fn test_producer_total_order3() {
let handles_map = U64Map::from_iter([
(0, Handle::Consumer),
(1, Handle::Consumer),
(2, Handle::Consumer),
(3, Handle::Producer),
(4, Handle::Consumer),
(5, Handle::Consumer),
(6, Handle::Producer),
]);
let chains = vec![
vec![0, 1, 3, 4, 6],
vec![0, 1, 3, 5, 6],
vec![0, 2, 3, 4, 6],
vec![0, 2, 3, 5, 6],
];
let result = validate_order(&handles_map, chains);
assert_eq!(result, Ok(()))
}
#[test]
fn test_producer_partial_order2() {
let handles_map = U64Map::from_iter([
(0, Handle::Consumer),
(1, Handle::Consumer),
(2, Handle::Consumer),
(3, Handle::Consumer),
(4, Handle::Producer),
(5, Handle::Consumer),
]);
let chains = vec![vec![0, 1, 2, 3, 4], vec![0, 5]];
let result = validate_order(&handles_map, chains);
assert_eq!(result, Err(BuildError::UnorderedProducer(vec![0, 5], 4)))
}
#[test]
fn test_producer_total_order4() {
let handles_map = U64Map::from_iter([
(0, Handle::Consumer),
(1, Handle::Consumer),
(2, Handle::Consumer),
(3, Handle::Consumer),
(4, Handle::Consumer),
(5, Handle::Consumer),
]);
let chains = vec![vec![0, 1, 2, 3, 4], vec![0, 5]];
let result = validate_order(&handles_map, chains);
assert_eq!(result, Ok(()))
}
#[test]
fn test_builder_empty_follows_disconnected_err() {
let result = DisruptorBuilder::new(32, || 0)
.add_handle(0, Handle::Consumer, Follows::LeadProducer)
.add_handle(1, Handle::Consumer, Follows::Handles(vec![]))
.build();
assert_eq!(result.err().unwrap(), BuildError::DisconnectedNode(1))
}
#[test]
fn test_builder_buffer_size_error() {
let result = DisruptorBuilder::new(0, || 0)
.add_handle(0, Handle::Consumer, Follows::LeadProducer)
.build();
assert_eq!(result.err().unwrap(), BuildError::BufferSize(0));
let result2 = DisruptorBuilder::new(12, || 0)
.add_handle(0, Handle::Consumer, Follows::LeadProducer)
.build();
assert_eq!(result2.err().unwrap(), BuildError::BufferSize(12))
}
#[test]
fn test_builder_overlapping_ids() {
let result = DisruptorBuilder::new(32, || 0)
.add_handle(0, Handle::Consumer, Follows::LeadProducer)
.add_handle(1, Handle::Consumer, Follows::LeadProducer)
.add_handle(1, Handle::Consumer, Follows::Handles(vec![0]))
.build();
assert_eq!(result.err().unwrap(), BuildError::OverlappingIDs(vec![1]))
}
#[test]
fn test_builder_empty_graph() {
let result = DisruptorBuilder::new(32, || 0).build();
assert_eq!(result.err().unwrap(), BuildError::EmptyGraph)
}
#[test]
fn test_builder_lead_not_followed() {
let result = DisruptorBuilder::new(32, || 0)
.add_handle(0, Handle::Consumer, Follows::Handles(vec![8]))
.add_handle(1, Handle::Consumer, Follows::Handles(vec![0]))
.build();
assert_eq!(result.err().unwrap(), BuildError::UnregisteredID(8))
}
}