use std::{collections::HashSet, sync::Arc};
use iroh::{
endpoint::Connection,
protocol::{AcceptError, ProtocolHandler, Router},
NodeAddr,
};
use tokio::sync::Mutex;
use super::{InboundRecorder, LocalEndpointFactory};
use crate::{
zakura::{Frame, ZakuraPeerId},
BoxError,
};
const GOSSIP_ALPN: &[u8] = b"/zakura/testkit/gossip/0";
const GOSSIP_MAX_FRAME: u32 = 64 * 1024;
const GOSSIP_STREAM_KIND: u16 = 2;
#[derive(Debug)]
struct GossipCore {
recorder: InboundRecorder,
seen: Mutex<HashSet<Vec<u8>>>,
conns: Mutex<Vec<Connection>>,
}
impl GossipCore {
fn new(recorder: InboundRecorder) -> Self {
Self {
recorder,
seen: Mutex::new(HashSet::new()),
conns: Mutex::new(Vec::new()),
}
}
async fn register_conn(&self, conn: Connection) {
self.conns.lock().await.push(conn);
}
async fn mark_new(&self, payload: &[u8]) -> bool {
self.seen.lock().await.insert(payload.to_vec())
}
fn record(&self, frame: &Frame) {
let _ = self.recorder.deliver(
ZakuraPeerId::new(vec![0; 32]).expect("test peer id is within bounds"),
GOSSIP_STREAM_KIND,
frame.clone(),
);
}
async fn broadcast(&self, frame: &Frame) {
let Ok(encoded) = frame.encode(GOSSIP_MAX_FRAME) else {
return;
};
let conns = self.conns.lock().await.clone();
for conn in conns {
if let Ok((mut send, _recv)) = conn.open_bi().await {
let _ = send.write_all(&encoded).await;
let _ = send.finish();
}
}
}
async fn on_received(self: &Arc<Self>, frame: Frame) {
if !self.mark_new(&frame.payload).await {
return;
}
self.record(&frame);
self.broadcast(&frame).await;
}
async fn serve(self: Arc<Self>, conn: Connection) {
while let Ok((_send, mut recv)) = conn.accept_bi().await {
let core = self.clone();
tokio::spawn(async move {
if let Ok(bytes) = recv.read_to_end(GOSSIP_MAX_FRAME as usize).await {
if let Ok(frame) = Frame::decode(&bytes, GOSSIP_MAX_FRAME) {
core.on_received(frame).await;
}
}
});
}
}
}
#[derive(Clone, Debug)]
struct GossipHandler {
core: Arc<GossipCore>,
}
impl ProtocolHandler for GossipHandler {
async fn accept(&self, connection: Connection) -> Result<(), AcceptError> {
self.core.register_conn(connection.clone()).await;
self.core.clone().serve(connection).await;
Ok(())
}
}
#[derive(Debug)]
pub struct GossipNode {
router: Router,
core: Arc<GossipCore>,
}
impl GossipNode {
pub async fn spawn(seed: u64) -> Result<Self, BoxError> {
let endpoint = LocalEndpointFactory::new().endpoint(seed).await?;
let core = Arc::new(GossipCore::new(InboundRecorder::new(1024)));
let router = Router::builder(endpoint)
.accept(GOSSIP_ALPN, GossipHandler { core: core.clone() })
.spawn();
Ok(Self { router, core })
}
pub async fn node_addr(&self) -> NodeAddr {
LocalEndpointFactory::node_addr(self.router.endpoint()).await
}
pub fn recorder(&self) -> InboundRecorder {
self.core.recorder.clone()
}
pub async fn connect(&self, peer: &GossipNode) -> Result<(), BoxError> {
let peer_addr = peer.node_addr().await;
let endpoint = self.router.endpoint();
endpoint.add_node_addr(peer_addr.clone())?;
let conn = endpoint.connect(peer_addr, GOSSIP_ALPN).await?;
self.core.register_conn(conn.clone()).await;
let core = self.core.clone();
tokio::spawn(async move { core.serve(conn).await });
Ok(())
}
pub async fn broadcast(&self, payload: Vec<u8>) -> Result<(), BoxError> {
let frame = Frame {
message_type: 1,
flags: 0,
payload: payload.clone(),
};
self.core.mark_new(&payload).await;
self.core.record(&frame);
self.core.broadcast(&frame).await;
Ok(())
}
pub async fn shutdown(&self) {
let _ = self.router.shutdown().await;
}
}
#[cfg(test)]
mod tests {
use super::super::{await_until, TEST_NET_TIMEOUT};
use super::*;
#[tokio::test]
async fn gossip_floods_to_every_node_across_a_line() -> Result<(), BoxError> {
let _guard = zakura_test::init();
let mut nodes = Vec::new();
for seed in 1..=5u64 {
nodes.push(GossipNode::spawn(seed).await?);
}
for pair in nodes.windows(2) {
pair[0].connect(&pair[1]).await?;
}
await_until("every node wired into the line", TEST_NET_TIMEOUT, || {
nodes.iter().enumerate().all(|(index, _)| {
let expected = if index == 0 || index == nodes.len() - 1 {
1
} else {
2
};
conn_count(&nodes[index]) >= expected
})
})
.await?;
let payload = b"flood-sub-hello".to_vec();
nodes[0].broadcast(payload.clone()).await?;
for (index, node) in nodes.iter().enumerate() {
let recorder = node.recorder();
let payload = payload.clone();
await_until(
format!("gossip reaches node {index}"),
TEST_NET_TIMEOUT,
|| recorder.contains_payload(GOSSIP_STREAM_KIND, &payload),
)
.await?;
}
for node in &nodes {
node.shutdown().await;
}
Ok(())
}
fn conn_count(node: &GossipNode) -> usize {
node.core
.conns
.try_lock()
.map(|conns| conns.len())
.unwrap_or(0)
}
}