From a6d3901f3e1d4077703e04f6d32f771da0ae60de Mon Sep 17 00:00:00 2001 From: Dylan Reid Date: Tue, 7 Apr 2026 10:43:57 -0700 Subject: [PATCH] misc: return errors from IOMMU address translation instead of panicking MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The address that is passed from the guest should be treated as untrusted. Currently an invalid address will panic the VMM. This only allows the guest to hurt itself, but we shouldn't have the VMM crashing. Instead let's return an error if possible or invalidate the queue if it happen during setup. The data flow from guest to translate_gva/translate_gpa is: 1. Guest writes a raw u64 address into a virtio descriptor in the shared descriptor table (guest memory). 2. The virtio-queue crate reads this descriptor via read_obj() and returns the addr field as-is in a GuestAddress — no validation. 3. Device code calls .translate_gva(access_platform, len) on the GuestAddress. 4. With IOMMU (access_platform is Some): the address is an IOVA that must be translated to a GPA via the IOMMU mapping table. If the guest provides an unmapped IOVA, translation returns Err. Previously, .unwrap() here panicked the VMM. 5. Without IOMMU (access_platform is None): translate_gva is a no-op (returns self). The raw address flows to GuestMemory::read_obj() which validates it — out-of-range addresses return Err(InvalidGuestAddress), so no host memory corruption is possible. Signed-off-by: Dylan Reid --- block/src/lib.rs | 9 ++-- net_util/src/ctrl_queue.rs | 9 ++-- net_util/src/queue_pair.rs | 15 ++++-- virtio-devices/src/console.rs | 14 +++-- virtio-devices/src/pmem.rs | 6 ++- virtio-devices/src/rng.rs | 5 +- .../src/transport/pci_common_config.rs | 33 +++++++++--- virtio-devices/src/vdpa.rs | 35 ++++++++----- virtio-devices/src/vsock/packet.rs | 31 +++++++---- vm-virtio/src/lib.rs | 52 ++++++++++++++----- 10 files changed, 153 insertions(+), 56 deletions(-) diff --git a/block/src/lib.rs b/block/src/lib.rs index 811ee974e..7a55cc449 100644 --- a/block/src/lib.rs +++ b/block/src/lib.rs @@ -312,7 +312,8 @@ impl Request { let hdr_desc_addr = hdr_desc .addr() - .translate_gva(access_platform, hdr_desc.len() as usize); + .translate_gva(access_platform, hdr_desc.len() as usize) + .map_err(|e| Error::GuestMemory(GuestMemoryError::IOError(e)))?; let mut req = Request { request_type: request_type(desc_chain.memory(), hdr_desc_addr)?, @@ -353,7 +354,8 @@ impl Request { req.data_descriptors.push(( desc.addr() - .translate_gva(access_platform, desc.len() as usize), + .translate_gva(access_platform, desc.len() as usize) + .map_err(|e| Error::GuestMemory(GuestMemoryError::IOError(e)))?, desc.len(), )); desc = desc_chain @@ -384,7 +386,8 @@ impl Request { req.status_addr = status_desc .addr() - .translate_gva(access_platform, status_desc.len() as usize); + .translate_gva(access_platform, status_desc.len() as usize) + .map_err(|e| Error::GuestMemory(GuestMemoryError::IOError(e)))?; Ok(req) } diff --git a/net_util/src/ctrl_queue.rs b/net_util/src/ctrl_queue.rs index e42b4c0ca..b14b38036 100644 --- a/net_util/src/ctrl_queue.rs +++ b/net_util/src/ctrl_queue.rs @@ -100,14 +100,16 @@ impl CtrlQueue { .read_obj( ctrl_desc .addr() - .translate_gva(access_platform, ctrl_desc.len() as usize), + .translate_gva(access_platform, ctrl_desc.len() as usize) + .map_err(|e| Error::GuestMemory(GuestMemoryError::IOError(e)))?, ) .map_err(Error::GuestMemory)?; let data_desc = desc_chain.next().ok_or(Error::NoDataDescriptor)?; let data_desc_addr = data_desc .addr() - .translate_gva(access_platform, data_desc.len() as usize); + .translate_gva(access_platform, data_desc.len() as usize) + .map_err(|e| Error::GuestMemory(GuestMemoryError::IOError(e)))?; let status_desc = desc_chain.next().ok_or(Error::NoStatusDescriptor)?; @@ -168,7 +170,8 @@ impl CtrlQueue { if ok { VIRTIO_NET_OK } else { VIRTIO_NET_ERR } as u8, status_desc .addr() - .translate_gva(access_platform, status_desc.len() as usize), + .translate_gva(access_platform, status_desc.len() as usize) + .map_err(|e| Error::GuestMemory(GuestMemoryError::IOError(e)))?, ) .map_err(Error::GuestMemory)?; // Per virtio spec 2.6.8, used_len is the number of bytes written diff --git a/net_util/src/queue_pair.rs b/net_util/src/queue_pair.rs index c0b8825e7..a56903181 100644 --- a/net_util/src/queue_pair.rs +++ b/net_util/src/queue_pair.rs @@ -69,7 +69,10 @@ impl TxVirtio { while let Some(desc) = next_desc { let desc_addr = desc .addr() - .translate_gva(access_platform, desc.len() as usize); + .translate_gva(access_platform, desc.len() as usize) + .map_err(|e| { + NetQueuePairError::GuestMemory(vm_memory::GuestMemoryError::IOError(e)) + })?; if !desc.is_write_only() && desc.len() > 0 { let buf = desc_chain .memory() @@ -207,7 +210,10 @@ impl RxVirtio { .memory() .checked_offset( desc.addr() - .translate_gva(access_platform, desc.len() as usize), + .translate_gva(access_platform, desc.len() as usize) + .map_err(|e| { + NetQueuePairError::GuestMemory(vm_memory::GuestMemoryError::IOError(e)) + })?, 10, ) .ok_or(NetQueuePairError::DescriptorInvalidHeader)?; @@ -217,7 +223,10 @@ impl RxVirtio { while let Some(desc) = next_desc { let desc_addr = desc .addr() - .translate_gva(access_platform, desc.len() as usize); + .translate_gva(access_platform, desc.len() as usize) + .map_err(|e| { + NetQueuePairError::GuestMemory(vm_memory::GuestMemoryError::IOError(e)) + })?; if desc.is_write_only() && desc.len() > 0 { let buf = desc_chain .memory() diff --git a/virtio-devices/src/console.rs b/virtio-devices/src/console.rs index 96282b522..d2d57b9e5 100644 --- a/virtio-devices/src/console.rs +++ b/virtio-devices/src/console.rs @@ -221,7 +221,10 @@ impl ConsoleEpollHandler { .write_slice( &source_slice[..], desc.addr() - .translate_gva(self.access_platform.as_deref(), desc.len() as usize), + .translate_gva(self.access_platform.as_deref(), desc.len() as usize) + .map_err(|e| { + Error::GuestMemoryWrite(vm_memory::GuestMemoryError::IOError(e)) + })?, ) .map_err(Error::GuestMemoryWrite)?; @@ -259,10 +262,11 @@ impl ConsoleEpollHandler { desc_chain .memory() .write_volatile_to( - desc.addr().translate_gva( - self.access_platform.as_deref(), - desc.len() as usize, - ), + desc.addr() + .translate_gva(self.access_platform.as_deref(), desc.len() as usize) + .map_err(|e| { + Error::GuestMemoryRead(vm_memory::GuestMemoryError::IOError(e)) + })?, &mut buf, desc.len() as usize, ) diff --git a/virtio-devices/src/pmem.rs b/virtio-devices/src/pmem.rs index fd995747c..3abec1c0f 100644 --- a/virtio-devices/src/pmem.rs +++ b/virtio-devices/src/pmem.rs @@ -123,7 +123,8 @@ impl Request { .memory() .read_obj( desc.addr() - .translate_gva(access_platform, desc.len() as usize), + .translate_gva(access_platform, desc.len() as usize) + .map_err(|e| Error::GuestMemory(GuestMemoryError::IOError(e)))?, ) .map_err(Error::GuestMemory)?; @@ -147,7 +148,8 @@ impl Request { type_: request_type, status_addr: status_desc .addr() - .translate_gva(access_platform, status_desc.len() as usize), + .translate_gva(access_platform, status_desc.len() as usize) + .map_err(|e| Error::GuestMemory(GuestMemoryError::IOError(e)))?, }) } } diff --git a/virtio-devices/src/rng.rs b/virtio-devices/src/rng.rs index 2133cb79e..6bb0269c5 100644 --- a/virtio-devices/src/rng.rs +++ b/virtio-devices/src/rng.rs @@ -78,7 +78,10 @@ impl RngEpollHandler { .memory() .read_volatile_from( desc.addr() - .translate_gva(self.access_platform.as_deref(), desc.len() as usize), + .translate_gva(self.access_platform.as_deref(), desc.len() as usize) + .map_err(|e| { + Error::GuestMemoryWrite(vm_memory::GuestMemoryError::IOError(e)) + })?, &mut self.random_file, desc.len() as usize, ) diff --git a/virtio-devices/src/transport/pci_common_config.rs b/virtio-devices/src/transport/pci_common_config.rs index c59c454b7..11e1d3ac7 100644 --- a/virtio-devices/src/transport/pci_common_config.rs +++ b/virtio-devices/src/transport/pci_common_config.rs @@ -287,15 +287,36 @@ impl VirtioPciCommonConfig { q.set_ready(ready); // Translate address of descriptor table and vrings. if ready && let Some(access_platform) = &self.access_platform { - let desc_table = access_platform + let desc_table = match access_platform .translate_gva(q.desc_table(), get_vring_size(VringType::Desc, q.size())) - .unwrap(); - let avail_ring = access_platform + { + Ok(addr) => addr, + Err(e) => { + error!("Failed to translate desc_table GVA: {e}"); + q.set_ready(false); + return; + } + }; + let avail_ring = match access_platform .translate_gva(q.avail_ring(), get_vring_size(VringType::Avail, q.size())) - .unwrap(); - let used_ring = access_platform + { + Ok(addr) => addr, + Err(e) => { + error!("Failed to translate avail_ring GVA: {e}"); + q.set_ready(false); + return; + } + }; + let used_ring = match access_platform .translate_gva(q.used_ring(), get_vring_size(VringType::Used, q.size())) - .unwrap(); + { + Ok(addr) => addr, + Err(e) => { + error!("Failed to translate used_ring GVA: {e}"); + q.set_ready(false); + return; + } + }; q.set_desc_table_address( Some((desc_table & 0xffff_ffff) as u32), Some((desc_table >> 32) as u32), diff --git a/virtio-devices/src/vdpa.rs b/virtio-devices/src/vdpa.rs index 7cd341518..a35c35eb5 100644 --- a/virtio-devices/src/vdpa.rs +++ b/virtio-devices/src/vdpa.rs @@ -88,6 +88,8 @@ pub enum Error { SetVringKick(#[source] vhost::Error), #[error("Failed to set vring size")] SetVringNum(#[source] vhost::Error), + #[error("Failed to translate address")] + TranslateAddress(#[source] std::io::Error), } pub type Result = std::result::Result; @@ -246,18 +248,27 @@ impl Vdpa { queue_max_size, queue_size, flags: 0u32, - desc_table_addr: queue.desc_table().translate_gpa( - self.common.access_platform.as_deref(), - queue_size as usize * std::mem::size_of::(), - ), - used_ring_addr: queue.used_ring().translate_gpa( - self.common.access_platform.as_deref(), - 4 + queue_size as usize * 8, - ), - avail_ring_addr: queue.avail_ring().translate_gpa( - self.common.access_platform.as_deref(), - 4 + queue_size as usize * 2, - ), + desc_table_addr: queue + .desc_table() + .translate_gpa( + self.common.access_platform.as_deref(), + queue_size as usize * std::mem::size_of::(), + ) + .map_err(Error::TranslateAddress)?, + used_ring_addr: queue + .used_ring() + .translate_gpa( + self.common.access_platform.as_deref(), + 4 + queue_size as usize * 8, + ) + .map_err(Error::TranslateAddress)?, + avail_ring_addr: queue + .avail_ring() + .translate_gpa( + self.common.access_platform.as_deref(), + 4 + queue_size as usize * 2, + ) + .map_err(Error::TranslateAddress)?, log_addr: None, }; diff --git a/virtio-devices/src/vsock/packet.rs b/virtio-devices/src/vsock/packet.rs index 57218a5b8..e6b4c5afb 100644 --- a/virtio-devices/src/vsock/packet.rs +++ b/virtio-devices/src/vsock/packet.rs @@ -142,7 +142,8 @@ impl VsockPacket { let guest_hdr_addr = head .addr() - .translate_gva(access_platform, VSOCK_PKT_HDR_SIZE); + .translate_gva(access_platform, VSOCK_PKT_HDR_SIZE) + .map_err(|_| VsockError::GuestMemory)?; // To avoid TOCTOU issues when reading/writing the VSock packet header in guest memory, // we need to copy the content of the header in the VMM's memory. @@ -178,8 +179,9 @@ impl VsockPacket { desc_chain.memory(), head.addr() .checked_add(VSOCK_PKT_HDR_SIZE as u64) - .unwrap() - .translate_gva(access_platform, buf_size), + .ok_or(VsockError::GuestMemory)? + .translate_gva(access_platform, buf_size) + .map_err(|_| VsockError::GuestMemory)?, buf_size, ) .ok_or(VsockError::GuestMemory)?; @@ -214,7 +216,10 @@ impl VsockPacket { let desc_len = desc.len() as usize; if desc_len > 0 && offset < total_len { let to_copy = std::cmp::min(desc_len, total_len - offset); - let desc_addr = desc.addr().translate_gva(access_platform, desc_len); + let desc_addr = desc + .addr() + .translate_gva(access_platform, desc_len) + .map_err(|_| VsockError::GuestMemory)?; desc_chain .memory() .read_slice(&mut owned[offset..offset + to_copy], desc_addr) @@ -242,7 +247,10 @@ impl VsockPacket { let buf_size = buf_desc.len() as usize; let buf_ptr = get_host_address_range( desc_chain.memory(), - buf_desc.addr().translate_gva(access_platform, buf_size), + buf_desc + .addr() + .translate_gva(access_platform, buf_size) + .map_err(|_| VsockError::GuestMemory)?, buf_size, ) .ok_or(VsockError::GuestMemory)?; @@ -283,7 +291,8 @@ impl VsockPacket { let guest_hdr_addr = head .addr() - .translate_gva(access_platform, VSOCK_PKT_HDR_SIZE); + .translate_gva(access_platform, VSOCK_PKT_HDR_SIZE) + .map_err(|_| VsockError::GuestMemory)?; // To avoid TOCTOU issues when reading/writing the VSock packet header in guest memory, // we need to copy the content of the header in the VMM's memory. @@ -313,7 +322,10 @@ impl VsockPacket { buf: Some(PacketBuffer::Borrowed { ptr: get_host_address_range( desc_chain.memory(), - buf_desc.addr().translate_gva(access_platform, buf_size), + buf_desc + .addr() + .translate_gva(access_platform, buf_size) + .map_err(|_| VsockError::GuestMemory)?, buf_size, ) .ok_or(VsockError::GuestMemory)?, @@ -330,8 +342,9 @@ impl VsockPacket { desc_chain.memory(), head.addr() .checked_add(VSOCK_PKT_HDR_SIZE as u64) - .unwrap() - .translate_gva(access_platform, buf_size), + .ok_or(VsockError::GuestMemory)? + .translate_gva(access_platform, buf_size) + .map_err(|_| VsockError::GuestMemory)?, buf_size, ) .ok_or(VsockError::GuestMemory)?, diff --git a/vm-virtio/src/lib.rs b/vm-virtio/src/lib.rs index c560e5c86..fbd94b2b7 100644 --- a/vm-virtio/src/lib.rs +++ b/vm-virtio/src/lib.rs @@ -101,32 +101,60 @@ pub trait AccessPlatform: Send + Sync + Debug { } pub trait Translatable { - fn translate_gva(&self, access_platform: Option<&dyn AccessPlatform>, len: usize) -> Self; - fn translate_gpa(&self, access_platform: Option<&dyn AccessPlatform>, len: usize) -> Self; + fn translate_gva( + &self, + access_platform: Option<&dyn AccessPlatform>, + len: usize, + ) -> std::result::Result + where + Self: Sized; + fn translate_gpa( + &self, + access_platform: Option<&dyn AccessPlatform>, + len: usize, + ) -> std::result::Result + where + Self: Sized; } impl Translatable for GuestAddress { - fn translate_gva(&self, access_platform: Option<&dyn AccessPlatform>, len: usize) -> Self { - GuestAddress(self.0.translate_gva(access_platform, len)) + fn translate_gva( + &self, + access_platform: Option<&dyn AccessPlatform>, + len: usize, + ) -> std::result::Result { + Ok(GuestAddress(self.0.translate_gva(access_platform, len)?)) } - fn translate_gpa(&self, access_platform: Option<&dyn AccessPlatform>, len: usize) -> Self { - GuestAddress(self.0.translate_gpa(access_platform, len)) + fn translate_gpa( + &self, + access_platform: Option<&dyn AccessPlatform>, + len: usize, + ) -> std::result::Result { + Ok(GuestAddress(self.0.translate_gpa(access_platform, len)?)) } } impl Translatable for u64 { - fn translate_gva(&self, access_platform: Option<&dyn AccessPlatform>, len: usize) -> Self { + fn translate_gva( + &self, + access_platform: Option<&dyn AccessPlatform>, + len: usize, + ) -> std::result::Result { if let Some(access_platform) = access_platform { - access_platform.translate_gva(*self, len as u64).unwrap() + access_platform.translate_gva(*self, len as u64) } else { - *self + Ok(*self) } } - fn translate_gpa(&self, access_platform: Option<&dyn AccessPlatform>, len: usize) -> Self { + fn translate_gpa( + &self, + access_platform: Option<&dyn AccessPlatform>, + len: usize, + ) -> std::result::Result { if let Some(access_platform) = access_platform { - access_platform.translate_gpa(*self, len as u64).unwrap() + access_platform.translate_gpa(*self, len as u64) } else { - *self + Ok(*self) } } }