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
17pub(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
99pub(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 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 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 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 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}