From f920526ecee5335ed9a0c9b376607408e608133f Mon Sep 17 00:00:00 2001 From: Johnathan Corgan Date: Sun, 22 Feb 2026 20:47:53 +0000 Subject: [PATCH] Add epoch-based peer restart detection to Noise IK handshake Each node generates a random 8-byte startup epoch, encrypted inside both Noise IK handshake messages (msg1 and msg2). When a peer's msg1 arrives with a different epoch than the stored value, the node tears down the stale session and processes the msg1 as a new connection, enabling near-instant restart detection instead of the 30-second dead timeout. Wire format impact: - msg1: 82 -> 106 bytes (added 24-byte encrypted epoch after ss DH) - msg2: 33 -> 57 bytes (added 24-byte encrypted epoch after se DH) - Wire msg1: 90 -> 114 bytes, wire msg2: 45 -> 69 bytes --- src/node/handlers/handshake.rs | 133 ++++++++++++++++++++++++-------- src/node/handlers/session.rs | 2 + src/node/lifecycle.rs | 2 +- src/node/mod.rs | 14 ++++ src/node/tests/handshake.rs | 18 ++--- src/node/tests/mod.rs | 6 +- src/node/tests/session.rs | 8 ++ src/node/tests/spanning_tree.rs | 2 +- src/node/tests/unit.rs | 8 +- src/node/wire.rs | 38 ++++----- src/noise/handshake.rs | 74 ++++++++++++++++-- src/noise/mod.rs | 14 +++- src/noise/tests.rs | 39 +++++++++- src/peer/active.rs | 14 ++++ src/peer/connection.rs | 53 +++++++++++-- 15 files changed, 338 insertions(+), 87 deletions(-) diff --git a/src/node/handlers/handshake.rs b/src/node/handlers/handshake.rs index 6cff91b..9591ee6 100644 --- a/src/node/handlers/handshake.rs +++ b/src/node/handlers/handshake.rs @@ -38,47 +38,64 @@ impl Node { // Check for existing connection from this address. // - // 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 already have an *inbound* link from this address, this could be: + // 1. A duplicate msg1 (our msg2 was lost) — resend msg2 + // 2. A restarted peer (different epoch) — tear down and reprocess + // // 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. + // + // Epoch-based restart detection: if the sender already has an inbound + // link AND is an active peer in self.peers, fall through to decrypt + // the msg1 and check the epoch. Otherwise, treat as duplicate. let addr_key = (packet.transport_id, packet.remote_addr.clone()); + let mut possible_restart = false; 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" - ), - } - } + // 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); + + if is_active_peer { + // Possible restart — fall through to decrypt and check epoch + possible_restart = true; } else { - debug!( - remote_addr = %packet.remote_addr, - "Duplicate msg1 but no stored msg2 to resend" - ); + // Genuinely pending handshake — resend 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(); + return; } - self.msg1_rate_limiter.complete_handshake(); - return; - } - // Outbound link to this address — cross-connection, allow msg1 - debug!( - transport_id = %packet.transport_id, - remote_addr = %packet.remote_addr, - existing_link_id = %existing_link_id, - "Cross-connection detected: have outbound, received inbound msg1" + } else { + // Outbound link to this address — cross-connection, allow msg1 + debug!( + transport_id = %packet.transport_id, + remote_addr = %packet.remote_addr, + existing_link_id = %existing_link_id, + "Cross-connection detected: have outbound, received inbound msg1" ); + } } // === CRYPTO COST PAID HERE === @@ -92,7 +109,7 @@ impl Node { let our_keypair = self.identity.keypair(); let noise_msg1 = &packet.data[header.noise_msg1_offset..]; - let msg2_response = match conn.receive_handshake_init(our_keypair, noise_msg1, packet.timestamp_ms) { + let msg2_response = match conn.receive_handshake_init(our_keypair, self.startup_epoch, noise_msg1, packet.timestamp_ms) { Ok(m) => m, Err(e) => { self.msg1_rate_limiter.complete_handshake(); @@ -114,6 +131,55 @@ impl Node { } }; + let peer_node_addr = *peer_identity.node_addr(); + + // Epoch-based restart detection and duplicate msg1 handling. + // + // If we fell through from the addr_to_link check above with + // possible_restart=true, we now have the decrypted epoch from msg1. + // Compare it against the stored epoch for this peer. + if possible_restart + && let Some(existing_peer) = self.peers.get(&peer_node_addr) + { + let new_epoch = conn.remote_epoch(); + let existing_epoch = existing_peer.remote_epoch(); + + match (existing_epoch, new_epoch) { + (Some(existing), Some(new)) if existing != new => { + // Epoch mismatch — peer restarted. Tear down stale session. + info!( + peer = %self.peer_display_name(&peer_node_addr), + "Peer restart detected (epoch mismatch), removing stale session" + ); + self.remove_active_peer(&peer_node_addr); + // Fall through to process as new connection + } + _ => { + // Same epoch (or no epoch stored) — duplicate msg1 from + // same session. Resend stored msg2. + if let Some(msg2) = existing_peer.handshake_msg2().map(|m| m.to_vec()) + && let Some(transport) = self.transports.get(&packet.transport_id) + { + match transport.send(&packet.remote_addr, &msg2).await { + Ok(_) => debug!( + peer = %self.peer_display_name(&peer_node_addr), + "Resent msg2 for duplicate msg1 (same epoch)" + ), + Err(e) => debug!( + peer = %self.peer_display_name(&peer_node_addr), + error = %e, + "Failed to resend msg2" + ), + } + } + self.msg1_rate_limiter.complete_handshake(); + return; + } + } + } + // If possible_restart was true but peer is no longer in self.peers + // (removed by another path), fall through to process as new connection. + // Note: we don't early-return if peer is already in self.peers here. // promote_connection handles cross-connection resolution via tie-breaker. @@ -560,6 +626,7 @@ impl Node { } })?.clone(); let link_stats = connection.link_stats().clone(); + let remote_epoch = connection.remote_epoch(); let peer_node_addr = *verified_identity.node_addr(); let is_outbound = connection.is_outbound(); @@ -601,6 +668,7 @@ impl Node { link_stats, is_outbound, &self.config.node.mmp, + remote_epoch, ); new_peer.set_tree_announce_min_interval_ms(self.config.node.tree.announce_min_interval_ms); @@ -683,6 +751,7 @@ impl Node { link_stats, is_outbound, &self.config.node.mmp, + remote_epoch, ); new_peer.set_tree_announce_min_interval_ms(self.config.node.tree.announce_min_interval_ms); diff --git a/src/node/handlers/session.rs b/src/node/handlers/session.rs index 28aceec..a5803bc 100644 --- a/src/node/handlers/session.rs +++ b/src/node/handlers/session.rs @@ -337,6 +337,7 @@ impl Node { // Create responder handshake and process msg1 let our_keypair = self.identity.keypair(); let mut handshake = HandshakeState::new_responder(our_keypair); + handshake.set_local_epoch(self.startup_epoch); if let Err(e) = handshake.read_message_1(&setup.handshake_payload) { debug!(error = %e, "Failed to process Noise IK msg1 in SessionSetup"); @@ -724,6 +725,7 @@ impl Node { // Create Noise IK initiator handshake let our_keypair = self.identity.keypair(); let mut handshake = HandshakeState::new_initiator(our_keypair, dest_pubkey); + 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 msg1 generation failed: {}", e), diff --git a/src/node/lifecycle.rs b/src/node/lifecycle.rs index 05bbd80..0a6c9dc 100644 --- a/src/node/lifecycle.rs +++ b/src/node/lifecycle.rs @@ -147,7 +147,7 @@ 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, current_time_ms) { + let noise_msg1 = match connection.start_handshake(our_keypair, self.startup_epoch, current_time_ms) { Ok(msg) => msg, Err(e) => { warn!( diff --git a/src/node/mod.rs b/src/node/mod.rs index a9ac41a..438238c 100644 --- a/src/node/mod.rs +++ b/src/node/mod.rs @@ -33,6 +33,7 @@ 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_SP}; use crate::{Config, ConfigError, Identity, IdentityError, NodeAddr, PeerIdentity}; +use rand::RngCore; use std::collections::{HashMap, VecDeque}; use std::fmt; use std::thread::JoinHandle; @@ -198,6 +199,10 @@ pub struct Node { /// This node's cryptographic identity. identity: Identity, + /// Random epoch generated at startup for peer restart detection. + /// Exchanged inside Noise handshake messages so peers can detect restarts. + startup_epoch: [u8; 8], + // === Configuration === /// Loaded configuration. config: Config, @@ -342,6 +347,9 @@ impl Node { let node_addr = *identity.node_addr(); let is_leaf_only = config.is_leaf_only(); + let mut startup_epoch = [0u8; 8]; + rand::thread_rng().fill_bytes(&mut startup_epoch); + let mut bloom_state = if is_leaf_only { BloomState::leaf_only(node_addr) } else { @@ -379,6 +387,7 @@ impl Node { Ok(Self { identity, + startup_epoch, config, state: NodeState::Created, is_leaf_only, @@ -427,6 +436,10 @@ impl Node { /// Create a node with a specific identity. pub fn with_identity(identity: Identity, config: Config) -> Self { let node_addr = *identity.node_addr(); + + let mut startup_epoch = [0u8; 8]; + rand::thread_rng().fill_bytes(&mut startup_epoch); + let tun_state = if config.tun.enabled { TunState::Configured } else { @@ -460,6 +473,7 @@ impl Node { Self { identity, + startup_epoch, config, state: NodeState::Created, is_leaf_only: false, diff --git a/src/node/tests/handshake.rs b/src/node/tests/handshake.rs index 9cec8e0..84c1a1e 100644 --- a/src/node/tests/handshake.rs +++ b/src/node/tests/handshake.rs @@ -65,7 +65,7 @@ async fn test_two_node_handshake_udp() { // Start handshake (generates Noise IK msg1) let our_keypair_a = node_a.identity.keypair(); - let noise_msg1 = conn_a.start_handshake(our_keypair_a, 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()); @@ -303,7 +303,7 @@ async fn test_run_rx_loop_handshake() { 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, 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()); @@ -488,7 +488,7 @@ 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, 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()); @@ -509,7 +509,7 @@ async fn test_cross_connection_both_initiate() { 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, 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()); @@ -611,7 +611,7 @@ 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, 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()); @@ -665,7 +665,7 @@ 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, 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()); @@ -710,7 +710,7 @@ 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, 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()); @@ -741,7 +741,7 @@ 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, 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()); @@ -827,7 +827,7 @@ async fn test_duplicate_msg2_dropped() { 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 fake_noise_msg2 = vec![0u8; 57]; // Noise IK msg2 is 57 bytes (33 ephem + 24 encrypted epoch) let wire_msg2 = build_msg2(sender_idx, receiver_idx, &fake_noise_msg2); let packet = ReceivedPacket { diff --git a/src/node/tests/mod.rs b/src/node/tests/mod.rs index 7361230..d856b7f 100644 --- a/src/node/tests/mod.rs +++ b/src/node/tests/mod.rs @@ -50,13 +50,15 @@ pub(super) fn make_completed_connection( // Run initiator side of handshake let our_keypair = node.identity.keypair(); - let msg1 = conn.start_handshake(our_keypair, current_time_ms).unwrap(); + let msg1 = conn.start_handshake(our_keypair, node.startup_epoch, current_time_ms).unwrap(); // Run responder side to generate msg2 let mut resp_conn = PeerConnection::inbound(LinkId::new(999), current_time_ms); let peer_keypair = peer_identity_full.keypair(); + let mut resp_epoch = [0u8; 8]; + rand::RngCore::fill_bytes(&mut rand::thread_rng(), &mut resp_epoch); let msg2 = resp_conn - .receive_handshake_init(peer_keypair, &msg1, current_time_ms) + .receive_handshake_init(peer_keypair, resp_epoch, &msg1, current_time_ms) .unwrap(); // Complete initiator handshake diff --git a/src/node/tests/session.rs b/src/node/tests/session.rs index 012d115..0bfb52b 100644 --- a/src/node/tests/session.rs +++ b/src/node/tests/session.rs @@ -1160,6 +1160,14 @@ fn make_noise_session( ); let mut responder = HandshakeState::new_responder(remote_identity.keypair()); + // Set epochs for both sides (required for handshake message encryption) + let mut init_epoch = [0u8; 8]; + rand::RngCore::fill_bytes(&mut rand::thread_rng(), &mut init_epoch); + initiator.set_local_epoch(init_epoch); + let mut resp_epoch = [0u8; 8]; + rand::RngCore::fill_bytes(&mut rand::thread_rng(), &mut resp_epoch); + responder.set_local_epoch(resp_epoch); + let msg1 = initiator.write_message_1().unwrap(); responder.read_message_1(&msg1).unwrap(); let msg2 = responder.write_message_2().unwrap(); diff --git a/src/node/tests/spanning_tree.rs b/src/node/tests/spanning_tree.rs index d86f642..78df33a 100644 --- a/src/node/tests/spanning_tree.rs +++ b/src/node/tests/spanning_tree.rs @@ -64,7 +64,7 @@ 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, 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()); diff --git a/src/node/tests/unit.rs b/src/node/tests/unit.rs index 296f80d..f44f833 100644 --- a/src/node/tests/unit.rs +++ b/src/node/tests/unit.rs @@ -450,7 +450,7 @@ fn test_promote_cleans_up_pending_outbound_to_same_peer() { PeerConnection::outbound(pending_link_id, peer_b_identity, pending_time_ms); let our_keypair = node.identity.keypair(); - let _msg1 = pending_conn.start_handshake(our_keypair, pending_time_ms).unwrap(); + let _msg1 = pending_conn.start_handshake(our_keypair, node.startup_epoch, pending_time_ms).unwrap(); let pending_index = node.index_allocator.allocate().unwrap(); pending_conn.set_our_index(pending_index); @@ -491,14 +491,16 @@ fn test_promote_cleans_up_pending_outbound_to_same_peer() { let our_keypair = node.identity.keypair(); let msg1 = completing_conn - .start_handshake(our_keypair, completing_time_ms) + .start_handshake(our_keypair, node.startup_epoch, completing_time_ms) .unwrap(); // B responds let mut resp_conn = PeerConnection::inbound(LinkId::new(999), completing_time_ms); let peer_keypair = peer_b_full.keypair(); + let mut resp_epoch = [0u8; 8]; + rand::RngCore::fill_bytes(&mut rand::thread_rng(), &mut resp_epoch); let msg2 = resp_conn - .receive_handshake_init(peer_keypair, &msg1, completing_time_ms) + .receive_handshake_init(peer_keypair, resp_epoch, &msg1, completing_time_ms) .unwrap(); completing_conn diff --git a/src/node/wire.rs b/src/node/wire.rs index 0d7f98e..103a5a0 100644 --- a/src/node/wire.rs +++ b/src/node/wire.rs @@ -11,11 +11,11 @@ //! //! ## Packet Types //! -//! | Phase | Type | Size | Description | -//! |-------|-----------------|-----------|--------------------------------| -//! | 0x0 | Encrypted frame | 32+ bytes | Post-handshake encrypted data | -//! | 0x1 | Noise IK msg1 | 90 bytes | Handshake initiation | -//! | 0x2 | Noise IK msg2 | 45 bytes | Handshake response | +//! | Phase | Type | Size | Description | +//! |-------|-----------------|------------|--------------------------------| +//! | 0x0 | Encrypted frame | 32+ bytes | Post-handshake encrypted data | +//! | 0x1 | Noise IK msg1 | 114 bytes | Handshake initiation | +//! | 0x2 | Noise IK msg2 | 69 bytes | Handshake response | use crate::utils::index::SessionIndex; use crate::noise::{HANDSHAKE_MSG1_SIZE, HANDSHAKE_MSG2_SIZE, TAG_SIZE}; @@ -43,10 +43,10 @@ pub const COMMON_PREFIX_SIZE: usize = 4; pub const ESTABLISHED_HEADER_SIZE: usize = 16; /// Size of Noise IK message 1 wire packet: prefix + sender_idx + noise_msg1. -pub const MSG1_WIRE_SIZE: usize = COMMON_PREFIX_SIZE + 4 + HANDSHAKE_MSG1_SIZE; // 90 bytes +pub const MSG1_WIRE_SIZE: usize = COMMON_PREFIX_SIZE + 4 + HANDSHAKE_MSG1_SIZE; // 114 bytes /// Size of Noise IK message 2 wire packet: prefix + sender_idx + receiver_idx + noise_msg2. -pub const MSG2_WIRE_SIZE: usize = COMMON_PREFIX_SIZE + 4 + 4 + HANDSHAKE_MSG2_SIZE; // 45 bytes +pub const MSG2_WIRE_SIZE: usize = COMMON_PREFIX_SIZE + 4 + 4 + HANDSHAKE_MSG2_SIZE; // 69 bytes /// Minimum size for encrypted frame: header + tag (no plaintext). pub const ENCRYPTED_MIN_SIZE: usize = ESTABLISHED_HEADER_SIZE + TAG_SIZE; // 32 bytes @@ -198,9 +198,9 @@ impl EncryptedHeader { /// Parsed Noise IK message 1 header (phase 0x1). /// -/// Wire format (90 bytes): +/// Wire format (114 bytes): /// ```text -/// [0x01][0x00][payload_len:2 LE][sender_idx:4 LE][noise_msg1:82] +/// [0x01][0x00][payload_len:2 LE][sender_idx:4 LE][noise_msg1:106] /// ``` #[derive(Clone, Debug)] pub struct Msg1Header { @@ -252,9 +252,9 @@ impl Msg1Header { /// Parsed Noise IK message 2 header (phase 0x2). /// -/// Wire format (45 bytes): +/// Wire format (69 bytes): /// ```text -/// [0x02][0x00][payload_len:2 LE][sender_idx:4 LE][receiver_idx:4 LE][noise_msg2:33] +/// [0x02][0x00][payload_len:2 LE][sender_idx:4 LE][receiver_idx:4 LE][noise_msg2:57] /// ``` #[derive(Clone, Debug)] pub struct Msg2Header { @@ -310,7 +310,7 @@ impl Msg2Header { /// Build a wire-format msg1 packet. /// -/// Format: `[0x01][0x00][payload_len:2 LE][sender_idx:4 LE][noise_msg1:82]` +/// Format: `[0x01][0x00][payload_len:2 LE][sender_idx:4 LE][noise_msg1:106]` pub fn build_msg1(sender_idx: SessionIndex, noise_msg1: &[u8]) -> Vec { debug_assert_eq!(noise_msg1.len(), HANDSHAKE_MSG1_SIZE); @@ -327,7 +327,7 @@ pub fn build_msg1(sender_idx: SessionIndex, noise_msg1: &[u8]) -> Vec { /// Build a wire-format msg2 packet. /// -/// Format: `[0x02][0x00][payload_len:2 LE][sender_idx:4 LE][receiver_idx:4 LE][noise_msg2:33]` +/// Format: `[0x02][0x00][payload_len:2 LE][sender_idx:4 LE][receiver_idx:4 LE][noise_msg2:57]` pub fn build_msg2(sender_idx: SessionIndex, receiver_idx: SessionIndex, noise_msg2: &[u8]) -> Vec { debug_assert_eq!(noise_msg2.len(), HANDSHAKE_MSG2_SIZE); @@ -542,8 +542,8 @@ mod tests { #[test] fn test_wire_sizes() { - assert_eq!(MSG1_WIRE_SIZE, 90); // 4 + 4 + 82 - assert_eq!(MSG2_WIRE_SIZE, 45); // 4 + 4 + 4 + 33 + assert_eq!(MSG1_WIRE_SIZE, 114); // 4 + 4 + 106 + assert_eq!(MSG2_WIRE_SIZE, 69); // 4 + 4 + 4 + 57 assert_eq!(ENCRYPTED_MIN_SIZE, 32); // 16 + 16 assert_eq!(COMMON_PREFIX_SIZE, 4); assert_eq!(ESTABLISHED_HEADER_SIZE, 16); @@ -610,8 +610,8 @@ mod tests { fn test_payload_len_in_msg1() { let packet = build_msg1(SessionIndex::new(1), &[0u8; HANDSHAKE_MSG1_SIZE]); let prefix = CommonPrefix::parse(&packet).unwrap(); - // payload_len = sender_idx(4) + noise_msg1(82) = 86 - assert_eq!(prefix.payload_len, 86); + // payload_len = sender_idx(4) + noise_msg1(106) = 110 + assert_eq!(prefix.payload_len, 110); } #[test] @@ -622,7 +622,7 @@ mod tests { &[0u8; HANDSHAKE_MSG2_SIZE], ); let prefix = CommonPrefix::parse(&packet).unwrap(); - // payload_len = sender_idx(4) + receiver_idx(4) + noise_msg2(33) = 41 - assert_eq!(prefix.payload_len, 41); + // payload_len = sender_idx(4) + receiver_idx(4) + noise_msg2(57) = 65 + assert_eq!(prefix.payload_len, 65); } } diff --git a/src/noise/handshake.rs b/src/noise/handshake.rs index 958f90f..c82de13 100644 --- a/src/noise/handshake.rs +++ b/src/noise/handshake.rs @@ -1,6 +1,7 @@ use super::{ CipherState, HandshakeProgress, HandshakeRole, NoiseError, NoiseSession, - HANDSHAKE_MSG1_SIZE, HANDSHAKE_MSG2_SIZE, PROTOCOL_NAME, PUBKEY_SIZE, + EPOCH_ENCRYPTED_SIZE, EPOCH_SIZE, HANDSHAKE_MSG1_SIZE, HANDSHAKE_MSG2_SIZE, + PROTOCOL_NAME, PUBKEY_SIZE, }; use hkdf::Hkdf; use rand::RngCore; @@ -120,6 +121,10 @@ pub struct HandshakeState { remote_ephemeral: Option, /// Secp256k1 context. secp: Secp256k1, + /// Our startup epoch for restart detection. + local_epoch: Option<[u8; 8]>, + /// Remote peer's startup epoch (learned during handshake). + remote_epoch: Option<[u8; 8]>, } impl HandshakeState { @@ -153,6 +158,8 @@ impl HandshakeState { remote_static: Some(remote_static), remote_ephemeral: None, secp, + local_epoch: None, + remote_epoch: None, }; // Mix in pre-message: <- s (responder's static is known) @@ -179,6 +186,8 @@ impl HandshakeState { remote_static: None, // Will learn from message 1 remote_ephemeral: None, secp, + local_epoch: None, + remote_epoch: None, }; // Mix in pre-message: <- s (our static, since we're responder) @@ -209,6 +218,16 @@ impl HandshakeState { self.remote_static.as_ref() } + /// Set the local startup epoch for restart detection. + pub fn set_local_epoch(&mut self, epoch: [u8; 8]) { + self.local_epoch = Some(epoch); + } + + /// Get the remote peer's startup epoch (available after processing their message). + pub fn remote_epoch(&self) -> Option<[u8; 8]> { + self.remote_epoch + } + /// Generate ephemeral keypair. fn generate_ephemeral(&mut self) { let mut rng = rand::thread_rng(); @@ -245,8 +264,9 @@ impl HandshakeState { /// Message 1 contains: /// - e: ephemeral public key (33 bytes) /// - encrypted s: our static public key encrypted (33 + 16 = 49 bytes) + /// - encrypted epoch: startup epoch for restart detection (8 + 16 = 24 bytes) /// - /// Total: 82 bytes + /// Total: 106 bytes pub fn write_message_1(&mut self) -> Result, NoiseError> { if self.role != HandshakeRole::Initiator { return Err(NoiseError::WrongState { @@ -262,6 +282,7 @@ impl HandshakeState { } let remote_static = self.remote_static.expect("initiator must have remote static"); + let epoch = self.local_epoch.expect("local epoch must be set before write_message_1"); // Generate ephemeral keypair self.generate_ephemeral(); @@ -287,6 +308,11 @@ impl HandshakeState { let ss = self.ecdh(&self.static_keypair.secret_key(), &remote_static); self.symmetric.mix_key(&ss); + // -> epoch: encrypt startup epoch for restart detection + let encrypted_epoch = self.symmetric.encrypt_and_hash(&epoch)?; + debug_assert_eq!(encrypted_epoch.len(), EPOCH_ENCRYPTED_SIZE); + message.extend_from_slice(&encrypted_epoch); + self.progress = HandshakeProgress::Message1Done; Ok(message) @@ -294,7 +320,7 @@ impl HandshakeState { /// Read message 1 (responder only). /// - /// Processes the initiator's first message and learns their identity. + /// Processes the initiator's first message and learns their identity and epoch. pub fn read_message_1(&mut self, message: &[u8]) -> Result<(), NoiseError> { if self.role != HandshakeRole::Responder { return Err(NoiseError::WrongState { @@ -327,7 +353,8 @@ impl HandshakeState { self.symmetric.mix_key(&es); // -> s: decrypt initiator's static - let encrypted_static = &message[PUBKEY_SIZE..]; + let encrypted_static_end = PUBKEY_SIZE + PUBKEY_SIZE + super::TAG_SIZE; + let encrypted_static = &message[PUBKEY_SIZE..encrypted_static_end]; let decrypted_static = self.symmetric.decrypt_and_hash(encrypted_static)?; let rs = PublicKey::from_slice(&decrypted_static).map_err(|_| NoiseError::InvalidPublicKey)?; @@ -337,6 +364,15 @@ impl HandshakeState { let ss = self.ecdh(&self.static_keypair.secret_key(), &rs); self.symmetric.mix_key(&ss); + // -> epoch: decrypt initiator's startup epoch + let encrypted_epoch = &message[encrypted_static_end..]; + debug_assert_eq!(encrypted_epoch.len(), EPOCH_ENCRYPTED_SIZE); + let decrypted_epoch = self.symmetric.decrypt_and_hash(encrypted_epoch)?; + debug_assert_eq!(decrypted_epoch.len(), EPOCH_SIZE); + let mut epoch = [0u8; EPOCH_SIZE]; + epoch.copy_from_slice(&decrypted_epoch); + self.remote_epoch = Some(epoch); + self.progress = HandshakeProgress::Message1Done; Ok(()) @@ -346,8 +382,9 @@ impl HandshakeState { /// /// Message 2 contains: /// - e: ephemeral public key (33 bytes) + /// - encrypted epoch: startup epoch for restart detection (8 + 16 = 24 bytes) /// - /// Total: 33 bytes + /// Total: 57 bytes pub fn write_message_2(&mut self) -> Result, NoiseError> { if self.role != HandshakeRole::Responder { return Err(NoiseError::WrongState { @@ -363,13 +400,17 @@ 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"); // Generate ephemeral keypair self.generate_ephemeral(); let ephemeral = self.ephemeral_keypair.as_ref().unwrap(); let e_pub = ephemeral.public_key().serialize(); + let mut message = Vec::with_capacity(HANDSHAKE_MSG2_SIZE); + // <- e: send ephemeral, mix into hash + message.extend_from_slice(&e_pub); self.symmetric.mix_hash(&e_pub); // <- ee: DH(e, re), mix into key @@ -380,9 +421,14 @@ impl HandshakeState { let se = self.ecdh(&self.static_keypair.secret_key(), &re); self.symmetric.mix_key(&se); + // <- epoch: encrypt startup epoch for restart detection + let encrypted_epoch = self.symmetric.encrypt_and_hash(&epoch)?; + debug_assert_eq!(encrypted_epoch.len(), EPOCH_ENCRYPTED_SIZE); + message.extend_from_slice(&encrypted_epoch); + self.progress = HandshakeProgress::Complete; - Ok(e_pub.to_vec()) + Ok(message) } /// Read message 2 (initiator only). @@ -409,9 +455,10 @@ impl HandshakeState { } // <- e: parse remote ephemeral, mix into hash - let re = PublicKey::from_slice(message).map_err(|_| NoiseError::InvalidPublicKey)?; + let e_pub = &message[..PUBKEY_SIZE]; + let re = PublicKey::from_slice(e_pub).map_err(|_| NoiseError::InvalidPublicKey)?; self.remote_ephemeral = Some(re); - self.symmetric.mix_hash(message); + self.symmetric.mix_hash(e_pub); // <- ee: DH(e, re), mix into key let ephemeral = self.ephemeral_keypair.as_ref().unwrap(); @@ -424,6 +471,15 @@ impl HandshakeState { let se = self.ecdh(&ephemeral.secret_key(), &rs); self.symmetric.mix_key(&se); + // <- epoch: decrypt responder's startup epoch + let encrypted_epoch = &message[PUBKEY_SIZE..]; + debug_assert_eq!(encrypted_epoch.len(), EPOCH_ENCRYPTED_SIZE); + let decrypted_epoch = self.symmetric.decrypt_and_hash(encrypted_epoch)?; + debug_assert_eq!(decrypted_epoch.len(), EPOCH_SIZE); + let mut epoch = [0u8; EPOCH_SIZE]; + epoch.copy_from_slice(&decrypted_epoch); + self.remote_epoch = Some(epoch); + self.progress = HandshakeProgress::Complete; Ok(()) @@ -473,6 +529,8 @@ impl fmt::Debug for HandshakeState { .field("has_ephemeral", &self.ephemeral_keypair.is_some()) .field("has_remote_static", &self.remote_static.is_some()) .field("has_remote_ephemeral", &self.remote_ephemeral.is_some()) + .field("has_local_epoch", &self.local_epoch.is_some()) + .field("has_remote_epoch", &self.remote_epoch.is_some()) .finish() } } diff --git a/src/noise/mod.rs b/src/noise/mod.rs index ac6d548..d3788c4 100644 --- a/src/noise/mod.rs +++ b/src/noise/mod.rs @@ -59,11 +59,17 @@ pub const TAG_SIZE: usize = 16; /// Size of a public key (compressed secp256k1). pub const PUBKEY_SIZE: usize = 33; -/// Size of handshake message 1: ephemeral (33) + encrypted static (33 + 16 tag). -pub const HANDSHAKE_MSG1_SIZE: usize = PUBKEY_SIZE + PUBKEY_SIZE + TAG_SIZE; +/// Size of the startup epoch (random bytes for restart detection). +pub const EPOCH_SIZE: usize = 8; -/// Size of handshake message 2: ephemeral only. -pub const HANDSHAKE_MSG2_SIZE: usize = PUBKEY_SIZE; +/// Size of encrypted epoch (epoch + AEAD tag). +pub const EPOCH_ENCRYPTED_SIZE: usize = EPOCH_SIZE + TAG_SIZE; + +/// Size of handshake message 1: ephemeral (33) + encrypted static (33 + 16 tag) + encrypted epoch (8 + 16 tag). +pub const HANDSHAKE_MSG1_SIZE: usize = PUBKEY_SIZE + PUBKEY_SIZE + TAG_SIZE + EPOCH_ENCRYPTED_SIZE; + +/// Size of handshake message 2: ephemeral (33) + encrypted epoch (8 + 16 tag). +pub const HANDSHAKE_MSG2_SIZE: usize = PUBKEY_SIZE + EPOCH_ENCRYPTED_SIZE; /// Replay window size in packets (matching WireGuard). pub const REPLAY_WINDOW_SIZE: usize = 2048; diff --git a/src/noise/tests.rs b/src/noise/tests.rs index 8b9749f..0eb7066 100644 --- a/src/noise/tests.rs +++ b/src/noise/tests.rs @@ -1,4 +1,5 @@ use super::*; +use rand::RngCore; use secp256k1::Parity; fn generate_keypair() -> secp256k1::Keypair { @@ -8,17 +9,27 @@ fn generate_keypair() -> secp256k1::Keypair { secp256k1::Keypair::from_secret_key(&secp, &secret_key) } +fn generate_epoch() -> [u8; 8] { + let mut epoch = [0u8; 8]; + rand::thread_rng().fill_bytes(&mut epoch); + epoch +} + #[test] fn test_full_handshake() { let initiator_keypair = generate_keypair(); let responder_keypair = generate_keypair(); + let initiator_epoch = generate_epoch(); + let responder_epoch = generate_epoch(); let responder_pub = responder_keypair.public_key(); // Initiator knows responder's static key // Responder does NOT know initiator's static key (IK pattern) let mut initiator = HandshakeState::new_initiator(initiator_keypair, responder_pub); + initiator.set_local_epoch(initiator_epoch); let mut responder = HandshakeState::new_responder(responder_keypair); + responder.set_local_epoch(responder_epoch); assert_eq!(initiator.role(), HandshakeRole::Initiator); assert_eq!(responder.role(), HandshakeRole::Responder); @@ -39,6 +50,9 @@ fn test_full_handshake() { &initiator_keypair.public_key() ); + // Responder learned initiator's epoch + assert_eq!(responder.remote_epoch(), Some(initiator_epoch)); + // Message 2: Responder -> Initiator let msg2 = responder.write_message_2().unwrap(); assert_eq!(msg2.len(), HANDSHAKE_MSG2_SIZE); @@ -49,6 +63,9 @@ fn test_full_handshake() { assert!(initiator.is_complete()); assert!(responder.is_complete()); + // Initiator learned responder's epoch + assert_eq!(initiator.remote_epoch(), Some(responder_epoch)); + // Handshake hashes should match assert_eq!(initiator.handshake_hash(), responder.handshake_hash()); @@ -77,7 +94,9 @@ fn test_multiple_messages() { let mut initiator = HandshakeState::new_initiator(initiator_keypair, responder_keypair.public_key()); + initiator.set_local_epoch(generate_epoch()); let mut responder = HandshakeState::new_responder(responder_keypair); + responder.set_local_epoch(generate_epoch()); let msg1 = initiator.write_message_1().unwrap(); responder.read_message_1(&msg1).unwrap(); @@ -105,6 +124,7 @@ fn test_wrong_role_errors() { let keypair2 = generate_keypair(); let mut initiator = HandshakeState::new_initiator(keypair1, keypair2.public_key()); + initiator.set_local_epoch(generate_epoch()); // Initiator can't read message 1 assert!(initiator @@ -119,6 +139,7 @@ fn test_wrong_role_errors() { fn test_invalid_pubkey_in_msg1() { let keypair = generate_keypair(); let mut responder = HandshakeState::new_responder(keypair); + responder.set_local_epoch(generate_epoch()); // Invalid pubkey bytes (first 33 bytes are zero) let invalid_msg = [0u8; HANDSHAKE_MSG1_SIZE]; @@ -133,7 +154,9 @@ fn test_decryption_failure_wrong_key() { // Session between 1 and 2 let mut init1 = HandshakeState::new_initiator(keypair1, keypair2.public_key()); + init1.set_local_epoch(generate_epoch()); let mut resp1 = HandshakeState::new_responder(keypair2); + resp1.set_local_epoch(generate_epoch()); let msg1 = init1.write_message_1().unwrap(); resp1.read_message_1(&msg1).unwrap(); @@ -144,7 +167,9 @@ fn test_decryption_failure_wrong_key() { // Session between 1 and 3 let mut init2 = HandshakeState::new_initiator(keypair1, keypair3.public_key()); + init2.set_local_epoch(generate_epoch()); let mut resp2 = HandshakeState::new_responder(keypair3); + resp2.set_local_epoch(generate_epoch()); let msg1 = init2.write_message_1().unwrap(); resp2.read_message_1(&msg1).unwrap(); @@ -178,7 +203,9 @@ fn test_session_remote_static() { let keypair2 = generate_keypair(); let mut init = HandshakeState::new_initiator(keypair1, keypair2.public_key()); + init.set_local_epoch(generate_epoch()); let mut resp = HandshakeState::new_responder(keypair2); + resp.set_local_epoch(generate_epoch()); let msg1 = init.write_message_1().unwrap(); resp.read_message_1(&msg1).unwrap(); @@ -196,8 +223,10 @@ fn test_session_remote_static() { #[test] fn test_message_sizes() { // Verify our size constants are correct - assert_eq!(HANDSHAKE_MSG1_SIZE, 33 + 33 + 16); // e + encrypted_s - assert_eq!(HANDSHAKE_MSG2_SIZE, 33); // e only + assert_eq!(EPOCH_SIZE, 8); + assert_eq!(EPOCH_ENCRYPTED_SIZE, 8 + 16); // epoch + AEAD tag + assert_eq!(HANDSHAKE_MSG1_SIZE, 33 + 33 + 16 + 24); // e + encrypted_s + encrypted_epoch + assert_eq!(HANDSHAKE_MSG2_SIZE, 33 + 24); // e + encrypted_epoch } #[test] @@ -207,12 +236,14 @@ fn test_responder_identity_discovery() { let responder_keypair = generate_keypair(); let mut responder = HandshakeState::new_responder(responder_keypair); + responder.set_local_epoch(generate_epoch()); // Before message 1: responder has no idea who's connecting assert!(responder.remote_static().is_none()); let mut initiator = HandshakeState::new_initiator(initiator_keypair, responder_keypair.public_key()); + initiator.set_local_epoch(generate_epoch()); let msg1 = initiator.write_message_1().unwrap(); // After processing message 1: responder knows initiator's identity @@ -330,7 +361,9 @@ fn test_session_replay_protection() { let keypair2 = generate_keypair(); let mut init = HandshakeState::new_initiator(keypair1, keypair2.public_key()); + init.set_local_epoch(generate_epoch()); let mut resp = HandshakeState::new_responder(keypair2); + resp.set_local_epoch(generate_epoch()); let msg1 = init.write_message_1().unwrap(); resp.read_message_1(&msg1).unwrap(); @@ -394,7 +427,9 @@ fn test_handshake_with_odd_parity_responder() { // Handshake using assumed-even key (as production code does) let mut initiator = HandshakeState::new_initiator(kp_a, assumed_even_b); + initiator.set_local_epoch(generate_epoch()); let mut responder = HandshakeState::new_responder(kp_b); + responder.set_local_epoch(generate_epoch()); let msg1 = initiator.write_message_1().unwrap(); responder.read_message_1(&msg1).unwrap(); diff --git a/src/peer/active.rs b/src/peer/active.rs index 6bf73d4..fde635b 100644 --- a/src/peer/active.rs +++ b/src/peer/active.rs @@ -127,6 +127,10 @@ pub struct ActivePeer { /// When this peer was last seen (any activity, Unix milliseconds). last_seen: u64, + // === Epoch (Restart Detection) === + /// Remote peer's startup epoch (from handshake). Used to detect restarts. + remote_epoch: Option<[u8; 8]>, + // === MMP === /// Per-peer MMP state (None for legacy peers without Noise sessions). mmp: Option, @@ -169,6 +173,7 @@ impl ActivePeer { link_stats: LinkStats::new(), authenticated_at, last_seen: authenticated_at, + remote_epoch: None, mmp: None, last_heartbeat_sent: None, handshake_msg2: None, @@ -207,6 +212,7 @@ impl ActivePeer { link_stats: LinkStats, is_initiator: bool, mmp_config: &MmpConfig, + remote_epoch: Option<[u8; 8]>, ) -> Self { Self { identity, @@ -230,6 +236,7 @@ impl ActivePeer { link_stats, authenticated_at, last_seen: authenticated_at, + remote_epoch, mmp: Some(MmpPeerState::new(mmp_config, is_initiator)), last_heartbeat_sent: None, handshake_msg2: None, @@ -380,6 +387,13 @@ impl ActivePeer { self.handshake_msg2 = None; } + // === Epoch Accessors === + + /// Get the remote peer's startup epoch (from handshake). + pub fn remote_epoch(&self) -> Option<[u8; 8]> { + self.remote_epoch + } + // === Tree Accessors === /// Get the peer's tree coordinates, if known. diff --git a/src/peer/connection.rs b/src/peer/connection.rs index f56d5c8..6f546e4 100644 --- a/src/peer/connection.rs +++ b/src/peer/connection.rs @@ -117,8 +117,12 @@ pub struct PeerConnection { /// Current source address (updated on packet receipt). source_addr: Option, + // === Epoch (Restart Detection) === + /// Remote peer's startup epoch (learned from handshake). + remote_epoch: Option<[u8; 8]>, + // === Handshake Resend === - /// Wire-format msg1 bytes for resend (initiator only, 90 bytes). + /// Wire-format msg1 bytes for resend (initiator only). handshake_msg1: Option>, /// Wire-format msg2 bytes for resend (responder only). @@ -156,6 +160,7 @@ impl PeerConnection { their_index: None, transport_id: None, source_addr: None, + remote_epoch: None, handshake_msg1: None, handshake_msg2: None, resend_count: 0, @@ -183,6 +188,7 @@ impl PeerConnection { their_index: None, transport_id: None, source_addr: None, + remote_epoch: None, handshake_msg1: None, handshake_msg2: None, resend_count: 0, @@ -214,6 +220,7 @@ impl PeerConnection { their_index: None, transport_id: Some(transport_id), source_addr: Some(source_addr), + remote_epoch: None, handshake_msg1: None, handshake_msg2: None, resend_count: 0, @@ -340,6 +347,13 @@ impl PeerConnection { self.source_addr = Some(addr); } + // === Epoch Accessors === + + /// Get the remote peer's startup epoch (available after handshake). + pub fn remote_epoch(&self) -> Option<[u8; 8]> { + self.remote_epoch + } + // === Handshake Resend === /// Store the wire-format msg1 bytes for resend and schedule the first resend. @@ -385,9 +399,11 @@ impl PeerConnection { /// Start the handshake as initiator and generate message 1. /// /// For outbound connections only. Returns the handshake message to send. + /// The epoch is our startup epoch, encrypted into msg1 for restart detection. pub fn start_handshake( &mut self, our_keypair: Keypair, + epoch: [u8; 8], current_time_ms: u64, ) -> Result, NoiseError> { if self.direction != LinkDirection::Outbound { @@ -411,6 +427,7 @@ impl PeerConnection { .pubkey_full(); let mut hs = noise::HandshakeState::new_initiator(our_keypair, remote_static); + hs.set_local_epoch(epoch); let msg1 = hs.write_message_1()?; self.noise_handshake = Some(hs); @@ -423,9 +440,11 @@ impl PeerConnection { /// Initialize responder and process incoming message 1. /// /// For inbound connections only. Returns the handshake message 2 to send. + /// The epoch is our startup epoch, encrypted into msg2 for restart detection. pub fn receive_handshake_init( &mut self, our_keypair: Keypair, + epoch: [u8; 8], message: &[u8], current_time_ms: u64, ) -> Result, NoiseError> { @@ -444,8 +463,9 @@ impl PeerConnection { } let mut hs = noise::HandshakeState::new_responder(our_keypair); + hs.set_local_epoch(epoch); - // Process message 1 (this reveals the initiator's identity) + // Process message 1 (this reveals the initiator's identity and epoch) hs.read_message_1(message)?; // Extract the discovered identity @@ -454,6 +474,9 @@ impl PeerConnection { .expect("remote static available after msg1"); self.expected_identity = Some(PeerIdentity::from_pubkey_full(remote_static)); + // Capture remote epoch from msg1 + self.remote_epoch = hs.remote_epoch(); + // Generate message 2 let msg2 = hs.write_message_2()?; @@ -488,6 +511,9 @@ impl PeerConnection { hs.read_message_2(message)?; + // Capture remote epoch from msg2 + self.remote_epoch = hs.remote_epoch(); + let session = hs.into_session()?; self.noise_session = Some(session); self.handshake_state = HandshakeState::Complete; @@ -557,6 +583,7 @@ impl fmt::Debug for PeerConnection { mod tests { use super::*; use crate::Identity; + use rand::RngCore; fn make_peer_identity() -> PeerIdentity { let identity = Identity::generate(); @@ -568,6 +595,12 @@ mod tests { identity.keypair() } + fn make_epoch() -> [u8; 8] { + let mut epoch = [0u8; 8]; + rand::thread_rng().fill_bytes(&mut epoch); + epoch + } + #[test] fn test_handshake_state_properties() { assert!(HandshakeState::Initial.is_in_progress()); @@ -611,6 +644,8 @@ mod tests { let initiator_keypair = initiator_identity.keypair(); let responder_keypair = responder_identity.keypair(); + let initiator_epoch = make_epoch(); + let responder_epoch = make_epoch(); // Use from_pubkey_full to preserve parity for ECDH let responder_peer_id = PeerIdentity::from_pubkey_full(responder_identity.pubkey_full()); @@ -621,12 +656,12 @@ mod tests { let mut responder_conn = PeerConnection::inbound(LinkId::new(2), 1000); // Initiator starts handshake - let msg1 = initiator_conn.start_handshake(initiator_keypair, 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 let msg2 = responder_conn - .receive_handshake_init(responder_keypair, &msg1, 1200) + .receive_handshake_init(responder_keypair, responder_epoch, &msg1, 1200) .unwrap(); assert_eq!(responder_conn.handshake_state(), HandshakeState::Complete); @@ -634,10 +669,16 @@ mod tests { let discovered = responder_conn.expected_identity().unwrap(); assert_eq!(discovered.pubkey(), initiator_identity.pubkey()); + // Responder learned initiator's epoch + assert_eq!(responder_conn.remote_epoch(), Some(initiator_epoch)); + // Initiator completes handshake initiator_conn.complete_handshake(&msg2, 1300).unwrap(); assert_eq!(initiator_conn.handshake_state(), HandshakeState::Complete); + // Initiator learned responder's epoch + assert_eq!(initiator_conn.remote_epoch(), Some(responder_epoch)); + // Both have sessions assert!(initiator_conn.has_session()); assert!(responder_conn.has_session()); @@ -683,11 +724,11 @@ mod tests { // Outbound can't receive_handshake_init let mut outbound = PeerConnection::outbound(LinkId::new(1), identity, 1000); assert!(outbound - .receive_handshake_init(keypair, &[0u8; 82], 1100) + .receive_handshake_init(keypair, make_epoch(), &[0u8; 106], 1100) .is_err()); // Inbound can't start_handshake let mut inbound = PeerConnection::inbound(LinkId::new(2), 1000); - assert!(inbound.start_handshake(keypair, 1100).is_err()); + assert!(inbound.start_handshake(keypair, make_epoch(), 1100).is_err()); } }