use std::{collections::VecDeque, sync::Arc};
use tokio::sync::mpsc;
use tokio_util::sync::CancellationToken;
use super::{events::*, scheduler::*, service::HeaderSyncPeerCommand, wire::*, *};
use crate::zakura::{
Edge, Flow, FramedRecv, Node, NodeKind, Pipe, PipeCx, PipeShape, SinkReject, ZakuraPeerId,
};
pub(super) struct HsLocal {
expected_headers: VecDeque<ExpectedHeadersResponse>,
commands: mpsc::UnboundedReceiver<HeaderSyncPeerCommand>,
new_block_meter: RateMeter,
}
impl HsLocal {
pub(super) fn new(
commands: mpsc::UnboundedReceiver<HeaderSyncPeerCommand>,
new_block_min_interval: Duration,
) -> Self {
Self {
expected_headers: VecDeque::new(),
commands,
new_block_meter: RateMeter::new(new_block_min_interval),
}
}
fn admit_new_block(&mut self) -> bool {
self.new_block_meter.try_take(Instant::now())
}
fn pop_expected_headers_response(&mut self) -> Option<ExpectedHeadersResponse> {
self.expected_headers.pop_front()
}
fn restore_expected_headers(&mut self, expected: ExpectedHeadersResponse) {
self.expected_headers.push_front(expected);
}
fn handle_command(&mut self, command: HeaderSyncPeerCommand) {
match command {
HeaderSyncPeerCommand::RecordExpectedHeaders(expected) => {
self.expected_headers.push_back(expected);
}
}
}
fn drain_ready_commands(&mut self) {
while let Ok(command) = self.commands.try_recv() {
self.handle_command(command);
}
}
}
#[derive(Clone)]
pub(super) struct HsEnv {
handle: HeaderSyncHandle,
}
impl HsEnv {
pub(super) fn new(handle: HeaderSyncHandle) -> Self {
Self { handle }
}
}
pub(super) const PIPE_SHAPE: PipeShape = PipeShape {
service: "header-sync",
nodes: &[
Node {
id: "guard",
kind: NodeKind::Guard,
},
Node {
id: "decode",
kind: NodeKind::Decode,
},
Node {
id: "correlate",
kind: NodeKind::Mutate,
},
Node {
id: "emit",
kind: NodeKind::Emit,
},
],
edges: &[
Edge {
from: "guard",
to: "correlate",
on: "Headers",
},
Edge {
from: "guard",
to: "decode",
on: "Control",
},
Edge {
from: "correlate",
to: "decode",
on: "Expected",
},
Edge {
from: "decode",
to: "emit",
on: "Ok",
},
],
};
pub(super) fn run_inbound(cx: &mut PipeCx<'_, HsLocal, HsEnv>, frame: Frame) -> Flow<()> {
if u8::try_from(frame.message_type).ok() == Some(MSG_HS_NEW_BLOCK)
&& !cx.local.admit_new_block()
{
metrics::counter!("sync.header.tip.new_block.predecode_throttled").increment(1);
return Flow::Done;
}
let expected = (u8::try_from(frame.message_type).ok() == Some(MSG_HS_HEADERS))
.then(|| cx.local.pop_expected_headers_response())
.flatten();
match deliver(&cx.env.handle, expected, cx.peer_id.clone(), frame) {
Flow::Reject(SinkReject::Local(error)) => {
if let Some(expected) = expected {
cx.local.restore_expected_headers(expected);
}
tracing::debug!(
?error,
peer_id = ?cx.peer_id,
"header-sync stream could not deliver frame locally"
);
Flow::Done
}
other => other,
}
}
pub(super) fn deliver(
handle: &HeaderSyncHandle,
expected: Option<ExpectedHeadersResponse>,
peer_id: ZakuraPeerId,
frame: Frame,
) -> Flow<()> {
if u8::try_from(frame.message_type).ok() == Some(MSG_HS_HEADERS) {
let Some(expected) = expected else {
let error = Arc::new(HeaderSyncWireError::UnsolicitedHeaders);
let _ = handle.try_send(HeaderSyncEvent::WireProtocolFailure {
peer: peer_id.clone(),
reason: HeaderSyncMisbehavior::UnsolicitedHeaders,
error: error.clone(),
});
let protocol_error =
std::io::Error::new(std::io::ErrorKind::InvalidData, error.to_string());
return Flow::Reject(SinkReject::protocol(protocol_error));
};
let msg = match HeaderSyncMessage::decode_frame(
frame,
HeaderSyncDecodeContext::for_headers_response(expected, expected.count),
) {
Ok(msg) => msg,
Err(error) => {
let protocol_error =
std::io::Error::new(std::io::ErrorKind::InvalidData, error.to_string());
let _ = handle.try_send(HeaderSyncEvent::WireProtocolFailure {
peer: peer_id.clone(),
reason: HeaderSyncMisbehavior::MalformedMessage,
error: Arc::new(error),
});
return Flow::Reject(SinkReject::protocol(protocol_error));
}
};
return forward(handle, HeaderSyncEvent::WireMessage { peer: peer_id, msg });
}
let msg = match decode_control_frame(frame) {
Ok(msg) => msg,
Err(error) => {
let protocol_error =
std::io::Error::new(std::io::ErrorKind::InvalidData, error.to_string());
let _ = handle.try_send(HeaderSyncEvent::WireDecodeFailed {
peer: peer_id,
error: Arc::new(error),
});
return Flow::Reject(SinkReject::protocol(protocol_error));
}
};
forward(handle, HeaderSyncEvent::WireMessage { peer: peer_id, msg })
}
pub(super) async fn run_peer(
mut pipe: Pipe<HsLocal, HsEnv>,
mut recv: FramedRecv,
cancel: CancellationToken,
) -> Result<(), SinkReject> {
enum Input {
Frame(Frame),
Command(HeaderSyncPeerCommand),
Done,
}
loop {
pipe.local_mut().drain_ready_commands();
let input = {
let local = pipe.local_mut();
tokio::select! {
biased;
() = cancel.cancelled() => Input::Done,
command = local.commands.recv() => match command {
Some(command) => Input::Command(command),
None => Input::Done,
},
frame = recv.recv() => match frame {
Some(frame) => Input::Frame(frame),
None => Input::Done,
},
}
};
match input {
Input::Done => return Ok(()),
Input::Frame(frame) => {
pipe.local_mut().drain_ready_commands();
match pipe.run_one(frame) {
Flow::Continue(()) | Flow::Done => {}
Flow::Reject(reject) => return Err(reject),
}
}
Input::Command(command) => pipe.local_mut().handle_command(command),
}
}
}
fn forward(handle: &HeaderSyncHandle, event: HeaderSyncEvent) -> Flow<()> {
match handle.try_send(event) {
Ok(()) => Flow::Continue(()),
Err(error) => Flow::Reject(SinkReject::local(format!(
"header-sync queue closed: {error}"
))),
}
}
fn decode_control_frame(frame: Frame) -> Result<HeaderSyncMessage, HeaderSyncWireError> {
if u8::try_from(frame.message_type).ok() == Some(MSG_HS_HEADERS) {
return Err(HeaderSyncWireError::UnsolicitedHeaders);
}
HeaderSyncMessage::decode_frame(frame, HeaderSyncDecodeContext::control())
}
#[cfg(test)]
mod tests {
use tokio::sync::watch;
use super::*;
use crate::zakura::{ServicePeerSnapshot, ZakuraHeaderSyncCandidateState};
const FRAME_FORKS: [&str; 2] = ["Headers", "Control"];
fn peer() -> ZakuraPeerId {
ZakuraPeerId::new(vec![5; 32]).expect("test peer id is within bounds")
}
fn test_handle() -> (HeaderSyncHandle, mpsc::Receiver<HeaderSyncEvent>) {
let (events, events_rx) = mpsc::channel(16);
let (lifecycle, _lifecycle_rx) = mpsc::unbounded_channel();
let (_tip_tx, tip) = watch::channel((block::Height(0), block::Hash([0; 32])));
let (_peers_tx, peers) = watch::channel(ServicePeerSnapshot::default());
let (_candidates_tx, candidates) =
watch::channel(ZakuraHeaderSyncCandidateState::default());
(
HeaderSyncHandle {
events,
lifecycle,
tip,
peers,
candidates,
},
events_rx,
)
}
fn saturated_events_handle() -> (HeaderSyncHandle, mpsc::Receiver<HeaderSyncEvent>) {
let (events, events_rx) = mpsc::channel(1);
events
.try_send(HeaderSyncEvent::PeerDisconnected(peer()))
.expect("the single events slot is free");
let (lifecycle, _lifecycle_rx) = mpsc::unbounded_channel();
let (_tip_tx, tip) = watch::channel((block::Height(0), block::Hash([0; 32])));
let (_peers_tx, peers) = watch::channel(ServicePeerSnapshot::default());
let (_candidates_tx, candidates) =
watch::channel(ZakuraHeaderSyncCandidateState::default());
(
HeaderSyncHandle {
events,
lifecycle,
tip,
peers,
candidates,
},
events_rx,
)
}
fn headers_frame(payload: Vec<u8>) -> Frame {
Frame {
message_type: u16::from(MSG_HS_HEADERS),
flags: 0,
payload,
}
}
#[test]
fn deliver_unsolicited_headers_rejects_without_expectation() {
let (handle, mut events) = test_handle();
let flow = deliver(&handle, None, peer(), headers_frame(Vec::new()));
assert!(matches!(flow, Flow::Reject(SinkReject::Protocol(_))));
match events.try_recv() {
Ok(HeaderSyncEvent::WireProtocolFailure { reason, .. }) => {
assert!(matches!(reason, HeaderSyncMisbehavior::UnsolicitedHeaders));
}
other => panic!("expected WireProtocolFailure(UnsolicitedHeaders), got {other:?}"),
}
}
#[test]
fn deliver_correlated_headers_decodes_against_expectation() {
let (handle, mut events) = test_handle();
let expected =
ExpectedHeadersResponse::new(block::Height(1), 1, true).expect("count is valid");
let flow = deliver(&handle, Some(expected), peer(), headers_frame(Vec::new()));
assert!(matches!(flow, Flow::Reject(SinkReject::Protocol(_))));
match events.try_recv() {
Ok(HeaderSyncEvent::WireProtocolFailure { reason, .. }) => {
assert!(matches!(reason, HeaderSyncMisbehavior::MalformedMessage));
}
other => panic!("expected WireProtocolFailure(MalformedMessage), got {other:?}"),
}
}
#[test]
fn local_correlation_queue_drains_commands_in_fifo_order() {
let (commands_tx, commands_rx) = mpsc::unbounded_channel();
let mut local = HsLocal::new(commands_rx, DEFAULT_HS_INBOUND_NEW_BLOCK_MIN_INTERVAL);
let first =
ExpectedHeadersResponse::new(block::Height(1), 1, false).expect("count is valid");
let second =
ExpectedHeadersResponse::new(block::Height(2), 2, false).expect("count is valid");
commands_tx
.send(HeaderSyncPeerCommand::RecordExpectedHeaders(first))
.expect("pipe is alive");
commands_tx
.send(HeaderSyncPeerCommand::RecordExpectedHeaders(second))
.expect("pipe is alive");
assert_eq!(local.pop_expected_headers_response(), None);
local.drain_ready_commands();
assert_eq!(local.pop_expected_headers_response(), Some(first));
assert_eq!(local.pop_expected_headers_response(), Some(second));
assert_eq!(local.pop_expected_headers_response(), None);
}
#[test]
fn new_block_flood_is_throttled_before_decode() {
use zakura_chain::serialization::ZcashDeserializeInto;
use zakura_test::vectors::{BLOCK_MAINNET_1_BYTES, BLOCK_MAINNET_2_BYTES};
let (handle, mut events) = test_handle();
let (_commands_tx, commands_rx) = mpsc::unbounded_channel();
let block_one: Arc<block::Block> = Arc::new(
BLOCK_MAINNET_1_BYTES
.zcash_deserialize_into()
.expect("block 1 vector parses"),
);
let block_two: Arc<block::Block> = Arc::new(
BLOCK_MAINNET_2_BYTES
.zcash_deserialize_into()
.expect("block 2 vector parses"),
);
let frame_one = HeaderSyncMessage::NewBlock(block_one.clone())
.encode_frame()
.expect("new block frame encodes");
let frame_two = HeaderSyncMessage::NewBlock(block_two.clone())
.encode_frame()
.expect("new block frame encodes");
let mut pipe = Pipe::new(
peer(),
HsLocal::new(commands_rx, DEFAULT_HS_INBOUND_NEW_BLOCK_MIN_INTERVAL),
HsEnv::new(handle),
crate::zakura::SessionGuard::oversize_only(MAX_HS_MESSAGE_BYTES as u32),
run_inbound,
&PIPE_SHAPE,
);
assert!(matches!(pipe.run_one(frame_one), Flow::Continue(())));
match events.try_recv() {
Ok(HeaderSyncEvent::WireMessage {
msg: HeaderSyncMessage::NewBlock(block),
..
}) => assert_eq!(block.hash(), block_one.hash()),
other => panic!("expected first NewBlock to be forwarded, got {other:?}"),
}
assert!(matches!(pipe.run_one(frame_two), Flow::Done));
assert!(
matches!(events.try_recv(), Err(mpsc::error::TryRecvError::Empty)),
"second NewBlock must be throttled before decode, not forwarded"
);
}
#[test]
fn saturated_events_queue_restores_solicited_expectation() {
use zakura_chain::{orchard, sapling, serialization::ZcashDeserializeInto};
use zakura_test::vectors::BLOCK_MAINNET_1_BYTES;
let (handle, _events_rx) = saturated_events_handle();
let (commands_tx, commands_rx) = mpsc::unbounded_channel();
let expected =
ExpectedHeadersResponse::new(block::Height(1), 1, true).expect("count is valid");
commands_tx
.send(HeaderSyncPeerCommand::RecordExpectedHeaders(expected))
.expect("pipe is alive");
let block_one: Arc<block::Block> = Arc::new(
BLOCK_MAINNET_1_BYTES
.zcash_deserialize_into()
.expect("block 1 vector parses"),
);
let solicited_headers = HeaderSyncMessage::Headers {
headers: vec![block_one.header.clone()],
body_sizes: vec![0],
tree_aux_roots: vec![BlockCommitmentRoots {
height: block::Height(1),
sapling_root: sapling::tree::NoteCommitmentTree::default().root(),
orchard_root: orchard::tree::NoteCommitmentTree::default().root(),
ironwood_root: zakura_chain::ironwood::tree::NoteCommitmentTree::default().root(),
sapling_tx: 0,
orchard_tx: 0,
ironwood_tx: 0,
auth_data_root: block::merkle::AuthDataRoot::from([0u8; 32]),
}],
}
.encode_frame()
.expect("headers frame encodes");
let mut pipe = Pipe::new(
peer(),
HsLocal::new(commands_rx, DEFAULT_HS_INBOUND_NEW_BLOCK_MIN_INTERVAL),
HsEnv::new(handle),
crate::zakura::SessionGuard::oversize_only(MAX_HS_MESSAGE_BYTES as u32),
run_inbound,
&PIPE_SHAPE,
);
pipe.local_mut().drain_ready_commands();
assert_eq!(
pipe.local_mut().pop_expected_headers_response(),
Some(expected),
"the solicited response expectation should be available after draining commands"
);
pipe.local_mut().restore_expected_headers(expected);
HeaderSyncMessage::decode_frame(
solicited_headers.clone(),
HeaderSyncDecodeContext::for_headers_response(expected, expected.count),
)
.expect("test Headers frame decodes against its expectation");
let flow = pipe.run_one(solicited_headers);
match flow {
Flow::Done => {}
Flow::Continue(()) => panic!("unexpected successful forward"),
Flow::Reject(SinkReject::Protocol(_)) => panic!("unexpected protocol reject"),
Flow::Reject(SinkReject::Local(_)) => panic!("unexpected local reject"),
}
assert_eq!(
pipe.local_mut().pop_expected_headers_response(),
Some(expected),
"a solicited Headers response dropped on reactor queue saturation must restore its expectation"
);
}
#[test]
fn pipe_shape_matches_runtime() {
PIPE_SHAPE
.validate()
.expect("header-sync PIPE_SHAPE edges name only real nodes");
let frame_forks: Vec<&str> = PIPE_SHAPE
.edges
.iter()
.filter(|edge| edge.from == "guard")
.map(|edge| edge.on)
.collect();
assert_eq!(
frame_forks.len(),
FRAME_FORKS.len(),
"guard has exactly the runtime frame-shape forks"
);
for fork in FRAME_FORKS {
assert!(
frame_forks.contains(&fork),
"guard edge missing for runtime fork {fork}"
);
}
assert!(
PIPE_SHAPE
.edges
.iter()
.any(|edge| edge.from == "correlate" && edge.to == "decode"),
"headers responses correlate before decode"
);
assert!(
PIPE_SHAPE
.nodes
.iter()
.any(|node| node.id == "emit" && matches!(node.kind, NodeKind::Emit)),
"the pipe terminates at a single `emit` node"
);
}
}