use std::{
fs::File,
sync::{
atomic::{AtomicU16, AtomicU32, AtomicU64, AtomicU8, Ordering},
Arc,
},
};
use crate::device::bus::{BusDevice, Request, RequestSize};
use memmap2::{Mmap, MmapMut, MmapOptions};
use tracing::warn;
use vfio_user::DmaMapFlags;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum AccessRights {
ReadOnly,
ReadWrite,
}
#[derive(thiserror::Error, Debug)]
pub enum DmaMapFlagsError {
#[error("Invalid DMA map flags: {value:?}")]
InvalidFlags { value: DmaMapFlags },
}
impl TryFrom<DmaMapFlags> for AccessRights {
type Error = DmaMapFlagsError;
fn try_from(value: DmaMapFlags) -> Result<Self, Self::Error> {
let readable = value.contains(DmaMapFlags::READ);
let writable = value.contains(DmaMapFlags::WRITE);
if (value & !DmaMapFlags::READ_WRITE).bits() != 0 {
warn!("Unknown DmaMapFlags set: {:0x}", value);
}
if readable && !writable {
Ok(Self::ReadOnly)
} else if readable && writable {
Ok(Self::ReadWrite)
} else {
Err(DmaMapFlagsError::InvalidFlags { value })
}
}
}
#[derive(Debug)]
enum Mapping {
ReadOnly(Mmap),
ReadWrite(MmapMut),
}
impl Mapping {
#[allow(clippy::missing_const_for_fn)] fn as_ptr(&self) -> *const u8 {
match self {
Self::ReadOnly(map) => map.as_ptr(),
Self::ReadWrite(map) => map.as_ptr(),
}
}
const fn is_writable(&self) -> bool {
match self {
Self::ReadWrite(_) => true,
Self::ReadOnly(_) => false,
}
}
}
#[derive(Debug)]
pub struct MemorySegment {
size: u64,
mapping: Arc<Mapping>,
}
impl MemorySegment {
pub fn new_from_fd(
fd: &File,
file_offset: u64,
size: u64,
access_rights: AccessRights,
) -> Result<Self, std::io::Error> {
Ok(Self {
size,
mapping: Arc::new({
let mut mmap = MmapOptions::new();
mmap.len(size.try_into().unwrap());
mmap.offset(file_offset);
match access_rights {
AccessRights::ReadOnly => unsafe { Mapping::ReadOnly(mmap.map(fd)?) },
AccessRights::ReadWrite => unsafe { Mapping::ReadWrite(mmap.map_mut(fd)?) },
}
}),
})
}
}
impl BusDevice for MemorySegment {
fn size(&self) -> u64 {
self.size
}
fn read(&self, req: Request) -> u64 {
assert!(
req.addr
.checked_add(req.size.into())
.is_some_and(|end| end <= self.size),
"address overflow or out of bounds"
);
let ptr = unsafe { self.mapping.as_ptr().add(req.addr.try_into().unwrap()) };
match req.size {
RequestSize::Size1 => {
let atomic = unsafe { &*(ptr as *const AtomicU8) };
atomic.load(Ordering::Relaxed).into()
}
RequestSize::Size2 => {
let atomic = unsafe { &*(ptr as *const AtomicU16) };
atomic.load(Ordering::Relaxed).into()
}
RequestSize::Size4 => {
let atomic = unsafe { &*(ptr as *const AtomicU32) };
atomic.load(Ordering::Relaxed).into()
}
RequestSize::Size8 => {
let atomic = unsafe { &*(ptr as *const AtomicU64) };
atomic.load(Ordering::Relaxed)
}
}
}
fn write(&self, req: Request, value: u64) {
assert!(
req.addr
.checked_add(req.size.into())
.is_some_and(|end| end <= self.size),
"address overflow or out of bounds"
);
if !self.mapping.is_writable() {
return;
}
let ptr = unsafe { self.mapping.as_ptr().add(req.addr.try_into().unwrap()) };
match req.size {
RequestSize::Size1 => {
let atomic = unsafe { &*(ptr as *const AtomicU8) };
atomic.store(value as u8, Ordering::Relaxed);
}
RequestSize::Size2 => {
let atomic = unsafe { &*(ptr as *const AtomicU16) };
atomic.store(value as u16, Ordering::Relaxed);
}
RequestSize::Size4 => {
let atomic = unsafe { &*(ptr as *const AtomicU32) };
atomic.store(value as u32, Ordering::Relaxed);
}
RequestSize::Size8 => {
let atomic = unsafe { &*(ptr as *const AtomicU64) };
atomic.store(value, Ordering::Relaxed);
}
}
}
}
#[cfg(test)]
mod tests {
use std::{
ffi::CString,
io::{Read, Seek},
os::fd::FromRawFd,
};
use super::*;
fn create_memfd(size: u64) -> Result<File, std::io::Error> {
let fd = unsafe { libc::memfd_create(CString::new("unittest").unwrap().as_ptr(), 0) };
if fd < 0 {
return Err(std::io::Error::last_os_error());
}
let file = unsafe { File::from_raw_fd(fd) };
file.set_len(size)?;
Ok(file)
}
#[test]
fn can_read_write() -> Result<(), std::io::Error> {
let memfd = create_memfd(0x1000)?;
let mseg = MemorySegment::new_from_fd(&memfd, 0, 0x1000, AccessRights::ReadWrite)?;
assert_eq!(mseg.read(Request::new(0, RequestSize::Size8)), 0);
mseg.write(Request::new(0, RequestSize::Size8), 0xcafed00dfeedface);
assert_eq!(
mseg.read(Request::new(0, RequestSize::Size8)),
0xcafed00dfeedface
);
Ok(())
}
#[test]
fn cant_write_to_read_only() -> Result<(), std::io::Error> {
let memfd = create_memfd(0x1000)?;
let mseg = MemorySegment::new_from_fd(&memfd, 0, 0x1000, AccessRights::ReadOnly)?;
mseg.write(Request::new(0, RequestSize::Size8), 0xcafed00dfeedface);
assert_eq!(mseg.read(Request::new(0, RequestSize::Size8)), 0);
Ok(())
}
#[test]
fn file_offset_is_respected() -> Result<(), std::io::Error> {
let mut memfd = create_memfd(0x2000)?;
let mseg = MemorySegment::new_from_fd(&memfd, 0x1000, 0x1000, AccessRights::ReadWrite)?;
let data = 0xcafed00dfeedface_u64.to_le_bytes();
mseg.write(
Request::new(0x10, RequestSize::Size8),
u64::from_le_bytes(data),
);
let mut check_data = [0; 8];
memfd.seek(std::io::SeekFrom::Start(0x1010))?;
memfd.read_exact(&mut check_data)?;
assert_eq!(check_data, data);
Ok(())
}
}