use std::collections::HashMap;
use byteorder::{LittleEndian, ReadBytesExt, WriteBytesExt};
use iroh::endpoint::{Connection, Endpoint, RecvStream, SendStream, VarInt};
use tokio::sync::Mutex;
use super::{LocalEndpointFactory, ZakuraTestNode};
use crate::{
zakura::{
legacy_gossip::ZAKURA_STREAM_GOSSIP, run_native_initiator_handshake, Frame, StreamPrelude,
ZakuraHandshakeConfig, ZakuraLocalLimits, ZakuraPeerId, FRAME_HEADER_BYTES,
LEGACY_GOSSIP_VERSION, P2P_V2_ALPN, STREAM_PRELUDE_MAGIC, ZAKURA_BLOCK_SYNC_STREAM_VERSION,
ZAKURA_CAP_HEADER_SYNC, ZAKURA_CAP_LEGACY_GOSSIP, ZAKURA_DISCOVERY_STREAM_VERSION,
ZAKURA_HEADER_SYNC_STREAM_VERSION, ZAKURA_STREAM_BLOCK_SYNC, ZAKURA_STREAM_DISCOVERY,
ZAKURA_STREAM_HEADER_SYNC,
},
BoxError, Config,
};
#[derive(Debug)]
pub struct HostilePeer {
endpoint: Endpoint,
connection: Connection,
limits: ZakuraLocalLimits,
held_streams: Vec<SendStream>,
ordered_streams: Mutex<HashMap<(u16, u16), (SendStream, RecvStream)>>,
}
impl HostilePeer {
pub async fn connect_native(victim: &ZakuraTestNode, seed: u64) -> Result<Self, BoxError> {
Self::connect_native_with_capabilities(
victim,
seed,
ZAKURA_CAP_LEGACY_GOSSIP | ZAKURA_CAP_HEADER_SYNC,
)
.await
}
pub async fn connect_native_with_capabilities(
victim: &ZakuraTestNode,
seed: u64,
capabilities: u64,
) -> Result<Self, BoxError> {
let limits = victim.limits().clone();
let endpoint = LocalEndpointFactory::with_transport_config(limits.transport_config())
.endpoint(seed)
.await?;
let victim_addr = victim.node_addr().await;
endpoint.add_node_addr(victim_addr.clone())?;
let connection = endpoint.connect(victim_addr, P2P_V2_ALPN).await?;
let mut config = ZakuraHandshakeConfig::for_network(&Config::default().network);
config.supported_capabilities = capabilities;
let local_peer_id = ZakuraPeerId::new(endpoint.node_id().as_bytes().to_vec())?;
run_native_initiator_handshake(&connection, &limits, &config, &local_peer_id).await?;
Ok(Self {
endpoint,
connection,
limits,
held_streams: Vec::new(),
ordered_streams: Mutex::new(HashMap::new()),
})
}
pub fn id(&self) -> Result<ZakuraPeerId, BoxError> {
Ok(ZakuraPeerId::new(
self.endpoint.node_id().as_bytes().to_vec(),
)?)
}
pub async fn send_frame(&self, stream_kind: u16, payload: Vec<u8>) -> Result<(), BoxError> {
self.send_raw_frame(
stream_kind,
Frame {
message_type: 1,
flags: 0,
payload,
},
)
.await
}
pub async fn send_raw_frame(&self, stream_kind: u16, frame: Frame) -> Result<(), BoxError> {
if matches!(
stream_kind,
ZAKURA_STREAM_GOSSIP
| ZAKURA_STREAM_DISCOVERY
| ZAKURA_STREAM_HEADER_SYNC
| ZAKURA_STREAM_BLOCK_SYNC
) {
return self.send_ordered_raw_frame(stream_kind, frame).await;
}
let (mut send, _recv) = self.connection.open_bi().await?;
self.write_prelude(&mut send, stream_kind).await?;
send.write_all(&frame.encode(self.limits.max_frame_bytes)?)
.await?;
let _ = send.finish();
Ok(())
}
async fn send_ordered_raw_frame(&self, stream_kind: u16, frame: Frame) -> Result<(), BoxError> {
self.send_ordered_raw_frame_with_version(
stream_kind,
Self::stream_version(stream_kind),
frame,
)
.await
}
pub async fn send_ordered_raw_frame_with_version(
&self,
stream_kind: u16,
stream_version: u16,
frame: Frame,
) -> Result<(), BoxError> {
let mut streams = self.ordered_streams.lock().await;
let (send, _recv) = match streams.entry((stream_kind, stream_version)) {
std::collections::hash_map::Entry::Occupied(entry) => entry.into_mut(),
std::collections::hash_map::Entry::Vacant(entry) => {
let (mut send, recv) = self.connection.open_bi().await?;
self.write_prelude_with_version(&mut send, stream_kind, stream_version)
.await?;
entry.insert((send, recv))
}
};
send.write_all(&frame.encode(self.limits.max_frame_bytes)?)
.await?;
Ok(())
}
pub async fn recv_ordered_frame(&self, stream_kind: u16) -> Result<Frame, BoxError> {
self.recv_ordered_frame_with_version(stream_kind, Self::stream_version(stream_kind))
.await
}
pub async fn recv_ordered_frame_with_version(
&self,
stream_kind: u16,
stream_version: u16,
) -> Result<Frame, BoxError> {
let mut streams = self.ordered_streams.lock().await;
let (_send, recv) = match streams.entry((stream_kind, stream_version)) {
std::collections::hash_map::Entry::Occupied(entry) => entry.into_mut(),
std::collections::hash_map::Entry::Vacant(entry) => {
let (mut send, recv) = self.connection.open_bi().await?;
self.write_prelude_with_version(&mut send, stream_kind, stream_version)
.await?;
entry.insert((send, recv))
}
};
Self::read_frame(recv, self.limits.max_frame_bytes).await
}
pub async fn finish_ordered_stream(
&self,
stream_kind: u16,
stream_version: u16,
) -> Result<(), BoxError> {
let Some((mut send, _recv)) = self
.ordered_streams
.lock()
.await
.remove(&(stream_kind, stream_version))
else {
return Ok(());
};
send.finish()?;
Ok(())
}
pub async fn reset_ordered_stream(
&self,
stream_kind: u16,
stream_version: u16,
) -> Result<(), BoxError> {
let Some((mut send, mut recv)) = self
.ordered_streams
.lock()
.await
.remove(&(stream_kind, stream_version))
else {
return Ok(());
};
let code = VarInt::from_u32(0);
send.reset(code)?;
recv.stop(code)?;
Ok(())
}
pub async fn reopen_ordered_stream(
&self,
stream_kind: u16,
stream_version: u16,
) -> Result<(), BoxError> {
self.reset_ordered_stream(stream_kind, stream_version)
.await?;
let mut streams = self.ordered_streams.lock().await;
let (mut send, recv) = self.connection.open_bi().await?;
self.write_prelude_with_version(&mut send, stream_kind, stream_version)
.await?;
streams.insert((stream_kind, stream_version), (send, recv));
Ok(())
}
pub async fn send_truncated_frame(&self, stream_kind: u16) -> Result<(), BoxError> {
let (mut send, _recv) = self.connection.open_bi().await?;
self.write_prelude(&mut send, stream_kind).await?;
let mut header = Vec::with_capacity(FRAME_HEADER_BYTES);
WriteBytesExt::write_u16::<LittleEndian>(&mut header, 1)?;
WriteBytesExt::write_u16::<LittleEndian>(&mut header, 0)?;
WriteBytesExt::write_u32::<LittleEndian>(&mut header, 8)?;
send.write_all(&header).await?;
send.write_all(&[1, 2, 3]).await?;
let _ = send.finish();
Ok(())
}
pub async fn send_frame_with_version(
&self,
stream_kind: u16,
stream_version: u16,
payload: Vec<u8>,
) -> Result<(), BoxError> {
let (mut send, _recv) = self.connection.open_bi().await?;
let prelude = StreamPrelude {
magic: STREAM_PRELUDE_MAGIC,
stream_kind,
stream_version,
request_id: None,
max_frame_bytes: self.limits.max_frame_bytes,
};
send.write_all(&prelude.encode()?).await?;
let frame = Frame {
message_type: 1,
flags: 0,
payload,
};
send.write_all(&frame.encode(self.limits.max_frame_bytes)?)
.await?;
let _ = send.finish();
Ok(())
}
pub async fn send_frame_with_request_id(
&self,
stream_kind: u16,
request_id: u64,
frame: Frame,
) -> Result<(), BoxError> {
let (mut send, _recv) = self.connection.open_bi().await?;
let prelude = StreamPrelude {
magic: STREAM_PRELUDE_MAGIC,
stream_kind,
stream_version: 1,
request_id: Some(request_id),
max_frame_bytes: self.limits.max_frame_bytes,
};
send.write_all(&prelude.encode()?).await?;
send.write_all(&frame.encode(self.limits.max_frame_bytes)?)
.await?;
let _ = send.finish();
Ok(())
}
pub async fn respond_to_next_request(&self, frames: Vec<Frame>) -> Result<(), BoxError> {
let (mut send, _recv) = self.connection.accept_bi().await?;
for frame in frames {
send.write_all(&frame.encode(self.limits.max_frame_bytes)?)
.await?;
}
let _ = send.finish();
Ok(())
}
pub async fn respond_to_next_request_with(
&self,
build_frames: impl FnOnce(u64) -> Vec<Frame>,
) -> Result<(), BoxError> {
let (mut send, mut recv) = self.connection.accept_bi().await?;
let prelude = Self::read_prelude(&mut recv).await?;
let request_id = prelude
.request_id
.ok_or_else(|| BoxError::from("request stream did not include a request id"))?;
for frame in build_frames(request_id) {
send.write_all(&frame.encode(self.limits.max_frame_bytes)?)
.await?;
}
let _ = send.finish();
Ok(())
}
pub async fn accept_next_request_without_response(&mut self) -> Result<(), BoxError> {
let (send, mut recv) = self.connection.accept_bi().await?;
let prelude = Self::read_prelude(&mut recv).await?;
prelude
.request_id
.ok_or_else(|| BoxError::from("request stream did not include a request id"))?;
self.held_streams.push(send);
Ok(())
}
pub async fn flood_stream(
&self,
stream_kind: u16,
label: char,
count: usize,
) -> Result<(), BoxError> {
let (mut send, _recv) = self.connection.open_bi().await?;
self.write_prelude(&mut send, stream_kind).await?;
for index in 0..count {
let frame = Frame {
message_type: 1,
flags: 0,
payload: format!("{label}-{index}").into_bytes(),
};
send.write_all(&frame.encode(self.limits.max_frame_bytes)?)
.await?;
}
let _ = send.finish();
Ok(())
}
pub async fn oversize_frame_declared_len(&self, stream_kind: u16) -> Result<(), BoxError> {
self.send_frame_header_with_declared_payload_len(
stream_kind,
self.limits.max_frame_bytes.saturating_add(1),
)
.await
}
pub async fn send_frame_header_with_declared_payload_len(
&self,
stream_kind: u16,
declared_payload_len: u32,
) -> Result<(), BoxError> {
let (mut send, _recv) = self.connection.open_bi().await?;
self.write_prelude(&mut send, stream_kind).await?;
let mut header = Vec::with_capacity(FRAME_HEADER_BYTES);
WriteBytesExt::write_u16::<LittleEndian>(&mut header, 1)?;
WriteBytesExt::write_u16::<LittleEndian>(&mut header, 0)?;
WriteBytesExt::write_u32::<LittleEndian>(&mut header, declared_payload_len)?;
send.write_all(&header).await?;
let _ = send.finish();
Ok(())
}
pub async fn open_and_never_send_prelude(&mut self) -> Result<(), BoxError> {
let (send, _recv) = self.connection.open_bi().await?;
self.held_streams.push(send);
tokio::time::sleep(self.limits.prelude_timeout + std::time::Duration::from_millis(20))
.await;
Ok(())
}
pub async fn churn_streams(&mut self, stream_kind: u16, count: usize) -> Result<(), BoxError> {
for _ in 0..count.min(4096) {
let (mut send, _recv) = self.connection.open_bi().await?;
self.write_prelude(&mut send, stream_kind).await?;
self.held_streams.push(send);
}
Ok(())
}
pub async fn shutdown(self) {
self.connection.close(VarInt::from_u32(0), b"hostile done");
self.endpoint.close().await;
}
async fn write_prelude(&self, send: &mut SendStream, stream_kind: u16) -> Result<(), BoxError> {
self.write_prelude_with_version(send, stream_kind, Self::stream_version(stream_kind))
.await
}
async fn write_prelude_with_version(
&self,
send: &mut SendStream,
stream_kind: u16,
stream_version: u16,
) -> Result<(), BoxError> {
let prelude = StreamPrelude {
magic: STREAM_PRELUDE_MAGIC,
stream_kind,
stream_version,
request_id: None,
max_frame_bytes: self.limits.max_frame_bytes,
};
send.write_all(&prelude.encode()?).await?;
Ok(())
}
fn stream_version(stream_kind: u16) -> u16 {
match stream_kind {
ZAKURA_STREAM_GOSSIP => LEGACY_GOSSIP_VERSION,
ZAKURA_STREAM_DISCOVERY => ZAKURA_DISCOVERY_STREAM_VERSION,
ZAKURA_STREAM_HEADER_SYNC => ZAKURA_HEADER_SYNC_STREAM_VERSION,
ZAKURA_STREAM_BLOCK_SYNC => ZAKURA_BLOCK_SYNC_STREAM_VERSION,
_ => 1,
}
}
pub fn selected_header_sync_version(capabilities: u64) -> Option<u16> {
(capabilities & ZAKURA_CAP_HEADER_SYNC != 0).then_some(ZAKURA_HEADER_SYNC_STREAM_VERSION)
}
async fn read_prelude(recv: &mut RecvStream) -> Result<StreamPrelude, BoxError> {
let mut bytes = vec![0; 4 + 2 + 2 + 1];
recv.read_exact(&mut bytes).await?;
match bytes[8] {
0 => {}
1 => {
let mut request_id = [0; 8];
recv.read_exact(&mut request_id).await?;
bytes.extend_from_slice(&request_id);
}
flag => return Err(format!("invalid request id flag: {flag}").into()),
}
let mut cap = [0; 4];
recv.read_exact(&mut cap).await?;
bytes.extend_from_slice(&cap);
Ok(StreamPrelude::decode(&bytes)?)
}
async fn read_frame(recv: &mut RecvStream, max_frame_bytes: u32) -> Result<Frame, BoxError> {
let mut header = vec![0; FRAME_HEADER_BYTES];
recv.read_exact(&mut header).await?;
let mut reader = std::io::Cursor::new(&header);
let _message_type = reader.read_u16::<LittleEndian>()?;
let _flags = reader.read_u16::<LittleEndian>()?;
let payload_len = usize::try_from(reader.read_u32::<LittleEndian>()?)?;
let frame_len = FRAME_HEADER_BYTES.saturating_add(payload_len);
if frame_len > usize::try_from(max_frame_bytes)? {
return Err(format!(
"frame payload length {payload_len} exceeds max_frame_bytes {max_frame_bytes}"
)
.into());
}
let mut payload = vec![0; payload_len];
recv.read_exact(&mut payload).await?;
header.extend_from_slice(&payload);
Ok(Frame::decode(&header, max_frame_bytes)?)
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Duration;
use iroh::protocol::{AcceptError, ProtocolHandler, Router};
const TEST_ALPN: &[u8] = b"/zakura/testkit/hostile-read-frame/0";
const MAX_FRAME_BYTES: u32 = 4096;
const DECLARED_PAYLOAD_LEN: u32 = 64 * 1024 * 1024;
#[derive(Clone, Debug)]
struct OversizeResponder;
impl ProtocolHandler for OversizeResponder {
async fn accept(&self, connection: Connection) -> Result<(), AcceptError> {
if let Ok((mut send, _recv)) = connection.open_bi().await {
let mut header = Vec::with_capacity(FRAME_HEADER_BYTES);
header.extend_from_slice(&1u16.to_le_bytes()); header.extend_from_slice(&0u16.to_le_bytes()); header.extend_from_slice(&DECLARED_PAYLOAD_LEN.to_le_bytes()); let _ = send.write_all(&header).await;
tokio::time::sleep(Duration::from_secs(8)).await;
}
Ok(())
}
}
#[tokio::test]
async fn read_frame_rejects_oversize_declared_len_before_allocating_payload(
) -> Result<(), BoxError> {
let server = LocalEndpointFactory::new().endpoint(4040).await?;
let router = Router::builder(server)
.accept(TEST_ALPN, OversizeResponder)
.spawn();
let server_addr = LocalEndpointFactory::node_addr(router.endpoint()).await;
let client = LocalEndpointFactory::new().endpoint(4041).await?;
client.add_node_addr(server_addr.clone())?;
let connection = client.connect(server_addr, TEST_ALPN).await?;
let (_send, mut recv) = connection.accept_bi().await?;
let outcome = tokio::time::timeout(
Duration::from_secs(2),
HostilePeer::read_frame(&mut recv, MAX_FRAME_BYTES),
)
.await
.expect(
"read_frame must reject the oversize declared length before allocating/reading the \
payload; a timeout here means payload_len was allocated/read before the \
max_frame_bytes check",
);
let err = outcome.expect_err("oversize declared frame length must be rejected");
let message = err.to_string();
assert!(
message.contains("max_frame_bytes"),
"expected a max_frame_bytes rejection, got: {message}",
);
connection.close(0u32.into(), b"done");
client.close().await;
Ok(())
}
}