Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 10 additions & 3 deletions doc/contributing/11_memory_models.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
177 changes: 145 additions & 32 deletions pyrit/memory/attack_analytics.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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,
Expand All @@ -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."""
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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.")
Expand All @@ -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))
Expand All @@ -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(
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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]],
Expand All @@ -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.
Expand All @@ -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))
Expand Down Expand Up @@ -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:
"""
Expand All @@ -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.
Expand All @@ -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.
Expand All @@ -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
Expand All @@ -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)
Expand All @@ -447,20 +545,24 @@ 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. "
"Consider explicitly configured WAL mode for concurrent workloads. "
"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):
Expand All @@ -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
Expand All @@ -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.
Expand All @@ -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 = (
Expand Down
Loading
Loading