use std::{
collections::HashMap,
sync::{
atomic::{AtomicU64, Ordering},
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,
session_id: u64,
direction: ServicePeerDirection,
inner: Arc<HeaderSyncPeerSessionInner>,
}
#[derive(Debug)]
struct HeaderSyncPeerSessionInner {
send: FramedSend,
cancel_token: CancellationToken,
commands: Option<mpsc::UnboundedSender<HeaderSyncPeerCommand>>,
next_request_id: AtomicU64,
}
impl HeaderSyncPeerSession {
fn new_with_commands(
session: &PeerStreamSession,
direction: ServicePeerDirection,
commands: mpsc::UnboundedSender<HeaderSyncPeerCommand>,
session_id: u64,
) -> Self {
Self::from_parts_with_direction_and_commands(
session.peer_id().clone(),
session_id,
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,
0,
direction,
send,
cancel_token,
None,
)
}
#[cfg(test)]
pub(crate) fn from_parts_with_direction_and_session_id(
peer_id: ZakuraPeerId,
direction: ServicePeerDirection,
send: FramedSend,
cancel_token: CancellationToken,
session_id: u64,
) -> Self {
Self::from_parts_with_direction_and_commands(
peer_id,
session_id,
direction,
send,
cancel_token,
None,
)
}
fn from_parts_with_direction_and_commands(
peer_id: ZakuraPeerId,
session_id: u64,
direction: ServicePeerDirection,
send: FramedSend,
cancel_token: CancellationToken,
commands: Option<mpsc::UnboundedSender<HeaderSyncPeerCommand>>,
) -> Self {
Self {
peer_id,
session_id,
direction,
inner: Arc::new(HeaderSyncPeerSessionInner {
send,
cancel_token,
commands,
next_request_id: AtomicU64::new(1),
}),
}
}
pub fn peer_id(&self) -> &ZakuraPeerId {
&self.peer_id
}
pub fn session_id(&self) -> u64 {
self.session_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 retire_expected_headers(
&self,
request_id: HeaderSyncRequestId,
) -> Result<(), OrderedSendError> {
let Some(commands) = &self.inner.commands else {
return Ok(());
};
commands
.send(HeaderSyncPeerCommand::Retire(request_id))
.map_err(|_| OrderedSendError::Closed)
}
fn next_request_id(&self) -> Result<HeaderSyncRequestId, OrderedSendError> {
let mut id = self.inner.next_request_id.load(Ordering::Relaxed);
loop {
let next_id = id.checked_add(1).ok_or_else(|| {
OrderedSendError::Encode("header-sync request ID counter exhausted".into())
})?;
match self.inner.next_request_id.compare_exchange_weak(
id,
next_id,
Ordering::Relaxed,
Ordering::Relaxed,
) {
Ok(_) => break,
Err(current_id) => id = current_id,
}
}
HeaderSyncRequestId::new(id).ok_or_else(|| {
OrderedSendError::Encode("header-sync request ID counter exhausted".into())
})
}
pub fn try_send_status(&self, status: HeaderSyncStatus) -> Result<(), OrderedSendError> {
self.try_send_message(HeaderSyncMessage::Status(status), None)
}
pub(super) fn prepare_get_headers(
&self,
start_height: block::Height,
count: u32,
want_tree_aux_roots: bool,
) -> Result<PreparedGetHeaders, OrderedSendError> {
let request_id = self.next_request_id()?;
let expected =
ExpectedHeadersResponse::new(request_id, start_height, count, want_tree_aux_roots)
.map_err(|error| OrderedSendError::Encode(Box::new(error)))?;
let frame = HeaderSyncMessage::GetHeaders {
start_height,
count,
want_tree_aux_roots,
}
.encode_frame(Some(request_id))
.map_err(|error| OrderedSendError::Encode(Box::new(error)))?;
let reservation = ExpectedHeadersReservation::new(self.inner.commands.clone(), expected)?;
Ok(PreparedGetHeaders {
request_id,
frame,
send: self.inner.send.clone(),
reservation,
})
}
pub fn try_send_headers_with_sizes_and_roots(
&self,
request_id: HeaderSyncRequestId,
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,
},
Some(request_id),
)
}
pub fn try_send_new_block(&self, block: Arc<block::Block>) -> Result<(), OrderedSendError> {
self.try_send_message(HeaderSyncMessage::NewBlock(block), None)
}
fn try_send_message(
&self,
msg: HeaderSyncMessage,
request_id: Option<HeaderSyncRequestId>,
) -> Result<(), OrderedSendError> {
let frame = msg
.encode_frame(request_id)
.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),
}
}
}
pub(super) struct PreparedGetHeaders {
request_id: HeaderSyncRequestId,
frame: Frame,
send: FramedSend,
reservation: ExpectedHeadersReservation,
}
impl PreparedGetHeaders {
pub(super) fn request_id(&self) -> HeaderSyncRequestId {
self.request_id
}
pub(super) async fn send(self) -> Result<HeaderSyncRequestId, OrderedSendError> {
let Self {
request_id,
frame,
send,
mut reservation,
} = self;
send.send(frame)
.await
.map_err(|_| OrderedSendError::Closed)?;
reservation.disarm();
Ok(request_id)
}
}
struct ExpectedHeadersReservation {
commands: Option<mpsc::UnboundedSender<HeaderSyncPeerCommand>>,
expected: ExpectedHeadersResponse,
armed: bool,
}
impl ExpectedHeadersReservation {
fn new(
commands: Option<mpsc::UnboundedSender<HeaderSyncPeerCommand>>,
expected: ExpectedHeadersResponse,
) -> Result<Self, OrderedSendError> {
if let Some(commands) = &commands {
commands
.send(HeaderSyncPeerCommand::Reserve(expected))
.map_err(|_| OrderedSendError::Closed)?;
}
Ok(Self {
commands,
expected,
armed: true,
})
}
fn disarm(&mut self) {
self.armed = false;
}
}
impl Drop for ExpectedHeadersReservation {
fn drop(&mut self) {
if self.armed {
if let Some(commands) = &self.commands {
let _ = commands.send(HeaderSyncPeerCommand::Cancel(self.expected));
}
}
}
}
#[derive(Debug)]
pub(super) enum HeaderSyncPeerCommand {
Reserve(ExpectedHeadersResponse),
Cancel(ExpectedHeadersResponse),
Retire(HeaderSyncRequestId),
}
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,
session_id,
request_id,
start,
count,
..
} => {
let _ = handle
.send(HeaderSyncEvent::HeaderRangeResponseFinished {
peer,
session_id,
request_id,
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,
session_id: u64,
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((session_id, recv, send)) =
peer.take_stream_with_session_id(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,
session_id,
);
{
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,
session_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_with_session_id(self.header_sync.clone(), session_id),
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 && record.session_id == session_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, 0, 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"
);
}
}
}
})
}
}
#[cfg(test)]
mod request_id_tests {
use super::*;
use crate::zakura::header_sync::{
requester::HeaderRequesterCommand,
state::{RangePriority, RangeRequest},
};
#[test]
fn request_id_exhaustion_remains_fail_closed() {
let (send, _recv) = crate::zakura::framed_channel(1);
let peer_id = ZakuraPeerId::new(vec![1; 32]).expect("test peer id is valid");
let session = HeaderSyncPeerSession::from_parts_with_direction(
peer_id,
ServicePeerDirection::Outbound,
send,
CancellationToken::new(),
);
session
.inner
.next_request_id
.store(u64::MAX, Ordering::Relaxed);
assert!(session.next_request_id().is_err());
assert!(session.next_request_id().is_err());
}
#[test]
fn requester_queue_failure_cancels_prepared_expectation() {
let (send, _recv) = crate::zakura::framed_channel(1);
let (commands_tx, mut commands_rx) = mpsc::unbounded_channel();
let peer_id = ZakuraPeerId::new(vec![5; 32]).expect("test peer id is valid");
let session = HeaderSyncPeerSession::from_parts_with_direction_and_commands(
peer_id,
1,
ServicePeerDirection::Outbound,
send,
CancellationToken::new(),
Some(commands_tx),
);
let range = RangeRequest {
start_height: block::Height(1),
count: 1,
anchor_hash: None,
finalized: false,
want_tree_aux_roots: true,
priority: RangePriority::Forward,
};
let prepared = session
.prepare_get_headers(range.start_height, range.count, range.want_tree_aux_roots)
.expect("valid test request is prepared");
let reserved = match commands_rx.try_recv().expect("reservation is published") {
HeaderSyncPeerCommand::Reserve(expected) => expected,
command => panic!("expected reservation command, got {command:?}"),
};
let (requester_tx, requester_rx) = mpsc::channel(1);
drop(requester_rx);
let rejected = match requester_tx.try_send(HeaderRequesterCommand { range, prepared }) {
Err(mpsc::error::TrySendError::Closed(command)) => command,
_ => panic!("closed requester queue rejects the prepared command"),
};
drop(rejected);
let cancelled = match commands_rx.try_recv().expect("cancellation is published") {
HeaderSyncPeerCommand::Cancel(expected) => expected,
command => panic!("expected cancellation command, got {command:?}"),
};
assert_eq!(cancelled, reserved);
}
}