Skip to main content

clt_database/storage/
database.rs

1use crate::io::FileSyncType;
2use crate::storage::checksum::ChecksumContext;
3use crate::storage::encryption::EncryptionContext;
4use crate::sync::Arc;
5use crate::{io::Completion, Buffer, CompletionError, LimboError, Result};
6use crate::{
7    turso_assert, turso_assert_eq, turso_assert_greater_than, turso_assert_greater_than_or_equal,
8    turso_assert_less_than_or_equal,
9};
10use tracing::{instrument, Level};
11
12#[derive(Debug, Clone)]
13pub enum EncryptionOrChecksum {
14    Encryption(EncryptionContext),
15    Checksum(ChecksumContext),
16    None,
17}
18
19#[derive(Debug, Clone)]
20pub struct IOContext {
21    encryption_or_checksum: EncryptionOrChecksum,
22}
23
24impl IOContext {
25    pub fn encryption_context(&self) -> Option<&EncryptionContext> {
26        match &self.encryption_or_checksum {
27            EncryptionOrChecksum::Encryption(ctx) => Some(ctx),
28            _ => None,
29        }
30    }
31
32    pub fn get_reserved_space_bytes(&self) -> u8 {
33        match &self.encryption_or_checksum {
34            EncryptionOrChecksum::Encryption(ctx) => ctx.required_reserved_bytes(),
35            EncryptionOrChecksum::Checksum(ctx) => ctx.required_reserved_bytes(),
36            EncryptionOrChecksum::None => Default::default(),
37        }
38    }
39
40    pub fn set_encryption(&mut self, encryption_ctx: EncryptionContext) {
41        self.encryption_or_checksum = EncryptionOrChecksum::Encryption(encryption_ctx);
42    }
43
44    pub fn encryption_or_checksum(&self) -> &EncryptionOrChecksum {
45        &self.encryption_or_checksum
46    }
47
48    pub fn reset_checksum(&mut self) {
49        self.encryption_or_checksum = EncryptionOrChecksum::None;
50    }
51}
52
53impl Default for IOContext {
54    fn default() -> Self {
55        #[cfg(clt_turso_feature = "checksum")]
56        let encryption_or_checksum = EncryptionOrChecksum::Checksum(ChecksumContext::default());
57        #[cfg(not(clt_turso_feature = "checksum"))]
58        let encryption_or_checksum = EncryptionOrChecksum::None;
59        Self {
60            encryption_or_checksum,
61        }
62    }
63}
64
65/// DatabaseStorage is an interface a database file that consists of pages.
66///
67/// The purpose of this trait is to abstract the upper layers of Limbo from
68/// the storage medium. A database can either be a file on disk, like in SQLite,
69/// or something like a remote page server service.
70pub trait DatabaseStorage: Send + Sync {
71    fn read_header(&self, c: Completion) -> Result<Completion>;
72
73    fn read_page(&self, page_idx: usize, io_ctx: &IOContext, c: Completion) -> Result<Completion>;
74    fn write_page(
75        &self,
76        page_idx: usize,
77        buffer: Arc<Buffer>,
78        io_ctx: &IOContext,
79        c: Completion,
80    ) -> Result<Completion>;
81    fn write_pages(
82        &self,
83        first_page_idx: usize,
84        page_size: usize,
85        buffers: Vec<Arc<Buffer>>,
86        io_ctx: &IOContext,
87        c: Completion,
88    ) -> Result<Completion>;
89    fn sync(&self, c: Completion, sync_type: FileSyncType) -> Result<Completion>;
90    fn size(&self) -> Result<u64>;
91    fn truncate(&self, len: usize, c: Completion) -> Result<Completion>;
92}
93
94#[derive(Clone)]
95pub struct DatabaseFile {
96    file: Arc<dyn crate::io::File>,
97}
98
99impl DatabaseStorage for DatabaseFile {
100    #[instrument(skip_all, level = Level::DEBUG)]
101    fn read_header(&self, c: Completion) -> Result<Completion> {
102        self.file.pread(0, c)
103    }
104
105    #[instrument(skip_all, level = Level::DEBUG)]
106    fn read_page(&self, page_idx: usize, io_ctx: &IOContext, c: Completion) -> Result<Completion> {
107        // casting to i64 to check some weird casting that could've happened before. This should be
108        // okay since page numbers should be u32
109        turso_assert_greater_than_or_equal!(page_idx as i64, 0);
110        let r = c.as_read();
111        let size = r.buf().len();
112        turso_assert_greater_than!(page_idx, 0);
113        if !(512..=65536).contains(&size) || size & (size - 1) != 0 {
114            return Err(LimboError::NotADB);
115        }
116        let Some(pos) = (page_idx as u64 - 1).checked_mul(size as u64) else {
117            return Err(LimboError::IntegerOverflow);
118        };
119
120        match &io_ctx.encryption_or_checksum {
121            EncryptionOrChecksum::Encryption(ctx) => {
122                let encryption_ctx = ctx.clone();
123                let read_buffer = r.buf_arc();
124                let original_c = c.clone();
125                let decrypt_complete =
126                    Box::new(move |res: Result<(Arc<Buffer>, i32), CompletionError>| {
127                        let (buf, bytes_read) = match res {
128                            Ok((buf, bytes_read)) => (buf, bytes_read),
129                            Err(err) => {
130                                tracing::error!(err = ?err);
131                                original_c.error(err);
132                                return original_c.get_error();
133                            }
134                        };
135                        turso_assert_greater_than!(
136                            bytes_read, 0,
137                            "database: expected to read data on success for encrypted page",
138                            { "page_idx": page_idx }
139                        );
140                        match encryption_ctx.decrypt_page(buf.as_slice(), page_idx) {
141                            Ok(decrypted_data) => {
142                                let original_buf = original_c.as_read().buf();
143                                original_buf.as_mut_slice().copy_from_slice(&decrypted_data);
144                                original_c.complete(bytes_read);
145                                original_c.get_error()
146                            }
147                            Err(e) => {
148                                tracing::error!(
149                                    "Failed to decrypt page data for page_id={page_idx}: {e}"
150                                );
151                                turso_assert!(
152                                    !original_c.failed(),
153                                    "Original completion already has an error"
154                                );
155                                original_c.error(CompletionError::DecryptionError { page_idx });
156                                original_c.get_error()
157                            }
158                        }
159                    });
160                let wrapped_completion = Completion::new_read(read_buffer, decrypt_complete);
161                self.file.pread(pos, wrapped_completion)
162            }
163            EncryptionOrChecksum::Checksum(ctx) => {
164                let checksum_ctx = ctx.clone();
165                let read_buffer = r.buf_arc();
166                let original_c = c.clone();
167
168                let verify_complete =
169                    Box::new(move |res: Result<(Arc<Buffer>, i32), CompletionError>| {
170                        let (buf, bytes_read) = match res {
171                            Ok((buf, bytes_read)) => (buf, bytes_read),
172                            Err(err) => {
173                                original_c.error(err);
174                                return original_c.get_error();
175                            }
176                        };
177                        if bytes_read <= 0 {
178                            tracing::trace!("Read page {page_idx} with {} bytes", bytes_read);
179                            original_c.complete(bytes_read);
180                            return original_c.get_error();
181                        }
182                        match checksum_ctx.verify_checksum(buf.as_mut_slice(), page_idx) {
183                            Ok(_) => {
184                                original_c.complete(bytes_read);
185                                original_c.get_error()
186                            }
187                            Err(e) => {
188                                tracing::error!(
189                                    "Failed to verify checksum for page_id={page_idx}: {e}"
190                                );
191                                turso_assert!(
192                                    !original_c.failed(),
193                                    "Original completion already has an error"
194                                );
195                                original_c.error(e);
196                                original_c.get_error()
197                            }
198                        }
199                    });
200
201                let wrapped_completion = Completion::new_read(read_buffer, verify_complete);
202                self.file.pread(pos, wrapped_completion)
203            }
204            EncryptionOrChecksum::None => self.file.pread(pos, c),
205        }
206    }
207
208    #[instrument(skip_all, level = Level::DEBUG)]
209    fn write_page(
210        &self,
211        page_idx: usize,
212        buffer: Arc<Buffer>,
213        io_ctx: &IOContext,
214        c: Completion,
215    ) -> Result<Completion> {
216        let buffer_size = buffer.len();
217        turso_assert_greater_than!(page_idx, 0);
218        turso_assert_greater_than_or_equal!(buffer_size, 512);
219        turso_assert_less_than_or_equal!(buffer_size, 65536);
220        turso_assert_eq!(buffer_size & (buffer_size - 1), 0);
221        let Some(pos) = (page_idx as u64 - 1).checked_mul(buffer_size as u64) else {
222            return Err(LimboError::IntegerOverflow);
223        };
224        let buffer = match &io_ctx.encryption_or_checksum {
225            EncryptionOrChecksum::Encryption(ctx) => encrypt_buffer(page_idx, buffer, ctx),
226            EncryptionOrChecksum::Checksum(ctx) => checksum_buffer(page_idx, buffer, ctx),
227            EncryptionOrChecksum::None => buffer,
228        };
229        self.file.pwrite(pos, buffer, c)
230    }
231
232    fn write_pages(
233        &self,
234        first_page_idx: usize,
235        page_size: usize,
236        buffers: Vec<Arc<Buffer>>,
237        io_ctx: &IOContext,
238        c: Completion,
239    ) -> Result<Completion> {
240        turso_assert_greater_than!(first_page_idx, 0);
241        turso_assert_greater_than_or_equal!(page_size, 512);
242        turso_assert_less_than_or_equal!(page_size, 65536);
243        turso_assert_eq!(page_size & (page_size - 1), 0);
244
245        let Some(pos) = (first_page_idx as u64 - 1).checked_mul(page_size as u64) else {
246            return Err(LimboError::IntegerOverflow);
247        };
248        let buffers = match &io_ctx.encryption_or_checksum() {
249            EncryptionOrChecksum::Encryption(ctx) => buffers
250                .into_iter()
251                .enumerate()
252                .map(|(i, buffer)| encrypt_buffer(first_page_idx + i, buffer, ctx))
253                .collect::<Vec<_>>(),
254            EncryptionOrChecksum::Checksum(ctx) => buffers
255                .into_iter()
256                .enumerate()
257                .map(|(i, buffer)| checksum_buffer(first_page_idx + i, buffer, ctx))
258                .collect::<Vec<_>>(),
259            EncryptionOrChecksum::None => buffers,
260        };
261        let c = self.file.pwritev(pos, buffers, c)?;
262        Ok(c)
263    }
264
265    #[instrument(skip_all, level = Level::DEBUG)]
266    fn sync(&self, c: Completion, sync_type: FileSyncType) -> Result<Completion> {
267        self.file.sync(c, sync_type)
268    }
269
270    #[instrument(skip_all, level = Level::DEBUG)]
271    fn size(&self) -> Result<u64> {
272        self.file.size()
273    }
274
275    #[instrument(skip_all, level = Level::DEBUG)]
276    fn truncate(&self, len: usize, c: Completion) -> Result<Completion> {
277        let c = self.file.truncate(len as u64, c)?;
278        Ok(c)
279    }
280}
281
282#[cfg(clt_turso_feature = "fs")]
283impl DatabaseFile {
284    pub fn new(file: Arc<dyn crate::io::File>) -> Self {
285        Self { file }
286    }
287}
288
289fn encrypt_buffer(page_idx: usize, buffer: Arc<Buffer>, ctx: &EncryptionContext) -> Arc<Buffer> {
290    let encrypted_data = ctx.encrypt_page(buffer.as_slice(), page_idx).unwrap();
291    Arc::new(Buffer::new(encrypted_data.to_vec()))
292}
293
294fn checksum_buffer(page_idx: usize, buffer: Arc<Buffer>, ctx: &ChecksumContext) -> Arc<Buffer> {
295    ctx.add_checksum_to_page(buffer.as_mut_slice(), page_idx)
296        .unwrap();
297    buffer
298}
299
300#[cfg(all(clt_turso_tests, clt_turso_feature = "checksum"))]
301mod tests {
302    use super::*;
303    use crate::File;
304    use crate::{io::IO, MemoryIO};
305
306    struct MockFile {
307        read_result: std::result::Result<i32, CompletionError>,
308    }
309
310    impl File for MockFile {
311        fn lock_file(&self, _exclusive: bool) -> Result<()> {
312            Ok(())
313        }
314
315        fn unlock_file(&self) -> Result<()> {
316            Ok(())
317        }
318
319        fn pread(&self, _pos: u64, c: Completion) -> Result<Completion> {
320            match self.read_result {
321                Ok(bytes_read) => c.complete(bytes_read),
322                Err(err) => c.error(err),
323            }
324            Ok(c)
325        }
326
327        fn pwrite(&self, _pos: u64, _buffer: Arc<Buffer>, c: Completion) -> Result<Completion> {
328            c.complete(0);
329            Ok(c)
330        }
331
332        fn sync(&self, c: Completion, _sync_type: FileSyncType) -> Result<Completion> {
333            c.complete(0);
334            Ok(c)
335        }
336
337        fn size(&self) -> Result<u64> {
338            Ok(0)
339        }
340
341        fn truncate(&self, _len: u64, c: Completion) -> Result<Completion> {
342            c.complete(0);
343            Ok(c)
344        }
345    }
346
347    #[test]
348    fn checksum_read_wrapper_propagates_callback_errors() {
349        let db_file = DatabaseFile {
350            file: Arc::new(MockFile { read_result: Ok(0) }),
351        };
352        let io_ctx = IOContext::default();
353        let page_idx = 1usize;
354        let expected = 4096usize;
355        let buf = Arc::new(Buffer::new_temporary(expected));
356        let original = Completion::new_read(buf, move |res| {
357            let (_, bytes_read) = res.expect("mock read should complete");
358            if bytes_read == 0 {
359                Some(CompletionError::ShortRead {
360                    page_idx,
361                    expected,
362                    actual: 0,
363                })
364            } else {
365                None
366            }
367        });
368
369        let wrapped = db_file
370            .read_page(page_idx, &io_ctx, original.clone())
371            .unwrap();
372        let io = MemoryIO::new();
373        let err = io
374            .wait_for_completion(wrapped)
375            .expect_err("wrapped completion must fail");
376        assert!(matches!(
377            err,
378            LimboError::CompletionError(CompletionError::ShortRead { .. })
379        ));
380        assert!(matches!(
381            original.get_error(),
382            Some(CompletionError::ShortRead { .. })
383        ));
384    }
385
386    #[test]
387    fn checksum_read_wrapper_propagates_transport_errors_to_original_completion() {
388        let db_file = DatabaseFile {
389            file: Arc::new(MockFile {
390                read_result: Err(CompletionError::Aborted),
391            }),
392        };
393        let io_ctx = IOContext::default();
394        let page_idx = 1usize;
395        let buf = Arc::new(Buffer::new_temporary(4096));
396        let original = Completion::new_read(buf, |_res| None);
397
398        let wrapped = db_file
399            .read_page(page_idx, &io_ctx, original.clone())
400            .unwrap();
401        let io = MemoryIO::new();
402        let err = io
403            .wait_for_completion(wrapped)
404            .expect_err("wrapped completion must fail");
405        assert!(matches!(
406            err,
407            LimboError::CompletionError(CompletionError::Aborted)
408        ));
409        assert_eq!(original.get_error(), Some(CompletionError::Aborted));
410    }
411}