use std::path::PathBuf;
use arcbox_virtio_core::error::VirtioError;
use arcbox_virtio_core::virtio_bindings;
#[derive(Debug, Clone)]
pub struct BlockConfig {
pub capacity: u64,
pub blk_size: u32,
pub path: PathBuf,
pub read_only: bool,
pub num_queues: u16,
}
impl Default for BlockConfig {
fn default() -> Self {
Self {
capacity: 0,
blk_size: 512,
path: PathBuf::new(),
read_only: false,
num_queues: 1,
}
}
}
pub const WRITE_ZEROES_FLAG_UNMAP: u32 = 1;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct DiscardWriteZeroesRange {
pub sector: u64,
pub num_sectors: u32,
pub flags: u32,
}
const RANGE_ENTRY_SIZE: usize = 16;
pub fn parse_range_list(
bytes: &[u8],
) -> std::result::Result<Vec<DiscardWriteZeroesRange>, VirtioError> {
if bytes.is_empty() || bytes.len() % RANGE_ENTRY_SIZE != 0 {
return Err(VirtioError::InvalidOperation(format!(
"range list size {} not a multiple of 16",
bytes.len()
)));
}
let mut ranges = Vec::with_capacity(bytes.len() / RANGE_ENTRY_SIZE);
for chunk in bytes.chunks_exact(RANGE_ENTRY_SIZE) {
ranges.push(DiscardWriteZeroesRange {
sector: u64::from_le_bytes(chunk[0..8].try_into().unwrap()),
num_sectors: u32::from_le_bytes(chunk[8..12].try_into().unwrap()),
flags: u32::from_le_bytes(chunk[12..16].try_into().unwrap()),
});
}
Ok(ranges)
}
impl DiscardWriteZeroesRange {
pub fn checked_byte_range(
self,
blk_size: u32,
capacity_sectors: u64,
) -> std::result::Result<(u64, u64), VirtioError> {
let sector_end = self
.sector
.checked_add(u64::from(self.num_sectors))
.ok_or_else(|| VirtioError::InvalidOperation("range sector overflow".into()))?;
if sector_end > capacity_sectors {
return Err(VirtioError::InvalidOperation(format!(
"range exceeds device capacity: {}..{} > {} sectors",
self.sector, sector_end, capacity_sectors
)));
}
let block_size = u64::from(blk_size);
let start = self
.sector
.checked_mul(block_size)
.ok_or_else(|| VirtioError::InvalidOperation("range byte offset overflow".into()))?;
let bytes = u64::from(self.num_sectors)
.checked_mul(block_size)
.ok_or_else(|| VirtioError::InvalidOperation("range byte length overflow".into()))?;
let end = start
.checked_add(bytes)
.ok_or_else(|| VirtioError::InvalidOperation("range byte end overflow".into()))?;
Ok((start, end))
}
}
pub fn checked_io_byte_range(
sector: u64,
byte_len: usize,
blk_size: u32,
capacity_sectors: u64,
) -> std::result::Result<(u64, u64), VirtioError> {
let block_size = u64::from(blk_size);
let capacity_bytes = capacity_sectors
.checked_mul(block_size)
.ok_or_else(|| VirtioError::InvalidOperation("device capacity overflow".into()))?;
let start = sector
.checked_mul(block_size)
.ok_or_else(|| VirtioError::InvalidOperation("I/O byte offset overflow".into()))?;
let len = u64::try_from(byte_len)
.map_err(|_| VirtioError::InvalidOperation("I/O byte length overflow".into()))?;
let end = start
.checked_add(len)
.ok_or_else(|| VirtioError::InvalidOperation("I/O byte end overflow".into()))?;
if end > capacity_bytes {
return Err(VirtioError::InvalidOperation(format!(
"I/O exceeds device capacity: bytes {}..{} > {}",
start, end, capacity_bytes
)));
}
Ok((start, end))
}
#[repr(u32)]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BlockRequestType {
In = virtio_bindings::virtio_blk::VIRTIO_BLK_T_IN,
Out = virtio_bindings::virtio_blk::VIRTIO_BLK_T_OUT,
Flush = virtio_bindings::virtio_blk::VIRTIO_BLK_T_FLUSH,
GetId = virtio_bindings::virtio_blk::VIRTIO_BLK_T_GET_ID,
Discard = virtio_bindings::virtio_blk::VIRTIO_BLK_T_DISCARD,
WriteZeroes = virtio_bindings::virtio_blk::VIRTIO_BLK_T_WRITE_ZEROES,
}
impl TryFrom<u32> for BlockRequestType {
type Error = VirtioError;
fn try_from(value: u32) -> std::result::Result<Self, Self::Error> {
use virtio_bindings::virtio_blk;
match value {
virtio_blk::VIRTIO_BLK_T_IN => Ok(Self::In),
virtio_blk::VIRTIO_BLK_T_OUT => Ok(Self::Out),
virtio_blk::VIRTIO_BLK_T_FLUSH => Ok(Self::Flush),
virtio_blk::VIRTIO_BLK_T_GET_ID => Ok(Self::GetId),
virtio_blk::VIRTIO_BLK_T_DISCARD => Ok(Self::Discard),
virtio_blk::VIRTIO_BLK_T_WRITE_ZEROES => Ok(Self::WriteZeroes),
_ => Err(VirtioError::InvalidOperation(format!(
"Unknown block request type: {value}"
))),
}
}
}
#[repr(u8)]
#[derive(Debug, Clone, Copy)]
pub enum BlockStatus {
Ok = virtio_bindings::virtio_blk::VIRTIO_BLK_S_OK as u8,
IoErr = virtio_bindings::virtio_blk::VIRTIO_BLK_S_IOERR as u8,
Unsupp = virtio_bindings::virtio_blk::VIRTIO_BLK_S_UNSUPP as u8,
}
#[repr(C)]
#[derive(Debug, Clone, Copy)]
pub struct BlockRequestHeader {
pub request_type: u32,
pub reserved: u32,
pub sector: u64,
}
impl BlockRequestHeader {
pub const SIZE: usize = 16;
#[must_use]
pub fn from_bytes(bytes: &[u8]) -> Option<Self> {
if bytes.len() < Self::SIZE {
return None;
}
Some(Self {
request_type: u32::from_le_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]),
reserved: u32::from_le_bytes([bytes[4], bytes[5], bytes[6], bytes[7]]),
sector: u64::from_le_bytes([
bytes[8], bytes[9], bytes[10], bytes[11], bytes[12], bytes[13], bytes[14],
bytes[15],
]),
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_request_header_parsing() {
let bytes = [
0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x10, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, ];
let header = BlockRequestHeader::from_bytes(&bytes).unwrap();
assert_eq!(header.request_type, 0);
assert_eq!(header.sector, 16);
}
#[test]
fn test_request_header_too_short() {
let bytes = [0x00, 0x00, 0x00];
let header = BlockRequestHeader::from_bytes(&bytes);
assert!(header.is_none());
}
#[test]
fn test_invalid_request_type() {
let result = BlockRequestType::try_from(999u32);
assert!(result.is_err());
}
#[test]
fn test_all_request_types() {
assert_eq!(BlockRequestType::try_from(0).unwrap(), BlockRequestType::In);
assert_eq!(
BlockRequestType::try_from(1).unwrap(),
BlockRequestType::Out
);
assert_eq!(
BlockRequestType::try_from(4).unwrap(),
BlockRequestType::Flush
);
assert_eq!(
BlockRequestType::try_from(8).unwrap(),
BlockRequestType::GetId
);
assert_eq!(
BlockRequestType::try_from(11).unwrap(),
BlockRequestType::Discard
);
assert_eq!(
BlockRequestType::try_from(13).unwrap(),
BlockRequestType::WriteZeroes
);
}
}