1use 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 current_region_index: usize,
42
43 current_region_offset: usize,
45
46 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 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 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 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 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 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 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 self.current_region_offset += count;
128 self.bytes_consumed += count;
129 return;
130 }
131
132 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
169pub struct Reader {
177 mem: GuestMemory,
178 regions: DescriptorChainRegions,
179}
180
181pub 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 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 pub fn peek_obj<T: FromBytes>(&self) -> io::Result<T> {
213 let mut obj = MaybeUninit::uninit();
214
215 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 Ok(unsafe { obj.assume_init() })
231 }
232
233 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 pub fn collect<C: FromIterator<io::Result<T>>, T: FromBytes>(&mut self) -> C {
244 self.iter().collect()
245 }
246
247 pub fn iter<T: FromBytes>(&mut self) -> ReaderIterator<T> {
251 ReaderIterator {
252 reader: self,
253 phantom: PhantomData,
254 }
255 }
256
257 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, };
270 }
271 self.regions.consume(read);
272 read
273 }
274
275 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 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 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 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 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 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 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 pub fn available_bytes(&self) -> usize {
421 self.regions.available_bytes()
422 }
423
424 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 pub fn get_remaining(&self) -> SmallVec<[VolatileSlice; 16]> {
437 self.regions.get_remaining(&self.mem)
438 }
439
440 pub fn consume(&mut self, amt: usize) {
443 self.regions.consume(amt)
444 }
445
446 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
458unsafe 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 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 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
516pub struct Writer {
524 mem: GuestMemory,
525 regions: DescriptorChainRegions,
526}
527
528impl Writer {
529 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 pub fn write_obj<T: Immutable + IntoBytes>(&mut self, val: T) -> io::Result<()> {
542 self.write_all(val.as_bytes())
543 }
544
545 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 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 pub fn available_bytes(&self) -> usize {
567 self.regions.available_bytes()
568 }
569
570 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, };
583 }
584 self.regions.consume(written);
585 written
586 }
587
588 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 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 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 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 pub fn get_remaining(&self) -> SmallVec<[VolatileSlice; 16]> {
724 self.regions.get_remaining(&self.mem)
725 }
726
727 pub fn consume_bytes(&mut self, amt: usize) {
730 self.regions.consume(amt)
731 }
732
733 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 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 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
792pub 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 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 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 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 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 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 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 let peek1 = reader.peek_obj::<Le16>().unwrap();
1529 assert_eq!(peek1, Le16::from(0xBEEF));
1530
1531 let peek2 = reader.peek_obj::<Le16>().unwrap();
1533 assert_eq!(peek2, Le16::from(0xBEEF));
1534
1535 let peek3 = reader.peek_obj::<Le32>().unwrap();
1537 assert_eq!(peek3, Le32::from(0xDEADBEEF));
1538
1539 let read1 = reader.read_obj::<Le16>().unwrap();
1541 assert_eq!(read1, Le16::from(0xBEEF));
1542
1543 reader
1545 .peek_obj::<Le32>()
1546 .expect_err("peek_obj past end of chain");
1547
1548 let read2 = reader.read_obj::<Le16>().unwrap();
1550 assert_eq!(read2, Le16::from(0xDEAD));
1551
1552 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 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 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}