diff --git a/dissect/database/sqlite3/sqlite3.py b/dissect/database/sqlite3/sqlite3.py index 0fe744c..7fdabcd 100644 --- a/dissect/database/sqlite3/sqlite3.py +++ b/dissect/database/sqlite3/sqlite3.py @@ -68,6 +68,7 @@ class SQLite3: fh: The path or file-like object to open a SQLite3 database on. wal: The path or file-like object to open a SQLite3 WAL file on. checkpoint: The checkpoint to apply from the WAL file. Can be a :class:`Checkpoint` object or an integer index. + validate_checksums: A boolean that sets whether to validate the checksum of frames when reading. Raises: ~dissect.database.sqlite3.exception.InvalidDatabase: If the file-like object does not look like a SQLite3 @@ -82,6 +83,8 @@ def __init__( fh: Path | BinaryIO, wal: WAL | Path | BinaryIO | None = None, checkpoint: Checkpoint | int | None = None, + *, + validate_checksums: bool = True, ): if isinstance(fh, Path): path = fh @@ -93,6 +96,7 @@ def __init__( self.path = path self.wal = None self.checkpoint = None + self.validate_checksums = validate_checksums self.header = c_sqlite3.header(self.fh) if self.header.magic != SQLITE3_HEADER_MAGIC: @@ -220,7 +224,7 @@ def raw_page(self, num: int) -> bytes: # Check if the latest valid instance of the page is committed (either the frame itself # is the commit frame or it is included in a commit's frames). If so, return that frame's data. for commit in reversed(self.wal.commits): - if (frame := commit.get(num)) and frame.valid: + if (frame := commit.get(num)) and frame.is_valid(validate_checksums=self.validate_checksums): data = frame.data break diff --git a/dissect/database/sqlite3/wal.py b/dissect/database/sqlite3/wal.py index ae01f34..04a538a 100644 --- a/dissect/database/sqlite3/wal.py +++ b/dissect/database/sqlite3/wal.py @@ -39,9 +39,24 @@ def __init__(self, fh: Path | BinaryIO): raise InvalidDatabase("Invalid WAL header magic") self.checksum_endian = "<" if self.header.magic == WAL_HEADER_MAGIC_LE else ">" - self.highest_page_num = max(fr.page_number for commit in self.commits for fr in commit.frames if fr.valid) + self._checksum_struct = struct.Struct(f"{self.checksum_endian}2I") self.frame = lru_cache(1024)(self.frame) + self.frame_size = len(c_sqlite3.wal_frame) + self.header.page_size + self.first_frame_offset = len(c_sqlite3.wal_header) + + # Only track the highest valid offset and its seed. + # Meaning: all frames with offset < _highest_valid_next_offset are considered valid. + # _highest_valid_next_offset initially points at the first frame; seed is checksum over header. + self._highest_valid_next_offset: int = self.first_frame_offset + self._highest_valid_seed: tuple[int, int] = self.header_checksum_seed + + # First offset that is known to fail checksum validation, or None. + self._checksum_failed_offset: int | None = None + + self.highest_page_num = max( + fr.page_number for commit in self.commits for fr in commit.frames if fr.is_valid_salt() + ) def close(self) -> None: """Close the WAL.""" @@ -50,8 +65,7 @@ def close(self) -> None: self.fh.close() def frame(self, frame_idx: int) -> Frame: - frame_size = len(c_sqlite3.wal_frame) + self.header.page_size - offset = len(c_sqlite3.wal_header) + frame_idx * frame_size + offset = self.first_frame_offset + frame_idx * self.frame_size return Frame(self, offset) def frames(self) -> Iterator[Frame]: @@ -63,6 +77,59 @@ def frames(self) -> Iterator[Frame]: except EOFError: # noqa: PERF203 break + def seed_for_offset(self, offset: int) -> tuple[int, int] | None: + """Return checksum seed after processing frames up to and including the frame at target_offset. + + Verify stored checksums for each frame as we walk. If a mismatch is found, update the WAL's + highest-known-valid-next-offset and return None. On success (no mismatches) update the + highest-known-valid-next-offset and seed and return the computed seed. + + References: + - https://sqlite.org/fileformat2.html#wal_file_format + - https://github.com/sqlite/sqlite/blob/master/src/wal.c#L995-L1047 + """ + # If the target offset is before the first frame, return the initial seed calculated from the WAL header. + if offset < self.first_frame_offset: + return self.header_checksum_seed + + # If the target offset is at or beyond the first known checksum failure, return None. + if self._checksum_failed_offset is not None and offset >= self._checksum_failed_offset: + return None + + # Start from the highest verified offset we know (saves re-checking earlier frames). + current_offset = self._highest_valid_next_offset + seed = self._highest_valid_seed + + while current_offset <= offset: + # Read frame header + self.fh.seek(current_offset) + frame_hdr_bytes = self.fh.read(len(c_sqlite3.wal_frame)) + if len(frame_hdr_bytes) < len(c_sqlite3.wal_frame): + raise EOFError("Incomplete frame header while calculating checksum") + + # Checksum first 16 bytes of frame header + seed = calculate_checksum(frame_hdr_bytes[:16], seed=seed, endian=self.checksum_endian) + + # Read and checksum page data + page_data = self.fh.read(self.header.page_size) + if len(page_data) < self.header.page_size: + raise EOFError("Incomplete page data while calculating checksum") + seed = calculate_checksum(page_data, seed=seed, endian=self.checksum_endian) + + # Compare computed seed to stored checksums in this frame header. + checksum1, checksum2 = self._checksum_struct.unpack(frame_hdr_bytes[-8:]) + if (seed[0], seed[1]) != (checksum1, checksum2): + self._checksum_failed_offset = current_offset + return None + + current_offset += self.frame_size + + # Update highest-known-valid-next-offset and seed to the next offset after target. + self._highest_valid_next_offset = current_offset + self._highest_valid_seed = seed + + return seed + @cached_property def commits(self) -> list[Commit]: """Return all commits in the WAL file. @@ -112,6 +179,11 @@ def checkpoints(self) -> list[Checkpoint]: return [checkpoints_map[salt] for salt in sorted(checkpoints_map.keys())] + @cached_property + def header_checksum_seed(self) -> tuple[int, int]: + """Cached initial checksum seed calculated from the WAL header first 24 bytes.""" + return calculate_checksum(self.header.dumps()[:24], endian=self.checksum_endian) + class Frame: def __init__(self, wal: WAL, offset: int): @@ -126,13 +198,40 @@ def __init__(self, wal: WAL, offset: int): def __repr__(self) -> str: return f"" - @property - def valid(self) -> bool: + def is_valid(self, validate_checksums: bool = True) -> bool: + """Return whether the frame is valid by comparing its salt values and optionally verifying the checksum. + + A frame is valid if: + - Its salt1 and salt2 values match those in the WAL header. + - Its checksum matches the calculated checksum. + + References: + - https://sqlite.org/fileformat2.html#wal_file_format + """ + return (self.is_valid_salt() and self.is_valid_checksum()) if validate_checksums else self.is_valid_salt() + + def is_valid_salt(self) -> bool: + """Return whether the frame's salt values match those in the WAL header. + + References: + - https://sqlite.org/fileformat2.html#wal_file_format + """ salt1_match = self.header.salt1 == self.wal.header.salt1 salt2_match = self.header.salt2 == self.wal.header.salt2 return salt1_match and salt2_match + def is_valid_checksum(self) -> bool: + """Return whether the frame's checksum matches the calculated checksum. + + Use WAL's highest valid offset to skip checks for already-verified frames. + """ + if self.offset < self.wal._highest_valid_next_offset: + return True + + seed = self.wal.seed_for_offset(self.offset) + return seed is not None + @property def data(self) -> bytes: self.fh.seek(self.offset + len(c_sqlite3.wal_frame)) @@ -188,8 +287,13 @@ class Commit(_FrameCollection): """ -def checksum(buf: bytes, endian: str = ">") -> tuple[int, int]: - s0 = s1 = 0 +def calculate_checksum(buf: bytes, seed: tuple[int, int] = (0, 0), endian: str = ">") -> tuple[int, int]: + """Calculate the checksum of a WAL header or frame. + + References: + - https://sqlite.org/fileformat2.html#checksum_algorithm + """ + s0, s1 = seed num_ints = len(buf) // 4 arr = struct.unpack(f"{endian}{num_ints}I", buf) diff --git a/tests/_data/sqlite3/big.sqlite b/tests/_data/sqlite3/big.sqlite new file mode 100644 index 0000000..bb2cb0f --- /dev/null +++ b/tests/_data/sqlite3/big.sqlite @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:49abb18da561d667cd02cd3f4aa5ad6f80740f29f9c2a24b6e7286cbe259ffe2 +size 69632 diff --git a/tests/_data/sqlite3/big.sqlite-wal b/tests/_data/sqlite3/big.sqlite-wal new file mode 100644 index 0000000..2df6fe5 --- /dev/null +++ b/tests/_data/sqlite3/big.sqlite-wal @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ada96856785eed8f8aca40d43f788cdb515eed6fcd40b77004134471c9160fd8 +size 8305952 diff --git a/tests/sqlite3/conftest.py b/tests/sqlite3/conftest.py index 9f86947..3579c0f 100644 --- a/tests/sqlite3/conftest.py +++ b/tests/sqlite3/conftest.py @@ -23,3 +23,13 @@ def sqlite_wal() -> Path: @pytest.fixture def empty_db() -> Path: return absolute_path("_data/sqlite3/empty.sqlite") + + +@pytest.fixture +def big_sqlite_db() -> Path: + return absolute_path("_data/sqlite3/big.sqlite") + + +@pytest.fixture +def big_sqlite_wal() -> Path: + return absolute_path("_data/sqlite3/big.sqlite-wal") diff --git a/tests/sqlite3/test_sqlite3.py b/tests/sqlite3/test_sqlite3.py index f5caaf9..e56a26b 100644 --- a/tests/sqlite3/test_sqlite3.py +++ b/tests/sqlite3/test_sqlite3.py @@ -17,11 +17,11 @@ [pytest.param(True, id="as_path"), pytest.param(False, id="as_fh")], ) def test_sqlite(sqlite_db: Path, open_as_path: bool) -> None: - db = sqlite3.SQLite3(sqlite_db if open_as_path else sqlite_db.open("rb")) + db = sqlite3.SQLite3(sqlite_db if open_as_path else sqlite_db.open("rb"), validate_checksums=False) _assert_sqlite_db(db) db.close() - with sqlite3.SQLite3(sqlite_db if open_as_path else sqlite_db.open("rb")) as db: + with sqlite3.SQLite3(sqlite_db if open_as_path else sqlite_db.open("rb"), validate_checksums=False) as db: _assert_sqlite_db(db) diff --git a/tests/sqlite3/test_wal.py b/tests/sqlite3/test_wal.py index e669561..8c71f71 100644 --- a/tests/sqlite3/test_wal.py +++ b/tests/sqlite3/test_wal.py @@ -10,6 +10,8 @@ if TYPE_CHECKING: from pathlib import Path + from pytest_benchmark.fixture import BenchmarkFixture + @pytest.mark.parametrize( ("db_as_path"), @@ -19,7 +21,7 @@ ("wal_as_path"), [pytest.param(True, id="wal_as_path"), pytest.param(False, id="wal_as_fh")], ) -def test_sqlite_wal(sqlite_db: Path, sqlite_wal: Path, db_as_path: bool, wal_as_path: bool) -> None: +def test_sqlite_wal_checkpoint(sqlite_db: Path, sqlite_wal: Path, db_as_path: bool, wal_as_path: bool) -> None: db = sqlite3.SQLite3( sqlite_db if db_as_path else sqlite_db.open("rb"), sqlite_wal if wal_as_path else sqlite_wal.open("rb"), @@ -48,6 +50,40 @@ def test_sqlite_wal(sqlite_db: Path, sqlite_wal: Path, db_as_path: bool, wal_as_ db.close() +@pytest.mark.parametrize( + ("db_as_path"), + [pytest.param(True, id="db_as_path"), pytest.param(False, id="db_as_fh")], +) +@pytest.mark.parametrize( + ("wal_as_path"), + [pytest.param(True, id="wal_as_path"), pytest.param(False, id="wal_as_fh")], +) +def test_sqlite_wal_checksum_validation(sqlite_db: Path, sqlite_wal: Path, db_as_path: bool, wal_as_path: bool) -> None: + # Test that the WAL checksum validation works as expected + # When validate_checksums=True, only entries before the last checkpoint are visible + db = sqlite3.SQLite3( + sqlite_db if db_as_path else sqlite_db.open("rb"), + sqlite_wal if wal_as_path else sqlite_wal.open("rb"), + validate_checksums=True, + ) + + _assert_valid_checksum(db) + + db.close() + + # When validate_checksums=False, entries after the last checkpoint are also visible + db = sqlite3.SQLite3( + sqlite_db if db_as_path else sqlite_db.open("rb"), + sqlite_wal if wal_as_path else sqlite_wal.open("rb"), + validate_checksums=False, + ) + + _assert_invalid_checksum(db) + + db.close() + + +# Assertion functions for test_sqlite_wal_checkpoint() def _assert_checkpoint_1(s: sqlite3.SQLite3) -> None: # After the first checkpoint the "after checkpoint" entries are present table = next(iter(s.tables())) @@ -165,6 +201,88 @@ def _assert_checkpoint_3(s: sqlite3.SQLite3) -> None: assert rows[9].value == 101 +# Assertion functions for test_sqlite_wal_checksum_validation() +def _assert_valid_checksum(s: sqlite3.SQLite3) -> None: + # If the checksum validation is correct, all entries BEFORE the last checkpoint should be present + table = next(iter(s.tables())) + rows = list(table.rows()) + + assert len(rows) == 11 + + assert rows[0].id == 1 + assert rows[0].name == "testing" + assert rows[0].value == 1337 + assert rows[1].id == 2 + assert rows[1].name == "omg" + assert rows[1].value == 7331 + assert rows[2].id == 3 + assert rows[2].name == "A" * 4100 + assert rows[2].value == 4100 + assert rows[3].id == 4 + assert rows[3].name == "B" * 4100 + assert rows[3].value == 4100 + assert rows[4].id == 5 + assert rows[4].name == "negative" + assert rows[4].value == -11644473429 + assert rows[5].id == 6 + assert rows[5].name == "after checkpoint" + assert rows[5].value == 42 + assert rows[6].id == 7 + assert rows[6].name == "after checkpoint" + assert rows[6].value == 43 + assert rows[7].id == 8 + assert rows[7].name == "after checkpoint" + assert rows[7].value == 44 + assert rows[8].id == 9 + assert rows[8].name == "after checkpoint" + assert rows[8].value == 45 + assert rows[9].id == 10 + assert rows[9].name == "second checkpoint" + assert rows[9].value == 100 + assert rows[10].id == 11 + assert rows[10].name == "second checkpoint" + assert rows[10].value == 101 + + +def _assert_invalid_checksum(s: sqlite3.SQLite3) -> None: + # If the checksum validation is incorrect, all entries AFTER the last checkpoint should be present + table = next(iter(s.tables())) + rows = list(table.rows()) + + assert len(rows) == 10 + + assert rows[0].id == 1 + assert rows[0].name == "testing" + assert rows[0].value == 1337 + assert rows[1].id == 2 + assert rows[1].name == "omg" + assert rows[1].value == 7331 + assert rows[2].id == 3 + assert rows[2].name == "A" * 4100 + assert rows[2].value == 4100 + assert rows[3].id == 4 + assert rows[3].name == "B" * 4100 + assert rows[3].value == 4100 + assert rows[4].id == 5 + assert rows[4].name == "negative" + assert rows[4].value == -11644473429 + assert rows[5].id == 6 + assert rows[5].name == "after checkpoint" + assert rows[5].value == 42 + assert rows[6].id == 8 + assert rows[6].name == "after checkpoint" + assert rows[6].value == 44 + assert rows[7].id == 9 + assert rows[7].name == "wow" + assert rows[7].value == 1234 + assert rows[8].id == 10 + assert rows[8].name == "second checkpoint" + assert rows[8].value == 100 + assert rows[9].id == 11 + assert rows[9].name == "second checkpoint" + assert rows[9].value == 101 + + def test_wal_page_count() -> None: """Test if we count the page numbers in the SQLite3 and WAL correctly. @@ -186,7 +304,7 @@ def test_wal_page_count() -> None: >>> con.commit() # Copy page_count.db* files before closing """ - db = sqlite3.SQLite3(absolute_path("_data/sqlite3/page_count.db")) + db = sqlite3.SQLite3(absolute_path("_data/sqlite3/page_count.db"), validate_checksums=False) table = db.table("t1") assert table.sql == "CREATE TABLE t1 (a, b)" @@ -198,3 +316,18 @@ def test_wal_page_count() -> None: assert db.wal.highest_page_num == 4 assert db.header.page_count == 2 assert db.page_count == 4 + + +@pytest.mark.parametrize( + ("validate"), + [pytest.param(True, id="True"), pytest.param(False, id="False")], +) +@pytest.mark.benchmark +def test_benchmark_wal_checksum_validation( + big_sqlite_db: Path, big_sqlite_wal: Path, validate: bool, benchmark: BenchmarkFixture +) -> None: + def benchy() -> None: + db = sqlite3.SQLite3(big_sqlite_db, big_sqlite_wal, validate_checksums=validate) + list(next(iter(db.tables()))) + + benchmark(benchy)