From 2293f7d2d56bee38c51b028e037ca45001544d0d Mon Sep 17 00:00:00 2001 From: Johnathan Corgan Date: Sun, 22 Feb 2026 22:03:32 +0000 Subject: [PATCH] Replace Noise IK with Noise XK at the FSP session layer The session-layer handshake now uses the 3-message XK pattern instead of the 2-message IK pattern, providing stronger initiator identity hiding. The initiator static key is deferred to msg3 and encrypted under the es+ee DH chain, so eavesdroppers cannot identify the initiator from the handshake. XK pattern: -> e, es (msg1) / <- e, ee + epoch (msg2) / -> s, se + epoch (msg3) Key changes: - Add XK handshake methods alongside existing IK methods in noise module - Add SessionMsg3 wire format and FSP_PHASE_MSG3 (0x03) prefix - Replace Responding state with AwaitingMsg3 in session state machine - Rewrite session handlers: handle_session_setup defers identity to msg3, handle_session_ack processes msg2 and sends msg3, new handle_session_msg3 completes the responder handshake and registers identity - Link-layer (FMP) continues to use Noise IK unchanged - Add comprehensive XK unit tests and update all integration tests --- src/control/queries.rs | 4 +- src/node/handlers/session.rs | 269 +++++++++++++++++-------- src/node/handlers/timeout.rs | 4 +- src/node/session.rs | 32 +-- src/node/session_wire.rs | 27 ++- src/node/tests/session.rs | 96 ++++++--- src/noise/handshake.rs | 375 +++++++++++++++++++++++++++++++++-- src/noise/mod.rs | 79 +++++--- src/noise/tests.rs | 297 +++++++++++++++++++++++++++ src/protocol/mod.rs | 6 +- src/protocol/session.rs | 125 ++++++++++++ 11 files changed, 1128 insertions(+), 186 deletions(-) diff --git a/src/control/queries.rs b/src/control/queries.rs index edcf760..08e33bc 100644 --- a/src/control/queries.rs +++ b/src/control/queries.rs @@ -185,8 +185,8 @@ pub fn show_sessions(node: &Node) -> Value { "established" } else if entry.is_initiating() { "initiating" - } else if entry.is_responding() { - "responding" + } else if entry.is_awaiting_msg3() { + "awaiting_msg3" } else { "unknown" }; diff --git a/src/node/handlers/session.rs b/src/node/handlers/session.rs index 5414cf9..ae99622 100644 --- a/src/node/handlers/session.rs +++ b/src/node/handlers/session.rs @@ -2,24 +2,26 @@ //! //! Handles locally-delivered session payloads from SessionDatagram envelopes. //! Dispatches based on FSP common prefix phase to specific handlers for -//! SessionSetup (Noise IK msg1), SessionAck (msg2), encrypted data, -//! and error signals (CoordsRequired, PathBroken). +//! SessionSetup (Noise XK msg1), SessionAck (msg2), SessionMsg3 (msg3), +//! encrypted data, and error signals (CoordsRequired, PathBroken). use crate::node::session::{EndToEndState, SessionEntry}; use crate::node::session_wire::{ build_fsp_header, fsp_prepend_inner_header, fsp_strip_inner_header, parse_encrypted_coords, FspCommonPrefix, FspEncryptedHeader, FSP_COMMON_PREFIX_SIZE, FSP_FLAG_CP, FSP_HEADER_SIZE, FSP_PHASE_ESTABLISHED, FSP_PHASE_MSG1, FSP_PHASE_MSG2, + FSP_PHASE_MSG3, }; use crate::protocol::{coords_wire_size, encode_coords}; use crate::upper::icmp::FIPS_OVERHEAD; use crate::node::{Node, NodeError}; -use crate::noise::{HandshakeState, HANDSHAKE_MSG1_SIZE, HANDSHAKE_MSG2_SIZE}; +use crate::noise::{HandshakeState, XK_HANDSHAKE_MSG1_SIZE, XK_HANDSHAKE_MSG2_SIZE, XK_HANDSHAKE_MSG3_SIZE}; use crate::mmp::report::ReceiverReport; use crate::mmp::{MAX_SESSION_REPORT_INTERVAL_MS, MIN_SESSION_REPORT_INTERVAL_MS}; use crate::protocol::{ CoordsRequired, FspInnerFlags, MtuExceeded, PathBroken, PathMtuNotification, SessionAck, - SessionDatagram, SessionMessageType, SessionReceiverReport, SessionSenderReport, SessionSetup, + SessionDatagram, SessionMessageType, SessionMsg3, SessionReceiverReport, SessionSenderReport, + SessionSetup, }; use crate::NodeAddr; use secp256k1::PublicKey; @@ -33,6 +35,7 @@ impl Node { /// /// - Phase 0x1 → SessionSetup (handshake msg1) /// - Phase 0x2 → SessionAck (handshake msg2) + /// - Phase 0x3 → SessionMsg3 (XK handshake msg3) /// - Phase 0x0 + U flag → plaintext error signal (CoordsRequired/PathBroken) /// - Phase 0x0 + !U → encrypted session message (data, reports, etc.) pub(in crate::node) async fn handle_session_payload( @@ -58,6 +61,9 @@ impl Node { FSP_PHASE_MSG2 => { self.handle_session_ack(src_addr, inner).await; } + FSP_PHASE_MSG3 => { + self.handle_session_msg3(src_addr, inner).await; + } FSP_PHASE_ESTABLISHED if prefix.is_unencrypted() => { // Plaintext error signals: read msg_type from first byte after prefix if inner.is_empty() { @@ -95,7 +101,7 @@ impl Node { /// Full FSP receive pipeline: /// 1. Parse FspEncryptedHeader (12 bytes) → counter, flags, header_bytes /// 2. If CP flag: parse cleartext coords, cache them - /// 3. Session lookup with Responding→Established transition + /// 3. Session lookup (must be Established) /// 4. AEAD decrypt with AAD = header_bytes /// 5. Strip FSP inner header → timestamp, msg_type, inner_flags /// 6. Dispatch by msg_type @@ -135,39 +141,31 @@ impl Node { let ciphertext = &payload[ciphertext_offset..]; - // Look up session entry, handle Responding→Established transition - let mut entry = match self.sessions.remove(src_addr) { - Some(e) => e, - None => { - debug!(src = %self.peer_display_name(src_addr), "Encrypted session message for unknown session"); + // Look up session entry — must be Established to decrypt + { + let entry = match self.sessions.get(src_addr) { + Some(e) => e, + None => { + debug!(src = %self.peer_display_name(src_addr), "Encrypted session message for unknown session"); + return; + } + }; + // Drop encrypted data if session is not yet established. + // With XK, the responder must wait for msg3 before it can decrypt. + if !entry.is_established() { + debug!( + src = %self.peer_display_name(src_addr), + "Encrypted message but session not established (awaiting handshake completion)" + ); return; } - }; - - if entry.is_responding() { - let old_state = entry.take_state(); - let handshake = match old_state { - Some(EndToEndState::Responding(hs)) => hs, - _ => { - debug!(src = %self.peer_display_name(src_addr), "Unexpected state during Responding transition"); - return; - } - }; - let noise_session = match handshake.into_session() { - Ok(s) => s, - Err(e) => { - debug!(error = %e, "Failed to create session from responding handshake"); - return; - } - }; - entry.set_state(EndToEndState::Established(noise_session)); - entry.set_coords_warmup_remaining(self.config.node.session.coords_warmup_packets); - entry.mark_established(Self::now_ms()); - entry.init_mmp(&self.config.node.session_mmp); - entry.clear_handshake_payload(); - info!(src = %self.peer_display_name(src_addr), "Session established (responder, on first encrypted message)"); } + let mut entry = match self.sessions.remove(src_addr) { + Some(e) => e, + None => return, + }; + // Decrypt with AAD = the 12-byte header let session = match entry.state_mut() { EndToEndState::Established(s) => s, @@ -278,10 +276,11 @@ impl Node { self.flush_pending_packets(src_addr).await; } - /// Handle an incoming SessionSetup (Noise IK msg1). + /// Handle an incoming SessionSetup (Noise XK msg1). /// /// The remote node wants to establish an end-to-end session with us. - /// We create a responder handshake, process msg1, send SessionAck with msg2. + /// We create an XK responder handshake, process msg1, send SessionAck with msg2, + /// and transition to AwaitingMsg3. async fn handle_session_setup(&mut self, src_addr: &NodeAddr, inner: &[u8]) { let setup = match SessionSetup::decode(inner) { Ok(s) => s, @@ -291,10 +290,10 @@ impl Node { } }; - if setup.handshake_payload.len() != HANDSHAKE_MSG1_SIZE { + if setup.handshake_payload.len() != XK_HANDSHAKE_MSG1_SIZE { debug!( len = setup.handshake_payload.len(), - expected = HANDSHAKE_MSG1_SIZE, + expected = XK_HANDSHAKE_MSG1_SIZE, "Invalid handshake payload size in SessionSetup" ); return; @@ -317,8 +316,8 @@ impl Node { src = %self.peer_display_name(src_addr), "Simultaneous session initiation: we lose, becoming responder" ); - } else if existing.is_responding() { - // Duplicate setup while we already responded — resend stored ack + } else if existing.is_awaiting_msg3() { + // Duplicate setup while we already sent msg2 — resend stored ack if let Some(payload) = existing.handshake_payload() { debug!(src = %self.peer_display_name(src_addr), "Duplicate SessionSetup, resending SessionAck"); let my_addr = *self.node_addr(); @@ -337,33 +336,25 @@ impl Node { } } - // Create responder handshake and process msg1 + // Create XK responder handshake and process msg1 let our_keypair = self.identity.keypair(); - let mut handshake = HandshakeState::new_responder(our_keypair); + let mut handshake = HandshakeState::new_xk_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"); + if let Err(e) = handshake.read_xk_message_1(&setup.handshake_payload) { + debug!(error = %e, "Failed to process Noise XK msg1 in SessionSetup"); return; } - // Extract the initiator's static public key (learned from msg1) - let remote_pubkey = match handshake.remote_static() { - Some(pk) => *pk, - None => { - debug!("No remote static key after processing msg1"); - return; - } - }; - - // Register the initiator's identity for future TUN → session routing - self.register_identity(*src_addr, remote_pubkey); + // XK: responder does NOT learn initiator's identity until msg3 + // Use a placeholder pubkey from src_addr for the session entry. + // The real pubkey will be registered when msg3 arrives. // Generate msg2 - let msg2 = match handshake.write_message_2() { + let msg2 = match handshake.write_xk_message_2() { Ok(m) => m, Err(e) => { - debug!(error = %e, "Failed to generate Noise IK msg2 for SessionAck"); + debug!(error = %e, "Failed to generate Noise XK msg2 for SessionAck"); return; } }; @@ -382,19 +373,22 @@ impl Node { return; } - // Store session entry in Responding state with ack payload for potential resend + // Store session entry in AwaitingMsg3 state with ack payload for potential resend. + // Use a dummy pubkey since we don't know the initiator's identity yet. + // We use our own pubkey as placeholder; it will be replaced in handle_session_msg3. + let placeholder_pubkey = self.identity.keypair().public_key(); let now_ms = Self::now_ms(); let resend_interval = self.config.node.rate_limit.handshake_resend_interval_ms; - let mut entry = SessionEntry::new(*src_addr, remote_pubkey, EndToEndState::Responding(handshake), now_ms, false); + let mut entry = SessionEntry::new(*src_addr, placeholder_pubkey, EndToEndState::AwaitingMsg3(handshake), now_ms, false); entry.set_handshake_payload(ack_payload, now_ms + resend_interval); self.sessions.insert(*src_addr, entry); - debug!(src = %self.peer_display_name(src_addr), "SessionSetup processed, SessionAck sent"); + debug!(src = %self.peer_display_name(src_addr), "SessionSetup processed (XK), SessionAck sent, awaiting msg3"); } - /// Handle an incoming SessionAck (Noise IK msg2). + /// Handle an incoming SessionAck (Noise XK msg2). /// - /// Completes our initiated handshake, transitions to Established. + /// Processes msg2, generates and sends msg3, then transitions to Established. async fn handle_session_ack(&mut self, src_addr: &NodeAddr, inner: &[u8]) { let ack = match SessionAck::decode(inner) { Ok(a) => a, @@ -404,10 +398,10 @@ impl Node { } }; - if ack.handshake_payload.len() != HANDSHAKE_MSG2_SIZE { + if ack.handshake_payload.len() != XK_HANDSHAKE_MSG2_SIZE { debug!( len = ack.handshake_payload.len(), - expected = HANDSHAKE_MSG2_SIZE, + expected = XK_HANDSHAKE_MSG2_SIZE, "Invalid handshake payload size in SessionAck" ); return; @@ -428,17 +422,44 @@ impl Node { self.sessions.insert(*src_addr, entry); return; } - let handshake = match entry.take_state() { + let mut handshake = match entry.take_state() { Some(EndToEndState::Initiating(hs)) => hs, _ => unreachable!("checked is_initiating above"), }; - // Complete the handshake - let session = match Self::complete_initiator_handshake(handshake, &ack.handshake_payload) { + // Process XK msg2: read_xk_message_2 (extracts responder's epoch) + if let Err(e) = handshake.read_xk_message_2(&ack.handshake_payload) { + debug!(error = %e, "Failed to process Noise XK msg2 in SessionAck"); + return; // Entry was already removed, don't put back a broken session + } + + // Generate XK msg3: write_xk_message_3 (sends encrypted static + epoch) + let msg3 = match handshake.write_xk_message_3() { + Ok(m) => m, + Err(e) => { + debug!(error = %e, "Failed to generate Noise XK msg3"); + return; + } + }; + + // Send SessionMsg3 (phase 0x3) + let msg3_wire = SessionMsg3::new(msg3); + let msg3_payload = msg3_wire.encode(); + let my_addr = *self.node_addr(); + let mut datagram = SessionDatagram::new(my_addr, *src_addr, msg3_payload) + .with_ttl(self.config.node.session.default_ttl); + + if let Err(e) = self.send_session_datagram(&mut datagram).await { + debug!(error = %e, dest = %self.peer_display_name(src_addr), "Failed to send SessionMsg3"); + return; + } + + // Complete the handshake: into_session() + let session = match handshake.into_session() { Ok(s) => s, Err(e) => { - debug!(error = %e, "Failed to complete session handshake"); - return; // Entry was already removed, don't put back a broken session + debug!(error = %e, "Failed to create session after XK msg3"); + return; } }; @@ -455,7 +476,92 @@ impl Node { // Flush any queued outbound packets for this destination self.flush_pending_packets(src_addr).await; - info!(src = %self.peer_display_name(src_addr), "Session established (initiator)"); + info!(src = %self.peer_display_name(src_addr), "Session established (initiator, XK)"); + } + + /// Handle an incoming SessionMsg3 (Noise XK msg3). + /// + /// The initiator reveals their encrypted static key. The responder + /// processes msg3, learns the initiator's identity, and transitions + /// to Established. + async fn handle_session_msg3(&mut self, src_addr: &NodeAddr, inner: &[u8]) { + let msg3 = match SessionMsg3::decode(inner) { + Ok(m) => m, + Err(e) => { + debug!(error = %e, "Malformed SessionMsg3"); + return; + } + }; + + if msg3.handshake_payload.len() != XK_HANDSHAKE_MSG3_SIZE { + debug!( + len = msg3.handshake_payload.len(), + expected = XK_HANDSHAKE_MSG3_SIZE, + "Invalid handshake payload size in SessionMsg3" + ); + return; + } + + // Remove the entry to take ownership of the handshake state + let mut entry = match self.sessions.remove(src_addr) { + Some(e) => e, + None => { + debug!(src = %self.peer_display_name(src_addr), "SessionMsg3 for unknown session"); + return; + } + }; + + // Must be in AwaitingMsg3 state + if !entry.is_awaiting_msg3() { + debug!(src = %self.peer_display_name(src_addr), "SessionMsg3 but session not in AwaitingMsg3 state"); + self.sessions.insert(*src_addr, entry); + return; + } + let mut handshake = match entry.take_state() { + Some(EndToEndState::AwaitingMsg3(hs)) => hs, + _ => unreachable!("checked is_awaiting_msg3 above"), + }; + + // Process XK msg3: read_xk_message_3 (extracts initiator's static key and epoch) + if let Err(e) = handshake.read_xk_message_3(&msg3.handshake_payload) { + debug!(error = %e, "Failed to process Noise XK msg3"); + return; // Entry was already removed + } + + // Extract the initiator's static public key (now available after msg3) + let remote_pubkey = match handshake.remote_static() { + Some(pk) => *pk, + None => { + debug!("No remote static key after processing XK msg3"); + return; + } + }; + + // Register the initiator's identity for future TUN → session routing + self.register_identity(*src_addr, remote_pubkey); + + // Complete the handshake + let session = match handshake.into_session() { + Ok(s) => s, + Err(e) => { + debug!(error = %e, "Failed to create session from XK handshake"); + return; + } + }; + + let now_ms = Self::now_ms(); + // Replace the placeholder pubkey with the real one + let mut new_entry = SessionEntry::new(*src_addr, remote_pubkey, EndToEndState::Established(session), now_ms, false); + new_entry.set_coords_warmup_remaining(self.config.node.session.coords_warmup_packets); + new_entry.mark_established(now_ms); + new_entry.init_mmp(&self.config.node.session_mmp); + new_entry.touch(now_ms); + self.sessions.insert(*src_addr, new_entry); + + // Flush any pending packets + self.flush_pending_packets(src_addr).await; + + info!(src = %self.peer_display_name(src_addr), "Session established (responder, XK)"); } // === Session-layer MMP report handlers === @@ -590,19 +696,6 @@ impl Node { } } - /// Complete an initiator-side Noise IK handshake given msg2. - fn complete_initiator_handshake( - mut handshake: HandshakeState, - msg2: &[u8], - ) -> Result { - handshake - .read_message_2(msg2) - .map_err(|e| format!("read_message_2 failed: {}", e))?; - handshake - .into_session() - .map_err(|e| format!("into_session failed: {}", e)) - } - /// Handle a CoordsRequired error signal from a transit router. /// /// The router couldn't route our packet because it lacks cached @@ -751,7 +844,7 @@ impl Node { /// Initiate an end-to-end session with a remote node. /// - /// Creates a Noise IK handshake as initiator, wraps msg1 in a + /// Creates a Noise XK handshake as initiator, wraps msg1 in a /// SessionSetup, encapsulates in a SessionDatagram, and routes /// toward the destination. pub(in crate::node) async fn initiate_session( @@ -766,13 +859,13 @@ impl Node { return Ok(()); } - // Create Noise IK initiator handshake + // Create Noise XK initiator handshake let our_keypair = self.identity.keypair(); - let mut handshake = HandshakeState::new_initiator(our_keypair, dest_pubkey); + let mut handshake = HandshakeState::new_xk_initiator(our_keypair, dest_pubkey); handshake.set_local_epoch(self.startup_epoch); - let msg1 = handshake.write_message_1().map_err(|e| NodeError::SendFailed { + let msg1 = handshake.write_xk_message_1().map_err(|e| NodeError::SendFailed { node_addr: dest_addr, - reason: format!("Noise msg1 generation failed: {}", e), + reason: format!("Noise XK msg1 generation failed: {}", e), })?; // Build SessionSetup with coordinates diff --git a/src/node/handlers/timeout.rs b/src/node/handlers/timeout.rs index 190d414..a724425 100644 --- a/src/node/handlers/timeout.rs +++ b/src/node/handlers/timeout.rs @@ -153,7 +153,7 @@ impl Node { /// Resend session-layer handshake messages and timeout stale handshakes. /// - /// For sessions in Initiating or Responding state: + /// For sessions in Initiating or AwaitingMsg3 state: /// - If the handshake has exceeded the timeout window, remove the session. /// - If a resend is due and under max resends, resend the stored payload /// wrapped in a fresh SessionDatagram (so routing can adapt). @@ -231,7 +231,7 @@ impl Node { /// Remove established sessions that have been idle too long. /// - /// Only targets sessions in the Established state. Initiating/Responding + /// Only targets sessions in the Established state. Initiating/AwaitingMsg3 /// sessions are handled by the handshake timeout. pub(in crate::node) fn purge_idle_sessions(&mut self, now_ms: u64) { let timeout_ms = self.config.node.session.idle_timeout_secs * 1000; diff --git a/src/node/session.rs b/src/node/session.rs index dd0c07e..264d78a 100644 --- a/src/node/session.rs +++ b/src/node/session.rs @@ -1,8 +1,9 @@ //! End-to-end session state. //! -//! Tracks Noise IK sessions between this node and remote endpoints. -//! Sessions are established via SessionSetup/SessionAck handshake -//! messages carried inside SessionDatagram envelopes through the mesh. +//! Tracks Noise XK sessions between this node and remote endpoints. +//! Sessions are established via a three-message XK handshake +//! (SessionSetup/SessionAck/SessionMsg3) carried inside SessionDatagram +//! envelopes through the mesh. use crate::config::SessionMmpConfig; use crate::mmp::MmpSessionState; @@ -12,10 +13,11 @@ use secp256k1::PublicKey; /// State machine for an end-to-end session. pub(crate) enum EndToEndState { - /// We initiated: sent SessionSetup with Noise IK msg1, awaiting SessionAck. + /// We initiated: sent SessionSetup with Noise XK msg1, awaiting SessionAck. Initiating(HandshakeState), - /// We are responding: received msg1, sent SessionAck with msg2. - Responding(HandshakeState), + /// XK responder: processed msg1, sent msg2, awaiting msg3. + /// Transitions to Established when msg3 arrives. + AwaitingMsg3(HandshakeState), /// Handshake complete, NoiseSession available for encrypt/decrypt. Established(NoiseSession), } @@ -31,9 +33,9 @@ impl EndToEndState { matches!(self, EndToEndState::Initiating(_)) } - /// Check if we are the responder (sent ack, waiting for data). - pub(crate) fn is_responding(&self) -> bool { - matches!(self, EndToEndState::Responding(_)) + /// Check if we are an XK responder awaiting msg3. + pub(crate) fn is_awaiting_msg3(&self) -> bool { + matches!(self, EndToEndState::AwaitingMsg3(_)) } } @@ -46,7 +48,7 @@ pub(crate) struct SessionEntry { /// Remote node's address (session table key). #[allow(dead_code)] remote_addr: NodeAddr, - /// Remote node's static public key (for Noise IK). + /// Remote node's static public key. remote_pubkey: PublicKey, /// Current session state. `None` only during state transitions. state: Option, @@ -65,7 +67,7 @@ pub(crate) struct SessionEntry { /// Initialized from config when session becomes Established; /// reset on CoordsRequired receipt. coords_warmup_remaining: u8, - /// Whether this node initiated the Noise IK handshake. + /// Whether this node initiated the Noise handshake. /// Used for spin bit role assignment in session-layer MMP. is_initiator: bool, /// Session-layer MMP state. Initialized on Established transition. @@ -153,9 +155,9 @@ impl SessionEntry { self.state.as_ref().is_some_and(|s| s.is_initiating()) } - /// Check if we are the responder (sent ack, waiting for data). - pub(crate) fn is_responding(&self) -> bool { - self.state.as_ref().is_some_and(|s| s.is_responding()) + /// Check if we are an XK responder awaiting msg3. + pub(crate) fn is_awaiting_msg3(&self) -> bool { + self.state.as_ref().is_some_and(|s| s.is_awaiting_msg3()) } /// Get creation time. @@ -196,7 +198,7 @@ impl SessionEntry { now_ms.wrapping_sub(self.session_start_ms) as u32 } - /// Whether this node initiated the Noise IK handshake. + /// Whether this node initiated the Noise handshake. #[cfg_attr(not(test), allow(dead_code))] pub(crate) fn is_initiator(&self) -> bool { self.is_initiator diff --git a/src/node/session_wire.rs b/src/node/session_wire.rs index f4668e6..1ddf5ca 100644 --- a/src/node/session_wire.rs +++ b/src/node/session_wire.rs @@ -17,8 +17,9 @@ //! |-------|--------|------------------|-----------------------------------| //! | 0x0 | 0 | Encrypted | Post-handshake encrypted data | //! | 0x0 | 1 | Plaintext error | CoordsRequired, PathBroken | -//! | 0x1 | - | Handshake msg1 | SessionSetup (Noise IK msg1) | -//! | 0x2 | - | Handshake msg2 | SessionAck (Noise IK msg2) | +//! | 0x1 | - | Handshake msg1 | SessionSetup (Noise XK msg1) | +//! | 0x2 | - | Handshake msg2 | SessionAck (Noise XK msg2) | +//! | 0x3 | - | Handshake msg3 | SessionMsg3 (Noise XK msg3) | use crate::protocol::{ProtocolError, decode_optional_coords}; use crate::tree::TreeCoordinate; @@ -36,9 +37,12 @@ pub const FSP_PHASE_ESTABLISHED: u8 = 0x0; /// Phase value for SessionSetup (Noise IK message 1). pub const FSP_PHASE_MSG1: u8 = 0x1; -/// Phase value for SessionAck (Noise IK message 2). +/// Phase value for SessionAck (Noise handshake message 2). pub const FSP_PHASE_MSG2: u8 = 0x2; +/// Phase value for XK message 3 (initiator's encrypted static). +pub const FSP_PHASE_MSG3: u8 = 0x3; + /// Size of the common packet prefix (all FSP message types). pub const FSP_COMMON_PREFIX_SIZE: usize = 4; @@ -245,7 +249,7 @@ pub fn build_fsp_encrypted(header: &[u8; FSP_HEADER_SIZE], ciphertext: &[u8]) -> /// Build a 4-byte common prefix for a handshake message. /// -/// `phase` should be `FSP_PHASE_MSG1` or `FSP_PHASE_MSG2`. +/// `phase` should be `FSP_PHASE_MSG1`, `FSP_PHASE_MSG2`, or `FSP_PHASE_MSG3`. /// Flags are zero during handshake. #[cfg_attr(not(test), allow(dead_code))] pub fn build_fsp_handshake_prefix(phase: u8, payload_len: u16) -> [u8; FSP_COMMON_PREFIX_SIZE] { @@ -480,6 +484,17 @@ mod tests { assert_eq!(u16::from_le_bytes([prefix[2], prefix[3]]), 50); } + #[test] + fn test_build_fsp_handshake_prefix_msg3() { + let prefix = build_fsp_handshake_prefix(FSP_PHASE_MSG3, 73); + assert_eq!(prefix[0], 0x03); // ver=0, phase=3 + assert_eq!(prefix[1], 0x00); // flags zero + assert_eq!(u16::from_le_bytes([prefix[2], prefix[3]]), 73); + + let parsed = FspCommonPrefix::parse(&prefix).unwrap(); + assert_eq!(parsed.phase, FSP_PHASE_MSG3); + } + // ===== Error Prefix Tests ===== #[test] @@ -578,5 +593,9 @@ mod tests { // SessionAck (phase 2) let prefix = FspCommonPrefix::parse(&[0x02, 0x00, 0x21, 0x00]).unwrap(); assert_eq!(prefix.phase, 2); + + // SessionMsg3 (phase 3) + let prefix = FspCommonPrefix::parse(&[0x03, 0x00, 0x49, 0x00]).unwrap(); + assert_eq!(prefix.phase, 3); } } diff --git a/src/node/tests/session.rs b/src/node/tests/session.rs index 0bfb52b..6c15839 100644 --- a/src/node/tests/session.rs +++ b/src/node/tests/session.rs @@ -65,7 +65,7 @@ fn test_session_entry_new_initiating() { assert!(entry.state().is_initiating()); assert!(!entry.state().is_established()); - assert!(!entry.state().is_responding()); + assert!(!entry.state().is_awaiting_msg3()); assert_eq!(entry.created_at(), 1000); assert_eq!(entry.last_activity(), 1000); } @@ -163,21 +163,21 @@ async fn test_session_direct_peer_handshake() { let count = process_available_packets(&mut nodes).await; assert!(count > 0, "Expected SessionSetup packet to arrive"); - // Node 1 should now have a session in Responding state + // Node 1 should now have a session in AwaitingMsg3 state (XK: identity not yet known) assert_eq!(nodes[1].node.session_count(), 1); assert!(nodes[1] .node .get_session(&node0_addr) .unwrap() .state() - .is_responding()); + .is_awaiting_msg3()); - // Process packets: SessionAck arrives at Node 0 + // Process packets: SessionAck arrives at Node 0, Node 0 sends SessionMsg3 tokio::time::sleep(Duration::from_millis(20)).await; let count = process_available_packets(&mut nodes).await; assert!(count > 0, "Expected SessionAck packet to arrive"); - // Node 0 should now be Established + // Node 0 should now be Established (transitions after sending msg3) assert!(nodes[0] .node .get_session(&node1_addr) @@ -185,6 +185,19 @@ async fn test_session_direct_peer_handshake() { .state() .is_established()); + // Process packets: SessionMsg3 arrives at Node 1 + tokio::time::sleep(Duration::from_millis(20)).await; + let count = process_available_packets(&mut nodes).await; + assert!(count > 0, "Expected SessionMsg3 packet to arrive"); + + // Node 1 should now be Established (transitions after processing msg3) + assert!(nodes[1] + .node + .get_session(&node0_addr) + .unwrap() + .state() + .is_established()); + cleanup_nodes(&mut nodes).await; } @@ -200,7 +213,7 @@ async fn test_session_direct_peer_data_transfer() { let node1_addr = *nodes[1].node.node_addr(); let node1_pubkey = nodes[1].node.identity().pubkey_full(); - // Establish session + // Establish session (XK: 3 messages — Setup, Ack, Msg3) nodes[0] .node .initiate_session(node1_addr, node1_pubkey) @@ -209,7 +222,9 @@ async fn test_session_direct_peer_data_transfer() { tokio::time::sleep(Duration::from_millis(20)).await; process_available_packets(&mut nodes).await; // Setup → Node 1 tokio::time::sleep(Duration::from_millis(20)).await; - process_available_packets(&mut nodes).await; // Ack → Node 0 + process_available_packets(&mut nodes).await; // Ack → Node 0, Node 0 sends Msg3 + tokio::time::sleep(Duration::from_millis(20)).await; + process_available_packets(&mut nodes).await; // Msg3 → Node 1 assert!(nodes[0] .node @@ -217,6 +232,12 @@ async fn test_session_direct_peer_data_transfer() { .unwrap() .state() .is_established()); + assert!(nodes[1] + .node + .get_session(&node0_addr) + .unwrap() + .state() + .is_established()); // Send data from Node 0 to Node 1 let test_data = b"Hello, FIPS session!"; @@ -231,14 +252,6 @@ async fn test_session_direct_peer_data_transfer() { let count = process_available_packets(&mut nodes).await; assert!(count > 0, "Expected encrypted data to arrive"); - // Node 1's session should now be Established (was Responding, transitions on first data) - assert!(nodes[1] - .node - .get_session(&node0_addr) - .unwrap() - .state() - .is_established()); - cleanup_nodes(&mut nodes).await; } @@ -273,7 +286,7 @@ async fn test_session_3node_forwarded_handshake() { tokio::time::sleep(Duration::from_millis(20)).await; process_available_packets(&mut nodes).await; - // Node 2 should have a Responding session + // Node 2 should have an AwaitingMsg3 session (XK: identity not yet known) assert!( nodes[2].node.get_session(&node0_addr).is_some(), "Node 2 should have a session entry for Node 0" @@ -283,17 +296,17 @@ async fn test_session_3node_forwarded_handshake() { .get_session(&node0_addr) .unwrap() .state() - .is_responding()); + .is_awaiting_msg3()); // Process: SessionAck: 2→1 (forwarded by transit B) tokio::time::sleep(Duration::from_millis(20)).await; process_available_packets(&mut nodes).await; - // Process: SessionAck: 1→0 (arrives at initiator A) + // Process: SessionAck: 1→0 (arrives at initiator A, sends SessionMsg3) tokio::time::sleep(Duration::from_millis(20)).await; process_available_packets(&mut nodes).await; - // Node 0 should now be Established + // Node 0 should now be Established (transitions after sending msg3) assert!(nodes[0] .node .get_session(&node2_addr) @@ -301,6 +314,22 @@ async fn test_session_3node_forwarded_handshake() { .state() .is_established()); + // Process: SessionMsg3: 0→1 (forwarded by transit B) + tokio::time::sleep(Duration::from_millis(20)).await; + process_available_packets(&mut nodes).await; + + // Process: SessionMsg3: 1→2 (arrives at responder C) + tokio::time::sleep(Duration::from_millis(20)).await; + process_available_packets(&mut nodes).await; + + // Node 2 should now be Established (transitions after processing msg3) + assert!(nodes[2] + .node + .get_session(&node0_addr) + .unwrap() + .state() + .is_established()); + // Transit node B should NOT have a session assert_eq!( nodes[1].node.session_count(), @@ -359,7 +388,7 @@ async fn test_session_3node_forwarded_data() { process_available_packets(&mut nodes).await; } - // Node 2 should have transitioned to Established on first data + // Node 2 should be Established (transitioned during XK handshake msg3) assert!(nodes[2] .node .get_session(&node0_addr) @@ -426,7 +455,7 @@ async fn test_session_ack_for_unknown_session() { // Fabricate a SessionAck and deliver directly let src_coords = nodes[1].node.tree_state().my_coords().clone(); let dest_coords = nodes[0].node.tree_state().my_coords().clone(); - let ack = SessionAck::new(src_coords, dest_coords).with_handshake(vec![0u8; 33]); + let ack = SessionAck::new(src_coords, dest_coords).with_handshake(vec![0u8; 57]); let datagram = SessionDatagram::new(node1_addr, node0_addr, ack.encode()); // Send through link layer @@ -575,7 +604,6 @@ async fn test_session_100_nodes() { // // For each session pair: // 1. Initiator sends one datagram to responder - // (this also transitions responder from Responding → Established) // 2. Responder sends one datagram back to initiator // // Batched per pair with draining between each. @@ -604,7 +632,7 @@ async fn test_session_100_nodes() { drain_to_quiescence(&mut nodes).await; // Reverse: responder → initiator - // (Responder should now be Established after receiving the forward datagram) + // (Responder should already be Established after XK msg3) let rev_payload = format!("rev-{}", pair_idx).into_bytes(); match nodes[dst] .node @@ -668,7 +696,7 @@ async fn test_session_100_nodes() { for (_, entry) in tn.node.sessions.iter() { if entry.state().is_established() { total_established += 1; - } else if entry.state().is_responding() { + } else if entry.state().is_awaiting_msg3() { total_responding += 1; all_est = false; } else { @@ -848,7 +876,7 @@ async fn test_session_100_nodes() { ); assert_eq!( send_reverse_err, 0, - "All reverse sends should succeed (responder Established after forward data)" + "All reverse sends should succeed (responder Established after XK msg3)" ); assert_eq!( fwd_delivered, send_forward_ok, @@ -943,12 +971,14 @@ async fn test_tun_outbound_established_session() { let src_fips = crate::FipsAddress::from_node_addr(&node0_addr); let dst_fips = crate::FipsAddress::from_node_addr(&node1_addr); - // Establish session + // Establish session (XK: 3 messages — Setup, Ack, Msg3) nodes[0].node.initiate_session(node1_addr, node1_pubkey).await.unwrap(); tokio::time::sleep(Duration::from_millis(20)).await; process_available_packets(&mut nodes).await; // Setup → Node 1 tokio::time::sleep(Duration::from_millis(20)).await; - process_available_packets(&mut nodes).await; // Ack → Node 0 + process_available_packets(&mut nodes).await; // Ack → Node 0, Node 0 sends Msg3 + tokio::time::sleep(Duration::from_millis(20)).await; + process_available_packets(&mut nodes).await; // Msg3 → Node 1 assert!(nodes[0].node.get_session(&node1_addr).unwrap().state().is_established()); @@ -1636,9 +1666,9 @@ async fn test_session_handshake_timeout() { assert!(!node.sessions.contains_key(&dest_addr), "Timed-out session should be removed"); } -/// Test that session handshake timeout removes stale Responding sessions. +/// Test that session handshake timeout removes stale AwaitingMsg3 sessions. #[tokio::test] -async fn test_session_responding_timeout() { +async fn test_session_awaiting_msg3_timeout() { use crate::noise::HandshakeState; let mut node = make_node(); @@ -1646,17 +1676,17 @@ async fn test_session_responding_timeout() { let identity_a = Identity::generate(); let identity_b = Identity::generate(); - let handshake = HandshakeState::new_responder( + let handshake = HandshakeState::new_xk_responder( identity_b.keypair(), ); let src_addr = *identity_a.node_addr(); - // Create a Responding session at time 1000 + // Create an AwaitingMsg3 session at time 1000 let entry = crate::node::session::SessionEntry::new( src_addr, identity_a.pubkey_full(), - EndToEndState::Responding(handshake), + EndToEndState::AwaitingMsg3(handshake), 1000, false, ); @@ -1668,5 +1698,5 @@ async fn test_session_responding_timeout() { let timeout_secs = node.config.node.rate_limit.handshake_timeout_secs; let after_timeout = 1000 + timeout_secs * 1000 + 1; node.resend_pending_session_handshakes(after_timeout).await; - assert!(!node.sessions.contains_key(&src_addr), "Timed-out Responding session should be removed"); + assert!(!node.sessions.contains_key(&src_addr), "Timed-out AwaitingMsg3 session should be removed"); } diff --git a/src/noise/handshake.rs b/src/noise/handshake.rs index c82de13..71b4975 100644 --- a/src/noise/handshake.rs +++ b/src/noise/handshake.rs @@ -1,7 +1,8 @@ use super::{ - CipherState, HandshakeProgress, HandshakeRole, NoiseError, NoiseSession, + CipherState, HandshakeProgress, HandshakeRole, NoiseError, NoisePattern, NoiseSession, EPOCH_ENCRYPTED_SIZE, EPOCH_SIZE, HANDSHAKE_MSG1_SIZE, HANDSHAKE_MSG2_SIZE, - PROTOCOL_NAME, PUBKEY_SIZE, + PROTOCOL_NAME_IK, PROTOCOL_NAME_XK, PUBKEY_SIZE, + XK_HANDSHAKE_MSG1_SIZE, XK_HANDSHAKE_MSG2_SIZE, XK_HANDSHAKE_MSG3_SIZE, }; use hkdf::Hkdf; use rand::RngCore; @@ -23,16 +24,16 @@ struct SymmetricState { impl SymmetricState { /// Initialize with protocol name. - fn initialize() -> Self { + fn initialize(protocol_name: &[u8]) -> Self { // If protocol name <= 32 bytes, pad with zeros // If > 32 bytes, hash it - let h = if PROTOCOL_NAME.len() <= 32 { + let h = if protocol_name.len() <= 32 { let mut h = [0u8; 32]; - h[..PROTOCOL_NAME.len()].copy_from_slice(PROTOCOL_NAME); + h[..protocol_name.len()].copy_from_slice(protocol_name); h } else { let mut hasher = Sha256::new(); - hasher.update(PROTOCOL_NAME); + hasher.update(protocol_name); hasher.finalize().into() }; @@ -101,8 +102,10 @@ impl SymmetricState { } } -/// Handshake state for Noise IK. +/// Handshake state for Noise IK and XK patterns. pub struct HandshakeState { + /// Which Noise pattern is being used. + pattern: NoisePattern, /// Our role in the handshake. role: HandshakeRole, /// Current progress. @@ -114,8 +117,10 @@ pub struct HandshakeState { /// Our ephemeral keypair (generated at handshake start). ephemeral_keypair: Option, /// Remote static public key. - /// For initiator: known before handshake (from config). - /// For responder: learned from message 1. + /// For IK initiator: known before handshake (from config). + /// For IK responder: learned from message 1. + /// For XK initiator: known before handshake (from config). + /// For XK responder: learned from message 3. remote_static: Option, /// Remote ephemeral public key (learned during handshake). remote_ephemeral: Option, @@ -144,15 +149,17 @@ impl HandshakeState { bytes } - /// Create a new handshake as initiator. + /// Create a new IK handshake as initiator. /// /// The initiator knows the responder's static key and will send first. + /// Used by FMP (link layer). pub fn new_initiator(static_keypair: Keypair, remote_static: PublicKey) -> Self { let secp = Secp256k1::new(); let mut state = Self { + pattern: NoisePattern::Ik, role: HandshakeRole::Initiator, progress: HandshakeProgress::Initial, - symmetric: SymmetricState::initialize(), + symmetric: SymmetricState::initialize(PROTOCOL_NAME_IK), static_keypair, ephemeral_keypair: None, remote_static: Some(remote_static), @@ -171,16 +178,17 @@ impl HandshakeState { state } - /// Create a new handshake as responder. + /// Create a new IK handshake as responder. /// /// The responder does NOT know the initiator's static key - it will be - /// learned from message 1. + /// learned from message 1. Used by FMP (link layer). pub fn new_responder(static_keypair: Keypair) -> Self { let secp = Secp256k1::new(); let mut state = Self { + pattern: NoisePattern::Ik, role: HandshakeRole::Responder, progress: HandshakeProgress::Initial, - symmetric: SymmetricState::initialize(), + symmetric: SymmetricState::initialize(PROTOCOL_NAME_IK), static_keypair, ephemeral_keypair: None, remote_static: None, // Will learn from message 1 @@ -198,6 +206,60 @@ impl HandshakeState { state } + /// Create a new XK handshake as initiator. + /// + /// The initiator knows the responder's static key. XK defers the + /// initiator's static key reveal to msg3. Used by FSP (session layer). + pub fn new_xk_initiator(static_keypair: Keypair, remote_static: PublicKey) -> Self { + let secp = Secp256k1::new(); + let mut state = Self { + pattern: NoisePattern::Xk, + role: HandshakeRole::Initiator, + progress: HandshakeProgress::Initial, + symmetric: SymmetricState::initialize(PROTOCOL_NAME_XK), + static_keypair, + ephemeral_keypair: None, + 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) + let normalized = Self::normalize_for_premessage(&remote_static); + state.symmetric.mix_hash(&normalized); + + state + } + + /// Create a new XK handshake as responder. + /// + /// The responder does NOT know the initiator's static key - it will be + /// learned from message 3. Used by FSP (session layer). + pub fn new_xk_responder(static_keypair: Keypair) -> Self { + let secp = Secp256k1::new(); + let mut state = Self { + pattern: NoisePattern::Xk, + role: HandshakeRole::Responder, + progress: HandshakeProgress::Initial, + symmetric: SymmetricState::initialize(PROTOCOL_NAME_XK), + static_keypair, + ephemeral_keypair: None, + remote_static: None, // Will learn from message 3 + remote_ephemeral: None, + secp, + local_epoch: None, + remote_epoch: None, + }; + + // Mix in pre-message: <- s (our static, since we're responder) + let normalized = Self::normalize_for_premessage(&state.static_keypair.public_key()); + state.symmetric.mix_hash(&normalized); + + state + } + /// Get our role. pub fn role(&self) -> HandshakeRole { self.role @@ -485,6 +547,290 @@ impl HandshakeState { Ok(()) } + // ======================================================================== + // XK Pattern Methods (Session Layer) + // ======================================================================== + + /// Write XK message 1 (initiator only). + /// + /// XK msg1: `-> e, es` + /// - e: ephemeral public key (33 bytes) + /// - es: DH(e_priv, rs_pub), mix_key + /// + /// Total: 33 bytes (ephemeral only — no static, no epoch) + pub fn write_xk_message_1(&mut self) -> Result, NoiseError> { + if self.role != HandshakeRole::Initiator { + return Err(NoiseError::WrongState { + expected: "initiator".to_string(), + got: "responder".to_string(), + }); + } + if self.progress != HandshakeProgress::Initial { + return Err(NoiseError::WrongState { + expected: HandshakeProgress::Initial.to_string(), + got: self.progress.to_string(), + }); + } + + let remote_static = self.remote_static.expect("initiator must have remote static"); + + // 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(XK_HANDSHAKE_MSG1_SIZE); + + // -> e: send ephemeral, mix into hash + message.extend_from_slice(&e_pub); + self.symmetric.mix_hash(&e_pub); + + // -> es: DH(e, rs), mix into key + let es = self.ecdh(&ephemeral.secret_key(), &remote_static); + self.symmetric.mix_key(&es); + + self.progress = HandshakeProgress::Message1Done; + + Ok(message) + } + + /// Read XK message 1 (responder only). + /// + /// Processes the initiator's first message. Does NOT learn initiator's + /// identity (that comes in msg3). + pub fn read_xk_message_1(&mut self, message: &[u8]) -> Result<(), NoiseError> { + if self.role != HandshakeRole::Responder { + return Err(NoiseError::WrongState { + expected: "responder".to_string(), + got: "initiator".to_string(), + }); + } + if self.progress != HandshakeProgress::Initial { + return Err(NoiseError::WrongState { + expected: HandshakeProgress::Initial.to_string(), + got: self.progress.to_string(), + }); + } + if message.len() != XK_HANDSHAKE_MSG1_SIZE { + return Err(NoiseError::MessageTooShort { + expected: XK_HANDSHAKE_MSG1_SIZE, + got: message.len(), + }); + } + + // -> e: parse remote ephemeral, mix into hash + let re = PublicKey::from_slice(&message[..PUBKEY_SIZE]) + .map_err(|_| NoiseError::InvalidPublicKey)?; + self.remote_ephemeral = Some(re); + self.symmetric.mix_hash(&message[..PUBKEY_SIZE]); + + // -> es: DH(s, re), mix into key + // (responder uses their static with initiator's ephemeral) + let es = self.ecdh(&self.static_keypair.secret_key(), &re); + self.symmetric.mix_key(&es); + + self.progress = HandshakeProgress::Message1Done; + + Ok(()) + } + + /// Write XK message 2 (responder only). + /// + /// XK msg2: `<- e, ee` + encrypted epoch + /// - e: ephemeral public key (33 bytes) + /// - ee: DH(e_priv, re_pub), mix_key + /// - encrypted epoch (24 bytes) + /// + /// Total: 57 bytes + pub fn write_xk_message_2(&mut self) -> Result, NoiseError> { + if self.role != HandshakeRole::Responder { + return Err(NoiseError::WrongState { + expected: "responder".to_string(), + got: "initiator".to_string(), + }); + } + if self.progress != HandshakeProgress::Message1Done { + return Err(NoiseError::WrongState { + expected: HandshakeProgress::Message1Done.to_string(), + got: self.progress.to_string(), + }); + } + + let re = self.remote_ephemeral.expect("should have remote ephemeral"); + let epoch = self.local_epoch.expect("local epoch must be set before write_xk_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(XK_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 + let ee = self.ecdh(&ephemeral.secret_key(), &re); + self.symmetric.mix_key(&ee); + + // <- 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::Message2Done; + + Ok(message) + } + + /// Read XK message 2 (initiator only). + /// + /// Processes the responder's message and extracts the responder's epoch. + /// Does NOT complete the handshake — msg3 still needed. + pub fn read_xk_message_2(&mut self, message: &[u8]) -> Result<(), NoiseError> { + if self.role != HandshakeRole::Initiator { + return Err(NoiseError::WrongState { + expected: "initiator".to_string(), + got: "responder".to_string(), + }); + } + if self.progress != HandshakeProgress::Message1Done { + return Err(NoiseError::WrongState { + expected: HandshakeProgress::Message1Done.to_string(), + got: self.progress.to_string(), + }); + } + if message.len() != XK_HANDSHAKE_MSG2_SIZE { + return Err(NoiseError::MessageTooShort { + expected: XK_HANDSHAKE_MSG2_SIZE, + got: message.len(), + }); + } + + // <- e: parse remote ephemeral, mix into hash + 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(e_pub); + + // <- ee: DH(e, re), mix into key + let ephemeral = self.ephemeral_keypair.as_ref().unwrap(); + let ee = self.ecdh(&ephemeral.secret_key(), &re); + self.symmetric.mix_key(&ee); + + // <- 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::Message2Done; + + Ok(()) + } + + /// Write XK message 3 (initiator only). + /// + /// XK msg3: `-> s, se` + encrypted epoch + /// - s: encrypt_and_hash(s_pub) — encrypted static (49 bytes) + /// - se: DH(s_priv, re_pub), mix_key + /// - encrypted epoch (24 bytes) + /// + /// Total: 73 bytes + pub fn write_xk_message_3(&mut self) -> Result, NoiseError> { + if self.role != HandshakeRole::Initiator { + return Err(NoiseError::WrongState { + expected: "initiator".to_string(), + got: "responder".to_string(), + }); + } + if self.progress != HandshakeProgress::Message2Done { + return Err(NoiseError::WrongState { + expected: HandshakeProgress::Message2Done.to_string(), + got: self.progress.to_string(), + }); + } + + let re = self.remote_ephemeral.expect("should have remote ephemeral after msg2"); + let epoch = self.local_epoch.expect("local epoch must be set before write_xk_message_3"); + + let mut message = Vec::with_capacity(XK_HANDSHAKE_MSG3_SIZE); + + // -> s: encrypt our static and send + let our_static = self.static_keypair.public_key().serialize(); + let encrypted_static = self.symmetric.encrypt_and_hash(&our_static)?; + message.extend_from_slice(&encrypted_static); + + // -> se: DH(s, re), mix into key + 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(message) + } + + /// Read XK message 3 (responder only). + /// + /// Processes the initiator's encrypted static key and epoch. + /// After this, the responder learns the initiator's identity. + pub fn read_xk_message_3(&mut self, message: &[u8]) -> Result<(), NoiseError> { + if self.role != HandshakeRole::Responder { + return Err(NoiseError::WrongState { + expected: "responder".to_string(), + got: "initiator".to_string(), + }); + } + if self.progress != HandshakeProgress::Message2Done { + return Err(NoiseError::WrongState { + expected: HandshakeProgress::Message2Done.to_string(), + got: self.progress.to_string(), + }); + } + if message.len() != XK_HANDSHAKE_MSG3_SIZE { + return Err(NoiseError::MessageTooShort { + expected: XK_HANDSHAKE_MSG3_SIZE, + got: message.len(), + }); + } + + // -> s: decrypt initiator's static + let encrypted_static_end = PUBKEY_SIZE + super::TAG_SIZE; + let encrypted_static = &message[..encrypted_static_end]; + let decrypted_static = self.symmetric.decrypt_and_hash(encrypted_static)?; + let rs = + PublicKey::from_slice(&decrypted_static).map_err(|_| NoiseError::InvalidPublicKey)?; + self.remote_static = Some(rs); + + // -> se: DH(e, rs), mix into key + // (responder uses their ephemeral with initiator's now-known static) + let ephemeral = self.ephemeral_keypair.as_ref().expect("should have ephemeral after msg2"); + let se = self.ecdh(&ephemeral.secret_key(), &rs); + self.symmetric.mix_key(&se); + + // -> 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::Complete; + + Ok(()) + } + /// Complete the handshake and return a NoiseSession. /// /// Must be called after the handshake is complete. @@ -524,6 +870,7 @@ impl HandshakeState { impl fmt::Debug for HandshakeState { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { f.debug_struct("HandshakeState") + .field("pattern", &self.pattern) .field("role", &self.role) .field("progress", &self.progress) .field("has_ephemeral", &self.ephemeral_keypair.is_some()) diff --git a/src/noise/mod.rs b/src/noise/mod.rs index d3788c4..d90b9b2 100644 --- a/src/noise/mod.rs +++ b/src/noise/mod.rs @@ -1,35 +1,39 @@ -//! Noise IK Protocol for Peer Authentication +//! Noise Protocol Implementations for FIPS //! -//! Implements the Noise Protocol Framework IK pattern using secp256k1 -//! for link-local peer authentication. This establishes encrypted -//! channels between direct peers over a transport. +//! Implements Noise Protocol Framework patterns using secp256k1: //! -//! The IK pattern assumes the initiator knows the responder's static -//! public key before the handshake. The responder learns the initiator's -//! identity from the encrypted payload in message 1. +//! - **IK pattern**: Used by FMP (link layer) for hop-by-hop peer authentication. +//! The initiator knows the responder's static key and sends its encrypted +//! static in msg1. Two-message handshake. //! -//! ## Handshake Pattern +//! - **XK pattern**: Used by FSP (session layer) for end-to-end sessions. +//! The initiator knows the responder's static key but defers revealing its +//! own identity until msg3, providing stronger identity hiding. Three-message +//! handshake. +//! +//! ## IK Handshake Pattern (Link Layer) //! -//! Pre-message (key known before handshake): //! ```text -//! <- s (responder's static known to initiator) +//! <- s (pre-message: responder's static known) +//! -> e, es, s, ss (msg1: ephemeral + encrypted static) +//! <- e, ee, se (msg2: ephemeral) //! ``` //! -//! Messages: -//! ```text -//! -> e, es, s, ss (initiator sends ephemeral + encrypted static) -//! <- e, ee, se (responder sends ephemeral) -//! ``` +//! ## XK Handshake Pattern (Session Layer) //! -//! After handshake, both parties derive symmetric keys for bidirectional -//! encrypted communication over the peer link. +//! ```text +//! <- s (pre-message: responder's static known) +//! -> e, es (msg1: ephemeral + DH with responder's static) +//! <- e, ee (msg2: ephemeral + DH) +//! -> s, se (msg3: encrypted static + DH) +//! ``` //! //! ## Separation of Concerns //! -//! This module handles **peer authentication** only - securing the direct -//! link between neighboring nodes. End-to-end FIPS session encryption -//! between arbitrary network addresses is a separate concern handled by -//! the session layer. +//! The IK pattern handles **link-layer peer authentication** — securing the +//! direct link between neighboring nodes. The XK pattern handles **session-layer +//! end-to-end encryption** between arbitrary network addresses, with stronger +//! initiator identity protection. mod handshake; mod replay; @@ -46,9 +50,13 @@ pub use handshake::HandshakeState; pub use replay::ReplayWindow; pub use session::NoiseSession; -/// Protocol name for Noise IK with secp256k1. +/// Protocol name for Noise IK with secp256k1 (link layer). /// Format: Noise_IK_secp256k1_ChaChaPoly_SHA256 -pub(crate) const PROTOCOL_NAME: &[u8] = b"Noise_IK_secp256k1_ChaChaPoly_SHA256"; +pub(crate) const PROTOCOL_NAME_IK: &[u8] = b"Noise_IK_secp256k1_ChaChaPoly_SHA256"; + +/// Protocol name for Noise XK with secp256k1 (session layer). +/// Format: Noise_XK_secp256k1_ChaChaPoly_SHA256 +pub(crate) const PROTOCOL_NAME_XK: &[u8] = b"Noise_XK_secp256k1_ChaChaPoly_SHA256"; /// Maximum message size for noise transport messages. pub const MAX_MESSAGE_SIZE: usize = 65535; @@ -65,12 +73,21 @@ pub const EPOCH_SIZE: usize = 8; /// 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). +/// Size of IK 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). +/// Size of IK handshake message 2: ephemeral (33) + encrypted epoch (8 + 16 tag). pub const HANDSHAKE_MSG2_SIZE: usize = PUBKEY_SIZE + EPOCH_ENCRYPTED_SIZE; +/// XK msg1: ephemeral only (33 bytes). +pub const XK_HANDSHAKE_MSG1_SIZE: usize = PUBKEY_SIZE; + +/// XK msg2: ephemeral (33) + encrypted epoch (8 + 16 tag) = 57 bytes. +pub const XK_HANDSHAKE_MSG2_SIZE: usize = PUBKEY_SIZE + EPOCH_ENCRYPTED_SIZE; + +/// XK msg3: encrypted static (33 + 16 tag) + encrypted epoch (8 + 16 tag) = 73 bytes. +pub const XK_HANDSHAKE_MSG3_SIZE: usize = PUBKEY_SIZE + TAG_SIZE + EPOCH_ENCRYPTED_SIZE; + /// Replay window size in packets (matching WireGuard). pub const REPLAY_WINDOW_SIZE: usize = 2048; @@ -129,6 +146,15 @@ impl fmt::Display for HandshakeRole { } } +/// Which Noise pattern is being used for this handshake. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum NoisePattern { + /// Noise IK: two-message handshake (link layer). + Ik, + /// Noise XK: three-message handshake (session layer). + Xk, +} + /// Handshake state machine states. #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub enum HandshakeProgress { @@ -136,6 +162,8 @@ pub enum HandshakeProgress { Initial, /// Message 1 sent/received, ready for message 2. Message1Done, + /// Message 2 sent/received, ready for message 3 (XK only). + Message2Done, /// Handshake complete, ready for transport. Complete, } @@ -145,6 +173,7 @@ impl fmt::Display for HandshakeProgress { match self { HandshakeProgress::Initial => write!(f, "initial"), HandshakeProgress::Message1Done => write!(f, "message1_done"), + HandshakeProgress::Message2Done => write!(f, "message2_done"), HandshakeProgress::Complete => write!(f, "complete"), } } diff --git a/src/noise/tests.rs b/src/noise/tests.rs index 0eb7066..a986d02 100644 --- a/src/noise/tests.rs +++ b/src/noise/tests.rs @@ -451,3 +451,300 @@ fn test_handshake_with_odd_parity_responder() { .unwrap(); assert_eq!(plaintext, b"parity test"); } + +// ===== XK Handshake Tests ===== + +#[test] +fn test_xk_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(); + + // XK: initiator knows responder's static, responder learns initiator's in msg3 + let mut initiator = HandshakeState::new_xk_initiator(initiator_keypair, responder_pub); + initiator.set_local_epoch(initiator_epoch); + let mut responder = HandshakeState::new_xk_responder(responder_keypair); + responder.set_local_epoch(responder_epoch); + + assert_eq!(initiator.role(), HandshakeRole::Initiator); + assert_eq!(responder.role(), HandshakeRole::Responder); + + // Initially, responder doesn't know initiator's identity + assert!(responder.remote_static().is_none()); + + // Message 1: Initiator -> Responder (e, es) + let msg1 = initiator.write_xk_message_1().unwrap(); + assert_eq!(msg1.len(), XK_HANDSHAKE_MSG1_SIZE); + assert_eq!(msg1.len(), 33); // ephemeral only + + responder.read_xk_message_1(&msg1).unwrap(); + + // After msg1: responder still doesn't know initiator's identity (XK property) + assert!(responder.remote_static().is_none()); + assert!(responder.remote_epoch().is_none()); + + // Message 2: Responder -> Initiator (e, ee + epoch) + let msg2 = responder.write_xk_message_2().unwrap(); + assert_eq!(msg2.len(), XK_HANDSHAKE_MSG2_SIZE); + assert_eq!(msg2.len(), 57); // 33 ephemeral + 24 encrypted epoch + + initiator.read_xk_message_2(&msg2).unwrap(); + + // After msg2: initiator learned responder's epoch + assert_eq!(initiator.remote_epoch(), Some(responder_epoch)); + // Neither side is complete yet + assert!(!initiator.is_complete()); + assert!(!responder.is_complete()); + + // Message 3: Initiator -> Responder (s, se + epoch) + let msg3 = initiator.write_xk_message_3().unwrap(); + assert_eq!(msg3.len(), XK_HANDSHAKE_MSG3_SIZE); + assert_eq!(msg3.len(), 73); // 49 encrypted static + 24 encrypted epoch + + responder.read_xk_message_3(&msg3).unwrap(); + + // Both should be complete now + assert!(initiator.is_complete()); + assert!(responder.is_complete()); + + // After msg3: responder now knows initiator's identity + assert!(responder.remote_static().is_some()); + assert_eq!( + responder.remote_static().unwrap(), + &initiator_keypair.public_key() + ); + + // Responder learned initiator's epoch from msg3 + assert_eq!(responder.remote_epoch(), Some(initiator_epoch)); + + // Handshake hashes should match + assert_eq!(initiator.handshake_hash(), responder.handshake_hash()); + + // Convert to sessions + let mut initiator_session = initiator.into_session().unwrap(); + let mut responder_session = responder.into_session().unwrap(); + + // Test bidirectional encryption + let plaintext = b"Hello via XK!"; + let ciphertext = initiator_session.encrypt(plaintext).unwrap(); + let decrypted = responder_session.decrypt(&ciphertext).unwrap(); + assert_eq!(decrypted, plaintext); + + let plaintext2 = b"XK reply!"; + let ciphertext2 = responder_session.encrypt(plaintext2).unwrap(); + let decrypted2 = initiator_session.decrypt(&ciphertext2).unwrap(); + assert_eq!(decrypted2, plaintext2); +} + +#[test] +fn test_xk_message_sizes() { + assert_eq!(XK_HANDSHAKE_MSG1_SIZE, 33); // ephemeral only + assert_eq!(XK_HANDSHAKE_MSG2_SIZE, 33 + 24); // ephemeral + encrypted epoch + assert_eq!(XK_HANDSHAKE_MSG3_SIZE, 33 + 16 + 24); // encrypted static + encrypted epoch +} + +#[test] +fn test_xk_identity_timing() { + // XK property: responder doesn't learn initiator identity until msg3 + let initiator_keypair = generate_keypair(); + let responder_keypair = generate_keypair(); + + let mut initiator = HandshakeState::new_xk_initiator(initiator_keypair, responder_keypair.public_key()); + initiator.set_local_epoch(generate_epoch()); + let mut responder = HandshakeState::new_xk_responder(responder_keypair); + responder.set_local_epoch(generate_epoch()); + + // Before any messages + assert!(responder.remote_static().is_none()); + + // After msg1 + let msg1 = initiator.write_xk_message_1().unwrap(); + responder.read_xk_message_1(&msg1).unwrap(); + assert!(responder.remote_static().is_none(), "XK: responder should NOT know identity after msg1"); + + // After msg2 + let msg2 = responder.write_xk_message_2().unwrap(); + initiator.read_xk_message_2(&msg2).unwrap(); + assert!(responder.remote_static().is_none(), "XK: responder should NOT know identity after msg2"); + + // After msg3 + let msg3 = initiator.write_xk_message_3().unwrap(); + responder.read_xk_message_3(&msg3).unwrap(); + assert!(responder.remote_static().is_some(), "XK: responder should know identity after msg3"); + assert_eq!(responder.remote_static().unwrap(), &initiator_keypair.public_key()); +} + +#[test] +fn test_xk_wrong_state_errors() { + let keypair1 = generate_keypair(); + let keypair2 = generate_keypair(); + + // Initiator can't read XK msg1 + let mut initiator = HandshakeState::new_xk_initiator(keypair1, keypair2.public_key()); + initiator.set_local_epoch(generate_epoch()); + assert!(initiator.read_xk_message_1(&[0u8; XK_HANDSHAKE_MSG1_SIZE]).is_err()); + + // Initiator can't write msg2 + assert!(initiator.write_xk_message_2().is_err()); + + // Initiator can't write msg3 before msg2 + assert!(initiator.write_xk_message_3().is_err()); + + // Responder can't write msg1 + let mut responder = HandshakeState::new_xk_responder(keypair2); + responder.set_local_epoch(generate_epoch()); + assert!(responder.write_xk_message_1().is_err()); + + // Responder can't read msg3 before msg2 + assert!(responder.read_xk_message_3(&[0u8; XK_HANDSHAKE_MSG3_SIZE]).is_err()); +} + +#[test] +fn test_xk_handshake_hash_differs_from_ik() { + // XK and IK should produce different handshake hashes (different protocol names) + let keypair1 = generate_keypair(); + let keypair2 = generate_keypair(); + let epoch1 = generate_epoch(); + let epoch2 = generate_epoch(); + + // Complete an IK handshake + let mut ik_init = HandshakeState::new_initiator(keypair1, keypair2.public_key()); + ik_init.set_local_epoch(epoch1); + let mut ik_resp = HandshakeState::new_responder(keypair2); + ik_resp.set_local_epoch(epoch2); + let msg1 = ik_init.write_message_1().unwrap(); + ik_resp.read_message_1(&msg1).unwrap(); + let msg2 = ik_resp.write_message_2().unwrap(); + ik_init.read_message_2(&msg2).unwrap(); + let ik_hash = ik_init.handshake_hash(); + + // Complete an XK handshake with the same keys + let mut xk_init = HandshakeState::new_xk_initiator(keypair1, keypair2.public_key()); + xk_init.set_local_epoch(epoch1); + let mut xk_resp = HandshakeState::new_xk_responder(keypair2); + xk_resp.set_local_epoch(epoch2); + let msg1 = xk_init.write_xk_message_1().unwrap(); + xk_resp.read_xk_message_1(&msg1).unwrap(); + let msg2 = xk_resp.write_xk_message_2().unwrap(); + xk_init.read_xk_message_2(&msg2).unwrap(); + let msg3 = xk_init.write_xk_message_3().unwrap(); + xk_resp.read_xk_message_3(&msg3).unwrap(); + let xk_hash = xk_init.handshake_hash(); + + assert_ne!(ik_hash, xk_hash, "IK and XK should produce different handshake hashes"); +} + +#[test] +fn test_xk_multiple_messages_after_handshake() { + let keypair1 = generate_keypair(); + let keypair2 = generate_keypair(); + + let mut initiator = HandshakeState::new_xk_initiator(keypair1, keypair2.public_key()); + initiator.set_local_epoch(generate_epoch()); + let mut responder = HandshakeState::new_xk_responder(keypair2); + responder.set_local_epoch(generate_epoch()); + + let msg1 = initiator.write_xk_message_1().unwrap(); + responder.read_xk_message_1(&msg1).unwrap(); + let msg2 = responder.write_xk_message_2().unwrap(); + initiator.read_xk_message_2(&msg2).unwrap(); + let msg3 = initiator.write_xk_message_3().unwrap(); + responder.read_xk_message_3(&msg3).unwrap(); + + let mut init_session = initiator.into_session().unwrap(); + let mut resp_session = responder.into_session().unwrap(); + + // Send many messages + for i in 0..100 { + let msg = format!("XK message {}", i); + let ct = init_session.encrypt(msg.as_bytes()).unwrap(); + let pt = resp_session.decrypt(&ct).unwrap(); + assert_eq!(pt, msg.as_bytes()); + } + + assert_eq!(init_session.send_nonce(), 100); + assert_eq!(resp_session.recv_nonce(), 100); +} + +#[test] +fn test_xk_with_odd_parity_responder() { + let secp = secp256k1::Secp256k1::new(); + + // Node B (responder) - odd parity key + let sk_b = secp256k1::SecretKey::from_slice( + &hex::decode("b102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1fb0") + .unwrap(), + ) + .unwrap(); + let kp_b = secp256k1::Keypair::from_secret_key(&secp, &sk_b); + let (xonly_b, parity_b) = kp_b.public_key().x_only_public_key(); + assert_eq!(parity_b, Parity::Odd, "Test requires odd-parity responder key"); + + // Node A (initiator) + let sk_a = secp256k1::SecretKey::from_slice( + &hex::decode("0102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1f20") + .unwrap(), + ) + .unwrap(); + let kp_a = secp256k1::Keypair::from_secret_key(&secp, &sk_a); + + // Simulate npub path: x-only → assumed even parity + let assumed_even_b = xonly_b.public_key(Parity::Even); + + let mut initiator = HandshakeState::new_xk_initiator(kp_a, assumed_even_b); + initiator.set_local_epoch(generate_epoch()); + let mut responder = HandshakeState::new_xk_responder(kp_b); + responder.set_local_epoch(generate_epoch()); + + let msg1 = initiator.write_xk_message_1().unwrap(); + responder.read_xk_message_1(&msg1).unwrap(); + let msg2 = responder.write_xk_message_2().unwrap(); + initiator.read_xk_message_2(&msg2).unwrap(); + let msg3 = initiator.write_xk_message_3().unwrap(); + responder.read_xk_message_3(&msg3).unwrap(); + + assert!(initiator.is_complete()); + assert!(responder.is_complete()); + + let mut sender = initiator.into_session().unwrap(); + let mut receiver = responder.into_session().unwrap(); + + let counter = sender.current_send_counter(); + let ciphertext = sender.encrypt(b"xk parity test").unwrap(); + let plaintext = receiver.decrypt_with_replay_check(&ciphertext, counter).unwrap(); + assert_eq!(plaintext, b"xk parity test"); +} + +#[test] +fn test_xk_invalid_msg1_size() { + let keypair = generate_keypair(); + let mut responder = HandshakeState::new_xk_responder(keypair); + responder.set_local_epoch(generate_epoch()); + + // Wrong size (IK msg1 size instead of XK) + assert!(responder.read_xk_message_1(&[0u8; HANDSHAKE_MSG1_SIZE]).is_err()); + // Too short + assert!(responder.read_xk_message_1(&[0u8; 10]).is_err()); +} + +#[test] +fn test_xk_invalid_msg3_size() { + let keypair1 = generate_keypair(); + let keypair2 = generate_keypair(); + + let mut initiator = HandshakeState::new_xk_initiator(keypair1, keypair2.public_key()); + initiator.set_local_epoch(generate_epoch()); + let mut responder = HandshakeState::new_xk_responder(keypair2); + responder.set_local_epoch(generate_epoch()); + + let msg1 = initiator.write_xk_message_1().unwrap(); + responder.read_xk_message_1(&msg1).unwrap(); + let _msg2 = responder.write_xk_message_2().unwrap(); + + // Responder is now in Message2Done, try wrong-size msg3 + assert!(responder.read_xk_message_3(&[0u8; 10]).is_err()); + assert!(responder.read_xk_message_3(&[0u8; XK_HANDSHAKE_MSG3_SIZE + 1]).is_err()); +} diff --git a/src/protocol/mod.rs b/src/protocol/mod.rs index 3600fc0..afe5e94 100644 --- a/src/protocol/mod.rs +++ b/src/protocol/mod.rs @@ -38,9 +38,9 @@ pub use filter::FilterAnnounce; pub use discovery::{LookupRequest, LookupResponse}; pub use session::{ CoordsRequired, FspFlags, FspInnerFlags, MtuExceeded, PathBroken, PathMtuNotification, - SessionAck, SessionFlags, SessionMessageType, SessionReceiverReport, SessionSenderReport, - SessionSetup, COORDS_REQUIRED_SIZE, MTU_EXCEEDED_SIZE, PATH_MTU_NOTIFICATION_SIZE, - SESSION_RECEIVER_REPORT_SIZE, SESSION_SENDER_REPORT_SIZE, + SessionAck, SessionFlags, SessionMessageType, SessionMsg3, SessionReceiverReport, + SessionSenderReport, SessionSetup, COORDS_REQUIRED_SIZE, MTU_EXCEEDED_SIZE, + PATH_MTU_NOTIFICATION_SIZE, SESSION_RECEIVER_REPORT_SIZE, SESSION_SENDER_REPORT_SIZE, }; pub(crate) use session::{coords_wire_size, decode_optional_coords, encode_coords}; diff --git a/src/protocol/session.rs b/src/protocol/session.rs index 71cbf2d..fbfcfcf 100644 --- a/src/protocol/session.rs +++ b/src/protocol/session.rs @@ -563,6 +563,97 @@ impl SessionAck { } } +// ============================================================================ +// Session Msg3 (XK Handshake Message 3) +// ============================================================================ + +/// XK handshake message 3 (initiator -> responder). +/// +/// Carries the initiator's encrypted static key and epoch. Sent by the +/// initiator after receiving msg2. The responder learns the initiator's +/// identity from this message. +/// +/// ## Wire Format +/// +/// | Offset | Field | Size | Description | +/// |--------|------------------|---------|-------------------------------------| +/// | 0 | flags | 1 byte | Reserved | +/// | 1 | handshake_len | 2 bytes | u16 LE, Noise payload length | +/// | 3 | handshake_payload| variable| Noise XK msg3 (73 bytes typical) | +#[derive(Clone, Debug)] +pub struct SessionMsg3 { + /// Reserved flags byte. + pub flags: u8, + /// Noise XK handshake message 3. + pub handshake_payload: Vec, +} + +impl SessionMsg3 { + /// Create a new SessionMsg3 with the given handshake payload. + pub fn new(handshake_payload: Vec) -> Self { + Self { + flags: 0, + handshake_payload, + } + } + + /// Encode as wire format (4-byte FSP prefix + flags + handshake). + /// + /// The 4-byte prefix: `[ver_phase:1][flags:1][payload_len:2 LE]` + /// where ver_phase = 0x03 (version 0, phase MSG3). + pub fn encode(&self) -> Vec { + // Build body first to compute payload_len + let mut body = Vec::new(); + body.push(self.flags); + let hs_len = self.handshake_payload.len() as u16; + body.extend_from_slice(&hs_len.to_le_bytes()); + body.extend_from_slice(&self.handshake_payload); + + // Prepend 4-byte FSP common prefix + let payload_len = body.len() as u16; + let mut buf = Vec::with_capacity(4 + body.len()); + buf.push(0x03); // version 0, phase 0x3 (MSG3) + buf.push(0x00); // flags (must be zero for handshake) + buf.extend_from_slice(&payload_len.to_le_bytes()); + buf.extend_from_slice(&body); + buf + } + + /// Decode from wire format (after 4-byte FSP prefix has been consumed). + pub fn decode(payload: &[u8]) -> Result { + if payload.is_empty() { + return Err(ProtocolError::MessageTooShort { + expected: 1, + got: 0, + }); + } + let flags = payload[0]; + let mut offset = 1; + + if payload.len() < offset + 2 { + return Err(ProtocolError::MessageTooShort { + expected: offset + 2, + got: payload.len(), + }); + } + let hs_len = u16::from_le_bytes([payload[offset], payload[offset + 1]]) as usize; + offset += 2; + + if payload.len() < offset + hs_len { + return Err(ProtocolError::MessageTooShort { + expected: offset + hs_len, + got: payload.len(), + }); + } + let handshake_payload = payload[offset..offset + hs_len].to_vec(); + + Ok(Self { + flags, + handshake_payload, + }) + } +} + // ============================================================================ // Session-Layer MMP Reports // ============================================================================ @@ -1549,4 +1640,38 @@ mod tests { fn test_mtu_exceeded_display() { assert_eq!(format!("{}", SessionMessageType::MtuExceeded), "MtuExceeded"); } + + // ===== SessionMsg3 Tests ===== + + #[test] + fn test_session_msg3_encode_decode() { + let handshake = vec![0xCC; 73]; // typical XK msg3 + let msg3 = SessionMsg3::new(handshake.clone()); + + let encoded = msg3.encode(); + // Verify FSP prefix: ver_phase=0x03 (version 0, phase MSG3) + assert_eq!(encoded[0], 0x03); + assert_eq!(encoded[1], 0x00); // flags = 0 for handshake + let payload_len = u16::from_le_bytes([encoded[2], encoded[3]]); + assert_eq!(payload_len as usize, encoded.len() - 4); + + // Decode (skip 4-byte FSP prefix) + let decoded = SessionMsg3::decode(&encoded[4..]).unwrap(); + assert_eq!(decoded.flags, 0); + assert_eq!(decoded.handshake_payload, handshake); + } + + #[test] + fn test_session_msg3_decode_too_short() { + assert!(SessionMsg3::decode(&[]).is_err()); + assert!(SessionMsg3::decode(&[0x00]).is_err()); // flags only, no hs_len + } + + #[test] + fn test_session_msg3_empty_handshake() { + let msg3 = SessionMsg3::new(vec![]); + let encoded = msg3.encode(); + let decoded = SessionMsg3::decode(&encoded[4..]).unwrap(); + assert!(decoded.handshake_payload.is_empty()); + } }