diff --git a/src/config/node.rs b/src/config/node.rs index fc8fbb41..d9036a92 100644 --- a/src/config/node.rs +++ b/src/config/node.rs @@ -60,6 +60,16 @@ pub struct RateLimitConfig { /// Stale handshake cleanup timeout in seconds (`node.rate_limit.handshake_timeout_secs`). #[serde(default = "RateLimitConfig::default_handshake_timeout_secs")] pub handshake_timeout_secs: u64, + /// Initial handshake resend interval in ms (`node.rate_limit.handshake_resend_interval_ms`). + /// Handshake messages are resent with exponential backoff within the timeout window. + #[serde(default = "RateLimitConfig::default_handshake_resend_interval_ms")] + pub handshake_resend_interval_ms: u64, + /// Handshake resend backoff multiplier (`node.rate_limit.handshake_resend_backoff`). + #[serde(default = "RateLimitConfig::default_handshake_resend_backoff")] + pub handshake_resend_backoff: f64, + /// Max handshake resends per attempt (`node.rate_limit.handshake_max_resends`). + #[serde(default = "RateLimitConfig::default_handshake_max_resends")] + pub handshake_max_resends: u32, } impl Default for RateLimitConfig { @@ -68,6 +78,9 @@ impl Default for RateLimitConfig { handshake_burst: 100, handshake_rate: 10.0, handshake_timeout_secs: 30, + handshake_resend_interval_ms: 1000, + handshake_resend_backoff: 2.0, + handshake_max_resends: 5, } } } @@ -76,6 +89,9 @@ impl RateLimitConfig { fn default_handshake_burst() -> u32 { 100 } fn default_handshake_rate() -> f64 { 10.0 } fn default_handshake_timeout_secs() -> u64 { 30 } + fn default_handshake_resend_interval_ms() -> u64 { 1000 } + fn default_handshake_resend_backoff() -> f64 { 2.0 } + fn default_handshake_max_resends() -> u32 { 5 } } /// Retry/backoff configuration (`node.retry.*`). diff --git a/src/node/handlers/handshake.rs b/src/node/handlers/handshake.rs index 04bb8c64..6cff91b7 100644 --- a/src/node/handlers/handshake.rs +++ b/src/node/handlers/handshake.rs @@ -38,21 +38,38 @@ impl Node { // Check for existing connection from this address. // - // If we already have an *inbound* link from this address, drop the msg1 - // (duplicate or replay). But if we have an *outbound* link to this address - // (we initiated to them AND they initiated to us), this is a cross-connection. - // Allow it to proceed — promote_connection() will resolve via tie-breaker. + // If we already have an *inbound* link from this address, this is a + // duplicate msg1 (our msg2 was probably lost). Resend msg2 if available. + // If we have an *outbound* link to this address (we initiated to them + // AND they initiated to us), this is a cross-connection — allow it. let addr_key = (packet.transport_id, packet.remote_addr.clone()); if let Some(&existing_link_id) = self.addr_to_link.get(&addr_key) && let Some(link) = self.links.get(&existing_link_id) { if link.direction() == LinkDirection::Inbound { + // Duplicate msg1 — try to resend stored msg2 + let msg2_bytes = self.find_stored_msg2(existing_link_id); + if let Some(msg2) = msg2_bytes { + if let Some(transport) = self.transports.get(&packet.transport_id) { + match transport.send(&packet.remote_addr, &msg2).await { + Ok(_) => debug!( + remote_addr = %packet.remote_addr, + "Resent msg2 for duplicate msg1" + ), + Err(e) => debug!( + remote_addr = %packet.remote_addr, + error = %e, + "Failed to resend msg2" + ), + } + } + } else { + debug!( + remote_addr = %packet.remote_addr, + "Duplicate msg1 but no stored msg2 to resend" + ); + } self.msg1_rate_limiter.complete_handshake(); - debug!( - transport_id = %packet.transport_id, - remote_addr = %packet.remote_addr, - "Already have inbound connection from this address" - ); return; } // Outbound link to this address — cross-connection, allow msg1 @@ -126,8 +143,11 @@ impl Node { self.addr_to_link.insert(addr_key, link_id); self.connections.insert(link_id, conn); - // Build and send msg2 response + // Build and send msg2 response, storing for potential resend let wire_msg2 = build_msg2(our_index, header.sender_idx, &msg2_response); + if let Some(conn) = self.connections.get_mut(&link_id) { + conn.set_handshake_msg2(wire_msg2.clone()); + } if let Some(transport) = self.transports.get(&packet.transport_id) { match transport.send(&packet.remote_addr, &wire_msg2).await { @@ -164,6 +184,10 @@ impl Node { Ok(result) => { match result { PromotionResult::Promoted(node_addr) => { + // Store msg2 on peer for resend on duplicate msg1 + if let Some(peer) = self.peers.get_mut(&node_addr) { + peer.set_handshake_msg2(wire_msg2.clone()); + } info!( peer = %self.peer_display_name(&node_addr), link_id = %link_id, @@ -178,6 +202,10 @@ impl Node { self.bloom_state.mark_update_needed(node_addr); } PromotionResult::CrossConnectionWon { loser_link_id, node_addr } => { + // Store msg2 on peer for resend on duplicate msg1 + if let Some(peer) = self.peers.get_mut(&node_addr) { + peer.set_handshake_msg2(wire_msg2.clone()); + } // Clean up the losing connection's link self.remove_link(&loser_link_id); info!( @@ -222,6 +250,28 @@ impl Node { self.msg1_rate_limiter.complete_handshake(); } + /// Find stored msg2 bytes for a given link (pre- or post-promotion). + /// + /// Checks the PeerConnection (if still pending) and then the ActivePeer + /// (if already promoted). + fn find_stored_msg2(&self, link_id: LinkId) -> Option> { + // Check pending connection first + if let Some(conn) = self.connections.get(&link_id) + && let Some(msg2) = conn.handshake_msg2() + { + return Some(msg2.to_vec()); + } + // Check promoted peer + for peer in self.peers.values() { + if peer.link_id() == link_id + && let Some(msg2) = peer.handshake_msg2() + { + return Some(msg2.to_vec()); + } + } + None + } + /// Handle handshake message 2 (phase 0x2). /// /// This completes an outbound handshake we initiated. diff --git a/src/node/handlers/rx_loop.rs b/src/node/handlers/rx_loop.rs index 7ff22a05..fcee5f55 100644 --- a/src/node/handlers/rx_loop.rs +++ b/src/node/handlers/rx_loop.rs @@ -79,6 +79,7 @@ impl Node { .duration_since(std::time::UNIX_EPOCH) .map(|d| d.as_millis() as u64) .unwrap_or(0); + self.resend_pending_handshakes(now_ms).await; self.purge_idle_sessions(now_ms); self.process_pending_retries(now_ms).await; self.check_tree_state().await; diff --git a/src/node/handlers/timeout.rs b/src/node/handlers/timeout.rs index ea7ca5fa..4e2bc0e0 100644 --- a/src/node/handlers/timeout.rs +++ b/src/node/handlers/timeout.rs @@ -1,6 +1,8 @@ -//! Timeout management for stale handshake connections and idle sessions. +//! Timeout management for stale handshake connections, idle sessions, +//! and handshake message resend scheduling. use crate::node::Node; +use crate::peer::HandshakeState; use crate::transport::LinkId; use tracing::{debug, info}; @@ -79,6 +81,76 @@ impl Node { self.remove_link(&link_id); } + /// Resend handshake messages for pending connections. + /// + /// For outbound connections in SentMsg1 state, resends the stored msg1 + /// with exponential backoff. Called periodically from the RX event loop. + pub(in crate::node) async fn resend_pending_handshakes(&mut self, now_ms: u64) { + if self.connections.is_empty() { + return; + } + + let max_resends = self.config.node.rate_limit.handshake_max_resends; + let interval_ms = self.config.node.rate_limit.handshake_resend_interval_ms; + let backoff = self.config.node.rate_limit.handshake_resend_backoff; + + // Collect resend candidates: outbound, in SentMsg1, with stored msg1, + // under max resends, and past the scheduled time. + 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() + }) + .filter_map(|(link_id, conn)| { + conn.handshake_msg1().map(|msg1| (*link_id, msg1.to_vec())) + }) + .collect(); + + for (link_id, msg1_bytes) in candidates { + // Get transport and address info from the connection + let (transport_id, remote_addr) = match self.connections.get(&link_id) { + Some(conn) => match (conn.transport_id(), conn.source_addr()) { + (Some(tid), Some(addr)) => (tid, addr.clone()), + _ => continue, + }, + None => continue, + }; + + // Send the stored msg1 + let sent = if let Some(transport) = self.transports.get(&transport_id) { + match transport.send(&remote_addr, &msg1_bytes).await { + Ok(_) => true, + Err(e) => { + debug!( + link_id = %link_id, + error = %e, + "Handshake msg1 resend failed" + ); + false + } + } + } else { + false + }; + + 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); + debug!( + link_id = %link_id, + resend = count, + "Resent handshake msg1" + ); + } + } + } + /// Remove established sessions that have been idle too long. /// /// Only targets sessions in the Established state. Initiating/Responding diff --git a/src/node/lifecycle.rs b/src/node/lifecycle.rs index 6bc67a0f..05bbd809 100644 --- a/src/node/lifecycle.rs +++ b/src/node/lifecycle.rs @@ -180,6 +180,10 @@ impl Node { "Peer connection initiated" ); + // Store msg1 for resend and schedule first resend + let resend_interval = self.config.node.rate_limit.handshake_resend_interval_ms; + 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.connections.insert(link_id, connection); diff --git a/src/node/tests/handshake.rs b/src/node/tests/handshake.rs index c50573a8..9cec8e02 100644 --- a/src/node/tests/handshake.rs +++ b/src/node/tests/handshake.rs @@ -689,3 +689,157 @@ async fn test_failed_connection_cleanup() { assert_eq!(node.link_count(), 0, "Failed link should be removed"); assert_eq!(node.index_allocator.count(), 0, "Session index should be freed"); } + +/// Test that msg1 bytes are stored on connection for resend. +#[tokio::test] +async fn test_msg1_stored_for_resend() { + use crate::node::wire::build_msg1; + + let mut node = make_node(); + let transport_id = TransportId::new(1); + + let peer_identity = make_peer_identity(); + let remote_addr = TransportAddr::from_string("10.0.0.2:4000"); + + let now_ms = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .map(|d| d.as_millis() as u64) + .unwrap_or(0); + let link_id = node.allocate_link_id(); + let mut conn = PeerConnection::outbound(link_id, peer_identity, now_ms); + + let our_index = node.index_allocator.allocate().unwrap(); + let our_keypair = node.identity.keypair(); + let noise_msg1 = conn.start_handshake(our_keypair, now_ms).unwrap(); + conn.set_our_index(our_index); + conn.set_transport_id(transport_id); + conn.set_source_addr(remote_addr.clone()); + + // Build wire msg1 and store it (as initiate_peer_connection does) + let wire_msg1 = build_msg1(our_index, &noise_msg1); + let resend_interval = node.config.node.rate_limit.handshake_resend_interval_ms; + conn.set_handshake_msg1(wire_msg1.clone(), now_ms + resend_interval); + + // Verify stored msg1 matches what was built + assert_eq!(conn.handshake_msg1().unwrap(), &wire_msg1); + assert_eq!(conn.resend_count(), 0); + assert!(conn.next_resend_at_ms() > now_ms); +} + +/// Test that resend scheduling respects max_resends and backoff. +#[tokio::test] +async fn test_resend_scheduling() { + let mut node = make_node(); + let transport_id = TransportId::new(1); + + let peer_identity = make_peer_identity(); + let remote_addr = TransportAddr::from_string("10.0.0.2:4000"); + + let now_ms = 100_000u64; // Use a fixed time for predictable testing + let link_id = node.allocate_link_id(); + let mut conn = PeerConnection::outbound(link_id, peer_identity, now_ms); + + let our_index = node.index_allocator.allocate().unwrap(); + let our_keypair = node.identity.keypair(); + let noise_msg1 = conn.start_handshake(our_keypair, now_ms).unwrap(); + conn.set_our_index(our_index); + conn.set_transport_id(transport_id); + conn.set_source_addr(remote_addr.clone()); + + // Store msg1 with first resend at now + 1000ms + let wire_msg1 = crate::node::wire::build_msg1(our_index, &noise_msg1); + 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), + ); + 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.connections.insert(link_id, conn); + + // Before resend time: nothing should happen (no transport = can't send, + // but the filter should exclude it because now < next_resend_at) + node.resend_pending_handshakes(now_ms + 500).await; + let conn = node.connections.get(&link_id).unwrap(); + assert_eq!(conn.resend_count(), 0, "No resend before scheduled time"); + + // At resend time: would resend if transport existed. Without transport, + // the send fails silently and resend_count stays at 0. + // This tests the filtering logic — the connection IS a candidate. + node.resend_pending_handshakes(now_ms + 1000).await; + // 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"); +} + +/// Test that msg2 is stored on PeerConnection for responder resend. +#[test] +fn test_msg2_stored_on_connection() { + let mut conn = PeerConnection::inbound(LinkId::new(1), 1000); + + assert!(conn.handshake_msg2().is_none()); + + let msg2_bytes = vec![0x01, 0x02, 0x03, 0x04]; + conn.set_handshake_msg2(msg2_bytes.clone()); + + assert_eq!(conn.handshake_msg2().unwrap(), &msg2_bytes); +} + +/// Test that resend_count and next_resend_at_ms track correctly. +#[test] +fn test_resend_count_tracking() { + let peer_identity = make_peer_identity(); + let mut conn = PeerConnection::outbound(LinkId::new(1), peer_identity, 1000); + + assert_eq!(conn.resend_count(), 0); + assert_eq!(conn.next_resend_at_ms(), 0); + + // Simulate storing msg1 and scheduling first resend + conn.set_handshake_msg1(vec![0x01], 2000); + assert_eq!(conn.resend_count(), 0); + assert_eq!(conn.next_resend_at_ms(), 2000); + + // Record first resend + conn.record_resend(4000); // next at 4000 (2s backoff) + assert_eq!(conn.resend_count(), 1); + assert_eq!(conn.next_resend_at_ms(), 4000); + + // Record second resend + conn.record_resend(8000); // next at 8000 (4s backoff) + assert_eq!(conn.resend_count(), 2); + assert_eq!(conn.next_resend_at_ms(), 8000); +} + +/// Test that duplicate msg2 is silently dropped when pending_outbound is already cleared. +#[tokio::test] +async fn test_duplicate_msg2_dropped() { + use crate::node::wire::build_msg2; + use crate::transport::ReceivedPacket; + + let mut node = make_node(); + let transport_id = TransportId::new(1); + + // No pending_outbound entry — simulate post-promotion state + let receiver_idx = SessionIndex::new(42); + let sender_idx = SessionIndex::new(99); + + // Build a fake msg2 packet + let fake_noise_msg2 = vec![0u8; 33]; // Noise IK msg2 is 33 bytes + let wire_msg2 = build_msg2(sender_idx, receiver_idx, &fake_noise_msg2); + + let packet = ReceivedPacket { + transport_id, + remote_addr: TransportAddr::from_string("10.0.0.2:4000"), + data: wire_msg2, + timestamp_ms: 1000, + }; + + // Should silently drop — no pending_outbound for this index + node.handle_msg2(packet).await; + // No panic, no state change — that's the test + assert_eq!(node.connection_count(), 0); + assert_eq!(node.peer_count(), 0); +} diff --git a/src/peer/active.rs b/src/peer/active.rs index 43f20922..f73637ca 100644 --- a/src/peer/active.rs +++ b/src/peer/active.rs @@ -130,6 +130,11 @@ pub struct ActivePeer { // === MMP === /// Per-peer MMP state (None for legacy peers without Noise sessions). mmp: Option, + + // === Handshake Resend === + /// Wire-format msg2 for resend on duplicate msg1 (responder only). + /// Cleared after the handshake timeout window. + handshake_msg2: Option>, } impl ActivePeer { @@ -161,6 +166,7 @@ impl ActivePeer { authenticated_at, last_seen: authenticated_at, mmp: None, + handshake_msg2: None, } } @@ -220,6 +226,7 @@ impl ActivePeer { authenticated_at, last_seen: authenticated_at, mmp: Some(MmpPeerState::new(mmp_config, is_initiator)), + handshake_msg2: None, } } @@ -350,6 +357,23 @@ impl ActivePeer { self.current_addr = Some(addr); } + // === Handshake Resend === + + /// Store wire-format msg2 for resend on duplicate msg1. + pub fn set_handshake_msg2(&mut self, msg2: Vec) { + self.handshake_msg2 = Some(msg2); + } + + /// Get stored msg2 bytes for resend. + pub fn handshake_msg2(&self) -> Option<&[u8]> { + self.handshake_msg2.as_deref() + } + + /// Clear stored msg2 (no longer needed after handshake window). + pub fn clear_handshake_msg2(&mut self) { + self.handshake_msg2 = None; + } + // === Tree Accessors === /// Get the peer's tree coordinates, if known. diff --git a/src/peer/connection.rs b/src/peer/connection.rs index 16629732..f56d5c8b 100644 --- a/src/peer/connection.rs +++ b/src/peer/connection.rs @@ -116,6 +116,19 @@ pub struct PeerConnection { /// Current source address (updated on packet receipt). source_addr: Option, + + // === Handshake Resend === + /// Wire-format msg1 bytes for resend (initiator only, 90 bytes). + handshake_msg1: Option>, + + /// Wire-format msg2 bytes for resend (responder only). + handshake_msg2: Option>, + + /// Number of resends performed so far. + resend_count: u32, + + /// When the next resend should fire (Unix ms). 0 = no resend scheduled. + next_resend_at_ms: u64, } impl PeerConnection { @@ -143,6 +156,10 @@ impl PeerConnection { their_index: None, transport_id: None, source_addr: None, + handshake_msg1: None, + handshake_msg2: None, + resend_count: 0, + next_resend_at_ms: 0, } } @@ -166,6 +183,10 @@ impl PeerConnection { their_index: None, transport_id: None, source_addr: None, + handshake_msg1: None, + handshake_msg2: None, + resend_count: 0, + next_resend_at_ms: 0, } } @@ -193,6 +214,10 @@ impl PeerConnection { their_index: None, transport_id: Some(transport_id), source_addr: Some(source_addr), + handshake_msg1: None, + handshake_msg2: None, + resend_count: 0, + next_resend_at_ms: 0, } } @@ -315,6 +340,46 @@ impl PeerConnection { self.source_addr = Some(addr); } + // === Handshake Resend === + + /// Store the wire-format msg1 bytes for resend and schedule the first resend. + pub fn set_handshake_msg1(&mut self, msg1: Vec, first_resend_at_ms: u64) { + self.handshake_msg1 = Some(msg1); + self.resend_count = 0; + self.next_resend_at_ms = first_resend_at_ms; + } + + /// Store the wire-format msg2 bytes for resend on duplicate msg1. + pub fn set_handshake_msg2(&mut self, msg2: Vec) { + self.handshake_msg2 = Some(msg2); + } + + /// Get the stored msg1 bytes (if any). + pub fn handshake_msg1(&self) -> Option<&[u8]> { + self.handshake_msg1.as_deref() + } + + /// Get the stored msg2 bytes (if any). + pub fn handshake_msg2(&self) -> Option<&[u8]> { + self.handshake_msg2.as_deref() + } + + /// Number of resends performed. + pub fn resend_count(&self) -> u32 { + self.resend_count + } + + /// When the next resend is scheduled (Unix ms). + pub fn next_resend_at_ms(&self) -> u64 { + self.next_resend_at_ms + } + + /// Record a resend and schedule the next one. + pub fn record_resend(&mut self, next_resend_at_ms: u64) { + self.resend_count += 1; + self.next_resend_at_ms = next_resend_at_ms; + } + // === Noise Handshake Operations === /// Start the handshake as initiator and generate message 1.