Apply rustfmt to noise-xx-only code

This commit is contained in:
Johnathan Corgan
2026-04-11 08:16:01 +00:00
parent 59b7ea765b
commit 34f840a77b
38 changed files with 1647 additions and 1012 deletions
+40 -14
View File
@@ -12,8 +12,8 @@ pub fn draw(frame: &mut Frame, app: &App, area: Rect) {
let data = match app.data.get(&Tab::Bloom) {
Some(d) => d,
None => {
let msg = Paragraph::new(" Waiting for data...")
.style(Style::default().fg(Color::DarkGray));
let msg =
Paragraph::new(" Waiting for data...").style(Style::default().fg(Color::DarkGray));
frame.render_widget(msg, area);
return;
}
@@ -22,7 +22,7 @@ pub fn draw(frame: &mut Frame, app: &App, area: Rect) {
let chunks = Layout::vertical([
Constraint::Length(7), // Bloom Filter State
Constraint::Length(15), // Bloom Announce Stats
Constraint::Min(3), // Peer Filters
Constraint::Min(3), // Peer Filters
])
.split(area);
@@ -64,20 +64,47 @@ fn draw_stats(frame: &mut Frame, data: &serde_json::Value, area: Rect) {
helpers::section_header("Inbound"),
helpers::kv_line("Received", &helpers::nested_u64(data, "stats", "received")),
helpers::kv_line("Accepted", &helpers::nested_u64(data, "stats", "accepted")),
helpers::kv_line("Decode Error", &helpers::nested_u64(data, "stats", "decode_error")),
helpers::kv_line(
"Decode Error",
&helpers::nested_u64(data, "stats", "decode_error"),
),
helpers::kv_line("Invalid", &helpers::nested_u64(data, "stats", "invalid")),
helpers::kv_line("Unknown Peer", &helpers::nested_u64(data, "stats", "unknown_peer")),
helpers::kv_line(
"Unknown Peer",
&helpers::nested_u64(data, "stats", "unknown_peer"),
),
helpers::kv_line("Stale", &helpers::nested_u64(data, "stats", "stale")),
Line::from(""),
helpers::section_header("Outbound"),
helpers::kv_line("Sent", &helpers::nested_u64(data, "stats", "sent")),
helpers::kv_line("Full Sends", &helpers::nested_u64(data, "stats", "full_sends")),
helpers::kv_line("Deltas Sent", &helpers::nested_u64(data, "stats", "deltas_sent")),
helpers::kv_line("NACKs Sent", &helpers::nested_u64(data, "stats", "nacks_sent")),
helpers::kv_line("NACKs Received", &helpers::nested_u64(data, "stats", "nacks_received")),
helpers::kv_line("Size Changes", &helpers::nested_u64(data, "stats", "size_changes")),
helpers::kv_line("Debounce Suppressed", &helpers::nested_u64(data, "stats", "debounce_suppressed")),
helpers::kv_line("Send Failed", &helpers::nested_u64(data, "stats", "send_failed")),
helpers::kv_line(
"Full Sends",
&helpers::nested_u64(data, "stats", "full_sends"),
),
helpers::kv_line(
"Deltas Sent",
&helpers::nested_u64(data, "stats", "deltas_sent"),
),
helpers::kv_line(
"NACKs Sent",
&helpers::nested_u64(data, "stats", "nacks_sent"),
),
helpers::kv_line(
"NACKs Received",
&helpers::nested_u64(data, "stats", "nacks_received"),
),
helpers::kv_line(
"Size Changes",
&helpers::nested_u64(data, "stats", "size_changes"),
),
helpers::kv_line(
"Debounce Suppressed",
&helpers::nested_u64(data, "stats", "debounce_suppressed"),
),
helpers::kv_line(
"Send Failed",
&helpers::nested_u64(data, "stats", "send_failed"),
),
];
let max_lines = inner.height as usize;
@@ -101,8 +128,7 @@ fn draw_peer_filters(frame: &mut Frame, data: &serde_json::Value, area: Rect) {
frame.render_widget(block, area);
if filters.is_empty() {
let msg =
Paragraph::new(" No peers").style(Style::default().fg(Color::DarkGray));
let msg = Paragraph::new(" No peers").style(Style::default().fg(Color::DarkGray));
frame.render_widget(msg, inner);
return;
}
+3 -8
View File
@@ -59,10 +59,8 @@ pub fn rle_decode(data: &[u8], expected_words: usize) -> Result<Vec<u64>, RleErr
let mut pos = 0;
while pos + 10 <= data.len() {
let count =
u16::from_le_bytes(data[pos..pos + 2].try_into().unwrap()) as usize;
let value =
u64::from_le_bytes(data[pos + 2..pos + 10].try_into().unwrap());
let count = u16::from_le_bytes(data[pos..pos + 2].try_into().unwrap()) as usize;
let value = u64::from_le_bytes(data[pos + 2..pos + 10].try_into().unwrap());
pos += 10;
if words.len() + count > expected_words {
@@ -180,10 +178,7 @@ mod tests {
// Expect wrong number of words
let result = rle_decode(&encoded, 20);
assert!(matches!(
result,
Err(RleError::DecodedSizeMismatch { .. })
));
assert!(matches!(result, Err(RleError::DecodedSizeMismatch { .. })));
}
#[test]
+8 -4
View File
@@ -3,8 +3,8 @@
use std::fmt;
use super::{
BloomError, DEFAULT_FILTER_SIZE_BITS, DEFAULT_HASH_COUNT, MAX_SIZE_CLASS,
MIN_SIZE_CLASS, SIZE_CLASS_BYTES,
BloomError, DEFAULT_FILTER_SIZE_BITS, DEFAULT_HASH_COUNT, MAX_SIZE_CLASS, MIN_SIZE_CLASS,
SIZE_CLASS_BYTES,
};
use crate::NodeAddr;
@@ -175,7 +175,9 @@ impl BloomFilter {
if target_bits >= self.num_bits {
return Err(BloomError::InvalidTargetSize(target_bits));
}
if !target_bits.is_power_of_two() || target_bits < SIZE_CLASS_BYTES[MIN_SIZE_CLASS as usize] * 8 {
if !target_bits.is_power_of_two()
|| target_bits < SIZE_CLASS_BYTES[MIN_SIZE_CLASS as usize] * 8
{
return Err(BloomError::InvalidTargetSize(target_bits));
}
@@ -214,7 +216,9 @@ impl BloomFilter {
if target_bits <= self.num_bits {
return Err(BloomError::InvalidTargetSize(target_bits));
}
if !target_bits.is_power_of_two() || target_bits > SIZE_CLASS_BYTES[MAX_SIZE_CLASS as usize] * 8 {
if !target_bits.is_power_of_two()
|| target_bits > SIZE_CLASS_BYTES[MAX_SIZE_CLASS as usize] * 8
{
return Err(BloomError::InvalidTargetSize(target_bits));
}
+1 -1
View File
@@ -2,7 +2,7 @@
use std::collections::{HashMap, HashSet};
use super::{size_class_to_bits, BloomFilter, MAX_SIZE_CLASS, MIN_SIZE_CLASS, V1_SIZE_CLASS};
use super::{BloomFilter, MAX_SIZE_CLASS, MIN_SIZE_CLASS, V1_SIZE_CLASS, size_class_to_bits};
use crate::NodeAddr;
/// State for managing Bloom filter announcements.
+18 -4
View File
@@ -373,13 +373,20 @@ fn test_bloom_filter_fold() {
// All inserted elements must still be found (no false negatives)
for i in 0..50 {
assert!(folded.contains(&make_node_addr(i)), "Node {} not found after fold", i);
assert!(
folded.contains(&make_node_addr(i)),
"Node {} not found after fold",
i
);
}
// Fill ratio should roughly double
let original_fill = filter.fill_ratio();
let folded_fill = folded.fill_ratio();
assert!(folded_fill > original_fill * 1.5, "Fill ratio didn't increase enough");
assert!(
folded_fill > original_fill * 1.5,
"Fill ratio didn't increase enough"
);
}
#[test]
@@ -416,7 +423,11 @@ fn test_bloom_filter_duplicate() {
// All elements still found at the larger size
for i in 0..50 {
assert!(duped.contains(&make_node_addr(i)), "Node {} not found after duplicate", i);
assert!(
duped.contains(&make_node_addr(i)),
"Node {} not found after duplicate",
i
);
}
}
@@ -438,7 +449,10 @@ fn test_bloom_filter_duplicate_to() {
#[test]
fn test_bloom_filter_duplicate_at_maximum() {
let filter = BloomFilter::with_params(32768 * 8, 5).unwrap();
assert!(matches!(filter.duplicate(), Err(BloomError::CannotDuplicate(_))));
assert!(matches!(
filter.duplicate(),
Err(BloomError::CannotDuplicate(_))
));
}
#[test]
+58 -43
View File
@@ -364,55 +364,70 @@ pub fn show_bloom(node: &Node) -> Value {
/// `show_mmp` — MMP metrics summary.
pub fn show_mmp(node: &Node) -> Value {
// Link-layer MMP per peer
let peers: Vec<Value> = node.peers().filter_map(|peer| {
let mmp = peer.mmp()?;
let addr = *peer.node_addr();
let metrics = &mmp.metrics;
let peers: Vec<Value> = node
.peers()
.filter_map(|peer| {
let mmp = peer.mmp()?;
let addr = *peer.node_addr();
let metrics = &mmp.metrics;
let mut link_layer = json!({
"loss_rate": metrics.loss_rate(),
"etx": metrics.etx,
"goodput_bps": metrics.goodput_bps,
});
let mut link_layer = json!({
"loss_rate": metrics.loss_rate(),
"etx": metrics.etx,
"goodput_bps": metrics.goodput_bps,
});
if let Some(smoothed_loss) = metrics.smoothed_loss() {
link_layer["smoothed_loss"] = json!(smoothed_loss);
}
if let Some(smoothed_etx) = metrics.smoothed_etx() {
link_layer["smoothed_etx"] = json!(smoothed_etx);
}
if let Some(srtt) = metrics.srtt_ms() {
link_layer["srtt_ms"] = json!(srtt);
if let Some(setx) = metrics.smoothed_etx() {
link_layer["lqi"] = json!(setx * (1.0 + srtt / 100.0));
if let Some(smoothed_loss) = metrics.smoothed_loss() {
link_layer["smoothed_loss"] = json!(smoothed_loss);
}
if let Some(smoothed_etx) = metrics.smoothed_etx() {
link_layer["smoothed_etx"] = json!(smoothed_etx);
}
if let Some(srtt) = metrics.srtt_ms() {
link_layer["srtt_ms"] = json!(srtt);
if let Some(setx) = metrics.smoothed_etx() {
link_layer["lqi"] = json!(setx * (1.0 + srtt / 100.0));
}
}
}
// Trend indicators
if metrics.rtt_trend.initialized() {
link_layer["rtt_trend"] = json!(trend_label(metrics.rtt_trend.short(), metrics.rtt_trend.long()));
}
if metrics.loss_trend.initialized() {
link_layer["loss_trend"] = json!(trend_label(metrics.loss_trend.short(), metrics.loss_trend.long()));
}
if metrics.goodput_trend.initialized() {
link_layer["goodput_trend"] = json!(trend_label(metrics.goodput_trend.short(), metrics.goodput_trend.long()));
}
if metrics.jitter_trend.initialized() {
link_layer["jitter_trend"] = json!(trend_label(metrics.jitter_trend.short(), metrics.jitter_trend.long()));
}
// Trend indicators
if metrics.rtt_trend.initialized() {
link_layer["rtt_trend"] = json!(trend_label(
metrics.rtt_trend.short(),
metrics.rtt_trend.long()
));
}
if metrics.loss_trend.initialized() {
link_layer["loss_trend"] = json!(trend_label(
metrics.loss_trend.short(),
metrics.loss_trend.long()
));
}
if metrics.goodput_trend.initialized() {
link_layer["goodput_trend"] = json!(trend_label(
metrics.goodput_trend.short(),
metrics.goodput_trend.long()
));
}
if metrics.jitter_trend.initialized() {
link_layer["jitter_trend"] = json!(trend_label(
metrics.jitter_trend.short(),
metrics.jitter_trend.long()
));
}
link_layer["delivery_ratio_forward"] = json!(metrics.delivery_ratio_forward);
link_layer["delivery_ratio_reverse"] = json!(metrics.delivery_ratio_reverse);
link_layer["ecn_ce_count"] = json!(metrics.last_ecn_ce_count());
link_layer["delivery_ratio_forward"] = json!(metrics.delivery_ratio_forward);
link_layer["delivery_ratio_reverse"] = json!(metrics.delivery_ratio_reverse);
link_layer["ecn_ce_count"] = json!(metrics.last_ecn_ce_count());
Some(json!({
"peer": hex::encode(addr.as_bytes()),
"display_name": node.peer_display_name(&addr),
"mode": format!("{}", mmp.mode()),
"link_layer": link_layer,
}))
}).collect();
Some(json!({
"peer": hex::encode(addr.as_bytes()),
"display_name": node.peer_display_name(&addr),
"mode": format!("{}", mmp.mode()),
"link_layer": link_layer,
}))
})
.collect();
// Session-layer MMP
let sessions: Vec<Value> = node
-1
View File
@@ -402,5 +402,4 @@ mod tests {
assert_eq!(compute_etx(0.0, 1.0), 100.0);
assert_eq!(compute_etx(1.0, 0.0), 100.0);
}
}
+1 -3
View File
@@ -22,9 +22,7 @@ pub mod report;
pub mod sender;
// Re-exports
pub use algorithms::{
DualEwma, JitterEstimator, OwdTrendDetector, SrttEstimator, compute_etx,
};
pub use algorithms::{DualEwma, JitterEstimator, OwdTrendDetector, SrttEstimator, compute_etx};
pub use metrics::MmpMetrics;
pub use receiver::ReceiverState;
pub use report::{ReceiverReport, SenderReport};
+1 -3
View File
@@ -177,9 +177,7 @@ impl ReceiverReport {
});
}
if format_version > 0
&& total_length < RECEIVER_REPORT_PAYLOAD as usize
{
if format_version > 0 && total_length < RECEIVER_REPORT_PAYLOAD as usize {
return Err(ProtocolError::MessageTooShort {
expected: RECEIVER_REPORT_PAYLOAD as usize,
got: total_length,
+9 -17
View File
@@ -4,9 +4,9 @@
//! including delta compression with NACK-based recovery and debounced
//! propagation to peers.
use crate::NodeAddr;
use crate::bloom::BloomFilter;
use crate::protocol::{FilterAnnounce, FilterNack};
use crate::NodeAddr;
use super::{Node, NodeError};
use std::collections::HashMap;
@@ -96,11 +96,10 @@ impl Node {
announce.filter.clone()
};
let (encoded, stats) =
announce.encode().map_err(|e| NodeError::SendFailed {
node_addr: *peer_addr,
reason: format!("FilterAnnounce encode failed: {}", e),
})?;
let (encoded, stats) = announce.encode().map_err(|e| NodeError::SendFailed {
node_addr: *peer_addr,
reason: format!("FilterAnnounce encode failed: {}", e),
})?;
// Send
if let Err(e) = self.send_encrypted_link_message(peer_addr, &encoded).await {
@@ -127,8 +126,7 @@ impl Node {
"Sent FilterAnnounce"
);
self.bloom_state.record_update_sent(*peer_addr, now_ms);
self.bloom_state
.record_sent_filter(*peer_addr, sent_filter);
self.bloom_state.record_sent_filter(*peer_addr, sent_filter);
if let Some(peer) = self.peers.get_mut(peer_addr) {
peer.clear_filter_update_needed();
}
@@ -230,9 +228,7 @@ impl Node {
expected_seq: expected_base,
};
let nack_encoded = nack.encode();
let _ = self
.send_encrypted_link_message(from, &nack_encoded)
.await;
let _ = self.send_encrypted_link_message(from, &nack_encoded).await;
self.stats_mut().bloom.nacks_sent += 1;
return;
}
@@ -251,9 +247,7 @@ impl Node {
let nack = FilterNack {
expected_seq: current_seq,
};
let _ = self
.send_encrypted_link_message(from, &nack.encode())
.await;
let _ = self.send_encrypted_link_message(from, &nack.encode()).await;
self.stats_mut().bloom.nacks_sent += 1;
return;
}
@@ -266,9 +260,7 @@ impl Node {
"Delta received but no stored filter, sending NACK"
);
let nack = FilterNack { expected_seq: 0 };
let _ = self
.send_encrypted_link_message(from, &nack.encode())
.await;
let _ = self.send_encrypted_link_message(from, &nack.encode()).await;
self.stats_mut().bloom.nacks_sent += 1;
return;
}
+20 -30
View File
@@ -20,11 +20,7 @@ impl Node {
/// 4. Lazy purge expired entries
/// 5. If we're the target, generate and send response
/// 6. If TTL > 0, forward to tree peers whose bloom filter matches
pub(in crate::node) async fn handle_lookup_request(
&mut self,
from: &NodeAddr,
payload: &[u8],
) {
pub(in crate::node) async fn handle_lookup_request(&mut self, from: &NodeAddr, payload: &[u8]) {
self.stats_mut().discovery.req_received += 1;
let request = match LookupRequest::decode(payload) {
@@ -52,10 +48,8 @@ impl Node {
}
// Record for reverse-path forwarding and dedup
self.recent_requests.insert(
request.request_id,
RecentRequest::new(*from, now_ms),
);
self.recent_requests
.insert(request.request_id, RecentRequest::new(*from, now_ms));
// Lazy purge expired entries
self.purge_expired_requests(now_ms);
@@ -76,7 +70,10 @@ impl Node {
if request.can_forward() {
// Transit-side rate limit: collapse rapid-fire lookups for the
// same target from misbehaving nodes generating fresh request_ids.
if !self.discovery_forward_limiter.should_forward(&request.target) {
if !self
.discovery_forward_limiter
.should_forward(&request.target)
{
self.stats_mut().discovery.req_forward_rate_limited += 1;
debug!(
request_id = request.request_id,
@@ -192,11 +189,8 @@ impl Node {
// Verify the proof signature
let (xonly, _parity) = target_pubkey.x_only_public_key();
let peer_id = PeerIdentity::from_pubkey(xonly);
let proof_data = LookupResponse::proof_bytes(
response.request_id,
&target,
&response.target_coords,
);
let proof_data =
LookupResponse::proof_bytes(response.request_id, &target, &response.target_coords);
if !peer_id.verify(&proof_data, &response.proof) {
self.stats_mut().discovery.resp_proof_failed += 1;
warn!(
@@ -220,12 +214,8 @@ impl Node {
"Discovery succeeded, proof verified, route cached"
);
self.coord_cache.insert_with_path_mtu(
target,
response.target_coords,
now_ms,
path_mtu,
);
self.coord_cache
.insert_with_path_mtu(target, response.target_coords, now_ms, path_mtu);
// Clean up pending lookup tracking
self.pending_lookups.remove(&target);
@@ -262,15 +252,11 @@ impl Node {
let our_coords = self.tree_state().my_coords().clone();
// Sign proof: Identity::sign hashes with SHA-256 internally
let proof_data = LookupResponse::proof_bytes(request.request_id, &request.target, &our_coords);
let proof_data =
LookupResponse::proof_bytes(request.request_id, &request.target, &our_coords);
let proof = self.identity().sign(&proof_data);
let response = LookupResponse::new(
request.request_id,
request.target,
our_coords,
proof,
);
let response = LookupResponse::new(request.request_id, request.target, our_coords, proof);
// Route toward origin via reverse path.
let next_hop_addr = if let Some(recent) = self.recent_requests.get(&request.request_id) {
@@ -297,7 +283,10 @@ impl Node {
);
let encoded = response.encode();
if let Err(e) = self.send_encrypted_link_message(&next_hop_addr, &encoded).await {
if let Err(e) = self
.send_encrypted_link_message(&next_hop_addr, &encoded)
.await
{
debug!(
next_hop = %self.peer_display_name(&next_hop_addr),
error = %e,
@@ -503,7 +492,8 @@ impl Node {
return;
}
self.pending_lookups.insert(*dest, PendingLookup::new(now_ms));
self.pending_lookups
.insert(*dest, PendingLookup::new(now_ms));
let ttl = self.config.node.discovery.ttl;
let sent = self.initiate_lookup(dest, ttl).await;
+12 -9
View File
@@ -1,8 +1,8 @@
//! Encrypted frame handling (hot path).
use crate::noise::NoiseError;
use crate::node::Node;
use crate::node::wire::{EncryptedHeader, strip_inner_header, FLAG_CE, FLAG_KEY_EPOCH};
use crate::node::wire::{EncryptedHeader, FLAG_CE, FLAG_KEY_EPOCH, strip_inner_header};
use crate::noise::NoiseError;
use crate::transport::ReceivedPacket;
use std::time::Instant;
use tracing::{debug, info, trace, warn};
@@ -53,8 +53,8 @@ impl Node {
// Check and perform cutover in a scoped borrow.
{
let peer = self.peers.get(&node_addr).unwrap();
let k_bit_flipped = received_k_bit != peer.current_k_bit()
&& peer.pending_new_session().is_some();
let k_bit_flipped =
received_k_bit != peer.current_k_bit() && peer.pending_new_session().is_some();
if k_bit_flipped {
let display_name = self.peer_display_name(&node_addr);
@@ -70,9 +70,10 @@ impl Node {
debug_assert!(
peer.transport_id().is_some()
&& peer.our_index().is_some()
&& self.peers_by_index.contains_key(
&(peer.transport_id().unwrap(), peer.our_index().unwrap().as_u32())
),
&& self.peers_by_index.contains_key(&(
peer.transport_id().unwrap(),
peer.our_index().unwrap().as_u32()
)),
"peers_by_index should contain pre-registered new index after K-bit flip"
);
}
@@ -160,12 +161,14 @@ impl Node {
);
}
peer.set_current_addr(packet.transport_id, packet.remote_addr.clone());
peer.link_stats_mut().record_recv(packet.data.len(), packet.timestamp_ms);
peer.link_stats_mut()
.record_recv(packet.data.len(), packet.timestamp_ms);
peer.touch(packet.timestamp_ms);
}
// Dispatch to link message handler
self.dispatch_link_message(&node_addr, link_message, ce_flag).await;
self.dispatch_link_message(&node_addr, link_message, ce_flag)
.await;
}
/// Log a decryption failure with replay suppression.
+123 -113
View File
@@ -5,14 +5,12 @@
//! - msg2 (responder → initiator): responder identity + epoch + negotiation
//! - msg3 (initiator → responder): initiator identity + epoch + negotiation
use crate::PeerIdentity;
use crate::node::wire::{Msg1Header, Msg2Header, Msg3Header, build_msg2, build_msg3};
use crate::node::{Node, NodeError};
use crate::peer::{
cross_connection_winner, ActivePeer, PeerConnection, PromotionResult,
};
use crate::peer::{ActivePeer, PeerConnection, PromotionResult, cross_connection_winner};
use crate::protocol::NegotiationPayload;
use crate::transport::{Link, LinkDirection, LinkId, ReceivedPacket};
use crate::node::wire::{build_msg2, build_msg3, Msg1Header, Msg2Header, Msg3Header};
use crate::PeerIdentity;
use std::time::Duration;
use tracing::{debug, info, warn};
@@ -66,8 +64,7 @@ impl Node {
{
if link.direction() == LinkDirection::Inbound {
// Check if this link belongs to an already-promoted active peer
let is_active_peer = self.peers.values()
.any(|p| p.link_id() == existing_link_id);
let is_active_peer = self.peers.values().any(|p| p.link_id() == existing_link_id);
if !is_active_peer {
// Genuinely pending handshake — resend msg2
@@ -202,7 +199,8 @@ impl Node {
// Clean up on failure
self.connections.remove(&link_id);
self.links.remove(&link_id);
self.addr_to_link.remove(&(packet.transport_id, packet.remote_addr));
self.addr_to_link
.remove(&(packet.transport_id, packet.remote_addr));
let _ = self.index_allocator.free(our_index);
self.msg1_rate_limiter.complete_handshake();
return;
@@ -212,10 +210,8 @@ impl Node {
// XX: handshake NOT complete yet — need msg3.
// Store in pending_inbound for msg3 dispatch.
self.pending_inbound.insert(
(packet.transport_id, our_index.as_u32()),
link_id,
);
self.pending_inbound
.insert((packet.transport_id, our_index.as_u32()), link_id);
self.msg1_rate_limiter.complete_handshake();
}
@@ -279,9 +275,7 @@ impl Node {
// Find peer with rekey in progress for this index
let peer_addr = self.peers.iter().find_map(|(addr, peer)| {
if peer.rekey_in_progress()
&& peer.rekey_our_index() == Some(header.receiver_idx)
{
if peer.rekey_in_progress() && peer.rekey_our_index() == Some(header.receiver_idx) {
Some(*addr)
} else {
None
@@ -293,20 +287,24 @@ impl Node {
// Complete the rekey handshake on the ActivePeer
// XX: complete_rekey_msg2 processes msg2 and generates msg3
let transport_id = self.peers.get(&peer_node_addr)
let transport_id = self
.peers
.get(&peer_node_addr)
.and_then(|p| p.transport_id());
let remote_addr = self.peers.get(&peer_node_addr)
let remote_addr = self
.peers
.get(&peer_node_addr)
.and_then(|p| p.current_addr().cloned());
if let Some(peer) = self.peers.get_mut(&peer_node_addr) {
match peer.complete_rekey_msg2(noise_msg2) {
Ok((msg3_bytes, session)) => {
let our_index = peer.rekey_our_index()
.unwrap_or(header.receiver_idx);
let our_index = peer.rekey_our_index().unwrap_or(header.receiver_idx);
// Send msg3 before setting pending session
let wire_msg3 = build_msg3(our_index, header.sender_idx, &msg3_bytes);
let msg3_sent = if let (Some(tid), Some(addr)) = (transport_id, &remote_addr)
let msg3_sent = if let (Some(tid), Some(addr)) =
(transport_id, &remote_addr)
&& let Some(transport) = self.transports.get(&tid)
{
match transport.send(addr, &wire_msg3).await {
@@ -334,10 +332,8 @@ impl Node {
peer.set_pending_session(session, our_index, header.sender_idx);
if let Some(tid) = transport_id {
self.peers_by_index.insert(
(tid, our_index.as_u32()),
peer_node_addr,
);
self.peers_by_index
.insert((tid, our_index.as_u32()), peer_node_addr);
}
debug!(
@@ -388,18 +384,19 @@ impl Node {
// Process Noise msg2 and generate msg3
let noise_msg2 = &packet.data[header.noise_msg2_offset..];
let (msg3_bytes, received_negotiation) = match conn.complete_handshake(noise_msg2, Some(&neg_payload), packet.timestamp_ms) {
Ok(result) => result,
Err(e) => {
warn!(
link_id = %link_id,
error = %e,
"Handshake completion failed"
);
conn.mark_failed();
return;
}
};
let (msg3_bytes, received_negotiation) =
match conn.complete_handshake(noise_msg2, Some(&neg_payload), packet.timestamp_ms) {
Ok(result) => result,
Err(e) => {
warn!(
link_id = %link_id,
error = %e,
"Handshake completion failed"
);
conn.mark_failed();
return;
}
};
// Process peer's FMP negotiation payload from msg2
if let Some(neg_bytes) = &received_negotiation {
@@ -506,15 +503,17 @@ impl Node {
let outbound_our_index = conn.our_index();
let outbound_session = conn.take_session();
let (outbound_session, outbound_our_index) =
match (outbound_session, outbound_our_index) {
(Some(s), Some(idx)) => (s, idx),
_ => {
warn!(peer = %self.peer_display_name(&peer_node_addr), "Incomplete outbound connection");
self.pending_outbound.remove(&key);
return;
}
};
let (outbound_session, outbound_our_index) = match (
outbound_session,
outbound_our_index,
) {
(Some(s), Some(idx)) => (s, idx),
_ => {
warn!(peer = %self.peer_display_name(&peer_node_addr), "Incomplete outbound connection");
self.pending_outbound.remove(&key);
return;
}
};
if let Some(peer) = self.peers.get_mut(&peer_node_addr) {
let suppressed = peer.replay_suppressed_count();
@@ -527,13 +526,12 @@ impl Node {
// Update peers_by_index: remove old inbound index, add outbound
let transport_id = peer.transport_id().unwrap();
if let Some(old_idx) = old_our_index {
self.peers_by_index.remove(&(transport_id, old_idx.as_u32()));
self.peers_by_index
.remove(&(transport_id, old_idx.as_u32()));
let _ = self.index_allocator.free(old_idx);
}
self.peers_by_index.insert(
(transport_id, outbound_our_index.as_u32()),
peer_node_addr,
);
self.peers_by_index
.insert((transport_id, outbound_our_index.as_u32()), peer_node_addr);
if suppressed > 0 {
debug!(
@@ -624,7 +622,10 @@ impl Node {
self.bloom_state.mark_update_needed(node_addr);
self.reset_discovery_backoff();
}
PromotionResult::CrossConnectionWon { loser_link_id, node_addr } => {
PromotionResult::CrossConnectionWon {
loser_link_id,
node_addr,
} => {
// Close the losing TCP connection (no-op for connectionless)
if let Some(loser_link) = self.links.get(&loser_link_id) {
let loser_tid = loser_link.transport_id();
@@ -636,10 +637,8 @@ impl Node {
// Clean up the losing connection's link
self.remove_link(&loser_link_id);
// Ensure addr_to_link points to the winning link
self.addr_to_link.insert(
(packet.transport_id, packet.remote_addr.clone()),
link_id,
);
self.addr_to_link
.insert((packet.transport_id, packet.remote_addr.clone()), link_id);
debug!(
peer = %self.peer_display_name(&node_addr),
loser_link_id = %loser_link_id,
@@ -723,23 +722,24 @@ impl Node {
// Process msg3 — learns initiator's identity and epoch
let noise_msg3 = &packet.data[header.noise_msg3_offset..];
let received_negotiation = match conn.complete_handshake_msg3(noise_msg3, packet.timestamp_ms) {
Ok(neg) => neg,
Err(e) => {
warn!(
link_id = %link_id,
error = %e,
"Msg3 processing failed"
);
// Clean up
self.connections.remove(&link_id);
self.remove_link(&link_id);
if let Some(idx) = self.connections.get(&link_id).and_then(|c| c.our_index()) {
let _ = self.index_allocator.free(idx);
let received_negotiation =
match conn.complete_handshake_msg3(noise_msg3, packet.timestamp_ms) {
Ok(neg) => neg,
Err(e) => {
warn!(
link_id = %link_id,
error = %e,
"Msg3 processing failed"
);
// Clean up
self.connections.remove(&link_id);
self.remove_link(&link_id);
if let Some(idx) = self.connections.get(&link_id).and_then(|c| c.our_index()) {
let _ = self.index_allocator.free(idx);
}
return;
}
return;
}
};
};
// Process peer's FMP negotiation payload from msg3
if let Some(neg_bytes) = &received_negotiation {
@@ -805,10 +805,8 @@ impl Node {
_ => {
// Same epoch (or no epoch stored).
// Check for rekey: session must be at least 30s old.
let session_age_secs = existing_peer
.session_established_at()
.elapsed()
.as_secs();
let session_age_secs =
existing_peer.session_established_at().elapsed().as_secs();
if self.config.node.rekey.enabled
&& existing_peer.has_session()
&& existing_peer.is_healthy()
@@ -929,7 +927,9 @@ impl Node {
}
// Promote the connection to active peer.
let wire_msg2 = self.connections.get(&link_id)
let wire_msg2 = self
.connections
.get(&link_id)
.and_then(|c| c.handshake_msg2().map(|m| m.to_vec()));
debug!(
@@ -944,7 +944,9 @@ impl Node {
match result {
PromotionResult::Promoted(node_addr) => {
// Store msg2 on peer for resend on duplicate msg1
if let (Some(peer), Some(msg2)) = (self.peers.get_mut(&node_addr), wire_msg2) {
if let (Some(peer), Some(msg2)) =
(self.peers.get_mut(&node_addr), wire_msg2)
{
peer.set_handshake_msg2(msg2);
}
debug!(
@@ -961,9 +963,14 @@ impl Node {
self.bloom_state.mark_update_needed(node_addr);
self.reset_discovery_backoff();
}
PromotionResult::CrossConnectionWon { loser_link_id, node_addr } => {
PromotionResult::CrossConnectionWon {
loser_link_id,
node_addr,
} => {
// Store msg2 on peer for resend on duplicate msg1
if let (Some(peer), Some(msg2)) = (self.peers.get_mut(&node_addr), wire_msg2) {
if let (Some(peer), Some(msg2)) =
(self.peers.get_mut(&node_addr), wire_msg2)
{
peer.set_handshake_msg2(msg2);
}
// Close the losing TCP connection (no-op for connectionless)
@@ -1055,16 +1062,15 @@ impl Node {
if let Some(peer) = self.peers.get_mut(&peer_node_addr) {
match peer.complete_rekey_msg3(noise_msg3) {
Ok(session) => {
let our_index = peer.rekey_responder_our_index()
let our_index = peer
.rekey_responder_our_index()
.unwrap_or(header.receiver_idx);
peer.set_pending_session(session, our_index, header.sender_idx);
peer.record_peer_rekey();
if let Some(transport_id) = peer.transport_id() {
self.peers_by_index.insert(
(transport_id, our_index.as_u32()),
peer_node_addr,
);
self.peers_by_index
.insert((transport_id, our_index.as_u32()), peer_node_addr);
}
debug!(
@@ -1131,33 +1137,35 @@ impl Node {
.take_session()
.ok_or(NodeError::NoSession(link_id))?;
let our_index = connection.our_index().ok_or_else(|| {
NodeError::PromotionFailed {
let our_index = connection
.our_index()
.ok_or_else(|| NodeError::PromotionFailed {
link_id,
reason: "missing our_index".into(),
}
})?;
let their_index = connection.their_index().ok_or_else(|| {
NodeError::PromotionFailed {
})?;
let their_index = connection
.their_index()
.ok_or_else(|| NodeError::PromotionFailed {
link_id,
reason: "missing their_index".into(),
}
})?;
let transport_id = connection.transport_id().ok_or_else(|| {
NodeError::PromotionFailed {
})?;
let transport_id = connection
.transport_id()
.ok_or_else(|| NodeError::PromotionFailed {
link_id,
reason: "missing transport_id".into(),
}
})?;
let current_addr = connection.source_addr().ok_or_else(|| {
NodeError::PromotionFailed {
})?;
let current_addr = connection
.source_addr()
.ok_or_else(|| NodeError::PromotionFailed {
link_id,
reason: "missing source_addr".into(),
}
})?.clone();
})?
.clone();
let link_stats = connection.link_stats().clone();
let remote_epoch = connection.remote_epoch();
let peer_profile = connection.peer_profile()
let peer_profile = connection
.peer_profile()
.unwrap_or(crate::protocol::NodeProfile::Full);
let peer_node_addr = *verified_identity.node_addr();
@@ -1168,11 +1176,8 @@ impl Node {
let existing_link_id = existing_peer.link_id();
// Determine which connection wins
let this_wins = cross_connection_winner(
self.identity.node_addr(),
&peer_node_addr,
is_outbound,
);
let this_wins =
cross_connection_winner(self.identity.node_addr(), &peer_node_addr, is_outbound);
if this_wins {
// This connection wins, replace the existing peer
@@ -1183,8 +1188,7 @@ impl Node {
if let (Some(old_tid), Some(old_idx)) =
(old_peer.transport_id(), old_peer.our_index())
{
self.peers_by_index
.remove(&(old_tid, old_idx.as_u32()));
self.peers_by_index.remove(&(old_tid, old_idx.as_u32()));
let _ = self.index_allocator.free(old_idx);
}
@@ -1204,7 +1208,9 @@ impl Node {
self.node_profile,
peer_profile,
);
new_peer.set_tree_announce_min_interval_ms(self.config.node.tree.announce_min_interval_ms);
new_peer.set_tree_announce_min_interval_ms(
self.config.node.tree.announce_min_interval_ms,
);
self.peers.insert(peer_node_addr, new_peer);
self.peers_by_index
@@ -1276,13 +1282,17 @@ impl Node {
// Normal promotion
if self.max_peers > 0 && self.peers.len() >= self.max_peers {
let _ = self.index_allocator.free(our_index);
return Err(NodeError::MaxPeersExceeded { max: self.max_peers });
return Err(NodeError::MaxPeersExceeded {
max: self.max_peers,
});
}
// Preserve tree announce rate-limit state from old peer (if reconnecting).
// Without this, reconnection resets the rate limit window to zero,
// allowing an immediate announce that can feed an announce loop.
let old_announce_ts = self.peers.get(&peer_node_addr)
let old_announce_ts = self
.peers
.get(&peer_node_addr)
.map(|p| p.last_tree_announce_sent_ms());
let mut new_peer = ActivePeer::with_session(
@@ -1301,7 +1311,8 @@ impl Node {
self.node_profile,
peer_profile,
);
new_peer.set_tree_announce_min_interval_ms(self.config.node.tree.announce_min_interval_ms);
new_peer
.set_tree_announce_min_interval_ms(self.config.node.tree.announce_min_interval_ms);
if let Some(ts) = old_announce_ts {
new_peer.set_last_tree_announce_sent_ms(ts);
}
@@ -1329,7 +1340,6 @@ impl Node {
Ok(PromotionResult::Promoted(peer_node_addr))
}
}
}
/// Process an FMP negotiation payload received from a peer.
+10 -8
View File
@@ -1,10 +1,13 @@
//! RX event loop and packet dispatch.
use crate::control::{commands, ControlSocket};
use crate::control::queries;
use crate::control::{ControlSocket, commands};
use crate::node::wire::{
COMMON_PREFIX_SIZE, CommonPrefix, FMP_VERSION, PHASE_ESTABLISHED, PHASE_MSG1, PHASE_MSG2,
PHASE_MSG3,
};
use crate::node::{Node, NodeError};
use crate::transport::ReceivedPacket;
use crate::node::wire::{CommonPrefix, PHASE_ESTABLISHED, PHASE_MSG1, PHASE_MSG2, PHASE_MSG3, FMP_VERSION, COMMON_PREFIX_SIZE};
use std::time::Duration;
use tracing::{debug, info, warn};
@@ -30,8 +33,7 @@ impl Node {
/// This method takes ownership of the packet_rx channel and runs
/// until the channel is closed (typically when stop() is called).
pub async fn run_rx_loop(&mut self) -> Result<(), NodeError> {
let mut packet_rx = self.packet_rx.take()
.ok_or(NodeError::NotStarted)?;
let mut packet_rx = self.packet_rx.take().ok_or(NodeError::NotStarted)?;
// Take the TUN outbound receiver, or create a dummy channel that never
// produces messages (when TUN is disabled). Holding the sender prevents
@@ -54,12 +56,12 @@ impl Node {
}
};
let mut tick = tokio::time::interval(Duration::from_secs(self.config.node.tick_interval_secs));
let mut tick =
tokio::time::interval(Duration::from_secs(self.config.node.tick_interval_secs));
// Set up control socket channel
let (control_tx, mut control_rx) = tokio::sync::mpsc::channel::<
crate::control::ControlMessage,
>(32);
let (control_tx, mut control_rx) =
tokio::sync::mpsc::channel::<crate::control::ControlMessage>(32);
if self.config.node.control.enabled {
let config = self.config.node.control.clone();
+173 -101
View File
@@ -5,26 +5,26 @@
//! SessionSetup (Noise XX msg1), SessionAck (msg2), SessionMsg3 (msg3),
//! encrypted data, and error signals (CoordsRequired, PathBroken).
use crate::node::session::{EndToEndState, SessionEntry};
use crate::node::session_wire::{
build_fsp_header, fsp_prepend_inner_header, fsp_strip_inner_header,
parse_encrypted_coords, FspCommonPrefix, FspEncryptedHeader, FSP_COMMON_PREFIX_SIZE,
FSP_FLAG_CP, FSP_FLAG_K, FSP_HEADER_SIZE, FSP_PHASE_ESTABLISHED, FSP_PHASE_MSG1,
FSP_PHASE_MSG2, FSP_PHASE_MSG3, FSP_PORT_HEADER_SIZE, FSP_PORT_IPV6_SHIM,
};
use crate::protocol::{coords_wire_size, encode_coords};
use crate::upper::icmp::FIPS_OVERHEAD;
use crate::node::{Node, NodeError};
use crate::noise::{HandshakeState, HANDSHAKE_MSG1_SIZE, HANDSHAKE_MSG2_SIZE, HANDSHAKE_MSG3_SIZE};
use crate::protocol::NegotiationPayload;
use crate::NodeAddr;
use crate::mmp::report::ReceiverReport;
use crate::mmp::{MAX_SESSION_REPORT_INTERVAL_MS, MIN_SESSION_REPORT_INTERVAL_MS};
use crate::node::session::{EndToEndState, SessionEntry};
use crate::node::session_wire::{
FSP_COMMON_PREFIX_SIZE, FSP_FLAG_CP, FSP_FLAG_K, FSP_HEADER_SIZE, FSP_PHASE_ESTABLISHED,
FSP_PHASE_MSG1, FSP_PHASE_MSG2, FSP_PHASE_MSG3, FSP_PORT_HEADER_SIZE, FSP_PORT_IPV6_SHIM,
FspCommonPrefix, FspEncryptedHeader, build_fsp_header, fsp_prepend_inner_header,
fsp_strip_inner_header, parse_encrypted_coords,
};
use crate::node::{Node, NodeError};
use crate::noise::{HANDSHAKE_MSG1_SIZE, HANDSHAKE_MSG2_SIZE, HANDSHAKE_MSG3_SIZE, HandshakeState};
use crate::protocol::NegotiationPayload;
use crate::protocol::{
CoordsRequired, FspInnerFlags, MtuExceeded, PathBroken, PathMtuNotification, SessionAck,
SessionDatagram, SessionMessageType, SessionMsg3, SessionReceiverReport, SessionSenderReport,
SessionSetup,
};
use crate::NodeAddr;
use crate::protocol::{coords_wire_size, encode_coords};
use crate::upper::icmp::FIPS_OVERHEAD;
use secp256k1::PublicKey;
use tracing::{debug, info, trace};
@@ -49,7 +49,10 @@ impl Node {
let prefix = match FspCommonPrefix::parse(payload) {
Some(p) => p,
None => {
debug!(len = payload.len(), "Session payload too short for FSP prefix");
debug!(
len = payload.len(),
"Session payload too short for FSP prefix"
);
return;
}
};
@@ -90,7 +93,8 @@ impl Node {
}
}
FSP_PHASE_ESTABLISHED => {
self.handle_encrypted_session_msg(src_addr, payload, path_mtu, ce_flag).await;
self.handle_encrypted_session_msg(src_addr, payload, path_mtu, ce_flag)
.await;
}
_ => {
debug!(phase = prefix.phase, "Unknown FSP phase");
@@ -107,12 +111,21 @@ impl Node {
/// 4. AEAD decrypt with AAD = header_bytes
/// 5. Strip FSP inner header → timestamp, msg_type, inner_flags
/// 6. Dispatch by msg_type
async fn handle_encrypted_session_msg(&mut self, src_addr: &NodeAddr, payload: &[u8], path_mtu: u16, ce_flag: bool) {
async fn handle_encrypted_session_msg(
&mut self,
src_addr: &NodeAddr,
payload: &[u8],
path_mtu: u16,
ce_flag: bool,
) {
// Parse the 12-byte encrypted header (includes the 4-byte prefix)
let header = match FspEncryptedHeader::parse(payload) {
Some(h) => h,
None => {
debug!(len = payload.len(), "Encrypted session message too short for FSP header");
debug!(
len = payload.len(),
"Encrypted session message too short for FSP header"
);
return;
}
};
@@ -167,8 +180,8 @@ impl Node {
let received_k_bit = header.flags & FSP_FLAG_K != 0;
{
let entry = self.sessions.get(src_addr).unwrap();
let k_bit_flipped = received_k_bit != entry.current_k_bit()
&& entry.pending_new_session().is_some();
let k_bit_flipped =
received_k_bit != entry.current_k_bit() && entry.pending_new_session().is_some();
if k_bit_flipped {
let display_name = self.peer_display_name(src_addr);
@@ -235,7 +248,8 @@ impl Node {
self.sessions.insert(*src_addr, entry);
// Strip FSP inner header (6 bytes)
let (timestamp, msg_type, inner_flags_byte, rest) = match fsp_strip_inner_header(&plaintext) {
let (timestamp, msg_type, inner_flags_byte, rest) = match fsp_strip_inner_header(&plaintext)
{
Some(parts) => parts,
None => {
debug!(src = %self.peer_display_name(src_addr), "Decrypted payload too short for FSP inner header");
@@ -248,9 +262,8 @@ impl Node {
&& let Some(mmp) = entry.mmp_mut()
{
let now = std::time::Instant::now();
mmp.receiver.record_recv(
header.counter, timestamp, plaintext.len(), ce_flag, now,
);
mmp.receiver
.record_recv(header.counter, timestamp, plaintext.len(), ce_flag, now);
let _inner_flags = FspInnerFlags::from_byte(inner_flags_byte);
}
@@ -278,9 +291,15 @@ impl Node {
FSP_PORT_IPV6_SHIM => {
use crate::FipsAddress;
let src_ipv6 = FipsAddress::from_node_addr(src_addr).to_ipv6().octets();
let dst_ipv6 = FipsAddress::from_node_addr(self.node_addr()).to_ipv6().octets();
let dst_ipv6 = FipsAddress::from_node_addr(self.node_addr())
.to_ipv6()
.octets();
match crate::upper::ipv6_shim::decompress_ipv6(service_payload, src_ipv6, dst_ipv6) {
match crate::upper::ipv6_shim::decompress_ipv6(
service_payload,
src_ipv6,
dst_ipv6,
) {
Some(mut packet) => {
if ce_flag {
mark_ipv6_ecn_ce(&mut packet);
@@ -538,7 +557,13 @@ impl Node {
let placeholder_pubkey = self.identity.keypair().public_key();
let now_ms = Self::now_ms();
let resend_interval = self.config.node.rate_limit.handshake_resend_interval_ms;
let mut entry = SessionEntry::new(*src_addr, placeholder_pubkey, EndToEndState::AwaitingMsg3(handshake), now_ms, false);
let mut entry = SessionEntry::new(
*src_addr,
placeholder_pubkey,
EndToEndState::AwaitingMsg3(handshake),
now_ms,
false,
);
entry.set_handshake_payload(ack_payload, now_ms + resend_interval);
self.sessions.insert(*src_addr, entry);
@@ -653,7 +678,10 @@ impl Node {
// Split msg2 into base XX part and optional negotiation payload
let (base_msg2, neg_bytes) = if ack.handshake_payload.len() > HANDSHAKE_MSG2_SIZE {
(&ack.handshake_payload[..HANDSHAKE_MSG2_SIZE], Some(&ack.handshake_payload[HANDSHAKE_MSG2_SIZE..]))
(
&ack.handshake_payload[..HANDSHAKE_MSG2_SIZE],
Some(&ack.handshake_payload[HANDSHAKE_MSG2_SIZE..]),
)
} else {
(ack.handshake_payload.as_slice(), None)
};
@@ -831,7 +859,10 @@ impl Node {
// Split msg3 into base XX part and optional negotiation payload
let (base_msg3, neg_bytes) = if msg3.handshake_payload.len() > HANDSHAKE_MSG3_SIZE {
(&msg3.handshake_payload[..HANDSHAKE_MSG3_SIZE], Some(&msg3.handshake_payload[HANDSHAKE_MSG3_SIZE..]))
(
&msg3.handshake_payload[..HANDSHAKE_MSG3_SIZE],
Some(&msg3.handshake_payload[HANDSHAKE_MSG3_SIZE..]),
)
} else {
(msg3.handshake_payload.as_slice(), None)
};
@@ -878,7 +909,13 @@ impl Node {
let now_ms = Self::now_ms();
// Replace the placeholder pubkey with the real one
let mut new_entry = SessionEntry::new(*src_addr, remote_pubkey, EndToEndState::Established(session), now_ms, false);
let mut new_entry = SessionEntry::new(
*src_addr,
remote_pubkey,
EndToEndState::Established(session),
now_ms,
false,
);
new_entry.set_coords_warmup_remaining(self.config.node.session.coords_warmup_packets);
new_entry.mark_established(now_ms);
new_entry.init_mmp(&self.config.node.session_mmp);
@@ -948,7 +985,8 @@ impl Node {
};
let now = std::time::Instant::now();
mmp.metrics.process_receiver_report(&rr, our_timestamp_ms, now);
mmp.metrics
.process_receiver_report(&rr, our_timestamp_ms, now);
// Feed SRTT back to sender/receiver report interval tuning (session-layer bounds)
if let Some(srtt_ms) = mmp.metrics.srtt_ms() {
@@ -970,7 +1008,8 @@ impl Node {
// Update reverse delivery ratio from our own receiver state, using per-interval deltas.
let our_recv_packets = mmp.receiver.cumulative_packets_recv();
let peer_highest = mmp.receiver.highest_counter();
mmp.metrics.update_reverse_delivery(our_recv_packets, peer_highest);
mmp.metrics
.update_reverse_delivery(our_recv_packets, peer_highest);
trace!(
src = %peer_name,
@@ -1045,7 +1084,10 @@ impl Node {
);
// Send standalone CoordsWarmup immediately (rate-limited)
if self.coords_response_rate_limiter.should_send(&msg.dest_addr) {
if self
.coords_response_rate_limiter
.should_send(&msg.dest_addr)
{
if let Some(entry) = self.sessions.get(&msg.dest_addr)
&& entry.is_established()
&& let Err(e) = self.send_coords_warmup(&msg.dest_addr).await
@@ -1103,7 +1145,10 @@ impl Node {
);
// Send standalone CoordsWarmup immediately (rate-limited)
if self.coords_response_rate_limiter.should_send(&msg.dest_addr) {
if self
.coords_response_rate_limiter
.should_send(&msg.dest_addr)
{
if let Some(entry) = self.sessions.get(&msg.dest_addr)
&& entry.is_established()
&& let Err(e) = self.send_coords_warmup(&msg.dest_addr).await
@@ -1209,16 +1254,17 @@ impl Node {
let our_keypair = self.identity.keypair();
let mut handshake = HandshakeState::new_initiator(our_keypair);
handshake.set_local_epoch(self.startup_epoch);
let msg1 = handshake.write_message_1().map_err(|e| NodeError::SendFailed {
node_addr: dest_addr,
reason: format!("Noise XX msg1 generation failed: {}", e),
})?;
let msg1 = handshake
.write_message_1()
.map_err(|e| NodeError::SendFailed {
node_addr: dest_addr,
reason: format!("Noise XX msg1 generation failed: {}", e),
})?;
// Build SessionSetup with coordinates
let our_coords = self.tree_state.my_coords().clone();
let dest_coords = self.get_dest_coords(&dest_addr);
let setup = SessionSetup::new(our_coords, dest_coords)
.with_handshake(msg1);
let setup = SessionSetup::new(our_coords, dest_coords).with_handshake(msg1);
let setup_payload = setup.encode();
// Wrap in SessionDatagram
@@ -1235,7 +1281,13 @@ impl Node {
// Store session entry with handshake payload for potential resend
let now_ms = Self::now_ms();
let resend_interval = self.config.node.rate_limit.handshake_resend_interval_ms;
let mut entry = SessionEntry::new(dest_addr, dest_pubkey, EndToEndState::Initiating(handshake), now_ms, true);
let mut entry = SessionEntry::new(
dest_addr,
dest_pubkey,
EndToEndState::Initiating(handshake),
now_ms,
true,
);
entry.set_handshake_payload(setup_payload, now_ms + resend_interval);
self.sessions.insert(dest_addr, entry);
@@ -1262,10 +1314,13 @@ impl Node {
let now_ms = Self::now_ms();
// First borrow: read session metadata (NLL releases before coord decision)
let entry = self.sessions.get(dest_addr).ok_or_else(|| NodeError::SendFailed {
node_addr: *dest_addr,
reason: "no session".into(),
})?;
let entry = self
.sessions
.get(dest_addr)
.ok_or_else(|| NodeError::SendFailed {
node_addr: *dest_addr,
reason: "no session".into(),
})?;
let wants_coords = entry.coords_warmup_remaining() > 0;
let timestamp = entry.session_timestamp(now_ms);
if !entry.is_established() {
@@ -1284,7 +1339,8 @@ impl Node {
// Build inner plaintext (doesn't depend on counter)
let msg_type = SessionMessageType::DataPacket.to_byte(); // 0x10
let inner_flags = FspInnerFlags::new().to_byte();
let inner_plaintext = fsp_prepend_inner_header(timestamp, msg_type, inner_flags, &port_payload);
let inner_plaintext =
fsp_prepend_inner_header(timestamp, msg_type, inner_flags, &port_payload);
// Determine whether coords fit within transport MTU.
// If not, send standalone CoordsWarmup before the data packet.
@@ -1292,7 +1348,8 @@ impl Node {
let src = self.tree_state.my_coords().clone();
let dst = self.get_dest_coords(dest_addr);
let coords_size = coords_wire_size(&src) + coords_wire_size(&dst);
let total_wire = FIPS_OVERHEAD as usize + FSP_PORT_HEADER_SIZE + coords_size + payload.len();
let total_wire =
FIPS_OVERHEAD as usize + FSP_PORT_HEADER_SIZE + coords_size + payload.len();
if total_wire <= self.transport_mtu() as usize {
(true, Some(src), Some(dst))
} else {
@@ -1308,9 +1365,7 @@ impl Node {
};
// Decrement warmup counter if we sent coords (piggybacked or standalone)
if wants_coords
&& let Some(entry) = self.sessions.get_mut(dest_addr)
{
if wants_coords && let Some(entry) = self.sessions.get_mut(dest_addr) {
entry.set_coords_warmup_remaining(entry.coords_warmup_remaining() - 1);
}
@@ -1323,10 +1378,13 @@ impl Node {
}
// Borrow session for counter + encryption (after potential standalone send)
let entry = self.sessions.get_mut(dest_addr).ok_or_else(|| NodeError::SendFailed {
node_addr: *dest_addr,
reason: "no session".into(),
})?;
let entry = self
.sessions
.get_mut(dest_addr)
.ok_or_else(|| NodeError::SendFailed {
node_addr: *dest_addr,
reason: "no session".into(),
})?;
let session = match entry.state_mut() {
EndToEndState::Established(s) => s,
_ => {
@@ -1343,12 +1401,12 @@ impl Node {
let header = build_fsp_header(counter, flags, payload_len);
// Encrypt with AAD binding to the FSP header
let ciphertext = session.encrypt_with_aad(&inner_plaintext, &header).map_err(|e| {
NodeError::SendFailed {
let ciphertext = session
.encrypt_with_aad(&inner_plaintext, &header)
.map_err(|e| NodeError::SendFailed {
node_addr: *dest_addr,
reason: format!("session encrypt failed: {}", e),
}
})?;
})?;
// Assemble: header(12) + [coords] + ciphertext
let mut fsp_payload = Vec::with_capacity(FSP_HEADER_SIZE + ciphertext.len() + 200);
@@ -1386,13 +1444,19 @@ impl Node {
dest_addr: &NodeAddr,
ipv6_packet: &[u8],
) -> Result<(), NodeError> {
let compressed = crate::upper::ipv6_shim::compress_ipv6(ipv6_packet)
.ok_or_else(|| NodeError::SendFailed {
let compressed = crate::upper::ipv6_shim::compress_ipv6(ipv6_packet).ok_or_else(|| {
NodeError::SendFailed {
node_addr: *dest_addr,
reason: "IPv6 header compression failed".into(),
})?;
self.send_session_data(dest_addr, FSP_PORT_IPV6_SHIM, FSP_PORT_IPV6_SHIM, &compressed)
.await
}
})?;
self.send_session_data(
dest_addr,
FSP_PORT_IPV6_SHIM,
FSP_PORT_IPV6_SHIM,
&compressed,
)
.await
}
/// Send a non-data session message (reports, notifications) over an established session.
@@ -1411,19 +1475,25 @@ impl Node {
let now_ms = Self::now_ms();
// Read spin bit and session timestamp from entry
let entry = self.sessions.get(dest_addr).ok_or_else(|| NodeError::SendFailed {
node_addr: *dest_addr,
reason: "no session".into(),
})?;
let entry = self
.sessions
.get(dest_addr)
.ok_or_else(|| NodeError::SendFailed {
node_addr: *dest_addr,
reason: "no session".into(),
})?;
let timestamp = entry.session_timestamp(now_ms);
let inner_flags = FspInnerFlags::new().to_byte();
// Get mutable access for encryption
let entry = self.sessions.get_mut(dest_addr).ok_or_else(|| NodeError::SendFailed {
node_addr: *dest_addr,
reason: "no session".into(),
})?;
let entry = self
.sessions
.get_mut(dest_addr)
.ok_or_else(|| NodeError::SendFailed {
node_addr: *dest_addr,
reason: "no session".into(),
})?;
// Read K-bit before mutable borrow of session state
let k_flags = if entry.current_k_bit() { FSP_FLAG_K } else { 0 };
@@ -1448,12 +1518,12 @@ impl Node {
let header = build_fsp_header(counter, k_flags, payload_len);
// Encrypt with AAD
let ciphertext = session.encrypt_with_aad(&inner_plaintext, &header).map_err(|e| {
NodeError::SendFailed {
let ciphertext = session
.encrypt_with_aad(&inner_plaintext, &header)
.map_err(|e| NodeError::SendFailed {
node_addr: *dest_addr,
reason: format!("session encrypt failed: {}", e),
}
})?;
})?;
// Assemble: header(12) + ciphertext (no coords)
let mut fsp_payload = Vec::with_capacity(FSP_HEADER_SIZE + ciphertext.len());
@@ -1483,27 +1553,30 @@ impl Node {
/// coordinates via `try_warm_coord_cache()` (same as CP-flagged data
/// packets). The encrypted inner payload is the 6-byte inner header
/// with no application data.
async fn send_coords_warmup(
&mut self,
dest_addr: &NodeAddr,
) -> Result<(), NodeError> {
async fn send_coords_warmup(&mut self, dest_addr: &NodeAddr) -> Result<(), NodeError> {
let now_ms = Self::now_ms();
let my_coords = self.tree_state.my_coords().clone();
let dest_coords = self.get_dest_coords(dest_addr);
// Read session metadata
let entry = self.sessions.get(dest_addr).ok_or_else(|| NodeError::SendFailed {
node_addr: *dest_addr,
reason: "no session".into(),
})?;
let entry = self
.sessions
.get(dest_addr)
.ok_or_else(|| NodeError::SendFailed {
node_addr: *dest_addr,
reason: "no session".into(),
})?;
let timestamp = entry.session_timestamp(now_ms);
// Get mutable access for encryption
let entry = self.sessions.get_mut(dest_addr).ok_or_else(|| NodeError::SendFailed {
node_addr: *dest_addr,
reason: "no session".into(),
})?;
let entry = self
.sessions
.get_mut(dest_addr)
.ok_or_else(|| NodeError::SendFailed {
node_addr: *dest_addr,
reason: "no session".into(),
})?;
let session = match entry.state_mut() {
EndToEndState::Established(s) => s,
_ => {
@@ -1526,12 +1599,12 @@ impl Node {
let header = build_fsp_header(counter, FSP_FLAG_CP, payload_len);
// Encrypt with AAD
let ciphertext = session.encrypt_with_aad(&inner_plaintext, &header).map_err(|e| {
NodeError::SendFailed {
let ciphertext = session
.encrypt_with_aad(&inner_plaintext, &header)
.map_err(|e| NodeError::SendFailed {
node_addr: *dest_addr,
reason: format!("session encrypt failed: {}", e),
}
})?;
})?;
// Assemble: header(12) + coords + ciphertext
let coords_size = coords_wire_size(&my_coords) + coords_wire_size(&dest_coords);
@@ -1598,7 +1671,8 @@ impl Node {
}
let encoded = datagram.encode();
self.send_encrypted_link_message(&next_hop_addr, &encoded).await?;
self.send_encrypted_link_message(&next_hop_addr, &encoded)
.await?;
self.stats_mut().forwarding.record_originated(encoded.len());
Ok(())
}
@@ -1705,19 +1779,20 @@ impl Node {
/// Send ICMPv6 Destination Unreachable back through TUN.
pub(in crate::node) fn send_icmpv6_dest_unreachable(&self, original_packet: &[u8]) {
use crate::upper::icmp::{build_dest_unreachable, should_send_icmp_error, DestUnreachableCode};
use crate::FipsAddress;
use crate::upper::icmp::{
DestUnreachableCode, build_dest_unreachable, should_send_icmp_error,
};
if !should_send_icmp_error(original_packet) {
return;
}
let our_ipv6 = FipsAddress::from_node_addr(self.node_addr()).to_ipv6();
if let Some(response) = build_dest_unreachable(
original_packet,
DestUnreachableCode::NoRoute,
our_ipv6,
) && let Some(tun_tx) = &self.tun_tx {
if let Some(response) =
build_dest_unreachable(original_packet, DestUnreachableCode::NoRoute, our_ipv6)
&& let Some(tun_tx) = &self.tun_tx
{
let _ = tun_tx.send(response);
}
}
@@ -1774,10 +1849,7 @@ impl Node {
return;
}
let queue = self
.pending_tun_packets
.entry(dest_addr)
.or_default();
let queue = self.pending_tun_packets.entry(dest_addr).or_default();
if queue.len() >= self.config.node.session.pending_packets_per_dest {
queue.pop_front(); // Drop oldest
}
+22 -18
View File
@@ -22,7 +22,9 @@ impl Node {
.unwrap_or(0);
let timeout_ms = self.config.node.rate_limit.handshake_timeout_secs * 1000;
let stale: Vec<LinkId> = self.connections.iter()
let stale: Vec<LinkId> = self
.connections
.iter()
.filter(|(_, conn)| conn.is_timed_out(now_ms, timeout_ms) || conn.is_failed())
.map(|(link_id, _)| *link_id)
.collect();
@@ -100,14 +102,17 @@ impl Node {
// Skip resend if the target peer is already promoted — a cross-connection
// was resolved via the inbound path and resending msg1 would start a new
// handshake on the peer, creating a session mismatch.
let candidates: Vec<(LinkId, Vec<u8>)> = self.connections.iter()
let candidates: Vec<(LinkId, Vec<u8>)> = self
.connections
.iter()
.filter(|(_, conn)| {
conn.is_outbound()
&& conn.handshake_state() == HandshakeState::SentMsg1
&& conn.resend_count() < max_resends
&& conn.next_resend_at_ms() > 0
&& now_ms >= conn.next_resend_at_ms()
&& !conn.expected_identity()
&& !conn
.expected_identity()
.map(|id| self.peers.contains_key(id.node_addr()))
.unwrap_or(false)
})
@@ -143,9 +148,7 @@ impl Node {
false
};
if sent
&& let Some(conn) = self.connections.get_mut(&link_id)
{
if sent && let Some(conn) = self.connections.get_mut(&link_id) {
let count = conn.resend_count() + 1;
let next = now_ms + (interval_ms as f64 * backoff.powi(count as i32)) as u64;
conn.record_resend(next);
@@ -176,10 +179,11 @@ impl Node {
let ttl = self.config.node.session.default_ttl;
// First pass: find timed-out sessions to remove
let timed_out: Vec<crate::NodeAddr> = self.sessions.iter()
let timed_out: Vec<crate::NodeAddr> = self
.sessions
.iter()
.filter(|(_, entry)| {
!entry.is_established()
&& now_ms.saturating_sub(entry.last_activity()) > timeout_ms
!entry.is_established() && now_ms.saturating_sub(entry.last_activity()) > timeout_ms
})
.map(|(addr, _)| *addr)
.collect();
@@ -193,7 +197,9 @@ impl Node {
// Second pass: collect resend candidates
let my_addr = *self.node_addr();
let candidates: Vec<(crate::NodeAddr, Vec<u8>)> = self.sessions.iter()
let candidates: Vec<(crate::NodeAddr, Vec<u8>)> = self
.sessions
.iter()
.filter(|(_, entry)| {
!entry.is_established()
&& entry.handshake_payload().is_some()
@@ -207,8 +213,7 @@ impl Node {
for (dest_addr, payload) in candidates {
use crate::protocol::SessionDatagram;
let mut datagram = SessionDatagram::new(my_addr, dest_addr, payload)
.with_ttl(ttl);
let mut datagram = SessionDatagram::new(my_addr, dest_addr, payload).with_ttl(ttl);
let sent = match self.send_session_datagram(&mut datagram).await {
Ok(_) => true,
Err(e) => {
@@ -221,9 +226,7 @@ impl Node {
}
};
if sent
&& let Some(entry) = self.sessions.get_mut(&dest_addr)
{
if sent && let Some(entry) = self.sessions.get_mut(&dest_addr) {
let count = entry.resend_count() + 1;
let next = now_ms + (interval_ms as f64 * backoff.powi(count as i32)) as u64;
entry.record_resend(next);
@@ -246,10 +249,11 @@ impl Node {
return; // disabled
}
let idle: Vec<_> = self.sessions.iter()
let idle: Vec<_> = self
.sessions
.iter()
.filter(|(_, entry)| {
entry.is_established()
&& now_ms.saturating_sub(entry.last_activity()) > timeout_ms
entry.is_established() && now_ms.saturating_sub(entry.last_activity()) > timeout_ms
})
.map(|(addr, _)| *addr)
.collect();
+127 -43
View File
@@ -1,11 +1,11 @@
//! Node lifecycle management: start, stop, and peer connection initiation.
use super::{Node, NodeError, NodeState};
use crate::node::wire::build_msg1;
use crate::peer::PeerConnection;
use crate::protocol::{Disconnect, DisconnectReason};
use crate::transport::{packet_channel, Link, LinkDirection, LinkId, TransportAddr, TransportId};
use crate::upper::tun::{run_tun_reader, shutdown_tun_interface, TunDevice, TunState};
use crate::node::wire::build_msg1;
use crate::transport::{Link, LinkDirection, LinkId, TransportAddr, TransportId, packet_channel};
use crate::upper::tun::{TunDevice, TunState, run_tun_reader, shutdown_tun_interface};
use crate::{NodeAddr, PeerIdentity};
use std::thread;
use std::time::Duration;
@@ -50,7 +50,10 @@ impl Node {
return;
}
debug!(count = peer_configs.len(), "Initiating static peer connections");
debug!(
count = peer_configs.len(),
"Initiating static peer connections"
);
for peer_config in peer_configs {
if let Err(e) = self.initiate_peer_connection(&peer_config).await {
@@ -67,14 +70,16 @@ impl Node {
/// Initiate a connection to a single peer.
///
/// Creates a link, starts the Noise handshake, and sends the first message.
pub(super) async fn initiate_peer_connection(&mut self, peer_config: &crate::config::PeerConfig) -> Result<(), NodeError> {
pub(super) async fn initiate_peer_connection(
&mut self,
peer_config: &crate::config::PeerConfig,
) -> Result<(), NodeError> {
// Parse the peer's npub to get their identity
let peer_identity = PeerIdentity::from_npub(&peer_config.npub).map_err(|e| {
NodeError::InvalidPeerNpub {
let peer_identity =
PeerIdentity::from_npub(&peer_config.npub).map_err(|e| NodeError::InvalidPeerNpub {
npub: peer_config.npub.clone(),
reason: e.to_string(),
}
})?;
})?;
let peer_node_addr = *peer_identity.node_addr();
@@ -158,7 +163,10 @@ impl Node {
(tid, TransportAddr::from_string(&addr.addr))
};
match self.initiate_connection(transport_id, remote_addr, Some(peer_identity)).await {
match self
.initiate_connection(transport_id, remote_addr, Some(peer_identity))
.await
{
Ok(()) => return Ok(()),
Err(e) => {
debug!(
@@ -195,7 +203,9 @@ impl Node {
remote_addr: TransportAddr,
peer_identity: Option<PeerIdentity>,
) -> Result<(), NodeError> {
let is_connection_oriented = self.transports.get(&transport_id)
let is_connection_oriented = self
.transports
.get(&transport_id)
.map(|t| t.transport_type().connection_oriented)
.unwrap_or(false);
@@ -265,7 +275,8 @@ impl Node {
Ok(())
} else {
// Connectionless: proceed with immediate handshake
self.start_handshake(link_id, transport_id, remote_addr, peer_identity).await
self.start_handshake(link_id, transport_id, remote_addr, peer_identity)
.await
}
}
@@ -305,16 +316,17 @@ impl Node {
// Start the Noise handshake and get message 1
let our_keypair = self.identity.keypair();
let noise_msg1 = match connection.start_handshake(our_keypair, self.startup_epoch, current_time_ms) {
Ok(msg) => msg,
Err(e) => {
// Clean up the index and link
let _ = self.index_allocator.free(our_index);
self.links.remove(&link_id);
self.addr_to_link.remove(&(transport_id, remote_addr));
return Err(NodeError::HandshakeFailed(e.to_string()));
}
};
let noise_msg1 =
match connection.start_handshake(our_keypair, self.startup_epoch, current_time_ms) {
Ok(msg) => msg,
Err(e) => {
// Clean up the index and link
let _ = self.index_allocator.free(our_index);
self.links.remove(&link_id);
self.addr_to_link.remove(&(transport_id, remote_addr));
return Err(NodeError::HandshakeFailed(e.to_string()));
}
};
// Set index and transport info on the connection
connection.set_our_index(our_index);
@@ -348,7 +360,8 @@ impl Node {
connection.set_handshake_msg1(wire_msg1.clone(), current_time_ms + resend_interval);
// Track in pending_outbound for msg2 dispatch
self.pending_outbound.insert((transport_id, our_index.as_u32()), link_id);
self.pending_outbound
.insert((transport_id, our_index.as_u32()), link_id);
self.connections.insert(link_id, connection);
// Send the wire format handshake message
@@ -431,7 +444,10 @@ impl Node {
// Anonymous discovery (shared-media beacon without identity).
// Identity will be learned from XX handshake msg2.
// Dedup by transport address — skip if link already exists.
if self.addr_to_link.contains_key(&(*transport_id, peer.addr.clone())) {
if self
.addr_to_link
.contains_key(&(*transport_id, peer.addr.clone()))
{
continue;
}
@@ -455,7 +471,10 @@ impl Node {
"Auto-connecting to anonymous discovered peer"
);
}
if let Err(e) = self.initiate_connection(transport_id, remote_addr, identity).await {
if let Err(e) = self
.initiate_connection(transport_id, remote_addr, identity)
.await
{
warn!(error = %e, "Failed to auto-connect to discovered peer");
}
}
@@ -516,12 +535,15 @@ impl Node {
);
// Start the handshake now that the transport is connected
if let Err(e) = self.start_handshake(
pending.link_id,
pending.transport_id,
pending.remote_addr.clone(),
pending.peer_identity,
).await {
if let Err(e) = self
.start_handshake(
pending.link_id,
pending.transport_id,
pending.remote_addr.clone(),
pending.peer_identity,
)
.await
{
warn!(
link_id = %pending.link_id,
error = %e,
@@ -620,10 +642,24 @@ impl Node {
// Calculate max MSS for TCP clamping
let effective_mtu = self.effective_ipv6_mtu();
let max_mss = effective_mtu.saturating_sub(40).saturating_sub(20); // IPv6 + TCP headers
info!("effective MTU: {} bytes", effective_mtu);
debug!(" max TCP MSS: {} bytes", max_mss);
// On macOS, create a shutdown pipe. Writing to it unblocks the
// reader thread's select() loop without closing the TUN fd
// (which would cause a double-close when TunDevice drops).
#[cfg(target_os = "macos")]
let (shutdown_read_fd, shutdown_write_fd) = {
let mut fds = [0i32; 2];
if unsafe { libc::pipe(fds.as_mut_ptr()) } < 0 {
return Err(NodeError::Tun(crate::upper::tun::TunError::Configure(
"failed to create shutdown pipe".into(),
)));
}
(fds[0], fds[1])
};
// Create writer (dups the fd for independent write access)
let (writer, tun_tx) = device.create_writer(max_mss)?;
@@ -641,8 +677,28 @@ impl Node {
// Spawn reader thread
let transport_mtu = self.transport_mtu();
#[cfg(target_os = "macos")]
let reader_handle = thread::spawn(move || {
run_tun_reader(device, mtu, our_addr, reader_tun_tx, outbound_tx, transport_mtu);
run_tun_reader(
device,
mtu,
our_addr,
reader_tun_tx,
outbound_tx,
transport_mtu,
shutdown_read_fd,
);
});
#[cfg(not(target_os = "macos"))]
let reader_handle = thread::spawn(move || {
run_tun_reader(
device,
mtu,
our_addr,
reader_tun_tx,
outbound_tx,
transport_mtu,
);
});
self.tun_state = TunState::Active;
@@ -651,6 +707,10 @@ impl Node {
self.tun_outbound_rx = Some(outbound_rx);
self.tun_reader_handle = Some(reader_handle);
self.tun_writer_handle = Some(writer_handle);
#[cfg(target_os = "macos")]
{
self.tun_shutdown_fd = Some(shutdown_write_fd);
}
}
Err(e) => {
self.tun_state = TunState::Failed;
@@ -667,11 +727,19 @@ impl Node {
let dns_channel_size = self.config.node.buffers.dns_channel;
let (identity_tx, identity_rx) = tokio::sync::mpsc::channel(dns_channel_size);
let dns_ttl = self.config.dns.ttl();
let base_hosts = crate::upper::hosts::HostMap::from_peer_configs(self.config.peers());
let hosts_path = std::path::PathBuf::from(crate::upper::hosts::DEFAULT_HOSTS_PATH);
let reloader = crate::upper::hosts::HostMapReloader::new(base_hosts, hosts_path);
let base_hosts =
crate::upper::hosts::HostMap::from_peer_configs(self.config.peers());
let hosts_path =
std::path::PathBuf::from(crate::upper::hosts::DEFAULT_HOSTS_PATH);
let reloader =
crate::upper::hosts::HostMapReloader::new(base_hosts, hosts_path);
info!(bind = %bind, hosts = reloader.hosts().len(), "DNS responder started for .fips domain (auto-reload enabled)");
let handle = tokio::spawn(crate::upper::dns::run_dns_responder(socket, identity_tx, dns_ttl, reloader));
let handle = tokio::spawn(crate::upper::dns::run_dns_responder(
socket,
identity_tx,
dns_ttl,
reloader,
));
self.dns_identity_rx = Some(identity_rx);
self.dns_task = Some(handle);
}
@@ -707,7 +775,8 @@ impl Node {
}
// Send disconnect notifications to all active peers before closing transports
self.send_disconnect_to_all_peers(DisconnectReason::Shutdown).await;
self.send_disconnect_to_all_peers(DisconnectReason::Shutdown)
.await;
// Shutdown transports (they're packet producers)
let transport_ids: Vec<_> = self.transports.keys().cloned().collect();
@@ -741,11 +810,21 @@ impl Node {
// Drop the tun_tx to signal the writer to stop
self.tun_tx.take();
// Delete the interface (causes reader to get EFAULT)
// Delete the interface (on Linux, causes reader to get EFAULT)
if let Err(e) = shutdown_tun_interface(&name).await {
warn!(name = %name, error = %e, "Failed to shutdown TUN interface");
}
// On macOS, signal the reader thread to exit by writing to the
// shutdown pipe. The reader's select() will wake up and break.
#[cfg(target_os = "macos")]
if let Some(fd) = self.tun_shutdown_fd.take() {
unsafe {
libc::write(fd, b"x".as_ptr() as *const libc::c_void, 1);
libc::close(fd);
}
}
// Wait for threads to finish
if let Some(handle) = self.tun_reader_handle.take() {
let _ = handle.join();
@@ -771,7 +850,9 @@ impl Node {
let plaintext = disconnect.encode();
// Collect node_addrs to avoid borrow conflict with send helper
let peer_addrs: Vec<NodeAddr> = self.peers.iter()
let peer_addrs: Vec<NodeAddr> = self
.peers
.iter()
.filter(|(_, peer)| peer.can_send() && peer.has_session())
.map(|(addr, _)| *addr)
.collect();
@@ -786,7 +867,10 @@ impl Node {
let mut sent = 0usize;
for node_addr in &peer_addrs {
match self.send_encrypted_link_message(node_addr, &plaintext).await {
match self
.send_encrypted_link_message(node_addr, &plaintext)
.await
{
Ok(()) => sent += 1,
Err(e) => {
debug!(
@@ -851,8 +935,8 @@ impl Node {
///
/// Removes the peer and suppresses auto-reconnect.
pub(crate) fn api_disconnect(&mut self, npub: &str) -> Result<serde_json::Value, String> {
let peer_identity = PeerIdentity::from_npub(npub)
.map_err(|e| format!("invalid npub '{npub}': {e}"))?;
let peer_identity =
PeerIdentity::from_npub(npub).map_err(|e| format!("invalid npub '{npub}': {e}"))?;
let node_addr = *peer_identity.node_addr();
if !self.peers.contains_key(&node_addr) {
+115 -88
View File
@@ -5,42 +5,43 @@
//! Bloom filters, coordinate caches, transports, links, and peers.
mod bloom;
mod discovery_rate_limit;
mod handlers;
mod lifecycle;
mod retry;
mod discovery_rate_limit;
mod rate_limit;
mod retry;
mod routing_error_rate_limit;
pub(crate) mod session;
pub(crate) mod session_wire;
pub(crate) mod wire;
pub(crate) mod stats;
mod tree;
#[cfg(test)]
mod tests;
mod tree;
pub(crate) mod wire;
use crate::bloom::BloomState;
use crate::protocol::NodeProfile;
use crate::cache::CoordCache;
use crate::utils::index::IndexAllocator;
use crate::node::session::SessionEntry;
use crate::peer::{ActivePeer, PeerConnection};
use self::discovery_rate_limit::{DiscoveryBackoff, DiscoveryForwardRateLimiter};
use self::rate_limit::HandshakeRateLimiter;
use self::routing_error_rate_limit::RoutingErrorRateLimiter;
use self::wire::{
FLAG_CE, FLAG_KEY_EPOCH, build_encrypted, build_established_header, prepend_inner_header,
};
use crate::bloom::BloomState;
use crate::cache::CoordCache;
use crate::node::session::SessionEntry;
use crate::peer::{ActivePeer, PeerConnection};
use crate::protocol::NodeProfile;
use crate::transport::ethernet::EthernetTransport;
use crate::transport::tcp::TcpTransport;
use crate::transport::tor::TorTransport;
use crate::transport::udp::UdpTransport;
use crate::transport::{
Link, LinkId, PacketRx, PacketTx, TransportAddr, TransportError, TransportHandle, TransportId,
};
use crate::transport::udp::UdpTransport;
use crate::transport::tcp::TcpTransport;
use crate::transport::tor::TorTransport;
#[cfg(target_os = "linux")]
use crate::transport::ethernet::EthernetTransport;
use crate::tree::TreeState;
use crate::upper::hosts::HostMap;
use crate::upper::icmp_rate_limit::IcmpRateLimiter;
use crate::upper::tun::{TunError, TunOutboundRx, TunState, TunTx};
use self::wire::{build_encrypted, build_established_header, prepend_inner_header, FLAG_CE, FLAG_KEY_EPOCH};
use crate::utils::index::IndexAllocator;
use crate::{Config, ConfigError, Identity, IdentityError, NodeAddr, PeerIdentity};
use rand::Rng;
use std::collections::{HashMap, VecDeque};
@@ -107,7 +108,11 @@ pub enum NodeError {
SendFailed { node_addr: NodeAddr, reason: String },
#[error("mtu exceeded forwarding to {node_addr}: packet {packet_size} > mtu {mtu}")]
MtuExceeded { node_addr: NodeAddr, packet_size: usize, mtu: u16 },
MtuExceeded {
node_addr: NodeAddr,
packet_size: usize,
mtu: u16,
},
#[error("config error: {0}")]
Config(#[from] ConfigError),
@@ -370,6 +375,10 @@ pub struct Node {
tun_reader_handle: Option<JoinHandle<()>>,
/// TUN writer thread handle.
tun_writer_handle: Option<JoinHandle<()>>,
/// Shutdown pipe: writing to this fd unblocks the TUN reader thread on macOS.
/// On Linux, deleting the interface via netlink serves the same purpose.
#[cfg(target_os = "macos")]
tun_shutdown_fd: Option<std::os::unix::io::RawFd>,
// === DNS Responder ===
/// Receiver for resolved identities from the DNS responder.
@@ -542,6 +551,8 @@ impl Node {
tun_outbound_rx: None,
tun_reader_handle: None,
tun_writer_handle: None,
#[cfg(target_os = "macos")]
tun_shutdown_fd: None,
dns_identity_rx: None,
dns_task: None,
index_allocator: IndexAllocator::new(),
@@ -554,10 +565,7 @@ impl Node {
coords_response_rate_limiter: RoutingErrorRateLimiter::with_interval(
std::time::Duration::from_millis(coords_response_interval_ms),
),
discovery_backoff: DiscoveryBackoff::with_params(
backoff_base_secs,
backoff_max_secs,
),
discovery_backoff: DiscoveryBackoff::with_params(backoff_base_secs, backoff_max_secs),
discovery_forward_limiter: DiscoveryForwardRateLimiter::with_interval(
std::time::Duration::from_secs(forward_min_interval_secs),
),
@@ -654,6 +662,8 @@ impl Node {
tun_outbound_rx: None,
tun_reader_handle: None,
tun_writer_handle: None,
#[cfg(target_os = "macos")]
tun_shutdown_fd: None,
dns_identity_rx: None,
dns_task: None,
index_allocator: IndexAllocator::new(),
@@ -711,21 +721,18 @@ impl Node {
}
// Create Ethernet transport instances
#[cfg(target_os = "linux")]
{
let eth_instances: Vec<_> = self
.config
.transports
.ethernet
.iter()
.map(|(name, config)| (name.map(|s| s.to_string()), config.clone()))
.collect();
let eth_instances: Vec<_> = self
.config
.transports
.ethernet
.iter()
.map(|(name, config)| (name.map(|s| s.to_string()), config.clone()))
.collect();
for (name, eth_config) in eth_instances {
let transport_id = self.allocate_transport_id();
let eth = EthernetTransport::new(transport_id, name, eth_config, packet_tx.clone());
transports.push(TransportHandle::Ethernet(eth));
}
for (name, eth_config) in eth_instances {
let transport_id = self.allocate_transport_id();
let eth = EthernetTransport::new(transport_id, name, eth_config, packet_tx.clone());
transports.push(TransportHandle::Ethernet(eth));
}
// Create TCP transport instances
@@ -794,7 +801,9 @@ impl Node {
#[cfg(any(not(feature = "ble"), test))]
if !ble_instances.is_empty() {
#[cfg(not(test))]
tracing::warn!("BLE transport configured but 'ble' feature not enabled at compile time");
tracing::warn!(
"BLE transport configured but 'ble' feature not enabled at compile time"
);
}
}
@@ -844,18 +853,9 @@ impl Node {
))
})?;
// Parse the MAC address
#[cfg(target_os = "linux")]
let mac = crate::transport::ethernet::parse_mac_string(mac_str).map_err(|e| {
NodeError::NoTransportForType(format!("invalid MAC in '{}': {}", addr_str, e))
})?;
#[cfg(not(target_os = "linux"))]
let mac: [u8; 6] = {
let _ = mac_str;
return Err(NodeError::NoTransportForType(
"Ethernet transport not available on this platform".into(),
));
};
Ok((transport_id, TransportAddr::from_bytes(&mac)))
}
@@ -864,13 +864,9 @@ impl Node {
/// (TransportId, TransportAddr) pair by finding the BLE transport
/// instance matching the adapter name.
#[cfg(target_os = "linux")]
fn resolve_ble_addr(
&self,
addr_str: &str,
) -> Result<(TransportId, TransportAddr), NodeError> {
fn resolve_ble_addr(&self, addr_str: &str) -> Result<(TransportId, TransportAddr), NodeError> {
let ta = TransportAddr::from_string(addr_str);
let adapter = crate::transport::ble::addr::adapter_from_addr(&ta)
.ok_or_else(|| {
let adapter = crate::transport::ble::addr::adapter_from_addr(&ta).ok_or_else(|| {
NodeError::NoTransportForType(format!(
"invalid BLE address format '{}': expected 'adapter/mac'",
addr_str
@@ -881,9 +877,7 @@ impl Node {
let transport_id = self
.transports
.iter()
.find(|(_, handle)| {
handle.transport_type().name == "ble" && handle.is_operational()
})
.find(|(_, handle)| handle.transport_type().name == "ble" && handle.is_operational())
.map(|(id, _)| *id)
.ok_or_else(|| {
NodeError::NoTransportForType(format!(
@@ -1095,9 +1089,10 @@ impl Node {
let now = std::time::Instant::now();
let should_log = match self.last_mesh_size_log {
None => true,
Some(last) => now.duration_since(last) >= std::time::Duration::from_secs(
self.config.node.mmp.log_interval_secs,
),
Some(last) => {
now.duration_since(last)
>= std::time::Duration::from_secs(self.config.node.mmp.log_interval_secs)
}
};
if should_log {
tracing::debug!(
@@ -1146,7 +1141,6 @@ impl Node {
self.tun_name.as_deref()
}
// === Resource Limits ===
/// Set the maximum number of connections (handshake phase).
@@ -1227,14 +1221,17 @@ impl Node {
/// Add a link.
pub fn add_link(&mut self, link: Link) -> Result<(), NodeError> {
if self.max_links > 0 && self.links.len() >= self.max_links {
return Err(NodeError::MaxLinksExceeded { max: self.max_links });
return Err(NodeError::MaxLinksExceeded {
max: self.max_links,
});
}
let link_id = link.link_id();
let transport_id = link.transport_id();
let remote_addr = link.remote_addr().clone();
self.links.insert(link_id, link);
self.addr_to_link.insert((transport_id, remote_addr), link_id);
self.addr_to_link
.insert((transport_id, remote_addr), link_id);
Ok(())
}
@@ -1249,8 +1246,14 @@ impl Node {
}
/// Find link ID by transport address.
pub fn find_link_by_addr(&self, transport_id: TransportId, addr: &TransportAddr) -> Option<LinkId> {
self.addr_to_link.get(&(transport_id, addr.clone())).copied()
pub fn find_link_by_addr(
&self,
transport_id: TransportId,
addr: &TransportAddr,
) -> Option<LinkId> {
self.addr_to_link
.get(&(transport_id, addr.clone()))
.copied()
}
/// Remove a link.
@@ -1396,11 +1399,14 @@ impl Node {
pub(crate) fn register_identity(&mut self, node_addr: NodeAddr, pubkey: secp256k1::PublicKey) {
let mut prefix = [0u8; 15];
prefix.copy_from_slice(&node_addr.as_bytes()[0..15]);
self.identity_cache.insert(prefix, (node_addr, pubkey, Self::now_ms()));
self.identity_cache
.insert(prefix, (node_addr, pubkey, Self::now_ms()));
// LRU eviction
let max = self.config.node.cache.identity_size;
if self.identity_cache.len() > max
&& let Some(oldest_key) = self.identity_cache.iter()
&& let Some(oldest_key) = self
.identity_cache
.iter()
.min_by_key(|(_, (_, _, ts))| *ts)
.map(|(k, _)| *k)
{
@@ -1409,7 +1415,10 @@ impl Node {
}
/// Look up a destination by FipsAddress prefix (bytes 1-15 of the IPv6 address).
pub(crate) fn lookup_by_fips_prefix(&mut self, prefix: &[u8; 15]) -> Option<(NodeAddr, secp256k1::PublicKey)> {
pub(crate) fn lookup_by_fips_prefix(
&mut self,
prefix: &[u8; 15],
) -> Option<(NodeAddr, secp256k1::PublicKey)> {
if let Some(entry) = self.identity_cache.get_mut(prefix) {
entry.2 = Self::now_ms(); // LRU touch
Some((entry.0, entry.1))
@@ -1448,9 +1457,7 @@ impl Node {
/// has declared us as their parent (making them our child).
pub(crate) fn is_tree_peer(&self, peer_addr: &NodeAddr) -> bool {
// Peer is our parent
if !self.tree_state.is_root()
&& self.tree_state.my_declaration().parent_id() == peer_addr
{
if !self.tree_state.is_root() && self.tree_state.my_declaration().parent_id() == peer_addr {
return true;
}
// Peer is our child (their declaration names us as parent)
@@ -1499,12 +1506,18 @@ impl Node {
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_millis() as u64)
.unwrap_or(0);
let dest_coords = self.coord_cache.get_and_touch(dest_node_addr, now_ms)?.clone();
let dest_coords = self
.coord_cache
.get_and_touch(dest_node_addr, now_ms)?
.clone();
// 3. Bloom filter candidates — requires dest_coords for loop-free selection
// 3. Bloom filter candidates — requires dest_coords for loop-free selection.
// If no candidate is strictly closer, fall through to tree routing.
let candidates: Vec<&ActivePeer> = self.destination_in_filters(dest_node_addr);
if !candidates.is_empty() {
return self.select_best_candidate(&candidates, &dest_coords);
if !candidates.is_empty()
&& let Some(peer) = self.select_best_candidate(&candidates, &dest_coords)
{
return Some(peer);
}
// 4. Greedy tree routing fallback (skip non-routing/leaf peers)
@@ -1605,7 +1618,8 @@ impl Node {
node_addr: &NodeAddr,
plaintext: &[u8],
) -> Result<(), NodeError> {
self.send_encrypted_link_message_with_ce(node_addr, plaintext, false).await
self.send_encrypted_link_message_with_ce(node_addr, plaintext, false)
.await
}
/// Like `send_encrypted_link_message` but allows setting the FMP CE flag.
@@ -1617,7 +1631,9 @@ impl Node {
plaintext: &[u8],
ce_flag: bool,
) -> Result<(), NodeError> {
let peer = self.peers.get_mut(node_addr)
let peer = self
.peers
.get_mut(node_addr)
.ok_or(NodeError::PeerNotFound(*node_addr))?;
let their_index = peer.their_index().ok_or_else(|| NodeError::SendFailed {
@@ -1628,10 +1644,13 @@ impl Node {
node_addr: *node_addr,
reason: "no transport_id".into(),
})?;
let remote_addr = peer.current_addr().cloned().ok_or_else(|| NodeError::SendFailed {
node_addr: *node_addr,
reason: "no current_addr".into(),
})?;
let remote_addr = peer
.current_addr()
.cloned()
.ok_or_else(|| NodeError::SendFailed {
node_addr: *node_addr,
reason: "no current_addr".into(),
})?;
// Prepend 4-byte session-relative timestamp (inner header)
let timestamp_ms = peer.session_elapsed_ms();
@@ -1644,10 +1663,12 @@ impl Node {
flags |= FLAG_KEY_EPOCH;
}
let session = peer.noise_session_mut().ok_or_else(|| NodeError::SendFailed {
node_addr: *node_addr,
reason: "no noise session".into(),
})?;
let session = peer
.noise_session_mut()
.ok_or_else(|| NodeError::SendFailed {
node_addr: *node_addr,
reason: "no noise session".into(),
})?;
// Inner plaintext: [timestamp:4 LE][msg_type][payload...]
let inner_plaintext = prepend_inner_header(timestamp_ms, plaintext);
@@ -1658,18 +1679,24 @@ impl Node {
let header = build_established_header(their_index, counter, flags, payload_len);
// Encrypt with AAD binding to the outer header
let ciphertext = session.encrypt_with_aad(&inner_plaintext, &header).map_err(|e| NodeError::SendFailed {
node_addr: *node_addr,
reason: format!("encryption failed: {}", e),
})?;
let ciphertext = session
.encrypt_with_aad(&inner_plaintext, &header)
.map_err(|e| NodeError::SendFailed {
node_addr: *node_addr,
reason: format!("encryption failed: {}", e),
})?;
let wire_packet = build_encrypted(&header, &ciphertext);
// Re-borrow peer for stats update after sending
let transport = self.transports.get(&transport_id)
let transport = self
.transports
.get(&transport_id)
.ok_or(NodeError::TransportNotFound(transport_id))?;
let bytes_sent = transport.send(&remote_addr, &wire_packet).await
let bytes_sent = transport
.send(&remote_addr, &wire_packet)
.await
.map_err(|e| match e {
TransportError::MtuExceeded { packet_size, mtu } => NodeError::MtuExceeded {
node_addr: *node_addr,
-1
View File
@@ -89,7 +89,6 @@ pub const FSP_FLAG_K: u8 = 0x02;
/// Unencrypted — payload is plaintext (error signals).
pub const FSP_FLAG_U: u8 = 0x04;
// ============================================================================
// Common Prefix
// ============================================================================
+14 -4
View File
@@ -284,8 +284,16 @@ async fn test_disconnect_clears_session() {
nodes[1].node.sessions.insert(node0_addr, entry);
}
assert_eq!(nodes[1].node.session_count(), 1, "Session should exist before disconnect");
assert_eq!(nodes[1].node.peer_count(), 1, "Peer should exist before disconnect");
assert_eq!(
nodes[1].node.session_count(),
1,
"Session should exist before disconnect"
);
assert_eq!(
nodes[1].node.peer_count(),
1,
"Peer should exist before disconnect"
);
// Node 0 sends Disconnect to node 1.
let disconnect = crate::protocol::Disconnect::new(DisconnectReason::Shutdown);
@@ -300,7 +308,8 @@ async fn test_disconnect_clears_session() {
// Peer must be gone.
assert_eq!(
nodes[1].node.peer_count(), 0,
nodes[1].node.peer_count(),
0,
"Peer should be removed after disconnect"
);
@@ -308,7 +317,8 @@ async fn test_disconnect_clears_session() {
// Before the fix, session_count() would still be 1 here because
// remove_active_peer didn't remove self.sessions[node0_addr].
assert_eq!(
nodes[1].node.session_count(), 0,
nodes[1].node.session_count(),
0,
"Session must be cleaned up when peer is removed (regression: issue #5)"
);
+241 -109
View File
@@ -5,9 +5,11 @@ use super::*;
#[tokio::test]
async fn test_two_node_handshake_udp() {
use crate::config::UdpConfig;
use crate::node::wire::{
build_encrypted, build_established_header, build_msg1, prepend_inner_header,
};
use crate::transport::udp::UdpTransport;
use crate::node::wire::{build_encrypted, build_established_header, build_msg1, prepend_inner_header};
use tokio::time::{timeout, Duration};
use tokio::time::{Duration, timeout};
// === Setup: Two nodes with UDP transports on localhost ===
@@ -26,10 +28,8 @@ async fn test_two_node_handshake_udp() {
let (packet_tx_a, mut packet_rx_a) = packet_channel(64);
let (packet_tx_b, mut packet_rx_b) = packet_channel(64);
let mut transport_a =
UdpTransport::new(transport_id_a, None, udp_config.clone(), packet_tx_a);
let mut transport_b =
UdpTransport::new(transport_id_b, None, udp_config, packet_tx_b);
let mut transport_a = UdpTransport::new(transport_id_a, None, udp_config.clone(), packet_tx_a);
let mut transport_b = UdpTransport::new(transport_id_b, None, udp_config, packet_tx_b);
transport_a.start_async().await.unwrap();
transport_b.start_async().await.unwrap();
@@ -49,23 +49,20 @@ async fn test_two_node_handshake_udp() {
// === Phase 1: Node A initiates handshake to Node B ===
// Create peer identity for B (must use full key for ECDH parity)
let peer_b_identity =
PeerIdentity::from_pubkey_full(node_b.identity.pubkey_full());
let peer_b_identity = PeerIdentity::from_pubkey_full(node_b.identity.pubkey_full());
let peer_b_node_addr = *peer_b_identity.node_addr();
let link_id_a = node_a.allocate_link_id();
let mut conn_a = PeerConnection::outbound(
link_id_a,
peer_b_identity,
1000,
);
let mut conn_a = PeerConnection::outbound(link_id_a, peer_b_identity, 1000);
// Allocate session index for A's outbound
let our_index_a = node_a.index_allocator.allocate().unwrap();
// Start handshake (generates Noise XX msg1)
let our_keypair_a = node_a.identity.keypair();
let noise_msg1 = conn_a.start_handshake(our_keypair_a, node_a.startup_epoch, 1000).unwrap();
let noise_msg1 = conn_a
.start_handshake(our_keypair_a, node_a.startup_epoch, 1000)
.unwrap();
conn_a.set_our_index(our_index_a);
conn_a.set_transport_id(transport_id_a);
conn_a.set_source_addr(remote_addr_b.clone());
@@ -82,10 +79,9 @@ async fn test_two_node_handshake_udp() {
);
node_a.links.insert(link_id_a, link_a);
node_a.connections.insert(link_id_a, conn_a);
node_a.pending_outbound.insert(
(transport_id_a, our_index_a.as_u32()),
link_id_a,
);
node_a
.pending_outbound
.insert((transport_id_a, our_index_a.as_u32()), link_id_a);
// Send msg1 from A to B over UDP
let transport = node_a.transports.get(&transport_id_a).unwrap();
@@ -103,14 +99,20 @@ async fn test_two_node_handshake_udp() {
node_b.handle_msg1(packet_b).await;
let peer_a_node_addr = *PeerIdentity::from_pubkey_full(
node_a.identity.pubkey_full(),
)
.node_addr();
let peer_a_node_addr =
*PeerIdentity::from_pubkey_full(node_a.identity.pubkey_full()).node_addr();
// XX: B has NOT promoted yet (needs msg3)
assert_eq!(node_b.peer_count(), 0, "Node B should have 0 peers after msg1 (XX awaits msg3)");
assert_eq!(node_b.connections.len(), 1, "Node B should have 1 pending connection awaiting msg3");
assert_eq!(
node_b.peer_count(),
0,
"Node B should have 0 peers after msg1 (XX awaits msg3)"
);
assert_eq!(
node_b.connections.len(),
1,
"Node B should have 1 pending connection awaiting msg3"
);
// === Phase 3: Node A receives msg2, sends msg3, promotes ===
@@ -122,7 +124,11 @@ async fn test_two_node_handshake_udp() {
node_a.handle_msg2(packet_a).await;
// Verify A promoted the outbound connection
assert_eq!(node_a.peer_count(), 1, "Node A should have 1 peer after msg2");
assert_eq!(
node_a.peer_count(),
1,
"Node A should have 1 peer after msg2"
);
let peer_b_on_a = node_a
.get_peer(&peer_b_node_addr)
.expect("Node A should have peer B");
@@ -152,7 +158,11 @@ async fn test_two_node_handshake_udp() {
node_b.handle_msg3(packet_b_msg3).await;
// Verify B promoted after msg3
assert_eq!(node_b.peer_count(), 1, "Node B should have 1 peer after msg3");
assert_eq!(
node_b.peer_count(),
1,
"Node B should have 1 peer after msg3"
);
let peer_a_on_b = node_b
.get_peer(&peer_a_node_addr)
.expect("Node B should have peer A");
@@ -255,8 +265,8 @@ async fn test_two_node_handshake_udp() {
#[tokio::test]
async fn test_run_rx_loop_handshake() {
use crate::config::UdpConfig;
use crate::transport::udp::UdpTransport;
use crate::node::wire::build_msg1;
use crate::transport::udp::UdpTransport;
use tokio::time::Duration;
// === Setup: Two nodes with UDP transports on localhost ===
@@ -276,10 +286,8 @@ async fn test_run_rx_loop_handshake() {
let (packet_tx_a, packet_rx_a) = packet_channel(64);
let (packet_tx_b, packet_rx_b) = packet_channel(64);
let mut transport_a =
UdpTransport::new(transport_id_a, None, udp_config.clone(), packet_tx_a);
let mut transport_b =
UdpTransport::new(transport_id_b, None, udp_config, packet_tx_b);
let mut transport_a = UdpTransport::new(transport_id_a, None, udp_config.clone(), packet_tx_a);
let mut transport_b = UdpTransport::new(transport_id_b, None, udp_config, packet_tx_b);
transport_a.start_async().await.unwrap();
transport_b.start_async().await.unwrap();
@@ -304,20 +312,17 @@ async fn test_run_rx_loop_handshake() {
// === Phase 1: Node A initiates handshake to Node B ===
let peer_b_identity =
PeerIdentity::from_pubkey_full(node_b.identity.pubkey_full());
let peer_b_identity = PeerIdentity::from_pubkey_full(node_b.identity.pubkey_full());
let peer_b_node_addr = *peer_b_identity.node_addr();
let link_id_a = node_a.allocate_link_id();
let mut conn_a = PeerConnection::outbound(
link_id_a,
peer_b_identity,
1000,
);
let mut conn_a = PeerConnection::outbound(link_id_a, peer_b_identity, 1000);
let our_index_a = node_a.index_allocator.allocate().unwrap();
let our_keypair_a = node_a.identity.keypair();
let noise_msg1 = conn_a.start_handshake(our_keypair_a, node_a.startup_epoch, 1000).unwrap();
let noise_msg1 = conn_a
.start_handshake(our_keypair_a, node_a.startup_epoch, 1000)
.unwrap();
conn_a.set_our_index(our_index_a);
conn_a.set_transport_id(transport_id_a);
conn_a.set_source_addr(remote_addr_b.clone());
@@ -333,10 +338,9 @@ async fn test_run_rx_loop_handshake() {
);
node_a.links.insert(link_id_a, link_a);
node_a.connections.insert(link_id_a, conn_a);
node_a.pending_outbound.insert(
(transport_id_a, our_index_a.as_u32()),
link_id_a,
);
node_a
.pending_outbound
.insert((transport_id_a, our_index_a.as_u32()), link_id_a);
// Send msg1 from A to B over real UDP
let transport = node_a.transports.get(&transport_id_a).unwrap();
@@ -369,8 +373,16 @@ async fn test_run_rx_loop_handshake() {
}
// XX: Node B has NOT promoted yet (needs msg3)
assert_eq!(node_b.peer_count(), 0, "Node B should have 0 peers after rx loop processed msg1 (XX awaits msg3)");
assert_eq!(node_b.connections.len(), 1, "Node B should have 1 pending connection");
assert_eq!(
node_b.peer_count(),
0,
"Node B should have 0 peers after rx loop processed msg1 (XX awaits msg3)"
);
assert_eq!(
node_b.connections.len(),
1,
"Node B should have 1 pending connection"
);
// === Phase 3: Run Node A's rx loop (processes msg2, sends msg3) ===
@@ -384,7 +396,11 @@ async fn test_run_rx_loop_handshake() {
}
// Verify Node A promoted after processing msg2
assert_eq!(node_a.peer_count(), 1, "Node A should have 1 peer after rx loop processed msg2");
assert_eq!(
node_a.peer_count(),
1,
"Node A should have 1 peer after rx loop processed msg2"
);
let peer_b_on_a = node_a
.get_peer(&peer_b_node_addr)
.expect("Node A should have peer B");
@@ -413,7 +429,11 @@ async fn test_run_rx_loop_handshake() {
// verified by test_two_node_handshake_udp which uses direct handler calls.
// This test verifies rx_loop correctly dispatches PHASE_MSG1 (Phase 2)
// and PHASE_MSG2 (Phase 3). B still has a pending connection awaiting msg3.
assert_eq!(node_b.connections.len(), 1, "Node B should still have pending connection awaiting msg3");
assert_eq!(
node_b.connections.len(),
1,
"Node B should still have pending connection awaiting msg3"
);
// Clean up transports
for (_, t) in node_a.transports.iter_mut() {
@@ -433,9 +453,9 @@ async fn test_run_rx_loop_handshake() {
#[tokio::test]
async fn test_cross_connection_both_initiate() {
use crate::config::UdpConfig;
use crate::transport::udp::UdpTransport;
use crate::node::wire::build_msg1;
use tokio::time::{timeout, Duration};
use crate::transport::udp::UdpTransport;
use tokio::time::{Duration, timeout};
// === Setup: Two nodes with UDP transports on localhost ===
@@ -454,10 +474,8 @@ async fn test_cross_connection_both_initiate() {
let (packet_tx_a, mut packet_rx_a) = packet_channel(64);
let (packet_tx_b, mut packet_rx_b) = packet_channel(64);
let mut transport_a =
UdpTransport::new(transport_id_a, None, udp_config.clone(), packet_tx_a);
let mut transport_b =
UdpTransport::new(transport_id_b, None, udp_config, packet_tx_b);
let mut transport_a = UdpTransport::new(transport_id_a, None, udp_config.clone(), packet_tx_a);
let mut transport_b = UdpTransport::new(transport_id_b, None, udp_config, packet_tx_b);
transport_a.start_async().await.unwrap();
transport_b.start_async().await.unwrap();
@@ -475,11 +493,9 @@ async fn test_cross_connection_both_initiate() {
.insert(transport_id_b, TransportHandle::Udp(transport_b));
// Peer identities (must use full key for ECDH parity)
let peer_b_identity =
PeerIdentity::from_pubkey_full(node_b.identity.pubkey_full());
let peer_b_identity = PeerIdentity::from_pubkey_full(node_b.identity.pubkey_full());
let peer_b_node_addr = *peer_b_identity.node_addr();
let peer_a_identity =
PeerIdentity::from_pubkey_full(node_a.identity.pubkey_full());
let peer_a_identity = PeerIdentity::from_pubkey_full(node_a.identity.pubkey_full());
let peer_a_node_addr = *peer_a_identity.node_addr();
// === Phase 1: Both nodes initiate handshakes (simulate auto_connect) ===
@@ -489,7 +505,9 @@ async fn test_cross_connection_both_initiate() {
let mut conn_a = PeerConnection::outbound(link_id_a_out, peer_b_identity, 1000);
let our_index_a = node_a.index_allocator.allocate().unwrap();
let our_keypair_a = node_a.identity.keypair();
let noise_msg1_a = conn_a.start_handshake(our_keypair_a, node_a.startup_epoch, 1000).unwrap();
let noise_msg1_a = conn_a
.start_handshake(our_keypair_a, node_a.startup_epoch, 1000)
.unwrap();
conn_a.set_our_index(our_index_a);
conn_a.set_transport_id(transport_id_a);
conn_a.set_source_addr(remote_addr_b.clone());
@@ -497,20 +515,29 @@ async fn test_cross_connection_both_initiate() {
let wire_msg1_a = build_msg1(our_index_a, &noise_msg1_a);
let link_a_out = Link::connectionless(
link_id_a_out, transport_id_a, remote_addr_b.clone(),
LinkDirection::Outbound, Duration::from_millis(100),
link_id_a_out,
transport_id_a,
remote_addr_b.clone(),
LinkDirection::Outbound,
Duration::from_millis(100),
);
node_a.links.insert(link_id_a_out, link_a_out);
node_a.addr_to_link.insert((transport_id_a, remote_addr_b.clone()), link_id_a_out);
node_a
.addr_to_link
.insert((transport_id_a, remote_addr_b.clone()), link_id_a_out);
node_a.connections.insert(link_id_a_out, conn_a);
node_a.pending_outbound.insert((transport_id_a, our_index_a.as_u32()), link_id_a_out);
node_a
.pending_outbound
.insert((transport_id_a, our_index_a.as_u32()), link_id_a_out);
// Node B initiates to Node A
let link_id_b_out = node_b.allocate_link_id();
let mut conn_b = PeerConnection::outbound(link_id_b_out, peer_a_identity, 1000);
let our_index_b = node_b.index_allocator.allocate().unwrap();
let our_keypair_b = node_b.identity.keypair();
let noise_msg1_b = conn_b.start_handshake(our_keypair_b, node_b.startup_epoch, 1000).unwrap();
let noise_msg1_b = conn_b
.start_handshake(our_keypair_b, node_b.startup_epoch, 1000)
.unwrap();
conn_b.set_our_index(our_index_b);
conn_b.set_transport_id(transport_id_b);
conn_b.set_source_addr(remote_addr_a.clone());
@@ -518,77 +545,130 @@ async fn test_cross_connection_both_initiate() {
let wire_msg1_b = build_msg1(our_index_b, &noise_msg1_b);
let link_b_out = Link::connectionless(
link_id_b_out, transport_id_b, remote_addr_a.clone(),
LinkDirection::Outbound, Duration::from_millis(100),
link_id_b_out,
transport_id_b,
remote_addr_a.clone(),
LinkDirection::Outbound,
Duration::from_millis(100),
);
node_b.links.insert(link_id_b_out, link_b_out);
node_b.addr_to_link.insert((transport_id_b, remote_addr_a.clone()), link_id_b_out);
node_b
.addr_to_link
.insert((transport_id_b, remote_addr_a.clone()), link_id_b_out);
node_b.connections.insert(link_id_b_out, conn_b);
node_b.pending_outbound.insert((transport_id_b, our_index_b.as_u32()), link_id_b_out);
node_b
.pending_outbound
.insert((transport_id_b, our_index_b.as_u32()), link_id_b_out);
// Both send msg1 over UDP
let transport = node_a.transports.get(&transport_id_a).unwrap();
transport.send(&remote_addr_b, &wire_msg1_a).await.expect("A send msg1");
transport
.send(&remote_addr_b, &wire_msg1_a)
.await
.expect("A send msg1");
let transport = node_b.transports.get(&transport_id_b).unwrap();
transport.send(&remote_addr_a, &wire_msg1_b).await.expect("B send msg1");
transport
.send(&remote_addr_a, &wire_msg1_b)
.await
.expect("B send msg1");
// === Phase 2: Both nodes receive the other's msg1 (XX: no promotion yet) ===
// B receives A's msg1
let packet_at_b = timeout(Duration::from_secs(1), packet_rx_b.recv())
.await.expect("Timeout").expect("Channel closed");
.await
.expect("Timeout")
.expect("Channel closed");
node_b.handle_msg1(packet_at_b).await;
// XX: B has NOT promoted yet (needs msg3 from A)
assert_eq!(node_b.peer_count(), 0, "Node B should have 0 peers after processing A's msg1 (XX)");
assert_eq!(
node_b.peer_count(),
0,
"Node B should have 0 peers after processing A's msg1 (XX)"
);
// A receives B's msg1
let packet_at_a = timeout(Duration::from_secs(1), packet_rx_a.recv())
.await.expect("Timeout").expect("Channel closed");
.await
.expect("Timeout")
.expect("Channel closed");
node_a.handle_msg1(packet_at_a).await;
// XX: A has NOT promoted yet (needs msg3 from B)
assert_eq!(node_a.peer_count(), 0, "Node A should have 0 peers after processing B's msg1 (XX)");
assert_eq!(
node_a.peer_count(),
0,
"Node A should have 0 peers after processing B's msg1 (XX)"
);
// === Phase 3: Both nodes receive msg2 + send msg3, initiator side promotes ===
// A receives B's msg2 (response to A's original msg1) → A sends msg3, A promotes
let msg2_at_a = timeout(Duration::from_secs(1), packet_rx_a.recv())
.await.expect("Timeout waiting for msg2 at A").expect("Channel closed");
.await
.expect("Timeout waiting for msg2 at A")
.expect("Channel closed");
node_a.handle_msg2(msg2_at_a).await;
// A promoted as initiator
assert_eq!(node_a.peer_count(), 1, "Node A should have 1 peer after processing msg2");
assert_eq!(
node_a.peer_count(),
1,
"Node A should have 1 peer after processing msg2"
);
// B receives A's msg2 (response to B's original msg1) → B sends msg3, B promotes
let msg2_at_b = timeout(Duration::from_secs(1), packet_rx_b.recv())
.await.expect("Timeout waiting for msg2 at B").expect("Channel closed");
.await
.expect("Timeout waiting for msg2 at B")
.expect("Channel closed");
node_b.handle_msg2(msg2_at_b).await;
// B promoted as initiator
assert_eq!(node_b.peer_count(), 1, "Node B should have 1 peer after processing msg2");
assert_eq!(
node_b.peer_count(),
1,
"Node B should have 1 peer after processing msg2"
);
// === Phase 4: Both nodes receive msg3, responder side completes ===
// Cross-connection resolution happens here (or in Phase 3 promotion).
// A receives B's msg3 (B completing A's inbound handshake)
let msg3_at_a = timeout(Duration::from_secs(1), packet_rx_a.recv())
.await.expect("Timeout waiting for msg3 at A").expect("Channel closed");
.await
.expect("Timeout waiting for msg3 at A")
.expect("Channel closed");
node_a.handle_msg3(msg3_at_a).await;
// B receives A's msg3 (A completing B's inbound handshake)
let msg3_at_b = timeout(Duration::from_secs(1), packet_rx_b.recv())
.await.expect("Timeout waiting for msg3 at B").expect("Channel closed");
.await
.expect("Timeout waiting for msg3 at B")
.expect("Channel closed");
node_b.handle_msg3(msg3_at_b).await;
// === Verification ===
// Both nodes should have exactly 1 peer each after cross-connection resolution
assert_eq!(node_a.peer_count(), 1, "Node A should have exactly 1 peer after cross-connection");
assert_eq!(node_b.peer_count(), 1, "Node B should have exactly 1 peer after cross-connection");
assert_eq!(
node_a.peer_count(),
1,
"Node A should have exactly 1 peer after cross-connection"
);
assert_eq!(
node_b.peer_count(),
1,
"Node B should have exactly 1 peer after cross-connection"
);
let peer_b_on_a = node_a.get_peer(&peer_b_node_addr).expect("A should have peer B");
let peer_a_on_b = node_b.get_peer(&peer_a_node_addr).expect("B should have peer A");
let peer_b_on_a = node_a
.get_peer(&peer_b_node_addr)
.expect("A should have peer B");
let peer_a_on_b = node_b
.get_peer(&peer_a_node_addr)
.expect("B should have peer A");
assert!(peer_b_on_a.has_session(), "Peer B on A should have session");
assert!(peer_a_on_b.has_session(), "Peer A on B should have session");
@@ -625,25 +705,35 @@ async fn test_stale_connection_cleanup() {
// Allocate session index and set transport info
let our_index = node.index_allocator.allocate().unwrap();
let our_keypair = node.identity.keypair();
let _noise_msg1 = conn.start_handshake(our_keypair, node.startup_epoch, past_time_ms).unwrap();
let _noise_msg1 = conn
.start_handshake(our_keypair, node.startup_epoch, past_time_ms)
.unwrap();
conn.set_our_index(our_index);
conn.set_transport_id(transport_id);
conn.set_source_addr(remote_addr.clone());
// Set up all the state that initiate_peer_connection would create
let link = Link::connectionless(
link_id, transport_id, remote_addr.clone(),
LinkDirection::Outbound, Duration::from_millis(100),
link_id,
transport_id,
remote_addr.clone(),
LinkDirection::Outbound,
Duration::from_millis(100),
);
node.links.insert(link_id, link);
node.addr_to_link.insert((transport_id, remote_addr.clone()), link_id);
node.addr_to_link
.insert((transport_id, remote_addr.clone()), link_id);
node.connections.insert(link_id, conn);
node.pending_outbound.insert((transport_id, our_index.as_u32()), link_id);
node.pending_outbound
.insert((transport_id, our_index.as_u32()), link_id);
// Verify state before timeout check
assert_eq!(node.connection_count(), 1);
assert_eq!(node.link_count(), 1);
assert!(node.pending_outbound.contains_key(&(transport_id, our_index.as_u32())));
assert!(
node.pending_outbound
.contains_key(&(transport_id, our_index.as_u32()))
);
assert_eq!(node.index_allocator.count(), 1);
// Connection was created at time 1000ms. check_timeouts uses SystemTime::now(),
@@ -651,13 +741,27 @@ async fn test_stale_connection_cleanup() {
node.check_timeouts();
// Verify everything was cleaned up
assert_eq!(node.connection_count(), 0, "Stale connection should be removed");
assert_eq!(
node.connection_count(),
0,
"Stale connection should be removed"
);
assert_eq!(node.link_count(), 0, "Stale link should be removed");
assert!(!node.pending_outbound.contains_key(&(transport_id, our_index.as_u32())),
"pending_outbound should be cleaned up");
assert_eq!(node.index_allocator.count(), 0, "Session index should be freed");
assert!(!node.addr_to_link.contains_key(&(transport_id, remote_addr)),
"addr_to_link should be cleaned up");
assert!(
!node
.pending_outbound
.contains_key(&(transport_id, our_index.as_u32())),
"pending_outbound should be cleaned up"
);
assert_eq!(
node.index_allocator.count(),
0,
"Session index should be freed"
);
assert!(
!node.addr_to_link.contains_key(&(transport_id, remote_addr)),
"addr_to_link should be cleaned up"
);
}
/// Test that failed connections are cleaned up by check_timeouts().
@@ -679,29 +783,44 @@ async fn test_failed_connection_cleanup() {
let our_index = node.index_allocator.allocate().unwrap();
let our_keypair = node.identity.keypair();
let _noise_msg1 = conn.start_handshake(our_keypair, node.startup_epoch, now_ms).unwrap();
let _noise_msg1 = conn
.start_handshake(our_keypair, node.startup_epoch, now_ms)
.unwrap();
conn.set_our_index(our_index);
conn.set_transport_id(transport_id);
conn.set_source_addr(remote_addr.clone());
conn.mark_failed(); // Simulate send failure
let link = Link::connectionless(
link_id, transport_id, remote_addr.clone(),
LinkDirection::Outbound, Duration::from_millis(100),
link_id,
transport_id,
remote_addr.clone(),
LinkDirection::Outbound,
Duration::from_millis(100),
);
node.links.insert(link_id, link);
node.addr_to_link.insert((transport_id, remote_addr.clone()), link_id);
node.addr_to_link
.insert((transport_id, remote_addr.clone()), link_id);
node.connections.insert(link_id, conn);
node.pending_outbound.insert((transport_id, our_index.as_u32()), link_id);
node.pending_outbound
.insert((transport_id, our_index.as_u32()), link_id);
assert_eq!(node.connection_count(), 1);
// Failed connections should be cleaned up immediately regardless of age
node.check_timeouts();
assert_eq!(node.connection_count(), 0, "Failed connection should be removed");
assert_eq!(
node.connection_count(),
0,
"Failed connection should be removed"
);
assert_eq!(node.link_count(), 0, "Failed link should be removed");
assert_eq!(node.index_allocator.count(), 0, "Session index should be freed");
assert_eq!(
node.index_allocator.count(),
0,
"Session index should be freed"
);
}
/// Test that msg1 bytes are stored on connection for resend.
@@ -724,7 +843,9 @@ async fn test_msg1_stored_for_resend() {
let our_index = node.index_allocator.allocate().unwrap();
let our_keypair = node.identity.keypair();
let noise_msg1 = conn.start_handshake(our_keypair, node.startup_epoch, now_ms).unwrap();
let noise_msg1 = conn
.start_handshake(our_keypair, node.startup_epoch, now_ms)
.unwrap();
conn.set_our_index(our_index);
conn.set_transport_id(transport_id);
conn.set_source_addr(remote_addr.clone());
@@ -755,7 +876,9 @@ async fn test_resend_scheduling() {
let our_index = node.index_allocator.allocate().unwrap();
let our_keypair = node.identity.keypair();
let noise_msg1 = conn.start_handshake(our_keypair, node.startup_epoch, now_ms).unwrap();
let noise_msg1 = conn
.start_handshake(our_keypair, node.startup_epoch, now_ms)
.unwrap();
conn.set_our_index(our_index);
conn.set_transport_id(transport_id);
conn.set_source_addr(remote_addr.clone());
@@ -765,12 +888,17 @@ async fn test_resend_scheduling() {
conn.set_handshake_msg1(wire_msg1, now_ms + 1000);
let link = Link::connectionless(
link_id, transport_id, remote_addr.clone(),
LinkDirection::Outbound, Duration::from_millis(100),
link_id,
transport_id,
remote_addr.clone(),
LinkDirection::Outbound,
Duration::from_millis(100),
);
node.links.insert(link_id, link);
node.addr_to_link.insert((transport_id, remote_addr), link_id);
node.pending_outbound.insert((transport_id, our_index.as_u32()), link_id);
node.addr_to_link
.insert((transport_id, remote_addr), link_id);
node.pending_outbound
.insert((transport_id, our_index.as_u32()), link_id);
node.connections.insert(link_id, conn);
// Before resend time: nothing should happen (no transport = can't send,
@@ -786,7 +914,11 @@ async fn test_resend_scheduling() {
// No transport registered, so send fails — count stays 0.
// That's the expected behavior (transport absence is a transient condition).
let conn = node.connections.get(&link_id).unwrap();
assert_eq!(conn.resend_count(), 0, "No transport means no resend recorded");
assert_eq!(
conn.resend_count(),
0,
"No transport means no resend recorded"
);
}
/// Test that msg2 is stored on PeerConnection for responder resend.
+6 -2
View File
@@ -68,10 +68,14 @@ pub(super) fn make_completed_connection(
.unwrap();
// Complete initiator handshake (XX: generates msg3)
let (msg3, _neg) = conn.complete_handshake(&msg2, None, current_time_ms).unwrap();
let (msg3, _neg) = conn
.complete_handshake(&msg2, None, current_time_ms)
.unwrap();
// Complete responder handshake (XX: processes msg3)
resp_conn.complete_handshake_msg3(&msg3, current_time_ms).unwrap();
resp_conn
.complete_handshake_msg3(&msg3, current_time_ms)
.unwrap();
// Set indices and transport info
let our_index = node.index_allocator.allocate().unwrap();
+324 -150
View File
@@ -3,8 +3,8 @@
use super::*;
use crate::node::session::EndToEndState;
use crate::node::tests::spanning_tree::{
cleanup_nodes, generate_random_edges, process_available_packets, run_tree_test,
run_tree_test_with_mtus, verify_tree_convergence, TestNode,
TestNode, cleanup_nodes, generate_random_edges, process_available_packets, run_tree_test,
run_tree_test_with_mtus, verify_tree_convergence,
};
use crate::protocol::{SessionAck, SessionDatagram};
@@ -142,12 +142,14 @@ async fn test_session_direct_peer_handshake() {
// Node 0 should have a session in Initiating state
assert_eq!(nodes[0].node.session_count(), 1);
assert!(nodes[0]
.node
.get_session(&node1_addr)
.unwrap()
.state()
.is_initiating());
assert!(
nodes[0]
.node
.get_session(&node1_addr)
.unwrap()
.state()
.is_initiating()
);
// Process packets: SessionSetup arrives at Node 1
tokio::time::sleep(Duration::from_millis(20)).await;
@@ -156,12 +158,14 @@ async fn test_session_direct_peer_handshake() {
// Node 1 should now have a session in AwaitingMsg3 state (XX: identity not yet known)
assert_eq!(nodes[1].node.session_count(), 1);
assert!(nodes[1]
.node
.get_session(&node0_addr)
.unwrap()
.state()
.is_awaiting_msg3());
assert!(
nodes[1]
.node
.get_session(&node0_addr)
.unwrap()
.state()
.is_awaiting_msg3()
);
// Process packets: SessionAck arrives at Node 0, Node 0 sends SessionMsg3
tokio::time::sleep(Duration::from_millis(20)).await;
@@ -169,12 +173,14 @@ async fn test_session_direct_peer_handshake() {
assert!(count > 0, "Expected SessionAck packet to arrive");
// Node 0 should now be Established (transitions after sending msg3)
assert!(nodes[0]
.node
.get_session(&node1_addr)
.unwrap()
.state()
.is_established());
assert!(
nodes[0]
.node
.get_session(&node1_addr)
.unwrap()
.state()
.is_established()
);
// Process packets: SessionMsg3 arrives at Node 1
tokio::time::sleep(Duration::from_millis(20)).await;
@@ -182,12 +188,14 @@ async fn test_session_direct_peer_handshake() {
assert!(count > 0, "Expected SessionMsg3 packet to arrive");
// Node 1 should now be Established (transitions after processing msg3)
assert!(nodes[1]
.node
.get_session(&node0_addr)
.unwrap()
.state()
.is_established());
assert!(
nodes[1]
.node
.get_session(&node0_addr)
.unwrap()
.state()
.is_established()
);
cleanup_nodes(&mut nodes).await;
}
@@ -217,18 +225,22 @@ async fn test_session_direct_peer_data_transfer() {
tokio::time::sleep(Duration::from_millis(20)).await;
process_available_packets(&mut nodes).await; // Msg3 → Node 1
assert!(nodes[0]
.node
.get_session(&node1_addr)
.unwrap()
.state()
.is_established());
assert!(nodes[1]
.node
.get_session(&node0_addr)
.unwrap()
.state()
.is_established());
assert!(
nodes[0]
.node
.get_session(&node1_addr)
.unwrap()
.state()
.is_established()
);
assert!(
nodes[1]
.node
.get_session(&node0_addr)
.unwrap()
.state()
.is_established()
);
// Send data from Node 0 to Node 1
let test_data = b"Hello, FIPS session!";
@@ -282,12 +294,14 @@ async fn test_session_3node_forwarded_handshake() {
nodes[2].node.get_session(&node0_addr).is_some(),
"Node 2 should have a session entry for Node 0"
);
assert!(nodes[2]
.node
.get_session(&node0_addr)
.unwrap()
.state()
.is_awaiting_msg3());
assert!(
nodes[2]
.node
.get_session(&node0_addr)
.unwrap()
.state()
.is_awaiting_msg3()
);
// Process: SessionAck: 2→1 (forwarded by transit B)
tokio::time::sleep(Duration::from_millis(20)).await;
@@ -298,12 +312,14 @@ async fn test_session_3node_forwarded_handshake() {
process_available_packets(&mut nodes).await;
// Node 0 should now be Established (transitions after sending msg3)
assert!(nodes[0]
.node
.get_session(&node2_addr)
.unwrap()
.state()
.is_established());
assert!(
nodes[0]
.node
.get_session(&node2_addr)
.unwrap()
.state()
.is_established()
);
// Process: SessionMsg3: 0→1 (forwarded by transit B)
tokio::time::sleep(Duration::from_millis(20)).await;
@@ -314,12 +330,14 @@ async fn test_session_3node_forwarded_handshake() {
process_available_packets(&mut nodes).await;
// Node 2 should now be Established (transitions after processing msg3)
assert!(nodes[2]
.node
.get_session(&node0_addr)
.unwrap()
.state()
.is_established());
assert!(
nodes[2]
.node
.get_session(&node0_addr)
.unwrap()
.state()
.is_established()
);
// Transit node B should NOT have a session
assert_eq!(
@@ -380,12 +398,14 @@ async fn test_session_3node_forwarded_data() {
}
// Node 2 should be Established (transitioned during XX handshake msg3)
assert!(nodes[2]
.node
.get_session(&node0_addr)
.unwrap()
.state()
.is_established());
assert!(
nodes[2]
.node
.get_session(&node0_addr)
.unwrap()
.state()
.is_established()
);
cleanup_nodes(&mut nodes).await;
}
@@ -511,12 +531,7 @@ async fn test_session_100_nodes() {
// Collect identities: (node_addr, pubkey) for all nodes
let all_info: Vec<(NodeAddr, secp256k1::PublicKey)> = nodes
.iter()
.map(|tn| {
(
*tn.node.node_addr(),
tn.node.identity().pubkey_full(),
)
})
.map(|tn| (*tn.node.node_addr(), tn.node.identity().pubkey_full()))
.collect();
// Each node picks one random target for its outbound session.
@@ -631,11 +646,7 @@ async fn test_session_100_nodes() {
// (Responder should already be Established after XX msg3)
let rev_payload = format!("rev-{}", pair_idx).into_bytes();
let rev_ipv6 = build_ipv6_packet(&dst_fips, &src_fips, &rev_payload);
match nodes[dst]
.node
.send_ipv6_packet(&src_addr, &rev_ipv6)
.await
{
match nodes[dst].node.send_ipv6_packet(&src_addr, &rev_ipv6).await {
Ok(()) => send_reverse_ok += 1,
Err(_) => send_reverse_err += 1,
}
@@ -714,10 +725,7 @@ async fn test_session_100_nodes() {
}
}
let session_counts: Vec<usize> = nodes
.iter()
.map(|tn| tn.node.session_count())
.collect();
let session_counts: Vec<usize> = nodes.iter().map(|tn| tn.node.session_count()).collect();
let total_sessions: usize = session_counts.iter().sum();
let min_sessions = *session_counts.iter().min().unwrap();
let max_sessions = *session_counts.iter().max().unwrap();
@@ -761,10 +769,8 @@ async fn test_session_100_nodes() {
};
// Coord cache stats
let coord_cache_sizes: Vec<usize> = nodes
.iter()
.map(|tn| tn.node.coord_cache().len())
.collect();
let coord_cache_sizes: Vec<usize> =
nodes.iter().map(|tn| tn.node.coord_cache().len()).collect();
let total_coord_entries: usize = coord_cache_sizes.iter().sum();
let min_coord = *coord_cache_sizes.iter().min().unwrap();
let max_coord = *coord_cache_sizes.iter().max().unwrap();
@@ -875,10 +881,7 @@ async fn test_session_100_nodes() {
// === Assertions ===
assert_eq!(
send_forward_err, 0,
"All forward sends should succeed"
);
assert_eq!(send_forward_err, 0, "All forward sends should succeed");
assert_eq!(
send_reverse_err, 0,
"All reverse sends should succeed (responder Established after XX msg3)"
@@ -906,7 +909,11 @@ async fn test_session_100_nodes() {
// ============================================================================
/// Build a minimal valid IPv6 packet with given source and destination addresses.
fn build_ipv6_packet(src: &crate::FipsAddress, dst: &crate::FipsAddress, payload: &[u8]) -> Vec<u8> {
fn build_ipv6_packet(
src: &crate::FipsAddress,
dst: &crate::FipsAddress,
payload: &[u8],
) -> Vec<u8> {
let payload_len = payload.len() as u16;
let mut packet = vec![0u8; 40 + payload.len()];
// Version (6) + traffic class high nibble
@@ -935,17 +942,14 @@ fn test_identity_cache_populated_on_promote() {
let transport_id = TransportId::new(1);
let link_id = LinkId::new(1);
let (conn, peer_identity) = make_completed_connection(
&mut node,
link_id,
transport_id,
1000,
);
let (conn, peer_identity) = make_completed_connection(&mut node, link_id, transport_id, 1000);
node.add_connection(conn).unwrap();
// Promote
let result = node.promote_connection(link_id, peer_identity, 2000).unwrap();
let result = node
.promote_connection(link_id, peer_identity, 2000)
.unwrap();
assert!(matches!(result, PromotionResult::Promoted(_)));
// Identity cache should contain the peer
@@ -953,7 +957,10 @@ fn test_identity_cache_populated_on_promote() {
let mut prefix = [0u8; 15];
prefix.copy_from_slice(&peer_addr.as_bytes()[0..15]);
let cached = node.lookup_by_fips_prefix(&prefix);
assert!(cached.is_some(), "Identity cache should contain promoted peer");
assert!(
cached.is_some(),
"Identity cache should contain promoted peer"
);
let (cached_addr, cached_pk) = cached.unwrap();
assert_eq!(cached_addr, peer_addr);
assert_eq!(cached_pk, peer_identity.pubkey_full());
@@ -977,7 +984,11 @@ async fn test_tun_outbound_established_session() {
let dst_fips = crate::FipsAddress::from_node_addr(&node1_addr);
// Establish session (XX: 3 messages — Setup, Ack, Msg3)
nodes[0].node.initiate_session(node1_addr, node1_pubkey).await.unwrap();
nodes[0]
.node
.initiate_session(node1_addr, node1_pubkey)
.await
.unwrap();
tokio::time::sleep(Duration::from_millis(20)).await;
process_available_packets(&mut nodes).await; // Setup → Node 1
tokio::time::sleep(Duration::from_millis(20)).await;
@@ -985,7 +996,14 @@ async fn test_tun_outbound_established_session() {
tokio::time::sleep(Duration::from_millis(20)).await;
process_available_packets(&mut nodes).await; // Msg3 → Node 1
assert!(nodes[0].node.get_session(&node1_addr).unwrap().state().is_established());
assert!(
nodes[0]
.node
.get_session(&node1_addr)
.unwrap()
.state()
.is_established()
);
// Install TUN receiver on Node 1
let (tun_tx, tun_rx) = std::sync::mpsc::channel();
@@ -1004,7 +1022,10 @@ async fn test_tun_outbound_established_session() {
// Verify plaintext arrived at Node 1's TUN
let delivered: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx.try_recv().ok()).collect();
assert_eq!(delivered.len(), 1, "Exactly one packet should be delivered");
assert_eq!(delivered[0], ipv6_packet, "Delivered packet should match original");
assert_eq!(
delivered[0], ipv6_packet,
"Delivered packet should match original"
);
cleanup_nodes(&mut nodes).await;
}
@@ -1040,17 +1061,35 @@ async fn test_tun_outbound_triggers_session_initiation() {
// Session should now be initiating
assert_eq!(nodes[0].node.session_count(), 1);
assert!(nodes[0].node.get_session(&node1_addr).unwrap().state().is_initiating());
assert!(
nodes[0]
.node
.get_session(&node1_addr)
.unwrap()
.state()
.is_initiating()
);
// Drain packets until session established and queued packet delivered
drain_to_quiescence(&mut nodes).await;
// Session should be established on Node 0
assert!(nodes[0].node.get_session(&node1_addr).unwrap().state().is_established());
assert!(
nodes[0]
.node
.get_session(&node1_addr)
.unwrap()
.state()
.is_established()
);
// Verify the queued packet was delivered to Node 1
let delivered: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx.try_recv().ok()).collect();
assert_eq!(delivered.len(), 1, "Queued packet should be delivered after handshake");
assert_eq!(
delivered.len(),
1,
"Queued packet should be delivered after handshake"
);
assert_eq!(delivered[0], ipv6_packet);
cleanup_nodes(&mut nodes).await;
@@ -1078,12 +1117,19 @@ async fn test_tun_outbound_unknown_destination() {
// Should receive ICMPv6 Destination Unreachable back on TUN
let delivered: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx.try_recv().ok()).collect();
assert_eq!(delivered.len(), 1, "Should receive ICMPv6 Destination Unreachable");
assert_eq!(
delivered.len(),
1,
"Should receive ICMPv6 Destination Unreachable"
);
// Verify it's an ICMPv6 Destination Unreachable (type 1, code 0)
// ICMPv6 header starts at byte 40, type at byte 40, code at byte 41
assert!(delivered[0].len() >= 48, "ICMPv6 response too short");
assert_eq!(delivered[0][6], 58, "Next header should be ICMPv6 (58)");
assert_eq!(delivered[0][40], 1, "ICMPv6 type should be Destination Unreachable (1)");
assert_eq!(
delivered[0][40], 1,
"ICMPv6 type should be Destination Unreachable (1)"
);
assert_eq!(delivered[0][41], 0, "ICMPv6 code should be No Route (0)");
cleanup_nodes(&mut nodes).await;
@@ -1122,7 +1168,14 @@ async fn test_tun_outbound_3node_forwarded() {
drain_to_quiescence(&mut nodes).await;
// Session should be established
assert!(nodes[0].node.get_session(&node2_addr).unwrap().state().is_established());
assert!(
nodes[0]
.node
.get_session(&node2_addr)
.unwrap()
.state()
.is_established()
);
// Verify packet delivered to Node 2
let delivered: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx.try_recv().ok()).collect();
@@ -1161,16 +1214,34 @@ async fn test_tun_outbound_pending_queue_flush() {
// First packet triggers session initiation, rest are queued
assert_eq!(nodes[0].node.session_count(), 1);
assert!(nodes[0].node.get_session(&node1_addr).unwrap().state().is_initiating());
assert!(
nodes[0]
.node
.get_session(&node1_addr)
.unwrap()
.state()
.is_initiating()
);
// Drain until session established and queued packets flushed
drain_to_quiescence(&mut nodes).await;
assert!(nodes[0].node.get_session(&node1_addr).unwrap().state().is_established());
assert!(
nodes[0]
.node
.get_session(&node1_addr)
.unwrap()
.state()
.is_established()
);
// All 5 packets should have been delivered
let delivered: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx.try_recv().ok()).collect();
assert_eq!(delivered.len(), 5, "All 5 queued packets should be delivered");
assert_eq!(
delivered.len(),
5,
"All 5 queued packets should be delivered"
);
for (i, pkt) in delivered.iter().enumerate() {
assert_eq!(*pkt, packets[i], "Packet {} should match", i);
}
@@ -1260,7 +1331,11 @@ fn test_purge_idle_sessions_keeps_active() {
let now_ms = 92_000;
node.purge_idle_sessions(now_ms);
assert_eq!(node.session_count(), 1, "Active session should survive purge");
assert_eq!(
node.session_count(),
1,
"Active session should survive purge"
);
}
#[test]
@@ -1286,7 +1361,11 @@ fn test_purge_idle_sessions_ignores_initiating() {
let now_ms = 1000 + 200_000;
node.purge_idle_sessions(now_ms);
assert_eq!(node.session_count(), 1, "Initiating session should not be purged by idle timeout");
assert_eq!(
node.session_count(),
1,
"Initiating session should not be purged by idle timeout"
);
}
#[test]
@@ -1317,8 +1396,10 @@ fn test_purge_idle_sessions_cleans_pending_packets() {
node.purge_idle_sessions(now_ms);
assert_eq!(node.session_count(), 0);
assert!(!node.pending_tun_packets.contains_key(&remote_addr),
"Pending packets should be cleaned up with idle session");
assert!(
!node.pending_tun_packets.contains_key(&remote_addr),
"Pending packets should be cleaned up with idle session"
);
}
#[test]
@@ -1344,7 +1425,11 @@ fn test_purge_idle_sessions_disabled_when_zero() {
let now_ms = 1000 + 1_000_000;
node.purge_idle_sessions(now_ms);
assert_eq!(node.session_count(), 1, "Sessions should not be purged when idle timeout is disabled");
assert_eq!(
node.session_count(),
1,
"Sessions should not be purged when idle timeout is disabled"
);
}
#[test]
@@ -1373,8 +1458,11 @@ fn test_purge_idle_sessions_mmp_activity_does_not_prevent_purge() {
let now_ms = 92_000;
node.purge_idle_sessions(now_ms);
assert_eq!(node.session_count(), 0,
"Session with MMP-only activity should be purged");
assert_eq!(
node.session_count(),
0,
"Session with MMP-only activity should be purged"
);
}
// ============================================================================
@@ -1398,8 +1486,11 @@ fn test_coords_warmup_counter_default_zero_on_new() {
true,
);
assert_eq!(entry.coords_warmup_remaining(), 0,
"Counter should be 0 for non-Established sessions");
assert_eq!(
entry.coords_warmup_remaining(),
0,
"Counter should be 0 for non-Established sessions"
);
}
#[test]
@@ -1450,15 +1541,20 @@ fn test_coords_warmup_counter_decrement() {
assert_eq!(entry.coords_warmup_remaining(), expected);
}
assert_eq!(entry.coords_warmup_remaining(), 0,
"Counter should reach 0 after N decrements");
assert_eq!(
entry.coords_warmup_remaining(),
0,
"Counter should reach 0 after N decrements"
);
}
#[test]
fn test_coords_warmup_config_default() {
let config = crate::config::Config::new();
assert_eq!(config.node.session.coords_warmup_packets, 5,
"Default coords_warmup_packets should be 5");
assert_eq!(
config.node.session.coords_warmup_packets, 5,
"Default coords_warmup_packets should be 5"
);
}
// ============================================================================
@@ -1477,11 +1573,13 @@ fn test_identity_cache_lru_eviction() {
// Insert first two with explicit timestamps to ensure deterministic ordering
let mut prefix1 = [0u8; 15];
prefix1.copy_from_slice(&id1.node_addr().as_bytes()[0..15]);
node.identity_cache.insert(prefix1, (*id1.node_addr(), id1.pubkey_full(), 1000));
node.identity_cache
.insert(prefix1, (*id1.node_addr(), id1.pubkey_full(), 1000));
let mut prefix2 = [0u8; 15];
prefix2.copy_from_slice(&id2.node_addr().as_bytes()[0..15]);
node.identity_cache.insert(prefix2, (*id2.node_addr(), id2.pubkey_full(), 2000));
node.identity_cache
.insert(prefix2, (*id2.node_addr(), id2.pubkey_full(), 2000));
assert_eq!(node.identity_cache_len(), 2);
@@ -1489,13 +1587,17 @@ fn test_identity_cache_lru_eviction() {
node.register_identity(*id3.node_addr(), id3.pubkey_full());
assert_eq!(node.identity_cache_len(), 2);
assert!(node.lookup_by_fips_prefix(&prefix1).is_none(),
"Oldest entry should have been evicted");
assert!(
node.lookup_by_fips_prefix(&prefix1).is_none(),
"Oldest entry should have been evicted"
);
let mut prefix3 = [0u8; 15];
prefix3.copy_from_slice(&id3.node_addr().as_bytes()[0..15]);
assert!(node.lookup_by_fips_prefix(&prefix3).is_some(),
"Newest entry should be present");
assert!(
node.lookup_by_fips_prefix(&prefix3).is_some(),
"Newest entry should be present"
);
}
#[test]
@@ -1644,12 +1746,18 @@ async fn test_session_handshake_timeout() {
let timeout_secs = node.config.node.rate_limit.handshake_timeout_secs;
let before_timeout = 1000 + timeout_secs * 1000 - 1;
node.resend_pending_session_handshakes(before_timeout).await;
assert!(node.sessions.contains_key(&dest_addr), "Session should survive before timeout");
assert!(
node.sessions.contains_key(&dest_addr),
"Session should survive before timeout"
);
// After timeout: session should be removed
let after_timeout = 1000 + timeout_secs * 1000 + 1;
node.resend_pending_session_handshakes(after_timeout).await;
assert!(!node.sessions.contains_key(&dest_addr), "Timed-out session should be removed");
assert!(
!node.sessions.contains_key(&dest_addr),
"Timed-out session should be removed"
);
}
/// Test that session handshake timeout removes stale AwaitingMsg3 sessions.
@@ -1662,9 +1770,7 @@ async fn test_session_awaiting_msg3_timeout() {
let identity_a = Identity::generate();
let identity_b = Identity::generate();
let handshake = HandshakeState::new_responder(
identity_b.keypair(),
);
let handshake = HandshakeState::new_responder(identity_b.keypair());
let src_addr = *identity_a.node_addr();
@@ -1684,7 +1790,10 @@ async fn test_session_awaiting_msg3_timeout() {
let timeout_secs = node.config.node.rate_limit.handshake_timeout_secs;
let after_timeout = 1000 + timeout_secs * 1000 + 1;
node.resend_pending_session_handshakes(after_timeout).await;
assert!(!node.sessions.contains_key(&src_addr), "Timed-out AwaitingMsg3 session should be removed");
assert!(
!node.sessions.contains_key(&src_addr),
"Timed-out AwaitingMsg3 session should be removed"
);
}
#[tokio::test]
@@ -1706,7 +1815,11 @@ async fn test_tun_outbound_path_mtu_generates_ptb() {
let dst_fips = crate::FipsAddress::from_node_addr(&node1_addr);
// Establish session (XX: 3 messages — Setup, Ack, Msg3)
nodes[0].node.initiate_session(node1_addr, node1_pubkey).await.unwrap();
nodes[0]
.node
.initiate_session(node1_addr, node1_pubkey)
.await
.unwrap();
tokio::time::sleep(Duration::from_millis(20)).await;
process_available_packets(&mut nodes).await;
tokio::time::sleep(Duration::from_millis(20)).await;
@@ -1714,7 +1827,14 @@ async fn test_tun_outbound_path_mtu_generates_ptb() {
tokio::time::sleep(Duration::from_millis(20)).await;
process_available_packets(&mut nodes).await;
assert!(nodes[0].node.get_session(&node1_addr).unwrap().state().is_established());
assert!(
nodes[0]
.node
.get_session(&node1_addr)
.unwrap()
.state()
.is_established()
);
// Simulate receipt of MtuExceeded by reducing PathMtuState to a value
// lower than the local transport MTU.
@@ -1723,7 +1843,8 @@ async fn test_tun_outbound_path_mtu_generates_ptb() {
{
let entry = nodes[0].node.get_session_mut(&node1_addr).unwrap();
let mmp = entry.mmp_mut().unwrap();
mmp.path_mtu.apply_notification(reduced_mtu, std::time::Instant::now());
mmp.path_mtu
.apply_notification(reduced_mtu, std::time::Instant::now());
assert_eq!(mmp.path_mtu.current_mtu(), reduced_mtu);
}
@@ -1736,14 +1857,24 @@ async fn test_tun_outbound_path_mtu_generates_ptb() {
let local_ipv6_mtu = nodes[0].node.effective_ipv6_mtu() as usize;
let oversized_payload = vec![0u8; reduced_ipv6_mtu - 39]; // 40-byte hdr + payload > reduced MTU
let ipv6_packet = build_ipv6_packet(&src_fips, &dst_fips, &oversized_payload);
assert!(ipv6_packet.len() > reduced_ipv6_mtu, "packet must exceed path MTU");
assert!(ipv6_packet.len() <= local_ipv6_mtu, "packet must fit local MTU");
assert!(
ipv6_packet.len() > reduced_ipv6_mtu,
"packet must exceed path MTU"
);
assert!(
ipv6_packet.len() <= local_ipv6_mtu,
"packet must fit local MTU"
);
nodes[0].node.handle_tun_outbound(ipv6_packet).await;
// Verify ICMPv6 Packet Too Big was generated
let ptb_messages: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx.try_recv().ok()).collect();
assert_eq!(ptb_messages.len(), 1, "Should generate exactly one ICMPv6 PTB");
assert_eq!(
ptb_messages.len(),
1,
"Should generate exactly one ICMPv6 PTB"
);
let ptb = &ptb_messages[0];
assert_eq!(ptb[0] >> 4, 6, "Should be IPv6");
@@ -1756,12 +1887,23 @@ async fn test_tun_outbound_path_mtu_generates_ptb() {
// address, causing a PMTUD blackhole.
let ptb_src = std::net::Ipv6Addr::from(<[u8; 16]>::try_from(&ptb[8..24]).unwrap());
let ptb_dst = std::net::Ipv6Addr::from(<[u8; 16]>::try_from(&ptb[24..40]).unwrap());
assert_eq!(ptb_src, dst_fips.to_ipv6(), "PTB source must be remote peer (original dst), not local node");
assert_eq!(ptb_dst, src_fips.to_ipv6(), "PTB destination must be local node (original src)");
assert_eq!(
ptb_src,
dst_fips.to_ipv6(),
"PTB source must be remote peer (original dst), not local node"
);
assert_eq!(
ptb_dst,
src_fips.to_ipv6(),
"PTB destination must be local node (original src)"
);
// Verify reported MTU (32-bit field at ICMPv6 header bytes 4-7)
let reported_mtu = u32::from_be_bytes([ptb[44], ptb[45], ptb[46], ptb[47]]);
assert_eq!(reported_mtu, reduced_ipv6_mtu as u32, "Reported MTU should match path IPv6 MTU");
assert_eq!(
reported_mtu, reduced_ipv6_mtu as u32,
"Reported MTU should match path IPv6 MTU"
);
// Verify a packet that fits within path MTU passes through (no PTB)
let (tun_tx2, tun_rx2) = std::sync::mpsc::channel();
@@ -1774,7 +1916,11 @@ async fn test_tun_outbound_path_mtu_generates_ptb() {
// No PTB should be generated for a fitting packet
let ptb_messages2: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx2.try_recv().ok()).collect();
assert_eq!(ptb_messages2.len(), 0, "Should not generate PTB for fitting packet");
assert_eq!(
ptb_messages2.len(),
0,
"Should not generate PTB for fitting packet"
);
cleanup_nodes(&mut nodes).await;
}
@@ -1817,10 +1963,19 @@ async fn test_multihop_pmtud_heterogeneous_mtu() {
nodes[0].node.register_identity(node2_addr, node2_pubkey);
// Establish session A→C via B (triggers routing through tree)
nodes[0].node.initiate_session(node2_addr, node2_pubkey).await.unwrap();
nodes[0]
.node
.initiate_session(node2_addr, node2_pubkey)
.await
.unwrap();
drain_to_quiescence(&mut nodes).await;
assert!(
nodes[0].node.get_session(&node2_addr).unwrap().state().is_established(),
nodes[0]
.node
.get_session(&node2_addr)
.unwrap()
.state()
.is_established(),
"Session A→C should be established"
);
@@ -1830,7 +1985,11 @@ async fn test_multihop_pmtud_heterogeneous_mtu() {
// With coords (~66 extra), the wire could exceed B's recv buffer.
for _ in 0..5 {
let small = build_ipv6_packet(&src_fips, &dst_fips, &[0u8; 10]);
nodes[0].node.send_ipv6_packet(&node2_addr, &small).await.unwrap();
nodes[0]
.node
.send_ipv6_packet(&node2_addr, &small)
.await
.unwrap();
}
drain_to_quiescence(&mut nodes).await;
@@ -1844,12 +2003,17 @@ async fn test_multihop_pmtud_heterogeneous_mtu() {
assert!(
ipv6_packet.len() <= local_effective_mtu,
"packet ({}) must fit A's local MTU ({})",
ipv6_packet.len(), local_effective_mtu
ipv6_packet.len(),
local_effective_mtu
);
// Send the oversized packet — B should fail to forward and send
// MtuExceeded signal back.
nodes[0].node.send_ipv6_packet(&node2_addr, &ipv6_packet).await.unwrap();
nodes[0]
.node
.send_ipv6_packet(&node2_addr, &ipv6_packet)
.await
.unwrap();
drain_to_quiescence(&mut nodes).await;
// Verify PathMtuState was updated on A
@@ -1874,7 +2038,8 @@ async fn test_multihop_pmtud_heterogeneous_mtu() {
let ptb_messages: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx2.try_recv().ok()).collect();
assert_eq!(
ptb_messages.len(), 1,
ptb_messages.len(),
1,
"Should generate ICMPv6 PTB for oversized packet after PathMtuState update"
);
@@ -1889,8 +2054,16 @@ async fn test_multihop_pmtud_heterogeneous_mtu() {
// address, causing a PMTUD blackhole.
let ptb_src = std::net::Ipv6Addr::from(<[u8; 16]>::try_from(&ptb[8..24]).unwrap());
let ptb_dst = std::net::Ipv6Addr::from(<[u8; 16]>::try_from(&ptb[24..40]).unwrap());
assert_eq!(ptb_src, dst_fips.to_ipv6(), "PTB source must be remote peer (original dst), not local node");
assert_eq!(ptb_dst, src_fips.to_ipv6(), "PTB destination must be local node (original src)");
assert_eq!(
ptb_src,
dst_fips.to_ipv6(),
"PTB source must be remote peer (original dst), not local node"
);
assert_eq!(
ptb_dst,
src_fips.to_ipv6(),
"PTB destination must be local node (original src)"
);
// Verify reported MTU is the path MTU (not local MTU)
let reported_mtu = u32::from_be_bytes([ptb[44], ptb[45], ptb[46], ptb[47]]);
@@ -1913,7 +2086,8 @@ async fn test_multihop_pmtud_heterogeneous_mtu() {
let ptb_messages3: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx3.try_recv().ok()).collect();
assert_eq!(
ptb_messages3.len(), 0,
ptb_messages3.len(),
0,
"Should not generate PTB for packet fitting within path MTU"
);
+26 -14
View File
@@ -69,7 +69,9 @@ pub(super) async fn initiate_handshake(nodes: &mut [TestNode], i: usize, j: usiz
let our_index = initiator.node.index_allocator.allocate().unwrap();
let our_keypair = initiator.node.identity().keypair();
let noise_msg1 = conn.start_handshake(our_keypair, initiator.node.startup_epoch, 1000).unwrap();
let noise_msg1 = conn
.start_handshake(our_keypair, initiator.node.startup_epoch, 1000)
.unwrap();
conn.set_our_index(our_index);
conn.set_transport_id(transport_id);
conn.set_source_addr(responder_addr.clone());
@@ -184,7 +186,12 @@ pub(super) fn print_tree_snapshot(label: &str, nodes: &[TestNode]) {
.count();
eprintln!(
" node[{}] root=node[{}] depth={} parent=node[{}] peers={} pending={}",
i, root_idx, ts.my_coords().depth(), parent_idx, tn.node.peer_count(), pending,
i,
root_idx,
ts.my_coords().depth(),
parent_idx,
tn.node.peer_count(),
pending,
);
}
} else if correct_root_count < nodes.len() {
@@ -209,7 +216,10 @@ pub(super) fn print_tree_snapshot(label: &str, nodes: &[TestNode]) {
///
/// Returns the number of packets processed.
pub(super) async fn process_available_packets(nodes: &mut [TestNode]) -> usize {
use crate::node::wire::{CommonPrefix, FMP_VERSION, PHASE_ESTABLISHED, PHASE_MSG1, PHASE_MSG2, PHASE_MSG3, COMMON_PREFIX_SIZE};
use crate::node::wire::{
COMMON_PREFIX_SIZE, CommonPrefix, FMP_VERSION, PHASE_ESTABLISHED, PHASE_MSG1, PHASE_MSG2,
PHASE_MSG3,
};
let mut count = 0;
for node in nodes.iter_mut() {
@@ -225,9 +235,7 @@ pub(super) async fn process_available_packets(nodes: &mut [TestNode]) -> usize {
PHASE_MSG1 => node.node.handle_msg1(packet).await,
PHASE_MSG2 => node.node.handle_msg2(packet).await,
PHASE_MSG3 => node.node.handle_msg3(packet).await,
PHASE_ESTABLISHED => {
node.node.handle_encrypted_frame(packet).await
}
PHASE_ESTABLISHED => node.node.handle_encrypted_frame(packet).await,
_ => {}
}
count += 1;
@@ -320,7 +328,11 @@ pub(super) async fn drain_all_packets(nodes: &mut [TestNode], verbose: bool) ->
///
/// First builds a random spanning tree to ensure connectivity,
/// then adds extra edges up to the target count.
pub(super) fn generate_random_edges(n: usize, target_edges: usize, seed: u64) -> Vec<(usize, usize)> {
pub(super) fn generate_random_edges(
n: usize,
target_edges: usize,
seed: u64,
) -> Vec<(usize, usize)> {
use rand::rngs::StdRng;
use rand::{RngExt, SeedableRng};
@@ -374,11 +386,7 @@ pub(super) fn verify_tree_convergence(nodes: &[TestNode]) {
assert!(n > 0);
// Find the expected root (smallest NodeAddr across all nodes)
let expected_root = nodes
.iter()
.map(|tn| *tn.node.node_addr())
.min()
.unwrap();
let expected_root = nodes.iter().map(|tn| *tn.node.node_addr()).min().unwrap();
// All nodes should agree on the root
for (i, tn) in nodes.iter().enumerate() {
@@ -628,12 +636,16 @@ pub(super) async fn run_tree_test_with_mtus(
assert!(
nodes[i].node.get_peer(&j_addr).is_some(),
"Node {} should have peer {} (node {})",
i, j_addr, j
i,
j_addr,
j
);
assert!(
nodes[j].node.get_peer(&i_addr).is_some(),
"Node {} should have peer {} (node {})",
j, i_addr, i
j,
i_addr,
i
);
}
+3 -1
View File
@@ -511,7 +511,9 @@ fn test_promote_cleans_up_pending_outbound_to_same_peer() {
let (msg3, _neg) = completing_conn
.complete_handshake(&msg2, None, completing_time_ms)
.unwrap();
resp_conn.complete_handshake_msg3(&msg3, completing_time_ms).unwrap();
resp_conn
.complete_handshake_msg3(&msg3, completing_time_ms)
.unwrap();
let completing_index = node.index_allocator.allocate().unwrap();
completing_conn.set_our_index(completing_index);
+22 -20
View File
@@ -18,8 +18,8 @@
//! | 0x2 | Noise XX msg2 | 118+ bytes | Handshake response |
//! | 0x3 | Noise XX msg3 | 85+ bytes | Handshake completion |
use crate::utils::index::SessionIndex;
use crate::noise::{HANDSHAKE_MSG1_SIZE, HANDSHAKE_MSG2_SIZE, HANDSHAKE_MSG3_SIZE, TAG_SIZE};
use crate::utils::index::SessionIndex;
// ============================================================================
// Constants
@@ -170,8 +170,7 @@ impl EncryptedHeader {
let payload_len = u16::from_le_bytes([data[2], data[3]]);
let receiver_idx = SessionIndex::from_le_bytes([data[4], data[5], data[6], data[7]]);
let counter = u64::from_le_bytes([
data[8], data[9], data[10], data[11],
data[12], data[13], data[14], data[15],
data[8], data[9], data[10], data[11], data[12], data[13], data[14], data[15],
]);
let mut header_bytes = [0u8; ESTABLISHED_HEADER_SIZE];
@@ -409,7 +408,11 @@ pub fn build_msg1(sender_idx: SessionIndex, noise_msg1: &[u8]) -> Vec<u8> {
///
/// Format: `[0x12][0x00][payload_len:2 LE][sender_idx:4 LE][receiver_idx:4 LE][noise_msg2:106+]`
/// The noise_msg2 may include an optional negotiation payload beyond the base XX msg2.
pub fn build_msg2(sender_idx: SessionIndex, receiver_idx: SessionIndex, noise_msg2: &[u8]) -> Vec<u8> {
pub fn build_msg2(
sender_idx: SessionIndex,
receiver_idx: SessionIndex,
noise_msg2: &[u8],
) -> Vec<u8> {
debug_assert!(noise_msg2.len() >= HANDSHAKE_MSG2_SIZE);
let payload_len = (4 + 4 + noise_msg2.len()) as u16; // sender + receiver + noise
@@ -429,7 +432,11 @@ pub fn build_msg2(sender_idx: SessionIndex, receiver_idx: SessionIndex, noise_ms
///
/// Format: `[0x13][0x00][payload_len:2 LE][sender_idx:4 LE][receiver_idx:4 LE][noise_msg3:73+]`
/// The noise_msg3 may include an optional negotiation payload beyond the base XX msg3.
pub fn build_msg3(sender_idx: SessionIndex, receiver_idx: SessionIndex, noise_msg3: &[u8]) -> Vec<u8> {
pub fn build_msg3(
sender_idx: SessionIndex,
receiver_idx: SessionIndex,
noise_msg3: &[u8],
) -> Vec<u8> {
debug_assert!(noise_msg3.len() >= HANDSHAKE_MSG3_SIZE);
let payload_len = (4 + 4 + noise_msg3.len()) as u16; // sender + receiver + noise
@@ -648,9 +655,9 @@ mod tests {
#[test]
fn test_wire_sizes() {
assert_eq!(MSG1_WIRE_SIZE, 41); // 4 + 4 + 33 (XX msg1)
assert_eq!(MSG2_WIRE_SIZE, 118); // 4 + 4 + 4 + 106 (XX msg2 minimum)
assert_eq!(MSG3_WIRE_SIZE, 85); // 4 + 4 + 4 + 73 (XX msg3 minimum)
assert_eq!(MSG1_WIRE_SIZE, 41); // 4 + 4 + 33 (XX msg1)
assert_eq!(MSG2_WIRE_SIZE, 118); // 4 + 4 + 4 + 106 (XX msg2 minimum)
assert_eq!(MSG3_WIRE_SIZE, 85); // 4 + 4 + 4 + 73 (XX msg3 minimum)
assert_eq!(ENCRYPTED_MIN_SIZE, 32); // 16 + 16
assert_eq!(COMMON_PREFIX_SIZE, 4);
assert_eq!(ESTABLISHED_HEADER_SIZE, 16);
@@ -692,22 +699,17 @@ mod tests {
#[test]
fn test_flags_byte() {
let header = build_established_header(
SessionIndex::new(1),
0,
FLAG_KEY_EPOCH | FLAG_CE,
100,
);
let header =
build_established_header(SessionIndex::new(1), 0, FLAG_KEY_EPOCH | FLAG_CE, 100);
assert_eq!(header[1], 0x03); // bits 0 and 1 set
let parsed = EncryptedHeader::parse(&[
header[0], header[1], header[2], header[3],
header[4], header[5], header[6], header[7],
header[8], header[9], header[10], header[11],
header[12], header[13], header[14], header[15],
// minimum: TAG_SIZE bytes of ciphertext
header[0], header[1], header[2], header[3], header[4], header[5], header[6], header[7],
header[8], header[9], header[10], header[11], header[12], header[13], header[14],
header[15], // minimum: TAG_SIZE bytes of ciphertext
0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
]).unwrap();
])
.unwrap();
assert_eq!(parsed.flags & FLAG_KEY_EPOCH, FLAG_KEY_EPOCH);
assert_eq!(parsed.flags & FLAG_CE, FLAG_CE);
}
+17 -8
View File
@@ -1,11 +1,11 @@
use super::{
CipherState, HandshakeProgress, HandshakeRole, NoiseError, NoisePattern, NoiseSession,
EPOCH_ENCRYPTED_SIZE, EPOCH_SIZE, HANDSHAKE_MSG1_SIZE, HANDSHAKE_MSG2_SIZE,
HANDSHAKE_MSG3_SIZE, PROTOCOL_NAME_XX, PUBKEY_SIZE,
CipherState, EPOCH_ENCRYPTED_SIZE, EPOCH_SIZE, HANDSHAKE_MSG1_SIZE, HANDSHAKE_MSG2_SIZE,
HANDSHAKE_MSG3_SIZE, HandshakeProgress, HandshakeRole, NoiseError, NoisePattern, NoiseSession,
PROTOCOL_NAME_XX, PUBKEY_SIZE,
};
use hkdf::Hkdf;
use rand::Rng;
use secp256k1::{ecdh::shared_secret_point, Keypair, PublicKey, Secp256k1, SecretKey};
use secp256k1::{Keypair, PublicKey, Secp256k1, SecretKey, ecdh::shared_secret_point};
use sha2::{Digest, Sha256};
use std::fmt;
@@ -333,7 +333,9 @@ impl HandshakeState {
}
let re = self.remote_ephemeral.expect("should have remote ephemeral");
let epoch = self.local_epoch.expect("local epoch must be set before write_message_2");
let epoch = self
.local_epoch
.expect("local epoch must be set before write_message_2");
// Generate ephemeral keypair
self.generate_ephemeral();
@@ -452,8 +454,12 @@ impl HandshakeState {
});
}
let re = self.remote_ephemeral.expect("should have remote ephemeral after msg2");
let epoch = self.local_epoch.expect("local epoch must be set before write_message_3");
let re = self
.remote_ephemeral
.expect("should have remote ephemeral after msg2");
let epoch = self
.local_epoch
.expect("local epoch must be set before write_message_3");
let mut message = Vec::with_capacity(HANDSHAKE_MSG3_SIZE);
@@ -510,7 +516,10 @@ impl HandshakeState {
// -> se: DH(e, rs), mix into key
// (responder uses their ephemeral with initiator's now-known static)
let ephemeral = self.ephemeral_keypair.as_ref().expect("should have ephemeral after msg2");
let ephemeral = self
.ephemeral_keypair
.as_ref()
.expect("should have ephemeral after msg2");
let se = self.ecdh(&ephemeral.secret_key(), &rs);
self.symmetric.mix_key(&se);
+58 -17
View File
@@ -132,21 +132,42 @@ fn test_identity_timing() {
// After msg1
let msg1 = initiator.write_message_1().unwrap();
responder.read_message_1(&msg1).unwrap();
assert!(initiator.remote_static().is_none(), "XX: initiator should NOT know identity after msg1");
assert!(responder.remote_static().is_none(), "XX: responder should NOT know identity after msg1");
assert!(
initiator.remote_static().is_none(),
"XX: initiator should NOT know identity after msg1"
);
assert!(
responder.remote_static().is_none(),
"XX: responder should NOT know identity after msg1"
);
// After msg2: initiator learns responder
let msg2 = responder.write_message_2().unwrap();
initiator.read_message_2(&msg2).unwrap();
assert!(initiator.remote_static().is_some(), "XX: initiator should know responder after msg2");
assert_eq!(initiator.remote_static().unwrap(), &responder_keypair.public_key());
assert!(responder.remote_static().is_none(), "XX: responder should NOT know initiator after msg2");
assert!(
initiator.remote_static().is_some(),
"XX: initiator should know responder after msg2"
);
assert_eq!(
initiator.remote_static().unwrap(),
&responder_keypair.public_key()
);
assert!(
responder.remote_static().is_none(),
"XX: responder should NOT know initiator after msg2"
);
// After msg3: responder learns initiator
let msg3 = initiator.write_message_3().unwrap();
responder.read_message_3(&msg3).unwrap();
assert!(responder.remote_static().is_some(), "XX: responder should know initiator after msg3");
assert_eq!(responder.remote_static().unwrap(), &initiator_keypair.public_key());
assert!(
responder.remote_static().is_some(),
"XX: responder should know initiator after msg3"
);
assert_eq!(
responder.remote_static().unwrap(),
&initiator_keypair.public_key()
);
}
#[test]
@@ -157,7 +178,11 @@ fn test_wrong_state_errors() {
// Initiator can't read msg1
let mut initiator = HandshakeState::new_initiator(keypair1);
initiator.set_local_epoch(generate_epoch());
assert!(initiator.read_message_1(&[0u8; HANDSHAKE_MSG1_SIZE]).is_err());
assert!(
initiator
.read_message_1(&[0u8; HANDSHAKE_MSG1_SIZE])
.is_err()
);
// Initiator can't write msg2
assert!(initiator.write_message_2().is_err());
@@ -171,7 +196,11 @@ fn test_wrong_state_errors() {
assert!(responder.write_message_1().is_err());
// Responder can't read msg3 before msg2
assert!(responder.read_message_3(&[0u8; HANDSHAKE_MSG3_SIZE]).is_err());
assert!(
responder
.read_message_3(&[0u8; HANDSHAKE_MSG3_SIZE])
.is_err()
);
}
#[test]
@@ -214,21 +243,23 @@ fn test_with_odd_parity() {
// Node A (initiator) - even parity key
let sk_a = secp256k1::SecretKey::from_slice(
&hex::decode("0102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1f20")
.unwrap(),
&hex::decode("0102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1f20").unwrap(),
)
.unwrap();
let kp_a = secp256k1::Keypair::from_secret_key(&secp, &sk_a);
// Node B (responder) - odd parity key
let sk_b = secp256k1::SecretKey::from_slice(
&hex::decode("b102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1fb0")
.unwrap(),
&hex::decode("b102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1fb0").unwrap(),
)
.unwrap();
let kp_b = secp256k1::Keypair::from_secret_key(&secp, &sk_b);
let (_, parity_b) = kp_b.public_key().x_only_public_key();
assert_eq!(parity_b, Parity::Odd, "Test requires odd-parity responder key");
assert_eq!(
parity_b,
Parity::Odd,
"Test requires odd-parity responder key"
);
let mut initiator = HandshakeState::new_initiator(kp_a);
initiator.set_local_epoch(generate_epoch());
@@ -250,7 +281,9 @@ fn test_with_odd_parity() {
let counter = sender.current_send_counter();
let ciphertext = sender.encrypt(b"xx parity test").unwrap();
let plaintext = receiver.decrypt_with_replay_check(&ciphertext, counter).unwrap();
let plaintext = receiver
.decrypt_with_replay_check(&ciphertext, counter)
.unwrap();
assert_eq!(plaintext, b"xx parity test");
}
@@ -276,7 +309,11 @@ fn test_invalid_msg_sizes() {
// Responder is now in Message2Done, try wrong-size msg3
assert!(responder.read_message_3(&[0u8; 10]).is_err());
assert!(responder.read_message_3(&[0u8; HANDSHAKE_MSG3_SIZE + 1]).is_err());
assert!(
responder
.read_message_3(&[0u8; HANDSHAKE_MSG3_SIZE + 1])
.is_err()
);
}
#[test]
@@ -440,7 +477,11 @@ fn test_replay_window_sequential() {
// All should be marked as seen
for i in 0..1000 {
assert!(!window.check(i), "Counter {} should be rejected as replay", i);
assert!(
!window.check(i),
"Counter {} should be rejected as replay",
i
);
}
assert_eq!(window.highest(), 999);
+23 -29
View File
@@ -5,11 +5,11 @@
use crate::bloom::BloomFilter;
use crate::mmp::{MmpConfig, MmpPeerState};
use crate::protocol::{NegotiationPayload, NodeProfile};
use crate::utils::index::SessionIndex;
use crate::noise::{HandshakeState as NoiseHandshakeState, NoiseError, NoiseSession};
use crate::protocol::{NegotiationPayload, NodeProfile};
use crate::transport::{LinkId, LinkStats, TransportAddr, TransportId};
use crate::tree::{ParentDeclaration, TreeCoordinate};
use crate::utils::index::SessionIndex;
use crate::{FipsAddress, NodeAddr, PeerIdentity};
use secp256k1::XOnlyPublicKey;
use std::fmt;
@@ -33,7 +33,10 @@ pub enum ConnectivityState {
impl ConnectivityState {
/// Check if the peer is usable for sending traffic.
pub fn can_send(&self) -> bool {
matches!(self, ConnectivityState::Connected | ConnectivityState::Stale)
matches!(
self,
ConnectivityState::Connected | ConnectivityState::Stale
)
}
/// Check if this is a terminal state requiring cleanup.
@@ -793,12 +796,7 @@ impl ActivePeer {
// === Filter Updates ===
/// Update peer's inbound filter.
pub fn update_filter(
&mut self,
filter: BloomFilter,
sequence: u64,
current_time_ms: u64,
) {
pub fn update_filter(&mut self, filter: BloomFilter, sequence: u64, current_time_ms: u64) {
self.inbound_filter = Some(filter);
self.filter_sequence = sequence;
self.filter_received_at = current_time_ms;
@@ -1013,12 +1011,11 @@ impl ActivePeer {
self.rekey_msg1_next_resend = 0;
self.rekey_in_progress = false;
// Return whichever index needs freeing
self.rekey_our_index.take()
.or_else(|| {
self.pending_new_session = None;
self.pending_their_index = None;
self.pending_our_index.take()
})
self.rekey_our_index.take().or_else(|| {
self.pending_new_session = None;
self.pending_their_index = None;
self.pending_our_index.take()
})
}
// === Rekey Handshake State (Initiator) ===
@@ -1053,7 +1050,8 @@ impl ActivePeer {
&mut self,
msg2_bytes: &[u8],
) -> Result<(Vec<u8>, NoiseSession), NoiseError> {
let mut hs = self.rekey_handshake
let mut hs = self
.rekey_handshake
.take()
.ok_or_else(|| NoiseError::WrongState {
expected: "rekey handshake in progress".to_string(),
@@ -1090,16 +1088,14 @@ impl ActivePeer {
///
/// Takes the stored responder handshake state, reads XX msg3, and returns
/// the completed NoiseSession.
pub fn complete_rekey_msg3(
&mut self,
msg3_bytes: &[u8],
) -> Result<NoiseSession, NoiseError> {
let mut hs = self.rekey_responder_handshake
.take()
.ok_or_else(|| NoiseError::WrongState {
expected: "rekey responder handshake awaiting msg3".to_string(),
got: "no responder handshake state".to_string(),
})?;
pub fn complete_rekey_msg3(&mut self, msg3_bytes: &[u8]) -> Result<NoiseSession, NoiseError> {
let mut hs =
self.rekey_responder_handshake
.take()
.ok_or_else(|| NoiseError::WrongState {
expected: "rekey responder handshake awaiting msg3".to_string(),
got: "no responder handshake state".to_string(),
})?;
// Split msg3 into base XX part and any extra (negotiation payload)
let base_size = crate::noise::HANDSHAKE_MSG3_SIZE;
@@ -1126,9 +1122,7 @@ impl ActivePeer {
/// Check if msg1 needs resending.
pub fn needs_msg1_resend(&self, now_ms: u64) -> bool {
self.rekey_in_progress
&& self.rekey_msg1.is_some()
&& now_ms >= self.rekey_msg1_next_resend
self.rekey_in_progress && self.rekey_msg1.is_some() && now_ms >= self.rekey_msg1_next_resend
}
/// Get msg1 bytes for resend (without consuming).
+25 -14
View File
@@ -4,11 +4,11 @@
//! PeerConnection tracks the Noise XX handshake state and transitions to
//! ActivePeer upon successful authentication.
use crate::protocol::NodeProfile;
use crate::utils::index::SessionIndex;
use crate::noise::{self, NoiseError, NoiseSession};
use crate::transport::{LinkDirection, LinkId, LinkStats, TransportAddr, TransportId};
use crate::PeerIdentity;
use crate::noise::{self, NoiseError, NoiseSession};
use crate::protocol::NodeProfile;
use crate::transport::{LinkDirection, LinkId, LinkStats, TransportAddr, TransportId};
use crate::utils::index::SessionIndex;
use secp256k1::Keypair;
use std::fmt;
@@ -708,7 +708,6 @@ impl PeerConnection {
pub fn is_timed_out(&self, current_time_ms: u64, timeout_ms: u64) -> bool {
self.idle_time(current_time_ms) > timeout_ms
}
}
impl fmt::Debug for PeerConnection {
@@ -801,24 +800,30 @@ mod tests {
let responder_peer_id = PeerIdentity::from_pubkey_full(responder_identity.pubkey_full());
// Create connections
let mut initiator_conn =
PeerConnection::outbound(LinkId::new(1), responder_peer_id, 1000);
let mut initiator_conn = PeerConnection::outbound(LinkId::new(1), responder_peer_id, 1000);
let mut responder_conn = PeerConnection::inbound(LinkId::new(2), 1000);
// Initiator starts XX handshake
let msg1 = initiator_conn.start_handshake(initiator_keypair, initiator_epoch, 1100).unwrap();
let msg1 = initiator_conn
.start_handshake(initiator_keypair, initiator_epoch, 1100)
.unwrap();
assert_eq!(initiator_conn.handshake_state(), HandshakeState::SentMsg1);
// Responder processes msg1 and sends msg2 (XX: does NOT complete yet)
let msg2 = responder_conn
.receive_handshake_init(responder_keypair, responder_epoch, &msg1, None, 1200)
.unwrap();
assert_eq!(responder_conn.handshake_state(), HandshakeState::ReceivedMsg1);
assert_eq!(
responder_conn.handshake_state(),
HandshakeState::ReceivedMsg1
);
// Responder does NOT know initiator's identity yet (XX property)
assert!(responder_conn.expected_identity().is_none());
// Initiator processes msg2 and generates msg3
let (msg3, _neg) = initiator_conn.complete_handshake(&msg2, None, 1300).unwrap();
let (msg3, _neg) = initiator_conn
.complete_handshake(&msg2, None, 1300)
.unwrap();
assert_eq!(initiator_conn.handshake_state(), HandshakeState::Complete);
// Initiator learned responder's identity from msg2
@@ -879,12 +884,18 @@ mod tests {
// Outbound can't receive_handshake_init
let mut outbound = PeerConnection::outbound(LinkId::new(1), identity, 1000);
assert!(outbound
.receive_handshake_init(keypair, make_epoch(), &[0u8; 33], None, 1100)
.is_err());
assert!(
outbound
.receive_handshake_init(keypair, make_epoch(), &[0u8; 33], None, 1100)
.is_err()
);
// Inbound can't start_handshake
let mut inbound = PeerConnection::inbound(LinkId::new(2), 1000);
assert!(inbound.start_handshake(keypair, make_epoch(), 1100).is_err());
assert!(
inbound
.start_handshake(keypair, make_epoch(), 1100)
.is_err()
);
}
}
+10 -27
View File
@@ -31,13 +31,7 @@ pub struct LookupRequest {
impl LookupRequest {
/// Create a new lookup request.
pub fn new(
request_id: u64,
target: NodeAddr,
origin: NodeAddr,
ttl: u8,
min_mtu: u16,
) -> Self {
pub fn new(request_id: u64, target: NodeAddr, origin: NodeAddr, ttl: u8, min_mtu: u16) -> Self {
Self {
request_id,
target,
@@ -49,12 +43,7 @@ impl LookupRequest {
}
/// Generate a new request with a random ID.
pub fn generate(
target: NodeAddr,
origin: NodeAddr,
ttl: u8,
min_mtu: u16,
) -> Self {
pub fn generate(target: NodeAddr, origin: NodeAddr, ttl: u8, min_mtu: u16) -> Self {
use rand::RngExt;
let request_id = rand::rng().random();
Self::new(request_id, target, origin, ttl, min_mtu)
@@ -152,11 +141,8 @@ impl LookupRequest {
"truncated TLV header in LookupRequest".to_string(),
));
}
let field_num =
u16::from_le_bytes(payload[pos..pos + 2].try_into().unwrap());
let length =
u16::from_le_bytes(payload[pos + 2..pos + 4].try_into().unwrap())
as usize;
let field_num = u16::from_le_bytes(payload[pos..pos + 2].try_into().unwrap());
let length = u16::from_le_bytes(payload[pos + 2..pos + 4].try_into().unwrap()) as usize;
pos += 4;
if pos + length > payload.len() {
return Err(ProtocolError::Malformed(format!(
@@ -322,11 +308,8 @@ impl LookupResponse {
"truncated TLV header in LookupResponse".to_string(),
));
}
let field_num =
u16::from_le_bytes(payload[pos..pos + 2].try_into().unwrap());
let length =
u16::from_le_bytes(payload[pos + 2..pos + 4].try_into().unwrap())
as usize;
let field_num = u16::from_le_bytes(payload[pos..pos + 2].try_into().unwrap());
let length = u16::from_le_bytes(payload[pos + 2..pos + 4].try_into().unwrap()) as usize;
pos += 4;
if pos + length > payload.len() {
return Err(ProtocolError::Malformed(format!(
@@ -498,8 +481,8 @@ mod tests {
let target = make_node_addr(10);
let origin = make_node_addr(20);
let request = LookupRequest::new(777, target, origin, 5, 0)
.with_tlv(9999, vec![0xFF, 0xFE, 0xFD]);
let request =
LookupRequest::new(777, target, origin, 5, 0).with_tlv(9999, vec![0xFF, 0xFE, 0xFD]);
let encoded = request.encode();
let mut decoded = LookupRequest::decode(&encoded[1..]).unwrap();
@@ -601,8 +584,8 @@ mod tests {
let coords = make_coords(&[42, 1, 0]);
let sig = make_test_sig();
let response = LookupResponse::new(999, target, coords, sig)
.with_tlv(9999, vec![0xFF, 0xFE, 0xFD]);
let response =
LookupResponse::new(999, target, coords, sig).with_tlv(9999, vec![0xFF, 0xFE, 0xFD]);
let encoded = response.encode();
let mut decoded = LookupResponse::decode(&encoded[1..]).unwrap();
+5 -13
View File
@@ -4,8 +4,8 @@
use super::error::ProtocolError;
use super::link::LinkMessageType;
use crate::bloom::codec::{rle_decode, rle_encode, CompressionStats};
use crate::bloom::BloomFilter;
use crate::bloom::codec::{CompressionStats, rle_decode, rle_encode};
/// Flag bit: this is a delta (XOR diff) update, not a full filter.
const FLAG_DELTA: u8 = 0x01;
@@ -55,12 +55,7 @@ impl FilterAnnounce {
}
/// Create a delta (XOR diff) FilterAnnounce.
pub fn delta(
diff: BloomFilter,
sequence: u64,
base_seq: u64,
size_class: u8,
) -> Self {
pub fn delta(diff: BloomFilter, sequence: u64, base_seq: u64, size_class: u8) -> Self {
Self {
filter: diff,
sequence,
@@ -160,9 +155,8 @@ impl FilterAnnounce {
let expected_words = expected_bytes / 8;
let compressed_data = &payload[pos..];
let words = rle_decode(compressed_data, expected_words).map_err(|e| {
ProtocolError::Malformed(format!("RLE decode error: {e}"))
})?;
let words = rle_decode(compressed_data, expected_words)
.map_err(|e| ProtocolError::Malformed(format!("RLE decode error: {e}")))?;
// Convert words to bytes for BloomFilter construction
let mut bytes = Vec::with_capacity(expected_bytes);
@@ -171,9 +165,7 @@ impl FilterAnnounce {
}
let filter = BloomFilter::from_bytes(bytes, crate::bloom::DEFAULT_HASH_COUNT)
.map_err(|e| {
ProtocolError::Malformed(format!("invalid bloom filter: {e}"))
})?;
.map_err(|e| ProtocolError::Malformed(format!("invalid bloom filter: {e}")))?;
Ok(Self {
filter,
+13 -14
View File
@@ -29,26 +29,25 @@ mod session;
mod tree;
// Re-export all public types at protocol:: level
pub use error::ProtocolError;
pub use link::{
Disconnect, DisconnectReason, HandshakeMessageType, LinkMessageType, SessionDatagram,
SESSION_DATAGRAM_HEADER_SIZE,
};
pub use tree::TreeAnnounce;
pub use filter::{FilterAnnounce, FilterNack};
pub use discovery::{LookupRequest, LookupResponse};
pub use error::ProtocolError;
pub use filter::{FilterAnnounce, FilterNack};
pub use link::{
Disconnect, DisconnectReason, HandshakeMessageType, LinkMessageType,
SESSION_DATAGRAM_HEADER_SIZE, SessionDatagram,
};
pub use negotiation::{
NegotiationPayload, NodeProfile, TlvEntry, NEGOTIATION_HEADER_SIZE,
FMP_FEAT_PROFILE_MASK, FMP_FEAT_PROVIDES_RR, FMP_FEAT_PROVIDES_SR,
FMP_FEAT_WANTS_RR, FMP_FEAT_WANTS_SR,
FMP_FEAT_PROFILE_MASK, FMP_FEAT_PROVIDES_RR, FMP_FEAT_PROVIDES_SR, FMP_FEAT_WANTS_RR,
FMP_FEAT_WANTS_SR, NEGOTIATION_HEADER_SIZE, NegotiationPayload, NodeProfile, TlvEntry,
};
pub use session::{
CoordsRequired, FspFlags, FspInnerFlags, MtuExceeded, PathBroken, PathMtuNotification,
SessionAck, SessionFlags, SessionMessageType, SessionMsg3, SessionReceiverReport,
SessionSenderReport, SessionSetup, COORDS_REQUIRED_SIZE, MTU_EXCEEDED_SIZE,
PATH_MTU_NOTIFICATION_SIZE, SESSION_RECEIVER_REPORT_SIZE, SESSION_SENDER_REPORT_SIZE,
COORDS_REQUIRED_SIZE, CoordsRequired, FspFlags, FspInnerFlags, MTU_EXCEEDED_SIZE, MtuExceeded,
PATH_MTU_NOTIFICATION_SIZE, PathBroken, PathMtuNotification, SESSION_RECEIVER_REPORT_SIZE,
SESSION_SENDER_REPORT_SIZE, SessionAck, SessionFlags, SessionMessageType, SessionMsg3,
SessionReceiverReport, SessionSenderReport, SessionSetup,
};
pub(crate) use session::{coords_wire_size, decode_optional_coords, encode_coords};
pub use tree::TreeAnnounce;
/// Protocol version for message compatibility.
pub const PROTOCOL_VERSION: u8 = 1;
+41 -41
View File
@@ -168,9 +168,7 @@ impl NegotiationPayload {
while offset < data.len() {
// Need at least 4 bytes for field_num + length
if offset + 4 > data.len() {
return Err(ProtocolError::Malformed(
"truncated TLV header".to_string(),
));
return Err(ProtocolError::Malformed("truncated TLV header".to_string()));
}
let field_num = u16::from_le_bytes(data[offset..offset + 2].try_into().unwrap());
@@ -273,10 +271,7 @@ impl NegotiationPayload {
/// Validate that two profiles form a valid link pairing.
///
/// At least one side must be `Full` or the link is rejected.
pub fn validate_profiles(
ours: NodeProfile,
theirs: NodeProfile,
) -> Result<(), ProtocolError> {
pub fn validate_profiles(ours: NodeProfile, theirs: NodeProfile) -> Result<(), ProtocolError> {
if ours != NodeProfile::Full && theirs != NodeProfile::Full {
return Err(ProtocolError::Malformed(format!(
"invalid profile pairing: {:?} <-> {:?} (at least one must be Full)",
@@ -285,7 +280,6 @@ impl NegotiationPayload {
}
Ok(())
}
}
#[cfg(test)]
@@ -364,8 +358,7 @@ mod tests {
#[test]
fn test_unknown_tlv_forward_compat() {
// Unknown field_nums should be preserved through encode/decode
let payload = NegotiationPayload::new(0, 1, 0)
.with_tlv(9999, vec![0xFF, 0xFE, 0xFD]);
let payload = NegotiationPayload::new(0, 1, 0).with_tlv(9999, vec![0xFF, 0xFE, 0xFD]);
let encoded = payload.encode();
let decoded = NegotiationPayload::decode(&encoded).unwrap();
@@ -396,8 +389,7 @@ mod tests {
#[test]
fn test_truncated_tlv() {
let payload = NegotiationPayload::new(0, 1, 0)
.with_tlv(1, vec![0xAA, 0xBB, 0xCC]);
let payload = NegotiationPayload::new(0, 1, 0).with_tlv(1, vec![0xAA, 0xBB, 0xCC]);
let mut encoded = payload.encode();
// Truncate the TLV value (remove last byte)
@@ -456,7 +448,11 @@ mod tests {
#[test]
fn test_fmp_payload_roundtrip() {
for profile in [NodeProfile::Full, NodeProfile::NonRouting, NodeProfile::Leaf] {
for profile in [
NodeProfile::Full,
NodeProfile::NonRouting,
NodeProfile::Leaf,
] {
let original = NegotiationPayload::fmp(1, 1, profile);
let encoded = original.encode();
let decoded = NegotiationPayload::decode(&encoded).unwrap();
@@ -479,45 +475,49 @@ mod tests {
#[test]
fn test_validate_profiles_valid() {
// F↔F
assert!(NegotiationPayload::validate_profiles(
NodeProfile::Full, NodeProfile::Full
).is_ok());
assert!(
NegotiationPayload::validate_profiles(NodeProfile::Full, NodeProfile::Full).is_ok()
);
// F↔N
assert!(NegotiationPayload::validate_profiles(
NodeProfile::Full, NodeProfile::NonRouting
).is_ok());
assert!(
NegotiationPayload::validate_profiles(NodeProfile::Full, NodeProfile::NonRouting)
.is_ok()
);
// N↔F
assert!(NegotiationPayload::validate_profiles(
NodeProfile::NonRouting, NodeProfile::Full
).is_ok());
assert!(
NegotiationPayload::validate_profiles(NodeProfile::NonRouting, NodeProfile::Full)
.is_ok()
);
// F↔L
assert!(NegotiationPayload::validate_profiles(
NodeProfile::Full, NodeProfile::Leaf
).is_ok());
assert!(
NegotiationPayload::validate_profiles(NodeProfile::Full, NodeProfile::Leaf).is_ok()
);
// L↔F
assert!(NegotiationPayload::validate_profiles(
NodeProfile::Leaf, NodeProfile::Full
).is_ok());
assert!(
NegotiationPayload::validate_profiles(NodeProfile::Leaf, NodeProfile::Full).is_ok()
);
}
#[test]
fn test_validate_profiles_invalid() {
// N↔N
assert!(NegotiationPayload::validate_profiles(
NodeProfile::NonRouting, NodeProfile::NonRouting
).is_err());
assert!(
NegotiationPayload::validate_profiles(NodeProfile::NonRouting, NodeProfile::NonRouting)
.is_err()
);
// N↔L
assert!(NegotiationPayload::validate_profiles(
NodeProfile::NonRouting, NodeProfile::Leaf
).is_err());
assert!(
NegotiationPayload::validate_profiles(NodeProfile::NonRouting, NodeProfile::Leaf)
.is_err()
);
// L↔N
assert!(NegotiationPayload::validate_profiles(
NodeProfile::Leaf, NodeProfile::NonRouting
).is_err());
assert!(
NegotiationPayload::validate_profiles(NodeProfile::Leaf, NodeProfile::NonRouting)
.is_err()
);
// L↔L
assert!(NegotiationPayload::validate_profiles(
NodeProfile::Leaf, NodeProfile::Leaf
).is_err());
assert!(
NegotiationPayload::validate_profiles(NodeProfile::Leaf, NodeProfile::Leaf).is_err()
);
}
}
+11 -15
View File
@@ -56,7 +56,6 @@ pub type DefaultBleTransport = BleTransport<io::BluerIo>;
#[cfg(any(not(feature = "ble"), test))]
pub type DefaultBleTransport = BleTransport<io::MockBleIo>;
// ============================================================================
// BLE Transport
// ============================================================================
@@ -892,16 +891,13 @@ mod tests {
fn make_transport(
io: MockBleIo,
) -> (BleTransport<MockBleIo>, tokio::sync::mpsc::Receiver<ReceivedPacket>) {
) -> (
BleTransport<MockBleIo>,
tokio::sync::mpsc::Receiver<ReceivedPacket>,
) {
let (tx, rx) = tokio::sync::mpsc::channel(64);
let config = BleConfig::default();
let transport = BleTransport::new(
TransportId::new(1),
None,
config,
io,
tx,
);
let transport = BleTransport::new(TransportId::new(1), None, config, io, tx);
(transport, rx)
}
@@ -945,8 +941,7 @@ mod tests {
// Probe connect must succeed for peers to reach the discovery buffer
let local = test_addr(1);
io.set_connect_handler(move |addr, _psm| {
let (stream, _peer) =
io::MockBleStream::pair(local.clone(), addr.clone(), 2048);
let (stream, _peer) = io::MockBleStream::pair(local.clone(), addr.clone(), 2048);
Ok(stream)
});
let (mut transport, _rx) = make_transport(io);
@@ -973,8 +968,7 @@ mod tests {
let io = MockBleIo::new("hci0", test_addr(1));
let local = test_addr(1);
io.set_connect_handler(move |addr, _psm| {
let (stream, _peer) =
io::MockBleStream::pair(local.clone(), addr.clone(), 2048);
let (stream, _peer) = io::MockBleStream::pair(local.clone(), addr.clone(), 2048);
Ok(stream)
});
let (mut transport, _rx) = make_transport(io);
@@ -1005,7 +999,9 @@ mod tests {
let io = MockBleIo::new("hci0", test_addr(1));
let (transport, _rx) = make_transport(io);
let addr = test_addr(2).to_transport_addr();
assert_eq!(transport.connection_state_sync(&addr), ConnectionState::None);
assert_eq!(
transport.connection_state_sync(&addr),
ConnectionState::None
);
}
}
+27 -15
View File
@@ -1,8 +1,9 @@
//! Ethernet Transport Implementation
//!
//! Provides raw Ethernet transport for FIPS peer communication using
//! AF_PACKET sockets with SOCK_DGRAM. Works on wired Ethernet and WiFi
//! interfaces (kernel mac80211 abstracts 802.11 transparently).
//! Provides raw Ethernet transport for FIPS peer communication. On Linux,
//! uses AF_PACKET/SOCK_DGRAM sockets; on macOS, uses BPF devices (`/dev/bpf*`).
//! Works on wired Ethernet and WiFi interfaces (kernel mac80211 abstracts
//! 802.11 transparently on Linux).
pub mod discovery;
pub mod socket;
@@ -13,10 +14,8 @@ use super::{
TransportId, TransportState, TransportType,
};
use crate::config::EthernetConfig;
use discovery::{
build_beacon, parse_beacon, DiscoveryBuffer, FRAME_TYPE_BEACON, FRAME_TYPE_DATA,
};
use socket::{AsyncPacketSocket, PacketSocket, ETHERNET_BROADCAST};
use discovery::{DiscoveryBuffer, FRAME_TYPE_BEACON, FRAME_TYPE_DATA, build_beacon, parse_beacon};
use socket::{AsyncPacketSocket, ETHERNET_BROADCAST, PacketSocket};
use stats::EthernetStats;
use std::sync::Arc;
@@ -220,16 +219,30 @@ impl EthernetTransport {
return Err(TransportError::NotStarted);
}
// Abort beacon task
if let Some(task) = self.beacon_task.take() {
task.abort();
let _ = task.await;
// Signal the socket to shut down. On macOS this writes to the
// shutdown pipe, waking the reader thread's select() immediately.
// On Linux this is a no-op (AsyncFd cancellation handles it).
if let Some(ref socket) = self.socket {
socket.shutdown();
}
// Abort receive task
// Abort tasks. On Linux, safe to await since all I/O is
// AsyncFd-based and cancellation-safe. On macOS, do NOT await —
// on a current_thread runtime the aborted task can't be polled
// while we're blocked on the JoinHandle, causing a deadlock.
if let Some(task) = self.beacon_task.take() {
task.abort();
#[cfg(not(target_os = "macos"))]
{
let _ = task.await;
}
}
if let Some(task) = self.recv_task.take() {
task.abort();
let _ = task.await;
#[cfg(not(target_os = "macos"))]
{
let _ = task.await;
}
}
// Drop socket
@@ -379,8 +392,7 @@ async fn ethernet_receive_loop(
continue;
}
// buf[1] is flags (reserved, ignored for now)
let payload_len =
u16::from_le_bytes([buf[2], buf[3]]) as usize;
let payload_len = u16::from_le_bytes([buf[2], buf[3]]) as usize;
if payload_len > len - 4 {
trace!(
"Data frame length field ({payload_len}) exceeds \
+4 -1
View File
@@ -403,7 +403,10 @@ mod tests {
let mut cursor = Cursor::new(frame);
let err = read_fmp_packet(&mut cursor, 1400).await.unwrap_err();
assert!(matches!(err, StreamError::HandshakeSizeMismatch { phase: 0x3, .. }));
assert!(matches!(
err,
StreamError::HandshakeSizeMismatch { phase: 0x3, .. }
));
}
#[tokio::test]
+36 -9
View File
@@ -423,7 +423,10 @@ fn test_evaluate_parent_stays_root_when_smallest() {
make_coords(&[1, 0]),
);
assert_eq!(state.evaluate_parent(&HashMap::new(), &HashSet::new()), None);
assert_eq!(
state.evaluate_parent(&HashMap::new(), &HashSet::new()),
None
);
}
#[test]
@@ -445,7 +448,10 @@ fn test_evaluate_parent_no_switch_when_already_best() {
state.recompute_coords();
// Now evaluate — should return None since peer1 is already our parent
assert_eq!(state.evaluate_parent(&HashMap::new(), &HashSet::new()), None);
assert_eq!(
state.evaluate_parent(&HashMap::new(), &HashSet::new()),
None
);
}
#[test]
@@ -453,7 +459,10 @@ fn test_evaluate_parent_no_peers() {
let my_node = make_node_addr(5);
let state = TreeState::new(my_node);
assert_eq!(state.evaluate_parent(&HashMap::new(), &HashSet::new()), None);
assert_eq!(
state.evaluate_parent(&HashMap::new(), &HashSet::new()),
None
);
}
#[test]
@@ -507,7 +516,10 @@ fn test_evaluate_parent_rejects_loop_candidate() {
);
// Should return None — the only candidate creates a loop
assert_eq!(state.evaluate_parent(&HashMap::new(), &HashSet::new()), None);
assert_eq!(
state.evaluate_parent(&HashMap::new(), &HashSet::new()),
None
);
}
#[test]
@@ -628,7 +640,10 @@ fn test_find_next_hop_chain() {
add_peer(&mut state, 2, &[2, 1, 5, 0]);
let dest = make_coords(&[2, 1, 5, 0]);
assert_eq!(state.find_next_hop(&dest, &HashSet::new()), Some(make_node_addr(2)));
assert_eq!(
state.find_next_hop(&dest, &HashSet::new()),
Some(make_node_addr(2))
);
}
#[test]
@@ -640,7 +655,10 @@ fn test_find_next_hop_chain_indirect() {
add_peer(&mut state, 1, &[1, 5, 0]);
let dest = make_coords(&[2, 1, 5, 0]);
assert_eq!(state.find_next_hop(&dest, &HashSet::new()), Some(make_node_addr(1)));
assert_eq!(
state.find_next_hop(&dest, &HashSet::new()),
Some(make_node_addr(1))
);
}
#[test]
@@ -651,7 +669,10 @@ fn test_find_next_hop_toward_root() {
add_peer(&mut state, 1, &[1, 0]);
let dest = make_coords(&[0]);
assert_eq!(state.find_next_hop(&dest, &HashSet::new()), Some(make_node_addr(1)));
assert_eq!(
state.find_next_hop(&dest, &HashSet::new()),
Some(make_node_addr(1))
);
}
#[test]
@@ -666,7 +687,10 @@ fn test_find_next_hop_sibling() {
add_peer(&mut state, 3, &[3, 0]);
let dest = make_coords(&[3, 0]);
assert_eq!(state.find_next_hop(&dest, &HashSet::new()), Some(make_node_addr(3)));
assert_eq!(
state.find_next_hop(&dest, &HashSet::new()),
Some(make_node_addr(3))
);
}
#[test]
@@ -731,7 +755,10 @@ fn test_find_next_hop_best_of_multiple() {
add_peer(&mut state, 3, &[3, 1, 0]);
let dest = make_coords(&[7, 3, 1, 0]);
assert_eq!(state.find_next_hop(&dest, &HashSet::new()), Some(make_node_addr(3)));
assert_eq!(
state.find_next_hop(&dest, &HashSet::new()),
Some(make_node_addr(3))
);
}
// === Cost-based parent selection tests ===