From 34f840a77b11aa75f07c114e8b57600c926449c9 Mon Sep 17 00:00:00 2001 From: Johnathan Corgan Date: Fri, 10 Apr 2026 09:06:54 +0000 Subject: [PATCH] Apply rustfmt to noise-xx-only code --- src/bin/fipstop/ui/bloom.rs | 54 +++- src/bloom/codec.rs | 11 +- src/bloom/filter.rs | 12 +- src/bloom/state.rs | 2 +- src/bloom/tests.rs | 22 +- src/control/queries.rs | 101 ++++--- src/mmp/algorithms.rs | 1 - src/mmp/mod.rs | 4 +- src/mmp/report.rs | 4 +- src/node/bloom.rs | 26 +- src/node/handlers/discovery.rs | 50 ++-- src/node/handlers/encrypted.rs | 21 +- src/node/handlers/handshake.rs | 236 ++++++++-------- src/node/handlers/rx_loop.rs | 18 +- src/node/handlers/session.rs | 274 +++++++++++------- src/node/handlers/timeout.rs | 40 +-- src/node/lifecycle.rs | 170 +++++++++--- src/node/mod.rs | 203 ++++++++------ src/node/session_wire.rs | 1 - src/node/tests/disconnect.rs | 18 +- src/node/tests/handshake.rs | 350 +++++++++++++++-------- src/node/tests/mod.rs | 8 +- src/node/tests/session.rs | 474 ++++++++++++++++++++++---------- src/node/tests/spanning_tree.rs | 40 ++- src/node/tests/unit.rs | 4 +- src/node/wire.rs | 42 +-- src/noise/handshake.rs | 25 +- src/noise/tests.rs | 75 +++-- src/peer/active.rs | 52 ++-- src/peer/connection.rs | 39 ++- src/protocol/discovery.rs | 37 +-- src/protocol/filter.rs | 18 +- src/protocol/mod.rs | 27 +- src/protocol/negotiation.rs | 82 +++--- src/transport/ble/mod.rs | 26 +- src/transport/ethernet/mod.rs | 42 ++- src/transport/tcp/stream.rs | 5 +- src/tree/tests.rs | 45 ++- 38 files changed, 1647 insertions(+), 1012 deletions(-) diff --git a/src/bin/fipstop/ui/bloom.rs b/src/bin/fipstop/ui/bloom.rs index 71c657c..ba1aa98 100644 --- a/src/bin/fipstop/ui/bloom.rs +++ b/src/bin/fipstop/ui/bloom.rs @@ -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; } diff --git a/src/bloom/codec.rs b/src/bloom/codec.rs index e809f86..fefd5b9 100644 --- a/src/bloom/codec.rs +++ b/src/bloom/codec.rs @@ -59,10 +59,8 @@ pub fn rle_decode(data: &[u8], expected_words: usize) -> Result, 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] diff --git a/src/bloom/filter.rs b/src/bloom/filter.rs index ec6c5b5..0e744cc 100644 --- a/src/bloom/filter.rs +++ b/src/bloom/filter.rs @@ -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)); } diff --git a/src/bloom/state.rs b/src/bloom/state.rs index e30636c..546d4f8 100644 --- a/src/bloom/state.rs +++ b/src/bloom/state.rs @@ -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. diff --git a/src/bloom/tests.rs b/src/bloom/tests.rs index 1cba0af..4942e52 100644 --- a/src/bloom/tests.rs +++ b/src/bloom/tests.rs @@ -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] diff --git a/src/control/queries.rs b/src/control/queries.rs index 6325d32..f1be1d3 100644 --- a/src/control/queries.rs +++ b/src/control/queries.rs @@ -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 = node.peers().filter_map(|peer| { - let mmp = peer.mmp()?; - let addr = *peer.node_addr(); - let metrics = &mmp.metrics; + let peers: Vec = 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 = node diff --git a/src/mmp/algorithms.rs b/src/mmp/algorithms.rs index cc11d86..ccb63f6 100644 --- a/src/mmp/algorithms.rs +++ b/src/mmp/algorithms.rs @@ -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); } - } diff --git a/src/mmp/mod.rs b/src/mmp/mod.rs index 2d9beac..10a1971 100644 --- a/src/mmp/mod.rs +++ b/src/mmp/mod.rs @@ -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}; diff --git a/src/mmp/report.rs b/src/mmp/report.rs index 2ba0cda..a2b32b8 100644 --- a/src/mmp/report.rs +++ b/src/mmp/report.rs @@ -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, diff --git a/src/node/bloom.rs b/src/node/bloom.rs index 97d0d54..e40864c 100644 --- a/src/node/bloom.rs +++ b/src/node/bloom.rs @@ -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; } diff --git a/src/node/handlers/discovery.rs b/src/node/handlers/discovery.rs index 65e9687..a5a26f8 100644 --- a/src/node/handlers/discovery.rs +++ b/src/node/handlers/discovery.rs @@ -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; diff --git a/src/node/handlers/encrypted.rs b/src/node/handlers/encrypted.rs index 375ba65..67c7416 100644 --- a/src/node/handlers/encrypted.rs +++ b/src/node/handlers/encrypted.rs @@ -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. diff --git a/src/node/handlers/handshake.rs b/src/node/handlers/handshake.rs index 143dfe2..e85b035 100644 --- a/src/node/handlers/handshake.rs +++ b/src/node/handlers/handshake.rs @@ -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. diff --git a/src/node/handlers/rx_loop.rs b/src/node/handlers/rx_loop.rs index 325f75d..d152860 100644 --- a/src/node/handlers/rx_loop.rs +++ b/src/node/handlers/rx_loop.rs @@ -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::(32); if self.config.node.control.enabled { let config = self.config.node.control.clone(); diff --git a/src/node/handlers/session.rs b/src/node/handlers/session.rs index bf90b79..fd16e40 100644 --- a/src/node/handlers/session.rs +++ b/src/node/handlers/session.rs @@ -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 } diff --git a/src/node/handlers/timeout.rs b/src/node/handlers/timeout.rs index c8c05c4..b88ce43 100644 --- a/src/node/handlers/timeout.rs +++ b/src/node/handlers/timeout.rs @@ -22,7 +22,9 @@ impl Node { .unwrap_or(0); let timeout_ms = self.config.node.rate_limit.handshake_timeout_secs * 1000; - let stale: Vec = self.connections.iter() + let stale: Vec = 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)> = self.connections.iter() + let candidates: Vec<(LinkId, Vec)> = 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 = self.sessions.iter() + let timed_out: 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(); @@ -193,7 +197,9 @@ impl Node { // Second pass: collect resend candidates let my_addr = *self.node_addr(); - let candidates: Vec<(crate::NodeAddr, Vec)> = self.sessions.iter() + let candidates: Vec<(crate::NodeAddr, Vec)> = 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(); diff --git a/src/node/lifecycle.rs b/src/node/lifecycle.rs index 93ae508..20c7e6b 100644 --- a/src/node/lifecycle.rs +++ b/src/node/lifecycle.rs @@ -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, ) -> 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 = self.peers.iter() + let peer_addrs: Vec = 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 { - 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) { diff --git a/src/node/mod.rs b/src/node/mod.rs index 0dd7e7e..0237cc1 100644 --- a/src/node/mod.rs +++ b/src/node/mod.rs @@ -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>, /// TUN writer thread handle. tun_writer_handle: Option>, + /// 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, // === 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 { - self.addr_to_link.get(&(transport_id, addr.clone())).copied() + pub fn find_link_by_addr( + &self, + transport_id: TransportId, + addr: &TransportAddr, + ) -> Option { + 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, diff --git a/src/node/session_wire.rs b/src/node/session_wire.rs index f675872..4e20336 100644 --- a/src/node/session_wire.rs +++ b/src/node/session_wire.rs @@ -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 // ============================================================================ diff --git a/src/node/tests/disconnect.rs b/src/node/tests/disconnect.rs index 96a0cf4..709855c 100644 --- a/src/node/tests/disconnect.rs +++ b/src/node/tests/disconnect.rs @@ -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)" ); diff --git a/src/node/tests/handshake.rs b/src/node/tests/handshake.rs index 5f9d443..bea4a82 100644 --- a/src/node/tests/handshake.rs +++ b/src/node/tests/handshake.rs @@ -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. diff --git a/src/node/tests/mod.rs b/src/node/tests/mod.rs index 02f5836..e537157 100644 --- a/src/node/tests/mod.rs +++ b/src/node/tests/mod.rs @@ -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(); diff --git a/src/node/tests/session.rs b/src/node/tests/session.rs index 06f3b8e..ec218f5 100644 --- a/src/node/tests/session.rs +++ b/src/node/tests/session.rs @@ -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 = nodes - .iter() - .map(|tn| tn.node.session_count()) - .collect(); + let session_counts: Vec = 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 = nodes - .iter() - .map(|tn| tn.node.coord_cache().len()) - .collect(); + let coord_cache_sizes: Vec = + 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 { +fn build_ipv6_packet( + src: &crate::FipsAddress, + dst: &crate::FipsAddress, + payload: &[u8], +) -> Vec { 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> = 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> = 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> = 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> = 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> = 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> = 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> = 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> = 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> = 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" ); diff --git a/src/node/tests/spanning_tree.rs b/src/node/tests/spanning_tree.rs index f14951e..0f8e60a 100644 --- a/src/node/tests/spanning_tree.rs +++ b/src/node/tests/spanning_tree.rs @@ -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 ); } diff --git a/src/node/tests/unit.rs b/src/node/tests/unit.rs index bb33448..e683e08 100644 --- a/src/node/tests/unit.rs +++ b/src/node/tests/unit.rs @@ -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); diff --git a/src/node/wire.rs b/src/node/wire.rs index a55ab48..0fe533b 100644 --- a/src/node/wire.rs +++ b/src/node/wire.rs @@ -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 { /// /// 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 { +pub fn build_msg2( + sender_idx: SessionIndex, + receiver_idx: SessionIndex, + noise_msg2: &[u8], +) -> Vec { 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 { +pub fn build_msg3( + sender_idx: SessionIndex, + receiver_idx: SessionIndex, + noise_msg3: &[u8], +) -> Vec { 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); } diff --git a/src/noise/handshake.rs b/src/noise/handshake.rs index 25f90ad..4badadb 100644 --- a/src/noise/handshake.rs +++ b/src/noise/handshake.rs @@ -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); diff --git a/src/noise/tests.rs b/src/noise/tests.rs index caf12ce..b8026c1 100644 --- a/src/noise/tests.rs +++ b/src/noise/tests.rs @@ -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); diff --git a/src/peer/active.rs b/src/peer/active.rs index c39a434..5aa9d3f 100644 --- a/src/peer/active.rs +++ b/src/peer/active.rs @@ -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, 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 { - 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 { + 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). diff --git a/src/peer/connection.rs b/src/peer/connection.rs index b1410f0..d458c8c 100644 --- a/src/peer/connection.rs +++ b/src/peer/connection.rs @@ -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() + ); } } diff --git a/src/protocol/discovery.rs b/src/protocol/discovery.rs index d01aec0..26a9c5e 100644 --- a/src/protocol/discovery.rs +++ b/src/protocol/discovery.rs @@ -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(); diff --git a/src/protocol/filter.rs b/src/protocol/filter.rs index e4a25f1..30d86d2 100644 --- a/src/protocol/filter.rs +++ b/src/protocol/filter.rs @@ -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, diff --git a/src/protocol/mod.rs b/src/protocol/mod.rs index d0c0c24..60756ee 100644 --- a/src/protocol/mod.rs +++ b/src/protocol/mod.rs @@ -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; diff --git a/src/protocol/negotiation.rs b/src/protocol/negotiation.rs index 440d211..9a3031d 100644 --- a/src/protocol/negotiation.rs +++ b/src/protocol/negotiation.rs @@ -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() + ); } - } diff --git a/src/transport/ble/mod.rs b/src/transport/ble/mod.rs index 884cdad..95a2f05 100644 --- a/src/transport/ble/mod.rs +++ b/src/transport/ble/mod.rs @@ -56,7 +56,6 @@ pub type DefaultBleTransport = BleTransport; #[cfg(any(not(feature = "ble"), test))] pub type DefaultBleTransport = BleTransport; - // ============================================================================ // BLE Transport // ============================================================================ @@ -892,16 +891,13 @@ mod tests { fn make_transport( io: MockBleIo, - ) -> (BleTransport, tokio::sync::mpsc::Receiver) { + ) -> ( + BleTransport, + tokio::sync::mpsc::Receiver, + ) { 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 + ); } - } diff --git a/src/transport/ethernet/mod.rs b/src/transport/ethernet/mod.rs index 5ce927e..b5aebc1 100644 --- a/src/transport/ethernet/mod.rs +++ b/src/transport/ethernet/mod.rs @@ -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 \ diff --git a/src/transport/tcp/stream.rs b/src/transport/tcp/stream.rs index 27e801d..626b76e 100644 --- a/src/transport/tcp/stream.rs +++ b/src/transport/tcp/stream.rs @@ -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] diff --git a/src/tree/tests.rs b/src/tree/tests.rs index 0c3f136..a2c46b9 100644 --- a/src/tree/tests.rs +++ b/src/tree/tests.rs @@ -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 ===