1use crate::runtime::server::IoError;
2
3use super::{ComputeStorage, StorageHandle, StorageId, StorageUtilization};
4use alloc::{
5 alloc::{Layout, alloc_zeroed, dealloc},
6 sync::Arc,
7 vec::Vec,
8};
9use core::{
10 fmt,
11 ops::{Deref, DerefMut, Range},
12 ptr::NonNull,
13};
14use hashbrown::HashMap;
15use ruda_core::backtrace::BackTrace;
16use spin::Mutex;
17
18#[derive(Default)]
22pub struct BytesStorage {
23 memory: HashMap<StorageId, Arc<AllocatedBytes>>,
24}
25
26impl fmt::Debug for BytesStorage {
27 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
28 f.write_str("BytesStorage")
29 }
30}
31
32#[derive(Clone, Debug)]
35pub struct BytesResource {
36 allocation: Arc<AllocatedBytes>,
37 range: Range<usize>,
38}
39
40#[derive(Clone, Debug, PartialEq, Eq)]
42pub enum BytesAccessError {
43 UnknownStorage(StorageId),
45 InvalidRange {
47 offset: u64,
49 size: u64,
51 allocation_size: usize,
53 },
54 BorrowConflict {
56 requested: Range<usize>,
58 existing: Range<usize>,
60 },
61}
62
63impl fmt::Display for BytesAccessError {
64 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
65 match self {
66 Self::UnknownStorage(id) => write!(f, "unknown or released storage {id}"),
67 Self::InvalidRange { offset, size, allocation_size } => write!(
68 f, "invalid storage range: offset={offset}, size={size}, allocation={allocation_size}",
69 ),
70 Self::BorrowConflict { requested, existing } => write!(
71 f, "storage range {requested:?} conflicts with live borrow {existing:?}",
72 ),
73 }
74 }
75}
76
77impl core::error::Error for BytesAccessError {}
78
79#[derive(Clone, Debug, PartialEq, Eq)]
80struct BorrowRegion {
81 range: Range<usize>,
82 writable: bool,
83}
84
85struct AllocatedBytes {
86 ptr: NonNull<u8>,
87 layout: Layout,
88 borrows: Mutex<Vec<BorrowRegion>>,
89}
90
91impl fmt::Debug for AllocatedBytes {
92 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
93 f.debug_struct("AllocatedBytes")
95 .field("size", &self.layout.size())
96 .field("alignment", &self.layout.align())
97 .finish_non_exhaustive()
98 }
99}
100
101unsafe impl Send for AllocatedBytes {}
106unsafe impl Sync for AllocatedBytes {}
108
109impl Drop for AllocatedBytes {
110 fn drop(&mut self) {
111 if self.layout.size() != 0 {
112 unsafe { dealloc(self.ptr.as_ptr(), self.layout) };
115 }
116 }
117}
118
119#[derive(Debug)]
122struct BorrowLease {
123 allocation: Arc<AllocatedBytes>,
124 region: BorrowRegion,
125}
126
127impl BorrowLease {
128 fn acquire(resource: &BytesResource, writable: bool) -> Result<Self, BytesAccessError> {
129 let region = BorrowRegion { range: resource.range.clone(), writable };
130 let mut borrows = resource.allocation.borrows.lock();
131 for existing in borrows.iter() {
132 let overlap = !region.range.is_empty()
133 && !existing.range.is_empty()
134 && region.range.start < existing.range.end
135 && existing.range.start < region.range.end;
136 if overlap && (region.writable || existing.writable) {
137 return Err(BytesAccessError::BorrowConflict {
138 requested: region.range,
139 existing: existing.range.clone(),
140 });
141 }
142 }
143 borrows.push(region.clone());
144 drop(borrows);
145 Ok(Self { allocation: resource.allocation.clone(), region })
146 }
147
148 fn ptr(&self) -> *mut u8 {
149 unsafe { self.allocation.ptr.as_ptr().add(self.region.range.start) }
152 }
153
154 fn len(&self) -> usize {
155 self.region.range.end - self.region.range.start
156 }
157
158 fn as_slice(&self) -> &[u8] {
159 unsafe { core::slice::from_raw_parts(self.ptr(), self.len()) }
162 }
163}
164
165impl Drop for BorrowLease {
166 fn drop(&mut self) {
167 let mut borrows = self.allocation.borrows.lock();
168 let index = borrows.iter().position(|region| region == &self.region)
171 .expect("live byte lease must have a registered range");
172 borrows.swap_remove(index);
173 }
174}
175
176#[derive(Debug)]
178pub struct BytesReadGuard {
179 lease: BorrowLease,
180}
181
182impl Deref for BytesReadGuard {
183 type Target = [u8];
184
185 fn deref(&self) -> &[u8] {
186 self.lease.as_slice()
187 }
188}
189
190impl AsRef<[u8]> for BytesReadGuard {
191 fn as_ref(&self) -> &[u8] { self }
192}
193
194#[derive(Debug)]
197pub struct BytesWriteGuard {
198 lease: BorrowLease,
199}
200
201impl Deref for BytesWriteGuard {
202 type Target = [u8];
203
204 fn deref(&self) -> &[u8] { self.lease.as_slice() }
205}
206
207impl DerefMut for BytesWriteGuard {
208 fn deref_mut(&mut self) -> &mut [u8] {
209 unsafe { core::slice::from_raw_parts_mut(self.lease.ptr(), self.lease.len()) }
212 }
213}
214
215impl AsRef<[u8]> for BytesWriteGuard {
216 fn as_ref(&self) -> &[u8] { self }
217}
218
219impl AsMut<[u8]> for BytesWriteGuard {
220 fn as_mut(&mut self) -> &mut [u8] { self }
221}
222
223impl BytesResource {
224 pub fn get_write_ptr_and_length(&self) -> (*mut u8, usize) {
231 let ptr = unsafe { self.allocation.ptr.as_ptr().add(self.range.start) };
233 (ptr, self.range.end - self.range.start)
234 }
235
236 pub fn try_write(&self) -> Result<BytesWriteGuard, BytesAccessError> {
238 Ok(BytesWriteGuard { lease: BorrowLease::acquire(self, true)? })
239 }
240
241 #[track_caller]
244 pub fn write(&self) -> BytesWriteGuard {
245 self.try_write().expect("conflicting byte-storage write")
246 }
247
248 pub fn try_read(&self) -> Result<BytesReadGuard, BytesAccessError> {
250 Ok(BytesReadGuard { lease: BorrowLease::acquire(self, false)? })
251 }
252
253 #[track_caller]
256 pub fn read(&self) -> BytesReadGuard {
257 self.try_read().expect("conflicting byte-storage read")
258 }
259}
260
261impl BytesStorage {
262 pub fn try_get(&self, handle: &StorageHandle) -> Result<BytesResource, BytesAccessError> {
264 let allocation = self.memory.get(&handle.id)
265 .ok_or(BytesAccessError::UnknownStorage(handle.id))?;
266 let invalid = || BytesAccessError::InvalidRange {
267 offset: handle.offset(), size: handle.size(), allocation_size: allocation.layout.size(),
268 };
269 let end = handle.offset().checked_add(handle.size()).ok_or_else(invalid)?;
270 let start = usize::try_from(handle.offset()).map_err(|_| invalid())?;
271 let end = usize::try_from(end).map_err(|_| invalid())?;
272 if end > allocation.layout.size() {
273 return Err(invalid());
274 }
275 Ok(BytesResource { allocation: allocation.clone(), range: start..end })
276 }
277}
278
279impl ComputeStorage for BytesStorage {
280 type Resource = BytesResource;
281
282 fn alignment(&self) -> usize { 4 }
283
284 fn get(&mut self, handle: &StorageHandle) -> Self::Resource {
285 self.try_get(handle).expect("invalid byte-storage handle")
286 }
287
288 #[cfg_attr(feature = "runtime-tracing", tracing::instrument(level = "trace", skip(self, size)))]
289 fn alloc(&mut self, size: u64) -> Result<StorageHandle, IoError> {
290 let too_big = || IoError::BufferTooBig { size, backtrace: BackTrace::capture() };
291 let size_usize = usize::try_from(size).map_err(|_| too_big())?;
292 let layout = Layout::from_size_align(size_usize, self.alignment())
294 .map_err(|_| too_big())?;
295 let id = StorageId::new();
296 let ptr = if size_usize == 0 {
297 NonNull::<u32>::dangling().cast::<u8>()
298 } else {
299 NonNull::new(unsafe { alloc_zeroed(layout) }).ok_or_else(too_big)?
301 };
302 self.memory.insert(id, Arc::new(AllocatedBytes {
303 ptr, layout, borrows: Mutex::new(Vec::new()),
304 }));
305 Ok(StorageHandle { id, utilization: StorageUtilization { offset: 0, size } })
306 }
307
308 #[cfg_attr(feature = "runtime-tracing", tracing::instrument(level = "trace", skip(self)))]
309 fn dealloc(&mut self, id: StorageId) {
310 self.memory.remove(&id);
312 }
313
314 fn flush(&mut self) {}
315}
316
317#[cfg(test)]
318mod tests {
319 use super::*;
320
321 #[test_log::test]
322 fn test_can_alloc_and_dealloc() {
323 let mut storage = BytesStorage::default();
324 let handle_1 = storage.alloc(64).unwrap();
325
326 assert_eq!(handle_1.size(), 64);
327 storage.dealloc(handle_1.id);
328 }
329
330 #[test_log::test]
331 fn test_slices() {
332 let mut storage = BytesStorage::default();
333 let handle_1 = storage.alloc(64).unwrap();
334 let handle_2 = StorageHandle::new(
335 handle_1.id,
336 StorageUtilization {
337 offset: 24,
338 size: 8,
339 },
340 );
341
342 storage
343 .get(&handle_1)
344 .write()
345 .iter_mut()
346 .enumerate()
347 .for_each(|(i, b)| {
348 *b = i as u8;
349 });
350
351 let bytes = storage.get(&handle_2).read().to_vec();
352
353 storage.dealloc(handle_1.id);
354 assert_eq!(bytes, &[24, 25, 26, 27, 28, 29, 30, 31]);
355 }
356
357 #[test_log::test]
359 fn test_read_after_alloc_without_write() {
360 let mut storage = BytesStorage::default();
361 let handle = storage.alloc(16).unwrap();
362 let resource = storage.get(&handle);
363 assert!(resource.read().iter().all(|&b| b == 0));
364 storage.dealloc(handle.id);
365 }
366
367 #[test_log::test]
369 fn test_zero_size_alloc_and_dealloc() {
370 let mut storage = BytesStorage::default();
371 let handle = storage.alloc(0).unwrap();
372 assert_eq!(handle.size(), 0);
373 storage.dealloc(handle.id);
374 }
375
376 #[test_log::test]
377 fn test_alloc_dealloc_realloc() {
378 let mut storage = BytesStorage::default();
379 let h1 = storage.alloc(32).unwrap();
380 storage.get(&h1).write()[0] = 0xAA;
381 storage.dealloc(h1.id);
382 let h2 = storage.alloc(32).unwrap();
383 storage.dealloc(h2.id);
384 }
385
386 #[test_log::test]
387 fn test_multiple_non_overlapping_regions() {
388 let mut storage = BytesStorage::default();
389 let base = storage.alloc(64).unwrap();
390
391 let regions: alloc::vec::Vec<_> = (0..4)
392 .map(|i| {
393 StorageHandle::new(
394 base.id,
395 StorageUtilization {
396 offset: i * 16,
397 size: 16,
398 },
399 )
400 })
401 .collect();
402
403 for (i, region) in regions.iter().enumerate() {
404 storage.get(region).write().fill(i as u8);
405 }
406 for (i, region) in regions.iter().enumerate() {
407 assert!(storage.get(region).read().iter().all(|&b| b == i as u8));
408 }
409 storage.dealloc(base.id);
410 }
411
412 #[test]
413 fn guard_outlives_resource_and_storage() {
414 let mut storage = BytesStorage::default();
415 let handle = storage.alloc(4).unwrap();
416 let weak = Arc::downgrade(storage.memory.get(&handle.id).unwrap());
417 let mut guard = storage.get(&handle).write();
418 storage.dealloc(handle.id);
419 assert!(matches!(storage.try_get(&handle), Err(BytesAccessError::UnknownStorage(_))));
420 drop(storage);
421 guard.copy_from_slice(&[10, 20, 30, 40]);
422 assert_eq!(&guard[..], &[10, 20, 30, 40]);
423 assert!(weak.upgrade().is_some());
424 drop(guard);
425 assert!(weak.upgrade().is_none());
426 }
427
428 #[test]
429 fn dropping_storage_frees_unborrowed_allocations() {
430 let mut storage = BytesStorage::default();
431 let handle = storage.alloc(8).unwrap();
432 let weak = Arc::downgrade(storage.memory.get(&handle.id).unwrap());
433 drop(storage);
434 assert!(weak.upgrade().is_none());
435 }
436
437 #[test]
438 fn cloned_resources_share_borrow_registry() {
439 let mut storage = BytesStorage::default();
440 let handle = storage.alloc(8).unwrap();
441 let first = storage.get(&handle);
442 let second = first.clone();
443 let mut writer = first.write();
444 writer[0] = 42;
445 assert!(second.try_read().is_err());
446 assert!(second.try_write().is_err());
447 drop(writer);
448 assert_eq!(second.read()[0], 42);
449 let a = first.read();
450 let b = second.read();
451 assert!(first.try_write().is_err());
452 drop(a);
453 assert!(first.try_write().is_err());
454 drop(b);
455 assert!(first.try_write().is_ok());
456 }
457
458 #[test]
459 fn independently_looked_up_overlapping_ranges_conflict() {
460 let mut storage = BytesStorage::default();
461 let base = storage.alloc(16).unwrap();
462 let left = storage.get(&StorageHandle::new(
463 base.id, StorageUtilization { offset: 0, size: 8 },
464 ));
465 let overlap = storage.get(&StorageHandle::new(
466 base.id, StorageUtilization { offset: 4, size: 8 },
467 ));
468 let _writer = left.write();
469 assert!(overlap.try_read().is_err());
470 assert!(overlap.try_write().is_err());
471 }
472
473 #[test]
474 fn disjoint_pooled_ranges_can_be_borrowed_together() {
475 let mut storage = BytesStorage::default();
476 let base = storage.alloc(16).unwrap();
477 let left = storage.get(&StorageHandle::new(
478 base.id, StorageUtilization { offset: 0, size: 8 },
479 ));
480 let right = storage.get(&StorageHandle::new(
481 base.id, StorageUtilization { offset: 8, size: 8 },
482 ));
483 let mut a = left.write();
484 let mut b = right.write();
485 a.fill(1);
486 b.fill(2);
487 assert_eq!(&a[..], &[1; 8]);
488 assert_eq!(&b[..], &[2; 8]);
489 drop((a, b));
490 let bytes = storage.get(&base).read();
491 assert_eq!(&bytes[..8], &[1; 8]);
492 assert_eq!(&bytes[8..], &[2; 8]);
493 }
494
495 #[test]
496 fn forged_ranges_are_rejected_before_pointer_arithmetic() {
497 let mut storage = BytesStorage::default();
498 let base = storage.alloc(8).unwrap();
499 for (offset, size) in [(9, 0), (7, 2), (0, 9), (u64::MAX, 2)] {
500 let bad = StorageHandle::new(base.id, StorageUtilization { offset, size });
501 assert!(matches!(storage.try_get(&bad), Err(BytesAccessError::InvalidRange { .. })));
502 }
503 let end = StorageHandle::new(base.id, StorageUtilization { offset: 8, size: 0 });
504 assert!(storage.try_get(&end).unwrap().read().is_empty());
505 }
506
507 #[test]
508 fn empty_ranges_do_not_conflict_with_live_nonempty_ranges() {
509 let mut storage = BytesStorage::default();
510 let base = storage.alloc(8).unwrap();
511 let _writer = storage.get(&base).write();
512 let empty = storage.get(&StorageHandle::new(
513 base.id, StorageUtilization { offset: 4, size: 0 },
514 ));
515 let a = empty.write();
516 let b = empty.write();
517 assert!(a.is_empty() && b.is_empty());
518 }
519
520 #[test]
521 fn zero_allocations_have_usable_empty_guards_and_correct_alignment() {
522 let mut storage = BytesStorage::default();
523 for size in [0, 1, 17] {
524 let handle = storage.alloc(size).unwrap();
525 let resource = storage.get(&handle);
526 let (ptr, len) = resource.get_write_ptr_and_length();
527 assert_eq!(ptr as usize % storage.alignment(), 0);
528 assert_eq!(len, size as usize);
529 assert_eq!(resource.read().len(), len);
530 storage.dealloc(handle.id);
531 }
532 assert!(storage.alloc(u64::MAX).is_err());
533 }
534
535 #[test]
536 #[cfg(feature = "runtime-std")]
537 fn active_write_blocks_cross_thread_access() {
538 let mut storage = BytesStorage::default();
539 let handle = storage.alloc(8).unwrap();
540 let first = storage.get(&handle);
541 let second = first.clone();
542 let writer = first.write();
543 std::thread::spawn(move || {
544 assert!(second.try_read().is_err());
545 assert!(second.try_write().is_err());
546 }).join().unwrap();
547 drop(writer);
548 assert!(first.try_read().is_ok());
549 }
550
551}