use bitflags::bitflags;
use libc::{c_void, iovec, EINVAL};
use libc::{sysconf, _SC_PAGESIZE};
use std::ffi::CString;
use std::fs::File;
use std::io::{IoSlice, Read, Write};
use std::mem::size_of;
use std::num::Wrapping;
use std::os::fd::OwnedFd;
use std::os::unix::{
io::{FromRawFd, RawFd},
net::{UnixListener, UnixStream},
};
use std::path::{Path, PathBuf};
use thiserror::Error;
use vfio_bindings::bindings::vfio::*;
use vm_memory::{ByteValued, FileOffset};
use vmm_sys_util::sock_ctrl_msg::ScmSocket;
#[macro_use]
extern crate serde_derive;
#[macro_use]
extern crate log;
#[allow(dead_code)]
#[repr(u16)]
#[derive(Clone, Copy, Debug, Default, enumn::N)]
pub enum Command {
#[default]
Unknown = 0,
Version = 1,
DmaMap = 2,
DmaUnmap = 3,
DeviceGetInfo = 4,
DeviceGetRegionInfo = 5,
GetRegionIoFds = 6,
GetIrqInfo = 7,
SetIrqs = 8,
RegionRead = 9,
RegionWrite = 10,
DmaRead = 11,
DmaWrite = 12,
DeviceReset = 13,
UserDirtyPages = 14,
}
#[allow(dead_code)]
#[repr(u32)]
#[derive(Clone, Copy, Debug, PartialEq, Eq, Default)]
enum HeaderFlags {
#[default]
Command = 0,
Reply = 1,
NoReply = 1 << 4,
Error = 1 << 5,
}
#[repr(C)]
#[derive(Default, Clone, Copy, Debug)]
struct Header {
message_id: u16,
command: u16,
message_size: u32,
flags: u32,
error: u32,
}
impl Header {
fn no_reply(&self) -> bool {
self.flags & HeaderFlags::NoReply as u32 != 0
}
}
#[repr(C)]
#[derive(Default, Clone, Copy, Debug)]
struct Version {
header: Header,
major: u16,
minor: u16,
}
#[derive(Serialize, Deserialize, Debug)]
struct MigrationCapabilities {
pgsize: u32,
}
const fn default_max_msg_fds() -> u32 {
1
}
const fn default_max_data_xfer_size() -> u32 {
1048576
}
#[inline(always)]
fn pagesize() -> u32 {
unsafe { sysconf(_SC_PAGESIZE) as u32 }
}
fn default_migration_capabilities() -> MigrationCapabilities {
MigrationCapabilities { pgsize: pagesize() }
}
bitflags! {
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct DmaMapFlags: u32 {
const READ = 1 << 0;
const WRITE = 1 << 1;
const READ_WRITE = Self::READ.bits() | Self::WRITE.bits();
const _ = !0;
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct DmaUnmapFlags: u32 {
const GET_DIRTY_PAGE_INFO = 1 << 1;
const UNMAP_ALL = 1 << 2;
const _ = !0;
}
}
#[repr(C)]
#[derive(Default, Clone, Copy, Debug)]
struct DmaMap {
header: Header,
argsz: u32,
flags: u32,
offset: u64,
address: u64,
size: u64,
}
#[repr(C)]
#[derive(Default, Clone, Copy, Debug)]
struct DmaUnmap {
header: Header,
argsz: u32,
flags: u32,
address: u64,
size: u64,
}
#[repr(C)]
#[derive(Default, Clone, Copy, Debug)]
struct DeviceGetInfo {
header: Header,
argsz: u32,
flags: u32,
num_regions: u32,
num_irqs: u32,
}
#[repr(C)]
#[derive(Default, Clone, Copy, Debug)]
struct DeviceGetRegionInfo {
header: Header,
region_info: vfio_region_info,
}
#[repr(C)]
#[derive(Default, Clone, Copy, Debug)]
struct RegionAccess {
header: Header,
offset: u64,
region: u32,
count: u32,
}
#[repr(C)]
#[derive(Default, Clone, Copy, Debug)]
struct GetIrqInfo {
header: Header,
argsz: u32,
flags: u32,
index: u32,
count: u32,
}
#[repr(C)]
#[derive(Default, Clone, Copy, Debug)]
struct SetIrqs {
header: Header,
argsz: u32,
flags: u32,
index: u32,
start: u32,
count: u32,
}
#[repr(C)]
#[derive(Default, Clone, Copy, Debug)]
struct DeviceReset {
header: Header,
}
unsafe impl ByteValued for Header {}
unsafe impl ByteValued for Version {}
unsafe impl ByteValued for DmaMap {}
unsafe impl ByteValued for DmaUnmap {}
unsafe impl ByteValued for DeviceGetInfo {}
unsafe impl ByteValued for DeviceGetRegionInfo {}
unsafe impl ByteValued for RegionAccess {}
unsafe impl ByteValued for GetIrqInfo {}
unsafe impl ByteValued for SetIrqs {}
unsafe impl ByteValued for DeviceReset {}
#[derive(Serialize, Deserialize, Debug)]
struct Capabilities {
#[serde(default = "default_max_msg_fds")]
max_msg_fds: u32,
#[serde(default = "default_max_data_xfer_size")]
max_data_xfer_size: u32,
#[serde(default = "default_migration_capabilities")]
migration: MigrationCapabilities,
}
#[derive(Serialize, Deserialize, Debug, Default)]
struct CapabilitiesData {
capabilities: Capabilities,
}
impl Default for Capabilities {
fn default() -> Self {
Self {
max_msg_fds: default_max_msg_fds(),
max_data_xfer_size: default_max_data_xfer_size(),
migration: default_migration_capabilities(),
}
}
}
pub struct Client {
stream: UnixStream,
next_message_id: Wrapping<u16>,
num_irqs: u32,
resettable: bool,
regions: Vec<Region>,
}
#[derive(Debug)]
pub struct Region {
pub flags: u32,
pub index: u32,
pub size: u64,
pub file_offset: Option<FileOffset>,
pub sparse_areas: Vec<vfio_region_sparse_mmap_area>,
}
#[derive(Clone, Copy, Debug)]
pub struct IrqInfo {
pub index: u32,
pub flags: u32,
pub count: u32,
}
#[derive(Error, Debug)]
pub enum Error {
#[error("Error connecting: {0}")]
Connect(#[source] std::io::Error),
#[error("Error serializing capabilities: {0}")]
SerializeCapabilites(#[source] serde_json::Error),
#[error("Error deserializing capabilities: {0}")]
DeserializeCapabilites(#[source] serde_json::Error),
#[error("Error writing to stream: {0}")]
StreamWrite(#[source] std::io::Error),
#[error("Error reading from stream: {0}")]
StreamRead(#[source] std::io::Error),
#[error("Error shutting down stream: {0}")]
StreamShutdown(#[source] std::io::Error),
#[error("Error writing with file descriptors: {0}")]
SendWithFd(#[source] vmm_sys_util::errno::Error),
#[error("Error reading with file descriptors: {0}")]
ReceiveWithFd(#[source] vmm_sys_util::errno::Error),
#[error("Not a PCI device")]
NotPciDevice,
#[error("Socket path already exists")]
SocketPathExists,
#[error("Error binding to socket: {0}")]
SocketBind(#[source] std::io::Error),
#[error("Error accepting connection: {0}")]
SocketAccept(#[source] std::io::Error),
#[error("Unknown command: {0}")]
UnknownCommand(u16),
#[error("Unsupported command: {0:?}")]
UnsupportedCommand(Command),
#[error("Unsupported feature")]
UnsupportedFeature,
#[error("Error from backend: {0:?}")]
Backend(#[source] std::io::Error),
#[error("Invalid input")]
InvalidInput,
#[error("No_reply bit unexpectedly set for command: {0:?}")]
UnexpectedNoReply(Command),
}
impl Client {
pub fn new(path: &Path) -> Result<Client, Error> {
let stream = UnixStream::connect(path).map_err(Error::Connect)?;
let mut client = Client {
next_message_id: Wrapping(0),
stream,
num_irqs: 0,
resettable: false,
regions: Vec::new(),
};
client.negotiate_version()?;
client.regions = client.get_regions()?;
Ok(client)
}
fn negotiate_version(&mut self) -> Result<(), Error> {
let caps = CapabilitiesData::default();
let version_data = serde_json::to_string(&caps).map_err(Error::SerializeCapabilites)?;
let version = Version {
header: Header {
message_id: self.next_message_id.0,
command: Command::Version as u16,
flags: HeaderFlags::Command as u32,
message_size: (size_of::<Version>() + version_data.len() + 1) as u32,
..Default::default()
},
major: 0,
minor: 1,
};
debug!("Command: {version:?}");
let version_data = CString::new(version_data.as_bytes()).unwrap();
let bufs = vec![
IoSlice::new(version.as_slice()),
IoSlice::new(version_data.as_bytes_with_nul()),
];
let _ = self
.stream
.write_vectored(&bufs)
.map_err(Error::StreamWrite)?;
debug!(
"Sent client version information: major = {} minor = {} capabilities = {:?}",
version.major, version.minor, caps.capabilities
);
self.next_message_id += Wrapping(1);
let mut server_version: Version = Version::default();
self.stream
.read_exact(server_version.as_mut_slice())
.map_err(Error::StreamRead)?;
debug!("Reply: {server_version:?}");
let mut server_version_data =
vec![0; server_version.header.message_size as usize - size_of::<Version>()];
self.stream
.read_exact(server_version_data.as_mut_slice())
.map_err(Error::StreamRead)?;
let server_caps: CapabilitiesData =
serde_json::from_slice(&server_version_data[0..server_version_data.len() - 1])
.map_err(Error::DeserializeCapabilites)?;
debug!(
"Received server version information: major = {} minor = {} capabilities = {:?}",
server_version.major, server_version.minor, server_caps.capabilities
);
Ok(())
}
pub fn dma_map(
&mut self,
offset: u64,
address: u64,
size: u64,
fd: RawFd,
) -> Result<(), Error> {
let dma_map = DmaMap {
header: Header {
message_id: self.next_message_id.0,
command: Command::DmaMap as u16,
flags: HeaderFlags::Command as u32,
message_size: size_of::<DmaMap>() as u32,
..Default::default()
},
argsz: (size_of::<DmaMap>() - size_of::<Header>()) as u32,
flags: DmaMapFlags::READ_WRITE.bits(),
offset,
address,
size,
};
debug!("Command: {dma_map:?}");
self.next_message_id += Wrapping(1);
self.stream
.send_with_fd(dma_map.as_slice(), fd)
.map_err(Error::SendWithFd)?;
let mut reply = Header::default();
self.stream
.read_exact(reply.as_mut_slice())
.map_err(Error::StreamRead)?;
debug!("Reply: {reply:?}");
Ok(())
}
pub fn dma_unmap(&mut self, address: u64, size: u64) -> Result<(), Error> {
let dma_unmap = DmaUnmap {
header: Header {
message_id: self.next_message_id.0,
command: Command::DmaUnmap as u16,
flags: HeaderFlags::Command as u32,
message_size: size_of::<DmaUnmap>() as u32,
..Default::default()
},
argsz: (size_of::<DmaUnmap>() - size_of::<Header>()) as u32,
flags: 0,
address,
size,
};
debug!("Command: {dma_unmap:?}");
self.next_message_id += Wrapping(1);
self.stream
.write_all(dma_unmap.as_slice())
.map_err(Error::StreamWrite)?;
let mut reply = DmaUnmap::default();
self.stream
.read_exact(reply.as_mut_slice())
.map_err(Error::StreamRead)?;
debug!("Reply: {reply:?}");
Ok(())
}
pub fn reset(&mut self) -> Result<(), Error> {
let reset = DeviceReset {
header: Header {
message_id: self.next_message_id.0,
command: Command::DeviceReset as u16,
flags: HeaderFlags::Command as u32,
message_size: size_of::<DeviceReset>() as u32,
..Default::default()
},
};
debug!("Command: {reset:?}");
self.next_message_id += Wrapping(1);
self.stream
.write_all(reset.as_slice())
.map_err(Error::StreamWrite)?;
let mut reply = Header::default();
self.stream
.read_exact(reply.as_mut_slice())
.map_err(Error::StreamRead)?;
debug!("Reply: {reply:?}");
Ok(())
}
fn get_regions(&mut self) -> Result<Vec<Region>, Error> {
let get_info = DeviceGetInfo {
header: Header {
message_id: self.next_message_id.0,
command: Command::DeviceGetInfo as u16,
flags: HeaderFlags::Command as u32,
message_size: size_of::<DeviceGetInfo>() as u32,
..Default::default()
},
argsz: size_of::<DeviceGetInfo>() as u32,
..Default::default()
};
debug!("Command: {get_info:?}");
self.next_message_id += Wrapping(1);
self.stream
.write_all(get_info.as_slice())
.map_err(Error::StreamWrite)?;
let mut reply = DeviceGetInfo::default();
self.stream
.read_exact(reply.as_mut_slice())
.map_err(Error::StreamRead)?;
debug!("Reply: {reply:?}");
self.num_irqs = reply.num_irqs;
if reply.flags & VFIO_DEVICE_FLAGS_PCI != VFIO_DEVICE_FLAGS_PCI {
return Err(Error::NotPciDevice);
}
self.resettable = reply.flags & VFIO_DEVICE_FLAGS_RESET != VFIO_DEVICE_FLAGS_RESET;
let num_regions = reply.num_regions;
let mut regions = Vec::new();
for index in 0..num_regions {
let (region_info, fd, sparse_areas) = self.get_region_info(index)?;
regions.push(Region {
flags: region_info.flags,
index: region_info.index,
size: region_info.size,
file_offset: fd.map(|fd| FileOffset::new(fd, region_info.offset)),
sparse_areas,
});
}
Ok(regions)
}
fn get_region_info(
&mut self,
index: u32,
) -> Result<
(
vfio_region_info,
Option<File>,
Vec<vfio_region_sparse_mmap_area>,
),
Error,
> {
let mut get_region_info = DeviceGetRegionInfo {
header: Header {
message_id: self.next_message_id.0,
command: Command::DeviceGetRegionInfo as u16,
flags: HeaderFlags::Command as u32,
message_size: std::mem::size_of::<DeviceGetRegionInfo>() as u32,
..Default::default()
},
region_info: vfio_region_info {
argsz: size_of::<vfio_region_info>() as u32,
index,
..Default::default()
},
};
debug!("Command: {get_region_info:?}");
self.next_message_id += Wrapping(1);
self.stream
.write_all(get_region_info.as_slice())
.map_err(Error::StreamWrite)?;
let mut reply = DeviceGetRegionInfo::default();
let (_, fd) = self
.stream
.recv_with_fd(reply.as_mut_slice())
.map_err(Error::ReceiveWithFd)?;
debug!("Reply: {reply:?}");
if reply.region_info.argsz > std::mem::size_of::<vfio_region_info>() as u32 {
get_region_info.region_info.argsz = reply.region_info.argsz;
debug!("Command: {get_region_info:?}");
self.next_message_id += Wrapping(1);
self.stream
.write_all(get_region_info.as_slice())
.map_err(Error::StreamWrite)?;
let mut reply = DeviceGetRegionInfo::default();
let (_, fd) = self
.stream
.recv_with_fd(reply.as_mut_slice())
.map_err(Error::ReceiveWithFd)?;
debug!("Reply: {reply:?}");
let cap_size = reply.region_info.argsz - std::mem::size_of::<vfio_region_info>() as u32;
assert_eq!(
cap_size,
reply.header.message_size - size_of::<DeviceGetRegionInfo>() as u32
);
let mut cap_data = vec![0; cap_size as usize];
self.stream
.read_exact(cap_data.as_mut_slice())
.map_err(Error::StreamRead)?;
let sparse_areas = Self::parse_region_caps(&cap_data, &reply.region_info)?;
Ok((reply.region_info, fd, sparse_areas))
} else {
Ok((reply.region_info, fd, Vec::new()))
}
}
fn parse_region_caps(
cap_data: &[u8],
region_info: &vfio_region_info,
) -> Result<Vec<vfio_region_sparse_mmap_area>, Error> {
let mut sparse_areas: Vec<vfio_region_sparse_mmap_area> = Vec::new();
let cap_size = cap_data.len() as u32;
let cap_header_size = size_of::<vfio_info_cap_header>() as u32;
let mmap_cap_size = size_of::<vfio_region_info_cap_sparse_mmap>() as u32;
let mmap_area_size = size_of::<vfio_region_sparse_mmap_area>() as u32;
let cap_data_ptr = cap_data.as_ptr();
let mut region_info_offset = region_info.cap_offset;
while region_info_offset != 0 {
let cap_offset = region_info_offset - size_of::<vfio_region_info>() as u32;
if cap_offset + cap_header_size > cap_size {
warn!(
"Unexpected end of cap data: 'cap_offset + cap_header_size > cap_size' \
cap_offset = {cap_offset}, cap_header_size = {cap_header_size}, cap_size = {cap_size}"
);
break;
}
let cap_ptr = unsafe { cap_data_ptr.offset(cap_offset as isize) };
let cap_header = unsafe { &*(cap_ptr as *const vfio_info_cap_header) };
match cap_header.id as u32 {
VFIO_REGION_INFO_CAP_SPARSE_MMAP => {
if cap_offset + mmap_cap_size > cap_size {
warn!(
"Unexpected end of cap data: 'cap_offset + mmap_cap_size > cap_size' \
cap_offset = {cap_offset}, mmap_cap_size = {mmap_cap_size}, cap_size = {cap_size}"
);
break;
}
let sparse_mmap = unsafe {
&*(cap_ptr as *mut u8 as *const vfio_region_info_cap_sparse_mmap)
};
let area_num = sparse_mmap.nr_areas;
if cap_offset + mmap_cap_size + area_num * mmap_area_size > cap_size {
warn!("Unexpected end of cap data: 'cap_offset + mmap_cap_size + area_num * mmap_area_size > cap_size' \
cap_offset = {cap_offset}, mmap_cap_size = {mmap_area_size}, area_num = {area_num}, mmap_area_size = {mmap_area_size}, cap_size = {cap_size}");
break;
}
let areas =
unsafe { sparse_mmap.areas.as_slice(sparse_mmap.nr_areas as usize) };
for area in areas.iter() {
sparse_areas.push(*area);
}
}
_ => {
warn!(
"Ignoring unsupported vfio region capability (id = '{}')",
cap_header.id
);
}
}
region_info_offset = cap_header.next;
}
Ok(sparse_areas)
}
pub fn region_read(&mut self, region: u32, offset: u64, data: &mut [u8]) -> Result<(), Error> {
let region_read = RegionAccess {
header: Header {
message_id: self.next_message_id.0,
command: Command::RegionRead as u16,
flags: HeaderFlags::Command as u32,
message_size: size_of::<RegionAccess>() as u32,
..Default::default()
},
offset,
count: data.len() as u32,
region,
};
debug!("Command: {region_read:?}");
self.next_message_id += Wrapping(1);
self.stream
.write_all(region_read.as_slice())
.map_err(Error::StreamWrite)?;
let mut reply = RegionAccess::default();
self.stream
.read_exact(reply.as_mut_slice())
.map_err(Error::StreamRead)?;
debug!("Reply: {reply:?}");
self.stream.read_exact(data).map_err(Error::StreamRead)?;
Ok(())
}
pub fn region_write(&mut self, region: u32, offset: u64, data: &[u8]) -> Result<(), Error> {
let region_write = RegionAccess {
header: Header {
message_id: self.next_message_id.0,
command: Command::RegionWrite as u16,
flags: HeaderFlags::Command as u32,
message_size: (size_of::<RegionAccess>() + data.len()) as u32,
..Default::default()
},
offset,
count: data.len() as u32,
region,
};
debug!("Command: {region_write:?}");
self.next_message_id += Wrapping(1);
let bufs = vec![IoSlice::new(region_write.as_slice()), IoSlice::new(data)];
let _ = self
.stream
.write_vectored(&bufs)
.map_err(Error::StreamWrite)?;
let mut reply = RegionAccess::default();
self.stream
.read_exact(reply.as_mut_slice())
.map_err(Error::StreamRead)?;
debug!("Reply: {reply:?}");
Ok(())
}
pub fn get_irq_info(&mut self, index: u32) -> Result<IrqInfo, Error> {
let get_irq_info = GetIrqInfo {
header: Header {
message_id: self.next_message_id.0,
command: Command::GetIrqInfo as u16,
flags: HeaderFlags::Command as u32,
message_size: size_of::<GetIrqInfo>() as u32,
..Default::default()
},
argsz: (size_of::<GetIrqInfo>() - size_of::<Header>()) as u32,
flags: 0,
index,
count: 0,
};
debug!("Command: {get_irq_info:?}");
self.next_message_id += Wrapping(1);
self.stream
.write_all(get_irq_info.as_slice())
.map_err(Error::StreamWrite)?;
let mut reply = GetIrqInfo::default();
self.stream
.read_exact(reply.as_mut_slice())
.map_err(Error::StreamRead)?;
debug!("Reply: {reply:?}");
Ok(IrqInfo {
index: reply.index,
flags: reply.flags,
count: reply.count,
})
}
pub fn set_irqs(
&mut self,
index: u32,
flags: u32,
start: u32,
count: u32,
fds: &[RawFd],
) -> Result<(), Error> {
let set_irqs = SetIrqs {
header: Header {
message_id: self.next_message_id.0,
command: Command::SetIrqs as u16,
flags: HeaderFlags::Command as u32,
message_size: size_of::<SetIrqs>() as u32,
..Default::default()
},
argsz: (size_of::<SetIrqs>() - size_of::<Header>()) as u32,
flags,
start,
index,
count,
};
debug!("Command: {set_irqs:?}");
self.next_message_id += Wrapping(1);
self.stream
.send_with_fds(&[set_irqs.as_slice()], fds)
.map_err(Error::SendWithFd)?;
let mut reply = Header::default();
self.stream
.read_exact(reply.as_mut_slice())
.map_err(Error::StreamRead)?;
debug!("Reply: {reply:?}");
Ok(())
}
pub fn region(&self, region_index: u32) -> Option<&Region> {
self.regions
.iter()
.find(|®ion| region.index == region_index)
}
pub fn resettable(&self) -> bool {
self.resettable
}
pub fn shutdown(&self) -> Result<(), Error> {
self.stream
.shutdown(std::net::Shutdown::Both)
.map_err(Error::StreamShutdown)
}
}
pub trait ServerBackend {
fn region_read(
&mut self,
_region: u32,
_offset: u64,
_data: &mut [u8],
) -> Result<(), std::io::Error>;
fn region_write(
&mut self,
_region: u32,
_offset: u64,
_data: &[u8],
) -> Result<(), std::io::Error>;
fn dma_map(
&mut self,
_flags: DmaMapFlags,
_offset: u64,
_address: u64,
_size: u64,
_fd: Option<File>,
) -> Result<(), std::io::Error>;
fn dma_unmap(
&mut self,
_flags: DmaUnmapFlags,
_address: u64,
_size: u64,
) -> Result<(), std::io::Error>;
fn reset(&mut self) -> Result<(), std::io::Error>;
fn set_irqs(
&mut self,
_index: u32,
_flags: u32,
_start: u32,
_count: u32,
_fds: Vec<File>,
) -> Result<(), std::io::Error>;
}
#[derive(Clone, Copy)]
pub struct SparseArea {
pub area: vfio_region_sparse_mmap_area,
}
#[derive(Clone)]
pub struct ServerRegion {
pub region_info: vfio_region_info,
pub sparse_areas: Vec<SparseArea>,
pub mmap_fd: Option<RawFd>,
}
pub struct Server {
listener: UnixListener,
path: Option<PathBuf>,
resettable: bool,
irqs: Vec<IrqInfo>,
regions: Vec<ServerRegion>,
}
impl Server {
pub fn new(
path: &Path,
resettable: bool,
irqs: Vec<IrqInfo>,
regions: Vec<ServerRegion>,
) -> Result<Server, Error> {
if path.exists() {
return Err(Error::SocketPathExists);
}
let listener = UnixListener::bind(path).map_err(Error::SocketBind)?;
Ok(Server {
listener,
path: Some(path.to_path_buf()),
resettable,
irqs,
regions,
})
}
pub fn from_owned_fd(
fd: OwnedFd,
resettable: bool,
irqs: Vec<IrqInfo>,
regions: Vec<ServerRegion>,
) -> Server {
let listener = UnixListener::from(fd);
Server {
listener,
path: None,
resettable,
irqs,
regions,
}
}
fn handle_command(
&self,
backend: &mut dyn ServerBackend,
stream: &mut UnixStream,
header: Header,
fds: Vec<File>,
) -> Result<(), Error> {
let command = Command::n(header.command).ok_or(Error::UnknownCommand(header.command))?;
match command {
Command::Unknown
| Command::GetRegionIoFds
| Command::DmaRead
| Command::DmaWrite
| Command::UserDirtyPages => {
return Err(Error::UnsupportedCommand(command));
}
Command::Version => {
if header.no_reply() {
return Err(Error::UnexpectedNoReply(Command::Version));
}
let mut client_version = Version {
header,
..Default::default()
};
stream
.read_exact(&mut client_version.as_mut_slice()[size_of::<Header>()..])
.map_err(Error::StreamRead)?;
let mut raw_version_data =
vec![0; header.message_size as usize - size_of::<Version>()];
stream
.read_exact(&mut raw_version_data)
.map_err(Error::StreamRead)?;
let client_version_data = CString::from_vec_with_nul(raw_version_data)
.unwrap()
.to_string_lossy()
.into_owned();
let client_capabilities: CapabilitiesData =
serde_json::from_str(&client_version_data)
.map_err(Error::DeserializeCapabilites)?;
info!(
"Received client version: major = {} minor = {} capabilities = {:?}",
client_version.major, client_version.minor, client_capabilities.capabilities,
);
let server_capabilities = CapabilitiesData::default();
let server_version_data = serde_json::to_string(&server_capabilities)
.map_err(Error::SerializeCapabilites)?;
let server_version = Version {
header: Header {
message_id: client_version.header.message_id,
command: Command::Version as u16,
flags: HeaderFlags::Reply as u32,
message_size: (size_of::<Version>() + server_version_data.len() + 1) as u32,
..Default::default()
},
major: 0,
minor: 0,
};
let server_version_data = CString::new(server_version_data.as_bytes()).unwrap();
let bufs = vec![
IoSlice::new(server_version.as_slice()),
IoSlice::new(server_version_data.as_bytes_with_nul()),
];
let _ = stream.write_vectored(&bufs).map_err(Error::StreamWrite)?;
info!(
"Sent server version: major = {} minor = {} capabilities = {:?}",
server_version.major, server_version.minor, server_capabilities.capabilities
);
}
Command::DmaMap => {
let mut cmd = DmaMap {
header,
..Default::default()
};
stream
.read_exact(&mut cmd.as_mut_slice()[size_of::<Header>()..])
.map_err(Error::StreamRead)?;
let mut fds = fds;
if fds.len() > 1 {
return Err(Error::InvalidInput);
}
backend
.dma_map(
DmaMapFlags::from_bits_truncate(cmd.flags),
cmd.offset,
cmd.address,
cmd.size,
fds.pop(),
)
.map_err(Error::Backend)?;
if header.no_reply() {
return Ok(());
}
let reply = Header {
message_id: cmd.header.message_id,
command: Command::DmaMap as u16,
flags: HeaderFlags::Reply as u32,
message_size: size_of::<Header>() as u32,
..Default::default()
};
stream
.write_all(reply.as_slice())
.map_err(Error::StreamWrite)?;
}
Command::DmaUnmap => {
let mut cmd = DmaUnmap {
header,
..Default::default()
};
stream
.read_exact(&mut cmd.as_mut_slice()[size_of::<Header>()..])
.map_err(Error::StreamRead)?;
backend
.dma_unmap(
DmaUnmapFlags::from_bits_truncate(cmd.flags),
cmd.address,
cmd.size,
)
.map_err(Error::Backend)?;
if header.no_reply() {
return Ok(());
}
let reply = DmaUnmap {
header: Header {
message_id: cmd.header.message_id,
command: Command::DmaUnmap as u16,
flags: HeaderFlags::Reply as u32,
message_size: size_of::<DmaUnmap>() as u32,
..Default::default()
},
argsz: cmd.argsz,
flags: cmd.flags,
address: cmd.address,
size: cmd.size,
};
stream
.write_all(reply.as_slice())
.map_err(Error::StreamWrite)?;
}
Command::DeviceGetInfo => {
if header.no_reply() {
return Err(Error::UnexpectedNoReply(Command::DeviceGetInfo));
}
let mut cmd = DeviceGetInfo {
header,
..Default::default()
};
stream
.read_exact(&mut cmd.as_mut_slice()[size_of::<Header>()..])
.map_err(Error::StreamRead)?;
let reply = DeviceGetInfo {
header: Header {
message_id: cmd.header.message_id,
command: Command::DeviceGetInfo as u16,
flags: HeaderFlags::Reply as u32,
message_size: size_of::<DeviceGetInfo>() as u32,
..Default::default()
},
argsz: size_of::<DeviceGetInfo>() as u32 - size_of::<Header>() as u32,
flags: VFIO_DEVICE_FLAGS_PCI
| if self.resettable {
VFIO_DEVICE_FLAGS_RESET
} else {
0
},
num_regions: self.regions.len() as u32,
num_irqs: self.irqs.len() as u32,
};
stream
.write_all(reply.as_slice())
.map_err(Error::StreamWrite)?;
}
Command::DeviceGetRegionInfo => {
if header.no_reply() {
return Err(Error::UnexpectedNoReply(Command::DeviceGetRegionInfo));
}
let mut cmd = DeviceGetRegionInfo {
header,
..Default::default()
};
stream
.read_exact(&mut cmd.as_mut_slice()[size_of::<Header>()..])
.map_err(Error::StreamRead)?;
if cmd.region_info.index as usize >= self.regions.len() {
return Err(Error::InvalidInput);
}
let server_region = &self.regions[cmd.region_info.index as usize];
let mut region_info = server_region.region_info;
let sparse_areas = &server_region.sparse_areas;
let mut cap_data: Vec<u8> = Vec::new();
if !sparse_areas.is_empty() {
let cap_header = vfio_info_cap_header {
id: VFIO_REGION_INFO_CAP_SPARSE_MMAP as u16,
version: 1,
next: 0,
};
cap_data.extend_from_slice(unsafe {
std::slice::from_raw_parts(
&cap_header as *const vfio_info_cap_header as *const u8,
size_of::<vfio_info_cap_header>(),
)
});
cap_data.extend_from_slice(&(sparse_areas.len() as u32).to_le_bytes());
cap_data.extend_from_slice(&0u32.to_le_bytes());
for sparse_area in sparse_areas {
cap_data.extend_from_slice(unsafe {
std::slice::from_raw_parts(
&sparse_area.area as *const vfio_region_sparse_mmap_area
as *const u8,
size_of::<vfio_region_sparse_mmap_area>(),
)
});
}
}
let send_cap_data = if !cap_data.is_empty() {
let full_argsz = (size_of::<vfio_region_info>() + cap_data.len()) as u32;
region_info.flags |= VFIO_REGION_INFO_FLAG_CAPS;
region_info.argsz = full_argsz;
region_info.cap_offset = size_of::<vfio_region_info>() as u32;
cmd.region_info.argsz >= full_argsz
} else {
false
};
let message_size = if send_cap_data {
(size_of::<DeviceGetRegionInfo>() + cap_data.len()) as u32
} else {
size_of::<DeviceGetRegionInfo>() as u32
};
let reply = DeviceGetRegionInfo {
header: Header {
message_id: cmd.header.message_id,
command: Command::DeviceGetRegionInfo as u16,
flags: HeaderFlags::Reply as u32,
message_size,
..Default::default()
},
region_info,
};
if send_cap_data {
let reply_bytes = reply.as_slice();
let mut buf = Vec::with_capacity(reply_bytes.len() + cap_data.len());
buf.extend_from_slice(reply_bytes);
buf.extend_from_slice(&cap_data);
if let Some(fd) = server_region.mmap_fd {
stream
.send_with_fds(&[&buf[..]], &[fd])
.map_err(Error::SendWithFd)?;
} else {
stream.write_all(&buf).map_err(Error::StreamWrite)?;
}
} else {
stream
.write_all(reply.as_slice())
.map_err(Error::StreamWrite)?;
}
}
Command::GetIrqInfo => {
if header.no_reply() {
return Err(Error::UnexpectedNoReply(Command::GetIrqInfo));
}
let mut cmd = GetIrqInfo {
header,
..Default::default()
};
stream
.read_exact(&mut cmd.as_mut_slice()[size_of::<Header>()..])
.map_err(Error::StreamRead)?;
if cmd.index as usize >= self.irqs.len() {
return Err(Error::InvalidInput);
}
let irq = &self.irqs[cmd.index as usize];
let reply = GetIrqInfo {
header: Header {
message_id: cmd.header.message_id,
command: Command::GetIrqInfo as u16,
flags: HeaderFlags::Reply as u32,
message_size: size_of::<GetIrqInfo>() as u32,
..Default::default()
},
argsz: (size_of::<GetIrqInfo>() - size_of::<Header>()) as u32,
index: irq.index,
flags: irq.flags,
count: irq.count,
};
stream
.write_all(reply.as_slice())
.map_err(Error::StreamWrite)?;
}
Command::SetIrqs => {
let mut cmd = SetIrqs {
header,
..Default::default()
};
stream
.read_exact(&mut cmd.as_mut_slice()[size_of::<Header>()..])
.map_err(Error::StreamRead)?;
if cmd.index as usize >= self.irqs.len() {
return Err(Error::InvalidInput);
}
if cmd.flags & VFIO_IRQ_SET_DATA_BOOL > 0 {
return Err(Error::UnsupportedFeature);
}
backend
.set_irqs(cmd.index, cmd.flags, cmd.start, cmd.count, fds)
.map_err(Error::Backend)?;
if header.no_reply() {
return Ok(());
}
let reply = Header {
message_id: cmd.header.message_id,
command: Command::SetIrqs as u16,
flags: HeaderFlags::Reply as u32,
message_size: size_of::<Header>() as u32,
..Default::default()
};
stream
.write_all(reply.as_slice())
.map_err(Error::StreamWrite)?;
}
Command::RegionRead => {
if header.no_reply() {
return Err(Error::UnexpectedNoReply(Command::RegionRead));
}
let mut cmd = RegionAccess {
header,
..Default::default()
};
stream
.read_exact(&mut cmd.as_mut_slice()[size_of::<Header>()..])
.map_err(Error::StreamRead)?;
let (region, offset, count) = (cmd.region, cmd.offset, cmd.count);
if region as usize >= self.regions.len() {
return Err(Error::InvalidInput);
}
let mut data = vec![0u8; count as usize];
backend
.region_read(region, offset, &mut data)
.map_err(Error::Backend)?;
let reply = RegionAccess {
header: Header {
message_id: cmd.header.message_id,
command: Command::RegionRead as u16,
flags: HeaderFlags::Reply as u32,
message_size: size_of::<RegionAccess>() as u32 + count,
..Default::default()
},
region,
offset,
count,
};
stream
.write_all(reply.as_slice())
.map_err(Error::StreamWrite)?;
stream.write_all(&data).map_err(Error::StreamWrite)?;
}
Command::RegionWrite => {
let mut cmd = RegionAccess {
header,
..Default::default()
};
stream
.read_exact(&mut cmd.as_mut_slice()[size_of::<Header>()..])
.map_err(Error::StreamRead)?;
let (region, offset, count) = (cmd.region, cmd.offset, cmd.count);
if region as usize >= self.regions.len() {
return Err(Error::InvalidInput);
}
let mut data = vec![0u8; count as usize];
stream.read_exact(&mut data).map_err(Error::StreamRead)?;
backend
.region_write(region, offset, &data)
.map_err(Error::Backend)?;
if header.no_reply() {
return Ok(());
}
let reply = RegionAccess {
header: Header {
message_id: cmd.header.message_id,
command: Command::RegionWrite as u16,
flags: HeaderFlags::Reply as u32,
message_size: size_of::<RegionAccess>() as u32,
..Default::default()
},
region,
offset,
count,
};
stream
.write_all(reply.as_slice())
.map_err(Error::StreamWrite)?;
}
Command::DeviceReset => {
backend.reset().map_err(Error::Backend)?;
if header.no_reply() {
return Ok(());
}
let reply = Header {
message_id: header.message_id,
command: Command::DeviceReset as u16,
flags: HeaderFlags::Reply as u32,
message_size: size_of::<Header>() as u32,
..Default::default()
};
stream
.write_all(reply.as_slice())
.map_err(Error::StreamWrite)?;
}
}
Ok(())
}
pub fn run(&self, backend: &mut dyn ServerBackend) -> Result<(), Error> {
let (mut stream, _) = self.listener.accept().map_err(Error::SocketAccept)?;
loop {
let mut header = Header::default();
let mut fds = vec![0; 16];
let mut iovecs = vec![iovec {
iov_base: header.as_mut_slice().as_mut_ptr() as *mut c_void,
iov_len: header.as_mut_slice().len(),
}];
let (bytes, fds_received) = unsafe {
stream
.recv_with_fds(&mut iovecs, &mut fds)
.map_err(Error::ReceiveWithFd)?
};
if bytes == 0 {
info!("Connection closed");
break;
}
fds.resize(fds_received, 0);
let fds: Vec<File> = fds
.iter()
.map(|fd| unsafe { File::from_raw_fd(*fd) })
.collect();
if let Err(e) = self.handle_command(backend, &mut stream, header, fds) {
error!("Error handling command: {:?}: {e}", header.command);
let reply = Header {
message_id: header.message_id,
command: header.command,
flags: HeaderFlags::Error as u32 | HeaderFlags::Reply as u32,
message_size: size_of::<Header>() as u32,
error: if matches!(e, Error::InvalidInput) {
EINVAL as u32
} else {
0
},
};
stream
.write_all(reply.as_slice())
.map_err(Error::StreamWrite)?;
}
}
Ok(())
}
}
impl Drop for Server {
fn drop(&mut self) {
if let Some(path) = &self.path {
if path.exists() {
let _ = std::fs::remove_file(path);
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::mem::size_of;
fn build_sparse_cap_data(areas: &[vfio_region_sparse_mmap_area]) -> Vec<u8> {
let mut cap_data: Vec<u8> = Vec::new();
let cap_header = vfio_info_cap_header {
id: VFIO_REGION_INFO_CAP_SPARSE_MMAP as u16,
version: 1,
next: 0,
};
cap_data.extend_from_slice(unsafe {
std::slice::from_raw_parts(
&cap_header as *const vfio_info_cap_header as *const u8,
size_of::<vfio_info_cap_header>(),
)
});
cap_data.extend_from_slice(&(areas.len() as u32).to_le_bytes());
cap_data.extend_from_slice(&0u32.to_le_bytes());
for area in areas {
cap_data.extend_from_slice(unsafe {
std::slice::from_raw_parts(
area as *const vfio_region_sparse_mmap_area as *const u8,
size_of::<vfio_region_sparse_mmap_area>(),
)
});
}
cap_data
}
#[test]
fn test_parse_sparse_mmap_caps() {
let areas = vec![
vfio_region_sparse_mmap_area {
offset: 0x0,
size: 0x1000,
},
vfio_region_sparse_mmap_area {
offset: 0x2000,
size: 0x3000,
},
];
let cap_data = build_sparse_cap_data(&areas);
let region_info = vfio_region_info {
argsz: (size_of::<vfio_region_info>() + cap_data.len()) as u32,
flags: VFIO_REGION_INFO_FLAG_CAPS,
cap_offset: size_of::<vfio_region_info>() as u32,
..Default::default()
};
let parsed = Client::parse_region_caps(&cap_data, ®ion_info).unwrap();
assert_eq!(parsed.len(), 2);
assert_eq!(parsed[0].offset, 0x0);
assert_eq!(parsed[0].size, 0x1000);
assert_eq!(parsed[1].offset, 0x2000);
assert_eq!(parsed[1].size, 0x3000);
}
#[test]
fn test_parse_empty_sparse_mmap_caps() {
let cap_data = build_sparse_cap_data(&[]);
let region_info = vfio_region_info {
argsz: (size_of::<vfio_region_info>() + cap_data.len()) as u32,
flags: VFIO_REGION_INFO_FLAG_CAPS,
cap_offset: size_of::<vfio_region_info>() as u32,
..Default::default()
};
let parsed = Client::parse_region_caps(&cap_data, ®ion_info).unwrap();
assert!(parsed.is_empty());
}
#[test]
fn test_no_caps_returns_empty() {
let region_info = vfio_region_info {
argsz: size_of::<vfio_region_info>() as u32,
flags: 0,
cap_offset: 0,
..Default::default()
};
let parsed = Client::parse_region_caps(&[], ®ion_info).unwrap();
assert!(parsed.is_empty());
}
}