Test that a crossed TCP handshake converges on one pair of keys

When two nodes dial each other over TCP and the smaller node completes
its inbound handshake before its outbound msg2 arrives, handle_msg2
queues msg3 on the outbound connection, swaps to the outbound session
and closes that same connection as the losing one, all in one call. The
send only queues msg3 for the connection's writer. Before the stream
writer drain, the close aborted the writer, msg3 was never written, and
the two ends kept different sessions while each reported a connected
peer.

The existing writer drain tests are transport-level, and nothing on
this line drove the crossed order at node level. This test holds the
handshake packets and delivers them in the order the loss was observed
in, then checks that each end sends on the index the other receives on
and that a frame from each end decrypts at the other.
This commit is contained in:
Johnathan Corgan
2026-10-01 17:34:55 +00:00
parent e9970a8b06
commit 12bed7ac5b
+186
View File
@@ -499,3 +499,189 @@ async fn api_disconnect_delivers_the_disconnect_before_closing() {
cleanup_nodes(&mut nodes).await;
}
/// Wait for the next handshake packet of `phase` at `tn`, holding back any
/// other packet that arrives first in `held`.
///
/// Returns `None` if no such packet arrives within `wait`, so a caller can
/// assert on what the node did without it rather than on the timeout.
async fn next_of_phase(
tn: &mut TestNode,
held: &mut Vec<ReceivedPacket>,
phase: u8,
wait: Duration,
) -> Option<ReceivedPacket> {
use crate::proto::fmp::wire::CommonPrefix;
let is_phase =
|p: &ReceivedPacket| CommonPrefix::parse(&p.data).map(|c| c.phase) == Some(phase);
if let Some(pos) = held.iter().position(is_phase) {
return Some(held.remove(pos));
}
let deadline = tokio::time::Instant::now() + wait;
loop {
let packet = tokio::time::timeout_at(deadline, tn.packet_rx.recv())
.await
.ok()??;
if is_phase(&packet) {
return Some(packet);
}
held.push(packet);
}
}
/// Hand a held packet to the node handler its phase selects.
async fn dispatch_packet(tn: &mut TestNode, packet: ReceivedPacket) {
use crate::proto::fmp::wire::{
CommonPrefix, PHASE_ESTABLISHED, PHASE_MSG1, PHASE_MSG2, PHASE_MSG3,
};
match CommonPrefix::parse(&packet.data).map(|c| c.phase) {
Some(PHASE_MSG1) => tn.node.handle_msg1(packet).await,
Some(PHASE_MSG2) => tn.node.handle_msg2(packet).await,
Some(PHASE_MSG3) => tn.node.handle_msg3(packet).await,
Some(PHASE_ESTABLISHED) => tn.node.handle_encrypted_frame(packet).await,
_ => {}
}
}
/// A crossed TCP handshake converges on one pair of keys when the winning
/// side queues its msg3 and closes that connection in the same call.
///
/// Both nodes dial at once. The smaller node completes its inbound handshake
/// first, then its outbound msg2 arrives: `handle_msg2` sends msg3 on the
/// outbound connection, swaps to the outbound session and closes that same
/// connection as the losing one. The send only queues msg3 for the
/// connection's writer, so the close must let the writer finish it; if the
/// close discards it, the larger node never swaps and the two ends keep
/// different sessions while each reports a connected peer.
#[tokio::test]
async fn crossed_tcp_handshake_converges_when_the_losing_connection_closes_after_msg3() {
use crate::proto::fmp::wire::{PHASE_MSG1, PHASE_MSG2, PHASE_MSG3};
use crate::proto::link::LinkMessageType;
const WAIT: Duration = Duration::from_secs(2);
let mut nodes = vec![make_test_node_tcp().await, make_test_node_tcp().await];
let (a, b) = if nodes[0].node.node_addr() < nodes[1].node.node_addr() {
(0, 1)
} else {
(1, 0)
};
let addr_a = *nodes[a].node.node_addr();
let addr_b = *nodes[b].node.node_addr();
let mut held: [Vec<ReceivedPacket>; 2] = Default::default();
// Both dial: B's connection reaches A's listener, A's reaches B's.
initiate_handshake(&mut nodes, b, a).await;
initiate_handshake(&mut nodes, a, b).await;
// A answers B's msg1 on B's connection.
let p = next_of_phase(&mut nodes[a], &mut held[a], PHASE_MSG1, WAIT)
.await
.expect("A should receive B's msg1");
nodes[a].node.handle_msg1(p).await;
// B completes its outbound handshake: sends msg3 and promotes A.
let p = next_of_phase(&mut nodes[b], &mut held[b], PHASE_MSG2, WAIT)
.await
.expect("B should receive A's msg2");
nodes[b].node.handle_msg2(p).await;
assert_eq!(nodes[b].node.peer_count(), 1, "B should promote A");
// B answers A's msg1 on A's connection.
let p = next_of_phase(&mut nodes[b], &mut held[b], PHASE_MSG1, WAIT)
.await
.expect("B should receive A's msg1");
nodes[b].node.handle_msg1(p).await;
// A completes its inbound handshake with B's msg3 and promotes B.
let p = next_of_phase(&mut nodes[a], &mut held[a], PHASE_MSG3, WAIT)
.await
.expect("A should receive B's msg3");
nodes[a].node.handle_msg3(p).await;
assert_eq!(nodes[a].node.peer_count(), 1, "A should promote B");
// A's outbound msg2 arrives second: in one call A queues msg3 on its
// outbound connection, swaps to the outbound session, and closes that
// connection as the loser.
let p = next_of_phase(&mut nodes[a], &mut held[a], PHASE_MSG2, WAIT)
.await
.expect("A should receive B's msg2");
nodes[a].node.handle_msg2(p).await;
// B processes A's msg3 if it was written before the close.
if let Some(p) = next_of_phase(&mut nodes[b], &mut held[b], PHASE_MSG3, WAIT).await {
nodes[b].node.handle_msg3(p).await;
}
// Settle whatever else was sent during the handshake.
for (tn, held) in nodes.iter_mut().zip(held.iter_mut()) {
for packet in held.drain(..) {
dispatch_packet(tn, packet).await;
}
}
tokio::time::sleep(Duration::from_millis(50)).await;
spanning_tree::process_available_packets(&mut nodes).await;
assert_eq!(nodes[a].node.peer_count(), 1, "A should hold one peer");
assert_eq!(nodes[b].node.peer_count(), 1, "B should hold one peer");
let peer_b_on_a = nodes[a].node.get_peer(&addr_b).expect("A should have B");
let peer_a_on_b = nodes[b].node.get_peer(&addr_a).expect("B should have A");
assert!(
peer_b_on_a.our_index().is_some() && peer_b_on_a.their_index().is_some(),
"both indices must be set, or the equalities below pass on None == None"
);
assert_eq!(
peer_b_on_a.their_index(),
peer_a_on_b.our_index(),
"A sends to B on an index B does not receive on: the ends diverged"
);
assert_eq!(
peer_a_on_b.their_index(),
peer_b_on_a.our_index(),
"B sends to A on an index A does not receive on: the ends diverged"
);
// A frame encrypted by each end must decrypt at the other.
let recv_before = [
peer_b_on_a.link_stats().packets_recv,
peer_a_on_b.link_stats().packets_recv,
];
let heartbeat = [LinkMessageType::Heartbeat as u8];
nodes[a]
.node
.send_encrypted_link_message(&addr_b, &heartbeat)
.await
.expect("A should send to B");
nodes[b]
.node
.send_encrypted_link_message(&addr_a, &heartbeat)
.await
.expect("B should send to A");
tokio::time::sleep(Duration::from_millis(50)).await;
spanning_tree::process_available_packets(&mut nodes).await;
let recv_b_on_a = nodes[a]
.node
.get_peer(&addr_b)
.expect("A should still have B")
.link_stats()
.packets_recv;
let recv_a_on_b = nodes[b]
.node
.get_peer(&addr_a)
.expect("B should still have A")
.link_stats()
.packets_recv;
assert!(
recv_b_on_a > recv_before[0],
"A decrypted no frame from B after the handshake"
);
assert!(
recv_a_on_b > recv_before[1],
"B decrypted no frame from A after the handshake"
);
cleanup_nodes(&mut nodes).await;
}