Skip to main content

clt_database/io/
memory.rs

1use super::{Buffer, Clock, Completion, File, OpenFlags, IO};
2use crate::io::clock::{DefaultClock, MonotonicInstant, WallClockInstant};
3use crate::io::FileSyncType;
4use crate::sync::{Mutex, RwLock};
5use crate::turso_assert;
6use crate::Result;
7use std::{
8    collections::{BTreeMap, HashMap},
9    sync::Arc,
10};
11use tracing::debug;
12
13pub struct MemoryIO {
14    files: Arc<Mutex<HashMap<String, Arc<MemoryFile>>>>,
15}
16
17// TODO: page size flag
18pub(super) const PAGE_SIZE: usize = 4096;
19pub(super) type MemPage = Box<[u8; PAGE_SIZE]>;
20
21struct MemStoreInner {
22    pages: BTreeMap<usize, MemPage>,
23    size: u64,
24}
25
26#[cfg(clt_turso_tests)]
27struct WritePause {
28    entered: std::sync::mpsc::Sender<()>,
29    release: std::sync::mpsc::Receiver<()>,
30}
31
32impl MemoryIO {
33    #[allow(clippy::arc_with_non_send_sync)]
34    pub fn new() -> Self {
35        debug!("Using IO backend 'memory'");
36        Self {
37            files: Arc::new(Mutex::new(HashMap::default())),
38        }
39    }
40}
41
42impl Default for MemoryIO {
43    fn default() -> Self {
44        Self::new()
45    }
46}
47
48impl Clock for MemoryIO {
49    fn current_time_monotonic(&self) -> MonotonicInstant {
50        DefaultClock.current_time_monotonic()
51    }
52
53    fn current_time_wall_clock(&self) -> WallClockInstant {
54        DefaultClock.current_time_wall_clock()
55    }
56}
57
58impl IO for MemoryIO {
59    fn open_file(&self, path: &str, flags: OpenFlags, _direct: bool) -> Result<Arc<dyn File>> {
60        let mut files = self.files.lock();
61        if !files.contains_key(path) && !flags.contains(OpenFlags::Create) {
62            return Err(crate::error::CompletionError::IOError(
63                std::io::ErrorKind::NotFound,
64                "open",
65            )
66            .into());
67        }
68        if !files.contains_key(path) {
69            files.insert(
70                path.to_string(),
71                Arc::new(MemoryFile {
72                    path: path.to_string(),
73                    store: MemStore::new(),
74                }),
75            );
76        }
77        Ok(files
78            .get(path)
79            .ok_or_else(|| {
80                crate::LimboError::InternalError("file should exist after insert".to_string())
81            })?
82            .clone())
83    }
84    fn remove_file(&self, path: &str) -> Result<()> {
85        let mut files = self.files.lock();
86        files.remove(path);
87        Ok(())
88    }
89
90    fn file_id(&self, path: &str) -> Result<super::FileId> {
91        Ok(super::FileId::from_path_hash(path))
92    }
93
94    fn supports_shared_wal_coordination(&self) -> bool {
95        false
96    }
97}
98
99/// In-memory page-backed storage shared by the synchronous [`MemoryIO`] and
100/// the yield-forcing [`super::MemoryYieldIO`] backends used for testing.
101pub(super) struct MemStore {
102    inner: RwLock<MemStoreInner>,
103    #[cfg(clt_turso_tests)]
104    next_write_pause: Mutex<Option<WritePause>>,
105}
106
107impl MemStore {
108    pub(super) fn new() -> Self {
109        Self {
110            inner: RwLock::new(MemStoreInner {
111                pages: BTreeMap::new(),
112                size: 0,
113            }),
114            #[cfg(clt_turso_tests)]
115            next_write_pause: Mutex::new(None),
116        }
117    }
118
119    #[cfg(clt_turso_tests)]
120    pub(super) fn pause_next_write(
121        &self,
122    ) -> (std::sync::mpsc::Receiver<()>, std::sync::mpsc::Sender<()>) {
123        let (entered_tx, entered_rx) = std::sync::mpsc::channel();
124        let (release_tx, release_rx) = std::sync::mpsc::channel();
125        assert!(
126            self.next_write_pause
127                .lock()
128                .replace(WritePause {
129                    entered: entered_tx,
130                    release: release_rx,
131                })
132                .is_none(),
133            "a write pause is already armed"
134        );
135        (entered_rx, release_tx)
136    }
137
138    #[cfg(clt_turso_tests)]
139    fn pause_test_write(&self) {
140        if let Some(pause) = self.next_write_pause.lock().take() {
141            pause.entered.send(()).unwrap();
142            pause.release.recv().unwrap();
143        }
144    }
145
146    fn get_or_allocate_page(inner: &mut MemStoreInner, page_no: usize) -> &mut MemPage {
147        inner
148            .pages
149            .entry(page_no)
150            .or_insert_with(|| Box::new([0; PAGE_SIZE]))
151    }
152
153    fn write_at_inner(inner: &mut MemStoreInner, pos: u64, data: &[u8]) -> usize {
154        let buf_len = data.len();
155        if buf_len == 0 {
156            return 0;
157        }
158        let mut offset = pos as usize;
159        let mut remaining = buf_len;
160        let mut buf_offset = 0;
161        while remaining > 0 {
162            let page_no = offset / PAGE_SIZE;
163            let page_offset = offset % PAGE_SIZE;
164            let bytes_to_write = remaining.min(PAGE_SIZE - page_offset);
165            let page = Self::get_or_allocate_page(inner, page_no);
166            page[page_offset..page_offset + bytes_to_write]
167                .copy_from_slice(&data[buf_offset..buf_offset + bytes_to_write]);
168            offset += bytes_to_write;
169            buf_offset += bytes_to_write;
170            remaining -= bytes_to_write;
171        }
172        inner.size = inner.size.max(pos + buf_len as u64);
173        buf_len
174    }
175
176    pub(super) fn size(&self) -> u64 {
177        self.inner.read().size
178    }
179
180    /// Read `pos..` into `buf`, zero-filling holes. Returns bytes read.
181    pub(super) fn read_into(&self, pos: u64, buf: &Buffer) -> i32 {
182        let buf_len = buf.len() as u64;
183        if buf_len == 0 {
184            return 0;
185        }
186        let inner = self.inner.read();
187        let file_size = inner.size;
188        if pos >= file_size {
189            return 0;
190        }
191        let read_len = buf_len.min(file_size - pos);
192        let dst = buf.as_mut_slice();
193        let mut offset = pos as usize;
194        let mut remaining = read_len as usize;
195        let mut buf_offset = 0;
196        while remaining > 0 {
197            let page_no = offset / PAGE_SIZE;
198            let page_offset = offset % PAGE_SIZE;
199            let bytes_to_read = remaining.min(PAGE_SIZE - page_offset);
200            if let Some(page) = inner.pages.get(&page_no) {
201                dst[buf_offset..buf_offset + bytes_to_read]
202                    .copy_from_slice(&page[page_offset..page_offset + bytes_to_read]);
203            } else {
204                dst[buf_offset..buf_offset + bytes_to_read].fill(0);
205            }
206            offset += bytes_to_read;
207            buf_offset += bytes_to_read;
208            remaining -= bytes_to_read;
209        }
210        read_len as i32
211    }
212
213    /// Write `data` at `pos`, allocating pages as needed. Returns bytes written
214    /// and grows the recorded file size if the write extends past it.
215    pub(super) fn write_at(&self, pos: u64, data: &[u8]) -> usize {
216        Self::write_at_inner(&mut self.inner.write(), pos, data)
217    }
218
219    /// Vectored write of `buffers` starting at `pos`. Returns total bytes written.
220    pub(super) fn writev(&self, pos: u64, buffers: &[Arc<Buffer>]) -> i32 {
221        let mut inner = self.inner.write();
222        let mut offset = pos;
223        let mut total_written = 0usize;
224        for buffer in buffers {
225            let written = Self::write_at_inner(&mut inner, offset, buffer.as_slice());
226            offset += written as u64;
227            total_written += written;
228            #[cfg(clt_turso_tests)]
229            self.pause_test_write();
230        }
231        total_written as i32
232    }
233
234    pub(super) fn truncate(&self, len: u64) {
235        let mut inner = self.inner.write();
236        if len < inner.size {
237            inner.pages.retain(|&k, _| k * PAGE_SIZE < len as usize);
238        }
239        inner.size = len;
240    }
241
242    pub(super) fn has_hole(&self, pos: usize, len: usize) -> bool {
243        let inner = self.inner.read();
244        let start_page = pos / PAGE_SIZE;
245        let end_page = ((pos + len.max(1)) - 1) / PAGE_SIZE;
246        for page_no in start_page..=end_page {
247            if inner.pages.contains_key(&page_no) {
248                return false;
249            }
250        }
251        true
252    }
253
254    pub(super) fn punch_hole(&self, pos: usize, len: usize) {
255        turso_assert!(
256            pos % PAGE_SIZE == 0 && len % PAGE_SIZE == 0,
257            "hole must be page aligned"
258        );
259        let mut inner = self.inner.write();
260        let start_page = pos / PAGE_SIZE;
261        let end_page = ((pos + len.max(1)) - 1) / PAGE_SIZE;
262        for page_no in start_page..=end_page {
263            inner.pages.remove(&page_no);
264        }
265    }
266}
267
268pub struct MemoryFile {
269    path: String,
270    store: MemStore,
271}
272
273crate::assert::assert_sync!(MemoryFile);
274
275impl File for MemoryFile {
276    fn lock_file(&self, _exclusive: bool) -> Result<()> {
277        Ok(())
278    }
279    fn unlock_file(&self) -> Result<()> {
280        Ok(())
281    }
282
283    fn pread(&self, pos: u64, c: Completion) -> Result<Completion> {
284        tracing::debug!("pread(path={}): pos={}", self.path, pos);
285        let n = self.store.read_into(pos, c.as_read().buf());
286        c.complete(n);
287        Ok(c)
288    }
289
290    fn pwrite(&self, pos: u64, buffer: Arc<Buffer>, c: Completion) -> Result<Completion> {
291        tracing::debug!(
292            "pwrite(path={}): pos={}, size={}",
293            self.path,
294            pos,
295            buffer.len()
296        );
297        let n = self.store.write_at(pos, buffer.as_slice());
298        c.complete(n as i32);
299        Ok(c)
300    }
301
302    fn sync(&self, c: Completion, _sync_type: FileSyncType) -> Result<Completion> {
303        tracing::debug!("sync(path={})", self.path);
304        // no-op
305        c.complete(0);
306        Ok(c)
307    }
308
309    fn truncate(&self, len: u64, c: Completion) -> Result<Completion> {
310        tracing::debug!("truncate(path={}): len={}", self.path, len);
311        self.store.truncate(len);
312        c.complete(0);
313        Ok(c)
314    }
315
316    fn pwritev(&self, pos: u64, buffers: Vec<Arc<Buffer>>, c: Completion) -> Result<Completion> {
317        tracing::debug!(
318            "pwritev(path={}): pos={}, buffers={:?}",
319            self.path,
320            pos,
321            buffers.iter().map(|x| x.len()).collect::<Vec<_>>()
322        );
323        let n = self.store.writev(pos, &buffers);
324        c.complete(n);
325        Ok(c)
326    }
327
328    fn size(&self) -> Result<u64> {
329        tracing::debug!("size(path={}): {}", self.path, self.store.size());
330        Ok(self.store.size())
331    }
332
333    fn has_hole(&self, pos: usize, len: usize) -> Result<bool> {
334        Ok(self.store.has_hole(pos, len))
335    }
336
337    fn punch_hole(&self, pos: usize, len: usize) -> Result<()> {
338        self.store.punch_hole(pos, len);
339        Ok(())
340    }
341}
342
343#[cfg(clt_turso_tests)]
344mod tests {
345    use super::*;
346    use std::{sync::mpsc, time::Duration};
347
348    #[test]
349    fn vectored_write_is_not_observed_partially() {
350        let store = Arc::new(MemStore::new());
351        store.write_at(0, &[0x11; PAGE_SIZE]);
352        let (write_entered, release_write) = store.pause_next_write();
353
354        let writer_store = store.clone();
355        let writer = std::thread::spawn(move || {
356            writer_store.writev(
357                0,
358                &[
359                    Arc::new(Buffer::new(vec![0xAA; PAGE_SIZE / 2])),
360                    Arc::new(Buffer::new(vec![0xAA; PAGE_SIZE / 2])),
361                ],
362            );
363        });
364        write_entered.recv().unwrap();
365
366        let (result_tx, result_rx) = mpsc::channel();
367        let reader = std::thread::spawn(move || {
368            let buffer = Buffer::new_temporary(PAGE_SIZE);
369            store.read_into(0, &buffer);
370            result_tx.send(buffer.as_slice().to_vec()).unwrap();
371        });
372
373        let early_result = result_rx.recv_timeout(Duration::from_secs(1));
374        release_write.send(()).unwrap();
375        writer.join().unwrap();
376        let bytes = match early_result {
377            Ok(bytes) => bytes,
378            Err(mpsc::RecvTimeoutError::Timeout) => {
379                result_rx.recv_timeout(Duration::from_secs(5)).unwrap()
380            }
381            Err(mpsc::RecvTimeoutError::Disconnected) => {
382                panic!("reader result channel disconnected")
383            }
384        };
385        reader.join().unwrap();
386        assert!(
387            bytes.iter().all(|&byte| byte == 0xAA),
388            "read returned a partially written file image; first old byte at offset {}",
389            bytes.iter().position(|&byte| byte == 0x11).unwrap()
390        );
391    }
392}