use std::io::{IoSliceMut, Read};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::{Duration, Instant};
use crossbeam_channel::{Receiver, RecvTimeoutError};
use crate::guest_mem::GuestMemWriter;
use crate::inline_conn::{self, InlineConn};
use crate::irq::IrqHandle;
use crate::queue::{self, RxQueueConfig};
const BATCH_SIZE: usize = 256;
const COALESCE_TIMEOUT: Duration = Duration::from_micros(200);
const DESCRIPTOR_BACKOFF: Duration = Duration::from_micros(100);
pub struct RxInjectThread {
pub rx: Receiver<Vec<u8>>,
pub conn_rx: Receiver<InlineConn>,
pub guest_mem: GuestMemWriter,
pub queue: RxQueueConfig,
pub irq: IrqHandle,
pub set_interrupt_status: Arc<dyn Fn() + Send + Sync>,
pub running: Arc<AtomicBool>,
pub event_idx_enabled: bool,
}
unsafe impl Send for RxInjectThread {}
impl RxInjectThread {
pub fn run(self) {
tracing::info!(
"rx-inject thread started (queue_size={}, event_idx={})",
self.queue.size,
self.event_idx_enabled,
);
let mut used_idx = self.guest_mem.read_u16(self.queue.used_gpa as usize + 2);
let mut old_used = used_idx;
let mut inline_conns: Vec<InlineConn> = Vec::new();
loop {
if !self.running.load(Ordering::Relaxed) {
break;
}
while let Ok(conn) = self.conn_rx.try_recv() {
tracing::info!(
"inline conn added: {}:{} -> {}:{}",
conn.remote_ip,
conn.remote_port,
conn.guest_ip,
conn.guest_port,
);
inline_conns.push(conn);
}
let mut batch = 0u16;
let loop_start = Instant::now();
if !inline_conns.is_empty() {
self.poll_inline_conns(&mut inline_conns, &mut used_idx, &mut batch);
}
let elapsed = loop_start.elapsed();
let remaining = COALESCE_TIMEOUT.saturating_sub(elapsed);
while (batch as usize) < BATCH_SIZE {
let timeout = if batch == 0 && inline_conns.is_empty() {
COALESCE_TIMEOUT
} else if remaining.is_zero() {
Duration::ZERO
} else {
remaining
};
let frame = match self.rx.recv_timeout(timeout) {
Ok(f) => f,
Err(RecvTimeoutError::Timeout) => break,
Err(RecvTimeoutError::Disconnected) => {
tracing::info!("rx-inject: channel disconnected, shutting down");
if batch > 0 {
self.flush_interrupt(old_used, used_idx);
}
return;
}
};
if queue::inject_one_frame(&self.guest_mem, &self.queue, &frame, &mut used_idx) {
batch += 1;
} else {
if batch > 0 {
self.flush_interrupt(old_used, used_idx);
old_used = used_idx;
batch = 0;
}
std::thread::sleep(DESCRIPTOR_BACKOFF);
if queue::inject_one_frame(&self.guest_mem, &self.queue, &frame, &mut used_idx)
{
batch += 1;
}
}
}
if batch > 0 {
self.flush_interrupt(old_used, used_idx);
old_used = used_idx;
}
}
if old_used != used_idx {
self.flush_interrupt(old_used, used_idx);
}
tracing::info!(
"rx-inject thread stopped ({} inline conns remaining)",
inline_conns.len(),
);
}
fn poll_inline_conns(
&self,
inline_conns: &mut Vec<InlineConn>,
used_idx: &mut u16,
batch: &mut u16,
) {
let q_size = self.queue.size as usize;
if q_size == 0 {
return;
}
inline_conns.retain(|c| !c.host_eof);
const PER_CONN_READS: u16 = 16;
const MAX_MERGE: usize = 16;
const MAX_FRAME_PAYLOAD: usize = 60000;
let desc_base = self.queue.desc_gpa as usize;
let mut head_indices: [u16; MAX_MERGE] = [0; MAX_MERGE];
let mut desc_ptrs: [*mut u8; MAX_MERGE] = [std::ptr::null_mut(); MAX_MERGE];
let mut desc_lens: [usize; MAX_MERGE] = [0; MAX_MERGE];
for conn in inline_conns.iter_mut() {
let mut per_conn = 0u16;
loop {
if (*batch as usize) >= BATCH_SIZE {
break;
}
if per_conn >= PER_CONN_READS {
break;
}
std::sync::atomic::fence(Ordering::Acquire);
let avail_idx = self.guest_mem.read_u16(self.queue.avail_gpa as usize + 2);
let available = avail_idx.wrapping_sub(*used_idx) as usize;
if available == 0 {
return;
}
let want = available.min(MAX_MERGE);
let mut count = 0usize;
let mut total_iov_cap = 0usize;
for i in 0..want {
let slot = (*used_idx).wrapping_add(i as u16) as usize % q_size;
let ring_off = self.queue.avail_gpa as usize + 4 + 2 * slot;
let head_idx = self.guest_mem.read_u16(ring_off);
let d_off = desc_base + (head_idx as usize) * 16;
let Some(desc_slice) = self.guest_mem.slice(d_off, 16) else {
break;
};
let addr_gpa =
u64::from_le_bytes(desc_slice[0..8].try_into().unwrap()) as usize;
let buf_len =
u32::from_le_bytes(desc_slice[8..12].try_into().unwrap()) as usize;
let flags = u16::from_le_bytes(desc_slice[12..14].try_into().unwrap());
if flags & 2 == 0 {
break;
}
let min_len = if i == 0 {
inline_conn::TOTAL_HDR_LEN + 1
} else {
1
};
if buf_len < min_len {
break;
}
let Some(off) = self.guest_mem.gpa_to_offset(addr_gpa, buf_len) else {
break;
};
let ptr = unsafe { self.guest_mem.ptr().add(off) };
let iov_cap = if i == 0 {
buf_len - inline_conn::TOTAL_HDR_LEN
} else {
buf_len
};
if i > 0 && total_iov_cap + iov_cap > MAX_FRAME_PAYLOAD {
break;
}
head_indices[i] = head_idx;
desc_ptrs[i] = ptr;
desc_lens[i] = buf_len;
total_iov_cap += iov_cap;
count += 1;
}
if count == 0 {
break;
}
let mut iovs: [IoSliceMut<'_>; MAX_MERGE] = std::array::from_fn(|_| {
IoSliceMut::new(&mut [])
});
for i in 0..count {
let (start, cap) = if i == 0 {
(
inline_conn::TOTAL_HDR_LEN,
desc_lens[i] - inline_conn::TOTAL_HDR_LEN,
)
} else {
(0, desc_lens[i])
};
let slice =
unsafe { std::slice::from_raw_parts_mut(desc_ptrs[i].add(start), cap) };
iovs[i] = IoSliceMut::new(slice);
}
let read_result = conn.stream.read_vectored(&mut iovs[..count]);
match read_result {
Ok(0) => {
tracing::debug!(
"inline {}:{}->{}:{} host EOF",
conn.remote_ip,
conn.remote_port,
conn.guest_ip,
conn.guest_port
);
conn.host_eof = true;
let first_buf =
unsafe { std::slice::from_raw_parts_mut(desc_ptrs[0], desc_lens[0]) };
inline_conn::write_fin_headers(first_buf, conn);
conn.our_seq.fetch_add(1, Ordering::Relaxed);
let used_entry_off =
self.queue.used_gpa as usize + 4 + ((*used_idx as usize) % q_size) * 8;
self.guest_mem
.write_u32(used_entry_off, head_indices[0] as u32);
self.guest_mem
.write_u32(used_entry_off + 4, inline_conn::TOTAL_HDR_LEN as u32);
std::sync::atomic::fence(Ordering::Release);
*used_idx = used_idx.wrapping_add(1);
self.guest_mem
.write_u16(self.queue.used_gpa as usize + 2, *used_idx);
*batch += 1;
break;
}
Ok(n) => {
let mut remaining = n;
let mut num_used = 0usize;
let mut per_desc_len = [0usize; MAX_MERGE];
for i in 0..count {
if remaining == 0 {
break;
}
let cap = if i == 0 {
desc_lens[i] - inline_conn::TOTAL_HDR_LEN
} else {
desc_lens[i]
};
let filled = remaining.min(cap);
per_desc_len[i] = filled;
remaining -= filled;
num_used = i + 1;
}
let first_buf =
unsafe { std::slice::from_raw_parts_mut(desc_ptrs[0], desc_lens[0]) };
inline_conn::write_inline_headers(first_buf, conn, n, num_used as u16);
conn.our_seq.fetch_add(n as u32, Ordering::Relaxed);
for i in 0..num_used {
let slot = (*used_idx).wrapping_add(i as u16) as usize % q_size;
let used_entry_off = self.queue.used_gpa as usize + 4 + slot * 8;
let entry_len = if i == 0 {
inline_conn::TOTAL_HDR_LEN + per_desc_len[0]
} else {
per_desc_len[i]
};
self.guest_mem
.write_u32(used_entry_off, head_indices[i] as u32);
self.guest_mem
.write_u32(used_entry_off + 4, entry_len as u32);
}
std::sync::atomic::fence(Ordering::Release);
*used_idx = used_idx.wrapping_add(num_used as u16);
self.guest_mem
.write_u16(self.queue.used_gpa as usize + 2, *used_idx);
*batch += num_used as u16;
per_conn += 1;
}
Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => {
break;
}
Err(e) => {
tracing::debug!("inline conn error: {e}");
conn.host_eof = true;
break;
}
}
}
}
}
fn flush_interrupt(&self, old_used: u16, used_idx: u16) {
let q_size = self.queue.size as usize;
let avail_event_off = self.queue.used_gpa as usize + 4 + q_size * 8;
let avail_idx = self.guest_mem.read_u16(self.queue.avail_gpa as usize + 2);
std::sync::atomic::fence(Ordering::Release);
self.guest_mem.write_u16(avail_event_off, avail_idx);
let fire = if self.event_idx_enabled {
std::sync::atomic::fence(Ordering::AcqRel);
crate::notify::should_notify(
&self.guest_mem,
self.queue.avail_gpa,
q_size as u16,
old_used,
used_idx,
)
} else {
old_used != used_idx
};
if fire {
(self.set_interrupt_status)();
self.irq.trigger();
}
}
}