use crate::error::{VirtioError, VirtioResult};
use crate::{VirtioDeviceID, constants::*};
use alloc::sync::Arc;
use alloc::vec::Vec;
use axaddrspace::{GuestMemoryAccessor, GuestPhysAddr};
#[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 DescriptorTable<T: GuestMemoryAccessor + Clone> {
pub base_addr: GuestPhysAddr,
pub size: u16,
accessor: Arc<T>,
}
impl<T: GuestMemoryAccessor + Clone> DescriptorTable<T> {
pub fn new(base_addr: GuestPhysAddr, size: u16, accessor: Arc<T>) -> Self {
Self {
base_addr,
size,
accessor,
}
}
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 fn total_size(&self) -> usize {
self.size as usize * core::mem::size_of::<VirtQueueDesc>()
}
pub fn is_valid(&self) -> bool {
self.base_addr.as_usize() != 0 && self.size > 0
}
pub fn read_desc(&self, index: u16) -> VirtioResult<VirtQueueDesc> {
if !self.is_valid() {
return Err(VirtioError::QueueNotReady);
}
let desc_addr = self.desc_addr(index).ok_or(VirtioError::InvalidQueue)?;
self.accessor
.read_obj(desc_addr)
.map_err(|_| VirtioError::InvalidAddress)
}
pub fn write_desc(&self, index: u16, desc: &VirtQueueDesc) -> VirtioResult<()> {
if !self.is_valid() {
return Err(VirtioError::QueueNotReady);
}
let desc_addr = self.desc_addr(index).ok_or(VirtioError::InvalidQueue)?;
self.accessor
.write_obj(desc_addr, *desc)
.map_err(|_| VirtioError::InvalidAddress)?;
Ok(())
}
pub fn follow_chain(&self, head_index: u16) -> 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)?;
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 chain_length(&self, head_index: u16) -> VirtioResult<u32> {
let descriptors = self.follow_chain(head_index)?;
Ok(descriptors.iter().map(|desc| desc.len).sum())
}
pub fn validate_chain(&self, head_index: u16) -> VirtioResult<bool> {
let descriptors = self.follow_chain(head_index)?;
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,
) -> VirtioResult<Vec<(GuestPhysAddr, usize, bool)>> {
let descriptors = self.follow_chain(head_index)?;
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) -> VirtioResult<GuestPhysAddr> {
let descriptors = self.follow_chain(head_index)?;
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 super::*;
use alloc::vec;
use memory_addr::PhysAddr;
#[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 accessor = Arc::new(translator);
let base = GuestPhysAddr::from(0x10usize);
let table: DescriptorTable<_> = DescriptorTable::new(base, 2, accessor.clone());
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).unwrap();
table.write_desc(1, &d1).unwrap();
let err = table.get_status_addr(0).unwrap_err();
assert!(matches!(err, VirtioError::InvalidQueue));
let mut d1_ok = d1;
d1_ok.len = 1;
table.write_desc(1, &d1_ok).unwrap();
let ok_addr = table.get_status_addr(0).unwrap();
assert_eq!(ok_addr.as_usize(), 0x200);
}
}