Skip to main content

ic_sqlite_vfs/stable/
memory_manager.rs

1//! Minimal fork of `ic-stable-structures` MemoryManager 0.7 layout.
2//!
3//! The fork keeps the existing on-stable-memory format, but removes unrelated
4//! stable data structures from this crate's dependency graph.
5
6use 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}