Switched to explicit SQLite transactions with rollback safety and batched delete operations.
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 23s
CI / Type Checking (push) Successful in 11s
CI / Spelling (push) Successful in 5s

This commit is contained in:
2026-03-23 09:16:46 -04:00
parent 093426ac61
commit 53b8bd3fae
2 changed files with 123 additions and 76 deletions
+102 -76
View File
@@ -14,10 +14,14 @@
"""Async SQLite database access for Crabstero's persistent storage."""
from typing import NamedTuple, Self
from contextlib import asynccontextmanager
from typing import TYPE_CHECKING, NamedTuple, Self
import aiosqlite
if TYPE_CHECKING:
from collections.abc import AsyncIterator
class StartWord(NamedTuple):
"""A Markov chain starting word row."""
@@ -106,6 +110,21 @@ class Database:
"""
self._connection = connection
@asynccontextmanager
async def _transaction(self) -> AsyncIterator[None]:
"""Begin an immediate write transaction.
Commits on success, rolls back on error.
"""
await self._connection.execute("BEGIN IMMEDIATE")
try:
yield
except BaseException:
await self._connection.rollback()
raise
else:
await self._connection.commit()
@classmethod
async def connect(cls, path: str) -> Self:
"""Open a SQLite database, configure it, and create the schema.
@@ -116,7 +135,7 @@ class Database:
:param path: The file path to the SQLite database.
:return: A new Database instance ready for use.
"""
connection = await aiosqlite.connect(path)
connection = await aiosqlite.connect(path, isolation_level=None)
# Set synchronous to NORMAL for a balance between safety and speed.
await connection.execute("PRAGMA synchronous=NORMAL")
@@ -143,22 +162,22 @@ class Database:
:param start_words: Starting word entries to insert.
:param transitions: Transition entries to insert.
"""
if start_words:
await self._connection.executemany(
"INSERT INTO markov_start_words"
" (channel_id, user_id, word)"
" VALUES (?, ?, ?)",
start_words,
)
if transitions:
await self._connection.executemany(
"INSERT INTO markov_transitions"
" (channel_id, user_id, word, next_word)"
" VALUES (?, ?, ?, ?)",
transitions,
)
if start_words or transitions:
await self._connection.commit()
async with self._transaction():
if start_words:
await self._connection.executemany(
"INSERT INTO markov_start_words"
" (channel_id, user_id, word)"
" VALUES (?, ?, ?)",
start_words,
)
if transitions:
await self._connection.executemany(
"INSERT INTO markov_transitions"
" (channel_id, user_id, word, next_word)"
" VALUES (?, ?, ?, ?)",
transitions,
)
async def remove_markov_data(
self,
@@ -173,28 +192,34 @@ class Database:
:param start_words: Starting word entries to remove.
:param transitions: Transition entries to remove.
"""
for channel_id, user_id, word in start_words:
await self._connection.execute(
"DELETE FROM markov_start_words"
" WHERE rowid = ("
" SELECT rowid FROM markov_start_words"
" WHERE channel_id = ? AND user_id = ? AND word = ?"
" LIMIT 1"
" )",
(channel_id, user_id, word),
)
for channel_id, user_id, word, next_word in transitions:
await self._connection.execute(
"DELETE FROM markov_transitions"
" WHERE rowid = ("
" SELECT rowid FROM markov_transitions"
" WHERE channel_id = ? AND user_id = ? AND word = ? AND next_word = ?"
" LIMIT 1"
" )",
(channel_id, user_id, word, next_word),
)
if start_words or transitions:
await self._connection.commit()
async with self._transaction():
if start_words:
await self._connection.executemany(
"DELETE FROM markov_start_words"
" WHERE rowid = ("
" SELECT rowid FROM markov_start_words"
" WHERE channel_id = ?"
" AND user_id = ?"
" AND word = ?"
" LIMIT 1"
" )",
start_words,
)
if transitions:
await self._connection.executemany(
"DELETE FROM markov_transitions"
" WHERE rowid = ("
" SELECT rowid"
" FROM markov_transitions"
" WHERE channel_id = ?"
" AND user_id = ?"
" AND word = ?"
" AND next_word = ?"
" LIMIT 1"
" )",
transitions,
)
async def get_random_start_word(self, channel_id: int) -> str | None:
"""Return a random starting word for a channel.
@@ -262,13 +287,13 @@ class Database:
:param images: Image entries to insert.
"""
if images:
await self._connection.executemany(
"INSERT INTO channel_images"
" (channel_id, user_id, url)"
" VALUES (?, ?, ?)",
images,
)
await self._connection.commit()
async with self._transaction():
await self._connection.executemany(
"INSERT INTO channel_images"
" (channel_id, user_id, url)"
" VALUES (?, ?, ?)",
images,
)
async def remove_images(self, images: list[ChannelImage]) -> None:
"""Remove one matching row per entry from the images table.
@@ -278,18 +303,19 @@ class Database:
:param images: Image entries to remove.
"""
for channel_id, user_id, url in images:
await self._connection.execute(
"DELETE FROM channel_images"
" WHERE rowid = ("
" SELECT rowid FROM channel_images"
" WHERE channel_id = ? AND user_id = ? AND url = ?"
" LIMIT 1"
" )",
(channel_id, user_id, url),
)
if images:
await self._connection.commit()
async with self._transaction():
await self._connection.executemany(
"DELETE FROM channel_images"
" WHERE rowid = ("
" SELECT rowid FROM channel_images"
" WHERE channel_id = ?"
" AND user_id = ?"
" AND url = ?"
" LIMIT 1"
" )",
images,
)
async def get_random_image(self, channel_id: int) -> str | None:
"""Return a random image URL for a given channel.
@@ -313,13 +339,13 @@ class Database:
:param entity_id: The Discord ID of the entity.
:param flag_name: The name of the flag to set.
"""
await self._connection.execute(
"INSERT OR IGNORE INTO flags"
" (entity_type, entity_id, flag_name)"
" VALUES (?, ?, ?)",
(entity_type, entity_id, flag_name),
)
await self._connection.commit()
async with self._transaction():
await self._connection.execute(
"INSERT OR IGNORE INTO flags"
" (entity_type, entity_id, flag_name)"
" VALUES (?, ?, ?)",
(entity_type, entity_id, flag_name),
)
async def clear_flag(
self, entity_type: str, entity_id: str, flag_name: str
@@ -330,14 +356,14 @@ class Database:
:param entity_id: The Discord ID of the entity.
:param flag_name: The name of the flag to clear.
"""
await self._connection.execute(
"DELETE FROM flags"
" WHERE entity_type = ?"
" AND entity_id = ?"
" AND flag_name = ?",
(entity_type, entity_id, flag_name),
)
await self._connection.commit()
async with self._transaction():
await self._connection.execute(
"DELETE FROM flags"
" WHERE entity_type = ?"
" AND entity_id = ?"
" AND flag_name = ?",
(entity_type, entity_id, flag_name),
)
async def is_flag_set(
self, entity_type: str, entity_id: str, flag_name: str
@@ -375,8 +401,8 @@ class Database:
:param channel_id: The Discord channel ID.
"""
await self._connection.execute(
"INSERT OR IGNORE INTO ingested_channels (channel_id) VALUES (?)",
(channel_id,),
)
await self._connection.commit()
async with self._transaction():
await self._connection.execute(
"INSERT OR IGNORE INTO ingested_channels (channel_id) VALUES (?)",
(channel_id,),
)