1use crate::config::STABLE_PAGE_SIZE;
7pub use crate::stable::memory_layout::MemoryId;
8use crate::stable::memory_layout::{
9 bucket_allocations_address, write_growing, BucketCache, BucketId, VirtualSegment,
10 BUCKETS_OFFSET_IN_BYTES, BUCKETS_OFFSET_IN_PAGES, BUCKET_SIZE_IN_PAGES, HEADER_RESERVED_BYTES,
11 HEADER_SIZE, LAYOUT_VERSION, MAGIC, MAX_NUM_BUCKETS, MAX_NUM_MEMORIES,
12 UNALLOCATED_BUCKET_MARKER,
13};
14use crate::stable::memory_manager_validation::{load_validated_layout, try_load_validated_layout};
15use crate::stable::raw_memory::{Memory, MemoryBackendIdentity};
16use std::cell::RefCell;
17use std::rc::Rc;
18
19#[derive(Clone, Copy, Debug, Eq, PartialEq)]
20pub(crate) struct MemoryIdentity {
21 backend: MemoryBackendIdentity,
22 id: MemoryId,
23}
24
25#[derive(Clone)]
26pub struct MemoryManager<M: Memory> {
27 inner: Rc<RefCell<MemoryManagerInner<M>>>,
28}
29
30#[derive(Debug, thiserror::Error)]
31pub enum MemoryManagerInitError {
32 #[error("bucket size must be greater than zero")]
33 BucketSizeIsZero,
34 #[error("non-empty memory does not contain a MemoryManager layout")]
35 NonMemoryManagerLayout,
36 #[error("{0}")]
37 InvalidLayout(String),
38}
39
40impl<M: Memory> MemoryManager<M> {
41 pub fn init(memory: M) -> Self {
42 Self::init_with_bucket_size(memory, BUCKET_SIZE_IN_PAGES as u16)
43 }
44
45 pub fn init_strict(memory: M) -> Result<Self, MemoryManagerInitError> {
46 Self::init_strict_with_bucket_size(memory, BUCKET_SIZE_IN_PAGES as u16)
47 }
48
49 pub fn init_with_bucket_size(memory: M, bucket_size_in_pages: u16) -> Self {
50 if bucket_size_in_pages == 0 {
51 panic!("bucket size must be greater than zero");
52 }
53 Self {
54 inner: Rc::new(RefCell::new(MemoryManagerInner::init(
55 memory,
56 bucket_size_in_pages,
57 ))),
58 }
59 }
60
61 pub fn init_strict_with_bucket_size(
62 memory: M,
63 bucket_size_in_pages: u16,
64 ) -> Result<Self, MemoryManagerInitError> {
65 if bucket_size_in_pages == 0 {
66 return Err(MemoryManagerInitError::BucketSizeIsZero);
67 }
68 Ok(Self {
69 inner: Rc::new(RefCell::new(MemoryManagerInner::init_strict(
70 memory,
71 bucket_size_in_pages,
72 )?)),
73 })
74 }
75
76 pub fn get(&self, id: MemoryId) -> VirtualMemory<M> {
77 VirtualMemory {
78 id,
79 memory_manager: Rc::clone(&self.inner),
80 cache: BucketCache::new(),
81 }
82 }
83}
84#[derive(Clone)]
85pub struct VirtualMemory<M: Memory> {
86 id: MemoryId,
87 memory_manager: Rc<RefCell<MemoryManagerInner<M>>>,
88 cache: BucketCache,
89}
90impl<M: Memory> VirtualMemory<M> {
91 pub(crate) fn identity(&self) -> MemoryIdentity {
92 let inner = self.memory_manager.borrow();
93 MemoryIdentity {
94 backend: inner.memory.identity(),
95 id: self.id,
96 }
97 }
98}
99
100impl<M: Memory> Memory for VirtualMemory<M> {
101 fn size(&self) -> u64 {
102 self.memory_manager.borrow().memory_size(self.id)
103 }
104
105 fn grow(&self, pages: u64) -> i64 {
106 self.memory_manager.borrow_mut().grow(self.id, pages)
107 }
108
109 fn read(&self, offset: u64, dst: &mut [u8]) {
110 self.memory_manager
111 .borrow()
112 .read(self.id, offset, dst, &self.cache);
113 }
114
115 unsafe fn read_unsafe(&self, offset: u64, dst: *mut u8, count: usize) {
116 self.memory_manager
117 .borrow()
118 .read_unsafe(self.id, offset, dst, count, &self.cache);
119 }
120
121 fn write(&self, offset: u64, src: &[u8]) {
122 self.memory_manager
123 .borrow()
124 .write(self.id, offset, src, &self.cache);
125 }
126}
127
128#[derive(Clone)]
129struct MemoryManagerInner<M: Memory> {
130 memory: M,
131 allocated_buckets: u16,
132 bucket_size_in_pages: u16,
133 memory_sizes_in_pages: [u64; MAX_NUM_MEMORIES as usize],
134 memory_buckets: Vec<Vec<BucketId>>,
135}
136impl<M: Memory> MemoryManagerInner<M> {
137 fn init(memory: M, bucket_size_in_pages: u16) -> Self {
138 if memory.size() == 0 {
139 return Self::new(memory, bucket_size_in_pages);
140 }
141
142 let mut magic = [0_u8; 3];
143 memory.read(0, &mut magic);
144 if &magic == MAGIC {
145 Self::load(memory)
146 } else {
147 Self::new(memory, bucket_size_in_pages)
148 }
149 }
150
151 fn init_strict(memory: M, bucket_size_in_pages: u16) -> Result<Self, MemoryManagerInitError> {
152 if memory.size() == 0 {
153 return Ok(Self::new(memory, bucket_size_in_pages));
154 }
155
156 let mut magic = [0_u8; 3];
157 memory.read(0, &mut magic);
158 if &magic != MAGIC {
159 return Err(MemoryManagerInitError::NonMemoryManagerLayout);
160 }
161 Self::try_load(memory)
162 }
163
164 fn new(memory: M, bucket_size_in_pages: u16) -> Self {
165 let manager = Self {
166 memory,
167 allocated_buckets: 0,
168 bucket_size_in_pages,
169 memory_sizes_in_pages: [0; MAX_NUM_MEMORIES as usize],
170 memory_buckets: vec![Vec::new(); MAX_NUM_MEMORIES as usize],
171 };
172 write_growing(
173 &manager.memory,
174 bucket_allocations_address(BucketId(0)),
175 &[UNALLOCATED_BUCKET_MARKER; MAX_NUM_BUCKETS as usize],
176 );
177 manager.save_header();
178 manager
179 }
180 fn load(memory: M) -> Self {
181 let mut header = vec![0_u8; HEADER_SIZE as usize];
182 memory.read(0, &mut header);
183 assert_eq!(&header[0..3], MAGIC, "Bad magic.");
184 assert_eq!(header[3], LAYOUT_VERSION, "Unsupported version.");
185 let layout = load_validated_layout(&memory, &header);
186
187 Self {
188 memory,
189 allocated_buckets: layout.allocated_buckets,
190 bucket_size_in_pages: layout.bucket_size_in_pages,
191 memory_sizes_in_pages: layout.memory_sizes_in_pages,
192 memory_buckets: layout.memory_buckets,
193 }
194 }
195
196 fn try_load(memory: M) -> Result<Self, MemoryManagerInitError> {
197 let mut header = vec![0_u8; HEADER_SIZE as usize];
198 memory.read(0, &mut header);
199 if &header[0..3] != MAGIC {
200 return Err(MemoryManagerInitError::NonMemoryManagerLayout);
201 }
202 if header[3] != LAYOUT_VERSION {
203 return Err(MemoryManagerInitError::InvalidLayout(
204 "Unsupported version.".to_string(),
205 ));
206 }
207 let layout = try_load_validated_layout(&memory, &header)
208 .map_err(|error| MemoryManagerInitError::InvalidLayout(error.to_string()))?;
209
210 Ok(Self {
211 memory,
212 allocated_buckets: layout.allocated_buckets,
213 bucket_size_in_pages: layout.bucket_size_in_pages,
214 memory_sizes_in_pages: layout.memory_sizes_in_pages,
215 memory_buckets: layout.memory_buckets,
216 })
217 }
218
219 fn save_header(&self) {
220 let mut header = [0_u8; HEADER_SIZE as usize];
221 header[0..3].copy_from_slice(MAGIC);
222 header[3] = LAYOUT_VERSION;
223 header[4..6].copy_from_slice(&self.allocated_buckets.to_le_bytes());
224 header[6..8].copy_from_slice(&self.bucket_size_in_pages.to_le_bytes());
225 let mut offset = 3 + 1 + 2 + 2 + HEADER_RESERVED_BYTES;
226 for size in self.memory_sizes_in_pages {
227 header[offset..offset + 8].copy_from_slice(&size.to_le_bytes());
228 offset += 8;
229 }
230 write_growing(&self.memory, 0, &header);
231 }
232
233 fn memory_size(&self, id: MemoryId) -> u64 {
234 self.memory_sizes_in_pages[id.0 as usize]
235 }
236
237 fn grow(&mut self, id: MemoryId, pages: u64) -> i64 {
238 let old_size = self.memory_size(id);
239 let Some(new_size) = old_size.checked_add(pages) else {
240 return -1;
241 };
242 let current_buckets = self.num_buckets_needed(old_size);
243 let required_buckets = self.num_buckets_needed(new_size);
244 let new_buckets = required_buckets - current_buckets;
245 let Some(target_allocated_buckets) =
246 new_buckets.checked_add(u64::from(self.allocated_buckets))
247 else {
248 return -1;
249 };
250 if target_allocated_buckets > MAX_NUM_BUCKETS {
251 return -1;
252 }
253 let Ok(new_buckets_len) = usize::try_from(new_buckets) else {
254 return -1;
255 };
256 let memory_bucket = &mut self.memory_buckets[id.0 as usize];
257 if memory_bucket.try_reserve(new_buckets_len).is_err() {
258 return -1;
259 }
260 let mut rollback_buckets = Vec::new();
261 if rollback_buckets.try_reserve(new_buckets_len).is_err() {
262 return -1;
263 }
264
265 let Some(data_pages) =
266 u64::from(self.bucket_size_in_pages).checked_mul(target_allocated_buckets)
267 else {
268 return -1;
269 };
270 let Some(pages_needed) = BUCKETS_OFFSET_IN_PAGES.checked_add(data_pages) else {
271 return -1;
272 };
273 let current_pages = self.memory.size();
274 if pages_needed > current_pages {
275 let previous = self.memory.grow(pages_needed - current_pages);
276 if previous < 0 {
277 return -1;
278 }
279 }
280
281 let mut rollback = AllocationRollback {
282 memory: std::ptr::addr_of!(self.memory),
283 buckets: rollback_buckets,
284 committed: false,
285 _memory: std::marker::PhantomData,
286 };
287 for _ in 0..new_buckets {
288 let bucket = BucketId(self.allocated_buckets);
289 memory_bucket.push(bucket);
290 write_growing(&self.memory, bucket_allocations_address(bucket), &[id.0]);
291 rollback.buckets.push(bucket);
292 self.allocated_buckets = self
293 .allocated_buckets
294 .checked_add(1)
295 .expect("allocated bucket count overflow");
296 }
297
298 self.memory_sizes_in_pages[id.0 as usize] = new_size;
299 self.save_header();
300 rollback.committed = true;
301 old_size as i64
302 }
303
304 fn read(&self, id: MemoryId, offset: u64, dst: &mut [u8], cache: &BucketCache) {
305 unsafe { self.read_unsafe(id, offset, dst.as_mut_ptr(), dst.len(), cache) }
306 }
307
308 unsafe fn read_unsafe(
309 &self,
310 id: MemoryId,
311 offset: u64,
312 dst: *mut u8,
313 count: usize,
314 cache: &BucketCache,
315 ) {
316 if count == 0 {
317 return;
318 }
319 self.assert_bounds(id, offset, count as u64, "read");
320 if let Some(real) = cache.get(VirtualSegment::new(offset, count as u64)) {
321 self.memory.read_unsafe(real, dst, count);
322 return;
323 }
324 let mut bytes_read = 0_u64;
325 self.for_each_bucket(id, offset, count as u64, cache, |address, len| {
326 self.memory
327 .read_unsafe(address, dst.add(bytes_read as usize), len as usize);
328 bytes_read += len;
329 });
330 }
331
332 fn write(&self, id: MemoryId, offset: u64, src: &[u8], cache: &BucketCache) {
333 if src.is_empty() {
334 return;
335 }
336 self.assert_bounds(id, offset, src.len() as u64, "write");
337 if let Some(real) = cache.get(VirtualSegment::new(offset, src.len() as u64)) {
338 self.memory.write(real, src);
339 return;
340 }
341 let mut written = 0_u64;
342 self.for_each_bucket(id, offset, src.len() as u64, cache, |address, len| {
343 self.memory
344 .write(address, &src[written as usize..(written + len) as usize]);
345 written += len;
346 });
347 }
348
349 fn for_each_bucket(
350 &self,
351 MemoryId(id): MemoryId,
352 offset: u64,
353 mut len: u64,
354 cache: &BucketCache,
355 mut f: impl FnMut(u64, u64),
356 ) {
357 let bucket_size = self.bucket_size_in_bytes();
358 let buckets = self.memory_buckets[id as usize].as_slice();
359 let mut bucket_idx = (offset / bucket_size) as usize;
360 let mut bucket_offset = offset % bucket_size;
361 while len > 0 {
362 let bucket = buckets.get(bucket_idx).expect("bucket idx out of bounds");
363 let bucket_address = self.bucket_address(*bucket);
364 let segment_len = (bucket_size - bucket_offset).min(len);
365 cache.store(
366 VirtualSegment::new(bucket_idx as u64 * bucket_size, bucket_size),
367 bucket_address,
368 );
369 f(bucket_address + bucket_offset, segment_len);
370 len -= segment_len;
371 bucket_idx += 1;
372 bucket_offset = 0;
373 }
374 }
375
376 fn assert_bounds(&self, id: MemoryId, offset: u64, len: u64, operation: &str) {
377 let end = offset
378 .checked_add(len)
379 .unwrap_or_else(|| panic!("{id:?}: {operation} out of bounds"));
380 let capacity = self
381 .memory_size(id)
382 .checked_mul(STABLE_PAGE_SIZE)
383 .unwrap_or_else(|| panic!("{id:?}: {operation} out of bounds"));
384 assert!(end <= capacity, "{id:?}: {operation} out of bounds");
385 }
386
387 fn bucket_size_in_bytes(&self) -> u64 {
388 u64::from(self.bucket_size_in_pages) * STABLE_PAGE_SIZE
389 }
390
391 fn num_buckets_needed(&self, pages: u64) -> u64 {
392 pages.div_ceil(u64::from(self.bucket_size_in_pages))
393 }
394
395 fn bucket_address(&self, id: BucketId) -> u64 {
396 BUCKETS_OFFSET_IN_BYTES + self.bucket_size_in_bytes() * u64::from(id.0)
397 }
398}
399
400struct AllocationRollback<'memory, M: Memory> {
401 memory: *const M,
402 buckets: Vec<BucketId>,
403 committed: bool,
404 _memory: std::marker::PhantomData<&'memory M>,
405}
406
407impl<M: Memory> Drop for AllocationRollback<'_, M> {
408 fn drop(&mut self) {
409 if self.committed || !std::thread::panicking() {
410 return;
411 }
412 for bucket in self.buckets.iter().copied() {
413 let memory = unsafe { &*self.memory };
414 write_growing(
415 memory,
416 bucket_allocations_address(bucket),
417 &[UNALLOCATED_BUCKET_MARKER],
418 );
419 }
420 }
421}