use std::collections::HashSet;
use tokio::sync::broadcast::{
self,
error::{RecvError, TryRecvError},
};
use tower::{timeout::Timeout, Service, ServiceExt};
use zakura_network::MAX_TX_INV_IN_SENT_MESSAGE;
use zakura_chain::transaction::UnminedTxId;
use zakura_network as zn;
use zakura_node_services::mempool::{MempoolChange, Request, Response};
use crate::{
components::sync::{PEER_GOSSIP_DELAY, TIPS_RESPONSE_TIMEOUT},
BoxError,
};
pub const MAX_CHANGES_BEFORE_SEND: usize = 10;
const MAX_TX_INV_IN_SENT_MESSAGE_USIZE: usize = MAX_TX_INV_IN_SENT_MESSAGE as usize;
pub(super) const MEMPOOL_CHANGE_CHANNEL_CAPACITY: usize = MAX_CHANGES_BEFORE_SEND * 4;
pub(crate) async fn run_mempool_transaction_id_gossip<ZN, ZM>(
mut receiver: broadcast::Receiver<MempoolChange>,
broadcast_network: ZN,
mut mempool: ZM,
) -> Result<(), BoxError>
where
ZN: Service<zn::Request, Response = zn::Response, Error = BoxError> + Send + Clone + 'static,
ZN::Future: Send,
ZM: Service<Request, Response = Response, Error = BoxError> + Send + Clone + 'static,
ZM::Future: Send + 'static,
{
info!("initializing transaction gossip task");
let mut broadcast_network = Timeout::new(broadcast_network, TIPS_RESPONSE_TIMEOUT);
let mut drain_pending_without_wakeup = false;
loop {
let mut combined_changes = 0;
if !drain_pending_without_wakeup {
combined_changes = 1;
loop {
match receiver.recv().await {
Ok(mempool_change) if mempool_change.is_added() => break,
Ok(_) => {
continue;
}
Err(RecvError::Lagged(skip_count)) => {
info!(
?skip_count,
"dropped mempool changes before gossiping, draining pending transaction IDs"
);
metrics::counter!("mempool.gossip.lagged.events.total")
.increment(skip_count);
break;
}
Err(closed @ RecvError::Closed) => Err(closed)?,
}
}
while combined_changes <= MAX_CHANGES_BEFORE_SEND {
match receiver.try_recv() {
Ok(mempool_change) if mempool_change.is_added() => {}
Ok(_) => {
continue;
}
Err(TryRecvError::Empty) => break,
Err(TryRecvError::Lagged(skip_count)) => {
info!(
?skip_count,
"dropped mempool changes before gossiping, draining pending transaction IDs"
);
metrics::counter!("mempool.gossip.lagged.events.total")
.increment(skip_count);
}
Err(closed @ TryRecvError::Closed) => Err(closed)?,
}
combined_changes += 1;
}
} else {
drain_pending_without_wakeup = false;
}
let advertised_count = advertise_pending_mempool_transaction_ids(
&mut mempool,
&mut broadcast_network,
combined_changes,
)
.await?;
if advertised_count == 0 {
continue;
}
drain_pending_without_wakeup = advertised_count == MAX_TX_INV_IN_SENT_MESSAGE;
tokio::time::sleep(PEER_GOSSIP_DELAY).await;
}
}
async fn advertise_pending_mempool_transaction_ids<ZN, ZM>(
mempool: &mut ZM,
broadcast_network: &mut Timeout<ZN>,
combined_changes: usize,
) -> Result<u64, BoxError>
where
ZN: Service<zn::Request, Response = zn::Response, Error = BoxError> + Send + Clone + 'static,
ZN::Future: Send,
ZM: Service<Request, Response = Response, Error = BoxError> + Send + Clone + 'static,
ZM::Future: Send + 'static,
{
let Response::TransactionIds(tx_ids) = mempool
.ready()
.await?
.call(Request::TakePendingGossipTransactionIds {
limit: MAX_TX_INV_IN_SENT_MESSAGE_USIZE,
})
.await?
else {
return Err(std::io::Error::other(
"mempool pending gossip request returned a different response variant",
)
.into());
};
let mut advertised_count = 0;
let mut chunk = HashSet::<UnminedTxId>::new();
for tx_id in tx_ids {
chunk.insert(tx_id);
if chunk.len() >= MAX_TX_INV_IN_SENT_MESSAGE_USIZE {
advertised_count +=
advertise_transaction_id_chunk(broadcast_network, &mut chunk, combined_changes)
.await?;
}
}
if !chunk.is_empty() {
advertised_count +=
advertise_transaction_id_chunk(broadcast_network, &mut chunk, combined_changes).await?;
}
if advertised_count > 0 {
metrics::counter!("mempool.gossip.pending.transactions.total").increment(advertised_count);
metrics::counter!("mempool.gossiped.transactions.total").increment(advertised_count);
}
Ok(advertised_count)
}
async fn advertise_transaction_id_chunk<ZN>(
broadcast_network: &mut Timeout<ZN>,
chunk: &mut HashSet<UnminedTxId>,
combined_changes: usize,
) -> Result<u64, BoxError>
where
ZN: Service<zn::Request, Response = zn::Response, Error = BoxError> + Send + Clone + 'static,
ZN::Future: Send,
{
let txs_len: u64 = chunk
.len()
.try_into()
.expect("transaction ID chunk length fits in u64");
let request = zn::Request::AdvertiseTransactionIds(std::mem::take(chunk), None);
info!(%request, changes = %combined_changes, "sending pending mempool transaction broadcast");
debug!(
?request,
changes = ?combined_changes,
"full list of pending mempool transactions in broadcast"
);
let _ = broadcast_network.ready().await?.call(request).await;
Ok(txs_len)
}
#[cfg(test)]
mod tests {
use std::{
collections::{HashSet, VecDeque},
sync::{Arc, Mutex},
time::Duration,
};
use tokio::sync::{broadcast, mpsc};
use tower::service_fn;
use zakura_chain::transaction;
use super::*;
fn test_tx_ids(count: usize, seed: u8) -> HashSet<UnminedTxId> {
(0..count)
.map(|index| {
let index: u64 = index
.try_into()
.expect("test transaction ID index fits in u64");
let mut bytes = [seed; 32];
bytes[..8].copy_from_slice(&index.to_le_bytes());
UnminedTxId::Legacy(transaction::Hash(bytes))
})
.collect()
}
fn mempool_service(
pending_batches: Vec<HashSet<UnminedTxId>>,
) -> (
impl Service<Request, Response = Response, Error = BoxError, Future: Send> + Clone,
mpsc::Receiver<usize>,
) {
let pending_batches = Arc::new(Mutex::new(VecDeque::from(pending_batches)));
let (limit_sender, limit_receiver) = mpsc::channel(16);
let service = service_fn(move |request| {
let pending_batches = pending_batches.clone();
let limit_sender = limit_sender.clone();
async move {
match request {
Request::TakePendingGossipTransactionIds { limit } => {
assert_eq!(
limit, MAX_TX_INV_IN_SENT_MESSAGE_USIZE,
"gossip task should bound each pending drain to one inv"
);
limit_sender
.send(limit)
.await
.expect("limit receiver should be open");
let tx_ids = pending_batches
.lock()
.expect("pending batch mutex should not be poisoned")
.pop_front()
.unwrap_or_default();
Ok(Response::TransactionIds(tx_ids))
}
unexpected_request => {
panic!("unexpected mempool request: {unexpected_request:?}")
}
}
}
});
(service, limit_receiver)
}
fn peer_set_service() -> (
impl Service<zn::Request, Response = zn::Response, Error = BoxError, Future: Send> + Clone,
mpsc::Receiver<zn::Request>,
) {
let (advertised_sender, advertised_receiver) = mpsc::channel(16);
let service = service_fn(move |request| {
let advertised_sender = advertised_sender.clone();
async move {
advertised_sender
.send(request)
.await
.expect("advertised request receiver should be open");
Ok(zn::Response::Nil)
}
});
(service, advertised_receiver)
}
async fn expect_advertised_transaction_ids(
advertised_receiver: &mut mpsc::Receiver<zn::Request>,
) -> HashSet<UnminedTxId> {
let advertised_request =
tokio::time::timeout(Duration::from_secs(1), advertised_receiver.recv())
.await
.expect("gossip task should advertise pending mempool txids")
.expect("peer set should advertise a request before the task exits");
let zn::Request::AdvertiseTransactionIds(advertised_tx_ids, None) = advertised_request
else {
panic!("unexpected advertised request: {advertised_request:?}");
};
advertised_tx_ids
}
#[tokio::test]
async fn added_mempool_gossip_drains_pending_transaction_ids() {
let _init_guard = zakura_test::init();
let pending_tx_ids = test_tx_ids(2, 1);
let (mempool, mut limit_receiver) = mempool_service(vec![pending_tx_ids.clone()]);
let (peer_set, mut advertised_receiver) = peer_set_service();
let (sender, receiver) = broadcast::channel(MEMPOOL_CHANGE_CHANNEL_CAPACITY);
sender
.send(MempoolChange::added(test_tx_ids(1, 2)))
.expect("receiver should be subscribed");
let gossip_task = tokio::spawn(run_mempool_transaction_id_gossip(
receiver, peer_set, mempool,
));
assert_eq!(
limit_receiver
.recv()
.await
.expect("gossip task should request pending txids"),
MAX_TX_INV_IN_SENT_MESSAGE_USIZE
);
assert_eq!(
expect_advertised_transaction_ids(&mut advertised_receiver).await,
pending_tx_ids,
"happy path should advertise the pending mempool txids",
);
gossip_task.abort();
}
#[tokio::test]
async fn lagged_mempool_gossip_drains_pending_transaction_ids() {
let _init_guard = zakura_test::init();
let pending_tx_ids = test_tx_ids(2, 1);
let dropped_tx_ids = test_tx_ids(2, 2);
let (mempool, mut limit_receiver) = mempool_service(vec![pending_tx_ids.clone()]);
let (peer_set, mut advertised_receiver) = peer_set_service();
let (sender, receiver) = broadcast::channel(1);
let mut lagged_events = dropped_tx_ids
.into_iter()
.map(|tx_id| MempoolChange::added([tx_id].into_iter().collect()));
sender
.send(
lagged_events
.next()
.expect("first lagged mempool change should exist"),
)
.expect("receiver should be subscribed");
sender
.send(
lagged_events
.next()
.expect("second lagged mempool change should exist"),
)
.expect("receiver should be subscribed");
let gossip_task = tokio::spawn(run_mempool_transaction_id_gossip(
receiver, peer_set, mempool,
));
assert_eq!(
limit_receiver
.recv()
.await
.expect("gossip task should request pending txids"),
MAX_TX_INV_IN_SENT_MESSAGE_USIZE
);
assert_eq!(
expect_advertised_transaction_ids(&mut advertised_receiver).await,
pending_tx_ids,
"lag recovery should advertise pending txids, not dropped channel payloads",
);
gossip_task.abort();
}
#[tokio::test(start_paused = true)]
async fn lagged_mempool_gossip_recovers_pending_transaction_ids_in_bounded_cycles() {
let _init_guard = zakura_test::init();
let first_batch = test_tx_ids(MAX_TX_INV_IN_SENT_MESSAGE_USIZE, 1);
let second_batch = test_tx_ids(2, 2);
let (mempool, mut limit_receiver) =
mempool_service(vec![first_batch.clone(), second_batch.clone()]);
let (peer_set, mut advertised_receiver) = peer_set_service();
let (sender, receiver) = broadcast::channel(1);
sender
.send(MempoolChange::added(
[UnminedTxId::Legacy(transaction::Hash([42; 32]))]
.into_iter()
.collect(),
))
.expect("receiver should be subscribed");
sender
.send(MempoolChange::added(
[UnminedTxId::Legacy(transaction::Hash([43; 32]))]
.into_iter()
.collect(),
))
.expect("receiver should be subscribed");
let gossip_task = tokio::spawn(run_mempool_transaction_id_gossip(
receiver, peer_set, mempool,
));
assert_eq!(
limit_receiver
.recv()
.await
.expect("first drain should request pending txids"),
MAX_TX_INV_IN_SENT_MESSAGE_USIZE
);
let advertised_tx_ids = expect_advertised_transaction_ids(&mut advertised_receiver).await;
assert_eq!(advertised_tx_ids, first_batch);
assert_eq!(
advertised_tx_ids.len(),
MAX_TX_INV_IN_SENT_MESSAGE_USIZE,
"first recovery cycle should be bounded to one inv-sized batch",
);
tokio::time::advance(PEER_GOSSIP_DELAY).await;
assert_eq!(
limit_receiver
.recv()
.await
.expect("second drain should happen without another wakeup"),
MAX_TX_INV_IN_SENT_MESSAGE_USIZE
);
let advertised_tx_ids = expect_advertised_transaction_ids(&mut advertised_receiver).await;
assert_eq!(advertised_tx_ids, second_batch);
assert!(
advertised_tx_ids.len() < MAX_TX_INV_IN_SENT_MESSAGE_USIZE,
"second recovery cycle should only advertise the remaining txids",
);
gossip_task.abort();
}
#[tokio::test]
async fn empty_pending_mempool_gossip_wakeup_does_not_advertise() {
let _init_guard = zakura_test::init();
let (mempool, mut limit_receiver) = mempool_service(vec![HashSet::new()]);
let (peer_set, mut advertised_receiver) = peer_set_service();
let (sender, receiver) = broadcast::channel(MEMPOOL_CHANGE_CHANNEL_CAPACITY);
sender
.send(MempoolChange::added(test_tx_ids(1, 1)))
.expect("receiver should be subscribed");
let gossip_task = tokio::spawn(run_mempool_transaction_id_gossip(
receiver, peer_set, mempool,
));
assert_eq!(
limit_receiver
.recv()
.await
.expect("gossip task should request pending txids"),
MAX_TX_INV_IN_SENT_MESSAGE_USIZE
);
assert!(
tokio::time::timeout(Duration::from_millis(50), advertised_receiver.recv())
.await
.is_err(),
"empty pending gossip wakeups should not advertise to peers",
);
assert!(
!gossip_task.is_finished(),
"gossip task should remain alive after an empty pending drain",
);
gossip_task.abort();
}
}