use std::{
collections::HashMap,
sync::{Arc, Mutex as StdMutex},
};
use tokio::{sync::mpsc, task};
use tokio_util::sync::CancellationToken;
use super::{events::*, pipe::*, wire::*, *};
use crate::zakura::{
handle_pipe_exit, spawn_supervised_pipe, BoxRunFuture, Flow, Frame, FramedRecv, FramedSend,
OrderedSendError, Peer, PeerStreamSession, Pipe, Service, ServicePeerDirection, SessionGuard,
Sink, SinkReject, Stream, StreamMode, ZakuraConnId, ZakuraPeerId, ZakuraSupervisorHandle,
ZAKURA_CAP_HEADER_SYNC,
};
const HEADER_SYNC_SERVICE_STREAMS: [Stream; 1] = [Stream {
kind: ZAKURA_STREAM_HEADER_SYNC,
version: ZAKURA_HEADER_SYNC_STREAM_VERSION,
frame_cap: (MAX_HS_MESSAGE_BYTES + FRAME_HEADER_BYTES) as u32,
capability: ZAKURA_CAP_HEADER_SYNC,
mode: StreamMode::Ordered,
}];
pub(crate) fn header_sync_streams() -> &'static [Stream] {
&HEADER_SYNC_SERVICE_STREAMS
}
#[derive(Clone, Debug)]
pub struct HeaderSyncPeerSession {
peer_id: ZakuraPeerId,
direction: ServicePeerDirection,
inner: Arc<HeaderSyncPeerSessionInner>,
}
#[derive(Debug)]
struct HeaderSyncPeerSessionInner {
send: FramedSend,
cancel_token: CancellationToken,
commands: Option<mpsc::UnboundedSender<HeaderSyncPeerCommand>>,
}
impl HeaderSyncPeerSession {
fn new_with_commands(
session: &PeerStreamSession,
direction: ServicePeerDirection,
commands: mpsc::UnboundedSender<HeaderSyncPeerCommand>,
) -> Self {
Self::from_parts_with_direction_and_commands(
session.peer_id().clone(),
direction,
session.sender(),
session.cancel_token(),
Some(commands),
)
}
#[cfg(test)]
pub(crate) fn from_parts(
peer_id: ZakuraPeerId,
send: FramedSend,
cancel_token: CancellationToken,
) -> Self {
Self::from_parts_with_direction(peer_id, ServicePeerDirection::Inbound, send, cancel_token)
}
#[cfg(test)]
pub(crate) fn from_parts_with_direction(
peer_id: ZakuraPeerId,
direction: ServicePeerDirection,
send: FramedSend,
cancel_token: CancellationToken,
) -> Self {
Self::from_parts_with_direction_and_commands(peer_id, direction, send, cancel_token, None)
}
#[cfg(test)]
fn from_parts_with_direction_and_commands(
peer_id: ZakuraPeerId,
direction: ServicePeerDirection,
send: FramedSend,
cancel_token: CancellationToken,
commands: Option<mpsc::UnboundedSender<HeaderSyncPeerCommand>>,
) -> Self {
Self {
peer_id,
direction,
inner: Arc::new(HeaderSyncPeerSessionInner {
send,
cancel_token,
commands,
}),
}
}
#[cfg(not(test))]
fn from_parts_with_direction_and_commands(
peer_id: ZakuraPeerId,
direction: ServicePeerDirection,
send: FramedSend,
cancel_token: CancellationToken,
commands: Option<mpsc::UnboundedSender<HeaderSyncPeerCommand>>,
) -> Self {
Self {
peer_id,
direction,
inner: Arc::new(HeaderSyncPeerSessionInner {
send,
cancel_token,
commands,
}),
}
}
pub fn peer_id(&self) -> &ZakuraPeerId {
&self.peer_id
}
pub fn direction(&self) -> ServicePeerDirection {
self.direction
}
pub fn cancel_token(&self) -> CancellationToken {
self.inner.cancel_token.clone()
}
pub fn outbound_capacity(&self) -> usize {
self.inner.send.capacity()
}
pub fn outbound_max_capacity(&self) -> usize {
self.inner.send.max_capacity()
}
pub fn try_send_status(&self, status: HeaderSyncStatus) -> Result<(), OrderedSendError> {
self.try_send_message(HeaderSyncMessage::Status(status))
}
pub fn try_send_get_headers(
&self,
start_height: block::Height,
count: u32,
want_tree_aux_roots: bool,
) -> Result<(), OrderedSendError> {
let expected = ExpectedHeadersResponse::new(start_height, count, want_tree_aux_roots)
.map_err(|error| OrderedSendError::Encode(Box::new(error)))?;
if let Some(commands) = &self.inner.commands {
self.try_send_message(HeaderSyncMessage::GetHeaders {
start_height,
count,
want_tree_aux_roots,
})?;
return commands
.send(HeaderSyncPeerCommand::RecordExpectedHeaders(expected))
.map_err(|_| OrderedSendError::Closed);
}
self.try_send_message(HeaderSyncMessage::GetHeaders {
start_height,
count,
want_tree_aux_roots,
})
}
pub fn try_send_headers(
&self,
headers: Vec<Arc<block::Header>>,
) -> Result<(), OrderedSendError> {
let body_sizes = vec![0; headers.len()];
let tree_aux_roots = Vec::new();
self.try_send_headers_with_sizes_and_roots(headers, body_sizes, tree_aux_roots)
}
pub fn try_send_headers_with_sizes_and_roots(
&self,
headers: Vec<Arc<block::Header>>,
body_sizes: Vec<u32>,
tree_aux_roots: Vec<BlockCommitmentRoots>,
) -> Result<(), OrderedSendError> {
self.try_send_message(HeaderSyncMessage::Headers {
headers,
body_sizes,
tree_aux_roots,
})
}
pub fn try_send_new_block(&self, block: Arc<block::Block>) -> Result<(), OrderedSendError> {
self.try_send_message(HeaderSyncMessage::NewBlock(block))
}
fn try_send_message(&self, msg: HeaderSyncMessage) -> Result<(), OrderedSendError> {
let frame = msg
.encode_frame()
.map_err(|error| OrderedSendError::Encode(Box::new(error)))?;
match self.inner.send.try_send(frame) {
Ok(()) => Ok(()),
Err(mpsc::error::TrySendError::Full(_frame)) => Err(OrderedSendError::Full),
Err(mpsc::error::TrySendError::Closed(_frame)) => Err(OrderedSendError::Closed),
}
}
}
#[derive(Debug)]
pub(super) enum HeaderSyncPeerCommand {
RecordExpectedHeaders(ExpectedHeadersResponse),
}
pub(crate) async fn drive_header_sync_actions(
mut actions: mpsc::Receiver<HeaderSyncAction>,
handle: HeaderSyncHandle,
_supervisor: ZakuraSupervisorHandle,
shutdown: CancellationToken,
) {
loop {
let action = tokio::select! {
_ = shutdown.cancelled() => return,
action = actions.recv() => {
let Some(action) = action else {
return;
};
action
}
};
match action {
#[cfg(test)]
HeaderSyncAction::SendMessage { .. } | HeaderSyncAction::ForwardNewBlock { .. } => {}
HeaderSyncAction::Misbehavior { peer, reason } => {
tracing::debug!(?peer, ?reason, "recorded Zakura header-sync peer violation");
}
HeaderSyncAction::NewBlockReceived { peer, hash, .. } => {
tracing::debug!(
?peer,
?hash,
"Zakura header-sync NewBlock body arrived before block-acceptance hook is wired"
);
}
HeaderSyncAction::QueryHeadersByHeightRange {
peer, start, count, ..
} => {
let _ = handle
.send(HeaderSyncEvent::HeaderRangeResponseFinished {
peer,
start_height: start,
requested_count: count,
returned_count: 0,
})
.await;
}
HeaderSyncAction::CommitHeaderRange {
peer,
start_height,
headers,
..
} => {
tracing::debug!(
?peer,
?start_height,
count = headers.len(),
"suppressing Zakura header range commit until state driver is wired"
);
}
HeaderSyncAction::QueryBestHeaderTip
| HeaderSyncAction::QueryMissingBlockBodies { .. }
| HeaderSyncAction::BodyGaps { .. }
| HeaderSyncAction::HeaderAdvanced { .. }
| HeaderSyncAction::HeaderReanchored { .. } => {}
}
}
}
#[derive(Debug)]
pub(crate) struct HeaderSyncService {
header_sync: HeaderSyncHandle,
peers: Arc<StdMutex<HashMap<ZakuraPeerId, HeaderSyncPeerRecord>>>,
}
#[derive(Debug)]
struct HeaderSyncPeerRecord {
conn_id: ZakuraConnId,
cancel_token: CancellationToken,
}
impl HeaderSyncService {
pub(crate) fn new(header_sync: HeaderSyncHandle) -> Self {
Self {
header_sync,
peers: Arc::new(StdMutex::new(HashMap::new())),
}
}
}
impl Service for HeaderSyncService {
fn name(&self) -> &'static str {
"header-sync"
}
fn streams(&self) -> &[Stream] {
header_sync_streams()
}
fn wants_peer(
&self,
_peer: &ZakuraPeerId,
_negotiated: u64,
direction: ServicePeerDirection,
) -> bool {
let snapshot = self.header_sync.peer_snapshot();
match direction {
ServicePeerDirection::Inbound => snapshot.inbound_slots_free > 0,
ServicePeerDirection::Outbound => snapshot.outbound_slots_free > 0,
}
}
fn add_peer(&self, mut peer: Peer) {
let Some((recv, send)) = peer.take_stream(ZAKURA_STREAM_HEADER_SYNC) else {
return;
};
let peer_id = peer.id.clone();
let session = PeerStreamSession::new(
peer_id.clone(),
ZAKURA_STREAM_HEADER_SYNC,
recv,
send,
peer.service_cancel_token(),
);
let service_cancel_token = session.cancel_token();
let connection_cancel_token = peer.cancel_token();
let close_cause = peer.close_cause();
let conn_id = peer.conn_id;
let (commands_tx, commands_rx) = mpsc::unbounded_channel();
let header_sync_session =
HeaderSyncPeerSession::new_with_commands(&session, peer.direction, commands_tx);
{
let mut peers = self
.peers
.lock()
.expect("header-sync peer map mutex is never poisoned");
if peers
.get(&peer_id)
.is_some_and(|record| record.conn_id > conn_id)
{
service_cancel_token.cancel();
return;
}
if let Some(old_record) = peers.insert(
peer_id.clone(),
HeaderSyncPeerRecord {
conn_id,
cancel_token: header_sync_session.cancel_token(),
},
) {
old_record.cancel_token.cancel();
}
}
let _ = self
.header_sync
.send_lifecycle(HeaderSyncEvent::PeerConnected(header_sync_session.clone()));
let (_session_peer, _stream_kind, recv, _send, _session_cancel) = session.into_parts();
let pipe = Pipe::new(
peer_id.clone(),
HsLocal::new(commands_rx, DEFAULT_HS_INBOUND_NEW_BLOCK_MIN_INTERVAL),
HsEnv::new(self.header_sync.clone()),
SessionGuard::oversize_only(header_sync_guard_max_bytes()),
run_inbound,
&PIPE_SHAPE,
);
let pipe_cancel_token = service_cancel_token.clone();
let protocol_connection_cancel_token = connection_cancel_token.clone();
let protocol_close_cause = close_cause.clone();
let pipe = async move {
handle_pipe_exit(
"header-sync",
&protocol_connection_cancel_token,
&protocol_close_cause,
run_peer(pipe, recv, pipe_cancel_token).await,
);
};
let teardown_handle = self.header_sync.clone();
let teardown_peers = self.peers.clone();
let teardown_peer = peer_id.clone();
let on_teardown = move || {
let should_notify = {
let mut peers = teardown_peers
.lock()
.expect("header-sync peer map mutex is never poisoned");
if peers
.get(&teardown_peer)
.is_some_and(|record| record.conn_id == conn_id)
{
peers.remove(&teardown_peer);
true
} else {
false
}
};
if should_notify {
let _ = teardown_handle
.send_lifecycle(HeaderSyncEvent::PeerDisconnected(teardown_peer));
}
};
let panic_connection_cancel_token = connection_cancel_token.clone();
let panic_close_cause = close_cause.clone();
let on_panic = move || {
panic_close_cause.record("service_panic");
panic_connection_cancel_token.cancel();
};
spawn_supervised_pipe(peer_id, service_cancel_token, on_teardown, on_panic, pipe);
}
fn remove_peer(&self, peer: &ZakuraPeerId, conn_id: ZakuraConnId) {
let removed = {
let mut peers = self
.peers
.lock()
.expect("header-sync peer map mutex is never poisoned");
if peers
.get(peer)
.is_some_and(|record| record.conn_id == conn_id)
{
peers.remove(peer)
} else {
None
}
};
if let Some(record) = removed {
record.cancel_token.cancel();
let _ = self
.header_sync
.send_lifecycle(HeaderSyncEvent::PeerDisconnected(peer.clone()));
}
}
fn deliver_frame(
&self,
peer_id: ZakuraPeerId,
stream_kind: u16,
frame: Frame,
) -> Result<(), SinkReject> {
if stream_kind != ZAKURA_STREAM_HEADER_SYNC {
return Ok(());
}
match deliver(&self.header_sync, None, peer_id, frame) {
Flow::Continue(()) | Flow::Done => Ok(()),
Flow::Reject(reject) => Err(reject),
}
}
}
fn header_sync_guard_max_bytes() -> u32 {
u32::try_from(MAX_HS_MESSAGE_BYTES)
.expect("MAX_HS_MESSAGE_BYTES is a 2 MiB constant that fits in u32")
}
#[derive(Debug)]
pub(crate) struct HeaderSyncPassthroughService {
inner: Arc<dyn Service>,
}
impl HeaderSyncPassthroughService {
pub(crate) fn new(inner: Arc<dyn Service>) -> Self {
Self { inner }
}
}
impl Service for HeaderSyncPassthroughService {
fn name(&self) -> &'static str {
"header-sync-passthrough"
}
fn streams(&self) -> &[Stream] {
header_sync_streams()
}
fn wants_peer(
&self,
peer: &ZakuraPeerId,
negotiated: u64,
direction: ServicePeerDirection,
) -> bool {
self.inner.wants_peer(peer, negotiated, direction)
}
fn add_peer(&self, mut peer: Peer) {
let Some((recv, _send)) = peer.take_stream(ZAKURA_STREAM_HEADER_SYNC) else {
return;
};
let inner = self.inner.clone();
let peer_id = peer.id.clone();
let cancel_token = peer.cancel_token();
task::spawn(async move {
let sink = Box::new(HeaderSyncPassthroughSink {
peer_id: peer_id.clone(),
inner,
cancel_token: cancel_token.clone(),
});
match sink.run(recv).await {
Ok(()) => {}
Err(SinkReject::Protocol(error)) => {
tracing::debug!(
?error,
?peer_id,
"header-sync passthrough rejected protocol-invalid frame"
);
cancel_token.cancel();
}
Err(SinkReject::Local(error)) => {
tracing::debug!(
?error,
?peer_id,
"header-sync passthrough could not deliver frame locally"
);
}
}
});
}
fn remove_peer(&self, _peer: &ZakuraPeerId, _conn_id: ZakuraConnId) {}
fn deliver_frame(
&self,
peer_id: ZakuraPeerId,
stream_kind: u16,
frame: Frame,
) -> Result<(), SinkReject> {
self.inner.deliver_frame(peer_id, stream_kind, frame)
}
}
#[derive(Debug)]
struct HeaderSyncPassthroughSink {
peer_id: ZakuraPeerId,
inner: Arc<dyn Service>,
cancel_token: CancellationToken,
}
impl Sink for HeaderSyncPassthroughSink {
fn run(self: Box<Self>, mut recv: FramedRecv) -> BoxRunFuture<'static, Result<(), SinkReject>> {
Box::pin(async move {
loop {
let frame = tokio::select! {
_ = self.cancel_token.cancelled() => return Ok(()),
frame = recv.recv() => {
let Some(frame) = frame else {
return Ok(());
};
frame
}
};
match self.inner.deliver_frame(
self.peer_id.clone(),
ZAKURA_STREAM_HEADER_SYNC,
frame,
) {
Ok(()) => {}
Err(SinkReject::Protocol(error)) => return Err(SinkReject::Protocol(error)),
Err(SinkReject::Local(error)) => {
tracing::debug!(
?error,
peer_id = ?self.peer_id,
"header-sync passthrough could not deliver frame locally"
);
}
}
}
})
}
}