1use 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#[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#[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}