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`. - Use decorators: `@on_command`, `@on_event`, `@on_route`, `@on_setup`, `@on_teardown`.
- Dynamic runtime registration is also supported via context registries. - Dynamic runtime registration is also supported via context registries.
- Storage: - 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 ## Build, Test, and Development Commands
- `uv sync`: install runtime and dev dependencies from `uv.lock`. - `uv sync`: install runtime and dev dependencies from `uv.lock`.
+1 -2
View File
@@ -22,8 +22,7 @@ owlbot:
#command_prefix: "!" #command_prefix: "!"
# Maximum seconds to wait for an event or command handler to complete. # Maximum seconds to wait for an event or command handler to complete.
# If a handler exceeds this, it is cancelled and any storage changes # If a handler exceeds this, it is cancelled.
# are rolled back.
# Default: 30.0 # Default: 30.0
#handler_timeout: 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. """Mark a function as a module setup hook.
The decorated function will be called during module loading with a The decorated function will be called during module loading with a
``ModuleContext``. Setup hooks run inside a storage transaction that ``ModuleContext``. Each storage operation within a setup hook
is committed on success and rolled back on failure. auto-commits independently.
Applied without parentheses:: Applied without parentheses::
+33 -52
View File
@@ -41,15 +41,13 @@ class ModuleStorage:
concurrent readers and a single writer can operate without "database is concurrent readers and a single writer can operate without "database is
locked" errors. locked" errors.
Transactions are managed at the handler level. The bot commits after each Each storage operation (execute, fetch_one, etc.) independently acquires a
handler succeeds, and rolls back if the handler throws an exception. Module connection from the pool, auto-commits on success, rolls back on failure,
developers don't need to think about commits for normal usage. and releases the connection immediately. No connection is held between
calls.
For finer control within a handler, use the transaction() context manager For operations that must succeed or fail together, use the transaction()
to group multiple operations that should succeed or fail together. context manager to group them into a single atomic unit.
Concurrency: The _checkout() context manager acquires a dedicated
connection from the pool for the current handler invocation.
""" """
def __init__(self, storage_dir: Path, module_name: str, pool_size: int = 4): def __init__(self, storage_dir: Path, module_name: str, pool_size: int = 4):
@@ -90,10 +88,16 @@ class ModuleStorage:
@asynccontextmanager @asynccontextmanager
async def transaction(self) -> AsyncIterator[ModuleStorage]: 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 Use this when you need multiple operations to succeed or fail together.
within a single handler. Commits on success, rolls back on exception. 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: Example:
async with ctx.storage.transaction(): async with ctx.storage.transaction():
@@ -102,15 +106,27 @@ class ModuleStorage:
# Both committed together, or both rolled back on error # Both committed together, or both rolled back on error
:return: This ModuleStorage instance. :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.") self._logger.debug("Explicit transaction started.")
try: try:
yield self yield self
await self._commit() await conn.commit()
self._logger.debug("Transaction committed.")
except BaseException: except BaseException:
await self._rollback() await conn.rollback()
self._logger.debug("Transaction rolled back.")
raise raise
finally:
self._txn_conn.reset(token)
self._release(conn)
async def execute( async def execute(
self, self,
@@ -255,10 +271,10 @@ class ModuleStorage:
async def _connection(self) -> AsyncIterator[aiosqlite.Connection]: async def _connection(self) -> AsyncIterator[aiosqlite.Connection]:
"""Async context manager that provides a connection. """Async context manager that provides a connection.
If already inside a ``_checkout``, yields the checked-out connection If already inside a ``transaction()``, yields the shared connection
without releasing it. Otherwise acquires a standalone connection from without committing (the transaction block handles that). Otherwise
the pool that auto-commits on success and rolls back on failure acquires a standalone connection from the pool that auto-commits on
before being released. success and rolls back on failure before being released.
""" """
existing = self._txn_conn.get() existing = self._txn_conn.get()
if existing is not None: if existing is not None:
@@ -275,41 +291,6 @@ class ModuleStorage:
finally: finally:
self._release(conn) 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: async def _close(self) -> None:
"""Close all pool connections (internal use by bot).""" """Close all pool connections (internal use by bot)."""
self._closed = True self._closed = True
+3 -9
View File
@@ -426,9 +426,9 @@ class ModuleLoader:
async def _run_module_setup(self, module_name: str) -> None: async def _run_module_setup(self, module_name: str) -> None:
"""Run a module's ``@on_setup`` hooks if any are defined. """Run a module's ``@on_setup`` hooks if any are defined.
All setup handlers run inside a single storage transaction that is Each storage operation within a setup handler auto-commits
committed on success. If any handler fails, the transaction is independently. If any handler fails, storage is closed and the
rolled back, storage is closed, and the module is fully cleaned up. module is fully cleaned up.
:param module_name: Name of the module whose setup to run. :param module_name: Name of the module whose setup to run.
:raises ModuleLoadError: If any @on_setup handler raises an exception. :raises ModuleLoadError: If any @on_setup handler raises an exception.
@@ -442,17 +442,11 @@ class ModuleLoader:
return return
module_ctx = self._module_contexts[module_name] module_ctx = self._module_contexts[module_name]
try:
async with module_ctx.storage._checkout():
try: try:
for setup_func in setup_funcs: for setup_func in setup_funcs:
logger.debug(f"Running @on_setup for module: {module_name}") logger.debug(f"Running @on_setup for module: {module_name}")
await setup_func(module_ctx) await setup_func(module_ctx)
await module_ctx.storage._commit()
logger.debug(f"Setup completed for module: {module_name}") logger.debug(f"Setup completed for module: {module_name}")
except Exception:
await module_ctx.storage._rollback()
raise
except Exception as e: except Exception as e:
await module_ctx.storage._close() await module_ctx.storage._close()
self._cleanup_module(module_name) self._cleanup_module(module_name)
+1 -13
View File
@@ -528,9 +528,6 @@ class CommandDispatcher:
module=module_ctx, module=module_ctx,
) )
# The checkout acquires a pooled connection for the
# duration of this command invocation.
async with module_ctx.storage._checkout():
try: try:
start = time.perf_counter() start = time.perf_counter()
await asyncio.wait_for( await asyncio.wait_for(
@@ -538,15 +535,8 @@ class CommandDispatcher:
) )
elapsed = (time.perf_counter() - start) * 1000 elapsed = (time.perf_counter() - start) * 1000
# Command succeeded, commit any database changes. logger.debug(f"Command '{command_info.name}' completed in {elapsed:.1f}ms.")
await module_ctx.storage._commit()
logger.debug(
f"Command '{command_info.name}' completed in {elapsed:.1f}ms."
)
except TimeoutError: except TimeoutError:
# Command timed out. Rollback any partial changes.
await module_ctx.storage._rollback()
logger.warning( logger.warning(
f"Command handler '{command_info.name}' " f"Command handler '{command_info.name}' "
f"from module '{command_info.module_name}' " f"from module '{command_info.module_name}' "
@@ -554,8 +544,6 @@ class CommandDispatcher:
f"{self._handler_timeout}s timeout." f"{self._handler_timeout}s timeout."
) )
except Exception as e: except Exception as e:
# Command raised an exception. Rollback any partial changes.
await module_ctx.storage._rollback()
logger.exception( logger.exception(
f"Command handler '{command_info.name}' " f"Command handler '{command_info.name}' "
f"from module '{command_info.module_name}' " f"from module '{command_info.module_name}' "
+1 -13
View File
@@ -374,7 +374,7 @@ class EventDispatcher:
module_name: str, module_name: str,
propagation: PropagationState, propagation: PropagationState,
) -> None: ) -> 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 handler: The handler function to call.
:param event: The event to pass to the handler. :param event: The event to pass to the handler.
@@ -395,30 +395,18 @@ class EventDispatcher:
_propagation=propagation, _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: try:
start = time.perf_counter() start = time.perf_counter()
await asyncio.wait_for(handler(ctx), timeout=self._handler_timeout) await asyncio.wait_for(handler(ctx), timeout=self._handler_timeout)
elapsed = (time.perf_counter() - start) * 1000 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.") logger.debug(f"Handler '{handler_name}' completed in {elapsed:.1f}ms.")
except TimeoutError: except TimeoutError:
# Handler took too long. Rollback any partial changes.
await module_ctx.storage._rollback()
logger.warning( logger.warning(
f"Handler '{handler_name}' from module '{module_name}' " f"Handler '{handler_name}' from module '{module_name}' "
f"cancelled after {self._handler_timeout}s timeout." f"cancelled after {self._handler_timeout}s timeout."
) )
except Exception as e: except Exception as e:
# Handler raised an exception. Rollback any partial changes.
await module_ctx.storage._rollback()
logger.exception( logger.exception(
f"Handler '{handler_name}' from module " f"Handler '{handler_name}' from module "
f"'{module_name}' raised exception: {e}" f"'{module_name}' raised exception: {e}"
+2 -11
View File
@@ -356,7 +356,7 @@ class RouteDispatcher:
"""Dispatches HTTP requests to registered module route handlers. """Dispatches HTTP requests to registered module route handlers.
Looks up routes in the RouteRegistry, validates methods, creates 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__( def __init__(
@@ -532,9 +532,6 @@ class RouteDispatcher:
f"Calling route handler: {route_info.full_path} from module: {module_name}" 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: try:
start = time.perf_counter() start = time.perf_counter()
result = await asyncio.wait_for( result = await asyncio.wait_for(
@@ -542,12 +539,8 @@ class RouteDispatcher:
) )
elapsed = (time.perf_counter() - start) * 1000 elapsed = (time.perf_counter() - start) * 1000
# Handler succeeded, commit any database changes.
await module_ctx.storage._commit()
mod_logger.debug( mod_logger.debug(
f"Route handler '{route_info.full_path}' " f"Route handler '{route_info.full_path}' completed in {elapsed:.1f}ms."
f"completed in {elapsed:.1f}ms."
) )
if result is None: if result is None:
@@ -564,14 +557,12 @@ class RouteDispatcher:
return web.Response(status=500) return web.Response(status=500)
except TimeoutError: except TimeoutError:
await module_ctx.storage._rollback()
mod_logger.warning( mod_logger.warning(
f"Route handler '{route_info.full_path}' timed out " f"Route handler '{route_info.full_path}' timed out "
f"after {self._handler_timeout}s." f"after {self._handler_timeout}s."
) )
return web.Response(status=500) return web.Response(status=500)
except Exception as e: except Exception as e:
await module_ctx.storage._rollback()
mod_logger.exception( mod_logger.exception(
f"Route handler '{route_info.full_path}' raised exception: {e}" f"Route handler '{route_info.full_path}' raised exception: {e}"
) )
+50 -5
View File
@@ -17,7 +17,7 @@
Tests cover the @on_command decorator, CommandEvent/CommandInfo dataclasses, Tests cover the @on_command decorator, CommandEvent/CommandInfo dataclasses,
CommandRegistry (trigger mapping, alias resolution, conflict detection, module CommandRegistry (trigger mapping, alias resolution, conflict detection, module
scanning), CommandDispatcher (parsing, permission checks, cooldowns, scanning), CommandDispatcher (parsing, permission checks, cooldowns,
transaction commit/rollback, built-in commands), ModuleCommands timeout handling, built-in commands), ModuleCommands
(ownership-scoped wrapper), and CommandContext. (ownership-scoped wrapper), and CommandContext.
""" """
@@ -755,10 +755,10 @@ class TestCommandDispatcherDispatch:
assert row is not None assert row is not None
assert row["name"] == "apple" 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 self, storage: ModuleStorage, caplog: pytest.LogCaptureFixture
) -> None: ) -> None:
"""A handler that raises has its database changes rolled back.""" """A handler that raises does not undo already-committed writes."""
await storage.execute( await storage.execute(
"CREATE TABLE IF NOT EXISTS items (id INTEGER, name TEXT)" "CREATE TABLE IF NOT EXISTS items (id INTEGER, name TEXT)"
) )
@@ -775,10 +775,11 @@ class TestCommandDispatcherDispatch:
await dispatcher.dispatch(_make_chat_event(body="!add")) await dispatcher.dispatch(_make_chat_event(body="!add"))
row = await storage.fetch_one("SELECT * FROM items WHERE id = ?", (1,)) 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 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 self, storage: ModuleStorage, caplog: pytest.LogCaptureFixture
) -> None: ) -> None:
"""A timed-out handler is cancelled and logged.""" """A timed-out handler is cancelled and logged."""
@@ -796,6 +797,50 @@ class TestCommandDispatcherDispatch:
for r in caplog.records 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: async def test_builtin_command_dispatched(self, storage: ModuleStorage) -> None:
"""A built-in command receives (event, owncast_client).""" """A built-in command receives (event, owncast_client)."""
captured: list[tuple[Any, Any]] = [] captured: list[tuple[Any, Any]] = []
+48 -13
View File
@@ -15,7 +15,7 @@
"""Unit tests for event registration, dispatching, and context infrastructure. """Unit tests for event registration, dispatching, and context infrastructure.
Tests cover the @on_event decorator, Priority enum, EventRegistry, 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 command dispatch phase), ModuleEvents (ownership checks), and
EventContext/PropagationState. EventContext/PropagationState.
""" """
@@ -653,8 +653,8 @@ class TestEventDispatcherDispatch:
assert call_order == ["stopper", "follower"] assert call_order == ["stopper", "follower"]
class TestEventDispatcherTransactions: class TestEventDispatcherStorage:
"""Tests _call_handler() commit/rollback behaviour in EventDispatcher.""" """Tests _call_handler() storage behaviour in EventDispatcher."""
async def test_success_commits(self, storage: ModuleStorage) -> None: async def test_success_commits(self, storage: ModuleStorage) -> None:
"""A successful handler's database changes are committed.""" """A successful handler's database changes are committed."""
@@ -682,16 +682,16 @@ class TestEventDispatcherTransactions:
assert row is not None assert row is not None
assert row[0] == "committed" 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 self, storage: ModuleStorage, caplog: pytest.LogCaptureFixture
) -> None: ) -> None:
"""A handler that raises has its database changes rolled back.""" """A handler that raises does not undo already-committed writes."""
# Pre-create the table outside the handler's transaction. # Pre-create the table outside the handler.
await storage.execute("CREATE TABLE IF NOT EXISTS test_tbl (val TEXT)") await storage.execute("CREATE TABLE IF NOT EXISTS test_tbl (val TEXT)")
async def handler(ctx: EventContext[Any]) -> None: async def handler(ctx: EventContext[Any]) -> None:
await ctx.storage.execute( await ctx.storage.execute(
"INSERT INTO test_tbl (val) VALUES (?)", ("rolled_back",) "INSERT INTO test_tbl (val) VALUES (?)", ("persisted",)
) )
raise RuntimeError("boom") raise RuntimeError("boom")
@@ -710,7 +710,8 @@ class TestEventDispatcherTransactions:
await dispatcher.dispatch(EventType.CHAT, _make_chat_event()) await dispatcher.dispatch(EventType.CHAT, _make_chat_event())
row = await storage.fetch_one("SELECT val FROM test_tbl") 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) assert any("raised exception: boom" in r.message for r in caplog.records)
async def test_exception_does_not_stop_next_handler( async def test_exception_does_not_stop_next_handler(
@@ -741,10 +742,10 @@ class TestEventDispatcherTransactions:
await dispatcher.dispatch(EventType.CHAT, _make_chat_event()) await dispatcher.dispatch(EventType.CHAT, _make_chat_event())
assert call_order == ["failing", "second"] assert call_order == ["failing", "second"]
async def test_timeout_rolls_back( async def test_timeout_is_logged(
self, storage: ModuleStorage, caplog: pytest.LogCaptureFixture self, storage: ModuleStorage, caplog: pytest.LogCaptureFixture
) -> None: ) -> 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: async def slow_handler(ctx: EventContext[Any]) -> None:
await asyncio.sleep(10) await asyncio.sleep(10)
@@ -767,12 +768,46 @@ class TestEventDispatcherTransactions:
for r in caplog.records for r in caplog.records
) )
async def test_timeout_rolls_back_db_changes(self, storage: ModuleStorage) -> None: async def test_timeout_does_not_roll_back_committed_writes(
"""A timed-out handler's database writes are rolled back.""" 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)") await storage.execute("CREATE TABLE IF NOT EXISTS test_tbl (val TEXT)")
async def slow_handler(ctx: EventContext[Any]) -> None: 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")
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 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) await asyncio.sleep(10)
module_ctx = _make_module_context(storage, "mod_a") module_ctx = _make_module_context(storage, "mod_a")
+68 -53
View File
@@ -17,6 +17,7 @@
from __future__ import annotations from __future__ import annotations
import asyncio import asyncio
import contextlib
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
import pytest import pytest
@@ -190,7 +191,7 @@ class TestFetchValue:
class TestAutoCommit: 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( async def test_auto_commits_on_success(
self, storage_with_table: ModuleStorage self, storage_with_table: ModuleStorage
@@ -226,48 +227,6 @@ class TestAutoCommit:
assert row["name"] == "a" 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: class TestTransaction:
"""transaction() context manager commits on clean exit, rolls back on exception.""" """transaction() context manager commits on clean exit, rolls back on exception."""
@@ -275,7 +234,7 @@ class TestTransaction:
self, storage_with_table: ModuleStorage self, storage_with_table: ModuleStorage
) -> None: ) -> None:
"""Writes inside transaction() persist after clean exit.""" """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( await storage_with_table.execute(
"INSERT INTO items (name, value) VALUES (?, ?)", ("a", 1) "INSERT INTO items (name, value) VALUES (?, ?)", ("a", 1)
) )
@@ -289,7 +248,7 @@ class TestTransaction:
) -> None: ) -> None:
"""Writes inside transaction() are discarded if an exception is raised.""" """Writes inside transaction() are discarded if an exception is raised."""
with pytest.raises(RuntimeError, match="boom"): # noqa: PT012 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( await storage_with_table.execute(
"INSERT INTO items (name, value) VALUES (?, ?)", "INSERT INTO items (name, value) VALUES (?, ?)",
("a", 1), ("a", 1),
@@ -299,6 +258,51 @@ class TestTransaction:
row = await storage_with_table.fetch_one("SELECT name FROM items WHERE id = 1") row = await storage_with_table.fetch_one("SELECT name FROM items WHERE id = 1")
assert row is None 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: class TestConnectionPool:
"""Connection pool reuses connections and respects pool_size.""" """Connection pool reuses connections and respects pool_size."""
@@ -311,21 +315,32 @@ class TestConnectionPool:
assert len(storage._all_connections) == 1 assert len(storage._all_connections) == 1
async def test_pool_grows_up_to_pool_size(self, tmp_path: Path) -> None: async def test_pool_grows_up_to_pool_size(self, tmp_path: Path) -> None:
"""Concurrent checkouts grow the pool up to pool_size.""" """Concurrent transactions grow the pool up to pool_size."""
async with ( gate = asyncio.Event()
ModuleStorage(tmp_path, "pool_test", pool_size=2) as s,
s._checkout(), async def hold_transaction(storage: ModuleStorage) -> None:
s._checkout(), 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 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: async def test_pool_exhaustion_blocks(self, tmp_path: Path) -> None:
"""When the pool is exhausted, acquire blocks until a connection is released.""" """When the pool is exhausted, acquire blocks until a connection is released."""
async with ( async with (
ModuleStorage(tmp_path, "pool_block", pool_size=1) as s, 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. # A second acquire should block, so wait_for should time out.
with pytest.raises(TimeoutError): with pytest.raises(TimeoutError):
await asyncio.wait_for(s._acquire(), timeout=0.1) await asyncio.wait_for(s._acquire(), timeout=0.1)