use alloc::string::String;
use alloc::sync::Arc;
use alloc::vec::Vec;
use crate::core::error::{Error, Result};
use crate::core::state::{ChunkReader, ChunkWriter, Sink, Source};
use crate::dev::ata::{Medium, Snapshot};
use super::queue::{Descriptor, Queue};
use super::{Backend, DEVICE_ID_BLOCK};
pub const SECTOR_SIZE: u64 = 512;
const HEADER_LEN: u64 = 16;
const CHUNK: u64 = 64 * 1024;
const T_IN: u32 = 0;
const T_OUT: u32 = 1;
const T_FLUSH: u32 = 4;
const T_GET_ID: u32 = 8;
const S_OK: u8 = 0;
const S_IOERR: u8 = 1;
const S_UNSUPP: u8 = 2;
const F_RO: u64 = 1 << 5;
const F_FLUSH: u64 = 1 << 9;
const ID_BYTES: usize = 20;
#[derive(Debug)]
pub struct VirtioBlk {
media: Arc<dyn Medium>,
serial: String,
read_only: bool,
bytes: u64,
}
impl VirtioBlk {
pub fn new(media: Arc<dyn Medium>, serial: String, read_only: bool) -> Result<VirtioBlk> {
let bytes = media.capacity();
if bytes == 0 || !bytes.is_multiple_of(SECTOR_SIZE) {
return Err(Error::Config {
at: String::from(super::BLK_CLASS_NAME),
message: alloc::format!(
"a virtio disk holds a whole number of {SECTOR_SIZE}-byte sectors, and \
{bytes} byte(s) is not a whole number of them"
),
});
}
Ok(VirtioBlk {
read_only: read_only || media.is_read_only(),
media,
serial,
bytes,
})
}
#[must_use]
pub fn capacity(&self) -> u64 {
self.bytes / SECTOR_SIZE
}
#[must_use]
pub fn medium(&self) -> &Arc<dyn Medium> {
&self.media
}
#[must_use]
pub fn is_read_only(&self) -> bool {
self.read_only
}
pub fn contents(&self) -> Result<Vec<u8>> {
let mut out = alloc::vec![0u8; self.bytes as usize];
self.media
.read_at(0, &mut out)
.map_err(|e| Error::State(alloc::format!("virtio disk: {e}")))?;
Ok(out)
}
fn range(&self, sector: u64, len: u64) -> Option<u64> {
let at = sector.checked_mul(SECTOR_SIZE)?;
let end = at.checked_add(len)?;
(end <= self.bytes).then_some(at)
}
fn read_in(&self, q: &Queue<'_>, chain: &[Descriptor], sector: u64, len: u64) -> (u8, u64) {
let Some(at) = self.range(sector, len) else {
return (S_IOERR, 0);
};
let mut buf = alloc::vec![0u8; CHUNK.min(len) as usize];
let mut done = 0u64;
while done < len {
let take = CHUNK.min(len - done) as usize;
let part = &mut buf[..take];
if self.media.read_at(at + done, part).is_err() {
return (S_IOERR, done);
}
match q.write_chain(chain, done, part) {
Ok(n) => {
done += n as u64;
if n < take {
return (S_IOERR, done);
}
}
Err(_) => return (S_IOERR, done),
}
}
(S_OK, done)
}
fn write_out(&self, q: &Queue<'_>, chain: &[Descriptor], sector: u64, len: u64) -> (u8, u64) {
if self.read_only {
return (S_IOERR, 0);
}
let Some(at) = self.range(sector, len) else {
return (S_IOERR, 0);
};
let mut buf = alloc::vec![0u8; CHUNK.min(len) as usize];
let mut done = 0u64;
while done < len {
let take = CHUNK.min(len - done) as usize;
let part = &mut buf[..take];
if q.read_chain(chain, HEADER_LEN + done, part).unwrap_or(0) < take {
return (S_IOERR, 0);
}
if self.media.write_at(at + done, part).is_err() {
return (S_IOERR, 0);
}
done += take as u64;
}
(S_OK, 0)
}
fn serve(&self, q: &Queue<'_>, chain: &[Descriptor]) -> (u8, u64) {
let writable = Queue::writable_len(chain);
let readable = Queue::readable_len(chain);
if writable == 0 || readable < HEADER_LEN {
return (S_OK, 0);
}
let mut header = [0u8; HEADER_LEN as usize];
if q.read_chain(chain, 0, &mut header).unwrap_or(0) < header.len() {
return (S_IOERR, 0);
}
let kind = u32::from_le_bytes([header[0], header[1], header[2], header[3]]);
let sector = u64::from_le_bytes([
header[8], header[9], header[10], header[11], header[12], header[13], header[14],
header[15],
]);
let payload = writable - 1;
match kind {
T_IN => self.read_in(q, chain, sector, payload),
T_OUT => self.write_out(q, chain, sector, readable - HEADER_LEN),
T_FLUSH => match self.media.flush() {
Ok(()) => (S_OK, 0),
Err(_) => (S_IOERR, 0),
},
T_GET_ID => {
let mut id = [0u8; ID_BYTES];
let serial = self.serial.as_bytes();
let take = serial.len().min(ID_BYTES);
id[..take].copy_from_slice(&serial[..take]);
let take = (payload as usize).min(ID_BYTES);
match q.write_chain(chain, 0, &id[..take]) {
Ok(n) => (S_OK, n as u64),
Err(_) => (S_IOERR, 0),
}
}
_ => (S_UNSUPP, 0),
}
}
}
impl Backend for VirtioBlk {
fn device_id(&self) -> u32 {
DEVICE_ID_BLOCK
}
fn queue_count(&self) -> usize {
1
}
fn features(&self) -> u64 {
F_FLUSH | if self.read_only { F_RO } else { 0 }
}
fn config_read(&self, offset: u64, dst: &mut [u8]) {
let capacity = self.capacity().to_le_bytes();
for (i, byte) in dst.iter_mut().enumerate() {
let at = offset + i as u64;
*byte = usize::try_from(at)
.ok()
.and_then(|at| capacity.get(at))
.copied()
.unwrap_or(0);
}
}
fn handle(&self, _queue: usize, q: &Queue<'_>, chain: &[Descriptor]) -> u32 {
let (status, written) = self.serve(q, chain);
let writable = Queue::writable_len(chain);
if writable == 0 {
return 0;
}
let _ = q.write_chain(chain, writable - 1, &[status]);
(written + 1) as u32
}
fn reset(&self) {
}
fn save(&self, w: &mut ChunkWriter<'_>) -> Result<()> {
match self.media.snapshot() {
Snapshot::Capture => w.write_bytes(&self.contents()?),
Snapshot::Reference => {
self.media
.flush()
.map_err(|e| Error::State(alloc::format!("virtio disk: {e}")))?;
w.write_bytes(self.media.describe().as_bytes())
}
Snapshot::Refuse => Err(Error::State(alloc::format!(
"this virtio disk's medium ({}) refuses to be snapshotted",
self.media.describe()
))),
}
}
fn load(&self, r: &mut ChunkReader<'_>) -> Result<()> {
let bytes: &[u8] = r.read_bytes()?;
match self.media.snapshot() {
Snapshot::Capture => {
if bytes.len() as u64 != self.bytes {
return Err(Error::State(alloc::format!(
"snapshot has a {}-byte disk, this device has {}",
bytes.len(),
self.bytes
)));
}
self.media
.write_at(0, bytes)
.map_err(|e| Error::State(alloc::format!("virtio disk: {e}")))
}
Snapshot::Reference => {
let want = self.media.describe();
if bytes != want.as_bytes() {
return Err(Error::State(alloc::format!(
"the snapshot references a different medium: it names `{}` and this \
disk holds `{want}`",
String::from_utf8_lossy(&bytes[..bytes.len().min(120)])
)));
}
Ok(())
}
Snapshot::Refuse => Err(Error::State(alloc::format!(
"this virtio disk's medium ({}) refuses to be snapshotted",
self.media.describe()
))),
}
}
}
#[cfg(test)]
mod tests;