ruda_runtime/runtime/memory_management/memory_pool/
handle.rs1use crate::runtime::memory_management::MemoryHandle;
2use alloc::sync::Arc;
3use core::sync::atomic::{AtomicUsize, Ordering};
4#[cfg(target_has_atomic = "64")]
5use core::sync::atomic::AtomicU64;
6#[cfg(not(target_has_atomic = "64"))]
7use spin::Mutex;
8
9#[derive(Debug)]
11pub struct ManagedMemoryHandle {
12 descriptor: Arc<ManagedMemoryDescriptor>,
13 handle_count: Arc<()>,
15}
16
17#[derive(Debug)]
19pub struct ManagedMemoryBinding {
20 descriptor: Arc<ManagedMemoryDescriptor>,
21}
22
23impl Clone for ManagedMemoryHandle {
24 fn clone(&self) -> Self {
25 Self {
26 descriptor: self.descriptor.clone(),
27 handle_count: self.handle_count.clone(),
28 }
29 }
30}
31
32pub(crate) struct ManagedMemoryDescriptor {
38 pub(crate) id: ManagedMemoryId,
39 pins: AtomicUsize,
40 pin_revision: AtomicUsize,
41 #[cfg(target_has_atomic = "64")]
42 location: AtomicU64,
43 #[cfg(not(target_has_atomic = "64"))]
44 location: Mutex<MemoryLocation>,
45}
46
47impl core::fmt::Debug for ManagedMemoryDescriptor {
48 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
49 f.debug_struct("ManagedMemoryDescriptor")
50 .field("id", &self.id)
51 .field("location", &self.location())
52 .finish()
53 }
54}
55
56#[derive(Debug, PartialEq, Eq, Clone, Copy, Hash)]
57pub struct ManagedMemoryId {
59 pub(crate) value: usize,
60}
61
62impl PartialEq for ManagedMemoryDescriptor {
63 fn eq(&self, other: &Self) -> bool {
64 self.id == other.id
65 }
66}
67
68impl Eq for ManagedMemoryDescriptor {}
69
70#[derive(Clone, Copy, Debug)]
71pub(crate) struct MemoryLocation {
73 pub pool: u8,
75 pub page: u16,
77 pub slice: u32,
79 pub init: u8,
81}
82
83impl ManagedMemoryDescriptor {
84 pub(crate) fn pin_revision(&self) -> usize {
85 let revision = self.pin_revision.load(Ordering::Acquire);
86 if revision == 0 {
87 self.pin_revision.compare_exchange(0, 1, Ordering::AcqRel, Ordering::Acquire)
88 .map(|_| 1).unwrap_or_else(|revision| revision)
89 } else {
90 revision
91 }
92 }
93
94 fn pin_changed(&self) {
95 let _ = self.pin_revision.fetch_update(Ordering::AcqRel, Ordering::Acquire, |revision| {
96 (revision != 0 && revision != usize::MAX).then(|| revision + 1)
97 });
98 }
99
100 pub(crate) fn update_location(&self, location: MemoryLocation) {
102 #[cfg(target_has_atomic = "64")]
103 {
104 self.location.store(location.to_bits(), Ordering::Release);
105 }
106 #[cfg(not(target_has_atomic = "64"))]
107 {
108 *self.location.lock() = location;
109 }
110 }
111
112 pub(crate) fn update_slice(&self, slice: u32) {
114 self.modify(|location| MemoryLocation { slice, ..location });
115 }
116
117 pub fn update_page(&self, page: u16) {
119 self.modify(|location| MemoryLocation { page, ..location });
120 }
121
122 pub(crate) fn location(&self) -> MemoryLocation {
124 #[cfg(target_has_atomic = "64")]
125 {
126 MemoryLocation::from_bits(self.location.load(Ordering::Acquire))
127 }
128 #[cfg(not(target_has_atomic = "64"))]
129 {
130 *self.location.lock()
131 }
132 }
133
134 pub(crate) fn slice(&self) -> usize {
135 self.location().slice as usize
136 }
137
138 pub(crate) fn page(&self) -> usize {
139 self.location().page as usize
140 }
141
142 fn modify(&self, update: impl Fn(MemoryLocation) -> MemoryLocation) {
143 #[cfg(target_has_atomic = "64")]
144 {
145 let _ = self.location.fetch_update(Ordering::AcqRel, Ordering::Acquire, |bits| {
146 Some(update(MemoryLocation::from_bits(bits)).to_bits())
147 });
148 }
149 #[cfg(not(target_has_atomic = "64"))]
150 {
151 let mut location = self.location.lock();
152 *location = update(*location);
153 }
154 }
155}
156
157impl MemoryLocation {
158 #[cfg(target_has_atomic = "64")]
159 fn to_bits(self) -> u64 {
160 self.pool as u64
161 | (self.page as u64) << 8
162 | (self.slice as u64) << 24
163 | (self.init as u64) << 56
164 }
165
166 #[cfg(target_has_atomic = "64")]
167 fn from_bits(bits: u64) -> Self {
168 Self {
169 pool: bits as u8,
170 page: (bits >> 8) as u16,
171 slice: (bits >> 24) as u32,
172 init: (bits >> 56) as u8,
173 }
174 }
175
176 pub(crate) fn new(pool: u8, page: u16, slice: u32) -> Self {
178 Self {
179 pool,
180 page,
181 slice,
182 init: 1,
183 }
184 }
185
186 pub(crate) fn uninit() -> Self {
188 Self {
189 pool: 0,
190 page: 0,
191 slice: 0,
192 init: 0,
193 }
194 }
195}
196
197impl ManagedMemoryHandle {
198 pub fn new() -> Self {
200 let value = Self::gen_id();
201
202 Self {
203 descriptor: Arc::new(ManagedMemoryDescriptor {
204 id: ManagedMemoryId { value },
205 pins: AtomicUsize::new(0),
206 pin_revision: AtomicUsize::new(0),
207 #[cfg(target_has_atomic = "64")]
208 location: AtomicU64::new(MemoryLocation::uninit().to_bits()),
209 #[cfg(not(target_has_atomic = "64"))]
210 location: Mutex::new(MemoryLocation::uninit()),
211 }),
212 handle_count: Arc::new(()),
213 }
214 }
215
216 pub(crate) fn descriptor(&self) -> &ManagedMemoryDescriptor {
218 &self.descriptor
219 }
220
221 pub fn can_mut(&self) -> bool {
223 Arc::strong_count(&self.handle_count) <= 2
224 }
225
226 pub fn is_free(&self) -> bool {
228 Arc::strong_count(&self.descriptor) <= 1
229 }
230
231 pub(crate) fn is_pinned(&self) -> bool {
232 self.descriptor.pins.load(core::sync::atomic::Ordering::Acquire) != 0
233 }
234
235 pub fn binding(self) -> ManagedMemoryBinding {
237 ManagedMemoryBinding {
238 descriptor: self.descriptor.clone(),
239 }
240 }
241
242 fn gen_id() -> usize {
243 static COUNTER: core::sync::atomic::AtomicUsize = core::sync::atomic::AtomicUsize::new(0);
244 let value = COUNTER.fetch_add(1, core::sync::atomic::Ordering::Relaxed);
245 if value == usize::MAX {
246 core::panic!("Memory ID overflowed");
247 }
248 value
249 }
250}
251
252impl ManagedMemoryBinding {
253 pub fn pin(&self) -> MemoryResourcePin {
255 self.descriptor.pins.fetch_add(1, core::sync::atomic::Ordering::AcqRel);
256 self.descriptor.pin_changed();
257 MemoryResourcePin { binding: self.clone() }
258 }
259 pub fn id(&self) -> ManagedMemoryId { self.descriptor.id }
262
263 pub(crate) fn descriptor(&self) -> &ManagedMemoryDescriptor {
265 &self.descriptor
266 }
267}
268
269#[derive(Debug)]
271pub struct MemoryResourcePin {
272 binding: ManagedMemoryBinding,
273}
274
275impl Clone for MemoryResourcePin {
276 fn clone(&self) -> Self { self.binding.pin() }
277}
278
279impl Drop for MemoryResourcePin {
280 fn drop(&mut self) {
281 self.binding.descriptor.pins.fetch_sub(1, core::sync::atomic::Ordering::AcqRel);
282 self.binding.descriptor.pin_changed();
283 }
284}
285
286impl Default for ManagedMemoryHandle {
287 fn default() -> Self {
288 Self::new()
289 }
290}
291
292impl Clone for ManagedMemoryBinding {
293 fn clone(&self) -> Self {
294 Self {
295 descriptor: self.descriptor.clone(),
296 }
297 }
298}
299
300impl MemoryHandle<ManagedMemoryBinding> for ManagedMemoryHandle {
301 fn can_mut(&self) -> bool {
302 self.can_mut()
303 }
304
305 fn binding(self) -> ManagedMemoryBinding {
306 self.binding()
307 }
308}
309
310pub fn optimal_align(shape: usize, elem_size: usize, buffer_align: usize) -> usize {
314 if shape == 1 {
315 elem_size
316 } else {
317 (shape * elem_size)
318 .next_power_of_two()
319 .clamp(16, buffer_align)
320 }
321}
322
323#[cfg(test)]
324mod tests {
325 use super::*;
326
327 #[test]
328 fn test_memory_id_mutability() {
329 let handle1 = ManagedMemoryHandle::new();
330 handle1.descriptor().update_slice(4);
331 assert_eq!(handle1.descriptor().slice(), 4);
332
333 let handle2 = ManagedMemoryHandle::new();
334 handle2
335 .clone()
336 .descriptor()
337 .update_location(handle1.descriptor().location());
338 assert_eq!(handle2.descriptor().slice(), 4);
339 }
340
341 #[test]
342 fn test_location_visible_through_shared_arc() {
343 let handle = ManagedMemoryHandle::new();
344 let handle2 = handle.clone();
345
346 let location = MemoryLocation::new(1, 2, 3);
347 handle.descriptor().update_location(location);
348
349 assert_eq!(handle2.descriptor().location().pool, 1);
350 assert_eq!(handle2.descriptor().location().page, 2);
351 assert_eq!(handle2.descriptor().location().slice, 3);
352 assert_eq!(handle2.descriptor().location().init, 1);
353
354 handle.descriptor().update_slice(42);
355 assert_eq!(handle2.descriptor().slice(), 42);
356 }
357
358 #[test]
359 #[cfg(feature = "runtime-std")]
360 fn concurrent_debug_reads_consistent_location_snapshots() {
361 let handle = ManagedMemoryHandle::new();
362 let writer = handle.clone();
363 let task = std::thread::spawn(move || {
364 for i in 1..=128_u32 {
365 writer.descriptor().update_location(MemoryLocation::new(i as u8, i as u16, i));
366 std::thread::yield_now();
367 }
368 });
369 for _ in 0..128 {
370 let _ = alloc::format!("{handle:?}");
371 let location = handle.descriptor().location();
372 assert_eq!(location.page as u32, location.slice);
373 assert_eq!(location.pool as u32, location.slice);
374 std::thread::yield_now();
375 }
376 task.join().unwrap();
377 assert_eq!(handle.descriptor().slice(), 128);
378 }
379
380 #[test]
381 #[cfg(feature = "runtime-std")]
382 fn concurrent_field_updates_do_not_overwrite_each_other() {
383 let handle = ManagedMemoryHandle::new();
384 let first = handle.clone();
385 let second = handle.clone();
386 let page_writer = std::thread::spawn(move || {
387 for page in 1..=128 { first.descriptor().update_page(page); }
388 });
389 let slice_writer = std::thread::spawn(move || {
390 for slice in 1..=128 { second.descriptor().update_slice(slice); }
391 });
392 page_writer.join().unwrap();
393 slice_writer.join().unwrap();
394 assert_eq!(handle.descriptor().page(), 128);
395 assert_eq!(handle.descriptor().slice(), 128);
396 }
397
398}