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