mirror of
https://github.com/jmcorgan/fips.git
synced 2026-07-22 07:48:26 +00:00
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
This commit is contained in:
@@ -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);
|
||||
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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!(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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());
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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<u8> {
|
||||
debug_assert_eq!(noise_msg1.len(), HANDSHAKE_MSG1_SIZE);
|
||||
|
||||
@@ -327,7 +327,7 @@ pub fn build_msg1(sender_idx: SessionIndex, noise_msg1: &[u8]) -> Vec<u8> {
|
||||
|
||||
/// 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<u8> {
|
||||
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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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<PublicKey>,
|
||||
/// Secp256k1 context.
|
||||
secp: Secp256k1<secp256k1::All>,
|
||||
/// 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<Vec<u8>, 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<Vec<u8>, 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()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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<MmpPeerState>,
|
||||
@@ -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.
|
||||
|
||||
@@ -117,8 +117,12 @@ pub struct PeerConnection {
|
||||
/// Current source address (updated on packet receipt).
|
||||
source_addr: Option<TransportAddr>,
|
||||
|
||||
// === 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<Vec<u8>>,
|
||||
|
||||
/// 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<Vec<u8>, 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<Vec<u8>, 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());
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user