ruda_runtime/runtime/memory_management/memory_pool/
handle.rs1use crate::runtime::memory_management::MemoryHandle;
2use alloc::sync::Arc;
3use spin::Mutex;
4
5#[derive(Debug)]
7pub struct ManagedMemoryHandle {
8 descriptor: Arc<ManagedMemoryDescriptor>,
9 handle_count: Arc<()>,
11}
12
13#[derive(Debug)]
15pub struct ManagedMemoryBinding {
16 descriptor: Arc<ManagedMemoryDescriptor>,
17}
18
19impl Clone for ManagedMemoryHandle {
20 fn clone(&self) -> Self {
21 Self {
22 descriptor: self.descriptor.clone(),
23 handle_count: self.handle_count.clone(),
24 }
25 }
26}
27
28pub(crate) struct ManagedMemoryDescriptor {
34 pub(crate) id: ManagedMemoryId,
35 location: Mutex<MemoryLocation>,
36}
37
38impl core::fmt::Debug for ManagedMemoryDescriptor {
39 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
40 f.debug_struct("ManagedMemoryDescriptor")
41 .field("id", &self.id)
42 .field("location", &self.location())
43 .finish()
44 }
45}
46
47#[derive(Debug, PartialEq, Eq, Clone, Copy, Hash)]
48pub struct ManagedMemoryId {
50 pub(crate) value: usize,
51}
52
53impl PartialEq for ManagedMemoryDescriptor {
54 fn eq(&self, other: &Self) -> bool {
55 self.id == other.id
56 }
57}
58
59impl Eq for ManagedMemoryDescriptor {}
60
61#[derive(Clone, Copy, Debug)]
62pub(crate) struct MemoryLocation {
64 pub pool: u8,
66 pub page: u16,
68 pub slice: u32,
70 pub init: u8,
72}
73
74impl ManagedMemoryDescriptor {
75 pub(crate) fn update_location(&self, location: MemoryLocation) {
77 *self.location.lock() = location;
78 }
79
80 pub(crate) fn update_slice(&self, slice: u32) {
82 self.location.lock().slice = slice;
83 }
84
85 pub fn update_page(&self, page: u16) {
87 self.location.lock().page = page;
88 }
89
90 pub(crate) fn location(&self) -> MemoryLocation {
92 *self.location.lock()
93 }
94
95 pub(crate) fn slice(&self) -> usize {
96 self.location.lock().slice as usize
97 }
98
99 pub(crate) fn page(&self) -> usize {
100 self.location.lock().page as usize
101 }
102}
103
104impl MemoryLocation {
105 pub(crate) fn new(pool: u8, page: u16, slice: u32) -> Self {
107 Self {
108 pool,
109 page,
110 slice,
111 init: 1,
112 }
113 }
114
115 pub(crate) fn uninit() -> Self {
117 Self {
118 pool: 0,
119 page: 0,
120 slice: 0,
121 init: 0,
122 }
123 }
124}
125
126impl ManagedMemoryHandle {
127 pub fn new() -> Self {
129 let value = Self::gen_id();
130
131 Self {
132 descriptor: Arc::new(ManagedMemoryDescriptor {
133 id: ManagedMemoryId { value },
134 location: Mutex::new(MemoryLocation::uninit()),
135 }),
136 handle_count: Arc::new(()),
137 }
138 }
139
140 pub(crate) fn descriptor(&self) -> &ManagedMemoryDescriptor {
142 &self.descriptor
143 }
144
145 pub fn can_mut(&self) -> bool {
147 Arc::strong_count(&self.handle_count) <= 2
148 }
149
150 pub fn is_free(&self) -> bool {
152 Arc::strong_count(&self.descriptor) <= 1
153 }
154
155 pub fn binding(self) -> ManagedMemoryBinding {
157 ManagedMemoryBinding {
158 descriptor: self.descriptor.clone(),
159 }
160 }
161
162 fn gen_id() -> usize {
163 static COUNTER: core::sync::atomic::AtomicUsize = core::sync::atomic::AtomicUsize::new(0);
164 let value = COUNTER.fetch_add(1, core::sync::atomic::Ordering::Relaxed);
165 if value == usize::MAX {
166 core::panic!("Memory ID overflowed");
167 }
168 value
169 }
170}
171
172impl ManagedMemoryBinding {
173 pub fn id(&self) -> ManagedMemoryId { self.descriptor.id }
176
177 pub(crate) fn descriptor(&self) -> &ManagedMemoryDescriptor {
179 &self.descriptor
180 }
181}
182
183impl Default for ManagedMemoryHandle {
184 fn default() -> Self {
185 Self::new()
186 }
187}
188
189impl Clone for ManagedMemoryBinding {
190 fn clone(&self) -> Self {
191 Self {
192 descriptor: self.descriptor.clone(),
193 }
194 }
195}
196
197impl MemoryHandle<ManagedMemoryBinding> for ManagedMemoryHandle {
198 fn can_mut(&self) -> bool {
199 self.can_mut()
200 }
201
202 fn binding(self) -> ManagedMemoryBinding {
203 self.binding()
204 }
205}
206
207pub fn optimal_align(shape: usize, elem_size: usize, buffer_align: usize) -> usize {
211 if shape == 1 {
212 elem_size
213 } else {
214 (shape * elem_size)
215 .next_power_of_two()
216 .clamp(16, buffer_align)
217 }
218}
219
220#[cfg(test)]
221mod tests {
222 use super::*;
223
224 #[test]
225 fn test_memory_id_mutability() {
226 let handle1 = ManagedMemoryHandle::new();
227 handle1.descriptor().update_slice(4);
228 assert_eq!(handle1.descriptor().slice(), 4);
229
230 let handle2 = ManagedMemoryHandle::new();
231 handle2
232 .clone()
233 .descriptor()
234 .update_location(handle1.descriptor().location());
235 assert_eq!(handle2.descriptor().slice(), 4);
236 }
237
238 #[test]
239 fn test_location_visible_through_shared_arc() {
240 let handle = ManagedMemoryHandle::new();
241 let handle2 = handle.clone();
242
243 let location = MemoryLocation::new(1, 2, 3);
244 handle.descriptor().update_location(location);
245
246 assert_eq!(handle2.descriptor().location().pool, 1);
247 assert_eq!(handle2.descriptor().location().page, 2);
248 assert_eq!(handle2.descriptor().location().slice, 3);
249 assert_eq!(handle2.descriptor().location().init, 1);
250
251 handle.descriptor().update_slice(42);
252 assert_eq!(handle2.descriptor().slice(), 42);
253 }
254
255 #[test]
256 #[cfg(feature = "runtime-std")]
257 fn concurrent_debug_reads_consistent_location_snapshots() {
258 let handle = ManagedMemoryHandle::new();
259 let writer = handle.clone();
260 let task = std::thread::spawn(move || {
261 for i in 1..=128_u32 {
262 writer.descriptor().update_location(MemoryLocation::new(i as u8, i as u16, i));
263 std::thread::yield_now();
264 }
265 });
266 for _ in 0..128 {
267 let _ = alloc::format!("{handle:?}");
268 let location = handle.descriptor().location();
269 assert_eq!(location.page as u32, location.slice);
270 assert_eq!(location.pool as u32, location.slice);
271 std::thread::yield_now();
272 }
273 task.join().unwrap();
274 assert_eq!(handle.descriptor().slice(), 128);
275 }
276
277 #[test]
278 #[cfg(feature = "runtime-std")]
279 fn concurrent_field_updates_do_not_overwrite_each_other() {
280 let handle = ManagedMemoryHandle::new();
281 let first = handle.clone();
282 let second = handle.clone();
283 let page_writer = std::thread::spawn(move || {
284 for page in 1..=128 { first.descriptor().update_page(page); }
285 });
286 let slice_writer = std::thread::spawn(move || {
287 for slice in 1..=128 { second.descriptor().update_slice(slice); }
288 });
289 page_writer.join().unwrap();
290 slice_writer.join().unwrap();
291 assert_eq!(handle.descriptor().page(), 128);
292 assert_eq!(handle.descriptor().slice(), 128);
293 }
294
295}