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
65pub 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 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}