diff --git a/doc/contributing/11_memory_models.md b/doc/contributing/11_memory_models.md index a7076265f5..bc497ad8a0 100644 --- a/doc/contributing/11_memory_models.md +++ b/doc/contributing/11_memory_models.md @@ -179,9 +179,16 @@ and result-ID selection; updated bounds use a half-open UTC interval. Each reque is copied and revalidated before acquiring a session, so changes to the caller's query during execution cannot mix different filters or axes in one report. `QueryControl` supplies a monotonic deadline and request-local cancellation signal. -Async session acquisition is bounded by that budget. SQLite's per-connection -busy timeout is bounded by the remaining budget before each statement and -restored afterward, so lock waits cannot use the full default busy timeout. +Async session acquisition is bounded by that budget. SQLite's native busy +sleeps are disabled on the request-owned connection because they accumulate +requested sleep durations rather than enforcing an absolute deadline. Analytics +retries only plain `SQLITE_BUSY` reads and report setup operations with paced +asynchronous waits against the shared budget, without restarting the read +transaction. Other database errors, including `SQLITE_BUSY_SNAPSHOT`, propagate. +The original busy timeout is restored and the request's progress handler removed +before releasing the connection, even under repeated task cancellation. A failed +reset discards the connection and propagates the error. Deadline +checks remain cooperative; OS scheduling can delay when a task observes expiry. SQL Server pool acquisition remains subject to its pool; aioodbc configures pyodbc query timeouts in its executor before statement cursors are created and restores them when the session closes. SQL Server reports require SNAPSHOT diff --git a/pyrit/memory/attack_analytics.py b/pyrit/memory/attack_analytics.py index 4641ae9729..4992d86240 100644 --- a/pyrit/memory/attack_analytics.py +++ b/pyrit/memory/attack_analytics.py @@ -15,10 +15,11 @@ import asyncio import json import math +import sqlite3 from contextlib import asynccontextmanager from dataclasses import dataclass from datetime import UTC, datetime -from typing import TYPE_CHECKING, Any, ClassVar, NotRequired, TypedDict +from typing import TYPE_CHECKING, Any, ClassVar, NotRequired, TypedDict, TypeVar from aiosqlite import Connection as SQLiteConnection from pydantic import ValidationError @@ -28,6 +29,7 @@ from pyrit.common.pagination import decode_keyset_cursor, encode_keyset_cursor, fingerprint_filters from pyrit.exceptions.analytics_exception import AnalyticsDataException, AnalyticsTimeoutException from pyrit.memory.attack_analytics_query import AttackAnalyticsQueryCompiler +from pyrit.memory.sqlite_memory import _finish_sqlite_cleanup_async from pyrit.models import ( AttackAnalyticsFacetQuery, AttackAnalyticsFilters, @@ -40,15 +42,19 @@ ) if TYPE_CHECKING: - from collections.abc import AsyncGenerator + from collections.abc import AsyncGenerator, Awaitable, Callable + from sqlalchemy import Executable, Result from sqlalchemy.engine import RowMapping - from sqlalchemy.ext.asyncio import AsyncSession + from sqlalchemy.ext.asyncio import AsyncConnection, AsyncSession from pyrit.memory.memory_interface import MemoryInterface from pyrit.memory.query_control import QueryControl +_ResultT = TypeVar("_ResultT") + + @dataclass class RawAnalyticsOption: """A typed stored key and optional source label, before SDK absence labels are applied.""" @@ -128,6 +134,7 @@ class AttackAnalyticsReader: # Frequent Python callbacks serialize concurrent SQLite readers on the GIL. SQLITE_PROGRESS_STEPS: ClassVar[int] = 100_000 + SQLITE_BUSY_RETRY_DELAY: ClassVar[float] = 0.01 MAX_COMPACT_PROFILES: ClassVar[int] = 4096 MAX_COMPACT_VALUE_LENGTH: ClassVar[int] = 4096 MAX_COMPACT_TOTAL_LENGTH: ClassVar[int] = 1_000_000 @@ -159,7 +166,9 @@ async def report_async( async with self._session_async(control=control, consistent=True) as (session, dialect, warnings): compiler = AttackAnalyticsQueryCompiler(dialect=dialect, filters=query.filters) counts: dict[str, int] = {} - for row in (await session.execute(compiler.totals())).mappings(): + for row in ( + await self._execute_async(session=session, statement=compiler.totals(), control=control) + ).mappings(): outcome, count = row["outcome"], row["count"] if not isinstance(outcome, str) or type(count) is not int or count < 0: raise AnalyticsDataException("Stored attack outcomes contain invalid counts.") @@ -171,17 +180,25 @@ async def report_async( has_more = False truncated = False profiles = ( - await self._read_compact_profiles_async(session=session, compiler=compiler, query=query) + await self._read_compact_profiles_async( + session=session, compiler=compiler, query=query, control=control + ) if use_compact_profiles else None ) if profiles is None: if query.compare_by is None: - records = (await session.execute(compiler.groups(query))).mappings().all() + records = ( + (await self._execute_async(session=session, statement=compiler.groups(query), control=control)) + .mappings() + .all() + ) has_more = len(records) > query.group_limit groups = [self._group(record) for record in records[: query.group_limit]] else: - for record in (await session.execute(compiler.matrix(query))).mappings(): + for record in ( + await self._execute_async(session=session, statement=compiler.matrix(query), control=control) + ).mappings(): truncated = bool(record["truncated"]) if record["record"] == "row": rows.append(self._option(record, index=0)) @@ -193,6 +210,7 @@ async def report_async( session=session, dialect=dialect, query=AttackAnalyticsResultsQuery(filters=query.filters, limit=query.result_limit), + control=control, ) control.check() return RawAnalyticsReport( @@ -223,7 +241,7 @@ async def results_async( """ query = AttackAnalyticsResultsQuery.model_validate(query.model_dump()) async with self._session_async(control=control) as (session, dialect, _): - result = await self._results_async(session=session, dialect=dialect, query=query) + result = await self._results_async(session=session, dialect=dialect, query=query, control=control) control.check() return result @@ -252,7 +270,11 @@ async def facets_async(self, *, query: AttackAnalyticsFacetQuery, control: Query ) async with self._session_async(control=control) as (session, dialect, _): compiler = AttackAnalyticsQueryCompiler(dialect=dialect, filters=filters) - records = (await session.execute(compiler.facet(query))).mappings().all() + records = ( + (await self._execute_async(session=session, statement=compiler.facet(query), control=control)) + .mappings() + .all() + ) control.check() return RawAnalyticsFacets( items=[self._option(record, index=0) for record in records[: query.limit]], @@ -261,7 +283,12 @@ async def facets_async(self, *, query: AttackAnalyticsFacetQuery, control: Query ) async def _read_compact_profiles_async( - self, *, session: AsyncSession, compiler: AttackAnalyticsQueryCompiler, query: AttackAnalyticsQuery + self, + *, + session: AsyncSession, + compiler: AttackAnalyticsQueryCompiler, + query: AttackAnalyticsQuery, + control: QueryControl, ) -> list[RawAnalyticsProfile] | None: """ Accept only a complete, size-bounded profile projection for SDK aggregation. @@ -283,7 +310,7 @@ async def _read_compact_profiles_async( ) if statement is None: return None - records = (await session.execute(statement)).mappings().all() + records = (await self._execute_async(session=session, statement=statement, control=control)).mappings().all() if len(records) > self.MAX_COMPACT_PROFILES or any(record["oversized"] for record in records): return None text_length = sum(len(value) for record in records for value in record.values() if isinstance(value, str)) @@ -340,6 +367,80 @@ async def _sqlite_busy_timeout_async(*, driver: SQLiteConnection, timeout_ms: in async with driver.execute(f"PRAGMA busy_timeout = {timeout_ms}"): pass + async def _execute_async( + self, *, session: AsyncSession, statement: Executable, control: QueryControl + ) -> Result[Any]: + """ + Execute a read with deadline-aware SQLite lock retries. + + Returns: + Result[Any]: Buffered rows from the same request-owned session. + """ + if session.get_bind().dialect.name != "sqlite": + return await session.execute(statement) + return await self._retry_sqlite_busy_async(operation=lambda: session.execute(statement), control=control) + + async def _restore_sqlite_settings_async( + self, + *, + connection: AsyncConnection, + driver: SQLiteConnection, + busy_timeout: int | None, + cancelled: asyncio.CancelledError | None, + ) -> None: + """ + Finish resets before pooling, even under repeated cancellation. + + Raises: + CancelledError: If task cancellation is observed after connection cleanup finishes. + Exception: If resetting or discarding the connection fails. + """ + + async def restore_async() -> None: + try: + await driver.set_progress_handler(lambda: 0, 0) + if busy_timeout is not None: + await self._sqlite_busy_timeout_async(driver=driver, timeout_ms=busy_timeout) + except (asyncio.CancelledError, Exception) as error: + await connection.invalidate(error) + raise + + try: + cancellation = await _finish_sqlite_cleanup_async(restore_async()) + except (asyncio.CancelledError, Exception) as error: + if cancelled is not None: + cause = error.__cause__ if isinstance(error, asyncio.CancelledError) else error + raise cancelled from cause + raise + if cancellation is not None and cancelled is None: + raise cancellation + + @classmethod + async def _retry_sqlite_busy_async( + cls, *, operation: Callable[[], Awaitable[_ResultT]], control: QueryControl + ) -> _ResultT: + """ + Retry plain SQLITE_BUSY reads without restarting the transaction. + + Returns: + _ResultT: The completed read or setup operation's result. + + Raises: + AnalyticsTimeoutException: If cancellation or the shared deadline is observed. + OperationalError: If the failure is not plain SQLITE_BUSY, including a stale snapshot. + """ + while True: + control.check() + try: + result = await operation() + control.check() + return result + except OperationalError as error: + if getattr(error.orig, "sqlite_errorcode", None) != sqlite3.SQLITE_BUSY: + raise + control.check() + await asyncio.sleep(min(cls.SQLITE_BUSY_RETRY_DELAY, control.remaining)) + @staticmethod async def _odbc_timeout_async(*, driver: Any, timeout: int) -> None: """ @@ -360,8 +461,11 @@ async def _session_async( """ Own an async session and its cancellation hooks until database work finishes. - SQLite's progress handler interrupts long statements cooperatively; its - busy timeout is bounded per statement so lock waits share the deadline. + SQLite's progress handler interrupts long statements cooperatively. Native + busy sleeps are disabled for this session: read/setup operations retry plain + SQLITE_BUSY asynchronously against the same absolute deadline, without + restarting the read transaction. Other lock errors propagate. The original + busy timeout is restored before releasing the connection. aioodbc receives the remaining whole-second timeout before each cursor is created. Database operations use the native async drivers, not the deprecated synchronous memory session API. @@ -381,6 +485,7 @@ async def _session_async( Raises: AnalyticsTimeoutException: If acquisition or execution outlasts the budget, or cancellation is observed. + CancelledError: If the requesting task is cancelled. AnalyticsDataException: If no live DBAPI connection is available. NotImplementedError: If the driver cannot supply the required interruption/timeout hook. OperationalError: If a database error occurs before expiry. @@ -407,6 +512,7 @@ async def _session_async( warnings: list[str] = [] old_timeout: int | None = None old_busy_timeout: int | None = None + cancelled: asyncio.CancelledError | None = None def before_statement( conn: Any, clauseelement: Any, multiparams: Any, params: Any, execution_options: Any @@ -423,16 +529,8 @@ def before_statement( def before_cursor_execute( conn: Any, cursor: Any, statement: str, parameters: Any, context: Any, executemany: bool ) -> None: - """Recheck the budget and bound SQLite waits at the statement boundary.""" + """Recheck the shared budget after SQL compilation.""" control.check() - if old_busy_timeout is not None: - timeout_ms = min(old_busy_timeout, max(1, math.ceil(control.remaining * 1000))) - conn.connection.run_async( - lambda aio_driver: self._sqlite_busy_timeout_async( - driver=aio_driver, - timeout_ms=timeout_ms, - ) - ) sync_connection = connection.sync_connection event.listen(sync_connection, "before_execute", before_statement) @@ -447,12 +545,14 @@ def before_cursor_execute( if row is None or not isinstance(row[0], int): raise AnalyticsDataException("The SQLite busy timeout could not be read.") old_busy_timeout = row[0] - await self._sqlite_busy_timeout_async( - driver=driver, timeout_ms=min(old_busy_timeout, max(1, math.ceil(control.remaining * 1000))) - ) + await self._sqlite_busy_timeout_async(driver=driver, timeout_ms=0) await driver.set_progress_handler(lambda: int(control.expired), self.SQLITE_PROGRESS_STEPS) if consistent: - mode = (await connection.exec_driver_sql("PRAGMA journal_mode")).scalar_one() + mode = ( + await self._retry_sqlite_busy_async( + operation=lambda: connection.exec_driver_sql("PRAGMA journal_mode"), control=control + ) + ).scalar_one() if mode not in {"wal", "memory"}: warnings.append( "SQLite rollback journaling can slow analytics during concurrent writes. " @@ -460,7 +560,9 @@ def before_cursor_execute( "Analytics has not changed database settings." ) if not driver.in_transaction: - await connection.exec_driver_sql("BEGIN") + await self._retry_sqlite_busy_async( + operation=lambda: connection.exec_driver_sql("BEGIN"), control=control + ) elif dialect == "mssql": old_timeout = getattr(driver, "timeout", None) if not isinstance(old_timeout, int): @@ -469,6 +571,9 @@ def before_cursor_execute( else: raise NotImplementedError(f"Attack analytics does not support {dialect!r}") yield session, dialect, warnings + except asyncio.CancelledError as error: + cancelled = error + raise except OperationalError as error: if control.expired: raise AnalyticsTimeoutException from error @@ -478,14 +583,14 @@ def before_cursor_execute( event.remove(sync_connection, "before_cursor_execute", before_cursor_execute) if not connection.invalidated and not connection.closed: if isinstance(driver, SQLiteConnection): - await driver.set_progress_handler(lambda: 0, 0) - if old_busy_timeout is not None: - await self._sqlite_busy_timeout_async(driver=driver, timeout_ms=old_busy_timeout) + await self._restore_sqlite_settings_async( + connection=connection, driver=driver, busy_timeout=old_busy_timeout, cancelled=cancelled + ) if old_timeout is not None: await self._odbc_timeout_async(driver=driver, timeout=old_timeout) async def _results_async( - self, *, session: AsyncSession, dialect: str, query: AttackAnalyticsResultsQuery + self, *, session: AsyncSession, dialect: str, query: AttackAnalyticsResultsQuery, control: QueryControl ) -> AttackAnalyticsResults: """ Validate a filter-bound seek cursor and materialize at most one visible metadata page. @@ -502,7 +607,15 @@ async def _results_async( if query.cursor is not None and after is None: raise ValueError("Invalid or stale analytics cursor. Reload results with the current filters.") compiler = AttackAnalyticsQueryCompiler(dialect=dialect, filters=query.filters) - records = (await session.execute(compiler.results(limit=query.limit, after=after))).mappings().all() + records = ( + ( + await self._execute_async( + session=session, statement=compiler.results(limit=query.limit, after=after), control=control + ) + ) + .mappings() + .all() + ) items = [self._result_row(record) for record in records[: query.limit]] has_more = len(records) > query.limit cursor = ( diff --git a/tests/unit/memory/test_attack_analytics_lock_retry.py b/tests/unit/memory/test_attack_analytics_lock_retry.py new file mode 100644 index 0000000000..9cc79a42b9 --- /dev/null +++ b/tests/unit/memory/test_attack_analytics_lock_retry.py @@ -0,0 +1,402 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +from __future__ import annotations + +import asyncio +import sqlite3 +import time +from typing import TYPE_CHECKING, Any, TypeVar +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from aiosqlite import Connection as SQLiteConnection +from sqlalchemy import text +from sqlalchemy.exc import OperationalError +from sqlalchemy.ext.asyncio import AsyncConnection, AsyncSession + +from pyrit.exceptions.analytics_exception import AnalyticsTimeoutException +from pyrit.memory import SQLiteMemory +from pyrit.memory.attack_analytics import AttackAnalyticsReader +from pyrit.memory.memory_models import Base +from pyrit.memory.query_control import QueryControl +from pyrit.models import ( + AttackAnalyticsDimension, + AttackAnalyticsFacetQuery, + AttackAnalyticsQuery, + AttackAnalyticsResultsQuery, + AttackOutcome, + AttackResult, +) + +if TYPE_CHECKING: + from collections.abc import AsyncGenerator, Awaitable, Callable + from pathlib import Path + + from sqlalchemy import CursorResult, Executable, Result + + +_ResultT = TypeVar("_ResultT") + + +def _sqlite_error(code: int) -> OperationalError: + original = sqlite3.OperationalError("database is locked") + original.sqlite_errorcode = code + return OperationalError("SELECT", {}, original) + + +@pytest.fixture +async def file_memory(tmp_path: Path) -> AsyncGenerator[SQLiteMemory, None]: + memory = SQLiteMemory.__new__(SQLiteMemory) + with patch.object(memory, "cleanup"): + memory.__init__(db_path=tmp_path / "analytics-retry.sqlite", skip_schema_migration=True) + memory.results_path = str(tmp_path) + try: + async with memory._get_async_engine().begin() as connection: + await connection.run_sync(Base.metadata.create_all) + await memory.add_attack_results_to_memory_async( + attack_results=[ + AttackResult( + conversation_id="lock-retry", + objective="Retry a locked analytics read", + operation="lock-test", + outcome=AttackOutcome.SUCCESS, + ) + ] + ) + yield memory + finally: + await memory.dispose_engine_async() + + +async def _read_async(*, reader: AttackAnalyticsReader, method: str, control: QueryControl) -> None: + if method in {"report", "matrix", "compact_report"}: + report = await reader.report_async( + query=AttackAnalyticsQuery( + group_by=AttackAnalyticsDimension(name="targeted_harm_category"), + compare_by=AttackAnalyticsDimension(name="operation") if method == "matrix" else None, + ), + control=control, + use_compact_profiles=method == "compact_report", + ) + assert report.counts == {"success": 1} + assert len(report.results.items) == 1 + if method == "compact_report": + assert report.profiles is not None + elif method == "matrix": + assert len(report.cells) == 1 + elif method == "results": + results = await reader.results_async(query=AttackAnalyticsResultsQuery(), control=control) + assert len(results.items) == 1 + assert results.items[0].operation == "lock-test" + else: + facets = await reader.facets_async( + query=AttackAnalyticsFacetQuery(dimension=AttackAnalyticsDimension(name="operation")), control=control + ) + assert [item.key.value for item in facets.items] == ["lock-test"] + + +async def test_sqlite_busy_retry_yields_between_attempts_async() -> None: + operation = AsyncMock(side_effect=[_sqlite_error(sqlite3.SQLITE_BUSY), _sqlite_error(sqlite3.SQLITE_BUSY), 42]) + control = QueryControl(deadline=time.monotonic() + 10) + with patch.object(asyncio, "sleep", new_callable=AsyncMock) as sleep: + assert await AttackAnalyticsReader._retry_sqlite_busy_async(operation=operation, control=control) == 42 + assert operation.await_count == 3 + assert sleep.await_count == 2 + assert all(call.args == (AttackAnalyticsReader.SQLITE_BUSY_RETRY_DELAY,) for call in sleep.await_args_list) + + +@pytest.mark.parametrize( + "error", + [ + _sqlite_error(sqlite3.SQLITE_LOCKED), + _sqlite_error(sqlite3.SQLITE_BUSY_SNAPSHOT), + _sqlite_error(sqlite3.SQLITE_BUSY_RECOVERY), + _sqlite_error(sqlite3.SQLITE_ERROR), + OperationalError("SELECT", {}, RuntimeError("database is locked")), + ], +) +async def test_sqlite_busy_retry_preserves_other_errors_async(error: OperationalError) -> None: + operation = AsyncMock(side_effect=error) + with patch.object(asyncio, "sleep", new_callable=AsyncMock) as sleep: + with pytest.raises(OperationalError) as raised: + await AttackAnalyticsReader._retry_sqlite_busy_async( + operation=operation, control=QueryControl(deadline=time.monotonic() + 10) + ) + assert raised.value is error + operation.assert_awaited_once() + sleep.assert_not_awaited() + + +async def test_sqlite_busy_retry_expiry_prevents_another_attempt_async() -> None: + control = QueryControl(deadline=time.monotonic() + 10) + operation = AsyncMock(side_effect=_sqlite_error(sqlite3.SQLITE_BUSY)) + + async def expire_async(delay: float) -> None: + assert 0 < delay <= control.remaining + control.deadline = 0 + + with patch.object(asyncio, "sleep", side_effect=expire_async): + with pytest.raises(AnalyticsTimeoutException): + await AttackAnalyticsReader._retry_sqlite_busy_async(operation=operation, control=control) + operation.assert_awaited_once() + + +async def test_sqlite_busy_retry_caps_sleep_to_remaining_budget_async() -> None: + operation = AsyncMock(side_effect=_sqlite_error(sqlite3.SQLITE_BUSY)) + control = MagicMock(spec=QueryControl) + control.remaining = 0.002 + control.check.side_effect = [None, None, AnalyticsTimeoutException] + with patch.object(asyncio, "sleep", new_callable=AsyncMock) as sleep: + with pytest.raises(AnalyticsTimeoutException): + await AttackAnalyticsReader._retry_sqlite_busy_async(operation=operation, control=control) + sleep.assert_awaited_once_with(0.002) + operation.assert_awaited_once() + + +async def test_sqlite_read_completed_after_expiry_is_not_returned_async() -> None: + control = QueryControl(deadline=time.monotonic() + 10) + + async def completed_async() -> int: + control.deadline = 0 + return 42 + + with pytest.raises(AnalyticsTimeoutException): + await AttackAnalyticsReader._retry_sqlite_busy_async(operation=completed_async, control=control) + + +async def test_non_sqlite_execution_does_not_retry_sqlite_error_codes_async() -> None: + error = _sqlite_error(sqlite3.SQLITE_BUSY) + session = MagicMock(spec=AsyncSession) + session.get_bind.return_value.dialect.name = "mssql" + session.execute = AsyncMock(side_effect=error) + reader = AttackAnalyticsReader(memory=MagicMock(spec=SQLiteMemory)) + with patch.object(reader, "_retry_sqlite_busy_async", new_callable=AsyncMock) as retry: + with pytest.raises(OperationalError) as raised: + await reader._execute_async( + session=session, + statement=text("SELECT 1"), + control=QueryControl(deadline=time.monotonic() + 10), + ) + assert raised.value is error + session.execute.assert_awaited_once() + retry.assert_not_awaited() + + +def _observe_busy(*, reader: AttackAnalyticsReader, observed: asyncio.Event) -> Callable[..., Awaitable[_ResultT]]: + retry = reader._retry_sqlite_busy_async + + async def observe_async(*, operation: Callable[[], Awaitable[_ResultT]], control: QueryControl) -> _ResultT: + async def observed_operation_async() -> _ResultT: + try: + return await operation() + except OperationalError as error: + if getattr(error.orig, "sqlite_errorcode", None) == sqlite3.SQLITE_BUSY: + observed.set() + raise + + return await retry(operation=observed_operation_async, control=control) + + return observe_async + + +@pytest.mark.parametrize("method", ["report", "matrix", "compact_report", "results", "facets"]) +async def test_sqlite_read_succeeds_when_writer_releases_lock_async(*, file_memory: SQLiteMemory, method: str) -> None: + reader = AttackAnalyticsReader(memory=file_memory) + observed = asyncio.Event() + engine = file_memory._get_async_engine() + async with engine.connect() as writer: + default_timeout = (await writer.exec_driver_sql("PRAGMA busy_timeout")).scalar_one() + await writer.exec_driver_sql("BEGIN EXCLUSIVE") + with patch.object( + reader, "_retry_sqlite_busy_async", side_effect=_observe_busy(reader=reader, observed=observed) + ): + task = asyncio.create_task( + _read_async(reader=reader, method=method, control=QueryControl(deadline=time.monotonic() + 10)) + ) + try: + await asyncio.wait_for(observed.wait(), timeout=5) + assert not task.done() + await writer.rollback() + await asyncio.wait_for(task, timeout=5) + async with engine.connect() as connection: + assert (await connection.exec_driver_sql("PRAGMA busy_timeout")).scalar_one() == default_timeout + finally: + task.cancel() + await asyncio.gather(task, return_exceptions=True) + await writer.rollback() + + +@pytest.mark.parametrize("method", ["report", "matrix", "compact_report", "results", "facets"]) +@pytest.mark.parametrize("cancel", ["control", "task"]) +async def test_sqlite_lock_retry_cancellation_restores_connection_async( + *, file_memory: SQLiteMemory, method: str, cancel: str +) -> None: + reader = AttackAnalyticsReader(memory=file_memory) + observed = asyncio.Event() + control = QueryControl(deadline=time.monotonic() + 10) + engine = file_memory._get_async_engine() + async with engine.connect() as writer: + default_timeout = (await writer.exec_driver_sql("PRAGMA busy_timeout")).scalar_one() + await writer.exec_driver_sql("BEGIN EXCLUSIVE") + with patch.object( + reader, "_retry_sqlite_busy_async", side_effect=_observe_busy(reader=reader, observed=observed) + ): + task = asyncio.create_task(_read_async(reader=reader, method=method, control=control)) + try: + await asyncio.wait_for(observed.wait(), timeout=5) + if cancel == "control": + control.cancel() + with pytest.raises(AnalyticsTimeoutException): + await asyncio.wait_for(task, timeout=1) + else: + task.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(task, timeout=1) + async with engine.connect() as connection: + assert (await connection.exec_driver_sql("PRAGMA busy_timeout")).scalar_one() == default_timeout + finally: + task.cancel() + await asyncio.gather(task, return_exceptions=True) + await writer.rollback() + await _read_async(reader=reader, method=method, control=QueryControl(deadline=time.monotonic() + 10)) + + +@pytest.mark.parametrize("statement", ["PRAGMA journal_mode", "BEGIN"]) +async def test_sqlite_consistent_report_retries_setup_without_restarting_transaction_async( + *, file_memory: SQLiteMemory, statement: str +) -> None: + original = AsyncConnection.exec_driver_sql + attempts = 0 + + async def execute_async(self: AsyncConnection, sql: str) -> CursorResult[Any]: + nonlocal attempts + if sql == statement: + attempts += 1 + if attempts == 1: + raise _sqlite_error(sqlite3.SQLITE_BUSY) + return await original(self, sql) + + with patch.object(AsyncConnection, "exec_driver_sql", execute_async): + await _read_async( + reader=AttackAnalyticsReader(memory=file_memory), + method="report", + control=QueryControl(deadline=time.monotonic() + 10), + ) + assert attempts == 2 + + +@pytest.mark.parametrize("timeout_ms", [0, 127, 5000]) +async def test_sqlite_busy_timeout_is_zero_only_during_analytics_session_async( + *, file_memory: SQLiteMemory, timeout_ms: int +) -> None: + engine = file_memory._get_async_engine() + async with engine.connect() as connection: + await connection.exec_driver_sql(f"PRAGMA busy_timeout = {timeout_ms}") + reader = AttackAnalyticsReader(memory=file_memory) + async with reader._session_async(control=QueryControl(deadline=time.monotonic() + 10)) as (session, _, _): + connection = await session.connection() + assert (await connection.exec_driver_sql("PRAGMA busy_timeout")).scalar_one() == 0 + async with engine.connect() as connection: + assert (await connection.exec_driver_sql("PRAGMA busy_timeout")).scalar_one() == timeout_ms + + +@pytest.mark.parametrize("cancel_in_body", [False, True]) +async def test_sqlite_timeout_restoration_finishes_under_repeated_cancellation_async( + *, file_memory: SQLiteMemory, cancel_in_body: bool +) -> None: + reader = AttackAnalyticsReader(memory=file_memory) + entered, restoring, release_restore = asyncio.Event(), asyncio.Event(), asyncio.Event() + finish_body = asyncio.Event() + original_timeout = reader._sqlite_busy_timeout_async + async with file_memory._get_async_engine().connect() as connection: + default_timeout = (await connection.exec_driver_sql("PRAGMA busy_timeout")).scalar_one() + + async def restore_async(*, driver: SQLiteConnection, timeout_ms: int) -> None: + if timeout_ms: + restoring.set() + await release_restore.wait() + await original_timeout(driver=driver, timeout_ms=timeout_ms) + + async def read_async() -> None: + async with reader._session_async(control=QueryControl(deadline=time.monotonic() + 10)): + entered.set() + await finish_body.wait() + + with patch.object(reader, "_sqlite_busy_timeout_async", side_effect=restore_async): + task = asyncio.create_task(read_async()) + try: + await asyncio.wait_for(entered.wait(), timeout=5) + if cancel_in_body: + task.cancel("original cancellation") + else: + finish_body.set() + await asyncio.wait_for(restoring.wait(), timeout=5) + for index in range(2): + task.cancel( + "original cancellation" if not cancel_in_body and index == 0 else "cancel restoration again" + ) + await asyncio.sleep(0) + assert not task.done() + release_restore.set() + with pytest.raises(asyncio.CancelledError, match="original cancellation"): + await asyncio.wait_for(task, timeout=5) + finally: + release_restore.set() + task.cancel() + await asyncio.gather(task, return_exceptions=True) + async with file_memory._get_async_engine().connect() as connection: + assert (await connection.exec_driver_sql("PRAGMA busy_timeout")).scalar_one() == default_timeout + + +@pytest.mark.parametrize("failure", ["progress_handler", "busy_timeout"]) +@pytest.mark.parametrize("cancelled", [False, True]) +async def test_sqlite_reset_failure_discards_connection_and_preserves_cancellation_async( + *, failure: str, cancelled: bool +) -> None: + reader = AttackAnalyticsReader(memory=MagicMock(spec=SQLiteMemory)) + connection = MagicMock(spec=AsyncConnection) + connection.invalidate = AsyncMock() + driver = MagicMock(spec=SQLiteConnection) + error = RuntimeError("connection reset failed") + driver.set_progress_handler = AsyncMock(side_effect=error if failure == "progress_handler" else None) + cancellation = asyncio.CancelledError("original cancellation") if cancelled else None + with patch.object( + reader, + "_sqlite_busy_timeout_async", + new_callable=AsyncMock, + side_effect=error if failure == "busy_timeout" else None, + ): + with pytest.raises(asyncio.CancelledError if cancelled else RuntimeError) as raised: + await reader._restore_sqlite_settings_async( + connection=connection, driver=driver, busy_timeout=5000, cancelled=cancellation + ) + assert raised.value is (cancellation if cancelled else error) + if cancelled: + assert raised.value.__cause__ is error + connection.invalidate.assert_awaited_once_with(error) + + +@pytest.mark.parametrize("method", ["report", "matrix", "compact_report", "results", "facets"]) +async def test_every_analytics_projection_retries_busy_in_the_same_session_async( + *, file_memory: SQLiteMemory, method: str +) -> None: + original = AsyncSession.execute + calls: dict[Executable, int] = {} + sessions: set[AsyncSession] = set() + + async def execute_async(self: AsyncSession, statement: Executable) -> Result[Any]: + sessions.add(self) + calls[statement] = calls.get(statement, 0) + 1 + if calls[statement] == 1: + raise _sqlite_error(sqlite3.SQLITE_BUSY) + return await original(self, statement) + + with patch.object(AsyncSession, "execute", execute_async): + await _read_async( + reader=AttackAnalyticsReader(memory=file_memory), + method=method, + control=QueryControl(deadline=time.monotonic() + 10), + ) + assert calls and all(count == 2 for count in calls.values()) + assert len(sessions) == 1