use crate::error::{NetError, Result};
use crate::packet::{parse_packet, L4Protocol};
use crate::protocol_expectation::ProtocolExpectation;
use crate::source_admission::{AdmissionAction, IpAddr, SourceAdmissionEngine};
use crate::transport::{
parse_quic_header_with_dcid_len, QuicAction, QuicConnParams, QuicConnectionTable, QuicHeaderType,
};
use std::sync::Arc;
use zenith_foundation::{FrameId, FramePool};
use zenith_linux::descriptor::XdpDesc;
use zenith_linux::umem::{UmemConfig, UmemManager};
use zenith_linux::xsk::XskSocket;
pub use zenith_linux::xsk::XskConfig;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum WorkerState {
Created,
Running,
Stopped,
Error,
}
#[derive(Debug, Clone, Copy, Default)]
pub struct WorkerStats {
pub rx_packets: u64,
pub tx_packets: u64,
pub rejected_packets: u64,
pub parse_errors: u64,
pub frame_allocs: u64,
pub frame_frees: u64,
pub quarantined_frames: u64,
pub cycles_completed: u64,
pub cpu_usage: f64,
pub queue_depth: u64,
pub latency_p99: u64,
pub throughput: u64,
pub unexpected_dropped: u64,
}
const MAX_BATCH_SIZE: usize = 64;
const FRAME_SIZE: usize = 4096;
pub struct Worker {
id: u32,
xsk: XskSocket,
frame_pool: FramePool,
admission: SourceAdmissionEngine,
state: WorkerState,
stats: WorkerStats,
domain_id: u32,
umem_buf: Vec<u8>,
umem_manager: Option<Arc<UmemManager>>,
rx_desc_buf: [XdpDesc; MAX_BATCH_SIZE],
rx_desc_count: usize,
quarantine_on_error: bool,
quarantine_on_deny: bool,
quic_table: QuicConnectionTable,
latency_samples: [u64; 64],
latency_sample_idx: usize,
latency_sample_count: usize,
last_cycle_end: Option<std::time::Instant>,
quic_stats: WorkerQuicStats,
alloc_generations: Vec<u64>,
protocol_expectation: Option<ProtocolExpectation>,
}
fn compute_p99(samples: &[u64; 64], valid: usize) -> u64 {
let valid = valid.min(64);
if valid == 0 {
return 0;
}
let mut sorted = *samples;
sorted.sort_unstable();
let rank = (valid * 99).div_ceil(100).clamp(1, valid);
sorted[64 - valid + (rank - 1)]
}
#[derive(Debug, Clone)]
pub struct UdpDatagram {
pub src_ip: IpAddr,
pub dst_ip: IpAddr,
pub src_port: u16,
pub dst_port: u16,
pub src_mac: [u8; 6],
pub dst_mac: [u8; 6],
pub payload: smallvec::SmallVec<[u8; 2048]>,
}
impl UdpDatagram {
pub fn src_socket_addr(&self) -> std::net::SocketAddr {
let ip = match self.src_ip {
IpAddr::V4(b) => std::net::IpAddr::V4(std::net::Ipv4Addr::from(b)),
IpAddr::V6(b) => std::net::IpAddr::V6(std::net::Ipv6Addr::from(b)),
IpAddr::Any => std::net::IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED),
};
std::net::SocketAddr::new(ip, self.src_port)
}
pub fn build_reply_frame(&self, payload: &[u8], buf: &mut [u8]) -> Option<usize> {
use crate::packet::{ETH_HEADER_LEN, UDP_HEADER_LEN};
const IPV4_HEADER_LEN: usize = 20;
const MAX_IP_TOTAL: usize = u16::MAX as usize;
let (local_ip, peer_ip) = match (self.dst_ip, self.src_ip) {
(IpAddr::V4(local), IpAddr::V4(peer)) => (local, peer),
_ => return None,
};
let total_len = ETH_HEADER_LEN
.checked_add(IPV4_HEADER_LEN)?
.checked_add(UDP_HEADER_LEN)?
.checked_add(payload.len())?;
let ip_total = IPV4_HEADER_LEN.checked_add(UDP_HEADER_LEN)?.checked_add(payload.len())?;
if total_len > buf.len() || ip_total > MAX_IP_TOTAL {
return None;
}
buf[0..6].copy_from_slice(&self.src_mac); buf[6..12].copy_from_slice(&self.dst_mac); buf[12] = 0x08;
buf[13] = 0x00;
let ip = ETH_HEADER_LEN;
buf[ip] = 0x45; buf[ip + 1] = 0; buf[ip + 2..ip + 4].copy_from_slice(&(ip_total as u16).to_be_bytes());
buf[ip + 4] = 0; buf[ip + 5] = 0;
buf[ip + 6] = 0x40; buf[ip + 7] = 0; buf[ip + 8] = 64; buf[ip + 9] = crate::packet::ip_proto::UDP;
buf[ip + 10] = 0; buf[ip + 11] = 0;
buf[ip + 12..ip + 16].copy_from_slice(&local_ip);
buf[ip + 16..ip + 20].copy_from_slice(&peer_ip);
let cksum = crate::packet::compute_ipv4_checksum(&buf[ip..ip + IPV4_HEADER_LEN]);
buf[ip + 10..ip + 12].copy_from_slice(&cksum.to_be_bytes());
let l4 = ip + IPV4_HEADER_LEN;
let udp_len = u16::try_from(UDP_HEADER_LEN.checked_add(payload.len())?).ok()?;
buf[l4..l4 + 2].copy_from_slice(&self.dst_port.to_be_bytes());
buf[l4 + 2..l4 + 4].copy_from_slice(&self.src_port.to_be_bytes());
buf[l4 + 4..l4 + 6].copy_from_slice(&udp_len.to_be_bytes());
buf[l4 + 6] = 0;
buf[l4 + 7] = 0;
buf[l4 + UDP_HEADER_LEN..total_len].copy_from_slice(payload);
Some(total_len)
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct WorkerQuicStats {
pub quic_packets: u64,
pub quic_new_connections: u64,
pub quic_amplification_blocked: u64,
pub quic_invalid_packets: u64,
}
impl Worker {
pub fn new(
id: u32,
xsk_config: XskConfig,
frame_pool_capacity: u32,
admission: SourceAdmissionEngine,
) -> Result<Self> {
let xsk = XskSocket::new(xsk_config)?;
let frame_pool = FramePool::new(
format!("worker-{}", id),
frame_pool_capacity,
FRAME_SIZE as u32,
);
let umem_buf = vec![0u8; (frame_pool_capacity as usize) * FRAME_SIZE];
Ok(Self {
id,
xsk,
frame_pool,
admission,
state: WorkerState::Created,
stats: WorkerStats::default(),
domain_id: id,
umem_buf,
umem_manager: None,
rx_desc_buf: [XdpDesc::zero(); MAX_BATCH_SIZE],
rx_desc_count: 0,
quarantine_on_error: false,
quarantine_on_deny: false,
quic_table: QuicConnectionTable::new(64),
latency_samples: [0u64; 64],
latency_sample_idx: 0,
latency_sample_count: 0,
last_cycle_end: None,
quic_stats: WorkerQuicStats::default(),
alloc_generations: vec![0u64; frame_pool_capacity as usize],
protocol_expectation: None,
})
}
pub fn with_frame_pool(
id: u32,
xsk_config: XskConfig,
frame_pool: FramePool,
admission: SourceAdmissionEngine,
) -> Result<Self> {
let capacity = frame_pool.capacity();
let xsk = XskSocket::new(xsk_config)?;
let umem_buf = vec![0u8; (capacity as usize) * FRAME_SIZE];
Ok(Self {
id,
xsk,
frame_pool,
admission,
state: WorkerState::Created,
stats: WorkerStats::default(),
domain_id: id,
umem_buf,
umem_manager: None,
rx_desc_buf: [XdpDesc::zero(); MAX_BATCH_SIZE],
rx_desc_count: 0,
quarantine_on_error: false,
quarantine_on_deny: false,
quic_table: QuicConnectionTable::new(64),
latency_samples: [0u64; 64],
latency_sample_idx: 0,
latency_sample_count: 0,
last_cycle_end: None,
quic_stats: WorkerQuicStats::default(),
alloc_generations: vec![0u64; capacity as usize],
protocol_expectation: None,
})
}
pub fn with_real_af_xdp(
id: u32,
xsk_config: XskConfig,
frame_pool_capacity: u32,
admission: SourceAdmissionEngine,
use_hugepage: bool,
lock_memory: bool,
) -> Result<Self> {
let umem_size_raw = (frame_pool_capacity as usize)
.checked_mul(FRAME_SIZE)
.ok_or(NetError::Internal("frame_pool_capacity * FRAME_SIZE overflow"))?;
let page_size = {
zenith_linux::page_size()
};
let umem_size = (umem_size_raw + page_size - 1) & !(page_size - 1);
let mut umem_mgr = UmemManager::new(UmemConfig {
size: umem_size,
hugepage: use_hugepage,
locked: lock_memory,
shared: false,
})?;
umem_mgr.create()?; let umem_arc: Arc<UmemManager> = Arc::new(umem_mgr);
let mut xsk = XskSocket::new(xsk_config)?;
xsk.create_socket()?; xsk.configure()?; xsk.bind(Arc::clone(&umem_arc))?;
let frame_pool = FramePool::new(
format!("worker-{id}"),
frame_pool_capacity,
FRAME_SIZE as u32,
);
let umem_buf = vec![0u8; (frame_pool_capacity as usize) * FRAME_SIZE];
Ok(Self {
id,
xsk,
frame_pool,
admission,
state: WorkerState::Created,
stats: WorkerStats::default(),
domain_id: id,
umem_buf,
umem_manager: Some(umem_arc),
rx_desc_buf: [XdpDesc::zero(); MAX_BATCH_SIZE],
rx_desc_count: 0,
quarantine_on_error: false,
quarantine_on_deny: false,
quic_table: QuicConnectionTable::new(64),
latency_samples: [0u64; 64],
latency_sample_idx: 0,
latency_sample_count: 0,
last_cycle_end: None,
quic_stats: WorkerQuicStats::default(),
alloc_generations: vec![0u64; frame_pool_capacity as usize],
protocol_expectation: None,
})
}
#[inline]
pub fn is_real_af_xdp(&self) -> bool {
self.xsk.get_fd().is_ok() && self.umem_manager.is_some()
}
#[cfg(all(feature = "linux", target_os = "linux"))]
pub fn register_xsk(&self, maps: &zenith_ebpf::BpfMaps<'_>) -> Result<()> {
let fd = self.xsk.get_fd()?;
let fd_u32 = u32::try_from(fd).map_err(|_| {
NetError::Internal("xsk fd negative (kernel ABI 异常)")
})?;
let queue_id = self.xsk.queue_id();
maps.update_xsk(queue_id, fd_u32)?;
tracing::debug!(
worker_id = self.id,
queue_id,
xsk_fd = fd,
"XSKMAP registered: queue -> xsk_fd"
);
Ok(())
}
#[inline]
pub fn umem_manager(&self) -> Option<&Arc<UmemManager>> {
self.umem_manager.as_ref()
}
#[inline]
pub fn quic_table(&self) -> &QuicConnectionTable {
&self.quic_table
}
#[inline]
pub fn quic_stats(&self) -> WorkerQuicStats {
self.quic_stats
}
#[inline]
pub fn quic_table_mut(&mut self) -> &mut QuicConnectionTable {
&mut self.quic_table
}
pub fn handle_quic_packet(
&mut self,
data: &[u8],
remote_addr: IpAddr,
remote_port: u16,
local_addr: IpAddr,
local_port: u16,
ip_version: crate::packet::IpVersion,
) -> QuicAction {
let hdr = match parse_quic_header_with_dcid_len(data, 8) {
Some(h) => h,
None => {
self.quic_stats.quic_invalid_packets += 1;
return QuicAction::InvalidPacket;
}
};
self.quic_stats.quic_packets += 1;
match hdr.header_type {
QuicHeaderType::Short => {
let conn_idx = match self.quic_table.find_by_dcid(&hdr.dcid[..hdr.dcid_len as usize]) {
Some(idx) => idx,
None => {
self.quic_stats.quic_invalid_packets += 1;
return QuicAction::InvalidPacket;
}
};
if let Some(c) = self.quic_table.get(conn_idx) {
c.observe_packet_number(hdr.packet_number);
c.record_rx_bytes(data.len() as u64);
}
QuicAction::DataReceived {
conn_idx,
stream_id: hdr.packet_number,
}
}
QuicHeaderType::VersionNegotiation | QuicHeaderType::Retry => {
QuicAction::InvalidPacket
}
QuicHeaderType::Long => {
if let Some(idx) = self
.quic_table
.find_by_dcid(&hdr.dcid[..hdr.dcid_len as usize])
&& let Some(c) = self.quic_table.get(idx) {
c.observe_packet_number(hdr.packet_number);
c.record_rx_bytes(data.len() as u64);
if hdr.has_token()
&& let Ok(ok) = c.verify_token(
&hdr.token[..hdr.token_len as usize],
)
&& !ok {
self.quic_stats.quic_amplification_blocked += 1;
return QuicAction::AmplificationBlocked { conn_idx: idx };
}
return QuicAction::DataReceived {
conn_idx: idx,
stream_id: hdr.packet_number,
};
}
match self.quic_table.allocate(QuicConnParams {
scid: &hdr.dcid[..hdr.dcid_len as usize],
dcid: &hdr.scid[..hdr.scid_len as usize],
remote_addr,
remote_port,
local_addr,
local_port,
ip_version,
}) {
Ok(idx) => {
if let Some(c) = self.quic_table.get(idx) {
c.observe_packet_number(hdr.packet_number);
c.record_rx_bytes(data.len() as u64);
if hdr.has_token() {
c.retry_token[..hdr.token_len as usize]
.copy_from_slice(&hdr.token[..hdr.token_len as usize]);
c.retry_token_len = hdr.token_len;
}
}
self.quic_stats.quic_new_connections += 1;
QuicAction::NewConnection { conn_idx: idx }
}
Err(_) => QuicAction::InvalidPacket,
}
}
}
}
#[inline]
pub fn id(&self) -> u32 {
self.id
}
#[inline]
pub fn state(&self) -> WorkerState {
self.state
}
#[inline]
pub fn stats(&self) -> WorkerStats {
self.stats
}
#[inline]
pub fn admission_mut(&mut self) -> &mut SourceAdmissionEngine {
&mut self.admission
}
#[inline]
pub fn admission(&self) -> &SourceAdmissionEngine {
&self.admission
}
#[inline]
pub fn set_quarantine_on_error(&mut self, enabled: bool) {
self.quarantine_on_error = enabled;
}
#[inline]
pub fn set_quarantine_on_deny(&mut self, enabled: bool) {
self.quarantine_on_deny = enabled;
}
#[inline]
pub fn quarantine_on_error(&self) -> bool {
self.quarantine_on_error
}
#[inline]
pub fn quarantine_on_deny(&self) -> bool {
self.quarantine_on_deny
}
#[inline]
pub fn set_protocol_expectation(&mut self, expectation: Option<ProtocolExpectation>) {
self.protocol_expectation = expectation;
}
#[inline]
pub fn protocol_expectation(&self) -> Option<&ProtocolExpectation> {
self.protocol_expectation.as_ref()
}
#[inline]
pub fn frame_pool(&self) -> &FramePool {
&self.frame_pool
}
#[inline]
pub fn frame_pool_mut(&mut self) -> &mut FramePool {
&mut self.frame_pool
}
pub fn start(&mut self) {
self.state = WorkerState::Running;
}
pub fn stop(&mut self) {
self.state = WorkerState::Stopped;
}
#[inline]
pub fn get_frame_data(&self, frame_idx: u32) -> Option<&[u8]> {
let idx = frame_idx as usize;
let offset = idx.checked_mul(FRAME_SIZE)?;
let end = offset.checked_add(FRAME_SIZE)?;
match &self.umem_manager {
Some(umem) => umem.slice(offset, FRAME_SIZE),
None => self.umem_buf.get(offset..end),
}
}
#[inline]
pub fn get_frame_data_mut(&mut self, frame_idx: u32) -> Option<&mut [u8]> {
let idx = frame_idx as usize;
let offset = idx.checked_mul(FRAME_SIZE)?;
let end = offset.checked_add(FRAME_SIZE)?;
match &self.umem_manager {
Some(umem) => umem.slice_mut(offset, FRAME_SIZE),
None => self.umem_buf.get_mut(offset..end),
}
}
pub fn simulate_rx_transfer<F>(&mut self, count: u32, data_len: usize, data_generator: F)
where
F: Fn(u32, &mut [u8]),
{
let mut descs: [XdpDesc; MAX_BATCH_SIZE] = [XdpDesc::zero(); MAX_BATCH_SIZE];
let mut desc_count = 0usize;
for _ in 0..(count as usize).min(MAX_BATCH_SIZE) {
if let Ok(desc) = self.xsk.fill_ring_mut().dequeue_batch(1)
&& let Some(d) = desc.first() {
let frame_idx = (d.addr / FRAME_SIZE as u64) as u32;
let Some(data) = self.get_frame_data_mut(frame_idx) else {
let _ = self.frame_pool.quarantine_by_id(
FrameId::new(frame_idx),
self.frame_generation(frame_idx),
"frame_idx_out_of_range",
);
self.stats.quarantined_frames += 1;
continue;
};
let len = data_len.min(FRAME_SIZE);
data_generator(frame_idx, &mut data[..len]);
descs[desc_count] = *d;
desc_count += 1;
}
}
if desc_count > 0 {
let _ = self.xsk.rx_ring_mut().enqueue_batch(&descs[..desc_count]);
}
}
pub fn simulate_tx_complete(&mut self, count: u32) {
if let Ok(descs) = self.xsk.tx_ring_mut().dequeue_batch(count)
&& !descs.is_empty() {
let _ = self.xsk.completion_ring_mut().enqueue_batch(&descs);
}
}
pub fn process_cycle(&mut self) -> Result<u32> {
let cycle_start = std::time::Instant::now();
let result = self.process_cycle_inner();
let processed = result.as_ref().ok().copied().unwrap_or(0);
self.finish_cycle_metrics(cycle_start, processed);
result
}
fn finish_cycle_metrics(&mut self, cycle_start: std::time::Instant, processed: u32) {
let elapsed_us = cycle_start.elapsed().as_micros() as u64;
self.latency_samples[self.latency_sample_idx] = elapsed_us;
self.latency_sample_idx = (self.latency_sample_idx + 1) % 64;
self.latency_sample_count = (self.latency_sample_count + 1).min(64);
self.stats.latency_p99 =
compute_p99(&self.latency_samples, self.latency_sample_count);
self.stats.queue_depth = self.rx_desc_count as u64;
let processed_pps = u64::from(processed).saturating_mul(1_000_000);
self.stats.throughput = processed_pps.checked_div(elapsed_us).unwrap_or(processed_pps);
let now = std::time::Instant::now();
if let Some(last) = self.last_cycle_end {
let wall_us = now.duration_since(last).as_micros() as u64;
self.stats.cpu_usage = if wall_us > 0 {
(elapsed_us as f64 / wall_us as f64).clamp(0.0, 1.0)
} else {
1.0
};
}
self.last_cycle_end = Some(now);
}
pub fn process_cycle_with_udp_sink(
&mut self,
sink: &mut dyn FnMut(UdpDatagram),
) -> Result<u32> {
let cycle_start = std::time::Instant::now();
let result = self.process_cycle_inner_with_sink(Some(sink));
let processed = result.as_ref().ok().copied().unwrap_or(0);
self.finish_cycle_metrics(cycle_start, processed);
result
}
fn process_cycle_inner(&mut self) -> Result<u32> {
self.process_cycle_inner_with_sink(None)
}
fn process_cycle_inner_with_sink(
&mut self,
sink: Option<&mut dyn FnMut(UdpDatagram)>,
) -> Result<u32> {
if self.state != WorkerState::Running {
return Err(NetError::WorkerState {
reason: "worker not running",
});
}
self.recycle_completed_frames()?;
let filled = self.fill_fill_ring()?;
let received = self.rx_batch_preallocated()?;
if received == 0 && filled == 0 {
self.stats.cycles_completed += 1;
return Ok(0);
}
let mut descs_copy: [XdpDesc; MAX_BATCH_SIZE] = [XdpDesc::zero(); MAX_BATCH_SIZE];
descs_copy[..self.rx_desc_count].copy_from_slice(&self.rx_desc_buf[..self.rx_desc_count]);
let (processed, allow_count, deny_count) = self
.process_packets_batch_with_sink(&descs_copy[..self.rx_desc_count], sink)?;
self.recycle_rejected_frames(deny_count)?;
self.submit_to_tx_ring(allow_count)?;
self.stats.cycles_completed += 1;
Ok(processed)
}
#[inline]
fn fill_fill_ring(&mut self) -> Result<u32> {
let available = self.xsk.fill_ring().available_space();
if available == 0 {
return Ok(0);
}
let to_fill = (available.min(MAX_BATCH_SIZE as u32)) as usize;
let mut descs: [XdpDesc; MAX_BATCH_SIZE] = [XdpDesc::zero(); MAX_BATCH_SIZE];
let mut tokens: [Option<zenith_foundation::FrameToken>; MAX_BATCH_SIZE] =
[const { None }; MAX_BATCH_SIZE];
let mut count = 0usize;
for (desc_slot, token_slot) in descs.iter_mut().zip(tokens.iter_mut()).take(to_fill) {
match self.frame_pool.allocate(self.domain_id) {
Ok(token) => {
let frame_idx = token.frame_id().value();
if let Some(g) = self.alloc_generations.get_mut(frame_idx as usize) {
*g = token.generation();
}
if let Some(data) = self.get_frame_data_mut(frame_idx) {
data.fill(0);
} else {
let _ = self.frame_pool.quarantine_by_id(
FrameId::new(frame_idx),
token.generation(),
"frame_idx_out_of_range",
);
self.stats.quarantined_frames += 1;
continue;
}
*desc_slot = XdpDesc {
addr: (frame_idx as u64) << 12,
len: 0,
options: 0,
};
*token_slot = Some(token);
count += 1;
}
Err(_) => break,
}
}
if count == 0 {
return Ok(0);
}
let enqueued = self
.xsk
.fill_ring_mut()
.enqueue_batch(&descs[..count])? as usize;
for (i, token_slot) in tokens.iter_mut().enumerate().take(count) {
if i < enqueued {
self.stats.frame_allocs += 1;
let _ = token_slot.take();
} else if let Some(token) = token_slot.take() {
let _ = self.frame_pool.release(token);
}
}
Ok(enqueued as u32)
}
#[inline]
fn rx_batch_preallocated(&mut self) -> Result<usize> {
let available = self.xsk.rx_ring().available_data();
if available == 0 {
self.rx_desc_count = 0;
return Ok(0);
}
let to_receive = available.min(MAX_BATCH_SIZE as u32);
let received = self
.xsk
.rx_ring_mut()
.dequeue_batch_to(&mut self.rx_desc_buf[..to_receive as usize])?;
self.rx_desc_count = received as usize;
self.stats.rx_packets += received as u64;
Ok(self.rx_desc_count)
}
fn process_packets_batch_with_sink(
&mut self,
descs: &[XdpDesc],
mut sink: Option<&mut dyn FnMut(UdpDatagram)>,
) -> Result<(u32, u32, u32)> {
if descs.is_empty() {
return Ok((0, 0, 0));
}
let mut processed = 0u32;
let mut allow_count = 0u32;
let mut deny_count = 0u32;
for desc in descs {
if desc.is_zero() {
continue;
}
let frame_idx = (desc.addr / FRAME_SIZE as u64) as u32;
let Some(data) = self.get_frame_data(frame_idx) else {
self.stats.parse_errors += 1;
deny_count += 1;
processed += 1;
let _ = self.frame_pool.quarantine_by_id(
FrameId::new(frame_idx),
self.frame_generation(frame_idx),
"frame_idx_out_of_range",
);
self.stats.quarantined_frames += 1;
let _ = self.xsk.completion_ring_mut().enqueue_batch(&[*desc]);
continue;
};
let parsed = match parse_packet(data) {
Ok(p) => p,
Err(_e) => {
self.stats.parse_errors += 1;
deny_count += 1;
processed += 1;
if self.quarantine_on_error {
let _ = self.frame_pool.quarantine_by_id(
FrameId::new(frame_idx),
self.frame_generation(frame_idx),
"parse_error",
);
self.stats.quarantined_frames += 1;
}
let _ = self.xsk.completion_ring_mut().enqueue_batch(&[*desc]);
continue;
}
};
let (src_ip, src_port, dst_port, protocol) = self.extract_flow_info(&parsed);
let action = self
.admission
.evaluate(src_ip, src_port, dst_port, protocol);
match action {
AdmissionAction::Allow => {
if let Some(ref expectation) = self.protocol_expectation {
let ttl = parsed.ipv4.map(|ip| ip.ttl()).unwrap_or(64);
let is_fragment = parsed.ipv4.map(|ip| {
ip.fragment_offset() > 0 || ip.more_fragments()
}).unwrap_or(false);
if !expectation.is_packet_expected(protocol, dst_port, ttl, is_fragment) {
self.stats.unexpected_dropped += 1;
let _ = self.xsk.completion_ring_mut().enqueue_batch(&[*desc]);
processed += 1;
continue;
}
}
if let Some(sink_fn) = sink.as_deref_mut()
&& protocol == crate::packet::ip_proto::UDP
&& let Some(payload) = Self::extract_udp_payload(&parsed)
{
let dst_ip = match parsed.ip_version {
crate::packet::IpVersion::V4 => parsed
.ipv4
.map(|i| IpAddr::V4(i.dst_ip()))
.unwrap_or(IpAddr::V4_WILDCARD),
crate::packet::IpVersion::V6 => parsed
.ipv6
.map(|i| IpAddr::V6(i.dst_ip()))
.unwrap_or(IpAddr::V6_WILDCARD),
crate::packet::IpVersion::Unknown => IpAddr::V4_WILDCARD,
};
sink_fn(UdpDatagram {
src_ip,
dst_ip,
src_port,
dst_port,
src_mac: parsed.eth.src_mac(),
dst_mac: parsed.eth.dst_mac(),
payload: smallvec::SmallVec::from_slice(payload),
});
}
let _ = self.xsk.tx_ring_mut().enqueue_batch(&[*desc]);
allow_count += 1;
processed += 1;
}
AdmissionAction::Deny => {
if self.quarantine_on_deny {
let _ = self.frame_pool.quarantine_by_id(
FrameId::new(frame_idx),
self.frame_generation(frame_idx),
"admission_deny",
);
self.stats.quarantined_frames += 1;
}
let _ = self.xsk.completion_ring_mut().enqueue_batch(&[*desc]);
self.stats.rejected_packets += 1;
deny_count += 1;
processed += 1;
}
}
}
Ok((processed, allow_count, deny_count))
}
fn extract_udp_payload<'a>(parsed: &crate::packet::ParsedPacket<'a>) -> Option<&'a [u8]> {
use crate::packet::{ETH_HEADER_LEN, IPV6_HEADER_LEN, IpVersion, UDP_HEADER_LEN, VLAN_TAG_LEN};
let udp = parsed.udp?;
let eth_len = ETH_HEADER_LEN
.checked_add(if parsed.vlan.is_some() { VLAN_TAG_LEN } else { 0 })?;
let ip_len = match parsed.ip_version {
IpVersion::V4 => parsed.ipv4?.header_length(),
IpVersion::V6 => IPV6_HEADER_LEN,
IpVersion::Unknown => return None,
};
let l4_start = eth_len.checked_add(ip_len)?;
let udp_total = u16::from_be(udp.length) as usize;
if udp_total < UDP_HEADER_LEN {
return None;
}
let payload_start = l4_start.checked_add(UDP_HEADER_LEN)?;
let payload_end = l4_start.checked_add(udp_total)?;
if payload_end > parsed.raw.len() || payload_start > payload_end {
return None;
}
Some(&parsed.raw[payload_start..payload_end])
}
#[inline]
fn extract_flow_info(
&self,
parsed: &crate::packet::ParsedPacket<'_>,
) -> (IpAddr, u16, u16, u8) {
let src_ip = match parsed.ip_version {
crate::packet::IpVersion::V4 => {
if let Some(ipv4) = parsed.ipv4 {
IpAddr::V4(ipv4.src_ip())
} else {
IpAddr::V4_WILDCARD
}
}
crate::packet::IpVersion::V6 => {
if let Some(ipv6) = parsed.ipv6 {
IpAddr::V6(ipv6.src_ip())
} else {
IpAddr::V6_WILDCARD
}
}
crate::packet::IpVersion::Unknown => IpAddr::V4_WILDCARD,
};
let src_port = match parsed.l4_proto {
L4Protocol::Tcp => parsed.tcp.map(|t| t.src_port()).unwrap_or(0),
L4Protocol::Udp => parsed.udp.map(|u| u.src_port()).unwrap_or(0),
_ => 0,
};
let dst_port = match parsed.l4_proto {
L4Protocol::Tcp => parsed.tcp.map(|t| t.dst_port()).unwrap_or(0),
L4Protocol::Udp => parsed.udp.map(|u| u.dst_port()).unwrap_or(0),
_ => 0,
};
let protocol = match parsed.ip_version {
crate::packet::IpVersion::V4 => {
parsed.ipv4.map(|i| i.protocol()).unwrap_or(0)
}
crate::packet::IpVersion::V6 => {
parsed.ipv6.map(|i| i.next_header()).unwrap_or(0)
}
crate::packet::IpVersion::Unknown => 0,
};
(src_ip, src_port, dst_port, protocol)
}
#[inline]
fn recycle_rejected_frames(&mut self, count: u32) -> Result<()> {
if count == 0 {
return Ok(());
}
self.recycle_completed_frames_inner(count)?;
Ok(())
}
#[inline]
fn submit_to_tx_ring(&mut self, count: u32) -> Result<()> {
if count == 0 {
return Ok(());
}
if self.is_real_af_xdp() {
self.xsk.notify_tx()?;
self.stats.tx_packets += count as u64;
} else {
self.simulate_tx_complete(count);
self.stats.tx_packets += count as u64;
self.recycle_completed_frames_inner(count)?;
}
Ok(())
}
#[inline]
fn recycle_completed_frames(&mut self) -> Result<()> {
let available = self.xsk.completion_ring().available_data();
if available == 0 {
return Ok(());
}
let count = available.min(MAX_BATCH_SIZE as u32);
self.recycle_completed_frames_inner(count)?;
Ok(())
}
#[inline]
fn recycle_completed_frames_inner(&mut self, count: u32) -> Result<()> {
if count == 0 {
return Ok(());
}
let mut buffer: [XdpDesc; MAX_BATCH_SIZE] = [XdpDesc::zero(); MAX_BATCH_SIZE];
let dequeued = self
.xsk
.completion_ring_mut()
.dequeue_batch_to(&mut buffer[..count as usize])?;
for desc in buffer.iter().take(dequeued as usize) {
if desc.is_zero() {
continue;
}
let frame_idx = (desc.addr / FRAME_SIZE as u64) as u32;
if let Some(frame_info) = self.frame_pool.get_frame_info(FrameId::new(frame_idx))
&& frame_info.state() == zenith_foundation::FrameState::Allocated {
let _ = self.frame_pool.release_by_id(
FrameId::new(frame_idx),
self.frame_generation(frame_idx),
);
self.stats.frame_frees += 1;
}
}
Ok(())
}
pub fn verify_conservation(&self) -> bool {
self.frame_pool.verify_conservation().is_ok()
}
pub fn frame_pool_info(&self) -> (u32, u32, u32) {
(
self.frame_pool.capacity(),
self.frame_pool.allocated_count(),
self.frame_pool.quarantined_count(),
)
}
#[inline]
fn frame_generation(&self, frame_idx: u32) -> u64 {
self.alloc_generations
.get(frame_idx as usize)
.copied()
.unwrap_or(0)
}
pub fn queue_tx_data(&mut self, data: &[u8]) -> Result<u32> {
if data.len() > FRAME_SIZE {
return Err(NetError::InvalidOperation {
reason: format!(
"tx data {} bytes exceeds frame size {}",
data.len(),
FRAME_SIZE
),
});
}
let token = self.frame_pool.allocate(self.domain_id)?;
let frame_idx = token.frame_id().value();
let buf = self.get_frame_data_mut(frame_idx).ok_or(NetError::Internal(
"allocated frame index out of umem range",
))?;
buf[..data.len()].copy_from_slice(data);
if data.len() < FRAME_SIZE {
buf[data.len()..].fill(0);
}
if let Some(g) = self.alloc_generations.get_mut(frame_idx as usize) {
*g = token.generation();
}
let desc = XdpDesc {
addr: (frame_idx as u64) << 12,
len: data.len() as u32,
options: 0,
};
self.xsk.tx_ring_mut().enqueue_batch(&[desc])?;
let _ = token;
Ok(frame_idx)
}
}
impl Drop for Worker {
fn drop(&mut self) {
self.state = WorkerState::Stopped;
}
}
impl std::fmt::Debug for Worker {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Worker")
.field("id", &self.id)
.field("state", &self.state)
.field("stats", &self.stats)
.field("domain_id", &self.domain_id)
.field("rx_desc_count", &self.rx_desc_count)
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::packet::*;
use crate::source_admission::{AdmissionRule, ProtoMatch};
#[test]
fn test_compute_p99_empty() {
assert_eq!(compute_p99(&[0u64; 64], 0), 0);
}
#[test]
fn test_compute_p99_partial_valid() {
let mut samples = [0u64; 64];
samples[0] = 10;
samples[1] = 20;
samples[2] = 30;
samples[3] = 100;
assert_eq!(compute_p99(&samples, 4), 100);
}
#[test]
fn test_compute_p99_full_ring() {
let mut samples = [0u64; 64];
for (i, s) in samples.iter_mut().enumerate() {
*s = (i as u64) + 1; }
assert_eq!(compute_p99(&samples, 64), 64);
}
#[test]
fn test_worker_stats_real_writes() {
let mut worker = Worker::new(
0,
create_test_config(),
64,
SourceAdmissionEngine::allow_all(),
)
.expect("worker 构造应成功");
worker.start();
worker.process_cycle().expect("周期 1 应成功");
std::thread::sleep(std::time::Duration::from_millis(1));
worker.process_cycle().expect("周期 2 应成功");
let stats = worker.stats();
assert_eq!(stats.cycles_completed, 2);
assert!(stats.latency_p99 > 0, "latency_p99 应非零: {stats:?}");
assert!(stats.cpu_usage > 0.0 && stats.cpu_usage <= 1.0,
"cpu_usage 应在 (0,1] 区间: {}", stats.cpu_usage);
let _ = stats.queue_depth;
let _ = stats.throughput;
}
fn create_test_config() -> XskConfig {
XskConfig {
ifindex: 0,
queue_id: 0,
zero_copy: false,
fill_ring_size: 256,
rx_ring_size: 256,
tx_ring_size: 256,
completion_ring_size: 256,
shared_umem: false,
frame_size: FRAME_SIZE as u32,
headroom: 0,
..Default::default()
}
}
fn create_test_packet(frame: &mut [u8], dst_mac: [u8; 6], src_mac: [u8; 6], src_ip: [u8; 4], dst_ip: [u8; 4], src_port: u16, dst_port: u16) -> usize {
let offset = ETH_HEADER_LEN;
frame[0..6].copy_from_slice(&dst_mac);
frame[6..12].copy_from_slice(&src_mac);
frame[12] = 0x08; frame[13] = 0x00;
frame[offset] = 0x45; frame[offset + 1] = 0x00; frame[offset + 2] = 0x00;
frame[offset + 3] = 40;
frame[offset + 4] = 0x00;
frame[offset + 5] = 0x01;
frame[offset + 6] = 0x40;
frame[offset + 7] = 0x00;
frame[offset + 8] = 64;
frame[offset + 9] = ip_proto::TCP;
frame[offset + 10] = 0x00;
frame[offset + 11] = 0x00;
frame[offset + 12..offset + 16].copy_from_slice(&src_ip);
frame[offset + 16..offset + 20].copy_from_slice(&dst_ip);
let tcp_offset = offset + IPV4_MIN_HEADER_LEN;
frame[tcp_offset] = (src_port >> 8) as u8;
frame[tcp_offset + 1] = (src_port & 0xFF) as u8;
frame[tcp_offset + 2] = (dst_port >> 8) as u8;
frame[tcp_offset + 3] = (dst_port & 0xFF) as u8;
frame[tcp_offset + 4..tcp_offset + 8].copy_from_slice(&[0x00, 0x00, 0x00, 0x01]);
frame[tcp_offset + 8..tcp_offset + 12].copy_from_slice(&[0x00, 0x00, 0x00, 0x00]);
frame[tcp_offset + 12] = 0x50;
frame[tcp_offset + 13] = 0x02;
frame[tcp_offset + 14] = 0x10;
frame[tcp_offset + 15] = 0x00;
frame[tcp_offset + 16] = 0x00;
frame[tcp_offset + 17] = 0x00;
ETH_HEADER_LEN + IPV4_MIN_HEADER_LEN + TCP_MIN_HEADER_LEN
}
#[test]
fn test_worker_creation() {
let config = create_test_config();
let admission = SourceAdmissionEngine::allow_all();
let worker = Worker::new(1, config, 256, admission);
assert!(worker.is_ok());
let worker = worker.unwrap();
assert_eq!(worker.id(), 1);
assert_eq!(worker.state(), WorkerState::Created);
assert_eq!(worker.frame_pool().capacity(), 256);
}
#[test]
fn test_worker_start_stop() {
let config = create_test_config();
let admission = SourceAdmissionEngine::allow_all();
let mut worker = Worker::new(1, config, 256, admission).unwrap();
assert_eq!(worker.state(), WorkerState::Created);
worker.start();
assert_eq!(worker.state(), WorkerState::Running);
worker.stop();
assert_eq!(worker.state(), WorkerState::Stopped);
}
#[test]
fn test_worker_not_running() {
let config = create_test_config();
let admission = SourceAdmissionEngine::allow_all();
let mut worker = Worker::new(1, config, 256, admission).unwrap();
let result = worker.process_cycle();
assert!(result.is_err());
}
#[test]
fn test_worker_process_cycle_empty() {
let config = create_test_config();
let admission = SourceAdmissionEngine::allow_all();
let mut worker = Worker::new(1, config, 256, admission).unwrap();
worker.start();
let result = worker.process_cycle();
assert!(result.is_ok());
let processed = result.unwrap();
assert_eq!(processed, 0);
let stats = worker.stats();
assert!(stats.frame_allocs > 0); assert_eq!(stats.rx_packets, 0);
}
#[test]
fn test_worker_with_admission_deny() {
let config = create_test_config();
let mut admission = SourceAdmissionEngine::deny_all();
admission
.add_rule(crate::source_admission::AdmissionRule {
id: 1,
src_ip: IpAddr::V4_WILDCARD,
prefix_len: 0,
src_port: 0,
dst_port: 8080,
proto: crate::source_admission::ProtoMatch::Tcp,
action: AdmissionAction::Allow,
enabled: true,
})
.unwrap();
let mut worker = Worker::new(1, config, 256, admission).unwrap();
worker.start();
let result = worker.process_cycle();
assert!(result.is_ok());
}
#[test]
fn test_worker_conservation_initial() {
let config = create_test_config();
let admission = SourceAdmissionEngine::allow_all();
let worker = Worker::new(1, config, 256, admission).unwrap();
assert!(worker.verify_conservation());
}
#[test]
fn test_worker_with_frame_pool() {
let config = create_test_config();
let pool = FramePool::new("external", 512, FRAME_SIZE as u32);
let admission = SourceAdmissionEngine::allow_all();
let worker = Worker::with_frame_pool(1, config, pool, admission);
assert!(worker.is_ok());
let worker = worker.unwrap();
let (capacity, allocated, quarantined) = worker.frame_pool_info();
assert_eq!(capacity, 512);
assert_eq!(allocated, 0);
assert_eq!(quarantined, 0);
}
#[test]
fn test_worker_full_cycle_with_data() {
let config = create_test_config();
let admission = SourceAdmissionEngine::allow_all();
let mut worker = Worker::new(1, config, 256, admission).unwrap();
worker.start();
let result = worker.process_cycle();
assert!(result.is_ok());
let packet_data = |_frame_idx: u32, data: &mut [u8]| {
create_test_packet(
data,
[0xAA, 0xBB, 0xCC, 0xDD, 0xEE, 0xFF],
[0x11, 0x22, 0x33, 0x44, 0x55, 0x66],
[10, 0, 0, 1],
[10, 0, 0, 2],
12345,
80,
);
};
worker.simulate_rx_transfer(4, FRAME_SIZE, packet_data);
let result = worker.process_cycle();
assert!(result.is_ok());
let processed = result.unwrap();
assert_eq!(processed, 4); }
#[test]
fn test_worker_admission_deny_specific() {
let config = create_test_config();
let mut admission = SourceAdmissionEngine::deny_all();
admission
.add_rule(crate::source_admission::AdmissionRule {
id: 1,
src_ip: IpAddr::V4([10, 0, 0, 1]),
prefix_len: 0,
src_port: 0,
dst_port: 8080,
proto: crate::source_admission::ProtoMatch::Tcp,
action: AdmissionAction::Allow,
enabled: true,
})
.unwrap();
let mut worker = Worker::new(1, config, 256, admission).unwrap();
worker.start();
let _ = worker.process_cycle();
let packet_data = |_frame_idx: u32, data: &mut [u8]| {
create_test_packet(
data,
[0xAA, 0xBB, 0xCC, 0xDD, 0xEE, 0xFF],
[0x11, 0x22, 0x33, 0x44, 0x55, 0x66],
[10, 0, 0, 1],
[10, 0, 0, 2],
12345,
80, );
};
worker.simulate_rx_transfer(2, FRAME_SIZE, packet_data);
let result = worker.process_cycle();
assert!(result.is_ok());
let processed = result.unwrap();
assert_eq!(processed, 2);
let stats = worker.stats();
assert_eq!(stats.rejected_packets, 2); }
#[test]
fn test_worker_multiple_cycles() {
let config = create_test_config();
let admission = SourceAdmissionEngine::allow_all();
let mut worker = Worker::new(1, config, 256, admission).unwrap();
worker.start();
for _ in 0..5 {
let result = worker.process_cycle();
assert!(result.is_ok());
}
let stats = worker.stats();
assert!(stats.cycles_completed >= 5);
}
#[test]
fn test_extract_flow_info_ipv4_tcp() {
let config = create_test_config();
let admission = SourceAdmissionEngine::allow_all();
let worker = Worker::new(1, config, 256, admission).unwrap();
let mut data = vec![0u8; 100];
create_test_packet(
&mut data,
[0xAA, 0xBB, 0xCC, 0xDD, 0xEE, 0xFF],
[0x11, 0x22, 0x33, 0x44, 0x55, 0x66],
[192, 168, 1, 100],
[10, 0, 0, 1],
54321,
443,
);
let parsed = parse_packet(&data).unwrap();
let (src_ip, src_port, dst_port, protocol) = worker.extract_flow_info(&parsed);
assert_eq!(src_ip, IpAddr::V4([192, 168, 1, 100]));
assert_eq!(src_port, 54321);
assert_eq!(dst_port, 443);
assert_eq!(protocol, 6); }
#[test]
fn test_extract_flow_info_udp() {
let config = create_test_config();
let admission = SourceAdmissionEngine::allow_all();
let worker = Worker::new(1, config, 256, admission).unwrap();
let mut data = vec![0u8; 100];
let offset = ETH_HEADER_LEN;
data[0..6].copy_from_slice(&[0xAA; 6]);
data[6..12].copy_from_slice(&[0xBB; 6]);
data[12] = 0x08;
data[13] = 0x00;
data[offset] = 0x45;
data[offset + 2] = 0x00;
data[offset + 3] = 28; data[offset + 8] = 64;
data[offset + 9] = ip_proto::UDP;
data[offset + 12..offset + 16].copy_from_slice(&[172, 16, 0, 1]);
data[offset + 16..offset + 20].copy_from_slice(&[8, 8, 8, 8]);
let udp_offset = offset + IPV4_MIN_HEADER_LEN;
data[udp_offset] = 0x00;
data[udp_offset + 1] = 53; data[udp_offset + 2] = 0x00;
data[udp_offset + 3] = 53; data[udp_offset + 4] = 0x00;
data[udp_offset + 5] = 8;
let parsed = parse_packet(&data).unwrap();
let (src_ip, src_port, dst_port, protocol) = worker.extract_flow_info(&parsed);
assert_eq!(src_ip, IpAddr::V4([172, 16, 0, 1]));
assert_eq!(src_port, 53);
assert_eq!(dst_port, 53);
assert_eq!(protocol, 17); }
fn make_test_datagram() -> UdpDatagram {
UdpDatagram {
src_ip: IpAddr::V4([172, 16, 0, 1]),
dst_ip: IpAddr::V4([10, 0, 0, 1]),
src_port: 12345,
dst_port: 443,
src_mac: [0xAA, 0xBB, 0xCC, 0xDD, 0xEE, 0xFF],
dst_mac: [0x11, 0x22, 0x33, 0x44, 0x55, 0x66],
payload: smallvec::SmallVec::from_slice(b"request"),
}
}
#[test]
fn test_build_reply_frame_layout() {
let dgram = make_test_datagram();
let payload = b"response-payload";
let mut buf = [0u8; 2048];
let len = dgram.build_reply_frame(payload, &mut buf).unwrap();
assert_eq!(len, 14 + 20 + 8 + payload.len());
assert_eq!(&buf[0..6], &[0xAA, 0xBB, 0xCC, 0xDD, 0xEE, 0xFF]); assert_eq!(&buf[6..12], &[0x11, 0x22, 0x33, 0x44, 0x55, 0x66]); assert_eq!(&buf[12..14], &[0x08, 0x00]);
let ip = 14;
assert_eq!(buf[ip], 0x45);
assert_eq!(u16::from_be_bytes([buf[ip + 2], buf[ip + 3]]), (20 + 8 + payload.len()) as u16);
assert_eq!(buf[ip + 6], 0x40, "DF 位必须置位(QUIC 禁止分片)");
assert_eq!(buf[ip + 8], 64);
assert_eq!(buf[ip + 9], 17);
assert_eq!(&buf[ip + 12..ip + 16], &[10, 0, 0, 1]); assert_eq!(&buf[ip + 16..ip + 20], &[172, 16, 0, 1]);
let mut sum: u32 = 0;
for pair in buf[ip..ip + 20].chunks_exact(2) {
sum += u32::from(u16::from_be_bytes([pair[0], pair[1]]));
}
while (sum >> 16) != 0 {
sum = (sum & 0xFFFF) + (sum >> 16);
}
assert_eq!(sum as u16, 0xFFFF, "IPv4 头校验和验证失败");
let l4 = ip + 20;
assert_eq!(u16::from_be_bytes([buf[l4], buf[l4 + 1]]), 443); assert_eq!(u16::from_be_bytes([buf[l4 + 2], buf[l4 + 3]]), 12345); assert_eq!(u16::from_be_bytes([buf[l4 + 4], buf[l4 + 5]]), (8 + payload.len()) as u16);
assert_eq!(&buf[l4 + 6..l4 + 8], &[0, 0]);
assert_eq!(&buf[l4 + 8..len], payload);
}
#[test]
fn test_build_reply_frame_fail_closed() {
let dgram = make_test_datagram();
let payload = b"x";
let mut tiny = [0u8; 41]; assert!(dgram.build_reply_frame(payload, &mut tiny).is_none());
let v6 = UdpDatagram {
src_ip: IpAddr::V6([0u8; 16]),
dst_ip: IpAddr::V6([0u8; 16]),
..make_test_datagram()
};
let mut buf = [0u8; 2048];
assert!(v6.build_reply_frame(payload, &mut buf).is_none());
let huge = vec![0u8; 70000];
assert!(dgram.build_reply_frame(&huge, &mut vec![0u8; 70100]).is_none());
}
#[test]
fn test_udp_datagram_src_socket_addr() {
let dgram = make_test_datagram();
let addr = dgram.src_socket_addr();
assert_eq!(addr.ip(), std::net::IpAddr::V4(std::net::Ipv4Addr::new(172, 16, 0, 1)));
assert_eq!(addr.port(), 12345);
}
#[test]
fn test_worker_rx_desc_preallocated() {
let config = create_test_config();
let admission = SourceAdmissionEngine::allow_all();
let mut worker = Worker::new(1, config, 256, admission).unwrap();
worker.start();
assert_eq!(worker.rx_desc_count, 0);
let result = worker.process_cycle();
assert!(result.is_ok());
assert_eq!(worker.rx_desc_count, 0);
}
#[test]
fn test_worker_quarantine_config() {
let config = create_test_config();
let admission = SourceAdmissionEngine::allow_all();
let mut worker = Worker::new(1, config, 256, admission).unwrap();
assert!(!worker.quarantine_on_error());
assert!(!worker.quarantine_on_deny());
worker.set_quarantine_on_error(true);
worker.set_quarantine_on_deny(true);
assert!(worker.quarantine_on_error());
assert!(worker.quarantine_on_deny());
}
#[test]
fn test_worker_id() {
let config = create_test_config();
let admission = SourceAdmissionEngine::allow_all();
let worker = Worker::new(42, config, 256, admission).unwrap();
assert_eq!(worker.id(), 42);
}
#[test]
fn test_worker_state_transitions() {
let config = create_test_config();
let admission = SourceAdmissionEngine::allow_all();
let mut worker = Worker::new(1, config, 256, admission).unwrap();
assert_eq!(worker.state(), WorkerState::Created);
worker.start();
assert_eq!(worker.state(), WorkerState::Running);
worker.stop();
assert_eq!(worker.state(), WorkerState::Stopped);
}
#[test]
fn test_worker_stats_initial() {
let config = create_test_config();
let admission = SourceAdmissionEngine::allow_all();
let worker = Worker::new(1, config, 256, admission).unwrap();
let stats = worker.stats();
assert_eq!(stats.rx_packets, 0);
assert_eq!(stats.tx_packets, 0);
assert_eq!(stats.rejected_packets, 0);
assert_eq!(stats.parse_errors, 0);
}
#[test]
fn test_worker_double_start() {
let config = create_test_config();
let admission = SourceAdmissionEngine::allow_all();
let mut worker = Worker::new(1, config, 256, admission).unwrap();
worker.start();
assert_eq!(worker.state(), WorkerState::Running);
worker.start();
assert_eq!(worker.state(), WorkerState::Running);
}
#[test]
fn test_worker_stop_before_start() {
let config = create_test_config();
let admission = SourceAdmissionEngine::allow_all();
let mut worker = Worker::new(1, config, 256, admission).unwrap();
worker.stop();
assert_eq!(worker.state(), WorkerState::Stopped);
}
#[test]
fn test_worker_frame_pool_info() {
let config = create_test_config();
let admission = SourceAdmissionEngine::allow_all();
let worker = Worker::new(1, config, 256, admission).unwrap();
let (capacity, allocated, quarantined) = worker.frame_pool_info();
assert_eq!(capacity, 256);
assert_eq!(allocated, 0);
assert_eq!(quarantined, 0);
}
#[test]
fn test_worker_admission_engine_link() {
let config = create_test_config();
let mut admission = SourceAdmissionEngine::deny_all();
admission
.add_rule(AdmissionRule {
id: 1,
src_ip: IpAddr::V4([10, 0, 0, 1]),
prefix_len: 0,
src_port: 0,
dst_port: 0,
proto: ProtoMatch::Any,
action: AdmissionAction::Allow,
enabled: true,
})
.unwrap();
let worker = Worker::new(1, config, 256, admission).unwrap();
let action = worker.admission().evaluate(
IpAddr::V4([10, 0, 0, 1]),
0,
0,
6,
);
assert_eq!(action, AdmissionAction::Allow);
let action2 = worker.admission().evaluate(
IpAddr::V4([10, 0, 0, 2]),
0,
0,
6,
);
assert_eq!(action2, AdmissionAction::Deny);
}
#[test]
fn test_worker_state_enum() {
let states = vec![
WorkerState::Created,
WorkerState::Running,
WorkerState::Stopped,
WorkerState::Error,
];
for state in states {
let debug = format!("{:?}", state);
assert!(!debug.is_empty());
}
}
#[test]
fn test_worker_with_deny_all_admission() {
let config = create_test_config();
let admission = SourceAdmissionEngine::deny_all();
let worker = Worker::new(1, config, 256, admission);
assert!(worker.is_ok());
}
#[test]
fn test_extract_flow_info_icmp() {
let config = create_test_config();
let admission = SourceAdmissionEngine::allow_all();
let worker = Worker::new(1, config, 256, admission).unwrap();
let mut data = vec![0u8; 100];
let offset = ETH_HEADER_LEN;
data[0..6].copy_from_slice(&[0xAA; 6]);
data[6..12].copy_from_slice(&[0xBB; 6]);
data[12] = 0x08;
data[13] = 0x00;
data[offset] = 0x45;
data[offset + 2] = 0x00;
data[offset + 3] = 28;
data[offset + 8] = 64;
data[offset + 9] = ip_proto::ICMP;
data[offset + 12..offset + 16].copy_from_slice(&[10, 0, 0, 1]);
data[offset + 16..offset + 20].copy_from_slice(&[10, 0, 0, 2]);
let parsed = parse_packet(&data).unwrap();
let (src_ip, src_port, dst_port, protocol) = worker.extract_flow_info(&parsed);
assert_eq!(src_ip, IpAddr::V4([10, 0, 0, 1]));
assert_eq!(src_port, 0);
assert_eq!(dst_port, 0);
assert_eq!(protocol, 1);
}
}