Fix for cannot change into wal mode from within a transaction

This commit is contained in:
Simon Willison 2023-06-25 15:52:22 -07:00
commit 1c0b8cbc04
2 changed files with 25 additions and 3 deletions

View file

@ -28,6 +28,7 @@ from typing import (
cast, cast,
Any, Any,
Callable, Callable,
ContextManager,
Dict, Dict,
Generator, Generator,
Iterable, Iterable,
@ -346,6 +347,25 @@ class Database:
"Close the SQLite connection, and the underlying database file" "Close the SQLite connection, and the underlying database file"
self.conn.close() self.conn.close()
@contextlib.contextmanager
def ensure_autocommit_off(self):
"""
Ensure autocommit is off for this database connection.
Example usage::
with db.ensure_autocommit_off():
# do stuff here
This will reset to the previous autocommit state at the end of the block.
"""
old_isolation_level = self.conn.isolation_level
try:
self.conn.isolation_level = None
yield
finally:
self.conn.isolation_level = old_isolation_level
@contextlib.contextmanager @contextlib.contextmanager
def tracer(self, tracer: Optional[Callable] = None): def tracer(self, tracer: Optional[Callable] = None):
""" """
@ -662,12 +682,14 @@ class Database:
Sets ``journal_mode`` to ``'wal'`` to enable Write-Ahead Log mode. Sets ``journal_mode`` to ``'wal'`` to enable Write-Ahead Log mode.
""" """
if self.journal_mode != "wal": if self.journal_mode != "wal":
self.execute("PRAGMA journal_mode=wal;") with self.ensure_autocommit_off():
self.execute("PRAGMA journal_mode=wal;")
def disable_wal(self): def disable_wal(self):
"Sets ``journal_mode`` back to ``'delete'`` to disable Write-Ahead Log mode." "Sets ``journal_mode`` back to ``'delete'`` to disable Write-Ahead Log mode."
if self.journal_mode != "delete": if self.journal_mode != "delete":
self.execute("PRAGMA journal_mode=delete;") with self.ensure_autocommit_off():
self.execute("PRAGMA journal_mode=delete;")
def _ensure_counts_table(self): def _ensure_counts_table(self):
with self.conn: with self.conn:

View file

@ -1429,7 +1429,7 @@ def test_enable_wal():
db = Database(dbname) db = Database(dbname)
db["t"].create({"pk": int}, pk="pk") db["t"].create({"pk": int}, pk="pk")
assert db.journal_mode == "delete" assert db.journal_mode == "delete"
result = runner.invoke(cli.cli, ["enable-wal"] + dbs) result = runner.invoke(cli.cli, ["enable-wal"] + dbs, catch_exceptions=False)
assert 0 == result.exit_code assert 0 == result.exit_code
for dbname in dbs: for dbname in dbs:
db = Database(dbname) db = Database(dbname)