use std::{
collections::{BTreeMap, VecDeque},
mem,
sync::mpsc::{self, Receiver, RecvTimeoutError, TryRecvError},
thread::{self, JoinHandle},
time::Instant,
};
use ed25519_dalek::VerifyingKey;
use crate::{
block_sync::messages::{BlockSyncMessage, BlockSyncRequest, BlockSyncResponse},
types::data_types::{BufferSize, ChainID, ViewNumber},
};
use super::{
messages::{Message, ProgressMessage},
network::Network,
};
pub(crate) fn start_polling<N: Network + 'static>(
mut network: N,
shutdown_signal: Receiver<()>,
) -> (
JoinHandle<()>,
Receiver<(VerifyingKey, ProgressMessage)>,
Receiver<(VerifyingKey, BlockSyncRequest)>,
Receiver<(VerifyingKey, BlockSyncResponse)>,
) {
let (to_progress_msg_receiver, progress_msg_receiver) = mpsc::channel();
let (to_sync_request_receiver, sync_request_receiver) = mpsc::channel();
let (to_sync_response_receiver, sync_response_receiver) = mpsc::channel();
let poller_thread = thread::spawn(move || loop {
match shutdown_signal.try_recv() {
Ok(()) => return,
Err(TryRecvError::Empty) => (),
Err(TryRecvError::Disconnected) => {
panic!("Poller thread disconnected from main thread")
}
}
if let Some((origin, msg)) = network.recv() {
match msg {
Message::ProgressMessage(p_msg) => {
let _ = to_progress_msg_receiver.send((origin, p_msg));
}
Message::BlockSyncMessage(s_msg) => match s_msg {
BlockSyncMessage::BlockSyncRequest(s_req) => {
let _ = to_sync_request_receiver.send((origin, s_req));
}
BlockSyncMessage::BlockSyncResponse(s_res) => {
let _ = to_sync_response_receiver.send((origin, s_res));
}
},
}
} else {
thread::yield_now()
}
});
(
poller_thread,
progress_msg_receiver,
sync_request_receiver,
sync_response_receiver,
)
}
pub(crate) struct ProgressMessageStub {
receiver: Receiver<(VerifyingKey, ProgressMessage)>,
msg_buffer: ProgressMessageBuffer,
}
impl ProgressMessageStub {
pub(crate) fn new(
receiver: Receiver<(VerifyingKey, ProgressMessage)>,
msg_buffer_capacity: BufferSize,
) -> ProgressMessageStub {
let msg_buffer: ProgressMessageBuffer = ProgressMessageBuffer::new(msg_buffer_capacity);
Self {
receiver,
msg_buffer,
}
}
pub(crate) fn recv(
&mut self,
chain_id: ChainID,
cur_view: ViewNumber,
deadline: Instant,
) -> Result<(VerifyingKey, ProgressMessage), ProgressMessageReceiveError> {
self.msg_buffer.remove_expired_msgs(cur_view);
if let Some((sender, msg)) = self.msg_buffer.get_msg(&cur_view) {
return Ok((sender, msg));
}
while Instant::now() < deadline {
match self.receiver.recv_timeout(deadline - Instant::now()) {
Ok((sender, msg)) => {
if msg.chain_id() != chain_id {
continue;
}
if msg.view().is_some_and(|view| view > cur_view) {
match msg.clone() {
ProgressMessage::HotStuffMessage(msg) => {
self.msg_buffer.insert(msg, sender);
}
ProgressMessage::PacemakerMessage(msg) => {
self.msg_buffer.insert(msg, sender);
}
ProgressMessage::BlockSyncAdvertiseMessage(_) => (),
}
}
let return_msg = match &msg {
ProgressMessage::HotStuffMessage(hotstuff_msg) => {
hotstuff_msg.view() == cur_view
}
ProgressMessage::PacemakerMessage(pacemaker_msg) => {
pacemaker_msg.view() >= cur_view
}
ProgressMessage::BlockSyncAdvertiseMessage(_) => true,
};
if return_msg {
return Ok((sender, msg));
}
}
Err(RecvTimeoutError::Timeout) => thread::yield_now(),
Err(RecvTimeoutError::Disconnected) => {
return Err(ProgressMessageReceiveError::Disconnected)
}
}
}
Err(ProgressMessageReceiveError::Timeout)
}
}
#[derive(Debug)]
pub(crate) enum ProgressMessageReceiveError {
Timeout,
Disconnected,
}
struct ProgressMessageBuffer {
buffer_capacity: BufferSize,
buffer: BTreeMap<ViewNumber, VecDeque<(VerifyingKey, ProgressMessage)>>,
buffer_size: BufferSize,
}
impl ProgressMessageBuffer {
fn new(buffer_capacity: BufferSize) -> Self {
Self {
buffer_capacity,
buffer: BTreeMap::new(),
buffer_size: BufferSize::new(0),
}
}
fn insert<M: Into<ProgressMessage> + Cacheable>(
&mut self,
msg: M,
sender: VerifyingKey,
) -> bool {
let bytes_requested = mem::size_of::<VerifyingKey>() as u64 + msg.size();
let new_buffer_size = self.buffer_size.int().checked_add(bytes_requested);
let buffer_will_be_overloaded =
new_buffer_size.is_none() || new_buffer_size.unwrap() > self.buffer_capacity.int();
let cache_message_if_buffer_will_be_overloaded = self.buffer.keys().max().is_none()
|| self
.buffer
.keys()
.max()
.is_some_and(|max_view| msg.view() < *max_view);
if buffer_will_be_overloaded && cache_message_if_buffer_will_be_overloaded {
self.remove_highest_viewed_msgs(bytes_requested);
};
if !buffer_will_be_overloaded
|| (buffer_will_be_overloaded && cache_message_if_buffer_will_be_overloaded)
{
let msg_queue = if let Some(msg_queue) = self.buffer.get_mut(&msg.view()) {
msg_queue
} else {
self.buffer.insert(msg.view(), VecDeque::new());
self.buffer.get_mut(&msg.view()).unwrap()
};
self.buffer_size += bytes_requested;
msg_queue.push_back((sender, msg.into()));
return true;
};
false
}
fn get_msg(&mut self, view: &ViewNumber) -> Option<(VerifyingKey, ProgressMessage)> {
self.buffer
.get_mut(view)
.map(|msg_queue| msg_queue.pop_front())
.flatten()
}
fn remove_highest_viewed_msgs(&mut self, bytes_to_remove: u64) {
let verifying_key_size = mem::size_of::<VerifyingKey>() as u64;
let mut bytes_removed = 0;
let mut views_removed = Vec::new();
let mut msg_queues_iter = self.buffer.iter_mut().rev();
while bytes_removed < bytes_to_remove {
if let Some((view, msg_queue)) = msg_queues_iter.next() {
let removals = msg_queue
.iter()
.rev()
.take_while(|(_, msg)| {
if bytes_removed < bytes_to_remove {
bytes_removed += msg.size() + verifying_key_size;
true
} else {
false
}
})
.count() as u64;
let _ = (0..removals).into_iter().for_each(|_| {
let _ = msg_queue.pop_back();
});
if msg_queue.is_empty() {
views_removed.push(*view)
}
} else {
break;
}
}
self.buffer_size -= bytes_removed;
views_removed.iter().for_each(|view| {
let _ = self.buffer.remove(view);
});
}
fn remove_expired_msgs(&mut self, cur_view: ViewNumber) {
self.buffer = self.buffer.split_off(&cur_view)
}
}
pub(crate) trait Cacheable {
fn view(&self) -> ViewNumber;
fn size(&self) -> u64;
}
pub(crate) struct BlockSyncClientStub {
responses: Receiver<(VerifyingKey, BlockSyncResponse)>,
}
impl BlockSyncClientStub {
pub(crate) fn new(
responses: Receiver<(VerifyingKey, BlockSyncResponse)>,
) -> BlockSyncClientStub {
BlockSyncClientStub { responses }
}
pub(crate) fn recv_response(
&self,
peer: VerifyingKey,
deadline: Instant,
) -> Result<BlockSyncResponse, BlockSyncResponseReceiveError> {
while Instant::now() < deadline {
match self.responses.recv_timeout(deadline - Instant::now()) {
Ok((sender, sync_response)) => {
if sender == peer {
return Ok(sync_response);
}
}
Err(RecvTimeoutError::Timeout) => thread::yield_now(),
Err(RecvTimeoutError::Disconnected) => {
return Err(BlockSyncResponseReceiveError::Disconnected)
}
}
}
Err(BlockSyncResponseReceiveError::Timeout)
}
}
#[derive(Debug)]
pub enum BlockSyncResponseReceiveError {
Disconnected,
Timeout,
}
pub(crate) struct BlockSyncServerStub {
requests: Receiver<(VerifyingKey, BlockSyncRequest)>,
}
impl BlockSyncServerStub {
pub(crate) fn new(requests: Receiver<(VerifyingKey, BlockSyncRequest)>) -> BlockSyncServerStub {
BlockSyncServerStub { requests }
}
pub(crate) fn recv_request(
&self,
) -> Result<(VerifyingKey, BlockSyncRequest), BlockSyncRequestReceiveError> {
match self.requests.try_recv() {
Ok((origin, request)) => Ok((origin, request)),
Err(TryRecvError::Disconnected) => Err(BlockSyncRequestReceiveError::Disconnected),
Err(TryRecvError::Empty) => Err(BlockSyncRequestReceiveError::NotAvailable),
}
}
}
#[derive(Debug)]
pub enum BlockSyncRequestReceiveError {
Disconnected,
NotAvailable,
}