Skip to main content

virtio_accel_device/
regions.rs

1//! Segmented byte ports over transport-neutral descriptor metadata.
2//!
3//! Each segmented port access scans at most the segment list and copies only the requested range;
4//! constructors reject empty, zero-length, and overflowing segment collections.
5
6use core::cmp::min;
7
8use virtio_accel_core::{BackendError, ByteSink, ByteSource};
9pub use virtio_accel_transport::{
10    ChainLayout, ChainLayoutError, ChainRegion, RegionDirection, validate_chain_layout,
11};
12
13#[derive(Clone, Copy, Debug, PartialEq, Eq)]
14pub enum SegmentedRegionError {
15    Empty,
16    ZeroLength,
17    LengthOverflow,
18}
19
20/// Borrowed segmented source used by unit tests, fuzz targets, and simple transport adapters.
21#[derive(Debug)]
22pub struct SegmentedSource<'segments, 'bytes> {
23    segments: &'segments [&'bytes [u8]],
24    len: u64,
25}
26
27impl<'segments, 'bytes> SegmentedSource<'segments, 'bytes> {
28    pub fn new(segments: &'segments [&'bytes [u8]]) -> Result<Self, SegmentedRegionError> {
29        let len = checked_segment_len(segments.iter().map(|segment| segment.len()))?;
30        Ok(Self { segments, len })
31    }
32}
33
34impl ByteSource for SegmentedSource<'_, '_> {
35    fn len(&self) -> u64 {
36        self.len
37    }
38
39    fn read_at(&self, offset: u64, target: &mut [u8]) -> Result<(), BackendError> {
40        checked_range(offset, target.len(), self.len)?;
41        if target.is_empty() {
42            return Ok(());
43        }
44
45        let mut skip = offset;
46        let mut written = 0;
47        for segment in self.segments {
48            let segment_len = segment.len() as u64;
49            if skip >= segment_len {
50                skip -= segment_len;
51                continue;
52            }
53
54            let start = skip as usize;
55            let count = min(segment.len() - start, target.len() - written);
56            target[written..written + count].copy_from_slice(&segment[start..start + count]);
57            written += count;
58            skip = 0;
59            if written == target.len() {
60                return Ok(());
61            }
62        }
63
64        Err(BackendError::OutOfBounds)
65    }
66
67    fn as_contiguous(&self) -> Option<&[u8]> {
68        (self.segments.len() == 1).then_some(self.segments[0])
69    }
70}
71
72/// Borrowed segmented sink used by unit tests, fuzz targets, and simple transport adapters.
73#[derive(Debug)]
74pub struct SegmentedSink<'segments, 'bytes> {
75    segments: &'segments mut [&'bytes mut [u8]],
76    len: u64,
77}
78
79impl<'segments, 'bytes> SegmentedSink<'segments, 'bytes> {
80    pub fn new(segments: &'segments mut [&'bytes mut [u8]]) -> Result<Self, SegmentedRegionError> {
81        let len = checked_segment_len(segments.iter().map(|segment| segment.len()))?;
82        Ok(Self { segments, len })
83    }
84}
85
86impl ByteSink for SegmentedSink<'_, '_> {
87    fn len(&self) -> u64 {
88        self.len
89    }
90
91    fn write_at(&mut self, offset: u64, source: &[u8]) -> Result<(), BackendError> {
92        checked_range(offset, source.len(), self.len)?;
93        if source.is_empty() {
94            return Ok(());
95        }
96
97        let mut skip = offset;
98        let mut read = 0;
99        for segment in self.segments.iter_mut() {
100            let segment = &mut **segment;
101            let segment_len = segment.len() as u64;
102            if skip >= segment_len {
103                skip -= segment_len;
104                continue;
105            }
106
107            let start = skip as usize;
108            let count = min(segment.len() - start, source.len() - read);
109            segment[start..start + count].copy_from_slice(&source[read..read + count]);
110            read += count;
111            skip = 0;
112            if read == source.len() {
113                return Ok(());
114            }
115        }
116
117        Err(BackendError::OutOfBounds)
118    }
119
120    fn as_contiguous_mut(&mut self) -> Option<&mut [u8]> {
121        if self.segments.len() == 1 {
122            Some(&mut *self.segments[0])
123        } else {
124            None
125        }
126    }
127}
128
129#[derive(Clone, Copy, Debug)]
130pub struct ReadableRegion<'a> {
131    source: &'a dyn ByteSource,
132    offset: u64,
133    len: u64,
134}
135
136impl<'a> ReadableRegion<'a> {
137    pub fn new(source: &'a dyn ByteSource, offset: u64, len: u64) -> Result<Self, BackendError> {
138        checked_range_u64(offset, len, source.len())?;
139        Ok(Self {
140            source,
141            offset,
142            len,
143        })
144    }
145}
146
147impl ByteSource for ReadableRegion<'_> {
148    fn len(&self) -> u64 {
149        self.len
150    }
151
152    fn read_at(&self, offset: u64, target: &mut [u8]) -> Result<(), BackendError> {
153        checked_range(offset, target.len(), self.len)?;
154        let source_offset = self
155            .offset
156            .checked_add(offset)
157            .ok_or(BackendError::OutOfBounds)?;
158        self.source.read_at(source_offset, target)
159    }
160
161    fn as_contiguous(&self) -> Option<&[u8]> {
162        let start = usize::try_from(self.offset).ok()?;
163        let len = usize::try_from(self.len).ok()?;
164        let end = start.checked_add(len)?;
165        self.source.as_contiguous()?.get(start..end)
166    }
167}
168
169#[derive(Debug)]
170pub struct WritableRegion<'a> {
171    sink: &'a mut dyn ByteSink,
172    offset: u64,
173    len: u64,
174}
175
176impl<'a> WritableRegion<'a> {
177    pub fn new(sink: &'a mut dyn ByteSink, offset: u64, len: u64) -> Result<Self, BackendError> {
178        checked_range_u64(offset, len, sink.len())?;
179        Ok(Self { sink, offset, len })
180    }
181}
182
183impl ByteSink for WritableRegion<'_> {
184    fn len(&self) -> u64 {
185        self.len
186    }
187
188    fn write_at(&mut self, offset: u64, source: &[u8]) -> Result<(), BackendError> {
189        checked_range(offset, source.len(), self.len)?;
190        let sink_offset = self
191            .offset
192            .checked_add(offset)
193            .ok_or(BackendError::OutOfBounds)?;
194        self.sink.write_at(sink_offset, source)
195    }
196
197    fn as_contiguous_mut(&mut self) -> Option<&mut [u8]> {
198        let start = usize::try_from(self.offset).ok()?;
199        let len = usize::try_from(self.len).ok()?;
200        let end = start.checked_add(len)?;
201        self.sink.as_contiguous_mut()?.get_mut(start..end)
202    }
203}
204
205fn checked_segment_len(
206    lengths: impl IntoIterator<Item = usize>,
207) -> Result<u64, SegmentedRegionError> {
208    let mut count = 0_usize;
209    let mut total = 0_u64;
210    for len in lengths {
211        count += 1;
212        if len == 0 {
213            return Err(SegmentedRegionError::ZeroLength);
214        }
215        total = total
216            .checked_add(len as u64)
217            .ok_or(SegmentedRegionError::LengthOverflow)?;
218    }
219    if count == 0 {
220        return Err(SegmentedRegionError::Empty);
221    }
222    Ok(total)
223}
224
225fn checked_range(offset: u64, bytes: usize, len: u64) -> Result<(), BackendError> {
226    let bytes = u64::try_from(bytes).map_err(|_| BackendError::OutOfBounds)?;
227    checked_range_u64(offset, bytes, len)
228}
229
230fn checked_range_u64(offset: u64, bytes: u64, len: u64) -> Result<(), BackendError> {
231    let end = offset.checked_add(bytes).ok_or(BackendError::OutOfBounds)?;
232    if end > len {
233        return Err(BackendError::OutOfBounds);
234    }
235    Ok(())
236}
237
238#[cfg(test)]
239mod tests {
240    use super::*;
241
242    #[test]
243    fn segmented_ports_cross_every_boundary() {
244        let bytes = *b"segmented";
245        for split in 1..bytes.as_slice().len() {
246            let source_segments = [&bytes[..split], &bytes[split..]];
247            let source = SegmentedSource::new(&source_segments).unwrap();
248            let mut decoded = [0_u8; 9];
249            source.read_at(0, &mut decoded).unwrap();
250            assert_eq!(decoded, bytes);
251
252            let mut first = [0_u8; 9];
253            let (left, right) = first.split_at_mut(split);
254            let mut sink_segments: [&mut [u8]; 2] = [left, right];
255            let mut sink = SegmentedSink::new(&mut sink_segments).unwrap();
256            sink.write_at(0, &bytes).unwrap();
257            assert_eq!(first, bytes);
258        }
259    }
260
261    #[test]
262    fn subregions_preserve_bounds_and_contiguous_fast_paths() {
263        let bytes = *b"01234567";
264        let region = ReadableRegion::new(&bytes, 2, 4).unwrap();
265        assert_eq!(region.as_contiguous(), Some(&b"2345"[..]));
266
267        let mut output = [0_u8; 8];
268        {
269            let mut region = WritableRegion::new(&mut output, 2, 4).unwrap();
270            assert_eq!(region.as_contiguous_mut().unwrap().len(), 4);
271            region.write_at(0, b"abcd").unwrap();
272        }
273        assert_eq!(&output, b"\0\0abcd\0\0");
274    }
275}