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