devices/virtio/
descriptor_utils.rs

1// Copyright 2019 The ChromiumOS Authors
2// Use of this source code is governed by a BSD-style license that can be
3// found in the LICENSE file.
4
5use std::cmp;
6use std::io;
7use std::io::Write;
8use std::iter::FromIterator;
9use std::marker::PhantomData;
10use std::mem::size_of;
11use std::mem::MaybeUninit;
12use std::ptr::copy_nonoverlapping;
13use std::sync::Arc;
14
15use anyhow::Context;
16use base::FileReadWriteAtVolatile;
17use base::FileReadWriteVolatile;
18use base::VolatileSlice;
19use cros_async::IoOptions;
20use cros_async::MemRegion;
21use cros_async::MemRegionIter;
22use data_model::Le16;
23use data_model::Le32;
24use data_model::Le64;
25use disk::AsyncDisk;
26use smallvec::SmallVec;
27use vm_memory::GuestAddress;
28use vm_memory::GuestMemory;
29use zerocopy::FromBytes;
30use zerocopy::Immutable;
31use zerocopy::IntoBytes;
32use zerocopy::KnownLayout;
33
34use super::DescriptorChain;
35use crate::virtio::SplitDescriptorChain;
36
37struct DescriptorChainRegions {
38    regions: SmallVec<[MemRegion; 2]>,
39
40    // Index of the current region in `regions`.
41    current_region_index: usize,
42
43    // Number of bytes consumed in the current region.
44    current_region_offset: usize,
45
46    // Total bytes consumed in the entire descriptor chain.
47    bytes_consumed: usize,
48}
49
50impl DescriptorChainRegions {
51    fn new(regions: SmallVec<[MemRegion; 2]>) -> Self {
52        DescriptorChainRegions {
53            regions,
54            current_region_index: 0,
55            current_region_offset: 0,
56            bytes_consumed: 0,
57        }
58    }
59
60    fn available_bytes(&self) -> usize {
61        // This is guaranteed not to overflow because the total length of the chain is checked
62        // during all creations of `DescriptorChain` (see `DescriptorChain::new()`).
63        self.get_remaining_regions()
64            .fold(0usize, |count, region| count + region.len)
65    }
66
67    fn bytes_consumed(&self) -> usize {
68        self.bytes_consumed
69    }
70
71    /// Returns all the remaining buffers in the `DescriptorChain`. Calling this function does not
72    /// consume any bytes from the `DescriptorChain`. Instead callers should use the `consume`
73    /// method to advance the `DescriptorChain`. Multiple calls to `get` with no intervening calls
74    /// to `consume` will return the same data.
75    fn get_remaining_regions(&self) -> MemRegionIter {
76        MemRegionIter::new(&self.regions[self.current_region_index..])
77            .skip_bytes(self.current_region_offset)
78    }
79
80    /// Like `get_remaining_regions` but guarantees that the combined length of all the returned
81    /// iovecs is not greater than `count`. The combined length of the returned iovecs may be less
82    /// than `count` but will always be greater than 0 as long as there is still space left in the
83    /// `DescriptorChain`.
84    fn get_remaining_regions_with_count(&self, count: usize) -> MemRegionIter {
85        MemRegionIter::new(&self.regions[self.current_region_index..])
86            .skip_bytes(self.current_region_offset)
87            .take_bytes(count)
88    }
89
90    /// Returns all the remaining buffers in the `DescriptorChain` as `VolatileSlice`s of the given
91    /// `GuestMemory`. Calling this function does not consume any bytes from the `DescriptorChain`.
92    /// Instead callers should use the `consume` method to advance the `DescriptorChain`. Multiple
93    /// calls to `get` with no intervening calls to `consume` will return the same data.
94    fn get_remaining<'mem>(&self, mem: &'mem GuestMemory) -> SmallVec<[VolatileSlice<'mem>; 16]> {
95        self.get_remaining_regions()
96            .filter_map(|region| {
97                mem.get_slice_at_addr(GuestAddress(region.offset), region.len)
98                    .ok()
99            })
100            .collect()
101    }
102
103    /// Like 'get_remaining_regions_with_count' except convert the offsets to volatile slices in
104    /// the 'GuestMemory' given by 'mem'.
105    fn get_remaining_with_count<'mem>(
106        &self,
107        mem: &'mem GuestMemory,
108        count: usize,
109    ) -> SmallVec<[VolatileSlice<'mem>; 16]> {
110        self.get_remaining_regions_with_count(count)
111            .filter_map(|region| {
112                mem.get_slice_at_addr(GuestAddress(region.offset), region.len)
113                    .ok()
114            })
115            .collect()
116    }
117
118    /// Consumes `count` bytes from the `DescriptorChain`. If `count` is larger than
119    /// `self.available_bytes()` then all remaining bytes in the `DescriptorChain` will be consumed.
120    fn consume(&mut self, mut count: usize) {
121        while let Some(region) = self.regions.get(self.current_region_index) {
122            let region_remaining = region.len - self.current_region_offset;
123            if count < region_remaining {
124                // The remaining count to consume is less than the remaining un-consumed length of
125                // the current region. Adjust the region offset without advancing to the next region
126                // and stop.
127                self.current_region_offset += count;
128                self.bytes_consumed += count;
129                return;
130            }
131
132            // The current region has been exhausted. Advance to the next region.
133            self.current_region_index += 1;
134            self.current_region_offset = 0;
135
136            self.bytes_consumed += region_remaining;
137            count -= region_remaining;
138        }
139    }
140
141    fn split_at(&mut self, offset: usize) -> DescriptorChainRegions {
142        let mut other = DescriptorChainRegions {
143            regions: self.regions.clone(),
144            current_region_index: self.current_region_index,
145            current_region_offset: self.current_region_offset,
146            bytes_consumed: self.bytes_consumed,
147        };
148        other.consume(offset);
149        other.bytes_consumed = 0;
150
151        let mut rem = offset;
152        let mut end = self.current_region_index;
153        for region in &mut self.regions[self.current_region_index..] {
154            if rem <= region.len {
155                region.len = rem;
156                break;
157            }
158
159            end += 1;
160            rem -= region.len;
161        }
162
163        self.regions.truncate(end + 1);
164
165        other
166    }
167}
168
169/// Provides high-level interface over the sequence of memory regions
170/// defined by readable descriptors in the descriptor chain.
171///
172/// Note that virtio spec requires driver to place any device-writable
173/// descriptors after any device-readable descriptors (2.6.4.2 in Virtio Spec v1.1).
174/// Reader will skip iterating over descriptor chain when first writable
175/// descriptor is encountered.
176pub struct Reader {
177    mem: GuestMemory,
178    regions: DescriptorChainRegions,
179}
180
181/// An iterator over `FromBytes` objects on readable descriptors in the descriptor chain.
182pub struct ReaderIterator<'a, T: FromBytes> {
183    reader: &'a mut Reader,
184    phantom: PhantomData<T>,
185}
186
187impl<T: FromBytes> Iterator for ReaderIterator<'_, T> {
188    type Item = io::Result<T>;
189
190    fn next(&mut self) -> Option<io::Result<T>> {
191        if self.reader.available_bytes() == 0 {
192            None
193        } else {
194            Some(self.reader.read_obj())
195        }
196    }
197}
198
199impl Reader {
200    /// Construct a new Reader wrapper over `readable_regions`.
201    pub fn new_from_regions(
202        mem: &GuestMemory,
203        readable_regions: SmallVec<[MemRegion; 2]>,
204    ) -> Reader {
205        Reader {
206            mem: mem.clone(),
207            regions: DescriptorChainRegions::new(readable_regions),
208        }
209    }
210
211    /// Reads an object from the descriptor chain buffer without consuming it.
212    pub fn peek_obj<T: FromBytes>(&self) -> io::Result<T> {
213        let mut obj = MaybeUninit::uninit();
214
215        // SAFETY: We pass a valid pointer and size of `obj`.
216        let copied = unsafe {
217            copy_regions_to_mut_ptr(
218                &self.mem,
219                self.get_remaining_regions(),
220                obj.as_mut_ptr() as *mut u8,
221                size_of::<T>(),
222            )?
223        };
224        if copied != size_of::<T>() {
225            return Err(io::Error::from(io::ErrorKind::UnexpectedEof));
226        }
227
228        // SAFETY: `FromBytes` guarantees any set of initialized bytes is a valid value for `T`, and
229        // we initialized all bytes in `obj` in the copy above.
230        Ok(unsafe { obj.assume_init() })
231    }
232
233    /// Reads and consumes an object from the descriptor chain buffer.
234    pub fn read_obj<T: FromBytes>(&mut self) -> io::Result<T> {
235        let obj = self.peek_obj::<T>()?;
236        self.consume(size_of::<T>());
237        Ok(obj)
238    }
239
240    /// Reads objects by consuming all the remaining data in the descriptor chain buffer and returns
241    /// them as a collection. Returns an error if the size of the remaining data is indivisible by
242    /// the size of an object of type `T`.
243    pub fn collect<C: FromIterator<io::Result<T>>, T: FromBytes>(&mut self) -> C {
244        self.iter().collect()
245    }
246
247    /// Creates an iterator for sequentially reading `FromBytes` objects from the `Reader`.
248    /// Unlike `collect`, this doesn't consume all the remaining data in the `Reader` and
249    /// doesn't require the objects to be stored in a separate collection.
250    pub fn iter<T: FromBytes>(&mut self) -> ReaderIterator<T> {
251        ReaderIterator {
252            reader: self,
253            phantom: PhantomData,
254        }
255    }
256
257    /// Reads data into a volatile slice up to the minimum of the slice's length or the number of
258    /// bytes remaining. Returns the number of bytes read.
259    pub fn read_to_volatile_slice(&mut self, slice: VolatileSlice) -> usize {
260        let mut read = 0usize;
261        let mut dst = slice;
262        for src in self.get_remaining() {
263            src.copy_to_volatile_slice(dst);
264            let copied = std::cmp::min(src.size(), dst.size());
265            read += copied;
266            dst = match dst.offset(copied) {
267                Ok(v) => v,
268                Err(_) => break, // The slice is fully consumed
269            };
270        }
271        self.regions.consume(read);
272        read
273    }
274
275    /// Reads data from the descriptor chain buffer and passes the `VolatileSlice`s to the callback
276    /// `cb`.
277    pub fn read_to_cb<C: FnOnce(&[VolatileSlice]) -> usize>(
278        &mut self,
279        cb: C,
280        count: usize,
281    ) -> usize {
282        let iovs = self.regions.get_remaining_with_count(&self.mem, count);
283        let written = cb(&iovs[..]);
284        self.regions.consume(written);
285        written
286    }
287
288    /// Reads data from the descriptor chain buffer into a writable object.
289    /// Returns the number of bytes read from the descriptor chain buffer.
290    /// The number of bytes read can be less than `count` if there isn't
291    /// enough data in the descriptor chain buffer.
292    pub fn read_to<F: FileReadWriteVolatile>(
293        &mut self,
294        mut dst: F,
295        count: usize,
296    ) -> io::Result<usize> {
297        let iovs = self.regions.get_remaining_with_count(&self.mem, count);
298        let written = dst.write_vectored_volatile(&iovs[..])?;
299        self.regions.consume(written);
300        Ok(written)
301    }
302
303    /// Reads data from the descriptor chain buffer into a File at offset `off`.
304    /// Returns the number of bytes read from the descriptor chain buffer.
305    /// The number of bytes read can be less than `count` if there isn't
306    /// enough data in the descriptor chain buffer.
307    pub fn read_to_at<F: FileReadWriteAtVolatile>(
308        &mut self,
309        dst: &F,
310        count: usize,
311        off: u64,
312    ) -> io::Result<usize> {
313        let iovs = self.regions.get_remaining_with_count(&self.mem, count);
314        let written = dst.write_vectored_at_volatile(&iovs[..], off)?;
315        self.regions.consume(written);
316        Ok(written)
317    }
318
319    /// Reads data from the descriptor chain similar to 'read_to' except reading 'count' or
320    /// returning an error if 'count' bytes can't be read.
321    pub fn read_exact_to<F: FileReadWriteVolatile>(
322        &mut self,
323        mut dst: F,
324        mut count: usize,
325    ) -> io::Result<()> {
326        while count > 0 {
327            match self.read_to(&mut dst, count) {
328                Ok(0) => {
329                    return Err(io::Error::new(
330                        io::ErrorKind::UnexpectedEof,
331                        "failed to fill whole buffer",
332                    ))
333                }
334                Ok(n) => count -= n,
335                Err(ref e) if e.kind() == io::ErrorKind::Interrupted => {}
336                Err(e) => return Err(e),
337            }
338        }
339
340        Ok(())
341    }
342
343    /// Reads data from the descriptor chain similar to 'read_to_at' except reading 'count' or
344    /// returning an error if 'count' bytes can't be read.
345    pub fn read_exact_to_at<F: FileReadWriteAtVolatile>(
346        &mut self,
347        dst: &F,
348        mut count: usize,
349        mut off: u64,
350    ) -> io::Result<()> {
351        while count > 0 {
352            match self.read_to_at(dst, count, off) {
353                Ok(0) => {
354                    return Err(io::Error::new(
355                        io::ErrorKind::UnexpectedEof,
356                        "failed to fill whole buffer",
357                    ))
358                }
359                Ok(n) => {
360                    count -= n;
361                    off += n as u64;
362                }
363                Err(ref e) if e.kind() == io::ErrorKind::Interrupted => {}
364                Err(e) => return Err(e),
365            }
366        }
367
368        Ok(())
369    }
370
371    /// Reads data from the descriptor chain buffer into an `AsyncDisk` at offset `off`.
372    /// Returns the number of bytes read from the descriptor chain buffer.
373    /// The number of bytes read can be less than `count` if there isn't
374    /// enough data in the descriptor chain buffer.
375    pub async fn read_to_at_fut<F: AsyncDisk + ?Sized>(
376        &mut self,
377        dst: &F,
378        count: usize,
379        off: u64,
380        options: IoOptions,
381    ) -> disk::Result<usize> {
382        let written = dst
383            .write_from_mem(
384                off,
385                Arc::new(self.mem.clone()),
386                self.regions.get_remaining_regions_with_count(count),
387                options,
388            )
389            .await?;
390        self.regions.consume(written);
391        Ok(written)
392    }
393
394    /// Reads exactly `count` bytes from the chain to the disk asynchronously or returns an error if
395    /// not enough data can be read.
396    pub async fn read_exact_to_at_fut<F: AsyncDisk + ?Sized>(
397        &mut self,
398        dst: &F,
399        mut count: usize,
400        mut off: u64,
401        options: IoOptions,
402    ) -> disk::Result<()> {
403        while count > 0 {
404            let nread = self.read_to_at_fut(dst, count, off, options).await?;
405            if nread == 0 {
406                return Err(disk::Error::ReadingData(io::Error::new(
407                    io::ErrorKind::UnexpectedEof,
408                    "failed to write whole buffer",
409                )));
410            }
411            count -= nread;
412            off += nread as u64;
413        }
414
415        Ok(())
416    }
417
418    /// Returns number of bytes available for reading.  May return an error if the combined
419    /// lengths of all the buffers in the DescriptorChain would cause an integer overflow.
420    pub fn available_bytes(&self) -> usize {
421        self.regions.available_bytes()
422    }
423
424    /// Returns number of bytes already read from the descriptor chain buffer.
425    pub fn bytes_read(&self) -> usize {
426        self.regions.bytes_consumed()
427    }
428
429    pub fn get_remaining_regions(&self) -> MemRegionIter {
430        self.regions.get_remaining_regions()
431    }
432
433    /// Returns a `&[VolatileSlice]` that represents all the remaining data in this `Reader`.
434    /// Calling this method does not actually consume any data from the `Reader` and callers should
435    /// call `consume` to advance the `Reader`.
436    pub fn get_remaining(&self) -> SmallVec<[VolatileSlice; 16]> {
437        self.regions.get_remaining(&self.mem)
438    }
439
440    /// Consumes `amt` bytes from the underlying descriptor chain. If `amt` is larger than the
441    /// remaining data left in this `Reader`, then all remaining data will be consumed.
442    pub fn consume(&mut self, amt: usize) {
443        self.regions.consume(amt)
444    }
445
446    /// Splits this `Reader` into two at the given offset in the `DescriptorChain` buffer. After the
447    /// split, `self` will be able to read up to `offset` bytes while the returned `Reader` can read
448    /// up to `available_bytes() - offset` bytes. If `offset > self.available_bytes()`, then the
449    /// returned `Reader` will not be able to read any bytes.
450    pub fn split_at(&mut self, offset: usize) -> Reader {
451        Reader {
452            mem: self.mem.clone(),
453            regions: self.regions.split_at(offset),
454        }
455    }
456}
457
458/// Copy up to `size` bytes from `src` into `dst`.
459///
460/// Returns the total number of bytes copied.
461///
462/// # Safety
463///
464/// The caller must ensure that it is safe to write `size` bytes of data into `dst`.
465///
466/// After the function returns, it is only safe to assume that the number of bytes indicated by the
467/// return value (which may be less than the requested `size`) have been initialized. Bytes beyond
468/// that point are not initialized by this function.
469unsafe fn copy_regions_to_mut_ptr(
470    mem: &GuestMemory,
471    src: MemRegionIter,
472    dst: *mut u8,
473    size: usize,
474) -> io::Result<usize> {
475    let mut copied = 0;
476    for src_region in src {
477        if copied >= size {
478            break;
479        }
480
481        let remaining = size - copied;
482        let count = cmp::min(remaining, src_region.len);
483
484        let vslice = mem
485            .get_slice_at_addr(GuestAddress(src_region.offset), count)
486            .map_err(|_e| io::Error::from(io::ErrorKind::InvalidData))?;
487
488        // SAFETY: `get_slice_at_addr()` verified that the region points to valid memory, and
489        // the `count` calculation ensures we will write at most `size` bytes into `dst`.
490        unsafe {
491            copy_nonoverlapping(vslice.as_ptr(), dst.add(copied), count);
492        }
493
494        copied += count;
495    }
496
497    Ok(copied)
498}
499
500impl io::Read for Reader {
501    fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
502        // SAFETY: We pass a valid pointer and size combination derived from `buf`.
503        let total = unsafe {
504            copy_regions_to_mut_ptr(
505                &self.mem,
506                self.regions.get_remaining_regions(),
507                buf.as_mut_ptr(),
508                buf.len(),
509            )?
510        };
511        self.regions.consume(total);
512        Ok(total)
513    }
514}
515
516/// Provides high-level interface over the sequence of memory regions
517/// defined by writable descriptors in the descriptor chain.
518///
519/// Note that virtio spec requires driver to place any device-writable
520/// descriptors after any device-readable descriptors (2.6.4.2 in Virtio Spec v1.1).
521/// Writer will start iterating the descriptors from the first writable one and will
522/// assume that all following descriptors are writable.
523pub struct Writer {
524    mem: GuestMemory,
525    regions: DescriptorChainRegions,
526}
527
528impl Writer {
529    /// Construct a new Writer wrapper over `writable_regions`.
530    pub fn new_from_regions(
531        mem: &GuestMemory,
532        writable_regions: SmallVec<[MemRegion; 2]>,
533    ) -> Writer {
534        Writer {
535            mem: mem.clone(),
536            regions: DescriptorChainRegions::new(writable_regions),
537        }
538    }
539
540    /// Writes an object to the descriptor chain buffer.
541    pub fn write_obj<T: Immutable + IntoBytes>(&mut self, val: T) -> io::Result<()> {
542        self.write_all(val.as_bytes())
543    }
544
545    /// Writes all objects produced by `iter` into the descriptor chain buffer. Unlike `consume`,
546    /// this doesn't require the values to be stored in an intermediate collection first. It also
547    /// allows callers to choose which elements in a collection to write, for example by using the
548    /// `filter` or `take` methods of the `Iterator` trait.
549    pub fn write_iter<T: Immutable + IntoBytes, I: Iterator<Item = T>>(
550        &mut self,
551        mut iter: I,
552    ) -> io::Result<()> {
553        iter.try_for_each(|v| self.write_obj(v))
554    }
555
556    /// Writes a collection of objects into the descriptor chain buffer.
557    pub fn consume<T: Immutable + IntoBytes, C: IntoIterator<Item = T>>(
558        &mut self,
559        vals: C,
560    ) -> io::Result<()> {
561        self.write_iter(vals.into_iter())
562    }
563
564    /// Returns number of bytes available for writing.  May return an error if the combined
565    /// lengths of all the buffers in the DescriptorChain would cause an overflow.
566    pub fn available_bytes(&self) -> usize {
567        self.regions.available_bytes()
568    }
569
570    /// Reads data into a volatile slice up to the minimum of the slice's length or the number of
571    /// bytes remaining. Returns the number of bytes read.
572    pub fn write_from_volatile_slice(&mut self, slice: VolatileSlice) -> usize {
573        let mut written = 0usize;
574        let mut src = slice;
575        for dst in self.get_remaining() {
576            src.copy_to_volatile_slice(dst);
577            let copied = std::cmp::min(src.size(), dst.size());
578            written += copied;
579            src = match src.offset(copied) {
580                Ok(v) => v,
581                Err(_) => break, // The slice is fully consumed
582            };
583        }
584        self.regions.consume(written);
585        written
586    }
587
588    /// Writes data to the descriptor chain buffer from a readable object.
589    /// Returns the number of bytes written to the descriptor chain buffer.
590    /// The number of bytes written can be less than `count` if
591    /// there isn't enough data in the descriptor chain buffer.
592    pub fn write_from<F: FileReadWriteVolatile>(
593        &mut self,
594        mut src: F,
595        count: usize,
596    ) -> io::Result<usize> {
597        let iovs = self.regions.get_remaining_with_count(&self.mem, count);
598        let read = src.read_vectored_volatile(&iovs[..])?;
599        self.regions.consume(read);
600        Ok(read)
601    }
602
603    /// Writes data to the descriptor chain buffer from a File at offset `off`.
604    /// Returns the number of bytes written to the descriptor chain buffer.
605    /// The number of bytes written can be less than `count` if
606    /// there isn't enough data in the descriptor chain buffer.
607    pub fn write_from_at<F: FileReadWriteAtVolatile>(
608        &mut self,
609        src: &F,
610        count: usize,
611        off: u64,
612    ) -> io::Result<usize> {
613        let iovs = self.regions.get_remaining_with_count(&self.mem, count);
614        let read = src.read_vectored_at_volatile(&iovs[..], off)?;
615        self.regions.consume(read);
616        Ok(read)
617    }
618
619    pub fn write_all_from<F: FileReadWriteVolatile>(
620        &mut self,
621        mut src: F,
622        mut count: usize,
623    ) -> io::Result<()> {
624        while count > 0 {
625            match self.write_from(&mut src, count) {
626                Ok(0) => {
627                    return Err(io::Error::new(
628                        io::ErrorKind::WriteZero,
629                        "failed to write whole buffer",
630                    ))
631                }
632                Ok(n) => count -= n,
633                Err(ref e) if e.kind() == io::ErrorKind::Interrupted => {}
634                Err(e) => return Err(e),
635            }
636        }
637
638        Ok(())
639    }
640
641    pub fn write_all_from_at<F: FileReadWriteAtVolatile>(
642        &mut self,
643        src: &F,
644        mut count: usize,
645        mut off: u64,
646    ) -> io::Result<()> {
647        while count > 0 {
648            match self.write_from_at(src, count, off) {
649                Ok(0) => {
650                    return Err(io::Error::new(
651                        io::ErrorKind::WriteZero,
652                        "failed to write whole buffer",
653                    ))
654                }
655                Ok(n) => {
656                    count -= n;
657                    off += n as u64;
658                }
659                Err(ref e) if e.kind() == io::ErrorKind::Interrupted => {}
660                Err(e) => return Err(e),
661            }
662        }
663        Ok(())
664    }
665    /// Writes data to the descriptor chain buffer from an `AsyncDisk` at offset `off`.
666    /// Returns the number of bytes written to the descriptor chain buffer.
667    /// The number of bytes written can be less than `count` if
668    /// there isn't enough data in the descriptor chain buffer.
669    pub async fn write_from_at_fut<F: AsyncDisk + ?Sized>(
670        &mut self,
671        src: &F,
672        count: usize,
673        off: u64,
674        options: IoOptions,
675    ) -> disk::Result<usize> {
676        let read = src
677            .read_to_mem(
678                off,
679                Arc::new(self.mem.clone()),
680                self.regions.get_remaining_regions_with_count(count),
681                options,
682            )
683            .await?;
684        self.regions.consume(read);
685        Ok(read)
686    }
687
688    pub async fn write_all_from_at_fut<F: AsyncDisk + ?Sized>(
689        &mut self,
690        src: &F,
691        mut count: usize,
692        mut off: u64,
693        options: IoOptions,
694    ) -> disk::Result<()> {
695        while count > 0 {
696            let nwritten = self.write_from_at_fut(src, count, off, options).await?;
697            if nwritten == 0 {
698                return Err(disk::Error::WritingData(io::Error::new(
699                    io::ErrorKind::UnexpectedEof,
700                    "failed to write whole buffer",
701                )));
702            }
703            count -= nwritten;
704            off += nwritten as u64;
705        }
706        Ok(())
707    }
708
709    /// Returns number of bytes already written to the descriptor chain buffer.
710    pub fn bytes_written(&self) -> usize {
711        self.regions.bytes_consumed()
712    }
713
714    pub fn get_remaining_regions(&self) -> MemRegionIter {
715        self.regions.get_remaining_regions()
716    }
717
718    /// Returns a `&[VolatileSlice]` that represents all the remaining data in this `Writer`.
719    /// Calling this method does not actually advance the current position of the `Writer` in the
720    /// buffer and callers should call `consume_bytes` to advance the `Writer`. Not calling
721    /// `consume_bytes` with the amount of data copied into the returned `VolatileSlice`s will
722    /// result in that that data being overwritten the next time data is written into the `Writer`.
723    pub fn get_remaining(&self) -> SmallVec<[VolatileSlice; 16]> {
724        self.regions.get_remaining(&self.mem)
725    }
726
727    /// Consumes `amt` bytes from the underlying descriptor chain. If `amt` is larger than the
728    /// remaining data left in this `Reader`, then all remaining data will be consumed.
729    pub fn consume_bytes(&mut self, amt: usize) {
730        self.regions.consume(amt)
731    }
732
733    /// Splits this `Writer` into two at the given offset in the `DescriptorChain` buffer. After the
734    /// split, `self` will be able to write up to `offset` bytes while the returned `Writer` can
735    /// write up to `available_bytes() - offset` bytes. If `offset > self.available_bytes()`, then
736    /// the returned `Writer` will not be able to write any bytes.
737    pub fn split_at(&mut self, offset: usize) -> Writer {
738        Writer {
739            mem: self.mem.clone(),
740            regions: self.regions.split_at(offset),
741        }
742    }
743}
744
745impl io::Write for Writer {
746    fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
747        let mut rem = buf;
748        let mut total = 0;
749        for b in self.regions.get_remaining(&self.mem) {
750            if rem.is_empty() {
751                break;
752            }
753
754            let count = cmp::min(rem.len(), b.size());
755            // SAFETY:
756            // Safe because we have already verified that `vs` points to valid memory.
757            unsafe {
758                copy_nonoverlapping(rem.as_ptr(), b.as_mut_ptr(), count);
759            }
760            rem = &rem[count..];
761            total += count;
762        }
763
764        self.regions.consume(total);
765        Ok(total)
766    }
767
768    fn flush(&mut self) -> io::Result<()> {
769        // Nothing to flush since the writes go straight into the buffer.
770        Ok(())
771    }
772}
773
774const VIRTQ_DESC_F_NEXT: u16 = 0x1;
775const VIRTQ_DESC_F_WRITE: u16 = 0x2;
776
777#[derive(Copy, Clone, PartialEq, Eq)]
778pub enum DescriptorType {
779    Readable,
780    Writable,
781}
782
783#[derive(Copy, Clone, Debug, FromBytes, Immutable, IntoBytes, KnownLayout)]
784#[repr(C)]
785struct virtq_desc {
786    addr: Le64,
787    len: Le32,
788    flags: Le16,
789    next: Le16,
790}
791
792/// Test utility function to create a descriptor chain in guest memory.
793pub fn create_descriptor_chain(
794    memory: &GuestMemory,
795    descriptor_array_addr: GuestAddress,
796    mut buffers_start_addr: GuestAddress,
797    descriptors: Vec<(DescriptorType, u32)>,
798    spaces_between_regions: u32,
799) -> anyhow::Result<DescriptorChain> {
800    let descriptors_len = descriptors.len();
801    for (index, (type_, size)) in descriptors.into_iter().enumerate() {
802        let mut flags = 0;
803        if let DescriptorType::Writable = type_ {
804            flags |= VIRTQ_DESC_F_WRITE;
805        }
806        if index + 1 < descriptors_len {
807            flags |= VIRTQ_DESC_F_NEXT;
808        }
809
810        let index = index as u16;
811        let desc = virtq_desc {
812            addr: buffers_start_addr.offset().into(),
813            len: size.into(),
814            flags: flags.into(),
815            next: (index + 1).into(),
816        };
817
818        let offset = size + spaces_between_regions;
819        buffers_start_addr = buffers_start_addr
820            .checked_add(offset as u64)
821            .context("Invalid buffers_start_addr)")?;
822
823        let _ = memory.write_obj_at_addr(
824            desc,
825            descriptor_array_addr
826                .checked_add(index as u64 * std::mem::size_of::<virtq_desc>() as u64)
827                .context("Invalid descriptor_array_addr")?,
828        );
829    }
830
831    let chain = SplitDescriptorChain::new(memory, descriptor_array_addr, 0x100, 0);
832    DescriptorChain::new(chain, memory, 0)
833}
834
835#[cfg(test)]
836mod tests {
837    use std::fs::File;
838    use std::io::Read;
839
840    use cros_async::Executor;
841    use tempfile::tempfile;
842    use tempfile::NamedTempFile;
843
844    use super::*;
845
846    #[test]
847    fn reader_test_simple_chain() {
848        use DescriptorType::*;
849
850        let memory_start_addr = GuestAddress(0x0);
851        let memory = GuestMemory::new(&[(memory_start_addr, 0x10000)]).unwrap();
852
853        let mut chain = create_descriptor_chain(
854            &memory,
855            GuestAddress(0x0),
856            GuestAddress(0x100),
857            vec![
858                (Readable, 8),
859                (Readable, 16),
860                (Readable, 18),
861                (Readable, 64),
862            ],
863            0,
864        )
865        .expect("create_descriptor_chain failed");
866        let reader = &mut chain.reader;
867        assert_eq!(reader.available_bytes(), 106);
868        assert_eq!(reader.bytes_read(), 0);
869
870        let mut buffer = [0u8; 64];
871        reader
872            .read_exact(&mut buffer)
873            .expect("read_exact should not fail here");
874
875        assert_eq!(reader.available_bytes(), 42);
876        assert_eq!(reader.bytes_read(), 64);
877
878        match reader.read(&mut buffer) {
879            Err(_) => panic!("read should not fail here"),
880            Ok(length) => assert_eq!(length, 42),
881        }
882
883        assert_eq!(reader.available_bytes(), 0);
884        assert_eq!(reader.bytes_read(), 106);
885    }
886
887    #[test]
888    fn writer_test_simple_chain() {
889        use DescriptorType::*;
890
891        let memory_start_addr = GuestAddress(0x0);
892        let memory = GuestMemory::new(&[(memory_start_addr, 0x10000)]).unwrap();
893
894        let mut chain = create_descriptor_chain(
895            &memory,
896            GuestAddress(0x0),
897            GuestAddress(0x100),
898            vec![
899                (Writable, 8),
900                (Writable, 16),
901                (Writable, 18),
902                (Writable, 64),
903            ],
904            0,
905        )
906        .expect("create_descriptor_chain failed");
907        let writer = &mut chain.writer;
908        assert_eq!(writer.available_bytes(), 106);
909        assert_eq!(writer.bytes_written(), 0);
910
911        let buffer = [0; 64];
912        writer
913            .write_all(&buffer)
914            .expect("write_all should not fail here");
915
916        assert_eq!(writer.available_bytes(), 42);
917        assert_eq!(writer.bytes_written(), 64);
918
919        match writer.write(&buffer) {
920            Err(_) => panic!("write should not fail here"),
921            Ok(length) => assert_eq!(length, 42),
922        }
923
924        assert_eq!(writer.available_bytes(), 0);
925        assert_eq!(writer.bytes_written(), 106);
926    }
927
928    #[test]
929    fn reader_test_incompatible_chain() {
930        use DescriptorType::*;
931
932        let memory_start_addr = GuestAddress(0x0);
933        let memory = GuestMemory::new(&[(memory_start_addr, 0x10000)]).unwrap();
934
935        let mut chain = create_descriptor_chain(
936            &memory,
937            GuestAddress(0x0),
938            GuestAddress(0x100),
939            vec![(Writable, 8)],
940            0,
941        )
942        .expect("create_descriptor_chain failed");
943        let reader = &mut chain.reader;
944        assert_eq!(reader.available_bytes(), 0);
945        assert_eq!(reader.bytes_read(), 0);
946
947        assert!(reader.read_obj::<u8>().is_err());
948
949        assert_eq!(reader.available_bytes(), 0);
950        assert_eq!(reader.bytes_read(), 0);
951    }
952
953    #[test]
954    fn writer_test_incompatible_chain() {
955        use DescriptorType::*;
956
957        let memory_start_addr = GuestAddress(0x0);
958        let memory = GuestMemory::new(&[(memory_start_addr, 0x10000)]).unwrap();
959
960        let mut chain = create_descriptor_chain(
961            &memory,
962            GuestAddress(0x0),
963            GuestAddress(0x100),
964            vec![(Readable, 8)],
965            0,
966        )
967        .expect("create_descriptor_chain failed");
968        let writer = &mut chain.writer;
969        assert_eq!(writer.available_bytes(), 0);
970        assert_eq!(writer.bytes_written(), 0);
971
972        assert!(writer.write_obj(0u8).is_err());
973
974        assert_eq!(writer.available_bytes(), 0);
975        assert_eq!(writer.bytes_written(), 0);
976    }
977
978    #[test]
979    fn reader_failing_io() {
980        use DescriptorType::*;
981
982        let memory_start_addr = GuestAddress(0x0);
983        let memory = GuestMemory::new(&[(memory_start_addr, 0x10000)]).unwrap();
984
985        let mut chain = create_descriptor_chain(
986            &memory,
987            GuestAddress(0x0),
988            GuestAddress(0x100),
989            vec![(Readable, 256), (Readable, 256)],
990            0,
991        )
992        .expect("create_descriptor_chain failed");
993
994        let reader = &mut chain.reader;
995
996        // Open a file in read-only mode so writes to it to trigger an I/O error.
997        let device_file = if cfg!(windows) { "NUL" } else { "/dev/zero" };
998        let mut ro_file = File::open(device_file).expect("failed to open device file");
999
1000        reader
1001            .read_exact_to(&mut ro_file, 512)
1002            .expect_err("successfully read more bytes than SharedMemory size");
1003
1004        // The write above should have failed entirely, so we end up not writing any bytes at all.
1005        assert_eq!(reader.available_bytes(), 512);
1006        assert_eq!(reader.bytes_read(), 0);
1007    }
1008
1009    #[test]
1010    fn writer_failing_io() {
1011        use DescriptorType::*;
1012
1013        let memory_start_addr = GuestAddress(0x0);
1014        let memory = GuestMemory::new(&[(memory_start_addr, 0x10000)]).unwrap();
1015
1016        let mut chain = create_descriptor_chain(
1017            &memory,
1018            GuestAddress(0x0),
1019            GuestAddress(0x100),
1020            vec![(Writable, 256), (Writable, 256)],
1021            0,
1022        )
1023        .expect("create_descriptor_chain failed");
1024
1025        let writer = &mut chain.writer;
1026
1027        let mut file = tempfile().unwrap();
1028
1029        file.set_len(384).unwrap();
1030
1031        writer
1032            .write_all_from(&mut file, 512)
1033            .expect_err("successfully wrote more bytes than in SharedMemory");
1034
1035        assert_eq!(writer.available_bytes(), 128);
1036        assert_eq!(writer.bytes_written(), 384);
1037    }
1038
1039    #[test]
1040    fn reader_writer_shared_chain() {
1041        use DescriptorType::*;
1042
1043        let memory_start_addr = GuestAddress(0x0);
1044        let memory = GuestMemory::new(&[(memory_start_addr, 0x10000)]).unwrap();
1045
1046        let mut chain = create_descriptor_chain(
1047            &memory,
1048            GuestAddress(0x0),
1049            GuestAddress(0x100),
1050            vec![
1051                (Readable, 16),
1052                (Readable, 16),
1053                (Readable, 96),
1054                (Writable, 64),
1055                (Writable, 1),
1056                (Writable, 3),
1057            ],
1058            0,
1059        )
1060        .expect("create_descriptor_chain failed");
1061        let reader = &mut chain.reader;
1062        let writer = &mut chain.writer;
1063
1064        assert_eq!(reader.bytes_read(), 0);
1065        assert_eq!(writer.bytes_written(), 0);
1066
1067        let mut buffer = Vec::with_capacity(200);
1068
1069        assert_eq!(
1070            reader
1071                .read_to_end(&mut buffer)
1072                .expect("read should not fail here"),
1073            128
1074        );
1075
1076        // The writable descriptors are only 68 bytes long.
1077        writer
1078            .write_all(&buffer[..68])
1079            .expect("write should not fail here");
1080
1081        assert_eq!(reader.available_bytes(), 0);
1082        assert_eq!(reader.bytes_read(), 128);
1083        assert_eq!(writer.available_bytes(), 0);
1084        assert_eq!(writer.bytes_written(), 68);
1085    }
1086
1087    #[test]
1088    fn reader_writer_shattered_object() {
1089        use DescriptorType::*;
1090
1091        let memory_start_addr = GuestAddress(0x0);
1092        let memory = GuestMemory::new(&[(memory_start_addr, 0x10000)]).unwrap();
1093
1094        let secret: Le32 = 0x12345678.into();
1095
1096        // Create a descriptor chain with memory regions that are properly separated.
1097        let mut chain_writer = create_descriptor_chain(
1098            &memory,
1099            GuestAddress(0x0),
1100            GuestAddress(0x100),
1101            vec![(Writable, 1), (Writable, 1), (Writable, 1), (Writable, 1)],
1102            123,
1103        )
1104        .expect("create_descriptor_chain failed");
1105        let writer = &mut chain_writer.writer;
1106        writer
1107            .write_obj(secret)
1108            .expect("write_obj should not fail here");
1109
1110        // Now create new descriptor chain pointing to the same memory and try to read it.
1111        let mut chain_reader = create_descriptor_chain(
1112            &memory,
1113            GuestAddress(0x0),
1114            GuestAddress(0x100),
1115            vec![(Readable, 1), (Readable, 1), (Readable, 1), (Readable, 1)],
1116            123,
1117        )
1118        .expect("create_descriptor_chain failed");
1119        let reader = &mut chain_reader.reader;
1120        match reader.read_obj::<Le32>() {
1121            Err(_) => panic!("read_obj should not fail here"),
1122            Ok(read_secret) => assert_eq!(read_secret, secret),
1123        }
1124    }
1125
1126    #[test]
1127    fn reader_unexpected_eof() {
1128        use DescriptorType::*;
1129
1130        let memory_start_addr = GuestAddress(0x0);
1131        let memory = GuestMemory::new(&[(memory_start_addr, 0x10000)]).unwrap();
1132
1133        let mut chain = create_descriptor_chain(
1134            &memory,
1135            GuestAddress(0x0),
1136            GuestAddress(0x100),
1137            vec![(Readable, 256), (Readable, 256)],
1138            0,
1139        )
1140        .expect("create_descriptor_chain failed");
1141
1142        let reader = &mut chain.reader;
1143
1144        let mut buf = vec![0; 1024];
1145
1146        assert_eq!(
1147            reader
1148                .read_exact(&mut buf[..])
1149                .expect_err("read more bytes than available")
1150                .kind(),
1151            io::ErrorKind::UnexpectedEof
1152        );
1153    }
1154
1155    #[test]
1156    fn split_border() {
1157        use DescriptorType::*;
1158
1159        let memory_start_addr = GuestAddress(0x0);
1160        let memory = GuestMemory::new(&[(memory_start_addr, 0x10000)]).unwrap();
1161
1162        let mut chain = create_descriptor_chain(
1163            &memory,
1164            GuestAddress(0x0),
1165            GuestAddress(0x100),
1166            vec![
1167                (Readable, 16),
1168                (Readable, 16),
1169                (Readable, 96),
1170                (Writable, 64),
1171                (Writable, 1),
1172                (Writable, 3),
1173            ],
1174            0,
1175        )
1176        .expect("create_descriptor_chain failed");
1177        let reader = &mut chain.reader;
1178
1179        let other = reader.split_at(32);
1180        assert_eq!(reader.available_bytes(), 32);
1181        assert_eq!(other.available_bytes(), 96);
1182    }
1183
1184    #[test]
1185    fn split_middle() {
1186        use DescriptorType::*;
1187
1188        let memory_start_addr = GuestAddress(0x0);
1189        let memory = GuestMemory::new(&[(memory_start_addr, 0x10000)]).unwrap();
1190
1191        let mut chain = create_descriptor_chain(
1192            &memory,
1193            GuestAddress(0x0),
1194            GuestAddress(0x100),
1195            vec![
1196                (Readable, 16),
1197                (Readable, 16),
1198                (Readable, 96),
1199                (Writable, 64),
1200                (Writable, 1),
1201                (Writable, 3),
1202            ],
1203            0,
1204        )
1205        .expect("create_descriptor_chain failed");
1206        let reader = &mut chain.reader;
1207
1208        let other = reader.split_at(24);
1209        assert_eq!(reader.available_bytes(), 24);
1210        assert_eq!(other.available_bytes(), 104);
1211    }
1212
1213    #[test]
1214    fn split_end() {
1215        use DescriptorType::*;
1216
1217        let memory_start_addr = GuestAddress(0x0);
1218        let memory = GuestMemory::new(&[(memory_start_addr, 0x10000)]).unwrap();
1219
1220        let mut chain = create_descriptor_chain(
1221            &memory,
1222            GuestAddress(0x0),
1223            GuestAddress(0x100),
1224            vec![
1225                (Readable, 16),
1226                (Readable, 16),
1227                (Readable, 96),
1228                (Writable, 64),
1229                (Writable, 1),
1230                (Writable, 3),
1231            ],
1232            0,
1233        )
1234        .expect("create_descriptor_chain failed");
1235        let reader = &mut chain.reader;
1236
1237        let other = reader.split_at(128);
1238        assert_eq!(reader.available_bytes(), 128);
1239        assert_eq!(other.available_bytes(), 0);
1240    }
1241
1242    #[test]
1243    fn split_beginning() {
1244        use DescriptorType::*;
1245
1246        let memory_start_addr = GuestAddress(0x0);
1247        let memory = GuestMemory::new(&[(memory_start_addr, 0x10000)]).unwrap();
1248
1249        let mut chain = create_descriptor_chain(
1250            &memory,
1251            GuestAddress(0x0),
1252            GuestAddress(0x100),
1253            vec![
1254                (Readable, 16),
1255                (Readable, 16),
1256                (Readable, 96),
1257                (Writable, 64),
1258                (Writable, 1),
1259                (Writable, 3),
1260            ],
1261            0,
1262        )
1263        .expect("create_descriptor_chain failed");
1264        let reader = &mut chain.reader;
1265
1266        let other = reader.split_at(0);
1267        assert_eq!(reader.available_bytes(), 0);
1268        assert_eq!(other.available_bytes(), 128);
1269    }
1270
1271    #[test]
1272    fn split_outofbounds() {
1273        use DescriptorType::*;
1274
1275        let memory_start_addr = GuestAddress(0x0);
1276        let memory = GuestMemory::new(&[(memory_start_addr, 0x10000)]).unwrap();
1277
1278        let mut chain = create_descriptor_chain(
1279            &memory,
1280            GuestAddress(0x0),
1281            GuestAddress(0x100),
1282            vec![
1283                (Readable, 16),
1284                (Readable, 16),
1285                (Readable, 96),
1286                (Writable, 64),
1287                (Writable, 1),
1288                (Writable, 3),
1289            ],
1290            0,
1291        )
1292        .expect("create_descriptor_chain failed");
1293        let reader = &mut chain.reader;
1294
1295        let other = reader.split_at(256);
1296        assert_eq!(
1297            other.available_bytes(),
1298            0,
1299            "Reader returned from out-of-bounds split still has available bytes"
1300        );
1301    }
1302
1303    #[test]
1304    fn read_full() {
1305        use DescriptorType::*;
1306
1307        let memory_start_addr = GuestAddress(0x0);
1308        let memory = GuestMemory::new(&[(memory_start_addr, 0x10000)]).unwrap();
1309
1310        let mut chain = create_descriptor_chain(
1311            &memory,
1312            GuestAddress(0x0),
1313            GuestAddress(0x100),
1314            vec![(Readable, 16), (Readable, 16), (Readable, 16)],
1315            0,
1316        )
1317        .expect("create_descriptor_chain failed");
1318        let reader = &mut chain.reader;
1319
1320        let mut buf = [0u8; 64];
1321        assert_eq!(
1322            reader.read(&mut buf[..]).expect("failed to read to buffer"),
1323            48
1324        );
1325    }
1326
1327    #[test]
1328    fn write_full() {
1329        use DescriptorType::*;
1330
1331        let memory_start_addr = GuestAddress(0x0);
1332        let memory = GuestMemory::new(&[(memory_start_addr, 0x10000)]).unwrap();
1333
1334        let mut chain = create_descriptor_chain(
1335            &memory,
1336            GuestAddress(0x0),
1337            GuestAddress(0x100),
1338            vec![(Writable, 16), (Writable, 16), (Writable, 16)],
1339            0,
1340        )
1341        .expect("create_descriptor_chain failed");
1342        let writer = &mut chain.writer;
1343
1344        let buf = [0xdeu8; 64];
1345        assert_eq!(
1346            writer.write(&buf[..]).expect("failed to write from buffer"),
1347            48
1348        );
1349    }
1350
1351    #[test]
1352    fn consume_collect() {
1353        use DescriptorType::*;
1354
1355        let memory_start_addr = GuestAddress(0x0);
1356        let memory = GuestMemory::new(&[(memory_start_addr, 0x10000)]).unwrap();
1357        let vs: Vec<Le64> = vec![
1358            0x0101010101010101.into(),
1359            0x0202020202020202.into(),
1360            0x0303030303030303.into(),
1361        ];
1362
1363        let mut write_chain = create_descriptor_chain(
1364            &memory,
1365            GuestAddress(0x0),
1366            GuestAddress(0x100),
1367            vec![(Writable, 24)],
1368            0,
1369        )
1370        .expect("create_descriptor_chain failed");
1371        let writer = &mut write_chain.writer;
1372        writer
1373            .consume(vs.clone())
1374            .expect("failed to consume() a vector");
1375
1376        let mut read_chain = create_descriptor_chain(
1377            &memory,
1378            GuestAddress(0x0),
1379            GuestAddress(0x100),
1380            vec![(Readable, 24)],
1381            0,
1382        )
1383        .expect("create_descriptor_chain failed");
1384        let reader = &mut read_chain.reader;
1385        let vs_read = reader
1386            .collect::<io::Result<Vec<Le64>>, _>()
1387            .expect("failed to collect() values");
1388        assert_eq!(vs, vs_read);
1389    }
1390
1391    #[test]
1392    fn get_remaining_region_with_count() {
1393        use DescriptorType::*;
1394
1395        let memory_start_addr = GuestAddress(0x0);
1396        let memory = GuestMemory::new(&[(memory_start_addr, 0x10000)]).unwrap();
1397
1398        let chain = create_descriptor_chain(
1399            &memory,
1400            GuestAddress(0x0),
1401            GuestAddress(0x100),
1402            vec![
1403                (Readable, 16),
1404                (Readable, 16),
1405                (Readable, 96),
1406                (Writable, 64),
1407                (Writable, 1),
1408                (Writable, 3),
1409            ],
1410            0,
1411        )
1412        .expect("create_descriptor_chain failed");
1413
1414        let Reader {
1415            mem: _,
1416            mut regions,
1417        } = chain.reader;
1418
1419        let drain = regions
1420            .get_remaining_regions_with_count(usize::MAX)
1421            .fold(0usize, |total, region| total + region.len);
1422        assert_eq!(drain, 128);
1423
1424        let exact = regions
1425            .get_remaining_regions_with_count(32)
1426            .fold(0usize, |total, region| total + region.len);
1427        assert!(exact > 0);
1428        assert!(exact <= 32);
1429
1430        let split = regions
1431            .get_remaining_regions_with_count(24)
1432            .fold(0usize, |total, region| total + region.len);
1433        assert!(split > 0);
1434        assert!(split <= 24);
1435
1436        regions.consume(64);
1437
1438        let first = regions
1439            .get_remaining_regions_with_count(8)
1440            .fold(0usize, |total, region| total + region.len);
1441        assert!(first > 0);
1442        assert!(first <= 8);
1443    }
1444
1445    #[test]
1446    fn get_remaining_with_count() {
1447        use DescriptorType::*;
1448
1449        let memory_start_addr = GuestAddress(0x0);
1450        let memory = GuestMemory::new(&[(memory_start_addr, 0x10000)]).unwrap();
1451
1452        let chain = create_descriptor_chain(
1453            &memory,
1454            GuestAddress(0x0),
1455            GuestAddress(0x100),
1456            vec![
1457                (Readable, 16),
1458                (Readable, 16),
1459                (Readable, 96),
1460                (Writable, 64),
1461                (Writable, 1),
1462                (Writable, 3),
1463            ],
1464            0,
1465        )
1466        .expect("create_descriptor_chain failed");
1467        let Reader {
1468            mem: _,
1469            mut regions,
1470        } = chain.reader;
1471
1472        let drain = regions
1473            .get_remaining_with_count(&memory, usize::MAX)
1474            .iter()
1475            .fold(0usize, |total, iov| total + iov.size());
1476        assert_eq!(drain, 128);
1477
1478        let exact = regions
1479            .get_remaining_with_count(&memory, 32)
1480            .iter()
1481            .fold(0usize, |total, iov| total + iov.size());
1482        assert!(exact > 0);
1483        assert!(exact <= 32);
1484
1485        let split = regions
1486            .get_remaining_with_count(&memory, 24)
1487            .iter()
1488            .fold(0usize, |total, iov| total + iov.size());
1489        assert!(split > 0);
1490        assert!(split <= 24);
1491
1492        regions.consume(64);
1493
1494        let first = regions
1495            .get_remaining_with_count(&memory, 8)
1496            .iter()
1497            .fold(0usize, |total, iov| total + iov.size());
1498        assert!(first > 0);
1499        assert!(first <= 8);
1500    }
1501
1502    #[test]
1503    fn reader_peek_obj() {
1504        use DescriptorType::*;
1505
1506        let memory_start_addr = GuestAddress(0x0);
1507        let memory = GuestMemory::new(&[(memory_start_addr, 0x10000)]).unwrap();
1508
1509        // Write test data to memory.
1510        memory
1511            .write_obj_at_addr(Le16::from(0xBEEF), GuestAddress(0x100))
1512            .unwrap();
1513        memory
1514            .write_obj_at_addr(Le16::from(0xDEAD), GuestAddress(0x200))
1515            .unwrap();
1516
1517        let mut chain_reader = create_descriptor_chain(
1518            &memory,
1519            GuestAddress(0x0),
1520            GuestAddress(0x100),
1521            vec![(Readable, 2), (Readable, 2)],
1522            0x100 - 2,
1523        )
1524        .expect("create_descriptor_chain failed");
1525        let reader = &mut chain_reader.reader;
1526
1527        // peek_obj() at the beginning of the chain should return the first object.
1528        let peek1 = reader.peek_obj::<Le16>().unwrap();
1529        assert_eq!(peek1, Le16::from(0xBEEF));
1530
1531        // peek_obj() again should return the same object, since it was not consumed.
1532        let peek2 = reader.peek_obj::<Le16>().unwrap();
1533        assert_eq!(peek2, Le16::from(0xBEEF));
1534
1535        // peek_obj() of an object spanning two descriptors should copy from both.
1536        let peek3 = reader.peek_obj::<Le32>().unwrap();
1537        assert_eq!(peek3, Le32::from(0xDEADBEEF));
1538
1539        // read_obj() should return the first object.
1540        let read1 = reader.read_obj::<Le16>().unwrap();
1541        assert_eq!(read1, Le16::from(0xBEEF));
1542
1543        // peek_obj() of a value that is larger than the rest of the chain should fail.
1544        reader
1545            .peek_obj::<Le32>()
1546            .expect_err("peek_obj past end of chain");
1547
1548        // read_obj() again should return the second object.
1549        let read2 = reader.read_obj::<Le16>().unwrap();
1550        assert_eq!(read2, Le16::from(0xDEAD));
1551
1552        // peek_obj() should fail at the end of the chain.
1553        reader
1554            .peek_obj::<Le16>()
1555            .expect_err("peek_obj past end of chain");
1556    }
1557
1558    #[test]
1559    fn region_reader_failing_io() {
1560        let ex = Executor::new().unwrap();
1561        ex.run_until(region_reader_failing_io_async(&ex)).unwrap();
1562    }
1563    async fn region_reader_failing_io_async(ex: &Executor) {
1564        use DescriptorType::*;
1565
1566        let memory_start_addr = GuestAddress(0x0);
1567        let memory = GuestMemory::new(&[(memory_start_addr, 0x10000)]).unwrap();
1568
1569        let mut chain = create_descriptor_chain(
1570            &memory,
1571            GuestAddress(0x0),
1572            GuestAddress(0x100),
1573            vec![(Readable, 256), (Readable, 256)],
1574            0,
1575        )
1576        .expect("create_descriptor_chain failed");
1577
1578        let reader = &mut chain.reader;
1579
1580        // Open a file in read-only mode so writes to it to trigger an I/O error.
1581        let named_temp_file = NamedTempFile::new().expect("failed to create temp file");
1582        let ro_file =
1583            File::open(named_temp_file.path()).expect("failed to open temp file read only");
1584        let async_ro_file = disk::SingleFileDisk::new(ro_file, ex).expect("Failed to crate SFD");
1585
1586        reader
1587            .read_exact_to_at_fut(&async_ro_file, 512, 0, Default::default())
1588            .await
1589            .expect_err("successfully read more bytes than SingleFileDisk size");
1590
1591        // The write above should have failed entirely, so we end up not writing any bytes at all.
1592        assert_eq!(reader.available_bytes(), 512);
1593        assert_eq!(reader.bytes_read(), 0);
1594    }
1595
1596    #[test]
1597    fn region_writer_failing_io() {
1598        let ex = Executor::new().unwrap();
1599        ex.run_until(region_writer_failing_io_async(&ex)).unwrap()
1600    }
1601    async fn region_writer_failing_io_async(ex: &Executor) {
1602        use DescriptorType::*;
1603
1604        let memory_start_addr = GuestAddress(0x0);
1605        let memory = GuestMemory::new(&[(memory_start_addr, 0x10000)]).unwrap();
1606
1607        let mut chain = create_descriptor_chain(
1608            &memory,
1609            GuestAddress(0x0),
1610            GuestAddress(0x100),
1611            vec![(Writable, 256), (Writable, 256)],
1612            0,
1613        )
1614        .expect("create_descriptor_chain failed");
1615
1616        let writer = &mut chain.writer;
1617
1618        let file = tempfile().expect("failed to create temp file");
1619
1620        file.set_len(384).unwrap();
1621        let async_file = disk::SingleFileDisk::new(file, ex).expect("Failed to crate SFD");
1622
1623        writer
1624            .write_all_from_at_fut(&async_file, 512, 0, Default::default())
1625            .await
1626            .expect_err("successfully wrote more bytes than in SingleFileDisk");
1627
1628        assert_eq!(writer.available_bytes(), 128);
1629        assert_eq!(writer.bytes_written(), 384);
1630    }
1631}