use alloc::vec::Vec;
use axvm_types::GuestPhysAddr;
use crate::{
VirtioDeviceID,
constants::*,
error::{VirtioError, VirtioResult},
memory::GuestMemory,
};
#[repr(C)]
#[derive(Debug, Clone, Copy)]
pub struct VirtQueueDesc {
pub base_addr: GuestPhysAddr,
pub len: u32,
pub flags: u16,
pub next: u16,
}
impl VirtQueueDesc {
pub fn new(base_addr: GuestPhysAddr, len: u32, flags: u16, next: u16) -> Self {
Self {
base_addr,
len,
flags,
next,
}
}
pub fn has_next(&self) -> bool {
(self.flags & VIRTQ_DESC_F_NEXT) != 0
}
pub fn is_write(&self) -> bool {
(self.flags & VIRTQ_DESC_F_WRITE) != 0
}
pub fn is_indirect(&self) -> bool {
(self.flags & VIRTQ_DESC_F_INDIRECT) != 0
}
pub fn guest_addr(&self) -> GuestPhysAddr {
self.base_addr
}
pub fn set_next(&mut self, has_next: bool) {
if has_next {
self.flags |= VIRTQ_DESC_F_NEXT;
} else {
self.flags &= !VIRTQ_DESC_F_NEXT;
}
}
pub fn set_write(&mut self, is_write: bool) {
if is_write {
self.flags |= VIRTQ_DESC_F_WRITE;
} else {
self.flags &= !VIRTQ_DESC_F_WRITE;
}
}
pub fn set_write_only(&mut self, is_write: bool) {
self.set_write(is_write);
}
pub fn is_write_only(&self) -> bool {
self.is_write()
}
pub fn set_indirect(&mut self, is_indirect: bool) {
if is_indirect {
self.flags |= VIRTQ_DESC_F_INDIRECT;
} else {
self.flags &= !VIRTQ_DESC_F_INDIRECT;
}
}
}
#[derive(Debug, Clone)]
pub struct DescriptorChain {
head: u16,
descriptors: Vec<VirtQueueDesc>,
}
impl DescriptorChain {
pub fn new(head: u16, descriptors: Vec<VirtQueueDesc>) -> Self {
Self { head, descriptors }
}
pub fn head(&self) -> u16 {
self.head
}
pub fn descriptors(&self) -> &[VirtQueueDesc] {
&self.descriptors
}
pub fn len(&self) -> usize {
self.descriptors.len()
}
pub fn is_empty(&self) -> bool {
self.descriptors.is_empty()
}
pub fn readable(&self) -> impl Iterator<Item = &VirtQueueDesc> {
self.descriptors.iter().filter(|d| !d.is_write())
}
pub fn writable(&self) -> impl Iterator<Item = &VirtQueueDesc> {
self.descriptors.iter().filter(|d| d.is_write())
}
pub fn readable_len(&self) -> VirtioResult<usize> {
sum_descriptor_lens(self.readable().map(|d| d.len as usize))
}
pub fn writable_len(&self) -> VirtioResult<usize> {
sum_descriptor_lens(self.writable().map(|d| d.len as usize))
}
}
fn sum_descriptor_lens(lens: impl Iterator<Item = usize>) -> VirtioResult<usize> {
let mut total = 0usize;
for v in lens {
total = total.checked_add(v).ok_or(VirtioError::InvalidDescriptor)?;
}
Ok(total)
}
#[derive(Debug, Clone)]
pub struct DescriptorTable {
pub base_addr: GuestPhysAddr,
pub size: u16,
}
impl DescriptorTable {
pub const fn new(base_addr: GuestPhysAddr, size: u16) -> Self {
Self { base_addr, size }
}
pub fn desc_addr(&self, index: u16) -> Option<GuestPhysAddr> {
if index >= self.size {
return None;
}
let offset = index as usize * core::mem::size_of::<VirtQueueDesc>();
Some(self.base_addr + offset)
}
pub(crate) const fn layout_size(size: u16) -> usize {
size as usize * core::mem::size_of::<VirtQueueDesc>()
}
pub fn total_size(&self) -> usize {
Self::layout_size(self.size)
}
pub fn is_valid(&self) -> bool {
self.base_addr.as_usize() != 0 && self.size > 0
}
pub fn read_desc(
&self,
index: u16,
memory: &mut dyn GuestMemory,
) -> VirtioResult<VirtQueueDesc> {
if !self.is_valid() {
return Err(VirtioError::QueueNotReady);
}
let desc_addr = self.desc_addr(index).ok_or(VirtioError::InvalidQueue)?;
let mut bytes = [0u8; 16];
memory.read(desc_addr, &mut bytes)?;
Ok(VirtQueueDesc {
base_addr: GuestPhysAddr::from(
u64::from_le_bytes(bytes[0..8].try_into().unwrap()) as usize
),
len: u32::from_le_bytes(bytes[8..12].try_into().unwrap()),
flags: u16::from_le_bytes(bytes[12..14].try_into().unwrap()),
next: u16::from_le_bytes(bytes[14..16].try_into().unwrap()),
})
}
pub fn write_desc(
&self,
index: u16,
desc: &VirtQueueDesc,
memory: &mut dyn GuestMemory,
) -> VirtioResult<()> {
if !self.is_valid() {
return Err(VirtioError::QueueNotReady);
}
let desc_addr = self.desc_addr(index).ok_or(VirtioError::InvalidQueue)?;
let mut bytes = [0u8; 16];
bytes[0..8].copy_from_slice(&(desc.base_addr.as_usize() as u64).to_le_bytes());
bytes[8..12].copy_from_slice(&desc.len.to_le_bytes());
bytes[12..14].copy_from_slice(&desc.flags.to_le_bytes());
bytes[14..16].copy_from_slice(&desc.next.to_le_bytes());
memory.write(desc_addr, &bytes)?;
Ok(())
}
pub fn follow_chain(
&self,
head_index: u16,
memory: &mut dyn GuestMemory,
) -> VirtioResult<Vec<VirtQueueDesc>> {
if !self.is_valid() {
return Err(VirtioError::QueueNotReady);
}
let mut descriptors = Vec::new();
let mut current_index = head_index;
loop {
if current_index >= self.size {
return Err(VirtioError::InvalidQueue);
}
let desc = self.read_desc(current_index, memory)?;
descriptors.push(desc);
if !desc.has_next() {
break;
}
current_index = desc.next;
if descriptors.len() > self.size as usize {
return Err(VirtioError::InvalidQueue);
}
}
Ok(descriptors)
}
pub fn descriptor_chain(
&self,
head: u16,
memory: &mut dyn GuestMemory,
) -> VirtioResult<DescriptorChain> {
if !self.is_valid() {
return Err(VirtioError::QueueNotReady);
}
if head >= self.size {
return Err(VirtioError::InvalidDescriptor);
}
let mut descriptors = Vec::new();
let mut current = head;
loop {
if current >= self.size {
return Err(VirtioError::InvalidDescriptor);
}
let desc = self.read_desc(current, memory)?;
if desc.is_indirect() {
return Err(VirtioError::NotSupported);
}
if desc
.base_addr
.as_usize()
.checked_add(desc.len as usize)
.is_none()
{
return Err(VirtioError::InvalidDescriptor);
}
descriptors.push(desc);
if !desc.has_next() {
break;
}
current = desc.next;
if descriptors.len() > self.size as usize {
return Err(VirtioError::InvalidDescriptor);
}
}
Ok(DescriptorChain::new(head, descriptors))
}
pub fn chain_length(&self, head_index: u16, memory: &mut dyn GuestMemory) -> VirtioResult<u32> {
let descriptors = self.follow_chain(head_index, memory)?;
Ok(descriptors.iter().map(|desc| desc.len).sum())
}
pub fn validate_chain(
&self,
head_index: u16,
memory: &mut dyn GuestMemory,
) -> VirtioResult<bool> {
let descriptors = self.follow_chain(head_index, memory)?;
if descriptors.is_empty() {
return Ok(false);
}
for (i, desc) in descriptors.iter().enumerate() {
if i == descriptors.len() - 1 && desc.has_next() {
return Ok(false);
}
if i < descriptors.len() - 1 && !desc.has_next() {
return Ok(false);
}
}
Ok(true)
}
pub fn get_data_buffers(
&self,
head_index: u16,
device_type: VirtioDeviceID,
memory: &mut dyn GuestMemory,
) -> VirtioResult<Vec<(GuestPhysAddr, usize, bool)>> {
let descriptors = self.follow_chain(head_index, memory)?;
if descriptors.len() < 2 && device_type == VirtioDeviceID::Block {
return Ok(Vec::new());
}
let mut buffers = Vec::new();
if device_type == VirtioDeviceID::Block {
for desc in &descriptors[1..descriptors.len() - 1] {
buffers.push((desc.base_addr, desc.len as usize, desc.is_write()));
}
} else {
for desc in &descriptors {
buffers.push((desc.base_addr, desc.len as usize, desc.is_write()));
}
}
Ok(buffers)
}
pub fn get_status_addr(
&self,
head_index: u16,
memory: &mut dyn GuestMemory,
) -> VirtioResult<GuestPhysAddr> {
let descriptors = self.follow_chain(head_index, memory)?;
if descriptors.is_empty() {
return Err(VirtioError::InvalidQueue);
}
let status_desc = &descriptors[descriptors.len() - 1];
if !status_desc.is_write() || status_desc.len < 1 {
return Err(VirtioError::InvalidQueue);
}
Ok(status_desc.base_addr)
}
}
#[cfg(test)]
mod tests {
use alloc::vec;
use ax_memory_addr::PhysAddr;
use axaddrspace::GuestMemoryAccessor;
use super::*;
#[derive(Clone)]
struct TestTranslator {
base_host_ptr: usize,
}
impl GuestMemoryAccessor for TestTranslator {
fn translate_and_get_limit(&self, guest_addr: GuestPhysAddr) -> Option<(PhysAddr, usize)> {
let offset = guest_addr.as_usize();
Some((PhysAddr::from(self.base_host_ptr + offset), usize::MAX))
}
}
#[test]
fn status_descriptor_len_must_be_at_least_one() {
let mut mem = vec![0u8; 4096];
let base_ptr = mem.as_mut_ptr() as usize;
let translator = TestTranslator {
base_host_ptr: base_ptr,
};
let mut memory = crate::AddressSpaceMemory::new(&translator);
let base = GuestPhysAddr::from(0x10usize);
let table = DescriptorTable::new(base, 2);
let mut d0 = VirtQueueDesc::new(GuestPhysAddr::from(0x100usize), 16, 0, 1);
d0.set_next(true);
let mut d1 = VirtQueueDesc::new(GuestPhysAddr::from(0x200usize), 0, 0, 0);
d1.set_write(true); d1.set_next(false);
table.write_desc(0, &d0, &mut memory).unwrap();
table.write_desc(1, &d1, &mut memory).unwrap();
let err = table.get_status_addr(0, &mut memory).unwrap_err();
assert!(matches!(err, VirtioError::InvalidQueue));
let mut d1_ok = d1;
d1_ok.len = 1;
table.write_desc(1, &d1_ok, &mut memory).unwrap();
let ok_addr = table.get_status_addr(0, &mut memory).unwrap();
assert_eq!(ok_addr.as_usize(), 0x200);
}
#[test]
fn layout_size_counts_descriptors() {
assert_eq!(DescriptorTable::layout_size(4), 64);
assert_eq!(DescriptorTable::layout_size(256), 4096);
}
}