Refactored storage layer to use per-operation auto-commit with explicit transaction API.
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 5s
CI / Tests (Python 3.12) (push) Successful in 20s
CI / Tests (Python 3.13) (push) Successful in 19s
CI / Tests (Python 3.14) (push) Successful in 17s
CI / Type Checking (push) Successful in 9s
CI / Spelling (push) Successful in 6s

This commit is contained in:
2026-03-09 20:03:38 -04:00
parent 87b037065f
commit 3186e6e392
12 changed files with 280 additions and 244 deletions
+1 -1
View File
@@ -30,7 +30,7 @@ Owlbot is a modular, event-driven chat bot for Owncast.
- Use decorators: `@on_command`, `@on_event`, `@on_route`, `@on_setup`, `@on_teardown`.
- Dynamic runtime registration is also supported via context registries.
- Storage:
- `api/storage.py` wraps `aiosqlite` with per-module DB files and handler-scoped transactions.
- `api/storage.py` wraps `aiosqlite` with per-module DB files, connection pooling, and auto-commit semantics.
## Build, Test, and Development Commands
- `uv sync`: install runtime and dev dependencies from `uv.lock`.
+1 -2
View File
@@ -22,8 +22,7 @@ owlbot:
#command_prefix: "!"
# Maximum seconds to wait for an event or command handler to complete.
# If a handler exceeds this, it is cancelled and any storage changes
# are rolled back.
# If a handler exceeds this, it is cancelled.
# Default: 30.0
#handler_timeout: 30.0
+1 -1
Submodule docs updated: 988cecbabb...2da7f1fc0f
+2 -2
View File
@@ -34,8 +34,8 @@ def on_setup(func: LifecycleHandler) -> LifecycleHandler:
"""Mark a function as a module setup hook.
The decorated function will be called during module loading with a
``ModuleContext``. Setup hooks run inside a storage transaction that
is committed on success and rolled back on failure.
``ModuleContext``. Each storage operation within a setup hook
auto-commits independently.
Applied without parentheses::
+33 -52
View File
@@ -41,15 +41,13 @@ class ModuleStorage:
concurrent readers and a single writer can operate without "database is
locked" errors.
Transactions are managed at the handler level. The bot commits after each
handler succeeds, and rolls back if the handler throws an exception. Module
developers don't need to think about commits for normal usage.
Each storage operation (execute, fetch_one, etc.) independently acquires a
connection from the pool, auto-commits on success, rolls back on failure,
and releases the connection immediately. No connection is held between
calls.
For finer control within a handler, use the transaction() context manager
to group multiple operations that should succeed or fail together.
Concurrency: The _checkout() context manager acquires a dedicated
connection from the pool for the current handler invocation.
For operations that must succeed or fail together, use the transaction()
context manager to group them into a single atomic unit.
"""
def __init__(self, storage_dir: Path, module_name: str, pool_size: int = 4):
@@ -90,10 +88,16 @@ class ModuleStorage:
@asynccontextmanager
async def transaction(self) -> AsyncIterator[ModuleStorage]:
"""Context manager for explicit transaction control within a handler.
"""Context manager for explicit transaction control.
Use this when you need multiple operations to succeed or fail together
within a single handler. Commits on success, rolls back on exception.
Use this when you need multiple operations to succeed or fail together.
Acquires a dedicated connection from the pool and shares it across all
operations within the block. Commits on success, rolls back on exception.
Nesting ``transaction()`` calls is not supported and raises
``RuntimeError``. SQLite only allows one writer at a time, so a
nested transaction would deadlock waiting for the outer connection's
write lock.
Example:
async with ctx.storage.transaction():
@@ -102,15 +106,27 @@ class ModuleStorage:
# Both committed together, or both rolled back on error
:return: This ModuleStorage instance.
:raises RuntimeError: If called while already inside a transaction.
"""
if self._txn_conn.get() is not None:
raise RuntimeError(
"transaction() cannot be nested. Already inside an active transaction."
)
conn = await self._acquire()
token = self._txn_conn.set(conn)
self._logger.debug("Explicit transaction started.")
try:
yield self
await self._commit()
await conn.commit()
self._logger.debug("Transaction committed.")
except BaseException:
await self._rollback()
await conn.rollback()
self._logger.debug("Transaction rolled back.")
raise
finally:
self._txn_conn.reset(token)
self._release(conn)
async def execute(
self,
@@ -255,10 +271,10 @@ class ModuleStorage:
async def _connection(self) -> AsyncIterator[aiosqlite.Connection]:
"""Async context manager that provides a connection.
If already inside a ``_checkout``, yields the checked-out connection
without releasing it. Otherwise acquires a standalone connection from
the pool that auto-commits on success and rolls back on failure
before being released.
If already inside a ``transaction()``, yields the shared connection
without committing (the transaction block handles that). Otherwise
acquires a standalone connection from the pool that auto-commits on
success and rolls back on failure before being released.
"""
existing = self._txn_conn.get()
if existing is not None:
@@ -275,41 +291,6 @@ class ModuleStorage:
finally:
self._release(conn)
@asynccontextmanager
async def _checkout(self) -> AsyncIterator[None]:
"""Check out a connection from the pool for the duration of a handler.
Sets a ContextVar so that all storage operations within the handler
reuse the same connection.
"""
conn = await self._acquire()
token = self._txn_conn.set(conn)
try:
yield
finally:
self._txn_conn.reset(token)
self._release(conn)
async def _commit(self) -> None:
"""Commit the current transaction (internal use by bot).
Called automatically after each handler completes successfully.
"""
conn = self._txn_conn.get()
if conn is not None:
await conn.commit()
self._logger.debug("Transaction committed.")
async def _rollback(self) -> None:
"""Rollback the current transaction (internal use by bot).
Called automatically if a handler throws an exception.
"""
conn = self._txn_conn.get()
if conn is not None:
await conn.rollback()
self._logger.debug("Transaction rolled back.")
async def _close(self) -> None:
"""Close all pool connections (internal use by bot)."""
self._closed = True
+7 -13
View File
@@ -426,9 +426,9 @@ class ModuleLoader:
async def _run_module_setup(self, module_name: str) -> None:
"""Run a module's ``@on_setup`` hooks if any are defined.
All setup handlers run inside a single storage transaction that is
committed on success. If any handler fails, the transaction is
rolled back, storage is closed, and the module is fully cleaned up.
Each storage operation within a setup handler auto-commits
independently. If any handler fails, storage is closed and the
module is fully cleaned up.
:param module_name: Name of the module whose setup to run.
:raises ModuleLoadError: If any @on_setup handler raises an exception.
@@ -443,16 +443,10 @@ class ModuleLoader:
module_ctx = self._module_contexts[module_name]
try:
async with module_ctx.storage._checkout():
try:
for setup_func in setup_funcs:
logger.debug(f"Running @on_setup for module: {module_name}")
await setup_func(module_ctx)
await module_ctx.storage._commit()
logger.debug(f"Setup completed for module: {module_name}")
except Exception:
await module_ctx.storage._rollback()
raise
for setup_func in setup_funcs:
logger.debug(f"Running @on_setup for module: {module_name}")
await setup_func(module_ctx)
logger.debug(f"Setup completed for module: {module_name}")
except Exception as e:
await module_ctx.storage._close()
self._cleanup_module(module_name)
+20 -32
View File
@@ -528,39 +528,27 @@ class CommandDispatcher:
module=module_ctx,
)
# The checkout acquires a pooled connection for the
# duration of this command invocation.
async with module_ctx.storage._checkout():
try:
start = time.perf_counter()
await asyncio.wait_for(
command_info.handler(cmd_ctx), timeout=self._handler_timeout
)
elapsed = (time.perf_counter() - start) * 1000
try:
start = time.perf_counter()
await asyncio.wait_for(
command_info.handler(cmd_ctx), timeout=self._handler_timeout
)
elapsed = (time.perf_counter() - start) * 1000
# Command succeeded, commit any database changes.
await module_ctx.storage._commit()
logger.debug(
f"Command '{command_info.name}' completed in {elapsed:.1f}ms."
)
except TimeoutError:
# Command timed out. Rollback any partial changes.
await module_ctx.storage._rollback()
logger.warning(
f"Command handler '{command_info.name}' "
f"from module '{command_info.module_name}' "
f"cancelled after "
f"{self._handler_timeout}s timeout."
)
except Exception as e:
# Command raised an exception. Rollback any partial changes.
await module_ctx.storage._rollback()
logger.exception(
f"Command handler '{command_info.name}' "
f"from module '{command_info.module_name}' "
f"raised exception: {e}"
)
logger.debug(f"Command '{command_info.name}' completed in {elapsed:.1f}ms.")
except TimeoutError:
logger.warning(
f"Command handler '{command_info.name}' "
f"from module '{command_info.module_name}' "
f"cancelled after "
f"{self._handler_timeout}s timeout."
)
except Exception as e:
logger.exception(
f"Command handler '{command_info.name}' "
f"from module '{command_info.module_name}' "
f"raised exception: {e}"
)
def _register_builtin_commands(self) -> None:
"""Register all built-in commands."""
+16 -28
View File
@@ -374,7 +374,7 @@ class EventDispatcher:
module_name: str,
propagation: PropagationState,
) -> None:
"""Call a single handler with timeout enforcement and transaction management.
"""Call a single handler with timeout enforcement.
:param handler: The handler function to call.
:param event: The event to pass to the handler.
@@ -395,34 +395,22 @@ class EventDispatcher:
_propagation=propagation,
)
# Call the handler with timeout enforcement to prevent
# runaway handlers from blocking everything.
# The checkout acquires a pooled connection for the
# duration of this handler invocation.
async with module_ctx.storage._checkout():
try:
start = time.perf_counter()
await asyncio.wait_for(handler(ctx), timeout=self._handler_timeout)
elapsed = (time.perf_counter() - start) * 1000
try:
start = time.perf_counter()
await asyncio.wait_for(handler(ctx), timeout=self._handler_timeout)
elapsed = (time.perf_counter() - start) * 1000
# Handler succeeded, commit any database changes.
await module_ctx.storage._commit()
logger.debug(f"Handler '{handler_name}' completed in {elapsed:.1f}ms.")
except TimeoutError:
# Handler took too long. Rollback any partial changes.
await module_ctx.storage._rollback()
logger.warning(
f"Handler '{handler_name}' from module '{module_name}' "
f"cancelled after {self._handler_timeout}s timeout."
)
except Exception as e:
# Handler raised an exception. Rollback any partial changes.
await module_ctx.storage._rollback()
logger.exception(
f"Handler '{handler_name}' from module "
f"'{module_name}' raised exception: {e}"
)
logger.debug(f"Handler '{handler_name}' completed in {elapsed:.1f}ms.")
except TimeoutError:
logger.warning(
f"Handler '{handler_name}' from module '{module_name}' "
f"cancelled after {self._handler_timeout}s timeout."
)
except Exception as e:
logger.exception(
f"Handler '{handler_name}' from module "
f"'{module_name}' raised exception: {e}"
)
class ModuleEvents:
+33 -42
View File
@@ -356,7 +356,7 @@ class RouteDispatcher:
"""Dispatches HTTP requests to registered module route handlers.
Looks up routes in the RouteRegistry, validates methods, creates
RouteContext, and calls the handler with timeout and transaction management.
RouteContext, and calls the handler with timeout enforcement.
"""
def __init__(
@@ -532,50 +532,41 @@ class RouteDispatcher:
f"Calling route handler: {route_info.full_path} from module: {module_name}"
)
# The checkout acquires a pooled connection for the
# duration of this route invocation.
async with module_ctx.storage._checkout():
try:
start = time.perf_counter()
result = await asyncio.wait_for(
route_info.handler(ctx), timeout=self._handler_timeout
)
elapsed = (time.perf_counter() - start) * 1000
try:
start = time.perf_counter()
result = await asyncio.wait_for(
route_info.handler(ctx), timeout=self._handler_timeout
)
elapsed = (time.perf_counter() - start) * 1000
# Handler succeeded, commit any database changes.
await module_ctx.storage._commit()
mod_logger.debug(
f"Route handler '{route_info.full_path}' completed in {elapsed:.1f}ms."
)
mod_logger.debug(
f"Route handler '{route_info.full_path}' "
f"completed in {elapsed:.1f}ms."
)
if result is None:
return web.Response(status=204) # No Content.
if isinstance(result, web.StreamResponse):
return result
if isinstance(result, dict):
return web.json_response(result)
# pragma: no branch — defensive against untyped handlers
mod_logger.error( # type: ignore[unreachable]
f"Route handler '{route_info.full_path}' returned "
f"unsupported type: {type(result).__name__}"
)
return web.Response(status=500)
if result is None:
return web.Response(status=204) # No Content.
if isinstance(result, web.StreamResponse):
return result
if isinstance(result, dict):
return web.json_response(result)
# pragma: no branch — defensive against untyped handlers
mod_logger.error( # type: ignore[unreachable]
f"Route handler '{route_info.full_path}' returned "
f"unsupported type: {type(result).__name__}"
)
return web.Response(status=500)
except TimeoutError:
await module_ctx.storage._rollback()
mod_logger.warning(
f"Route handler '{route_info.full_path}' timed out "
f"after {self._handler_timeout}s."
)
return web.Response(status=500)
except Exception as e:
await module_ctx.storage._rollback()
mod_logger.exception(
f"Route handler '{route_info.full_path}' raised exception: {e}"
)
return web.Response(status=500)
except TimeoutError:
mod_logger.warning(
f"Route handler '{route_info.full_path}' timed out "
f"after {self._handler_timeout}s."
)
return web.Response(status=500)
except Exception as e:
mod_logger.exception(
f"Route handler '{route_info.full_path}' raised exception: {e}"
)
return web.Response(status=500)
class ModuleRoutes:
+50 -5
View File
@@ -17,7 +17,7 @@
Tests cover the @on_command decorator, CommandEvent/CommandInfo dataclasses,
CommandRegistry (trigger mapping, alias resolution, conflict detection, module
scanning), CommandDispatcher (parsing, permission checks, cooldowns,
transaction commit/rollback, built-in commands), ModuleCommands
timeout handling, built-in commands), ModuleCommands
(ownership-scoped wrapper), and CommandContext.
"""
@@ -755,10 +755,10 @@ class TestCommandDispatcherDispatch:
assert row is not None
assert row["name"] == "apple"
async def test_exception_rolls_back(
async def test_exception_does_not_roll_back_committed_writes(
self, storage: ModuleStorage, caplog: pytest.LogCaptureFixture
) -> None:
"""A handler that raises has its database changes rolled back."""
"""A handler that raises does not undo already-committed writes."""
await storage.execute(
"CREATE TABLE IF NOT EXISTS items (id INTEGER, name TEXT)"
)
@@ -775,10 +775,11 @@ class TestCommandDispatcherDispatch:
await dispatcher.dispatch(_make_chat_event(body="!add"))
row = await storage.fetch_one("SELECT * FROM items WHERE id = ?", (1,))
assert row is None
assert row is not None
assert row["name"] == "apple"
assert "raised exception: boom" in caplog.text
async def test_timeout_rolls_back(
async def test_timeout_is_logged(
self, storage: ModuleStorage, caplog: pytest.LogCaptureFixture
) -> None:
"""A timed-out handler is cancelled and logged."""
@@ -796,6 +797,50 @@ class TestCommandDispatcherDispatch:
for r in caplog.records
)
async def test_timeout_does_not_roll_back_committed_writes(
self, storage: ModuleStorage
) -> None:
"""A timed-out handler's already-committed writes persist."""
await storage.execute(
"CREATE TABLE IF NOT EXISTS items (id INTEGER, name TEXT)"
)
async def handler(ctx: CommandContext) -> None:
await ctx.storage.execute(
"INSERT INTO items (id, name) VALUES (?, ?)", (1, "apple")
)
await asyncio.sleep(10)
dispatcher, _, _ = self._make_dispatcher(storage, handler_timeout=0.05)
dispatcher.register("slow", handler, module_name="mod_a")
await dispatcher.dispatch(_make_chat_event(body="!slow"))
row = await storage.fetch_one("SELECT * FROM items WHERE id = ?", (1,))
assert row is not None
assert row["name"] == "apple"
async def test_timeout_rolls_back_explicit_transaction(
self, storage: ModuleStorage
) -> None:
"""A timed-out handler's uncommitted transaction is rolled back."""
await storage.execute(
"CREATE TABLE IF NOT EXISTS items (id INTEGER, name TEXT)"
)
async def handler(ctx: CommandContext) -> None:
async with ctx.storage.transaction():
await ctx.storage.execute(
"INSERT INTO items (id, name) VALUES (?, ?)", (1, "apple")
)
await asyncio.sleep(10)
dispatcher, _, _ = self._make_dispatcher(storage, handler_timeout=0.05)
dispatcher.register("slow", handler, module_name="mod_a")
await dispatcher.dispatch(_make_chat_event(body="!slow"))
row = await storage.fetch_one("SELECT * FROM items WHERE id = ?", (1,))
assert row is None
async def test_builtin_command_dispatched(self, storage: ModuleStorage) -> None:
"""A built-in command receives (event, owncast_client)."""
captured: list[tuple[Any, Any]] = []
+48 -13
View File
@@ -15,7 +15,7 @@
"""Unit tests for event registration, dispatching, and context infrastructure.
Tests cover the @on_event decorator, Priority enum, EventRegistry,
EventDispatcher (priority sorting, propagation control, timeout/rollback,
EventDispatcher (priority sorting, propagation control, timeout handling,
command dispatch phase), ModuleEvents (ownership checks), and
EventContext/PropagationState.
"""
@@ -653,8 +653,8 @@ class TestEventDispatcherDispatch:
assert call_order == ["stopper", "follower"]
class TestEventDispatcherTransactions:
"""Tests _call_handler() commit/rollback behaviour in EventDispatcher."""
class TestEventDispatcherStorage:
"""Tests _call_handler() storage behaviour in EventDispatcher."""
async def test_success_commits(self, storage: ModuleStorage) -> None:
"""A successful handler's database changes are committed."""
@@ -682,16 +682,16 @@ class TestEventDispatcherTransactions:
assert row is not None
assert row[0] == "committed"
async def test_exception_rolls_back(
async def test_exception_does_not_roll_back_committed_writes(
self, storage: ModuleStorage, caplog: pytest.LogCaptureFixture
) -> None:
"""A handler that raises has its database changes rolled back."""
# Pre-create the table outside the handler's transaction.
"""A handler that raises does not undo already-committed writes."""
# Pre-create the table outside the handler.
await storage.execute("CREATE TABLE IF NOT EXISTS test_tbl (val TEXT)")
async def handler(ctx: EventContext[Any]) -> None:
await ctx.storage.execute(
"INSERT INTO test_tbl (val) VALUES (?)", ("rolled_back",)
"INSERT INTO test_tbl (val) VALUES (?)", ("persisted",)
)
raise RuntimeError("boom")
@@ -710,7 +710,8 @@ class TestEventDispatcherTransactions:
await dispatcher.dispatch(EventType.CHAT, _make_chat_event())
row = await storage.fetch_one("SELECT val FROM test_tbl")
assert row is None
assert row is not None
assert row[0] == "persisted"
assert any("raised exception: boom" in r.message for r in caplog.records)
async def test_exception_does_not_stop_next_handler(
@@ -741,10 +742,10 @@ class TestEventDispatcherTransactions:
await dispatcher.dispatch(EventType.CHAT, _make_chat_event())
assert call_order == ["failing", "second"]
async def test_timeout_rolls_back(
async def test_timeout_is_logged(
self, storage: ModuleStorage, caplog: pytest.LogCaptureFixture
) -> None:
"""A handler that exceeds the timeout is cancelled and changes rolled back."""
"""A handler that exceeds the timeout is cancelled and logged."""
async def slow_handler(ctx: EventContext[Any]) -> None:
await asyncio.sleep(10)
@@ -767,12 +768,16 @@ class TestEventDispatcherTransactions:
for r in caplog.records
)
async def test_timeout_rolls_back_db_changes(self, storage: ModuleStorage) -> None:
"""A timed-out handler's database writes are rolled back."""
async def test_timeout_does_not_roll_back_committed_writes(
self, storage: ModuleStorage
) -> None:
"""A timed-out handler's already-committed writes persist."""
await storage.execute("CREATE TABLE IF NOT EXISTS test_tbl (val TEXT)")
async def slow_handler(ctx: EventContext[Any]) -> None:
await ctx.storage.execute("INSERT INTO test_tbl (val) VALUES (?)", ("tmp",))
await ctx.storage.execute(
"INSERT INTO test_tbl (val) VALUES (?)", ("persisted",)
)
await asyncio.sleep(10)
module_ctx = _make_module_context(storage, "mod_a")
@@ -788,6 +793,36 @@ class TestEventDispatcherTransactions:
dispatcher.register(slow_handler, (EventType.CHAT,), "mod_a")
await dispatcher.dispatch(EventType.CHAT, _make_chat_event())
row = await storage.fetch_one("SELECT val FROM test_tbl")
assert row is not None
assert row[0] == "persisted"
async def test_timeout_rolls_back_explicit_transaction(
self, storage: ModuleStorage
) -> None:
"""A timed-out handler's uncommitted transaction is rolled back."""
await storage.execute("CREATE TABLE IF NOT EXISTS test_tbl (val TEXT)")
async def slow_handler(ctx: EventContext[Any]) -> None:
async with ctx.storage.transaction():
await ctx.storage.execute(
"INSERT INTO test_tbl (val) VALUES (?)", ("uncommitted",)
)
await asyncio.sleep(10)
module_ctx = _make_module_context(storage, "mod_a")
async def command_dispatch(event: ChatEvent) -> None:
pass
dispatcher = EventDispatcher(
command_dispatch=command_dispatch,
get_module_context=lambda name: module_ctx,
handler_timeout=0.05,
)
dispatcher.register(slow_handler, (EventType.CHAT,), "mod_a")
await dispatcher.dispatch(EventType.CHAT, _make_chat_event())
row = await storage.fetch_one("SELECT val FROM test_tbl")
assert row is None
+68 -53
View File
@@ -17,6 +17,7 @@
from __future__ import annotations
import asyncio
import contextlib
from typing import TYPE_CHECKING, Any
import pytest
@@ -190,7 +191,7 @@ class TestFetchValue:
class TestAutoCommit:
"""Outside _checkout(), each operation auto-commits or auto-rolls-back."""
"""Each operation auto-commits on success or auto-rolls-back on failure."""
async def test_auto_commits_on_success(
self, storage_with_table: ModuleStorage
@@ -226,48 +227,6 @@ class TestAutoCommit:
assert row["name"] == "a"
class TestCheckout:
"""_checkout() pins a single connection for all operations."""
async def test_commit_persists(self, storage_with_table: ModuleStorage) -> None:
"""Writes inside _checkout() become visible after _commit()."""
async with storage_with_table._checkout():
await storage_with_table.execute(
"INSERT INTO items (name, value) VALUES (?, ?)", ("a", 1)
)
await storage_with_table._commit()
row = await storage_with_table.fetch_one("SELECT name FROM items WHERE id = 1")
assert row is not None
assert row["name"] == "a"
async def test_rollback_discards(self, storage_with_table: ModuleStorage) -> None:
"""Writes inside _checkout() are discarded after _rollback()."""
async with storage_with_table._checkout():
await storage_with_table.execute(
"INSERT INTO items (name, value) VALUES (?, ?)", ("a", 1)
)
await storage_with_table._rollback()
row = await storage_with_table.fetch_one("SELECT name FROM items WHERE id = 1")
assert row is None
async def test_uncommitted_writes_invisible_from_other_connection(
self, storage_with_table: ModuleStorage
) -> None:
"""Uncommitted writes are not visible from a separate connection."""
async with (
ModuleStorage(storage_with_table._db_path.parent, "test_module") as s2,
storage_with_table._checkout(),
):
await storage_with_table.execute(
"INSERT INTO items (name, value) VALUES (?, ?)", ("a", 1)
)
# Before commit, s2 should not see the row.
row = await s2.fetch_one("SELECT name FROM items WHERE id = 1")
assert row is None
class TestTransaction:
"""transaction() context manager commits on clean exit, rolls back on exception."""
@@ -275,7 +234,7 @@ class TestTransaction:
self, storage_with_table: ModuleStorage
) -> None:
"""Writes inside transaction() persist after clean exit."""
async with storage_with_table._checkout(), storage_with_table.transaction():
async with storage_with_table.transaction():
await storage_with_table.execute(
"INSERT INTO items (name, value) VALUES (?, ?)", ("a", 1)
)
@@ -289,7 +248,7 @@ class TestTransaction:
) -> None:
"""Writes inside transaction() are discarded if an exception is raised."""
with pytest.raises(RuntimeError, match="boom"): # noqa: PT012
async with storage_with_table._checkout(), storage_with_table.transaction():
async with storage_with_table.transaction():
await storage_with_table.execute(
"INSERT INTO items (name, value) VALUES (?, ?)",
("a", 1),
@@ -299,6 +258,51 @@ class TestTransaction:
row = await storage_with_table.fetch_one("SELECT name FROM items WHERE id = 1")
assert row is None
async def test_nested_transaction_raises(
self, storage_with_table: ModuleStorage
) -> None:
"""Nesting transaction() calls raises RuntimeError."""
async with storage_with_table.transaction():
with pytest.raises(RuntimeError, match="cannot be nested"):
async with storage_with_table.transaction():
pass # pragma: no cover
async def test_uncommitted_transaction_invisible_from_other_connection(
self, storage_with_table: ModuleStorage
) -> None:
"""Uncommitted writes inside a transaction are not visible externally."""
async with (
ModuleStorage(storage_with_table._db_path.parent, "test_module") as s2,
storage_with_table.transaction(),
):
await storage_with_table.execute(
"INSERT INTO items (name, value) VALUES (?, ?)", ("a", 1)
)
# Before commit, s2 should not see the row.
row = await s2.fetch_one("SELECT name FROM items WHERE id = 1")
assert row is None
async def test_rolls_back_on_cancellation(
self, storage_with_table: ModuleStorage
) -> None:
"""A cancelled transaction rolls back uncommitted writes."""
async def slow_txn() -> None:
async with storage_with_table.transaction():
await storage_with_table.execute(
"INSERT INTO items (name, value) VALUES (?, ?)", ("a", 1)
)
await asyncio.sleep(10)
task = asyncio.create_task(slow_txn())
await asyncio.sleep(0.05)
task.cancel()
with contextlib.suppress(asyncio.CancelledError):
await task
row = await storage_with_table.fetch_one("SELECT name FROM items WHERE id = 1")
assert row is None
class TestConnectionPool:
"""Connection pool reuses connections and respects pool_size."""
@@ -311,21 +315,32 @@ class TestConnectionPool:
assert len(storage._all_connections) == 1
async def test_pool_grows_up_to_pool_size(self, tmp_path: Path) -> None:
"""Concurrent checkouts grow the pool up to pool_size."""
async with (
ModuleStorage(tmp_path, "pool_test", pool_size=2) as s,
s._checkout(),
s._checkout(),
):
"""Concurrent transactions grow the pool up to pool_size."""
gate = asyncio.Event()
async def hold_transaction(storage: ModuleStorage) -> None:
async with storage.transaction():
gate.set()
await asyncio.sleep(0.5)
async with ModuleStorage(tmp_path, "pool_test", pool_size=2) as s:
task = asyncio.create_task(hold_transaction(s))
await gate.wait()
# First task holds one connection; acquire a second directly.
conn = await asyncio.wait_for(s._acquire(), timeout=1.0)
assert len(s._all_connections) == 2
s._release(conn)
task.cancel()
with contextlib.suppress(asyncio.CancelledError):
await task
async def test_pool_exhaustion_blocks(self, tmp_path: Path) -> None:
"""When the pool is exhausted, acquire blocks until a connection is released."""
async with (
ModuleStorage(tmp_path, "pool_block", pool_size=1) as s,
s._checkout(),
s.transaction(),
):
# Pool is exhausted (pool_size=1, one checked out).
# Pool is exhausted (pool_size=1, one held by transaction).
# A second acquire should block, so wait_for should time out.
with pytest.raises(TimeoutError):
await asyncio.wait_for(s._acquire(), timeout=0.1)