mirror of
https://github.com/EasyTier/EasyTier.git
synced 2026-09-04 18:15:39 +00:00
refactor(core): migrate packet processing from pnet to smoltcp (#2456)
Replace pnet_packet parsing and mutation across gateway packet paths with the existing smoltcp wire APIs. Preserve length validation, fragmentation classification, TCP flags, and checksum behavior while removing the core pnet_packet feature dependency. Reject stale non-initiator OSPF sync sessions: only initiator requests may create missing sessions, and a rejection clears the old initiator role only when the remote session generation is unchanged. This fixes an unbounded RPC storm caused by a delayed route sync recreating a session after both peers relinquished the initiator role, with regression tests for session creation and response reordering.
This commit is contained in:
@@ -23,7 +23,7 @@ use std::{
|
||||
time::Duration,
|
||||
};
|
||||
|
||||
use pnet_packet::{Packet, ip::IpNextHeaderProtocols, ipv4::Ipv4Packet, tcp::TcpPacket};
|
||||
use smoltcp::wire::{IpProtocol, Ipv4Packet, TcpPacket};
|
||||
use tokio::{
|
||||
select,
|
||||
sync::{Mutex, mpsc},
|
||||
@@ -74,7 +74,7 @@ use self::{
|
||||
deadline::{DataPlaneDeadline, DataPlaneIoDeadline},
|
||||
error::DataPlaneResult,
|
||||
flow::{FlowKey, FlowKind, FlowLease, FlowTable},
|
||||
packet::PeerPacketRoute,
|
||||
packet::{PeerPacketRoute, tcp_flags},
|
||||
resource::{DataPlaneConsumers, DataPlaneIoGuard, DataPlaneLease},
|
||||
route::{
|
||||
DataPlaneRoutePolicy, DataPlaneTcpRoute, DataPlaneTcpRouteInput,
|
||||
@@ -218,15 +218,16 @@ where
|
||||
|| x == PacketType::DataWithQuicSrcModified as u8
|
||||
)
|
||||
{
|
||||
if let Some(ipv4) = Ipv4Packet::new(packet.payload()) {
|
||||
if let Ok(ipv4) = Ipv4Packet::new_checked(packet.payload()) {
|
||||
let (tcp_src_port, tcp_dst_port, tcp_flags) =
|
||||
if ipv4.get_next_level_protocol() == IpNextHeaderProtocols::Tcp {
|
||||
TcpPacket::new(ipv4.payload())
|
||||
if ipv4.next_header() == IpProtocol::Tcp {
|
||||
TcpPacket::new_checked(ipv4.payload())
|
||||
.ok()
|
||||
.map(|tcp| {
|
||||
(
|
||||
Some(tcp.get_source()),
|
||||
Some(tcp.get_destination()),
|
||||
Some(tcp.get_flags()),
|
||||
Some(tcp.src_port()),
|
||||
Some(tcp.dst_port()),
|
||||
Some(tcp_flags(&tcp)),
|
||||
)
|
||||
})
|
||||
.unwrap_or((None, None, None))
|
||||
@@ -237,9 +238,9 @@ where
|
||||
packet_type = hdr.packet_type,
|
||||
from_peer_id = hdr.from_peer_id.get(),
|
||||
to_peer_id = hdr.to_peer_id.get(),
|
||||
ipv4_src = %ipv4.get_source(),
|
||||
ipv4_dst = %ipv4.get_destination(),
|
||||
next_protocol = ?ipv4.get_next_level_protocol(),
|
||||
ipv4_src = %ipv4.src_addr(),
|
||||
ipv4_dst = %ipv4.dst_addr(),
|
||||
next_protocol = ?ipv4.next_header(),
|
||||
?tcp_src_port,
|
||||
?tcp_dst_port,
|
||||
?tcp_flags,
|
||||
|
||||
@@ -2,12 +2,10 @@
|
||||
|
||||
use std::net::{IpAddr, Ipv4Addr, SocketAddr};
|
||||
|
||||
use pnet_packet::{
|
||||
Packet, ip::IpNextHeaderProtocols, ipv4::Ipv4Packet, tcp::TcpPacket, udp::UdpPacket,
|
||||
};
|
||||
use smoltcp::wire::{IPV4_HEADER_LEN, IpProtocol, Ipv4Packet, TcpPacket, UdpPacket};
|
||||
|
||||
use crate::{
|
||||
gateway::proxy::ip_reassembler::{IpReassembler, SmolIpv4Packet},
|
||||
gateway::proxy::ip_reassembler::IpReassembler,
|
||||
packet::{PacketType, ZCPacket},
|
||||
};
|
||||
|
||||
@@ -45,21 +43,21 @@ pub(crate) enum PeerPacketRoute {
|
||||
},
|
||||
}
|
||||
fn classify_peer_ipv4_payload(payload: &[u8]) -> ClassifiedPeerPacket {
|
||||
let Some(ipv4) = Ipv4Packet::new(payload) else {
|
||||
let Ok(ipv4) = Ipv4Packet::new_checked(payload) else {
|
||||
return ClassifiedPeerPacket::Unsupported;
|
||||
};
|
||||
if ipv4.get_version() != 4 {
|
||||
if ipv4.version() != 4 || usize::from(ipv4.header_len()) < IPV4_HEADER_LEN {
|
||||
return ClassifiedPeerPacket::Unsupported;
|
||||
}
|
||||
|
||||
match ipv4.get_next_level_protocol() {
|
||||
IpNextHeaderProtocols::Tcp => {
|
||||
let Some(tcp) = TcpPacket::new(ipv4.payload()) else {
|
||||
match ipv4.next_header() {
|
||||
IpProtocol::Tcp => {
|
||||
let Ok(tcp) = TcpPacket::new_checked(ipv4.payload()) else {
|
||||
return ClassifiedPeerPacket::Unsupported;
|
||||
};
|
||||
let entry = FlowKey {
|
||||
dst: SocketAddr::new(ipv4.get_source().into(), tcp.get_source()),
|
||||
src: SocketAddr::new(ipv4.get_destination().into(), tcp.get_destination()),
|
||||
dst: SocketAddr::new(ipv4.src_addr().into(), tcp.src_port()),
|
||||
src: SocketAddr::new(ipv4.dst_addr().into(), tcp.dst_port()),
|
||||
kind: FlowKind::Tcp,
|
||||
};
|
||||
let listen_entry = FlowKey {
|
||||
@@ -70,23 +68,22 @@ fn classify_peer_ipv4_payload(payload: &[u8]) -> ClassifiedPeerPacket {
|
||||
ClassifiedPeerPacket::Tcp {
|
||||
entry,
|
||||
listen_entry,
|
||||
flags: tcp.get_flags(),
|
||||
flags: tcp_flags(&tcp),
|
||||
}
|
||||
}
|
||||
IpNextHeaderProtocols::Udp => {
|
||||
let smol_ipv4 = SmolIpv4Packet::new_unchecked(ipv4.packet());
|
||||
if IpReassembler::is_packet_fragmented(&smol_ipv4) {
|
||||
IpProtocol::Udp => {
|
||||
if IpReassembler::is_packet_fragmented(&ipv4) {
|
||||
return ClassifiedPeerPacket::FragmentedUdp {
|
||||
source: ipv4.get_source(),
|
||||
source: ipv4.src_addr(),
|
||||
};
|
||||
}
|
||||
let Some(udp) = UdpPacket::new(ipv4.payload()) else {
|
||||
let Ok(udp) = UdpPacket::new_checked(ipv4.payload()) else {
|
||||
return ClassifiedPeerPacket::Unsupported;
|
||||
};
|
||||
ClassifiedPeerPacket::Udp {
|
||||
entry: FlowKey {
|
||||
dst: SocketAddr::new(ipv4.get_source().into(), udp.get_source()),
|
||||
src: SocketAddr::new(ipv4.get_destination().into(), udp.get_destination()),
|
||||
dst: SocketAddr::new(ipv4.src_addr().into(), udp.src_port()),
|
||||
src: SocketAddr::new(ipv4.dst_addr().into(), udp.dst_port()),
|
||||
kind: FlowKind::Udp,
|
||||
},
|
||||
}
|
||||
@@ -94,6 +91,17 @@ fn classify_peer_ipv4_payload(payload: &[u8]) -> ClassifiedPeerPacket {
|
||||
_ => ClassifiedPeerPacket::Unsupported,
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn tcp_flags<T: AsRef<[u8]>>(tcp: &TcpPacket<T>) -> u8 {
|
||||
u8::from(tcp.fin())
|
||||
| (u8::from(tcp.syn()) << 1)
|
||||
| (u8::from(tcp.rst()) << 2)
|
||||
| (u8::from(tcp.psh()) << 3)
|
||||
| (u8::from(tcp.ack()) << 4)
|
||||
| (u8::from(tcp.urg()) << 5)
|
||||
| (u8::from(tcp.ece()) << 6)
|
||||
| (u8::from(tcp.cwr()) << 7)
|
||||
}
|
||||
impl<V> FlowTable<V> {
|
||||
pub fn route_peer_packet(
|
||||
&self,
|
||||
@@ -158,37 +166,32 @@ impl<V> FlowTable<V> {
|
||||
mod tests {
|
||||
use std::net::{IpAddr, Ipv4Addr, SocketAddr};
|
||||
|
||||
use pnet_packet::{
|
||||
MutablePacket,
|
||||
ip::IpNextHeaderProtocols,
|
||||
ipv4::MutableIpv4Packet,
|
||||
tcp::{MutableTcpPacket, TcpFlags},
|
||||
udp::MutableUdpPacket,
|
||||
};
|
||||
|
||||
use super::*;
|
||||
use crate::packet::{PacketType, ZCPacket};
|
||||
|
||||
fn ipv4_packet(protocol: pnet_packet::ip::IpNextHeaderProtocol, payload_len: usize) -> Vec<u8> {
|
||||
const TCP_SYN: u8 = 0x02;
|
||||
|
||||
fn ipv4_packet(protocol: IpProtocol, payload_len: usize) -> Vec<u8> {
|
||||
let mut packet = vec![0; 20 + payload_len];
|
||||
let packet_len = packet.len() as u16;
|
||||
let mut ipv4 = MutableIpv4Packet::new(&mut packet).unwrap();
|
||||
let mut ipv4 = Ipv4Packet::new_unchecked(&mut packet);
|
||||
ipv4.set_version(4);
|
||||
ipv4.set_header_length(5);
|
||||
ipv4.set_total_length(packet_len);
|
||||
ipv4.set_source(Ipv4Addr::new(10, 1, 1, 2));
|
||||
ipv4.set_destination(Ipv4Addr::new(10, 2, 2, 3));
|
||||
ipv4.set_next_level_protocol(protocol);
|
||||
ipv4.set_header_len(20);
|
||||
ipv4.set_total_len(packet_len);
|
||||
ipv4.set_src_addr(Ipv4Addr::new(10, 1, 1, 2));
|
||||
ipv4.set_dst_addr(Ipv4Addr::new(10, 2, 2, 3));
|
||||
ipv4.set_next_header(protocol);
|
||||
packet
|
||||
}
|
||||
#[test]
|
||||
fn classifies_tcp_and_listen_keys() {
|
||||
let mut packet = ipv4_packet(IpNextHeaderProtocols::Tcp, 20);
|
||||
let mut ipv4 = MutableIpv4Packet::new(&mut packet).unwrap();
|
||||
let mut tcp = MutableTcpPacket::new(ipv4.payload_mut()).unwrap();
|
||||
tcp.set_source(1234);
|
||||
tcp.set_destination(4321);
|
||||
tcp.set_flags(TcpFlags::SYN);
|
||||
let mut packet = ipv4_packet(IpProtocol::Tcp, 20);
|
||||
let mut ipv4 = Ipv4Packet::new_unchecked(&mut packet);
|
||||
let mut tcp = TcpPacket::new_unchecked(ipv4.payload_mut());
|
||||
tcp.set_src_port(1234);
|
||||
tcp.set_dst_port(4321);
|
||||
tcp.set_header_len(20);
|
||||
tcp.set_syn(true);
|
||||
|
||||
assert_eq!(
|
||||
classify_peer_ipv4_payload(&packet),
|
||||
@@ -203,17 +206,18 @@ mod tests {
|
||||
dst: SocketAddr::new(IpAddr::V4(Ipv4Addr::UNSPECIFIED), 0),
|
||||
kind: FlowKind::TcpListen,
|
||||
},
|
||||
flags: TcpFlags::SYN,
|
||||
flags: TCP_SYN,
|
||||
}
|
||||
);
|
||||
}
|
||||
#[test]
|
||||
fn classifies_udp_and_fragmented_udp() {
|
||||
let mut packet = ipv4_packet(IpNextHeaderProtocols::Udp, 8);
|
||||
let mut ipv4 = MutableIpv4Packet::new(&mut packet).unwrap();
|
||||
let mut udp = MutableUdpPacket::new(ipv4.payload_mut()).unwrap();
|
||||
udp.set_source(1234);
|
||||
udp.set_destination(4321);
|
||||
let mut packet = ipv4_packet(IpProtocol::Udp, 8);
|
||||
let mut ipv4 = Ipv4Packet::new_unchecked(&mut packet);
|
||||
let mut udp = UdpPacket::new_unchecked(ipv4.payload_mut());
|
||||
udp.set_src_port(1234);
|
||||
udp.set_dst_port(4321);
|
||||
udp.set_len(8);
|
||||
assert_eq!(
|
||||
classify_peer_ipv4_payload(&packet),
|
||||
ClassifiedPeerPacket::Udp {
|
||||
@@ -225,10 +229,8 @@ mod tests {
|
||||
}
|
||||
);
|
||||
|
||||
let mut fragmented = ipv4_packet(IpNextHeaderProtocols::Udp, 8);
|
||||
MutableIpv4Packet::new(&mut fragmented)
|
||||
.unwrap()
|
||||
.set_fragment_offset(1);
|
||||
let mut fragmented = ipv4_packet(IpProtocol::Udp, 8);
|
||||
Ipv4Packet::new_unchecked(&mut fragmented).set_frag_offset(8);
|
||||
assert_eq!(
|
||||
classify_peer_ipv4_payload(&fragmented),
|
||||
ClassifiedPeerPacket::FragmentedUdp {
|
||||
@@ -243,18 +245,35 @@ mod tests {
|
||||
ClassifiedPeerPacket::Unsupported
|
||||
);
|
||||
assert_eq!(
|
||||
classify_peer_ipv4_payload(&ipv4_packet(IpNextHeaderProtocols::Icmp, 8)),
|
||||
classify_peer_ipv4_payload(&ipv4_packet(IpProtocol::Icmp, 8)),
|
||||
ClassifiedPeerPacket::Unsupported
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_ipv4_header_shorter_than_minimum() {
|
||||
let mut packet = ipv4_packet(IpProtocol::Tcp, 20);
|
||||
let mut ipv4 = Ipv4Packet::new_unchecked(&mut packet);
|
||||
ipv4.set_header_len(16);
|
||||
let mut tcp = TcpPacket::new_unchecked(ipv4.payload_mut());
|
||||
tcp.set_src_port(1234);
|
||||
tcp.set_dst_port(4321);
|
||||
tcp.set_header_len(20);
|
||||
|
||||
assert_eq!(
|
||||
classify_peer_ipv4_payload(&packet),
|
||||
ClassifiedPeerPacket::Unsupported
|
||||
);
|
||||
}
|
||||
#[test]
|
||||
fn flow_table_routes_tcp_exact_and_listen_fallback() {
|
||||
let mut packet = ipv4_packet(IpNextHeaderProtocols::Tcp, 20);
|
||||
let mut ipv4 = MutableIpv4Packet::new(&mut packet).unwrap();
|
||||
let mut tcp = MutableTcpPacket::new(ipv4.payload_mut()).unwrap();
|
||||
tcp.set_source(1234);
|
||||
tcp.set_destination(4321);
|
||||
tcp.set_flags(TcpFlags::SYN);
|
||||
let mut packet = ipv4_packet(IpProtocol::Tcp, 20);
|
||||
let mut ipv4 = Ipv4Packet::new_unchecked(&mut packet);
|
||||
let mut tcp = TcpPacket::new_unchecked(ipv4.payload_mut());
|
||||
tcp.set_src_port(1234);
|
||||
tcp.set_dst_port(4321);
|
||||
tcp.set_header_len(20);
|
||||
tcp.set_syn(true);
|
||||
|
||||
let exact = FlowKey {
|
||||
src: "10.2.2.3:4321".parse().unwrap(),
|
||||
@@ -272,7 +291,7 @@ mod tests {
|
||||
table.route_peer_ipv4_payload(&packet, false),
|
||||
PeerPacketRoute::Unmatched {
|
||||
entry: exact.clone(),
|
||||
tcp_flags: Some(TcpFlags::SYN),
|
||||
tcp_flags: Some(TCP_SYN),
|
||||
}
|
||||
);
|
||||
|
||||
@@ -281,7 +300,7 @@ mod tests {
|
||||
table.route_peer_ipv4_payload(&packet, true),
|
||||
PeerPacketRoute::Deliver {
|
||||
entry: listen,
|
||||
tcp_flags: Some(TcpFlags::SYN),
|
||||
tcp_flags: Some(TCP_SYN),
|
||||
}
|
||||
);
|
||||
|
||||
@@ -290,16 +309,14 @@ mod tests {
|
||||
table.route_peer_ipv4_payload(&packet, true),
|
||||
PeerPacketRoute::Deliver {
|
||||
entry: exact,
|
||||
tcp_flags: Some(TcpFlags::SYN),
|
||||
tcp_flags: Some(TCP_SYN),
|
||||
}
|
||||
);
|
||||
}
|
||||
#[test]
|
||||
fn flow_table_routes_fragmented_udp_by_source_ip() {
|
||||
let mut packet = ipv4_packet(IpNextHeaderProtocols::Udp, 8);
|
||||
MutableIpv4Packet::new(&mut packet)
|
||||
.unwrap()
|
||||
.set_fragment_offset(1);
|
||||
let mut packet = ipv4_packet(IpProtocol::Udp, 8);
|
||||
Ipv4Packet::new_unchecked(&mut packet).set_frag_offset(8);
|
||||
let table = FlowTable::default();
|
||||
|
||||
assert_eq!(
|
||||
@@ -328,11 +345,12 @@ mod tests {
|
||||
}
|
||||
#[test]
|
||||
fn flow_table_routes_loopback_modified_source_packets() {
|
||||
let mut payload = ipv4_packet(IpNextHeaderProtocols::Tcp, 20);
|
||||
let mut ipv4 = MutableIpv4Packet::new(&mut payload).unwrap();
|
||||
let mut tcp = MutableTcpPacket::new(ipv4.payload_mut()).unwrap();
|
||||
tcp.set_source(1234);
|
||||
tcp.set_destination(4321);
|
||||
let mut payload = ipv4_packet(IpProtocol::Tcp, 20);
|
||||
let mut ipv4 = Ipv4Packet::new_unchecked(&mut payload);
|
||||
let mut tcp = TcpPacket::new_unchecked(ipv4.payload_mut());
|
||||
tcp.set_src_port(1234);
|
||||
tcp.set_dst_port(4321);
|
||||
tcp.set_header_len(20);
|
||||
let entry = FlowKey {
|
||||
src: "10.2.2.3:4321".parse().unwrap(),
|
||||
dst: "10.1.1.2:1234".parse().unwrap(),
|
||||
@@ -359,8 +377,7 @@ mod tests {
|
||||
#[test]
|
||||
fn flow_table_passes_non_loopback_or_malformed_modified_source_packets() {
|
||||
let table = FlowTable::<()>::default();
|
||||
let mut non_loopback =
|
||||
ZCPacket::new_with_payload(&ipv4_packet(IpNextHeaderProtocols::Tcp, 20));
|
||||
let mut non_loopback = ZCPacket::new_with_payload(&ipv4_packet(IpProtocol::Tcp, 20));
|
||||
non_loopback.fill_peer_manager_hdr(7, 8, PacketType::DataWithKcpSrcModified as u8);
|
||||
assert_eq!(
|
||||
table.route_peer_packet(&non_loopback, false),
|
||||
|
||||
@@ -5,7 +5,7 @@ use std::{
|
||||
sync::{Arc, Weak},
|
||||
};
|
||||
|
||||
use pnet_packet::ipv4::Ipv4Packet;
|
||||
use smoltcp::wire::Ipv4Packet;
|
||||
use tokio::{
|
||||
sync::{Mutex, mpsc},
|
||||
task::JoinSet,
|
||||
@@ -54,11 +54,11 @@ impl SmoltcpPlane {
|
||||
|
||||
forward_tasks.spawn(async move {
|
||||
while let Some(data) = stack_stream.recv().await {
|
||||
let Some(ipv4) = Ipv4Packet::new(&data) else {
|
||||
let Ok(ipv4) = Ipv4Packet::new_checked(&data) else {
|
||||
tracing::error!(?data, "smoltcp emitted a non-IPv4 packet");
|
||||
continue;
|
||||
};
|
||||
let destination = ipv4.get_destination();
|
||||
let destination = ipv4.dst_addr();
|
||||
let Some(peer_manager) = peer_manager.upgrade() else {
|
||||
tracing::debug!("smoltcp-to-peer bridge lost PeerManager");
|
||||
return;
|
||||
|
||||
@@ -1,11 +1,6 @@
|
||||
use std::net::{IpAddr, Ipv4Addr, SocketAddr};
|
||||
|
||||
use pnet_packet::{
|
||||
MutablePacket,
|
||||
ip::IpNextHeaderProtocols,
|
||||
ipv4::{self, MutableIpv4Packet},
|
||||
tcp::{self, MutableTcpPacket, TcpFlags},
|
||||
};
|
||||
use smoltcp::wire::{IpAddress, IpProtocol, Ipv4Packet, TcpPacket};
|
||||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
|
||||
use super::*;
|
||||
@@ -177,28 +172,25 @@ fn build_tcp_packet(src: SocketAddr, dst: SocketAddr) -> Vec<u8> {
|
||||
};
|
||||
|
||||
{
|
||||
let mut ip_packet = MutableIpv4Packet::new(&mut buf).unwrap();
|
||||
let mut ip_packet = Ipv4Packet::new_unchecked(&mut buf);
|
||||
ip_packet.set_version(4);
|
||||
ip_packet.set_header_length(5);
|
||||
ip_packet.set_total_length(40);
|
||||
ip_packet.set_ttl(64);
|
||||
ip_packet.set_next_level_protocol(IpNextHeaderProtocols::Tcp);
|
||||
ip_packet.set_source(src_ip);
|
||||
ip_packet.set_destination(dst_ip);
|
||||
ip_packet.set_header_len(20);
|
||||
ip_packet.set_total_len(40);
|
||||
ip_packet.set_hop_limit(64);
|
||||
ip_packet.set_next_header(IpProtocol::Tcp);
|
||||
ip_packet.set_src_addr(src_ip);
|
||||
ip_packet.set_dst_addr(dst_ip);
|
||||
|
||||
let mut tcp_packet = MutableTcpPacket::new(ip_packet.payload_mut()).unwrap();
|
||||
tcp_packet.set_source(src.port());
|
||||
tcp_packet.set_destination(dst.port());
|
||||
tcp_packet.set_data_offset(5);
|
||||
tcp_packet.set_flags(TcpFlags::SYN | TcpFlags::ACK);
|
||||
tcp_packet.set_window(65535);
|
||||
tcp_packet.set_checksum(tcp::ipv4_checksum(
|
||||
&tcp_packet.to_immutable(),
|
||||
&src_ip,
|
||||
&dst_ip,
|
||||
));
|
||||
let mut tcp_packet = TcpPacket::new_unchecked(ip_packet.payload_mut());
|
||||
tcp_packet.set_src_port(src.port());
|
||||
tcp_packet.set_dst_port(dst.port());
|
||||
tcp_packet.set_header_len(20);
|
||||
tcp_packet.set_syn(true);
|
||||
tcp_packet.set_ack(true);
|
||||
tcp_packet.set_window_len(65535);
|
||||
tcp_packet.fill_checksum(&IpAddress::Ipv4(src_ip), &IpAddress::Ipv4(dst_ip));
|
||||
|
||||
ip_packet.set_checksum(ipv4::checksum(&ip_packet.to_immutable()));
|
||||
ip_packet.fill_checksum();
|
||||
}
|
||||
|
||||
buf
|
||||
@@ -207,20 +199,20 @@ fn build_tcp_packet(src: SocketAddr, dst: SocketAddr) -> Vec<u8> {
|
||||
fn build_udp_followup_fragment(src: Ipv4Addr, dst: Ipv4Addr) -> Vec<u8> {
|
||||
let mut buf = vec![0u8; 28];
|
||||
{
|
||||
let mut ip_packet = MutableIpv4Packet::new(&mut buf).unwrap();
|
||||
let mut ip_packet = Ipv4Packet::new_unchecked(&mut buf);
|
||||
ip_packet.set_version(4);
|
||||
ip_packet.set_header_length(5);
|
||||
ip_packet.set_total_length(28);
|
||||
ip_packet.set_ttl(64);
|
||||
ip_packet.set_next_level_protocol(IpNextHeaderProtocols::Udp);
|
||||
ip_packet.set_fragment_offset(1);
|
||||
ip_packet.set_source(src);
|
||||
ip_packet.set_destination(dst);
|
||||
ip_packet.set_header_len(20);
|
||||
ip_packet.set_total_len(28);
|
||||
ip_packet.set_hop_limit(64);
|
||||
ip_packet.set_next_header(IpProtocol::Udp);
|
||||
ip_packet.set_frag_offset(8);
|
||||
ip_packet.set_src_addr(src);
|
||||
ip_packet.set_dst_addr(dst);
|
||||
ip_packet
|
||||
.payload_mut()
|
||||
.copy_from_slice(&[0xde, 0xad, 0xbe, 0xef, 0xca, 0xfe, 0xba, 0xbe]);
|
||||
|
||||
ip_packet.set_checksum(ipv4::checksum(&ip_packet.to_immutable()));
|
||||
ip_packet.fill_checksum();
|
||||
}
|
||||
|
||||
buf
|
||||
|
||||
@@ -1,12 +1,9 @@
|
||||
use std::{future::Future, net::Ipv4Addr};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use pnet_packet::{
|
||||
MutablePacket, Packet,
|
||||
icmp::{self, IcmpPacket, IcmpTypes, MutableIcmpPacket},
|
||||
ip::IpNextHeaderProtocols,
|
||||
ipv4::{self, Ipv4Flags, Ipv4Packet, MutableIpv4Packet},
|
||||
udp::{self, MutableUdpPacket, UdpPacket},
|
||||
use smoltcp::wire::{
|
||||
IPV4_HEADER_LEN, Icmpv4Message, Icmpv4Packet, IpAddress, IpProtocol, Ipv4Packet,
|
||||
UDP_HEADER_LEN, UdpPacket,
|
||||
};
|
||||
|
||||
use crate::{
|
||||
@@ -119,43 +116,45 @@ where
|
||||
if packet.peer_manager_header().is_none() {
|
||||
return false;
|
||||
}
|
||||
let Some(ip_packet) = Ipv4Packet::new(packet.payload()) else {
|
||||
if packet.payload().len() < IPV4_HEADER_LEN {
|
||||
return false;
|
||||
};
|
||||
if ip_packet.get_version() != 4 || ip_packet.get_destination() != fake_ip {
|
||||
}
|
||||
let ip_packet = Ipv4Packet::new_unchecked(packet.payload());
|
||||
if ip_packet.version() != 4 || ip_packet.dst_addr() != fake_ip {
|
||||
return false;
|
||||
}
|
||||
|
||||
let ip_header_length = ip_packet.get_header_length() as usize * 4;
|
||||
let ip_total_length = ip_packet.get_total_length() as usize;
|
||||
if ip_header_length < MutableIpv4Packet::minimum_packet_size()
|
||||
let ip_header_length = ip_packet.header_len() as usize;
|
||||
let ip_total_length = ip_packet.total_len() as usize;
|
||||
if ip_header_length < IPV4_HEADER_LEN
|
||||
|| ip_header_length > ip_total_length
|
||||
|| ip_total_length != packet.payload().len()
|
||||
|| ip_packet.get_fragment_offset() != 0
|
||||
|| ip_packet.get_flags() & Ipv4Flags::MoreFragments != 0
|
||||
|| ip_packet.frag_offset() != 0
|
||||
|| ip_packet.more_frags()
|
||||
{
|
||||
return false;
|
||||
}
|
||||
|
||||
let protocol = ip_packet.get_next_level_protocol();
|
||||
let source_ip = ip_packet.get_source();
|
||||
let destination_ip = ip_packet.get_destination();
|
||||
let protocol = ip_packet.next_header();
|
||||
let source_ip = ip_packet.src_addr();
|
||||
let destination_ip = ip_packet.dst_addr();
|
||||
|
||||
match protocol {
|
||||
IpNextHeaderProtocols::Udp => {
|
||||
IpProtocol::Udp => {
|
||||
let ip_payload = &packet.payload()[ip_header_length..ip_total_length];
|
||||
let Some(udp_packet) = UdpPacket::new(ip_payload) else {
|
||||
return false;
|
||||
};
|
||||
let udp_length = udp_packet.get_length() as usize;
|
||||
if udp_length != ip_payload.len() || udp_length < UdpPacket::minimum_packet_size() {
|
||||
if ip_payload.len() < UDP_HEADER_LEN {
|
||||
return false;
|
||||
}
|
||||
if udp_packet.get_destination() != 53 {
|
||||
let udp_packet = UdpPacket::new_unchecked(ip_payload);
|
||||
let udp_length = udp_packet.len() as usize;
|
||||
if udp_length != ip_payload.len() || udp_length < UDP_HEADER_LEN {
|
||||
return false;
|
||||
}
|
||||
let source_port = udp_packet.get_source();
|
||||
let destination_port = udp_packet.get_destination();
|
||||
if udp_packet.dst_port() != 53 {
|
||||
return false;
|
||||
}
|
||||
let source_port = udp_packet.src_port();
|
||||
let destination_port = udp_packet.dst_port();
|
||||
let query = MagicDnsQuery {
|
||||
source: std::net::SocketAddr::from((source_ip, source_port)),
|
||||
payload: udp_packet.payload().to_vec(),
|
||||
@@ -175,30 +174,26 @@ where
|
||||
return false;
|
||||
}
|
||||
}
|
||||
IpNextHeaderProtocols::Icmp => {
|
||||
let Some(icmp_packet) = IcmpPacket::new(&packet.payload()[ip_header_length..]) else {
|
||||
return false;
|
||||
};
|
||||
if icmp_packet.get_icmp_type() != IcmpTypes::EchoRequest {
|
||||
return false;
|
||||
}
|
||||
let Some(mut icmp_packet) =
|
||||
MutableIcmpPacket::new(&mut packet.mut_payload()[ip_header_length..])
|
||||
IpProtocol::Icmp => {
|
||||
let Ok(icmp_packet) = Icmpv4Packet::new_checked(&packet.payload()[ip_header_length..])
|
||||
else {
|
||||
return false;
|
||||
};
|
||||
icmp_packet.set_icmp_type(IcmpTypes::EchoReply);
|
||||
icmp_packet.set_checksum(icmp::checksum(&icmp_packet.to_immutable()));
|
||||
if icmp_packet.msg_type() != Icmpv4Message::EchoRequest {
|
||||
return false;
|
||||
}
|
||||
let mut icmp_packet =
|
||||
Icmpv4Packet::new_unchecked(&mut packet.mut_payload()[ip_header_length..]);
|
||||
icmp_packet.set_msg_type(Icmpv4Message::EchoReply);
|
||||
icmp_packet.fill_checksum();
|
||||
}
|
||||
_ => return false,
|
||||
}
|
||||
|
||||
let Some(mut ip_packet) = MutableIpv4Packet::new(packet.mut_payload()) else {
|
||||
return false;
|
||||
};
|
||||
ip_packet.set_source(destination_ip);
|
||||
ip_packet.set_destination(source_ip);
|
||||
ip_packet.set_checksum(ipv4::checksum(&ip_packet.to_immutable()));
|
||||
let mut ip_packet = Ipv4Packet::new_unchecked(packet.mut_payload());
|
||||
ip_packet.set_src_addr(destination_ip);
|
||||
ip_packet.set_dst_addr(source_ip);
|
||||
ip_packet.fill_checksum();
|
||||
let payload_length = packet.payload().len() as u32;
|
||||
let Some(header) = packet.mut_peer_manager_header() else {
|
||||
return false;
|
||||
@@ -218,7 +213,7 @@ fn apply_udp_response(
|
||||
ip_header_length: usize,
|
||||
response: &[u8],
|
||||
) -> bool {
|
||||
let Some(udp_length) = UdpPacket::minimum_packet_size().checked_add(response.len()) else {
|
||||
let Some(udp_length) = UDP_HEADER_LEN.checked_add(response.len()) else {
|
||||
return false;
|
||||
};
|
||||
let Some(ip_length) = ip_header_length.checked_add(udp_length) else {
|
||||
@@ -237,26 +232,21 @@ fn apply_udp_response(
|
||||
if packet.mut_inner().capacity() < inner_length {
|
||||
packet
|
||||
.mut_inner()
|
||||
.truncate(header_length + ip_header_length + UdpPacket::minimum_packet_size());
|
||||
.truncate(header_length + ip_header_length + UDP_HEADER_LEN);
|
||||
}
|
||||
packet.mut_inner().resize(inner_length, 0);
|
||||
|
||||
let Some(mut ip_packet) = MutableIpv4Packet::new(packet.mut_payload()) else {
|
||||
return false;
|
||||
};
|
||||
ip_packet.set_total_length(ip_length as u16);
|
||||
let Some(mut udp_packet) = MutableUdpPacket::new(ip_packet.payload_mut()) else {
|
||||
return false;
|
||||
};
|
||||
udp_packet.set_length(udp_length as u16);
|
||||
udp_packet.set_source(destination_port);
|
||||
udp_packet.set_destination(source_port);
|
||||
let mut ip_packet = Ipv4Packet::new_unchecked(packet.mut_payload());
|
||||
ip_packet.set_total_len(ip_length as u16);
|
||||
let mut udp_packet = UdpPacket::new_unchecked(ip_packet.payload_mut());
|
||||
udp_packet.set_len(udp_length as u16);
|
||||
udp_packet.set_src_port(destination_port);
|
||||
udp_packet.set_dst_port(source_port);
|
||||
udp_packet.payload_mut().copy_from_slice(response);
|
||||
udp_packet.set_checksum(udp::ipv4_checksum(
|
||||
&udp_packet.to_immutable(),
|
||||
&destination_ip,
|
||||
&source_ip,
|
||||
));
|
||||
udp_packet.fill_checksum(
|
||||
&IpAddress::Ipv4(destination_ip),
|
||||
&IpAddress::Ipv4(source_ip),
|
||||
);
|
||||
true
|
||||
}
|
||||
|
||||
@@ -267,17 +257,17 @@ mod tests {
|
||||
fn udp_query(payload: &[u8], destination_port: u16) -> ZCPacket {
|
||||
let mut bytes = vec![0; 20 + 8 + payload.len()];
|
||||
{
|
||||
let mut ip = MutableIpv4Packet::new(&mut bytes).unwrap();
|
||||
let mut ip = Ipv4Packet::new_unchecked(&mut bytes);
|
||||
ip.set_version(4);
|
||||
ip.set_header_length(5);
|
||||
ip.set_total_length((20 + 8 + payload.len()) as u16);
|
||||
ip.set_next_level_protocol(IpNextHeaderProtocols::Udp);
|
||||
ip.set_source("10.0.0.2".parse().unwrap());
|
||||
ip.set_destination("100.100.100.101".parse().unwrap());
|
||||
let mut udp = MutableUdpPacket::new(ip.payload_mut()).unwrap();
|
||||
udp.set_source(53000);
|
||||
udp.set_destination(destination_port);
|
||||
udp.set_length((8 + payload.len()) as u16);
|
||||
ip.set_header_len(20);
|
||||
ip.set_total_len((20 + 8 + payload.len()) as u16);
|
||||
ip.set_next_header(IpProtocol::Udp);
|
||||
ip.set_src_addr("10.0.0.2".parse().unwrap());
|
||||
ip.set_dst_addr("100.100.100.101".parse().unwrap());
|
||||
let mut udp = UdpPacket::new_unchecked(ip.payload_mut());
|
||||
udp.set_src_port(53000);
|
||||
udp.set_dst_port(destination_port);
|
||||
udp.set_len((8 + payload.len()) as u16);
|
||||
udp.payload_mut().copy_from_slice(payload);
|
||||
}
|
||||
ZCPacket::new_with_payload(&bytes)
|
||||
@@ -286,15 +276,15 @@ mod tests {
|
||||
fn icmp_echo_request() -> ZCPacket {
|
||||
let mut bytes = vec![0; 20 + 8];
|
||||
{
|
||||
let mut ip = MutableIpv4Packet::new(&mut bytes).unwrap();
|
||||
let mut ip = Ipv4Packet::new_unchecked(&mut bytes);
|
||||
ip.set_version(4);
|
||||
ip.set_header_length(5);
|
||||
ip.set_total_length(28);
|
||||
ip.set_next_level_protocol(IpNextHeaderProtocols::Icmp);
|
||||
ip.set_source("10.0.0.2".parse().unwrap());
|
||||
ip.set_destination("100.100.100.101".parse().unwrap());
|
||||
let mut icmp = MutableIcmpPacket::new(ip.payload_mut()).unwrap();
|
||||
icmp.set_icmp_type(IcmpTypes::EchoRequest);
|
||||
ip.set_header_len(20);
|
||||
ip.set_total_len(28);
|
||||
ip.set_next_header(IpProtocol::Icmp);
|
||||
ip.set_src_addr("10.0.0.2".parse().unwrap());
|
||||
ip.set_dst_addr("100.100.100.101".parse().unwrap());
|
||||
let mut icmp = Icmpv4Packet::new_unchecked(ip.payload_mut());
|
||||
icmp.set_msg_type(Icmpv4Message::EchoRequest);
|
||||
}
|
||||
ZCPacket::new_with_payload(&bytes)
|
||||
}
|
||||
@@ -314,18 +304,15 @@ mod tests {
|
||||
.await;
|
||||
|
||||
assert!(handled);
|
||||
let ip = Ipv4Packet::new(packet.payload()).unwrap();
|
||||
let ip = Ipv4Packet::new_checked(packet.payload()).unwrap();
|
||||
assert_eq!(
|
||||
ip.get_source(),
|
||||
ip.src_addr(),
|
||||
"100.100.100.101".parse::<Ipv4Addr>().unwrap()
|
||||
);
|
||||
assert_eq!(
|
||||
ip.get_destination(),
|
||||
"10.0.0.2".parse::<Ipv4Addr>().unwrap()
|
||||
);
|
||||
let udp = UdpPacket::new(ip.payload()).unwrap();
|
||||
assert_eq!(udp.get_source(), 53);
|
||||
assert_eq!(udp.get_destination(), 53000);
|
||||
assert_eq!(ip.dst_addr(), "10.0.0.2".parse::<Ipv4Addr>().unwrap());
|
||||
let udp = UdpPacket::new_checked(ip.payload()).unwrap();
|
||||
assert_eq!(udp.src_port(), 53);
|
||||
assert_eq!(udp.dst_port(), 53000);
|
||||
assert_eq!(udp.payload(), b"response");
|
||||
assert_eq!(packet.get_dst_peer_id(), Some(42));
|
||||
assert_eq!(
|
||||
@@ -337,9 +324,7 @@ mod tests {
|
||||
#[tokio::test]
|
||||
async fn packet_engine_rejects_invalid_ipv4_header_without_mutation() {
|
||||
let mut packet = udp_query(b"query", 53);
|
||||
MutableIpv4Packet::new(packet.mut_payload())
|
||||
.unwrap()
|
||||
.set_header_length(15);
|
||||
Ipv4Packet::new_unchecked(packet.mut_payload()).set_header_len(60);
|
||||
let original = packet.payload().to_vec();
|
||||
|
||||
assert!(
|
||||
@@ -373,10 +358,8 @@ mod tests {
|
||||
#[tokio::test]
|
||||
async fn packet_engine_rejects_inconsistent_udp_length_without_mutation() {
|
||||
let mut packet = udp_query(b"query", 53);
|
||||
let mut ip = MutableIpv4Packet::new(packet.mut_payload()).unwrap();
|
||||
MutableUdpPacket::new(ip.payload_mut())
|
||||
.unwrap()
|
||||
.set_length(8);
|
||||
let mut ip = Ipv4Packet::new_unchecked(packet.mut_payload());
|
||||
UdpPacket::new_unchecked(ip.payload_mut()).set_len(8);
|
||||
let original = packet.payload().to_vec();
|
||||
|
||||
assert!(
|
||||
@@ -394,9 +377,7 @@ mod tests {
|
||||
#[tokio::test]
|
||||
async fn packet_engine_rejects_fragmented_packets_without_mutation() {
|
||||
let mut packet = udp_query(b"query", 53);
|
||||
MutableIpv4Packet::new(packet.mut_payload())
|
||||
.unwrap()
|
||||
.set_flags(Ipv4Flags::MoreFragments);
|
||||
Ipv4Packet::new_unchecked(packet.mut_payload()).set_more_frags(true);
|
||||
let original = packet.payload().to_vec();
|
||||
|
||||
assert!(
|
||||
@@ -442,13 +423,13 @@ mod tests {
|
||||
.await;
|
||||
|
||||
assert!(handled);
|
||||
let ip = Ipv4Packet::new(packet.payload()).unwrap();
|
||||
let ip = Ipv4Packet::new_checked(packet.payload()).unwrap();
|
||||
assert_eq!(
|
||||
ip.get_source(),
|
||||
ip.src_addr(),
|
||||
"100.100.100.101".parse::<Ipv4Addr>().unwrap()
|
||||
);
|
||||
let icmp = pnet_packet::icmp::IcmpPacket::new(ip.payload()).unwrap();
|
||||
assert_eq!(icmp.get_icmp_type(), IcmpTypes::EchoReply);
|
||||
let icmp = Icmpv4Packet::new_checked(ip.payload()).unwrap();
|
||||
assert_eq!(icmp.msg_type(), Icmpv4Message::EchoReply);
|
||||
assert_eq!(packet.get_dst_peer_id(), Some(7));
|
||||
}
|
||||
|
||||
|
||||
@@ -1,21 +1,14 @@
|
||||
use std::{net::Ipv4Addr, sync::Arc, time::Duration};
|
||||
|
||||
use dashmap::DashMap;
|
||||
use pnet_packet::{
|
||||
Packet,
|
||||
icmp::{self, IcmpCode, IcmpTypes, MutableIcmpPacket, echo_reply::MutableEchoReplyPacket},
|
||||
ip::IpNextHeaderProtocols,
|
||||
ipv4::Ipv4Packet,
|
||||
};
|
||||
use quanta::Instant;
|
||||
use smoltcp::wire::{IPV4_HEADER_LEN, Icmpv4Message, Icmpv4Packet, Ipv4Packet};
|
||||
|
||||
use crate::packet::{PacketType, ZCPacket};
|
||||
|
||||
use super::{
|
||||
cidr_table::ProxyCidrTable,
|
||||
ip_reassembler::{
|
||||
ComposeIpv4PacketArgs, IpProtocol, IpReassembler, SmolIpv4Packet, compose_ipv4_packet,
|
||||
},
|
||||
ip_reassembler::{ComposeIpv4PacketArgs, IpProtocol, IpReassembler, compose_ipv4_packet},
|
||||
};
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
|
||||
@@ -51,6 +44,8 @@ struct IcmpNatEntry {
|
||||
started_at: Instant,
|
||||
}
|
||||
|
||||
const ICMP_ECHO_HEADER_LEN: usize = 8;
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct IcmpProxyEngine {
|
||||
cidr_table: Arc<ProxyCidrTable>,
|
||||
@@ -84,15 +79,17 @@ impl IcmpProxyEngine {
|
||||
if header.packet_type != PacketType::Data as u8 || header.is_no_proxy() {
|
||||
return IcmpProxyAction::Pass;
|
||||
}
|
||||
let Some(ipv4) = Ipv4Packet::new(packet.payload()) else {
|
||||
let Ok(ipv4) = Ipv4Packet::new_checked(packet.payload()) else {
|
||||
return IcmpProxyAction::Pass;
|
||||
};
|
||||
if ipv4.get_version() != 4 || ipv4.get_next_level_protocol() != IpNextHeaderProtocols::Icmp
|
||||
if ipv4.version() != 4
|
||||
|| usize::from(ipv4.header_len()) < IPV4_HEADER_LEN
|
||||
|| ipv4.next_header() != IpProtocol::Icmp
|
||||
{
|
||||
return IcmpProxyAction::Pass;
|
||||
}
|
||||
|
||||
let mapped_destination = ipv4.get_destination();
|
||||
let mapped_destination = ipv4.dst_addr();
|
||||
let real_destination = self.cidr_table.lookup_v4(mapped_destination);
|
||||
let is_local_no_tun = context.no_tun && mapped_destination == virtual_ipv4;
|
||||
if real_destination.is_none() && !header.is_exit_node() && !is_local_no_tun {
|
||||
@@ -100,33 +97,27 @@ impl IcmpProxyEngine {
|
||||
}
|
||||
|
||||
let reassembled;
|
||||
let smol_ipv4 = SmolIpv4Packet::new_unchecked(ipv4.packet());
|
||||
let request = if IpReassembler::is_packet_fragmented(&smol_ipv4) {
|
||||
let Ok(smol_ipv4) = SmolIpv4Packet::new_checked(ipv4.packet()) else {
|
||||
return IcmpProxyAction::Pass;
|
||||
};
|
||||
reassembled = self.reassembler.add_fragment(&smol_ipv4);
|
||||
let request_bytes = if IpReassembler::is_packet_fragmented(&ipv4) {
|
||||
reassembled = self.reassembler.add_fragment(&ipv4);
|
||||
let Some(reassembled) = reassembled.as_ref() else {
|
||||
return IcmpProxyAction::Pass;
|
||||
};
|
||||
let Some(request) = icmp::echo_request::EchoRequestPacket::new(reassembled) else {
|
||||
return IcmpProxyAction::Pass;
|
||||
};
|
||||
request
|
||||
reassembled.as_slice()
|
||||
} else {
|
||||
let Some(request) = icmp::echo_request::EchoRequestPacket::new(ipv4.payload()) else {
|
||||
return IcmpProxyAction::Pass;
|
||||
};
|
||||
request
|
||||
ipv4.payload()
|
||||
};
|
||||
if request.get_icmp_type() != IcmpTypes::EchoRequest {
|
||||
if request_bytes.len() < ICMP_ECHO_HEADER_LEN {
|
||||
return IcmpProxyAction::Pass;
|
||||
}
|
||||
let request = Icmpv4Packet::new_unchecked(request_bytes);
|
||||
if request.msg_type() != Icmpv4Message::EchoRequest {
|
||||
return IcmpProxyAction::Pass;
|
||||
}
|
||||
|
||||
if is_local_no_tun {
|
||||
return self.local_reply(
|
||||
mapped_destination,
|
||||
ipv4.get_source(),
|
||||
ipv4.src_addr(),
|
||||
header.to_peer_id.get(),
|
||||
header.from_peer_id.get(),
|
||||
&request,
|
||||
@@ -136,15 +127,15 @@ impl IcmpProxyEngine {
|
||||
let real_destination = real_destination.unwrap_or(mapped_destination);
|
||||
let key = IcmpNatKey {
|
||||
real_destination,
|
||||
identifier: request.get_identifier(),
|
||||
sequence: request.get_sequence_number(),
|
||||
identifier: request.echo_ident(),
|
||||
sequence: request.echo_seq_no(),
|
||||
};
|
||||
self.nat_table.insert(
|
||||
key,
|
||||
IcmpNatEntry {
|
||||
source_peer_id: header.from_peer_id.get(),
|
||||
local_peer_id: header.to_peer_id.get(),
|
||||
source_ip: ipv4.get_source(),
|
||||
source_ip: ipv4.src_addr(),
|
||||
mapped_destination,
|
||||
started_at: Instant::now(),
|
||||
},
|
||||
@@ -152,35 +143,35 @@ impl IcmpProxyEngine {
|
||||
|
||||
IcmpProxyAction::SendToSocket {
|
||||
destination: real_destination,
|
||||
packet: request.packet().to_vec(),
|
||||
packet: request.as_ref().to_vec(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn handle_socket_response(&self, peer_ip: Ipv4Addr, packet: &mut [u8]) -> Vec<ZCPacket> {
|
||||
let Some(ipv4) = Ipv4Packet::new(packet) else {
|
||||
let Ok(ipv4) = Ipv4Packet::new_checked(&*packet) else {
|
||||
return Vec::new();
|
||||
};
|
||||
let Some(reply) = icmp::echo_reply::EchoReplyPacket::new(ipv4.payload()) else {
|
||||
if usize::from(ipv4.header_len()) < IPV4_HEADER_LEN
|
||||
|| ipv4.payload().len() < ICMP_ECHO_HEADER_LEN
|
||||
{
|
||||
return Vec::new();
|
||||
};
|
||||
if reply.get_icmp_type() != IcmpTypes::EchoReply {
|
||||
}
|
||||
let reply = Icmpv4Packet::new_unchecked(ipv4.payload());
|
||||
if reply.msg_type() != Icmpv4Message::EchoReply {
|
||||
return Vec::new();
|
||||
}
|
||||
let key = IcmpNatKey {
|
||||
real_destination: peer_ip,
|
||||
identifier: reply.get_identifier(),
|
||||
sequence: reply.get_sequence_number(),
|
||||
identifier: reply.echo_ident(),
|
||||
sequence: reply.echo_seq_no(),
|
||||
};
|
||||
let Some((_, entry)) = self.nat_table.remove(&key) else {
|
||||
return Vec::new();
|
||||
};
|
||||
let Some(payload_len) = packet
|
||||
.len()
|
||||
.checked_sub(ipv4.get_header_length() as usize * 4)
|
||||
else {
|
||||
let Some(payload_len) = packet.len().checked_sub(ipv4.header_len() as usize) else {
|
||||
return Vec::new();
|
||||
};
|
||||
let ip_id = ipv4.get_identification();
|
||||
let ip_id = ipv4.ident();
|
||||
let mut responses = Vec::new();
|
||||
let _ = compose_ipv4_packet(
|
||||
ComposeIpv4PacketArgs {
|
||||
@@ -226,17 +217,16 @@ impl IcmpProxyEngine {
|
||||
destination: Ipv4Addr,
|
||||
source_peer_id: u32,
|
||||
destination_peer_id: u32,
|
||||
request: &icmp::echo_request::EchoRequestPacket<'_>,
|
||||
request: &Icmpv4Packet<&[u8]>,
|
||||
) -> IcmpProxyAction {
|
||||
let mut buffer = vec![0_u8; request.packet().len() + 20];
|
||||
let mut reply = MutableEchoReplyPacket::new(&mut buffer[20..]).unwrap();
|
||||
reply.set_icmp_type(IcmpTypes::EchoReply);
|
||||
reply.set_icmp_code(IcmpCode::new(0));
|
||||
reply.set_identifier(request.get_identifier());
|
||||
reply.set_sequence_number(request.get_sequence_number());
|
||||
reply.set_payload(request.payload());
|
||||
let mut reply = MutableIcmpPacket::new(&mut buffer[20..]).unwrap();
|
||||
reply.set_checksum(icmp::checksum(&reply.to_immutable()));
|
||||
let mut buffer = vec![0_u8; request.as_ref().len() + 20];
|
||||
let mut reply = Icmpv4Packet::new_unchecked(&mut buffer[20..]);
|
||||
reply.set_msg_type(Icmpv4Message::EchoReply);
|
||||
reply.set_msg_code(0);
|
||||
reply.set_echo_ident(request.echo_ident());
|
||||
reply.set_echo_seq_no(request.echo_seq_no());
|
||||
reply.data_mut().copy_from_slice(request.data());
|
||||
reply.fill_checksum();
|
||||
|
||||
let payload_len = buffer.len() - 20;
|
||||
let mut responses = Vec::new();
|
||||
@@ -267,12 +257,6 @@ impl IcmpProxyEngine {
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use pnet_packet::{
|
||||
MutablePacket as _,
|
||||
icmp::{MutableIcmpPacket, echo_request::MutableEchoRequestPacket},
|
||||
ipv4::{self, MutableIpv4Packet},
|
||||
};
|
||||
|
||||
use super::*;
|
||||
use crate::gateway::proxy::cidr_table::{ProxyCidrRule, ProxyCidrSnapshot};
|
||||
|
||||
@@ -283,25 +267,24 @@ mod tests {
|
||||
) -> ZCPacket {
|
||||
let mut bytes = vec![0_u8; 20 + 8 + payload.len()];
|
||||
{
|
||||
let mut request = MutableEchoRequestPacket::new(&mut bytes[20..]).unwrap();
|
||||
request.set_icmp_type(IcmpTypes::EchoRequest);
|
||||
request.set_identifier(7);
|
||||
request.set_sequence_number(11);
|
||||
request.set_payload(payload);
|
||||
let mut icmp = MutableIcmpPacket::new(&mut bytes[20..]).unwrap();
|
||||
icmp.set_checksum(icmp::checksum(&icmp.to_immutable()));
|
||||
let mut request = Icmpv4Packet::new_unchecked(&mut bytes[20..]);
|
||||
request.set_msg_type(Icmpv4Message::EchoRequest);
|
||||
request.set_echo_ident(7);
|
||||
request.set_echo_seq_no(11);
|
||||
request.data_mut().copy_from_slice(payload);
|
||||
request.fill_checksum();
|
||||
}
|
||||
{
|
||||
let packet_len = bytes.len() as u16;
|
||||
let mut ipv4 = MutableIpv4Packet::new(&mut bytes).unwrap();
|
||||
let mut ipv4 = Ipv4Packet::new_unchecked(&mut bytes);
|
||||
ipv4.set_version(4);
|
||||
ipv4.set_header_length(5);
|
||||
ipv4.set_total_length(packet_len);
|
||||
ipv4.set_ttl(64);
|
||||
ipv4.set_next_level_protocol(IpNextHeaderProtocols::Icmp);
|
||||
ipv4.set_source(source);
|
||||
ipv4.set_destination(destination);
|
||||
ipv4.set_checksum(ipv4::checksum(&ipv4.to_immutable()));
|
||||
ipv4.set_header_len(20);
|
||||
ipv4.set_total_len(packet_len);
|
||||
ipv4.set_hop_limit(64);
|
||||
ipv4.set_next_header(IpProtocol::Icmp);
|
||||
ipv4.set_src_addr(source);
|
||||
ipv4.set_dst_addr(destination);
|
||||
ipv4.fill_checksum();
|
||||
}
|
||||
let mut packet = ZCPacket::new_with_payload(&bytes);
|
||||
packet.fill_peer_manager_hdr(101, 202, PacketType::Data as u8);
|
||||
@@ -357,16 +340,32 @@ mod tests {
|
||||
let header = reply.peer_manager_header().unwrap();
|
||||
assert_eq!(header.from_peer_id.get(), 202);
|
||||
assert_eq!(header.to_peer_id.get(), 101);
|
||||
let ipv4 = Ipv4Packet::new(reply.payload()).unwrap();
|
||||
assert_eq!(ipv4.get_source(), "10.0.0.1".parse::<Ipv4Addr>().unwrap());
|
||||
assert_eq!(
|
||||
ipv4.get_destination(),
|
||||
"10.0.0.2".parse::<Ipv4Addr>().unwrap()
|
||||
);
|
||||
let reply = icmp::echo_reply::EchoReplyPacket::new(ipv4.payload()).unwrap();
|
||||
assert_eq!(reply.get_identifier(), 7);
|
||||
assert_eq!(reply.get_sequence_number(), 11);
|
||||
assert_eq!(reply.payload(), b"ping");
|
||||
let ipv4 = Ipv4Packet::new_checked(reply.payload()).unwrap();
|
||||
assert_eq!(ipv4.src_addr(), "10.0.0.1".parse::<Ipv4Addr>().unwrap());
|
||||
assert_eq!(ipv4.dst_addr(), "10.0.0.2".parse::<Ipv4Addr>().unwrap());
|
||||
let reply = Icmpv4Packet::new_checked(ipv4.payload()).unwrap();
|
||||
assert_eq!(reply.echo_ident(), 7);
|
||||
assert_eq!(reply.echo_seq_no(), 11);
|
||||
assert_eq!(reply.data(), b"ping");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn peer_packet_rejects_ipv4_header_shorter_than_minimum() {
|
||||
let engine = engine(None);
|
||||
let mut packet = echo_request("10.0.0.2".parse().unwrap(), "8.0.0.1".parse().unwrap());
|
||||
Ipv4Packet::new_unchecked(packet.mut_payload()).set_header_len(16);
|
||||
|
||||
assert!(matches!(
|
||||
engine.handle_peer_packet(
|
||||
&packet,
|
||||
IcmpProxyContext {
|
||||
virtual_ipv4: Some("8.0.0.1".parse().unwrap()),
|
||||
no_tun: true,
|
||||
..Default::default()
|
||||
}
|
||||
),
|
||||
IcmpProxyAction::Pass
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -391,20 +390,19 @@ mod tests {
|
||||
panic!("expected socket request");
|
||||
};
|
||||
assert_eq!(destination, "127.0.0.42".parse::<Ipv4Addr>().unwrap());
|
||||
let request = icmp::echo_request::EchoRequestPacket::new(&request).unwrap();
|
||||
assert_eq!(request.payload(), b"ping");
|
||||
let request = Icmpv4Packet::new_checked(&request).unwrap();
|
||||
assert_eq!(request.data(), b"ping");
|
||||
|
||||
let mut response = echo_request(destination, "10.0.0.1".parse().unwrap())
|
||||
.payload()
|
||||
.to_vec();
|
||||
{
|
||||
let mut ipv4 = MutableIpv4Packet::new(&mut response).unwrap();
|
||||
let mut reply = MutableEchoReplyPacket::new(ipv4.payload_mut()).unwrap();
|
||||
reply.set_icmp_type(IcmpTypes::EchoReply);
|
||||
let mut icmp = MutableIcmpPacket::new(ipv4.payload_mut()).unwrap();
|
||||
icmp.set_checksum(icmp::checksum(&icmp.to_immutable()));
|
||||
ipv4.set_source(destination);
|
||||
ipv4.set_checksum(ipv4::checksum(&ipv4.to_immutable()));
|
||||
let mut ipv4 = Ipv4Packet::new_unchecked(&mut response);
|
||||
let mut reply = Icmpv4Packet::new_unchecked(ipv4.payload_mut());
|
||||
reply.set_msg_type(Icmpv4Message::EchoReply);
|
||||
reply.fill_checksum();
|
||||
ipv4.set_src_addr(destination);
|
||||
ipv4.fill_checksum();
|
||||
}
|
||||
let replies = engine.handle_socket_response(destination, &mut response);
|
||||
let [reply] = replies.as_slice() else {
|
||||
@@ -414,15 +412,54 @@ mod tests {
|
||||
assert_eq!(header.from_peer_id.get(), 202);
|
||||
assert_eq!(header.to_peer_id.get(), 101);
|
||||
assert!(header.is_no_proxy());
|
||||
let ipv4 = Ipv4Packet::new(reply.payload()).unwrap();
|
||||
assert_eq!(
|
||||
ipv4.get_source(),
|
||||
"10.10.10.42".parse::<Ipv4Addr>().unwrap()
|
||||
);
|
||||
assert_eq!(
|
||||
ipv4.get_destination(),
|
||||
"10.0.0.2".parse::<Ipv4Addr>().unwrap()
|
||||
let ipv4 = Ipv4Packet::new_checked(reply.payload()).unwrap();
|
||||
assert_eq!(ipv4.src_addr(), "10.10.10.42".parse::<Ipv4Addr>().unwrap());
|
||||
assert_eq!(ipv4.dst_addr(), "10.0.0.2".parse::<Ipv4Addr>().unwrap());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn socket_response_rejects_ipv4_header_shorter_than_minimum() {
|
||||
let engine = engine(Some(ProxyCidrRule {
|
||||
cidr: "127.0.0.0/24".parse().unwrap(),
|
||||
mapped_cidr: Some("10.10.10.0/24".parse().unwrap()),
|
||||
}));
|
||||
let destination = "127.0.0.42".parse().unwrap();
|
||||
let request = echo_request("10.0.0.2".parse().unwrap(), "10.10.10.42".parse().unwrap());
|
||||
assert!(matches!(
|
||||
engine.handle_peer_packet(
|
||||
&request,
|
||||
IcmpProxyContext {
|
||||
virtual_ipv4: Some("10.0.0.1".parse().unwrap()),
|
||||
..Default::default()
|
||||
},
|
||||
),
|
||||
IcmpProxyAction::SendToSocket { .. }
|
||||
));
|
||||
|
||||
let key = IcmpNatKey {
|
||||
real_destination: destination,
|
||||
identifier: 7,
|
||||
sequence: 11,
|
||||
};
|
||||
let mut response = echo_request(destination, "10.0.0.1".parse().unwrap())
|
||||
.payload()
|
||||
.to_vec();
|
||||
{
|
||||
let mut ipv4 = Ipv4Packet::new_unchecked(&mut response);
|
||||
ipv4.set_header_len(16);
|
||||
let mut reply = Icmpv4Packet::new_unchecked(ipv4.payload_mut());
|
||||
reply.set_msg_type(Icmpv4Message::EchoReply);
|
||||
reply.set_msg_code(0);
|
||||
reply.set_echo_ident(7);
|
||||
reply.set_echo_seq_no(11);
|
||||
}
|
||||
|
||||
assert!(
|
||||
engine
|
||||
.handle_socket_response(destination, &mut response)
|
||||
.is_empty()
|
||||
);
|
||||
assert!(engine.nat_table.contains_key(&key));
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -485,24 +522,23 @@ mod tests {
|
||||
.payload()
|
||||
.to_vec();
|
||||
{
|
||||
let mut ipv4 = MutableIpv4Packet::new(&mut response).unwrap();
|
||||
let mut reply = MutableEchoReplyPacket::new(ipv4.payload_mut()).unwrap();
|
||||
reply.set_icmp_type(IcmpTypes::EchoReply);
|
||||
let mut icmp = MutableIcmpPacket::new(ipv4.payload_mut()).unwrap();
|
||||
icmp.set_checksum(icmp::checksum(&icmp.to_immutable()));
|
||||
ipv4.set_source(destination);
|
||||
let mut ipv4 = Ipv4Packet::new_unchecked(&mut response);
|
||||
let mut reply = Icmpv4Packet::new_unchecked(ipv4.payload_mut());
|
||||
reply.set_msg_type(Icmpv4Message::EchoReply);
|
||||
reply.fill_checksum();
|
||||
ipv4.set_src_addr(destination);
|
||||
// Raw sockets may return a buffer with bytes beyond the IPv4 total
|
||||
// length. The native implementation composes from the received
|
||||
// buffer length, so keep that case covered without changing the
|
||||
// existing in-place composer in this refactor.
|
||||
ipv4.set_total_length(1220);
|
||||
ipv4.set_checksum(ipv4::checksum(&ipv4.to_immutable()));
|
||||
ipv4.set_total_len(1220);
|
||||
ipv4.fill_checksum();
|
||||
}
|
||||
let ipv4 = Ipv4Packet::new(&response).unwrap();
|
||||
let echo_reply = icmp::echo_reply::EchoReplyPacket::new(ipv4.payload()).unwrap();
|
||||
assert_eq!(echo_reply.get_icmp_type(), IcmpTypes::EchoReply);
|
||||
assert_eq!(echo_reply.get_identifier(), 7);
|
||||
assert_eq!(echo_reply.get_sequence_number(), 11);
|
||||
let ipv4 = Ipv4Packet::new_checked(&response).unwrap();
|
||||
let echo_reply = Icmpv4Packet::new_checked(ipv4.payload()).unwrap();
|
||||
assert_eq!(echo_reply.msg_type(), Icmpv4Message::EchoReply);
|
||||
assert_eq!(echo_reply.echo_ident(), 7);
|
||||
assert_eq!(echo_reply.echo_seq_no(), 11);
|
||||
|
||||
let replies = engine.handle_socket_response(destination, &mut response);
|
||||
assert_eq!(replies.len(), 3);
|
||||
|
||||
@@ -4,8 +4,8 @@ use std::{
|
||||
};
|
||||
|
||||
use dashmap::DashMap;
|
||||
pub use smoltcp::wire::IpProtocol;
|
||||
use smoltcp::wire::Ipv4Packet;
|
||||
pub use smoltcp::wire::{IpProtocol, Ipv4Packet as SmolIpv4Packet};
|
||||
|
||||
#[derive(Debug, Hash, PartialEq, Eq, Clone)]
|
||||
struct IpReassemblerKey {
|
||||
|
||||
@@ -1,10 +1,8 @@
|
||||
use std::net::Ipv4Addr;
|
||||
|
||||
use cidr::Ipv4Inet;
|
||||
use pnet_packet::{
|
||||
ip::IpNextHeaderProtocols,
|
||||
ipv4::{self, Ipv4Flags, Ipv4Packet, MutableIpv4Packet},
|
||||
udp::{self, MutableUdpPacket, UdpPacket},
|
||||
use smoltcp::wire::{
|
||||
IPV4_HEADER_LEN, IpAddress, IpProtocol, Ipv4Packet, UDP_HEADER_LEN, UdpPacket,
|
||||
};
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
|
||||
@@ -128,36 +126,37 @@ pub struct UdpPacketSummary {
|
||||
|
||||
impl UdpPacketSummary {
|
||||
pub fn parse(packet: &[u8]) -> Option<Self> {
|
||||
let ipv4_packet = Ipv4Packet::new(packet)?;
|
||||
if ipv4_packet.get_version() != 4
|
||||
|| ipv4_packet.get_next_level_protocol() != IpNextHeaderProtocols::Udp
|
||||
{
|
||||
if packet.len() < IPV4_HEADER_LEN {
|
||||
return None;
|
||||
}
|
||||
let ipv4_packet = Ipv4Packet::new_unchecked(packet);
|
||||
if ipv4_packet.version() != 4 || ipv4_packet.next_header() != IpProtocol::Udp {
|
||||
return None;
|
||||
}
|
||||
|
||||
let header_len = usize::from(ipv4_packet.get_header_length()) * 4;
|
||||
let total_len = usize::from(ipv4_packet.get_total_length());
|
||||
if header_len < Ipv4Packet::minimum_packet_size()
|
||||
|| total_len < header_len + UdpPacket::minimum_packet_size()
|
||||
let header_len = usize::from(ipv4_packet.header_len());
|
||||
let total_len = usize::from(ipv4_packet.total_len());
|
||||
if header_len < IPV4_HEADER_LEN
|
||||
|| total_len < header_len + UDP_HEADER_LEN
|
||||
|| total_len > packet.len()
|
||||
{
|
||||
return None;
|
||||
}
|
||||
|
||||
let udp_packet = UdpPacket::new(&packet[header_len..total_len])?;
|
||||
let udp_len = usize::from(udp_packet.get_length());
|
||||
if udp_len < UdpPacket::minimum_packet_size() || header_len + udp_len != total_len {
|
||||
let udp_packet = UdpPacket::new_unchecked(&packet[header_len..total_len]);
|
||||
let udp_len = usize::from(udp_packet.len());
|
||||
if udp_len < UDP_HEADER_LEN || header_len + udp_len != total_len {
|
||||
return None;
|
||||
}
|
||||
|
||||
Some(Self {
|
||||
src: ipv4_packet.get_source(),
|
||||
dst: ipv4_packet.get_destination(),
|
||||
src_port: udp_packet.get_source(),
|
||||
dst_port: udp_packet.get_destination(),
|
||||
src: ipv4_packet.src_addr(),
|
||||
dst: ipv4_packet.dst_addr(),
|
||||
src_port: udp_packet.src_port(),
|
||||
dst_port: udp_packet.dst_port(),
|
||||
ip_len: total_len,
|
||||
udp_len,
|
||||
payload_len: udp_len - UdpPacket::minimum_packet_size(),
|
||||
payload_len: udp_len - UDP_HEADER_LEN,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -234,30 +233,29 @@ fn parse_udp_broadcast(
|
||||
packet: &[u8],
|
||||
config: &BroadcastRelayConfig,
|
||||
) -> Result<ParsedUdpBroadcastPacket, UdpBroadcastPacketRejection> {
|
||||
let ipv4_packet = Ipv4Packet::new(packet).ok_or(UdpBroadcastPacketRejection::MalformedIpv4)?;
|
||||
if ipv4_packet.get_version() != 4
|
||||
|| ipv4_packet.get_next_level_protocol() != IpNextHeaderProtocols::Udp
|
||||
{
|
||||
if packet.len() < IPV4_HEADER_LEN {
|
||||
return Err(UdpBroadcastPacketRejection::MalformedIpv4);
|
||||
}
|
||||
let ipv4_packet = Ipv4Packet::new_unchecked(packet);
|
||||
if ipv4_packet.version() != 4 || ipv4_packet.next_header() != IpProtocol::Udp {
|
||||
return Err(UdpBroadcastPacketRejection::NotUdpIpv4);
|
||||
}
|
||||
|
||||
if ipv4_packet.get_fragment_offset() != 0
|
||||
|| ipv4_packet.get_flags() & Ipv4Flags::MoreFragments != 0
|
||||
{
|
||||
if ipv4_packet.frag_offset() != 0 || ipv4_packet.more_frags() {
|
||||
return Err(UdpBroadcastPacketRejection::Fragmented);
|
||||
}
|
||||
|
||||
let header_len = usize::from(ipv4_packet.get_header_length()) * 4;
|
||||
let total_len = usize::from(ipv4_packet.get_total_length());
|
||||
if header_len < Ipv4Packet::minimum_packet_size()
|
||||
|| total_len < header_len + UdpPacket::minimum_packet_size()
|
||||
let header_len = usize::from(ipv4_packet.header_len());
|
||||
let total_len = usize::from(ipv4_packet.total_len());
|
||||
if header_len < IPV4_HEADER_LEN
|
||||
|| total_len < header_len + UDP_HEADER_LEN
|
||||
|| total_len > packet.len()
|
||||
{
|
||||
return Err(UdpBroadcastPacketRejection::BadIpv4Length);
|
||||
}
|
||||
|
||||
let src = ipv4_packet.get_source();
|
||||
let dst = ipv4_packet.get_destination();
|
||||
let src = ipv4_packet.src_addr();
|
||||
let dst = ipv4_packet.dst_addr();
|
||||
if should_ignore_interface_addr(src) {
|
||||
return Err(UdpBroadcastPacketRejection::IgnoredSource);
|
||||
}
|
||||
@@ -275,10 +273,9 @@ fn parse_udp_broadcast(
|
||||
return Err(UdpBroadcastPacketRejection::LoopbackDestination);
|
||||
}
|
||||
|
||||
let udp_packet = UdpPacket::new(&packet[header_len..total_len])
|
||||
.ok_or(UdpBroadcastPacketRejection::MalformedUdp)?;
|
||||
let udp_len = usize::from(udp_packet.get_length());
|
||||
if udp_len < UdpPacket::minimum_packet_size() || header_len + udp_len != total_len {
|
||||
let udp_packet = UdpPacket::new_unchecked(&packet[header_len..total_len]);
|
||||
let udp_len = usize::from(udp_packet.len());
|
||||
if udp_len < UDP_HEADER_LEN || header_len + udp_len != total_len {
|
||||
return Err(UdpBroadcastPacketRejection::BadUdpLength);
|
||||
}
|
||||
|
||||
@@ -301,27 +298,23 @@ pub fn normalize_udp_broadcast_packet(
|
||||
let mut normalized = packet[..packet_len].to_vec();
|
||||
|
||||
{
|
||||
let mut ipv4_packet = MutableIpv4Packet::new(&mut normalized)
|
||||
.ok_or(UdpBroadcastPacketRejection::MalformedIpv4)?;
|
||||
ipv4_packet.set_source(virtual_ipv4);
|
||||
ipv4_packet.set_destination(destination);
|
||||
ipv4_packet.set_total_length(packet_len as u16);
|
||||
let mut ipv4_packet = Ipv4Packet::new_unchecked(&mut normalized);
|
||||
ipv4_packet.set_src_addr(virtual_ipv4);
|
||||
ipv4_packet.set_dst_addr(destination);
|
||||
ipv4_packet.set_total_len(packet_len as u16);
|
||||
ipv4_packet.set_checksum(0);
|
||||
}
|
||||
|
||||
{
|
||||
let mut udp_packet = MutableUdpPacket::new(&mut normalized[header_len..packet_len])
|
||||
.ok_or(UdpBroadcastPacketRejection::MalformedUdp)?;
|
||||
udp_packet.set_checksum(0);
|
||||
let checksum = udp::ipv4_checksum(&udp_packet.to_immutable(), &virtual_ipv4, &destination);
|
||||
udp_packet.set_checksum(checksum);
|
||||
let mut udp_packet = UdpPacket::new_unchecked(&mut normalized[header_len..packet_len]);
|
||||
udp_packet.fill_checksum(
|
||||
&IpAddress::Ipv4(virtual_ipv4),
|
||||
&IpAddress::Ipv4(destination),
|
||||
);
|
||||
}
|
||||
|
||||
{
|
||||
let mut ipv4_packet = MutableIpv4Packet::new(&mut normalized)
|
||||
.ok_or(UdpBroadcastPacketRejection::MalformedIpv4)?;
|
||||
let checksum = ipv4::checksum(&ipv4_packet.to_immutable());
|
||||
ipv4_packet.set_checksum(checksum);
|
||||
Ipv4Packet::new_unchecked(&mut normalized).fill_checksum();
|
||||
}
|
||||
|
||||
Ok(NormalizedPacket {
|
||||
@@ -373,7 +366,6 @@ impl UdpBroadcastRelayStats {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use pnet_packet::{MutablePacket, Packet};
|
||||
|
||||
fn config() -> BroadcastRelayConfig {
|
||||
BroadcastRelayConfig::new(
|
||||
@@ -385,47 +377,38 @@ mod tests {
|
||||
fn build_udp_packet(src: Ipv4Addr, dst: Ipv4Addr, payload: &[u8]) -> Vec<u8> {
|
||||
let mut packet = vec![0; 20 + 8 + payload.len()];
|
||||
{
|
||||
let mut ipv4_packet = MutableIpv4Packet::new(&mut packet).unwrap();
|
||||
let mut ipv4_packet = Ipv4Packet::new_unchecked(&mut packet);
|
||||
ipv4_packet.set_version(4);
|
||||
ipv4_packet.set_header_length(5);
|
||||
ipv4_packet.set_total_length((20 + 8 + payload.len()) as u16);
|
||||
ipv4_packet.set_ttl(64);
|
||||
ipv4_packet.set_next_level_protocol(IpNextHeaderProtocols::Udp);
|
||||
ipv4_packet.set_source(src);
|
||||
ipv4_packet.set_destination(dst);
|
||||
ipv4_packet.set_header_len(20);
|
||||
ipv4_packet.set_total_len((20 + 8 + payload.len()) as u16);
|
||||
ipv4_packet.set_hop_limit(64);
|
||||
ipv4_packet.set_next_header(IpProtocol::Udp);
|
||||
ipv4_packet.set_src_addr(src);
|
||||
ipv4_packet.set_dst_addr(dst);
|
||||
}
|
||||
|
||||
{
|
||||
let mut udp_packet = MutableUdpPacket::new(&mut packet[20..]).unwrap();
|
||||
udp_packet.set_source(12345);
|
||||
udp_packet.set_destination(37020);
|
||||
udp_packet.set_length((8 + payload.len()) as u16);
|
||||
let mut udp_packet = UdpPacket::new_unchecked(&mut packet[20..]);
|
||||
udp_packet.set_src_port(12345);
|
||||
udp_packet.set_dst_port(37020);
|
||||
udp_packet.set_len((8 + payload.len()) as u16);
|
||||
udp_packet.payload_mut().copy_from_slice(payload);
|
||||
let checksum = udp::ipv4_checksum(&udp_packet.to_immutable(), &src, &dst);
|
||||
udp_packet.set_checksum(checksum);
|
||||
udp_packet.fill_checksum(&IpAddress::Ipv4(src), &IpAddress::Ipv4(dst));
|
||||
}
|
||||
|
||||
{
|
||||
let mut ipv4_packet = MutableIpv4Packet::new(&mut packet).unwrap();
|
||||
let checksum = ipv4::checksum(&ipv4_packet.to_immutable());
|
||||
ipv4_packet.set_checksum(checksum);
|
||||
}
|
||||
Ipv4Packet::new_unchecked(&mut packet).fill_checksum();
|
||||
|
||||
packet
|
||||
}
|
||||
|
||||
fn assert_valid_checksums(packet: &[u8]) {
|
||||
let ipv4_packet = Ipv4Packet::new(packet).unwrap();
|
||||
assert_eq!(ipv4::checksum(&ipv4_packet), ipv4_packet.get_checksum());
|
||||
let udp_packet = UdpPacket::new(ipv4_packet.payload()).unwrap();
|
||||
assert_eq!(
|
||||
udp::ipv4_checksum(
|
||||
&udp_packet,
|
||||
&ipv4_packet.get_source(),
|
||||
&ipv4_packet.get_destination()
|
||||
),
|
||||
udp_packet.get_checksum()
|
||||
);
|
||||
let ipv4_packet = Ipv4Packet::new_checked(packet).unwrap();
|
||||
assert!(ipv4_packet.verify_checksum());
|
||||
let udp_packet = UdpPacket::new_checked(ipv4_packet.payload()).unwrap();
|
||||
assert!(udp_packet.verify_checksum(
|
||||
&IpAddress::Ipv4(ipv4_packet.src_addr()),
|
||||
&IpAddress::Ipv4(ipv4_packet.dst_addr()),
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -433,11 +416,11 @@ mod tests {
|
||||
let packet = build_udp_packet(Ipv4Addr::new(192, 168, 1, 7), Ipv4Addr::BROADCAST, b"hello");
|
||||
|
||||
let normalized = normalize_udp_broadcast_packet(&packet, &config()).unwrap();
|
||||
let ipv4_packet = Ipv4Packet::new(&normalized.packet).unwrap();
|
||||
let ipv4_packet = Ipv4Packet::new_checked(&normalized.packet).unwrap();
|
||||
|
||||
assert_eq!(normalized.destination, Ipv4Addr::BROADCAST);
|
||||
assert_eq!(ipv4_packet.get_source(), Ipv4Addr::new(10, 144, 144, 1));
|
||||
assert_eq!(ipv4_packet.get_destination(), Ipv4Addr::BROADCAST);
|
||||
assert_eq!(ipv4_packet.src_addr(), Ipv4Addr::new(10, 144, 144, 1));
|
||||
assert_eq!(ipv4_packet.dst_addr(), Ipv4Addr::BROADCAST);
|
||||
assert_eq!(&ipv4_packet.payload()[8..], b"hello");
|
||||
assert_valid_checksums(&normalized.packet);
|
||||
}
|
||||
@@ -451,14 +434,11 @@ mod tests {
|
||||
);
|
||||
|
||||
let normalized = normalize_udp_broadcast_packet(&packet, &config()).unwrap();
|
||||
let ipv4_packet = Ipv4Packet::new(&normalized.packet).unwrap();
|
||||
let ipv4_packet = Ipv4Packet::new_checked(&normalized.packet).unwrap();
|
||||
|
||||
assert_eq!(normalized.destination, Ipv4Addr::new(10, 144, 144, 255));
|
||||
assert_eq!(ipv4_packet.get_source(), Ipv4Addr::new(10, 144, 144, 1));
|
||||
assert_eq!(
|
||||
ipv4_packet.get_destination(),
|
||||
Ipv4Addr::new(10, 144, 144, 255)
|
||||
);
|
||||
assert_eq!(ipv4_packet.src_addr(), Ipv4Addr::new(10, 144, 144, 1));
|
||||
assert_eq!(ipv4_packet.dst_addr(), Ipv4Addr::new(10, 144, 144, 255));
|
||||
assert_eq!(&ipv4_packet.payload()[8..], b"directed");
|
||||
assert_valid_checksums(&normalized.packet);
|
||||
}
|
||||
@@ -469,11 +449,11 @@ mod tests {
|
||||
let packet = build_udp_packet(Ipv4Addr::new(192, 168, 1, 7), multicast, b"multicast");
|
||||
|
||||
let normalized = normalize_udp_broadcast_packet(&packet, &config()).unwrap();
|
||||
let ipv4_packet = Ipv4Packet::new(&normalized.packet).unwrap();
|
||||
let ipv4_packet = Ipv4Packet::new_checked(&normalized.packet).unwrap();
|
||||
|
||||
assert_eq!(normalized.destination, multicast);
|
||||
assert_eq!(ipv4_packet.get_source(), Ipv4Addr::new(10, 144, 144, 1));
|
||||
assert_eq!(ipv4_packet.get_destination(), multicast);
|
||||
assert_eq!(ipv4_packet.src_addr(), Ipv4Addr::new(10, 144, 144, 1));
|
||||
assert_eq!(ipv4_packet.dst_addr(), multicast);
|
||||
assert_eq!(&ipv4_packet.payload()[8..], b"multicast");
|
||||
assert_valid_checksums(&normalized.packet);
|
||||
}
|
||||
@@ -502,8 +482,8 @@ mod tests {
|
||||
b"fragment",
|
||||
);
|
||||
{
|
||||
let mut ipv4_packet = MutableIpv4Packet::new(&mut packet).unwrap();
|
||||
ipv4_packet.set_flags(Ipv4Flags::MoreFragments);
|
||||
let mut ipv4_packet = Ipv4Packet::new_unchecked(&mut packet);
|
||||
ipv4_packet.set_more_frags(true);
|
||||
}
|
||||
|
||||
assert_eq!(
|
||||
@@ -540,9 +520,7 @@ mod tests {
|
||||
fn rejects_non_udp_ipv4_packets() {
|
||||
let mut packet =
|
||||
build_udp_packet(Ipv4Addr::new(192, 168, 1, 7), Ipv4Addr::BROADCAST, b"tcp");
|
||||
MutableIpv4Packet::new(&mut packet)
|
||||
.unwrap()
|
||||
.set_next_level_protocol(IpNextHeaderProtocols::Tcp);
|
||||
Ipv4Packet::new_unchecked(&mut packet).set_next_header(IpProtocol::Tcp);
|
||||
|
||||
assert_eq!(
|
||||
normalize_udp_broadcast_packet(&packet, &config()),
|
||||
|
||||
@@ -899,6 +899,9 @@ struct SyncedRouteInfo {
|
||||
foreign_network: DashMap<ForeignNetworkRouteInfoKey, ForeignNetworkRouteInfoEntry>,
|
||||
group_trust_map: DashMap<PeerId, HashMap<String, Vec<u8>>>,
|
||||
group_trust_map_cache: DashMap<PeerId, Arc<Vec<String>>>, // cache for group trust map, should sync with group_trust_map
|
||||
// Serializes every read-modify-write of the derived group maps. Both
|
||||
// proof verification and credential grants update these maps.
|
||||
group_trust_update_lock: parking_lot::Mutex<()>,
|
||||
|
||||
// Aggregated trusted credential pubkeys from all admin nodes
|
||||
// Maps pubkey bytes -> TrustedCredentialPubkey
|
||||
@@ -928,7 +931,8 @@ impl Debug for SyncedRouteInfo {
|
||||
|
||||
#[allow(dead_code)]
|
||||
impl SyncedRouteInfo {
|
||||
fn set_peer_groups(&self, peer_id: PeerId, groups: HashMap<String, Vec<u8>>) {
|
||||
// Must be called with group_trust_update_lock held.
|
||||
fn set_peer_groups_locked(&self, peer_id: PeerId, groups: HashMap<String, Vec<u8>>) {
|
||||
if groups.is_empty() {
|
||||
self.group_trust_map.remove(&peer_id);
|
||||
self.group_trust_map_cache.remove(&peer_id);
|
||||
@@ -941,7 +945,8 @@ impl SyncedRouteInfo {
|
||||
.insert(peer_id, Arc::new(group_names));
|
||||
}
|
||||
|
||||
fn get_proof_groups(&self, peer_id: PeerId) -> HashMap<String, Vec<u8>> {
|
||||
// Must be called with group_trust_update_lock held.
|
||||
fn get_proof_groups_locked(&self, peer_id: PeerId) -> HashMap<String, Vec<u8>> {
|
||||
self.group_trust_map
|
||||
.get(&peer_id)
|
||||
.map(|groups| {
|
||||
@@ -1161,6 +1166,7 @@ impl SyncedRouteInfo {
|
||||
peer_infos: &OrderedHashMap<PeerId, RoutePeerInfo>,
|
||||
all_trusted: &HashMap<Vec<u8>, TrustedCredentialPubkey>,
|
||||
) {
|
||||
let _group_trust_lock = self.group_trust_update_lock.lock();
|
||||
for (_, info) in peer_infos.iter() {
|
||||
if info.noise_static_pubkey.is_empty() {
|
||||
continue;
|
||||
@@ -1169,11 +1175,11 @@ impl SyncedRouteInfo {
|
||||
let Some(credential) = all_trusted.get(&info.noise_static_pubkey) else {
|
||||
continue;
|
||||
};
|
||||
let mut group_map = self.get_proof_groups(info.peer_id);
|
||||
let mut group_map = self.get_proof_groups_locked(info.peer_id);
|
||||
for group in &credential.groups {
|
||||
group_map.entry(group.clone()).or_default();
|
||||
}
|
||||
self.set_peer_groups(info.peer_id, group_map);
|
||||
self.set_peer_groups_locked(info.peer_id, group_map);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1262,19 +1268,23 @@ impl SyncedRouteInfo {
|
||||
}
|
||||
}
|
||||
|
||||
{
|
||||
let _group_trust_lock = self.group_trust_update_lock.lock();
|
||||
for peer_id in &peer_ids {
|
||||
self.group_trust_map.remove(peer_id);
|
||||
self.group_trust_map_cache.remove(peer_id);
|
||||
}
|
||||
shrink_dashmap(&self.group_trust_map, None);
|
||||
shrink_dashmap(&self.group_trust_map_cache, None);
|
||||
}
|
||||
for peer_id in &peer_ids {
|
||||
self.raw_peer_infos.remove(peer_id);
|
||||
self.group_trust_map.remove(peer_id);
|
||||
self.group_trust_map_cache.remove(peer_id);
|
||||
}
|
||||
self.foreign_network
|
||||
.retain(|k, _| !peer_ids.contains(&k.peer_id));
|
||||
|
||||
shrink_dashmap(&self.raw_peer_infos, None);
|
||||
shrink_dashmap(&self.foreign_network, None);
|
||||
shrink_dashmap(&self.group_trust_map, None);
|
||||
shrink_dashmap(&self.group_trust_map_cache, None);
|
||||
|
||||
self.version.inc();
|
||||
}
|
||||
|
||||
@@ -1652,6 +1662,7 @@ impl SyncedRouteInfo {
|
||||
local_group_declarations: &[PeerGroupIdentity],
|
||||
trust_admin_groups_without_proof: bool,
|
||||
) {
|
||||
let _group_trust_lock = self.group_trust_update_lock.lock();
|
||||
let local_group_declarations = local_group_declarations
|
||||
.iter()
|
||||
.map(|g| (g.group_name.as_str(), g.group_secret.as_str()))
|
||||
@@ -1719,14 +1730,59 @@ impl SyncedRouteInfo {
|
||||
}
|
||||
}
|
||||
|
||||
fn verify_and_update_current_group_trusts(
|
||||
&self,
|
||||
received_peer_infos: &[RoutePeerInfo],
|
||||
local_group_declarations: &[PeerGroupIdentity],
|
||||
trust_admin_groups_without_proof: bool,
|
||||
) {
|
||||
let peer_ids: HashSet<_> = received_peer_infos
|
||||
.iter()
|
||||
.map(|info| info.peer_id)
|
||||
.collect();
|
||||
let peer_infos_guard = self.peer_infos.read();
|
||||
let current_peer_infos: Vec<_> = peer_ids
|
||||
.iter()
|
||||
.filter_map(|peer_id| peer_infos_guard.get(peer_id).cloned())
|
||||
.collect();
|
||||
self.verify_and_update_group_trusts(
|
||||
¤t_peer_infos,
|
||||
local_group_declarations,
|
||||
trust_admin_groups_without_proof,
|
||||
);
|
||||
// Keep the authoritative peer-info snapshot locked until its derived
|
||||
// ACL cache has been updated, so another sync cannot interleave a
|
||||
// newer peer version and then be overwritten by this one.
|
||||
drop(peer_infos_guard);
|
||||
}
|
||||
|
||||
fn verify_and_update_all_current_group_trusts(
|
||||
&self,
|
||||
local_group_declarations: &[PeerGroupIdentity],
|
||||
trust_admin_groups_without_proof: bool,
|
||||
) {
|
||||
let peer_infos_guard = self.peer_infos.read();
|
||||
let current_peer_infos: Vec<_> = peer_infos_guard
|
||||
.iter()
|
||||
.map(|(_, info)| info.clone())
|
||||
.collect();
|
||||
self.verify_and_update_group_trusts(
|
||||
¤t_peer_infos,
|
||||
local_group_declarations,
|
||||
trust_admin_groups_without_proof,
|
||||
);
|
||||
drop(peer_infos_guard);
|
||||
}
|
||||
|
||||
fn update_my_group_trusts(&self, my_peer_id: PeerId, groups: &[PeerGroupInfo]) {
|
||||
let _group_trust_lock = self.group_trust_update_lock.lock();
|
||||
let mut my_group_map = HashMap::new();
|
||||
|
||||
for group in groups.iter() {
|
||||
my_group_map.insert(group.group_name.clone(), group.group_proof.clone());
|
||||
}
|
||||
|
||||
self.set_peer_groups(my_peer_id, my_group_map);
|
||||
self.set_peer_groups_locked(my_peer_id, my_group_map);
|
||||
}
|
||||
|
||||
/// Collect trusted credential pubkeys from admin nodes (network_secret holders)
|
||||
@@ -1835,6 +1891,13 @@ type SessionId = u64;
|
||||
|
||||
type AtomicSessionId = atomic_shim::AtomicU64;
|
||||
|
||||
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
||||
struct SyncRequestSnapshot {
|
||||
my_session_id: SessionId,
|
||||
state_revision: u64,
|
||||
is_initiator: bool,
|
||||
}
|
||||
|
||||
struct SessionTask {
|
||||
my_peer_id: PeerId,
|
||||
task: Arc<std::sync::Mutex<Option<JoinHandle<()>>>>,
|
||||
@@ -1921,6 +1984,13 @@ impl VersionAndTouchTime {
|
||||
}
|
||||
}
|
||||
|
||||
// A responder session sends no keepalives when there is no new route data,
|
||||
// so the only sign that the initiator still owns the session is its
|
||||
// periodic inbound syncs (forced at least every ~10s by session_task).
|
||||
// Treat a longer silence as the initiator having lost this session (e.g.
|
||||
// restarted) without telling us.
|
||||
const INITIATOR_SESSION_LIVENESS_TIMEOUT: Duration = Duration::from_secs(45);
|
||||
|
||||
// if we need to sync route info with one peer, we create a SyncRouteSession with that peer.
|
||||
#[derive(Debug)]
|
||||
#[allow(dead_code)]
|
||||
@@ -1937,6 +2007,11 @@ struct SyncRouteSession {
|
||||
|
||||
last_sync_succ_timestamp: AtomicCell<Option<SystemTime>>,
|
||||
|
||||
// Last time any sync interaction (inbound request or successful
|
||||
// response) confirmed the peer still holds this session. Drives the
|
||||
// responder liveness timeout; initialized at session creation.
|
||||
last_contact_instant: AtomicCell<Instant>,
|
||||
|
||||
my_session_id: AtomicSessionId,
|
||||
dst_session_id: AtomicSessionId,
|
||||
|
||||
@@ -1944,6 +2019,11 @@ struct SyncRouteSession {
|
||||
we_are_initiator: AtomicBool,
|
||||
dst_is_initiator: AtomicBool,
|
||||
|
||||
// Serializes the compound session state observed by an outbound RPC.
|
||||
// A response may only commit while this revision still matches the
|
||||
// request snapshot captured before the await.
|
||||
state_revision: AtomicU64,
|
||||
|
||||
need_sync_initiator_info: AtomicBool,
|
||||
|
||||
rpc_tx_count: AtomicU32,
|
||||
@@ -1968,11 +2048,14 @@ impl SyncRouteSession {
|
||||
|
||||
last_sync_succ_timestamp: AtomicCell::new(None),
|
||||
|
||||
last_contact_instant: AtomicCell::new(Instant::now()),
|
||||
|
||||
my_session_id: AtomicSessionId::new(rand::random()),
|
||||
dst_session_id: AtomicSessionId::new(0),
|
||||
|
||||
we_are_initiator: AtomicBool::new(false),
|
||||
dst_is_initiator: AtomicBool::new(false),
|
||||
state_revision: AtomicU64::new(0),
|
||||
|
||||
need_sync_initiator_info: AtomicBool::new(false),
|
||||
|
||||
@@ -2111,24 +2194,117 @@ impl SyncRouteSession {
|
||||
}
|
||||
|
||||
fn update_initiator_flag(&self, is_initiator: bool) {
|
||||
self.we_are_initiator.store(is_initiator, Ordering::Relaxed);
|
||||
let _session_lock = self.lock.lock();
|
||||
if self.we_are_initiator.load(Ordering::Relaxed) != is_initiator {
|
||||
self.we_are_initiator.store(is_initiator, Ordering::Relaxed);
|
||||
self.state_revision.fetch_add(1, Ordering::Relaxed);
|
||||
}
|
||||
self.need_sync_initiator_info.store(true, Ordering::Relaxed);
|
||||
}
|
||||
|
||||
// return whether session id is updated
|
||||
fn update_dst_session_id(&self, session_id: SessionId) {
|
||||
if session_id != self.dst_session_id.load(Ordering::Relaxed) {
|
||||
// Must be called with the session lock held.
|
||||
fn update_remote_state_locked(&self, session_id: SessionId, is_initiator: bool) {
|
||||
let session_id_changed = session_id != self.dst_session_id.load(Ordering::Relaxed);
|
||||
let initiator_changed = is_initiator != self.dst_is_initiator.load(Ordering::Relaxed);
|
||||
|
||||
if session_id_changed {
|
||||
tracing::warn!(?self, ?session_id, "session id mismatch, clear saved info.");
|
||||
self.dst_session_id.store(session_id, Ordering::Relaxed);
|
||||
self.dst_saved_conn_info_version.clear();
|
||||
self.dst_saved_peer_info_versions.clear();
|
||||
self.dst_saved_foreign_network_versions.clear();
|
||||
|
||||
// update_dst_session_id is always called with session lock held, so clear
|
||||
// last_sync_succ_timestamp and unreachable_peers non-atomic is safe.
|
||||
self.last_sync_succ_timestamp.store(None);
|
||||
self.unreachable_peers_for_peer_info.lock().clear();
|
||||
self.unreachable_peers_for_conn_info.lock().clear();
|
||||
}
|
||||
|
||||
if initiator_changed {
|
||||
self.dst_is_initiator.store(is_initiator, Ordering::Relaxed);
|
||||
}
|
||||
|
||||
if session_id_changed || initiator_changed {
|
||||
self.state_revision.fetch_add(1, Ordering::Relaxed);
|
||||
}
|
||||
}
|
||||
|
||||
// Must be called with the session lock held. A different generation may
|
||||
// only take over through an initiator request after the previous remote
|
||||
// initiator role has been relinquished or expired.
|
||||
fn admit_inbound_locked(&self, session_id: SessionId, is_initiator: bool) -> bool {
|
||||
let current_session_id = self.dst_session_id.load(Ordering::Relaxed);
|
||||
if session_id == current_session_id {
|
||||
self.update_remote_state_locked(session_id, is_initiator);
|
||||
return true;
|
||||
}
|
||||
|
||||
if current_session_id == 0 && !is_initiator && self.we_are_initiator.load(Ordering::Relaxed)
|
||||
{
|
||||
// The responder can send its first reverse sync before our
|
||||
// initiating RPC response arrives and teaches us its session ID.
|
||||
self.update_remote_state_locked(session_id, false);
|
||||
return true;
|
||||
}
|
||||
|
||||
if !is_initiator || self.dst_is_initiator.load(Ordering::Relaxed) {
|
||||
return false;
|
||||
}
|
||||
|
||||
self.update_remote_state_locked(session_id, true);
|
||||
true
|
||||
}
|
||||
|
||||
// Must be called with the session lock held.
|
||||
fn request_snapshot_locked(&self) -> SyncRequestSnapshot {
|
||||
SyncRequestSnapshot {
|
||||
my_session_id: self.my_session_id.load(Ordering::Relaxed),
|
||||
state_revision: self.state_revision.load(Ordering::Relaxed),
|
||||
is_initiator: self.we_are_initiator.load(Ordering::Relaxed),
|
||||
}
|
||||
}
|
||||
|
||||
// Must be called with the session lock held.
|
||||
fn request_is_current_locked(&self, snapshot: SyncRequestSnapshot) -> bool {
|
||||
snapshot.my_session_id == self.my_session_id.load(Ordering::Relaxed)
|
||||
&& snapshot.state_revision == self.state_revision.load(Ordering::Relaxed)
|
||||
&& snapshot.is_initiator == self.we_are_initiator.load(Ordering::Relaxed)
|
||||
}
|
||||
|
||||
// Must be called with the session lock held.
|
||||
fn clear_dst_initiator_if_session_unchanged_locked(
|
||||
&self,
|
||||
expected_dst_session_id: SessionId,
|
||||
) -> bool {
|
||||
if self.dst_session_id.load(Ordering::Relaxed) != expected_dst_session_id {
|
||||
return false;
|
||||
}
|
||||
|
||||
self.update_remote_state_locked(expected_dst_session_id, false);
|
||||
true
|
||||
}
|
||||
|
||||
// Responder sessions have no outbound keepalive, so the only sign that
|
||||
// the initiator still owns the session is its periodic inbound syncs.
|
||||
// If no sync interaction arrived within INITIATOR_SESSION_LIVENESS_TIMEOUT,
|
||||
// assume the initiator lost this session without telling us and clear the
|
||||
// stale role so the election loop can re-establish the edge. Only the
|
||||
// flag is cleared; the election loop takes over from there. Returns true
|
||||
// if the initiator role was cleared.
|
||||
fn clear_stale_dst_initiator(&self) -> bool {
|
||||
let _session_lock = self.lock.lock();
|
||||
if !self.dst_is_initiator.load(Ordering::Relaxed)
|
||||
|| self.we_are_initiator.load(Ordering::Relaxed)
|
||||
{
|
||||
return false;
|
||||
}
|
||||
|
||||
if self.last_contact_instant.load().elapsed() < INITIATOR_SESSION_LIVENESS_TIMEOUT {
|
||||
return false;
|
||||
}
|
||||
|
||||
let session_id = self.dst_session_id.load(Ordering::Relaxed);
|
||||
self.update_remote_state_locked(session_id, false);
|
||||
true
|
||||
}
|
||||
|
||||
fn clean_dst_saved_map(&self) {
|
||||
@@ -2256,6 +2432,7 @@ impl PeerRouteServiceImpl {
|
||||
foreign_network: DashMap::new(),
|
||||
group_trust_map: DashMap::new(),
|
||||
group_trust_map_cache: DashMap::new(),
|
||||
group_trust_update_lock: parking_lot::Mutex::new(()),
|
||||
trusted_credential_pubkeys: DashMap::new(),
|
||||
non_reusable_credential_owners: DashMap::new(),
|
||||
suppressed_non_reusable_credential_peers: DashMap::new(),
|
||||
@@ -2312,8 +2489,32 @@ impl PeerRouteServiceImpl {
|
||||
self.sessions.get(&dst_peer_id).map(|x| x.value().clone())
|
||||
}
|
||||
|
||||
fn is_current_session(&self, dst_peer_id: PeerId, expected: &Arc<SyncRouteSession>) -> bool {
|
||||
self.sessions
|
||||
.get(&dst_peer_id)
|
||||
.is_some_and(|current| Arc::ptr_eq(current.value(), expected))
|
||||
}
|
||||
|
||||
// Must be called with the expected session lock held. remove_session
|
||||
// takes the same lock before changing the map, so this check remains
|
||||
// valid until the caller finishes committing the RPC result.
|
||||
fn sync_request_is_current_locked(
|
||||
&self,
|
||||
dst_peer_id: PeerId,
|
||||
expected: &Arc<SyncRouteSession>,
|
||||
snapshot: SyncRequestSnapshot,
|
||||
) -> bool {
|
||||
self.is_current_session(dst_peer_id, expected)
|
||||
&& expected.request_is_current_locked(snapshot)
|
||||
}
|
||||
|
||||
fn remove_session(&self, dst_peer_id: PeerId) {
|
||||
self.sessions.remove(&dst_peer_id);
|
||||
let Some(session) = self.get_session(dst_peer_id) else {
|
||||
return;
|
||||
};
|
||||
let _session_lock = session.lock.lock();
|
||||
self.sessions
|
||||
.remove_if(&dst_peer_id, |_, current| Arc::ptr_eq(current, &session));
|
||||
shrink_dashmap(&self.sessions, None);
|
||||
}
|
||||
|
||||
@@ -2804,18 +3005,11 @@ impl PeerRouteServiceImpl {
|
||||
let trust_admin_groups_without_proof =
|
||||
self.context.network_identity().network_secret.is_none();
|
||||
|
||||
let peer_infos: Vec<_> = self
|
||||
.synced_route_info
|
||||
.peer_infos
|
||||
.read()
|
||||
.iter()
|
||||
.map(|(_, info)| info.clone())
|
||||
.collect();
|
||||
self.synced_route_info.verify_and_update_group_trusts(
|
||||
&peer_infos,
|
||||
&self.context.acl_group_declarations(),
|
||||
trust_admin_groups_without_proof,
|
||||
);
|
||||
self.synced_route_info
|
||||
.verify_and_update_all_current_group_trusts(
|
||||
&self.context.acl_group_declarations(),
|
||||
trust_admin_groups_without_proof,
|
||||
);
|
||||
|
||||
let untrusted = self.refresh_credential_trusts_with_current_topology();
|
||||
self.disconnect_untrusted_peers(&untrusted).await;
|
||||
@@ -3028,6 +3222,8 @@ impl PeerRouteServiceImpl {
|
||||
|
||||
let next_last_sync_succ_timestamp =
|
||||
self.synced_route_info.get_next_last_sync_succ_timestamp();
|
||||
let request_snapshot = session.request_snapshot_locked();
|
||||
let expected_dst_session_id = session.dst_session_id.load(Ordering::Relaxed);
|
||||
let (peer_infos, conn_info, foreign_network) =
|
||||
self.build_sync_request(&session, dst_peer_id);
|
||||
if peer_infos.is_none()
|
||||
@@ -3064,8 +3260,8 @@ impl PeerRouteServiceImpl {
|
||||
|
||||
let sync_route_info_req = SyncRouteInfoRequest {
|
||||
my_peer_id,
|
||||
my_session_id: session.my_session_id.load(Ordering::Relaxed),
|
||||
is_initiator: session.we_are_initiator.load(Ordering::Relaxed),
|
||||
my_session_id: request_snapshot.my_session_id,
|
||||
is_initiator: request_snapshot.is_initiator,
|
||||
peer_infos: peer_infos.clone().map(|x| RoutePeerInfos { items: x }),
|
||||
conn_info: conn_info.clone(),
|
||||
foreign_network_infos: foreign_network.clone(),
|
||||
@@ -3100,6 +3296,16 @@ impl PeerRouteServiceImpl {
|
||||
next_last_sync_succ_timestamp
|
||||
);
|
||||
|
||||
if !self.sync_request_is_current_locked(dst_peer_id, &session, request_snapshot) {
|
||||
tracing::debug!(
|
||||
?my_peer_id,
|
||||
?dst_peer_id,
|
||||
?request_snapshot,
|
||||
"discard stale route sync response"
|
||||
);
|
||||
return false;
|
||||
}
|
||||
|
||||
match ret.as_ref() {
|
||||
Err(e) => {
|
||||
tracing::error!(
|
||||
@@ -3115,7 +3321,17 @@ impl PeerRouteServiceImpl {
|
||||
}
|
||||
Ok(resp) => {
|
||||
if let Some(err) = resp.error {
|
||||
if err == Error::DuplicatePeerId as i32 {
|
||||
if err == Error::Stopped as i32 && !sync_route_info_req.is_initiator {
|
||||
let cleared = session.clear_dst_initiator_if_session_unchanged_locked(
|
||||
expected_dst_session_id,
|
||||
);
|
||||
tracing::debug!(
|
||||
?my_peer_id,
|
||||
?dst_peer_id,
|
||||
?cleared,
|
||||
"stale non-initiator route sync rejected"
|
||||
);
|
||||
} else if err == Error::DuplicatePeerId as i32 {
|
||||
if !self.context.feature_flags().is_public_server {
|
||||
panic!("duplicate peer id");
|
||||
}
|
||||
@@ -3127,12 +3343,9 @@ impl PeerRouteServiceImpl {
|
||||
}
|
||||
} else {
|
||||
session.rpc_tx_count.fetch_add(1, Ordering::Relaxed);
|
||||
session.last_contact_instant.store(Instant::now());
|
||||
|
||||
session
|
||||
.dst_is_initiator
|
||||
.store(resp.is_initiator, Ordering::Relaxed);
|
||||
|
||||
session.update_dst_session_id(resp.session_id);
|
||||
session.update_remote_state_locked(resp.session_id, resp.is_initiator);
|
||||
|
||||
if let Some(peer_infos) = &peer_infos {
|
||||
session.update_dst_saved_peer_info_version(peer_infos, dst_peer_id);
|
||||
@@ -3399,6 +3612,23 @@ impl RouteSessionManager {
|
||||
}
|
||||
}
|
||||
|
||||
// Detect responder sessions whose initiator silently lost the
|
||||
// session (e.g. restarted and elected someone else). Without
|
||||
// this, a responder with no new route data could wait forever
|
||||
// for an initiator that no longer syncs with us. Clearing the
|
||||
// stale role makes the peer an initiator candidate again below.
|
||||
for peer_id in session_peers.iter() {
|
||||
if let Some(session) = service_impl.get_session(*peer_id)
|
||||
&& session.clear_stale_dst_initiator()
|
||||
{
|
||||
tracing::warn!(
|
||||
?peer_id,
|
||||
my_peer_id = ?service_impl.my_peer_id,
|
||||
"initiator route sync liveness timeout, clearing stale initiator role"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
// find peer_ids that are not initiators.
|
||||
let mut initiator_candidates = Vec::new();
|
||||
for peer_id in peers.iter().copied() {
|
||||
@@ -3573,7 +3803,17 @@ impl RouteSessionManager {
|
||||
}
|
||||
|
||||
let my_peer_id = service_impl.my_peer_id;
|
||||
let session = self.get_or_start_session(from_peer_id)?;
|
||||
let session = if let Some(session) = service_impl.get_session(from_peer_id) {
|
||||
session
|
||||
} else if is_initiator {
|
||||
self.get_or_start_session(from_peer_id)?
|
||||
} else {
|
||||
tracing::debug!(
|
||||
?from_peer_id,
|
||||
"ignore stale route sync from non-initiator without a session"
|
||||
);
|
||||
return Err(Error::Stopped);
|
||||
};
|
||||
|
||||
let from_identity_type = service_impl
|
||||
.get_peer_identity_type_from_interface(from_peer_id)
|
||||
@@ -3599,9 +3839,20 @@ impl RouteSessionManager {
|
||||
|
||||
let _session_lock = session.lock.lock();
|
||||
|
||||
session.rpc_rx_count.fetch_add(1, Ordering::Relaxed);
|
||||
if !service_impl.is_current_session(from_peer_id, &session)
|
||||
|| !session.admit_inbound_locked(from_session_id, is_initiator)
|
||||
{
|
||||
tracing::debug!(
|
||||
?from_peer_id,
|
||||
?from_session_id,
|
||||
?is_initiator,
|
||||
"ignore route sync from a stale session generation"
|
||||
);
|
||||
return Err(Error::Stopped);
|
||||
}
|
||||
|
||||
session.update_dst_session_id(from_session_id);
|
||||
session.rpc_rx_count.fetch_add(1, Ordering::Relaxed);
|
||||
session.last_contact_instant.store(Instant::now());
|
||||
|
||||
let mut need_update_route_table = false;
|
||||
let mut untrusted_peers = Vec::new();
|
||||
@@ -3637,7 +3888,7 @@ impl RouteSessionManager {
|
||||
)?;
|
||||
service_impl
|
||||
.synced_route_info
|
||||
.verify_and_update_group_trusts(
|
||||
.verify_and_update_current_group_trusts(
|
||||
pi,
|
||||
&service_impl.context.acl_group_declarations(),
|
||||
trust_admin_groups_without_proof,
|
||||
@@ -3694,9 +3945,6 @@ impl RouteSessionManager {
|
||||
service_impl.route_table
|
||||
);
|
||||
|
||||
session
|
||||
.dst_is_initiator
|
||||
.store(is_initiator, Ordering::Relaxed);
|
||||
let is_initiator = session.we_are_initiator.load(Ordering::Relaxed);
|
||||
let session_id = session.my_session_id.load(Ordering::Relaxed);
|
||||
|
||||
@@ -4385,6 +4633,29 @@ mod tests {
|
||||
PeerRouteServiceImpl::new(my_peer_id, Arc::new(NoopPeerContext::default()))
|
||||
}
|
||||
|
||||
async fn test_route_with_admin_peer(
|
||||
context: ArcPeerContext,
|
||||
) -> (Arc<PeerRoute>, Arc<PeerRpcManager>) {
|
||||
let peer_rpc = Arc::new(PeerRpcManager::new(TestPeerRpcTransport));
|
||||
let route = PeerRoute::new(
|
||||
1,
|
||||
context,
|
||||
Arc::new(TestPublicIpv6Runtime),
|
||||
peer_rpc.clone(),
|
||||
);
|
||||
*route.service_impl.interface.lock().await = Some(Box::new(CountingInterface {
|
||||
my_peer_id: 1,
|
||||
peers: Arc::new(Mutex::new(vec![2])),
|
||||
peer_identity_types: Arc::new(Mutex::new(HashMap::from([(
|
||||
2,
|
||||
Some(PeerIdentityType::Admin),
|
||||
)]))),
|
||||
list_peers_calls: Arc::new(AtomicU32::new(0)),
|
||||
get_peer_identity_type_calls: Arc::new(AtomicU32::new(0)),
|
||||
}));
|
||||
(route, peer_rpc)
|
||||
}
|
||||
|
||||
fn peer(peer_id: PeerId) -> OspfPeerInfo {
|
||||
OspfPeerInfo {
|
||||
peer_id,
|
||||
@@ -4522,6 +4793,381 @@ mod tests {
|
||||
assert_eq!(get_peer_identity_type_calls.load(Ordering::Relaxed), 4);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stopped_sync_clears_matching_remote_initiator() {
|
||||
let session = SyncRouteSession::new(1, 2);
|
||||
let _session_lock = session.lock.lock();
|
||||
session.update_remote_state_locked(10, true);
|
||||
|
||||
assert!(session.clear_dst_initiator_if_session_unchanged_locked(10));
|
||||
assert!(!session.dst_is_initiator.load(Ordering::Relaxed));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stopped_sync_preserves_newer_remote_initiator() {
|
||||
let session = SyncRouteSession::new(1, 2);
|
||||
let _session_lock = session.lock.lock();
|
||||
session.update_remote_state_locked(10, true);
|
||||
|
||||
session.update_remote_state_locked(11, true);
|
||||
|
||||
assert!(!session.clear_dst_initiator_if_session_unchanged_locked(10));
|
||||
assert!(session.dst_is_initiator.load(Ordering::Relaxed));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn remote_state_change_invalidates_outbound_snapshot() {
|
||||
let service_impl = Arc::new(test_service_impl(1));
|
||||
let session = service_impl.get_or_create_session(2);
|
||||
let _session_lock = session.lock.lock();
|
||||
session.update_remote_state_locked(10, false);
|
||||
let snapshot = session.request_snapshot_locked();
|
||||
|
||||
assert!(session.admit_inbound_locked(20, true));
|
||||
assert!(!service_impl.sync_request_is_current_locked(2, &session, snapshot));
|
||||
assert_eq!(session.dst_session_id.load(Ordering::Relaxed), 20);
|
||||
assert!(session.dst_is_initiator.load(Ordering::Relaxed));
|
||||
assert_eq!(session.rpc_tx_count.load(Ordering::Relaxed), 0);
|
||||
assert!(session.dst_saved_peer_info_versions.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn replaced_session_invalidates_outbound_snapshot() {
|
||||
let service_impl = Arc::new(test_service_impl(1));
|
||||
let old_session = service_impl.get_or_create_session(2);
|
||||
let snapshot = {
|
||||
let _session_lock = old_session.lock.lock();
|
||||
old_session.request_snapshot_locked()
|
||||
};
|
||||
|
||||
service_impl.remove_session(2);
|
||||
let new_session = service_impl.get_or_create_session(2);
|
||||
let _old_session_lock = old_session.lock.lock();
|
||||
|
||||
assert!(!Arc::ptr_eq(&old_session, &new_session));
|
||||
assert!(!service_impl.sync_request_is_current_locked(2, &old_session, snapshot));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn active_remote_initiator_rejects_different_generation() {
|
||||
let session = SyncRouteSession::new(1, 2);
|
||||
{
|
||||
let _session_lock = session.lock.lock();
|
||||
assert!(session.admit_inbound_locked(20, true));
|
||||
assert!(!session.admit_inbound_locked(10, true));
|
||||
assert_eq!(session.dst_session_id.load(Ordering::Relaxed), 20);
|
||||
}
|
||||
|
||||
session
|
||||
.last_contact_instant
|
||||
.store(Instant::now() - INITIATOR_SESSION_LIVENESS_TIMEOUT - Duration::from_secs(1));
|
||||
assert!(session.clear_stale_dst_initiator());
|
||||
|
||||
let _session_lock = session.lock.lock();
|
||||
assert!(session.admit_inbound_locked(10, true));
|
||||
assert_eq!(session.dst_session_id.load(Ordering::Relaxed), 10);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn local_initiator_accepts_initial_responder_generation() {
|
||||
let session = SyncRouteSession::new(1, 2);
|
||||
session.update_initiator_flag(true);
|
||||
|
||||
let _session_lock = session.lock.lock();
|
||||
assert!(session.admit_inbound_locked(20, false));
|
||||
assert_eq!(session.dst_session_id.load(Ordering::Relaxed), 20);
|
||||
assert!(!session.dst_is_initiator.load(Ordering::Relaxed));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stale_responder_session_clears_initiator_after_liveness_timeout() {
|
||||
let session = SyncRouteSession::new(1, 2);
|
||||
{
|
||||
let _session_lock = session.lock.lock();
|
||||
session.update_remote_state_locked(10, true);
|
||||
}
|
||||
session
|
||||
.last_contact_instant
|
||||
.store(Instant::now() - INITIATOR_SESSION_LIVENESS_TIMEOUT - Duration::from_secs(1));
|
||||
|
||||
assert!(session.clear_stale_dst_initiator());
|
||||
assert!(!session.dst_is_initiator.load(Ordering::Relaxed));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn healthy_responder_session_keeps_initiator() {
|
||||
let session = SyncRouteSession::new(1, 2);
|
||||
{
|
||||
let _session_lock = session.lock.lock();
|
||||
session.update_remote_state_locked(10, true);
|
||||
}
|
||||
|
||||
assert!(!session.clear_stale_dst_initiator());
|
||||
assert!(session.dst_is_initiator.load(Ordering::Relaxed));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn initiator_session_ignores_liveness_timeout() {
|
||||
let session = SyncRouteSession::new(1, 2);
|
||||
session.update_initiator_flag(true);
|
||||
{
|
||||
let _session_lock = session.lock.lock();
|
||||
session.update_remote_state_locked(10, true);
|
||||
}
|
||||
session
|
||||
.last_contact_instant
|
||||
.store(Instant::now() - INITIATOR_SESSION_LIVENESS_TIMEOUT - Duration::from_secs(1));
|
||||
|
||||
assert!(!session.clear_stale_dst_initiator());
|
||||
assert!(session.dst_is_initiator.load(Ordering::Relaxed));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn initiator_sync_creates_session_and_marks_contact() {
|
||||
let peer_rpc = Arc::new(PeerRpcManager::new(TestPeerRpcTransport));
|
||||
let route = PeerRoute::new(
|
||||
1,
|
||||
Arc::new(NoopPeerContext::default()),
|
||||
Arc::new(TestPublicIpv6Runtime),
|
||||
peer_rpc,
|
||||
);
|
||||
let peers = Arc::new(Mutex::new(vec![2]));
|
||||
*route.service_impl.interface.lock().await = Some(Box::new(CountingInterface {
|
||||
my_peer_id: 1,
|
||||
peers,
|
||||
peer_identity_types: Arc::new(Mutex::new(HashMap::from([(
|
||||
2,
|
||||
Some(PeerIdentityType::Admin),
|
||||
)]))),
|
||||
list_peers_calls: Arc::new(AtomicU32::new(0)),
|
||||
get_peer_identity_type_calls: Arc::new(AtomicU32::new(0)),
|
||||
}));
|
||||
|
||||
route
|
||||
.session_mgr
|
||||
.do_sync_route_info(2, 1, true, None, None, None, None)
|
||||
.await
|
||||
.expect("initiator sync should succeed");
|
||||
|
||||
let session = route
|
||||
.service_impl
|
||||
.get_session(2)
|
||||
.expect("initiator sync should create the session");
|
||||
assert!(session.dst_is_initiator.load(Ordering::Relaxed));
|
||||
assert!(
|
||||
session.last_contact_instant.load().elapsed() < Duration::from_secs(5),
|
||||
"inbound sync should refresh the liveness timestamp"
|
||||
);
|
||||
|
||||
route.stop().await;
|
||||
assert!(route.service_impl.sessions.is_empty());
|
||||
assert_eq!(route.task_count(), 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn stale_non_initiator_sync_does_not_create_session() {
|
||||
let peer_rpc = Arc::new(PeerRpcManager::new(TestPeerRpcTransport));
|
||||
let route = PeerRoute::new(
|
||||
1,
|
||||
Arc::new(NoopPeerContext::default()),
|
||||
Arc::new(TestPublicIpv6Runtime),
|
||||
peer_rpc,
|
||||
);
|
||||
let peers = Arc::new(Mutex::new(vec![2]));
|
||||
*route.service_impl.interface.lock().await = Some(Box::new(CountingInterface {
|
||||
my_peer_id: 1,
|
||||
peers,
|
||||
peer_identity_types: Arc::new(Mutex::new(HashMap::from([(
|
||||
2,
|
||||
Some(PeerIdentityType::Admin),
|
||||
)]))),
|
||||
list_peers_calls: Arc::new(AtomicU32::new(0)),
|
||||
get_peer_identity_type_calls: Arc::new(AtomicU32::new(0)),
|
||||
}));
|
||||
|
||||
let result = route
|
||||
.session_mgr
|
||||
.do_sync_route_info(2, 1, false, None, None, None, None)
|
||||
.await;
|
||||
|
||||
assert!(matches!(result, Err(Error::Stopped)));
|
||||
assert!(route.service_impl.sessions.is_empty());
|
||||
assert_eq!(route.task_count(), 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn stale_non_initiator_sync_preserves_newer_session_generation() {
|
||||
let (route, _peer_rpc) =
|
||||
test_route_with_admin_peer(Arc::new(NoopPeerContext::default())).await;
|
||||
let session = route.service_impl.get_or_create_session(2);
|
||||
|
||||
route
|
||||
.session_mgr
|
||||
.do_sync_route_info(2, 22, true, None, None, None, None)
|
||||
.await
|
||||
.expect("new initiator generation should be accepted");
|
||||
let contact = session.last_contact_instant.load();
|
||||
let rx_count = session.rpc_rx_count.load(Ordering::Relaxed);
|
||||
|
||||
let stale_peer_info = RoutePeerInfo {
|
||||
peer_id: 3,
|
||||
version: 1,
|
||||
..Default::default()
|
||||
};
|
||||
let result = route
|
||||
.session_mgr
|
||||
.do_sync_route_info(
|
||||
2,
|
||||
11,
|
||||
false,
|
||||
Some(vec![stale_peer_info.clone()]),
|
||||
Some(vec![raw_route_peer_info(&stale_peer_info)]),
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
assert!(matches!(result, Err(Error::Stopped)));
|
||||
assert_eq!(session.dst_session_id.load(Ordering::Relaxed), 22);
|
||||
assert!(session.dst_is_initiator.load(Ordering::Relaxed));
|
||||
assert_eq!(session.last_contact_instant.load(), contact);
|
||||
assert_eq!(session.rpc_rx_count.load(Ordering::Relaxed), rx_count);
|
||||
assert!(
|
||||
!route
|
||||
.service_impl
|
||||
.synced_route_info
|
||||
.peer_infos
|
||||
.read()
|
||||
.contains_key(&3)
|
||||
);
|
||||
|
||||
route.stop().await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn stale_peer_info_does_not_restore_removed_acl_group() {
|
||||
let context = Arc::new(NoopPeerContext::new(CoreNetworkIdentity::new_credential(
|
||||
"default".to_owned(),
|
||||
)));
|
||||
let (route, _peer_rpc) = test_route_with_admin_peer(context).await;
|
||||
route.service_impl.get_or_create_session(2);
|
||||
|
||||
let peer_info = |version, groups| RoutePeerInfo {
|
||||
peer_id: 3,
|
||||
peer_route_id: 30,
|
||||
version,
|
||||
groups,
|
||||
..Default::default()
|
||||
};
|
||||
let sync_peer_info = |info: RoutePeerInfo| {
|
||||
let raw = raw_route_peer_info(&info);
|
||||
(Some(vec![info]), Some(vec![raw]))
|
||||
};
|
||||
|
||||
let v1 = peer_info(
|
||||
1,
|
||||
vec![PeerGroupInfo {
|
||||
group_name: "legacy".to_owned(),
|
||||
group_proof: Vec::new(),
|
||||
}],
|
||||
);
|
||||
let (peer_infos, raw_peer_infos) = sync_peer_info(v1.clone());
|
||||
route
|
||||
.session_mgr
|
||||
.do_sync_route_info(2, 22, true, peer_infos, raw_peer_infos, None, None)
|
||||
.await
|
||||
.expect("v1 peer info should be accepted");
|
||||
assert_eq!(route.service_impl.get_peer_groups(3).as_ref(), &["legacy"]);
|
||||
|
||||
let v2 = peer_info(2, Vec::new());
|
||||
let (peer_infos, raw_peer_infos) = sync_peer_info(v2);
|
||||
route
|
||||
.session_mgr
|
||||
.do_sync_route_info(2, 22, true, peer_infos, raw_peer_infos, None, None)
|
||||
.await
|
||||
.expect("v2 peer info should be accepted");
|
||||
assert!(route.service_impl.get_peer_groups(3).is_empty());
|
||||
|
||||
let (peer_infos, raw_peer_infos) = sync_peer_info(v1);
|
||||
route
|
||||
.session_mgr
|
||||
.do_sync_route_info(2, 22, true, peer_infos, raw_peer_infos, None, None)
|
||||
.await
|
||||
.expect("stale peer info should be ignored");
|
||||
|
||||
assert_eq!(
|
||||
route
|
||||
.service_impl
|
||||
.synced_route_info
|
||||
.peer_infos
|
||||
.read()
|
||||
.get(&3)
|
||||
.expect("peer info should remain present")
|
||||
.version,
|
||||
2
|
||||
);
|
||||
assert!(route.service_impl.get_peer_groups(3).is_empty());
|
||||
|
||||
route.stop().await;
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn credential_group_refresh_does_not_restore_removed_proof_group() {
|
||||
let service_impl = Arc::new(test_service_impl(1));
|
||||
let mut peer_infos = OrderedHashMap::new();
|
||||
peer_infos.insert(
|
||||
3,
|
||||
RoutePeerInfo {
|
||||
peer_id: 3,
|
||||
version: 3,
|
||||
noise_static_pubkey: vec![3; 32],
|
||||
..Default::default()
|
||||
},
|
||||
);
|
||||
let all_trusted = HashMap::from([(
|
||||
vec![3; 32],
|
||||
TrustedCredentialPubkey {
|
||||
groups: vec!["credential".to_owned()],
|
||||
..Default::default()
|
||||
},
|
||||
)]);
|
||||
|
||||
let group_trust_lock = service_impl
|
||||
.synced_route_info
|
||||
.group_trust_update_lock
|
||||
.lock();
|
||||
service_impl
|
||||
.synced_route_info
|
||||
.set_peer_groups_locked(3, HashMap::from([("legacy".to_owned(), vec![1])]));
|
||||
|
||||
let (started_tx, started_rx) = std::sync::mpsc::channel();
|
||||
let credential_refresh = std::thread::spawn({
|
||||
let service_impl = service_impl.clone();
|
||||
move || {
|
||||
started_tx.send(()).unwrap();
|
||||
service_impl
|
||||
.synced_route_info
|
||||
.update_credential_groups(&peer_infos, &all_trusted);
|
||||
}
|
||||
});
|
||||
started_rx.recv().unwrap();
|
||||
|
||||
service_impl
|
||||
.synced_route_info
|
||||
.set_peer_groups_locked(3, HashMap::new());
|
||||
drop(group_trust_lock);
|
||||
credential_refresh.join().unwrap();
|
||||
|
||||
assert_eq!(
|
||||
service_impl
|
||||
.synced_route_info
|
||||
.group_trust_map
|
||||
.get(&3)
|
||||
.as_deref(),
|
||||
Some(&HashMap::from([("credential".to_owned(), Vec::new())]))
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn stop_waits_for_in_flight_route_sync_before_draining_sessions() {
|
||||
let peer_rpc = Arc::new(PeerRpcManager::new(TestPeerRpcTransport));
|
||||
@@ -4542,7 +5188,7 @@ mod tests {
|
||||
let session_mgr = route.session_mgr.clone();
|
||||
async move {
|
||||
session_mgr
|
||||
.do_sync_route_info(2, 1, false, None, None, None, None)
|
||||
.do_sync_route_info(2, 1, true, None, None, None, None)
|
||||
.await
|
||||
}
|
||||
});
|
||||
|
||||
Reference in New Issue
Block a user