use std::collections::HashMap;
use std::sync::{Arc, Mutex, RwLock};
use arcbox_virtio::{DeviceStatus, QueueConfig, VirtioDevice};
use crate::error::{Result, VmmError};
use crate::irq::{Irq, IrqChip};
use crate::memory::MemoryManager;
mod checksum;
mod debug;
mod dispatch;
mod mmio_state;
pub(crate) mod net_worker;
mod poll;
mod tree;
#[cfg(test)]
mod tests;
use checksum::finalize_virtio_net_checksum;
pub use debug::{DeviceDebug, QueueDebug};
pub use mmio_state::{MAX_VIRTQUEUES, MmioDevice, VirtioMmioState, virtio_mmio};
pub use tree::DeviceTreeEntry;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct DeviceId(u32);
impl DeviceId {
#[must_use]
pub const fn new(id: u32) -> Self {
Self(id)
}
#[must_use]
pub const fn raw(&self) -> u32 {
self.0
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum DeviceType {
Serial,
VirtioBlock,
VirtioNet,
VirtioConsole,
VirtioFs,
VirtioVsock,
VirtioRng,
VirtioBalloon,
Other,
}
#[derive(Debug)]
pub struct DeviceInfo {
pub id: DeviceId,
pub device_type: DeviceType,
pub name: String,
pub mmio_base: Option<u64>,
pub mmio_size: u64,
pub irq: Option<Irq>,
}
pub struct RegisteredDevice {
pub info: DeviceInfo,
pub mmio_state: Option<Arc<RwLock<VirtioMmioState>>>,
pub virtio_device: Option<Arc<Mutex<dyn VirtioDevice>>>,
}
pub type DeviceIrqCallback = Arc<dyn Fn(Irq, bool) -> Result<()> + Send + Sync>;
pub struct DeviceManager {
devices: HashMap<DeviceId, RegisteredDevice>,
next_id: u32,
mmio_map: HashMap<u64, DeviceId>,
guest_ram_base: Option<*mut u8>,
guest_ram_size: usize,
guest_ram_gpa: u64,
irq_callback: Option<DeviceIrqCallback>,
primary_net_device_id: Option<DeviceId>,
primary_net: Option<Arc<Mutex<arcbox_virtio::net::VirtioNet>>>,
bridge_net_device_id: Option<DeviceId>,
bridge_net: Option<Arc<Mutex<arcbox_virtio::net::VirtioNet>>>,
vsock_connections:
std::sync::Arc<std::sync::Mutex<crate::vsock_manager::VsockConnectionManager>>,
vsock: Option<Arc<Mutex<arcbox_virtio::vsock::VirtioVsock>>>,
console: Option<Arc<Mutex<arcbox_virtio::console::VirtioConsole>>>,
debug_console_socket: Option<Arc<Mutex<arcbox_virtio::console::SocketConsole>>>,
blk_workers: Mutex<HashMap<DeviceId, crate::blk_worker::BlkWorkerHandle>>,
net_rx_worker: net_worker::NetRxWorkerSlot,
}
unsafe impl Send for DeviceManager {}
unsafe impl Sync for DeviceManager {}
impl DeviceManager {
#[must_use]
pub fn new() -> Self {
Self {
devices: HashMap::new(),
next_id: 0,
mmio_map: HashMap::new(),
guest_ram_base: None,
guest_ram_size: 0,
guest_ram_gpa: 0,
irq_callback: None,
primary_net_device_id: None,
primary_net: None,
bridge_net_device_id: None,
bridge_net: None,
vsock_connections: std::sync::Arc::new(std::sync::Mutex::new(
crate::vsock_manager::VsockConnectionManager::new(),
)),
vsock: None,
console: None,
debug_console_socket: None,
blk_workers: Mutex::new(HashMap::new()),
net_rx_worker: net_worker::NetRxWorkerSlot::new(),
}
}
pub unsafe fn set_guest_memory(&mut self, base: *mut u8, size: usize, gpa_base: u64) {
self.guest_ram_base = Some(base);
self.guest_ram_size = size;
self.guest_ram_gpa = gpa_base;
}
pub fn set_irq_callback(&mut self, callback: DeviceIrqCallback) {
self.irq_callback = Some(callback);
}
pub fn set_net_host_fd(&mut self, fd: std::os::unix::io::RawFd, device_id: DeviceId) {
use arcbox_virtio::net::NetPort;
self.primary_net_device_id = Some(device_id);
self.net_rx_worker.set_host_fd(fd);
if let Some(primary) = self.primary_net.as_ref() {
let port = NetPort {
host_fd: fd,
last_avail_tx: std::sync::atomic::AtomicU16::new(0),
};
if let Ok(dev) = primary.lock() {
if dev.bind_port(port).is_err() {
tracing::warn!("primary_net port already bound — ignoring rebind");
}
}
} else {
tracing::error!("set_net_host_fd called before set_primary_net");
}
}
pub fn set_primary_net(
&mut self,
device_id: DeviceId,
device: Arc<Mutex<arcbox_virtio::net::VirtioNet>>,
) {
self.primary_net_device_id = Some(device_id);
if let Some(ctx) = self.build_device_ctx(device_id) {
if let Ok(mut dev) = device.lock() {
dev.bind_ctx(ctx);
}
} else {
tracing::warn!(
"set_primary_net: DeviceCtx not built (guest_mem or irq_callback missing) — \
primary NIC hot paths will be no-ops"
);
}
self.primary_net = Some(device);
}
pub fn primary_net_device_id(&self) -> Option<DeviceId> {
self.primary_net_device_id
}
pub fn primary_net(&self) -> Option<&Arc<Mutex<arcbox_virtio::net::VirtioNet>>> {
self.primary_net.as_ref()
}
pub fn set_vsock(
&mut self,
device_id: DeviceId,
device: Arc<Mutex<arcbox_virtio::vsock::VirtioVsock>>,
) {
if let Some(ctx) = self.build_device_ctx(device_id) {
if let Ok(mut dev) = device.lock() {
dev.bind_ctx(ctx);
dev.bind_connection_manager(self.vsock_connections.clone());
}
} else {
tracing::warn!(
"set_vsock: DeviceCtx not built (guest_mem or irq_callback missing) — \
vsock TX hot path will fall back to QueueConfig plumbing"
);
}
self.vsock = Some(device);
}
pub fn vsock(&self) -> Option<&Arc<Mutex<arcbox_virtio::vsock::VirtioVsock>>> {
self.vsock.as_ref()
}
pub fn set_console(&mut self, device: Arc<Mutex<arcbox_virtio::console::VirtioConsole>>) {
self.console = Some(device);
}
pub fn set_debug_console_socket(
&mut self,
socket: Arc<Mutex<arcbox_virtio::console::SocketConsole>>,
) {
self.debug_console_socket = Some(socket);
}
pub fn debug_console_socket(
&self,
) -> Option<&Arc<Mutex<arcbox_virtio::console::SocketConsole>>> {
self.debug_console_socket.as_ref()
}
pub fn set_bridge_net(
&mut self,
device_id: DeviceId,
device: Arc<Mutex<arcbox_virtio::net::VirtioNet>>,
) {
self.bridge_net_device_id = Some(device_id);
if let Some(ctx) = self.build_device_ctx(device_id) {
if let Ok(mut dev) = device.lock() {
dev.bind_ctx(ctx);
}
} else {
tracing::warn!(
"set_bridge_net: DeviceCtx not built (guest_mem or irq_callback missing) — \
bridge hot paths will be no-ops"
);
}
self.bridge_net = Some(device);
}
fn build_device_ctx(&self, device_id: DeviceId) -> Option<arcbox_virtio::DeviceCtx> {
let ram_base = self.guest_ram_base?;
if self.guest_ram_size == 0 {
return None;
}
let device = self.devices.get(&device_id)?;
let irq = device.info.irq?;
let mmio_arc = device.mmio_state.as_ref()?.clone();
let irq_callback = self.irq_callback.as_ref()?.clone();
let mem = unsafe {
arcbox_virtio::GuestMemWriter::new(
ram_base,
self.guest_ram_size,
self.guest_ram_gpa as usize,
)
};
let raise_irq: Arc<dyn Fn(u32) + Send + Sync> = Arc::new(move |reason: u32| {
if let Ok(mut s) = mmio_arc.write() {
s.trigger_interrupt(reason);
}
let _ = irq_callback(irq, true);
});
Some(arcbox_virtio::DeviceCtx {
mem: Arc::new(mem),
raise_irq,
})
}
pub fn set_bridge_host_fd(&mut self, fd: std::os::unix::io::RawFd, _device_id: DeviceId) {
use arcbox_virtio::net::NetPort;
let Some(bridge) = self.bridge_net.as_ref() else {
tracing::error!("set_bridge_host_fd called before set_bridge_net");
return;
};
let port = NetPort {
host_fd: fd,
last_avail_tx: std::sync::atomic::AtomicU16::new(0),
};
if let Ok(dev) = bridge.lock() {
if dev.bind_port(port).is_err() {
tracing::warn!("bridge_net port already bound — ignoring rebind");
}
}
}
pub fn bridge_net(&self) -> Option<&Arc<Mutex<arcbox_virtio::net::VirtioNet>>> {
self.bridge_net.as_ref()
}
pub fn guest_ram_base_ptr(&self) -> Option<*mut u8> {
self.guest_ram_base
}
pub fn guest_ram_size(&self) -> usize {
self.guest_ram_size
}
pub fn guest_ram_gpa(&self) -> u64 {
self.guest_ram_gpa
}
pub fn get_registered_device(&self, id: DeviceId) -> Option<&RegisteredDevice> {
self.devices.get(&id)
}
pub fn set_blk_worker(&self, device_id: DeviceId, handle: crate::blk_worker::BlkWorkerHandle) {
self.blk_workers
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.insert(device_id, handle);
}
pub fn clear_blk_workers(&self) {
self.blk_workers
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.clear();
}
pub fn set_net_rx_hooks(
&mut self,
irq_callback: Arc<dyn Fn(crate::irq::Irq, bool) -> crate::error::Result<()> + Send + Sync>,
exit_vcpus: Arc<dyn Fn() + Send + Sync>,
) {
self.net_rx_worker.set_hooks(irq_callback, exit_vcpus);
}
pub fn set_running(&mut self, running: Arc<std::sync::atomic::AtomicBool>) {
self.net_rx_worker.set_running(running);
}
pub fn set_rx_inject_channel(&mut self, rx: crossbeam_channel::Receiver<Vec<u8>>) {
self.net_rx_worker.set_rx_inject_channel(rx);
}
pub fn set_inline_conn_channel(
&mut self,
rx: crossbeam_channel::Receiver<arcbox_net_inject::inline_conn::InlineConn>,
) {
self.net_rx_worker.set_inline_conn_channel(rx);
}
pub fn take_net_rx_worker_handle(&self) -> Option<std::thread::JoinHandle<()>> {
self.net_rx_worker.take_handle()
}
pub(super) fn maybe_spawn_net_rx_worker(
&self,
device_id: DeviceId,
mmio_arc: &Arc<RwLock<VirtioMmioState>>,
) {
if self.primary_net_device_id != Some(device_id) {
return;
}
let device = match self.devices.get(&device_id) {
Some(d) if d.info.device_type == DeviceType::VirtioNet => d,
_ => return,
};
let Some(irq) = device.info.irq else {
tracing::warn!("net-io: device has no IRQ");
return;
};
let Some(guest_base) = self.guest_ram_base else {
tracing::warn!("net-io: guest_ram_base not set");
return;
};
self.net_rx_worker.try_spawn(
mmio_arc,
irq,
guest_base,
self.guest_ram_size,
self.guest_ram_gpa,
);
}
pub fn irq_callback_clone(&self) -> Option<DeviceIrqCallback> {
self.irq_callback.clone()
}
pub fn vsock_connections(
&self,
) -> std::sync::Arc<std::sync::Mutex<crate::vsock_manager::VsockConnectionManager>> {
self.vsock_connections.clone()
}
pub fn sync_irq_level(&self, device_id: DeviceId) {
let Some(device) = self.devices.get(&device_id) else {
return;
};
let Some(irq) = device.info.irq else {
return;
};
let Some(ref mmio_arc) = device.mmio_state else {
return;
};
let Ok(mmio) = mmio_arc.read() else {
return;
};
if mmio.status & DeviceStatus::DRIVER_OK == 0 {
tracing::trace!(
"sync_irq_level: device {:?} not DRIVER_OK (status={:#x}), skipping",
device.info.device_type,
mmio.status,
);
return;
}
let level = mmio.interrupt_status != 0;
tracing::trace!(
"sync_irq_level: device {:?} irq={} interrupt_status={} -> SPI level={}",
device.info.device_type,
irq,
mmio.interrupt_status,
level,
);
if let Some(ref cb) = self.irq_callback {
let _ = cb(irq, level);
}
}
pub fn trigger_irq_callback(&self, irq: Irq, level: bool) {
let device_ready = self.devices.values().any(|d| {
d.info.irq == Some(irq)
&& d.mmio_state
.as_ref()
.and_then(|s| s.read().ok())
.is_some_and(|s| s.status & DeviceStatus::DRIVER_OK != 0)
});
if !device_ready {
return;
}
if let Some(ref cb) = self.irq_callback {
let _ = cb(irq, level);
}
}
pub fn register(
&mut self,
device_type: DeviceType,
name: impl Into<String>,
) -> Result<DeviceId> {
let id = DeviceId::new(self.next_id);
self.next_id += 1;
let info = DeviceInfo {
id,
device_type,
name: name.into(),
mmio_base: None,
mmio_size: 0,
irq: None,
};
self.devices.insert(
id,
RegisteredDevice {
info,
mmio_state: None,
virtio_device: None,
},
);
Ok(id)
}
pub fn register_virtio(
&mut self,
device_type: DeviceType,
name: impl Into<String>,
virtio_device_id: u32,
features: u64,
memory_manager: &mut MemoryManager,
irq_chip: &IrqChip,
) -> Result<DeviceId> {
let id = DeviceId::new(self.next_id);
self.next_id += 1;
let mmio_base = memory_manager.allocate_mmio(virtio_mmio::MMIO_SIZE, &name.into())?;
let irq = irq_chip.allocate_level_irq()?;
let name_str = format!("{}", id.0);
let info = DeviceInfo {
id,
device_type,
name: name_str,
mmio_base: Some(mmio_base),
mmio_size: virtio_mmio::MMIO_SIZE,
irq: Some(irq),
};
let mmio_state = Arc::new(RwLock::new(VirtioMmioState::new(
virtio_device_id,
features,
)));
self.mmio_map.insert(mmio_base, id);
self.devices.insert(
id,
RegisteredDevice {
info,
mmio_state: Some(mmio_state),
virtio_device: None,
},
);
tracing::info!(
"Registered VirtIO device {} at MMIO {:#x}, IRQ {}",
id.0,
mmio_base,
irq
);
Ok(id)
}
pub fn register_virtio_device<D: VirtioDevice + 'static>(
&mut self,
device_type: DeviceType,
name: impl Into<String>,
device: D,
memory_manager: &mut MemoryManager,
irq_chip: &IrqChip,
) -> Result<(DeviceId, Arc<Mutex<D>>)> {
let id = DeviceId::new(self.next_id);
self.next_id += 1;
let virtio_device_id = device.device_id() as u32;
let features = device.features();
let name_str = name.into();
let mmio_base = memory_manager.allocate_mmio(virtio_mmio::MMIO_SIZE, &name_str)?;
let irq = irq_chip.allocate_level_irq()?;
let info = DeviceInfo {
id,
device_type,
name: name_str.clone(),
mmio_base: Some(mmio_base),
mmio_size: virtio_mmio::MMIO_SIZE,
irq: Some(irq),
};
let mmio_state = Arc::new(RwLock::new(VirtioMmioState::new(
virtio_device_id,
features,
)));
let virtio_device: Arc<Mutex<D>> = Arc::new(Mutex::new(device));
let virtio_device_erased: Arc<Mutex<dyn VirtioDevice>> = virtio_device.clone();
self.mmio_map.insert(mmio_base, id);
self.devices.insert(
id,
RegisteredDevice {
info,
mmio_state: Some(mmio_state),
virtio_device: Some(virtio_device_erased),
},
);
tracing::info!(
"Registered VirtIO device '{}' (type {:?}) at MMIO {:#x}, IRQ {}",
name_str,
device_type,
mmio_base,
irq
);
Ok((id, virtio_device))
}
#[must_use]
pub fn get(&self, id: DeviceId) -> Option<&DeviceInfo> {
self.devices.get(&id).map(|d| &d.info)
}
#[must_use]
pub fn get_mmio_state(&self, id: DeviceId) -> Option<Arc<RwLock<VirtioMmioState>>> {
self.devices.get(&id).and_then(|d| d.mmio_state.clone())
}
#[must_use]
pub fn get_virtio_device(&self, id: DeviceId) -> Option<Arc<Mutex<dyn VirtioDevice>>> {
self.devices.get(&id).and_then(|d| d.virtio_device.clone())
}
pub fn trigger_interrupt(&self, id: DeviceId, reason: u32) -> Result<()> {
let device = self
.devices
.get(&id)
.ok_or_else(|| VmmError::Device(format!("Device {} not found", id.0)))?;
if let Some(state) = &device.mmio_state {
let mut state = state
.write()
.map_err(|e| VmmError::Device(format!("Failed to lock device state: {e}")))?;
state.trigger_interrupt(reason);
}
Ok(())
}
pub fn raise_interrupt_for(&self, device_type: DeviceType, reason: u32) {
for (id, dev) in &self.devices {
if dev.info.device_type == device_type {
if let Some(ref mmio_arc) = dev.mmio_state {
if let Ok(mut s) = mmio_arc.write() {
s.trigger_interrupt(reason);
}
}
self.sync_irq_level(*id);
break;
}
}
}
pub fn raise_interrupt_for_device(&self, device_id: DeviceId, reason: u32) {
if let Some(dev) = self.devices.get(&device_id) {
if let Some(ref mmio_arc) = dev.mmio_state {
if let Ok(mut s) = mmio_arc.write() {
s.trigger_interrupt(reason);
}
}
self.sync_irq_level(device_id);
}
}
pub fn bridge_device_id(&self) -> Option<DeviceId> {
self.bridge_net_device_id
}
pub fn iter(&self) -> impl Iterator<Item = &DeviceInfo> {
self.devices.values().map(|d| &d.info)
}
}
impl Default for DeviceManager {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
const _: () = {
#[allow(dead_code)]
fn assert_send<T: Send>() {}
#[allow(dead_code)]
fn assert_sync<T: Sync>() {}
fn _check() {
assert_send::<DeviceManager>();
assert_sync::<DeviceManager>();
}
};