Use zerocopy Ref for header parsing and Arc for XSK packet sharing

Three coordinated changes to reduce per-packet allocations:

1. packet_parser: Replace manual byte-index parsing with typed header
   structs (EthHdr, Ipv4Hdr, Ipv6Hdr, TcpHdr, UdpHdr) using
   #[repr(C, packed)] + zerocopy derives. Ref::from_prefix() validates
   bounds and alignment in one call, eliminating manual index arithmetic
   and unsafe magic-number offset accesses.

2. xsk_manager: ML engine now receives a &[u8] slice directly from UMEM
   before any heap allocation. The single to_vec() copy that follows is
   wrapped in Arc so that Suricata and the forward-TX channel share the
   same allocation. Previously: 2 copies (to_vec + clone) when Suricata
   is active. Now: 1 copy always.

3. suricata/engine: inject() now takes Arc<Vec<u8>>; the internal mirror
   channel is changed to Sender<Arc<Vec<u8>>>. The sendto() call reads
   through the Arc deref with no extra allocation.

https://claude.ai/code/session_0138PxtKH73hqxv7h1oaoSdS
This commit is contained in:
Claude 2026-06-16 06:36:06 +00:00
parent b847374341
commit df610e778e
No known key found for this signature in database
5 changed files with 160 additions and 134 deletions

5
Cargo.lock generated
View File

@ -1546,6 +1546,7 @@ dependencies = [
"url",
"uuid",
"xsk-rs",
"zerocopy 0.8.33",
]
[[package]]
@ -2693,9 +2694,9 @@ dependencies = [
[[package]]
name = "sysinfo"
version = "0.39.3"
version = "0.39.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "21d0d938c10fcda3e897e28aaddf4ab462375d411f4378cd63b1c945f69aba96"
checksum = "14311e7e9a03114cd4b65eedd54e8fed2945e17f08586ae97ef53bc0669f9581"
dependencies = [
"libc",
"memchr",

View File

@ -21,6 +21,7 @@ parking_lot = "0.12.5"
rust-embed = "8.7.2"
serde = { workspace = true }
serde_json = "1.0.143"
zerocopy = { workspace = true }
sysinfo = "0.39.2"
thiserror = "2.0.3"
tokio = { version = "1.40.0", features = ["full", "macros"] }

View File

@ -61,8 +61,8 @@ impl XskManager {
let combined_queue_count = config.combined_queue_count;
for queue_id in 0..combined_queue_count {
let (ingress_to_egress_tx, ingress_to_egress_rx) = bounded(config.channel_size);
let (egress_to_ingress_tx, egress_to_ingress_rx) = bounded(config.channel_size);
let (ingress_to_egress_tx, ingress_to_egress_rx) = bounded::<Arc<Vec<u8>>>(config.channel_size);
let (egress_to_ingress_tx, egress_to_ingress_rx) = bounded::<Arc<Vec<u8>>>(config.channel_size);
let ingress_xsk = XskPair::new(
config.clone(),
@ -258,8 +258,8 @@ impl XskPair {
pub fn run(
mut self,
forward_tx: Sender<Vec<u8>>,
forward_rx: Receiver<Vec<u8>>,
forward_tx: Sender<Arc<Vec<u8>>>,
forward_rx: Receiver<Arc<Vec<u8>>>,
) -> Result<oneshot::Sender<()>, EbpfError> {
let (shutdown_tx, shutdown_rx) = oneshot::channel();
@ -363,26 +363,26 @@ impl XskPair {
Ok(nb_completed)
}
fn process_rx_queue(&mut self, forward_tx: &Sender<Vec<u8>>) -> Result<usize, EbpfError> {
fn process_rx_queue(&mut self, forward_tx: &Sender<Arc<Vec<u8>>>) -> Result<usize, EbpfError> {
let mut rx_descs = vec![FrameDesc::default(); 256];
let rx_count = unsafe { self.rx.consume(&mut rx_descs) };
if rx_count > 0 {
for rx_desc in rx_descs.iter().take(rx_count) {
let lengths = rx_desc.lengths();
let packet_len = lengths.data() as usize;
let packet_len = rx_desc.lengths().data() as usize;
let data = unsafe { self.umem.data(rx_desc) };
let packet_data = data.contents()[..packet_len].to_vec();
if let Some(ref se) = self.suricata_engine {
se.inject(packet_data.clone());
}
let packet_slice = &data.contents()[..packet_len];
// ML engine reads directly from UMEM — no copy for this consumer
if let Some(ref engine) = self.engine {
engine.process_packet(&packet_data, self.direction == Direction::Ingress);
engine.process_packet(packet_slice, self.direction == Direction::Ingress);
}
// Single copy shared between Suricata and forward_tx via Arc
let packet_data = Arc::new(packet_slice.to_vec());
if let Some(ref se) = self.suricata_engine {
se.inject(Arc::clone(&packet_data));
}
if let Err(e) = forward_tx.try_send(packet_data) {
match e {
crossbeam::channel::TrySendError::Full(_) => {
@ -455,7 +455,7 @@ impl XskPair {
}
}
fn process_tx_queue(&mut self, forward_rx: &Receiver<Vec<u8>>) -> Result<usize, EbpfError> {
fn process_tx_queue(&mut self, forward_rx: &Receiver<Arc<Vec<u8>>>) -> Result<usize, EbpfError> {
// Drain completed TX frames first to maximise pool availability.
let _ = self.process_comp_queue();
@ -474,7 +474,7 @@ impl XskPair {
// Consume at most min(pool_size, 64) packets so we never over-commit.
let max_to_send = pool_size.min(64);
let mut packets_to_send = Vec::with_capacity(max_to_send);
let mut packets_to_send: Vec<Arc<Vec<u8>>> = Vec::with_capacity(max_to_send);
while let Ok(packet) = forward_rx.try_recv() {
packets_to_send.push(packet);
if packets_to_send.len() >= max_to_send {

View File

@ -22,7 +22,7 @@ const SURICATA_LOG: &str = "/tmp/suricata.log";
const CHANNEL_CAP: usize = 4096;
pub struct SuricataEngine {
tx: Sender<Vec<u8>>,
tx: Sender<Arc<Vec<u8>>>,
child: std::sync::Mutex<Child>,
}
@ -133,7 +133,7 @@ impl SuricataEngine {
})
.map_err(|e| SuricataError::ProcessSpawnFailed { reason: e.to_string() })?;
let (tx, rx) = bounded::<Vec<u8>>(CHANNEL_CAP);
let (tx, rx) = bounded::<Arc<Vec<u8>>>(CHANNEL_CAP);
thread::Builder::new()
.name("suricata-mirror".into())
@ -182,7 +182,7 @@ impl SuricataEngine {
}
/* Non-blocking: drops silently when the channel is full under load. */
pub fn inject(&self, data: Vec<u8>) {
pub fn inject(&self, data: Arc<Vec<u8>>) {
match self.tx.try_send(data) {
Ok(()) => {}
Err(crossbeam::channel::TrySendError::Full(_)) => {

View File

@ -1,166 +1,190 @@
use std::mem;
use std::time;
use common::model::event::{Event, IPv4Event, IPv6Event, TcpFlags};
use network_types::ip::IpProto;
use zerocopy::byteorder::{BigEndian, U16, U32};
use zerocopy::{FromBytes, Immutable, IntoBytes, KnownLayout, Ref};
#[repr(C, packed)]
#[derive(FromBytes, IntoBytes, KnownLayout, Immutable, Clone, Copy)]
struct EthHdr {
_dst_mac: [u8; 6],
_src_mac: [u8; 6],
ether_type: U16<BigEndian>,
}
#[repr(C, packed)]
#[derive(FromBytes, IntoBytes, KnownLayout, Immutable, Clone, Copy)]
struct Ipv4Hdr {
version_ihl: u8,
_dscp_ecn: u8,
total_len: U16<BigEndian>,
_ident: U16<BigEndian>,
_flags_frag: U16<BigEndian>,
_ttl: u8,
protocol: u8,
_checksum: U16<BigEndian>,
src_addr: U32<BigEndian>,
dst_addr: U32<BigEndian>,
}
#[repr(C, packed)]
#[derive(FromBytes, IntoBytes, KnownLayout, Immutable, Clone, Copy)]
struct Ipv6Hdr {
_version_tc_fl: U32<BigEndian>,
payload_len: U16<BigEndian>,
next_hdr: u8,
_hop_limit: u8,
src_addr: [u8; 16],
dst_addr: [u8; 16],
}
#[repr(C, packed)]
#[derive(FromBytes, IntoBytes, KnownLayout, Immutable, Clone, Copy)]
struct TcpHdr {
src_port: U16<BigEndian>,
dst_port: U16<BigEndian>,
_seq_num: U32<BigEndian>,
_ack_num: U32<BigEndian>,
data_off_flags: U16<BigEndian>,
window: U16<BigEndian>,
_checksum: U16<BigEndian>,
_urgent_ptr: U16<BigEndian>,
}
#[repr(C, packed)]
#[derive(FromBytes, IntoBytes, KnownLayout, Immutable, Clone, Copy)]
struct UdpHdr {
src_port: U16<BigEndian>,
dst_port: U16<BigEndian>,
_length: U16<BigEndian>,
_checksum: U16<BigEndian>,
}
const ETH_HDR_LEN: usize = 14;
const IPV6_HDR_LEN: usize = 40;
pub fn parse_packet(packet_data: &[u8]) -> Option<(Event, usize)> {
if packet_data.len() < 14 {
return None;
}
let eth_type = u16::from_be_bytes([packet_data[12], packet_data[13]]);
let (eth_hdr, ip_rest) = Ref::<&[u8], EthHdr>::from_prefix(packet_data).ok()?;
let timestamp_us = time::SystemTime::now()
.duration_since(time::UNIX_EPOCH)
.ok()?
.as_micros() as u64;
match eth_type {
0x0800 => parse_ipv4(packet_data, timestamp_us),
0x86DD => parse_ipv6(packet_data, timestamp_us),
match eth_hdr.ether_type.get() {
0x0800 => parse_ipv4(ip_rest, timestamp_us),
0x86DD => parse_ipv6(ip_rest, timestamp_us),
_ => None,
}
}
fn parse_ipv4(packet_data: &[u8], timestamp_us: u64) -> Option<(Event, usize)> {
if packet_data.len() < 34 {
return None;
}
let ip_header = &packet_data[14..];
let protocol_byte = ip_header[9];
let protocol = unsafe { mem::transmute::<u8, IpProto>(protocol_byte) };
let src_ip = u32::from_be_bytes([ip_header[12], ip_header[13], ip_header[14], ip_header[15]]);
let dst_ip = u32::from_be_bytes([ip_header[16], ip_header[17], ip_header[18], ip_header[19]]);
fn parse_ipv4(ip_bytes: &[u8], timestamp_us: u64) -> Option<(Event, usize)> {
let (ipv4_hdr, _) = Ref::<&[u8], Ipv4Hdr>::from_prefix(ip_bytes).ok()?;
let protocol_byte = ipv4_hdr.protocol;
if protocol_byte != 6 && protocol_byte != 17 {
return None;
}
let protocol = unsafe { std::mem::transmute::<u8, IpProto>(protocol_byte) };
let ihl = (ip_header[0] & 0x0F) as usize * 4;
let total_len = u16::from_be_bytes([ip_header[2], ip_header[3]]) as u32;
let src_ip = ipv4_hdr.src_addr.get();
let dst_ip = ipv4_hdr.dst_addr.get();
let total_len = ipv4_hdr.total_len.get() as u32;
let ihl = (ipv4_hdr.version_ihl & 0x0F) as usize * 4;
if packet_data.len() < 14 + ihl + 4 {
if ip_bytes.len() < ihl + 4 {
return None;
}
let transport_bytes = &ip_bytes[ihl..];
let transport_header = &ip_header[ihl..];
let src_port = u16::from_be_bytes([transport_header[0], transport_header[1]]);
let dst_port = u16::from_be_bytes([transport_header[2], transport_header[3]]);
let (tcp_flags, tcp_window_size, header_length) = if protocol_byte == 6 {
if packet_data.len() < 14 + ihl + 20 {
return None;
}
let data_offset = (transport_header[12] >> 4) as u16 * 4;
let flags = TcpFlags::from_byte(transport_header[13]);
let window = u16::from_be_bytes([transport_header[14], transport_header[15]]);
(flags, window, data_offset)
} else if protocol_byte == 17 {
(TcpFlags::default(), 0, 8)
let (src_port, dst_port, tcp_flags, tcp_window_size, header_length) = if protocol_byte == 6 {
let (tcp_hdr, _) = Ref::<&[u8], TcpHdr>::from_prefix(transport_bytes).ok()?;
let dof = tcp_hdr.data_off_flags.get();
let data_offset = (dof >> 12) as u16 * 4;
let flags = TcpFlags::from_byte((dof & 0xFF) as u8);
(tcp_hdr.src_port.get(), tcp_hdr.dst_port.get(), flags, tcp_hdr.window.get(), data_offset)
} else {
(TcpFlags::default(), 0, 0)
let (udp_hdr, _) = Ref::<&[u8], UdpHdr>::from_prefix(transport_bytes).ok()?;
(udp_hdr.src_port.get(), udp_hdr.dst_port.get(), TcpFlags::default(), 0u16, 8u16)
};
let payload_length = total_len.saturating_sub(ihl as u32 + header_length as u32);
let payload_start = 14 + ihl + header_length as usize;
let payload_start = ETH_HDR_LEN + ihl + header_length as usize;
let event = IPv4Event {
protocol,
src_ip,
dst_ip,
src_port,
dst_port,
packet_length: total_len,
payload_length,
header_length,
timestamp_us,
tcp_flags,
tcp_window_size,
is_forward: false,
};
Some((Event::IPv4(event), payload_start))
Some((
Event::IPv4(IPv4Event {
protocol,
src_ip,
dst_ip,
src_port,
dst_port,
packet_length: total_len,
payload_length,
header_length,
timestamp_us,
tcp_flags,
tcp_window_size,
is_forward: false,
}),
payload_start,
))
}
fn parse_ipv6(packet_data: &[u8], timestamp_us: u64) -> Option<(Event, usize)> {
if packet_data.len() < 54 {
return None;
}
let ip_header = &packet_data[14..];
let protocol_byte = ip_header[6];
let protocol = unsafe { mem::transmute::<u8, IpProto>(protocol_byte) };
let mut source_ip_bytes = [0u8; 16];
source_ip_bytes.copy_from_slice(&ip_header[8..24]);
let src_ip = u128::from_be_bytes(source_ip_bytes);
let mut dest_ip_bytes = [0u8; 16];
dest_ip_bytes.copy_from_slice(&ip_header[24..40]);
let dst_ip = u128::from_be_bytes(dest_ip_bytes);
fn parse_ipv6(ip_bytes: &[u8], timestamp_us: u64) -> Option<(Event, usize)> {
let (ipv6_hdr, transport_bytes) = Ref::<&[u8], Ipv6Hdr>::from_prefix(ip_bytes).ok()?;
let protocol_byte = ipv6_hdr.next_hdr;
if protocol_byte != 6 && protocol_byte != 17 {
return None;
}
let protocol = unsafe { std::mem::transmute::<u8, IpProto>(protocol_byte) };
let payload_len = u16::from_be_bytes([ip_header[4], ip_header[5]]) as u32;
let src_ip = u128::from_be_bytes(ipv6_hdr.src_addr);
let dst_ip = u128::from_be_bytes(ipv6_hdr.dst_addr);
let payload_len = ipv6_hdr.payload_len.get() as u32;
let total_len = payload_len + 40;
if packet_data.len() < 54 + 4 {
if transport_bytes.len() < 4 {
return None;
}
let transport_header = &ip_header[40..];
let src_port = u16::from_be_bytes([transport_header[0], transport_header[1]]);
let dst_port = u16::from_be_bytes([transport_header[2], transport_header[3]]);
let (tcp_flags, tcp_window_size, header_length) = if protocol_byte == 6 {
if packet_data.len() < 54 + 20 {
return None;
}
let data_offset = (transport_header[12] >> 4) as u16 * 4;
let flags = TcpFlags::from_byte(transport_header[13]);
let window = u16::from_be_bytes([transport_header[14], transport_header[15]]);
(flags, window, data_offset)
} else if protocol_byte == 17 {
(TcpFlags::default(), 0, 8)
let (src_port, dst_port, tcp_flags, tcp_window_size, header_length) = if protocol_byte == 6 {
let (tcp_hdr, _) = Ref::<&[u8], TcpHdr>::from_prefix(transport_bytes).ok()?;
let dof = tcp_hdr.data_off_flags.get();
let data_offset = (dof >> 12) as u16 * 4;
let flags = TcpFlags::from_byte((dof & 0xFF) as u8);
(tcp_hdr.src_port.get(), tcp_hdr.dst_port.get(), flags, tcp_hdr.window.get(), data_offset)
} else {
(TcpFlags::default(), 0, 0)
let (udp_hdr, _) = Ref::<&[u8], UdpHdr>::from_prefix(transport_bytes).ok()?;
(udp_hdr.src_port.get(), udp_hdr.dst_port.get(), TcpFlags::default(), 0u16, 8u16)
};
let payload_length = total_len.saturating_sub(40 + header_length as u32);
let payload_start = 14 + 40 + header_length as usize;
let payload_start = ETH_HDR_LEN + IPV6_HDR_LEN + header_length as usize;
let event = IPv6Event {
protocol,
src_ip,
dst_ip,
src_port,
dst_port,
packet_length: total_len,
payload_length,
header_length,
timestamp_us,
tcp_flags,
tcp_window_size,
is_forward: false,
};
Some((Event::IPv6(event), payload_start))
Some((
Event::IPv6(IPv6Event {
protocol,
src_ip,
dst_ip,
src_port,
dst_port,
packet_length: total_len,
payload_length,
header_length,
timestamp_us,
tcp_flags,
tcp_window_size,
is_forward: false,
}),
payload_start,
))
}
pub fn format_ipv4(addr: u32) -> String {
let bytes = addr.to_be_bytes();
format!("{}.{}.{}.{}", bytes[0], bytes[1], bytes[2], bytes[3],)
format!("{}.{}.{}.{}", bytes[0], bytes[1], bytes[2], bytes[3])
}
pub fn format_ipv6(addr: u128) -> String {
@ -182,6 +206,6 @@ pub fn format_ipv6(addr: u128) -> String {
bytes[12],
bytes[13],
bytes[14],
bytes[15]
bytes[15],
)
}