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