From 81ab214e5a177601cbc851389962b5b814ee5433 Mon Sep 17 00:00:00 2001 From: Roman Lutz Date: Thu, 8 Oct 2026 00:21:58 -0700 Subject: [PATCH 1/8] FEAT: Add bounded async attack-result analytics SDK Interpret saved outcome counts through an async SDK with shared loop-bound admission, cooperative cancellation, and deterministic cleanup. Match bounded Unicode profile aggregation to the complete SQL fallback, with exact drilldowns, focused compatibility coverage, and SDK lifecycle documentation. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- doc/code/analytics/0_attack_results.md | 222 +++++++ doc/code/framework.md | 3 + doc/myst.yml | 1 + pyrit/analytics/__init__.py | 2 + pyrit/analytics/_execution.py | 249 +++++++ pyrit/analytics/_profile_aggregation.py | 250 +++++++ pyrit/analytics/attack_result_analytics.py | 366 +++++++++++ .../analytics/test_attack_result_analytics.py | 621 ++++++++++++++++++ tests/unit/analytics/test_execution.py | 380 +++++++++++ .../analytics/test_profile_aggregation.py | 414 ++++++++++++ .../unit/common/test_lazy_package_imports.py | 12 + 11 files changed, 2520 insertions(+) create mode 100644 doc/code/analytics/0_attack_results.md create mode 100644 pyrit/analytics/_execution.py create mode 100644 pyrit/analytics/_profile_aggregation.py create mode 100644 pyrit/analytics/attack_result_analytics.py create mode 100644 tests/unit/analytics/test_attack_result_analytics.py create mode 100644 tests/unit/analytics/test_execution.py create mode 100644 tests/unit/analytics/test_profile_aggregation.py diff --git a/doc/code/analytics/0_attack_results.md b/doc/code/analytics/0_attack_results.md new file mode 100644 index 0000000000..e229de40e0 --- /dev/null +++ b/doc/code/analytics/0_attack_results.md @@ -0,0 +1,222 @@ +# Attack-result analytics + +`AttackResultAnalytics` is an async Python SDK for saved attack outcomes. It reads +aggregate counts and lightweight metadata through the configured memory backend, +without loading conversations, messages, media, scores, or full `AttackResult` +objects. This API does not change History's existing result-selection behavior. +REST and GUI integration are separate work, not part of this SDK. + +## Counting policy + +Every distinct saved `attack_result_id` counts once. Different result IDs sharing +a conversation, including failed attempts followed by successful retries, remain +separate. Updating an existing ID changes its current outcome, not the result count. +Selecting a scenario is a cohort filter, not a switch to scenario-unit statistics. + +ASR uses `successes / (successes + failures)`, reusing the existing `AttackStats` +outcome policy. With no successes or failures, ASR is `None`. Errors and +undetermined outcomes remain visible but do not enter that denominator. +`total_results` includes all four outcomes; `outcome_shares` and `decided_share` +use that whole selected cohort. Empty outcome shares are zero, and an empty +cohort's decided share is `None`. All proportions are between 0 and 1. +Overall ASR is calculated from overall counts, never by averaging subgroup rates. + +This differs from latest-execution-unit scenario success statistics. A failed +attempt followed by a successful retry yields 50% raw-result ASR even if the +scenario's latest-unit success rate is 100%. No result-role filtering or inference +from conversation presence is applied. + +Outcome restrictions apply to the entire report and its result page. +`outcome_filter_applied` asks a renderer to annotate the rate as **ASR*** and +explain the restricted population. For example, success-only results have a 100% +ASR when nonempty, not evidence of a 100% unfiltered success rate. Selecting all +four outcomes normalizes to unrestricted selection and removes that annotation. + +## Python usage + +Initialize PyRIT before constructing the SDK. The constructor uses the configured +`CentralMemory` or accepts an explicit initialized `memory` instance. It performs +no queries, migration, repair, or database-setting changes. + +```python +from pyrit.analytics import AttackResultAnalytics +from pyrit.models import ( + AttackAnalyticsDimension, + AttackAnalyticsFacetQuery, + AttackAnalyticsFilter, + AttackAnalyticsFilters, + AttackAnalyticsQuery, + AttackAnalyticsResultsQuery, + AttackAnalyticsValue, +) + +async with AttackResultAnalytics() as analytics: + report = await analytics.query_async( + query=AttackAnalyticsQuery( + filters=AttackAnalyticsFilters( + dimensions=[ + AttackAnalyticsFilter( + dimension=AttackAnalyticsDimension(name="operation"), + values=[AttackAnalyticsValue(value="operation-a")], + ) + ] + ), + group_by=AttackAnalyticsDimension(name="targeted_harm_category"), + compare_by=AttackAnalyticsDimension(name="attack_type"), + ) + ) + print(report.summary.total_results, report.summary.success_rate) + + if report.drilldown_unavailable_reason: + print(report.drilldown_unavailable_reason) + elif report.cells: + cell = report.cells[0] + narrowed = report.filters.model_copy( + update={"dimensions": [*report.filters.dimensions, *cell.drilldown_filters]} + ) + page = await analytics.results_async(query=AttackAnalyticsResultsQuery(filters=narrowed)) + if page.has_more: + next_page = await analytics.results_async( + query=AttackAnalyticsResultsQuery(filters=narrowed, cursor=page.next_cursor) + ) + + options = await analytics.facets_async( + query=AttackAnalyticsFacetQuery( + filters=report.filters, + dimension=AttackAnalyticsDimension(name="converter_type", converter_direction="response"), + search="base64", + ) + ) +``` + +Omit `compare_by` for a one-dimensional breakdown. Supported dimensions include +operation, operator, targeted harm category, attack type, converter type, +objective target, model, scenario, and custom labels. Custom labels require a +literal `label_key`; dots in that key do not select nested JSON. Converter +dimensions select the request or response pipeline. + +Within a predicate, values use ANY matching by default. Converter predicates can +also request ALL matching. Separate predicates are AND-combined, even when they +refer to the same dimension. A chart drill-down therefore **appends** its one +group predicate or two cell predicates. Replacing an existing converter ANY +predicate could broaden the cohort instead of drilling into it. + +Queries allow at most 16 predicates, 500 values overall, and 100 values in one +predicate. A final legal click can reach the limit. The resulting report and +result pages remain valid, but `drilldown_unavailable_reason` explains why another +click would exceed the budget. Do not drop predicates to make room. +The SDK snapshots and revalidates nested request models before admission, +including UTC normalization of updated-date bounds. + +### Keys, labels, and overlapping groups + +Harm categories and converter types are multi-valued. Their groups overlap, but +duplicate memberships count each result only once per key or cell and once in the +overall cohort. `groups_overlap` warns against summing groups to reconstruct the +cohort total. Empty heatmap cells explicitly contain zero counts and `None` ASR. + +Use typed keys, not display labels, when filtering: + +- `VALUE` preserves real metadata, including blank strings and the literal + `"Unknown"`. Blank display text is labeled `"(Blank)"`. +- `MISSING` is labeled `"Not recorded"` and is not a string sentinel. +- `NO_CONVERTERS` is labeled `"No converters"` and distinguishes a known empty + pipeline from missing converter metadata. + +Case-insensitive categorical keys use the reader's backend semantics. In SQLite, +the SQL function and SDK profiles both use Python `str.lower`, including Unicode. +They do not use ASCII-only lowercasing or `casefold`. Display labels retain original +spellings, choosing the binary-smallest label for equivalent keys. + +Objective-target keys are persisted frozen-v1 behavioral evaluation hashes. +Deployment changes need not split behaviorally equivalent targets. The +`target_identifier_hash` in a result row is instead its exact content hash for +inspection. Neither key is recomputed from current registry state. Scenario keys +identify saved runs, not potentially repeated scenario names. The SDK adds short +identity suffixes to target/scenario display labels and forwards reader warnings. + +## Freshness and projection limits + +`query_async` reads overall counts, chart inputs, and the first result page in one +short consistent database view. The report and that page share `computed_at`. +Later `results_async` pages and `facets_async` lookups are separate fresh reads with +their own timestamps. Neither operation recalculates reports. No long-lived +snapshot freezes results between calls. + +Result cursors are bound to cohort filters and distinct-result selection. Invalid +or stale cursors raise an explicit error. A facet excludes predicates only on its +exact dimension, retaining other label keys, pipeline directions, outcome filters, +and date restrictions so alternatives remain discoverable. + +Defaults display 15 groups, 25 results, and 50 facet options. Groups are bounded +at 50, result/facet pages at 100, and each heatmap axis at 20. These are output +limits, not sampling: every matching saved result contributes to the statistics. +Updated-date filters are half-open intervals over last modification timestamps, +not attack execution-time trends. + +## Execution ownership and shutdown + +Reuse a long-lived SDK owner on one event loop, then await `close_async()` or exit +its async context before disposing or replacing memory. Do not create and close an +SDK context per request in an application with concurrent callers. + +Facades using the **same memory object** share one controller, even if constructed +independently. This prevents multiplying that backend's capacity by constructing +more facades. A controller rejects use or shutdown from another event loop rather +than creating another pool. Separate memory objects and processes have separate +budgets; applications should share one initialized memory object per backend. + +| Lane | Active operations | Additional queued calls | Queue wait | Execution budget | +|---|---:|---:|---:|---:| +| Reports | 5 | 10 | 1 second | 5 seconds | +| Result pages and facets, combined | 2 | 10 | 1 second | 1 second | + +These defaults are overload safeguards, **not latency guarantees**. Queues are +FIFO. Full or expired admission raises `AnalyticsBusyException`; execution expiry +raises `AnalyticsTimeoutException`. Native async readers own the `QueryControl` +database deadlines, connection acquisition, interruption, and session restoration. +The SDK adds no second SQL timeout mechanism, sync DB facade, or event-loop thread. + +Cancelling a caller or reaching its response deadline signals cooperative +cancellation. It does **not** release a still-running query's slot. Capacity +remains occupied until the operation and its session cleanup actually finish. +Closing rejects queued/new calls, signals active work, and waits for that cleanup, +even when it outlasts the response deadline. Cancelling `close_async` is propagated +only after draining; repeated cancellation cannot open an overlapping controller. + +Closing one bound facade closes the shared controller for every facade attached +to it. Those facades are terminal. Construct a new SDK owner only after closing +completes to begin another lifetime. Closing an unused facade does not affect +other facades. Memory engines remain caller-owned and are not disposed by the SDK. + +There is no in-flight coalescing, access-scope parameter, or persistent result +cache. Every request has independent result objects and cancellation. The SDK is +not an authorization layer; applications must enforce access policy themselves. + +## Bounded profile aggregation and backend limits + +For eligible SQLite category, attack-type, and converter-type charts, memory +probes at most 4,097 pre-counted metadata profiles. It accepts at most 4,096, +subject to 4,096-character source limits and a 1,000,000-character combined-text +limit. The SDK expands their deduplicated memberships and sums saved-outcome +weights, yielding between bounded batches. Converter arrays already contain the +reader's canonical names, including supported legacy names and missing members. + +Overflow selects the **complete SQL fallback**, never partial totals or sampling. +`profiles=None` means SQL supplied the chart; `profiles=[]` is a valid empty +cohort. The combined-text check happens after fetching the probe. These caps bound +transferred profiles and SDK work, not SQL scans, intermediate SQL work, or peak +network bytes. + +SQLite reports preserve reader warnings about rollback journaling and never enable +WAL themselves. Choose journaling and deployment settings explicitly outside the +SDK. SQL Server uses the reader's SQL path and requires its configured SNAPSHOT +support. Cross-backend collation details, malformed historical metadata, and live +Azure SQL validation remain storage/deployment concerns. + +Focused tests compare the SDK/profile path with the actual SQLite SQL fallback, +including Unicode, typed absence, duplicate/legacy memberships, and cap boundaries. +They also exercise weighted aggregation at the profile cap and deterministic +admission/cancellation/cleanup lifetimes. They are not a 100,000-row REST benchmark +or evidence of production latency. End-to-end application performance and runtime +integration require their own validation. diff --git a/doc/code/framework.md b/doc/code/framework.md index 09313f9db0..37a42bdc8e 100644 --- a/doc/code/framework.md +++ b/doc/code/framework.md @@ -328,6 +328,8 @@ The below talks about responsibilities of most modules in the PyRIT library - **Does not own**: live, in-attack decisions — any decision made *during* an attack is a scorer's job. Analytics only operates on stored results, after the fact. - Today it includes `ConversationAnalytics` (inspecting conversation history), `analyze_results` / `AttackStats` (aggregating outcomes across techniques), and text-matching strategies (`ExactTextMatching`, `ApproximateTextMatching`). - Shared analytics contracts (filters, dimensions, typed values, reports, facets, result pages, and `AttackStats`) live in `pyrit.models.analytics`. They validate data without querying memory or calculating statistics. `AttackResultSelection` defines selection modes without changing existing callers. +- [`AttackResultAnalytics`](./analytics/0_attack_results.md) provides async saved-result reports, lightweight result pages, and facet lookups. It reuses raw-outcome `AttackStats` policy, supplies display labels and exact additional drill-down predicates, and counts every saved result ID rather than latest scenario execution units. Memory owns cohort SQL and consistent projections; the SDK owns their interpretation. +- Analytics facades sharing a memory instance share one loop-bound, bounded report/quick-query controller. Close that shared SDK lifetime before replacing or disposing memory; cancellation does not free a running query's slot before its actual session cleanup. SDK lifecycle is independent of backend runtime integration. - Filter-bound cursor and label-normalization helpers live in `pyrit.common.pagination`. The backend pagination module retains compatibility exports, including History's invalid-cursor first-page fallback. ## Auth @@ -353,6 +355,7 @@ The below talks about responsibilities of most modules in the PyRIT library - Components should access memory through `CentralMemory` rather than passing state directly between each other. - Memory backends are swappable too (e.g. SQLite or Azure SQL) without changing the components that use them. - Memory loads and locks observation evidence for model-owned validation, and owns atomic writes and reference cleanup. +- Analytics readers return saved-outcome counts and lightweight metadata, or bounded typed profiles with a complete SQL fallback. They own native async sessions, consistent read views, and `QueryControl` database deadlines/interruption, not rate calculations or display labels. - **Does not own**: business logic or decisions. Memory stores and retrieves state; it doesn't decide what to send, how to score, or when to branch — components do that and persist results here. ## [Models](../contributing/11_memory_models) diff --git a/doc/myst.yml b/doc/myst.yml index 7e8a8c59f8..62ce6032b6 100644 --- a/doc/myst.yml +++ b/doc/myst.yml @@ -184,6 +184,7 @@ project: - file: code/registry/1_class_registry.ipynb - file: code/registry/2_instance_registry.ipynb - file: code/output/0_output.ipynb + - file: code/analytics/0_attack_results.md - file: api/index.md children: - file: api/pyrit_analytics.md diff --git a/pyrit/analytics/__init__.py b/pyrit/analytics/__init__.py index 4e3b933261..9d2ea1ac5e 100644 --- a/pyrit/analytics/__init__.py +++ b/pyrit/analytics/__init__.py @@ -9,6 +9,7 @@ from pyrit.common.lazy_imports import get_lazy_dir, resolve_lazy_export if TYPE_CHECKING: + from pyrit.analytics.attack_result_analytics import AttackResultAnalytics from pyrit.analytics.conversation_analytics import ConversationAnalytics from pyrit.analytics.result_analysis import ( AttackStats, @@ -21,6 +22,7 @@ _LAZY_EXPORTS: dict[str, str | tuple[str, str | None]] = { "analyze_results": "pyrit.analytics.result_analysis", "ApproximateTextMatching": "pyrit.analytics.text_matching", + "AttackResultAnalytics": "pyrit.analytics.attack_result_analytics", "AttackStats": "pyrit.analytics.result_analysis", "ConversationAnalytics": "pyrit.analytics.conversation_analytics", "ExactTextMatching": "pyrit.analytics.text_matching", diff --git a/pyrit/analytics/_execution.py b/pyrit/analytics/_execution.py new file mode 100644 index 0000000000..cd61c4d7dd --- /dev/null +++ b/pyrit/analytics/_execution.py @@ -0,0 +1,249 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Bound native async analytics work through its actual resource cleanup.""" + +from __future__ import annotations + +import asyncio +import logging +from collections import deque +from dataclasses import dataclass, field +from sys import float_info +from time import monotonic +from typing import TYPE_CHECKING, TypeVar + +from pyrit.common.task_utils import gather_with_cleanup_async +from pyrit.exceptions.analytics_exception import AnalyticsBusyException, AnalyticsTimeoutException +from pyrit.memory.query_control import QueryControl + +if TYPE_CHECKING: + from collections.abc import Awaitable, Callable + +logger = logging.getLogger(__name__) +T = TypeVar("T") + + +@dataclass(eq=False) +class _Waiter: + ready: asyncio.Future[None] + deadline: float + + +@dataclass(eq=False) +class _Operation: + """Keep the task strongly owned even after its requesting coroutine has left.""" + + control: QueryControl + finished: asyncio.Future[None] + abandoned: bool = False + task: Awaitable[object] | None = None + + +@dataclass +class _Lane: + limit: int + timeout: float + active: int = 0 + queued: deque[_Waiter] = field(default_factory=deque) + running: set[_Operation] = field(default_factory=set) + + +class AnalyticsExecution: + """ + Reserve independent report and quick-query capacity on one event loop. + + Admission is bounded before creating an operation task. A response deadline or + caller cancellation signals ``QueryControl`` but does not cancel that task: + native drivers and session cleanup must finish before its slot is released. + There is no coalescing, result cache, thread pool, or SQL timeout policy here. + The reader owns database interruption and session cleanup. + """ + + def __init__( + self, + *, + report_workers: int = 5, + quick_workers: int = 2, + max_queue: int = 10, + queue_timeout: float = 1.0, + report_timeout: float = 5.0, + quick_timeout: float = 1.0, + ) -> None: + """ + Bind a controller to the running loop without starting database work. + + Args: + report_workers (int): Concurrent report slots, including cleanup. + quick_workers (int): Separate slots shared by result pages and facets. + max_queue (int): Maximum waiting requests per lane, excluding active slots. + queue_timeout (float): Maximum admission wait in seconds. + report_timeout (float): Report execution budget after admission, in seconds. + quick_timeout (float): Result-page/facet execution budget, in seconds. + + Raises: + ValueError: If a count is not a positive integer or a timeout is not positive and finite. + RuntimeError: If construction occurs outside a running event loop. + """ + for name, value in ( + ("report_workers", report_workers), + ("quick_workers", quick_workers), + ("max_queue", max_queue), + ): + if type(value) is not int or value <= 0: + raise ValueError(f"{name} must be a positive integer.") + for name, timeout in ( + ("queue_timeout", queue_timeout), + ("report_timeout", report_timeout), + ("quick_timeout", quick_timeout), + ): + if isinstance(timeout, bool) or not isinstance(timeout, (int, float)) or not 0 < timeout <= float_info.max: + raise ValueError(f"{name} must be positive and finite.") + self._loop = asyncio.get_running_loop() + self._lanes = { + True: _Lane(limit=report_workers, timeout=report_timeout), + False: _Lane(limit=quick_workers, timeout=quick_timeout), + } + self._max_queue = max_queue + self._queue_timeout = queue_timeout + self._closing = False + self._closed = False + self._close_task: asyncio.Task[None] | None = None + + @property + def is_closed(self) -> bool: + """Whether shutdown has drained all started operations, not just their callers.""" + return self._closed + + async def run_async(self, *, report: bool, task: Callable[[QueryControl], Awaitable[T]]) -> T: + """ + Admit one operation and await its response without blocking the owning loop. + + Args: + report (bool): Use the report lane rather than the reserved quick lane. + task (Callable[[QueryControl], Awaitable[T]]): Native async work that honors + the shared control and does not return before releasing its resources. + + Returns: + T: This operation's result. Separate calls never share response objects. + + Raises: + AnalyticsBusyException: If the lane is full, admission expires, or shutdown has begun. + AnalyticsTimeoutException: If execution exceeds its budget. + CancelledError: If the caller leaves; running work retains its slot until cleanup finishes. + RuntimeError: If used from a different event loop. + """ + self._check_loop() + lane = self._lanes[report] + await self._admit_async(lane) + work = _Operation( + control=QueryControl(deadline=monotonic() + lane.timeout), + finished=self._loop.create_future(), + ) + lane.running.add(work) + operation = self._loop.create_task(self._execute_async(task=task, control=work.control)) + work.task = operation + operation.add_done_callback(lambda completed: self._complete(lane=lane, work=work, task=completed)) + try: + done, _ = await asyncio.wait({operation}, timeout=work.control.remaining) + if not done: + work.abandoned = True + work.control.cancel() + raise AnalyticsTimeoutException + return operation.result() + except asyncio.CancelledError: + work.abandoned = True + work.control.cancel() + raise + + async def close_async(self) -> None: + """ + Reject admission, signal cancellation, and drain actual work on the owning loop. + + This is idempotent and terminal. It can outlast response deadlines because + returning early would allow a replacement controller to overlap database + operations still cleaning up. Cancelling this await, even repeatedly, is + propagated only after draining. The memory backend itself is not disposed. + """ + self._check_loop() + if self._close_task is None: + self._closing = True + for lane in self._lanes.values(): + while lane.queued: + waiter = lane.queued.popleft() + if not waiter.ready.done(): + waiter.ready.set_exception(AnalyticsBusyException()) + for work in lane.running: + work.control.cancel() + self._close_task = self._loop.create_task(self._drain_async()) + cancellation: asyncio.CancelledError | None = None + while not self._close_task.done(): + try: + await asyncio.shield(self._close_task) + except asyncio.CancelledError as error: + cancellation = error + self._close_task.result() + if cancellation is not None: + raise cancellation + + def _check_loop(self) -> None: + if asyncio.get_running_loop() is not self._loop: + raise RuntimeError("Analytics must be used and closed on its owning event loop.") + + async def _admit_async(self, lane: _Lane) -> None: + if self._closing: + raise AnalyticsBusyException + if lane.active < lane.limit: + lane.active += 1 + return + if len(lane.queued) >= self._max_queue: + raise AnalyticsBusyException + waiter = _Waiter(ready=self._loop.create_future(), deadline=monotonic() + self._queue_timeout) + lane.queued.append(waiter) + try: + try: + await asyncio.wait_for(waiter.ready, timeout=self._queue_timeout) + except TimeoutError as error: + raise AnalyticsBusyException from error + if self._closing or monotonic() >= waiter.deadline: + raise AnalyticsBusyException + except BaseException: + if waiter in lane.queued: + lane.queued.remove(waiter) + elif not waiter.ready.cancelled() and waiter.ready.exception() is None: + self._release(lane) + raise + + def _release(self, lane: _Lane) -> None: + lane.active -= 1 + while lane.queued and not self._closing: + waiter = lane.queued.popleft() + if waiter.ready.done(): + continue + if monotonic() >= waiter.deadline: + waiter.ready.set_exception(AnalyticsBusyException()) + continue + lane.active += 1 + waiter.ready.set_result(None) + break + + def _complete(self, *, lane: _Lane, work: _Operation, task: asyncio.Task[T]) -> None: + if not task.cancelled(): + error = task.exception() + if work.abandoned and error is not None and not isinstance(error, AnalyticsTimeoutException): + logger.error("Analytics operation failed after its caller left.", exc_info=error) + lane.running.remove(work) + work.task = None + self._release(lane) + work.finished.set_result(None) + + async def _drain_async(self) -> None: + await gather_with_cleanup_async(work.finished for lane in self._lanes.values() for work in lane.running) + self._closed = True + + @staticmethod + async def _execute_async(*, task: Callable[[QueryControl], Awaitable[T]], control: QueryControl) -> T: + control.check() + result = await task(control) + control.check() + return result diff --git a/pyrit/analytics/_profile_aggregation.py b/pyrit/analytics/_profile_aggregation.py new file mode 100644 index 0000000000..aeeadcae07 --- /dev/null +++ b/pyrit/analytics/_profile_aggregation.py @@ -0,0 +1,250 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Finish bounded SQLite metadata profiles using the complete SQL path's membership rules.""" + +from __future__ import annotations + +import asyncio +import json +from collections import defaultdict +from dataclasses import dataclass +from typing import TYPE_CHECKING, ClassVar + +from pydantic import ValidationError + +from pyrit.exceptions.analytics_exception import AnalyticsDataException +from pyrit.memory.attack_analytics import RawAnalyticsGroup, RawAnalyticsOption +from pyrit.models import AttackAnalyticsDimensionName, AttackAnalyticsValue, AttackAnalyticsValueKind + +if TYPE_CHECKING: + from pyrit.memory.attack_analytics import RawAnalyticsProfile, RawAnalyticsReport + from pyrit.memory.query_control import QueryControl + from pyrit.models import AttackAnalyticsDimension, AttackAnalyticsQuery + +ValueKey = tuple[str, str] + + +@dataclass +class _Profile: + """One saved-outcome weight with deduplicated memberships, not hydrated results.""" + + values: list[dict[ValueKey, RawAnalyticsOption]] + outcome: str + weight: int + + +class ProfileAggregation: + """ + Expand only the reader-approved SQLite profiles, never an unbounded result set. + + Memory caps profile count and text size and falls back to complete SQL grouping + on overflow. Converter sources already contain canonical names, including + legacy-name precedence. This layer does not normalize identifier objects, + select results, change totals, or recompute target evaluation identities. + """ + + _CHECK_INTERVAL: ClassVar[int] = 128 + + @classmethod + async def populate_async( + cls, *, report: RawAnalyticsReport, query: AttackAnalyticsQuery, control: QueryControl + ) -> None: + """ + Populate chart projections while leaving overall counts and result rows untouched. + + ``profiles=None`` means SQL already supplied the chart; [] is a valid empty + cohort. Decoded options are cached only within this call. Bounded batches + yield to the event loop so cancellation and other callers can progress. + + Args: + report (RawAnalyticsReport): Reader output whose profiles passed every cap. + query (AttackAnalyticsQuery): The same grouping and limits used by the reader. + control (QueryControl): Shared execution budget, including SDK CPU work. + """ + if report.profiles is None: + return + dimensions = [query.group_by] + ([query.compare_by] if query.compare_by is not None else []) + profiles: list[_Profile] = [] + axes: list[dict[ValueKey, RawAnalyticsOption]] = [{} for _ in dimensions] + cache: list[dict[str | None, dict[ValueKey, RawAnalyticsOption]]] = [{} for _ in dimensions] + for index, record in enumerate(report.profiles): + if index % cls._CHECK_INTERVAL == 0: + await cls._checkpoint_async(control) + profile = cls._profile(record=record, dimensions=dimensions, cache=cache) + profiles.append(profile) + for axis, values in zip(axes, profile.values, strict=True): + for key, option in values.items(): + cls._add_option(axis=axis, key=key, option=option) + if query.compare_by is None: + await cls._groups_async(report=report, query=query, profiles=profiles, axis=axes[0], control=control) + else: + await cls._matrix_async(report=report, query=query, profiles=profiles, axes=axes, control=control) + control.check() + + @classmethod + def _profile( + cls, + *, + record: RawAnalyticsProfile, + dimensions: list[AttackAnalyticsDimension], + cache: list[dict[str | None, dict[ValueKey, RawAnalyticsOption]]], + ) -> _Profile: + if record["oversized"]: + raise AnalyticsDataException("Oversized metadata profiles require complete SQL aggregation.") + sources = [record["source0"]] + if len(dimensions) == 2: + if "source1" not in record: + raise AnalyticsDataException("A comparison profile is missing its second dimension.") + sources.append(record["source1"]) + values = [] + for raw, dimension, options in zip(sources, dimensions, cache, strict=True): + if raw not in options: + options[raw] = cls._values(raw=raw, dimension=dimension) + values.append(options[raw]) + return _Profile(values=values, outcome=record["outcome"], weight=record["weight"]) + + @classmethod + def _values(cls, *, raw: str | None, dimension: AttackAnalyticsDimension) -> dict[ValueKey, RawAnalyticsOption]: + """ + Decode scalar attack names or canonical category/converter name arrays. + + JSON null and null members mean missing metadata. Empty harm arrays are + missing; empty converter arrays mean a known pipeline with no converters. + Scalar attack names and real blank strings are not JSON absence markers. + + Returns: + dict[ValueKey, RawAnalyticsOption]: One representative option per typed key. + + Raises: + AnalyticsDataException: If an unsupported dimension or malformed array reaches this path. + """ + if dimension.name is AttackAnalyticsDimensionName.ATTACK_TYPE: + return cls._options([raw]) + if dimension.name not in { + AttackAnalyticsDimensionName.TARGETED_HARM_CATEGORY, + AttackAnalyticsDimensionName.CONVERTER_TYPE, + }: + raise AnalyticsDataException("This dimension does not support compact profile aggregation.") + if raw is None or raw == "null": + return cls._options([None]) + try: + values = json.loads(raw) + except json.JSONDecodeError as error: + raise AnalyticsDataException("Stored category/converter metadata is invalid JSON.") from error + if not isinstance(values, list): + raise AnalyticsDataException("Stored category/converter metadata is not an array.") + if not values: + kind = ( + AttackAnalyticsValueKind.NO_CONVERTERS + if dimension.name is AttackAnalyticsDimensionName.CONVERTER_TYPE + else AttackAnalyticsValueKind.MISSING + ) + return {(kind.value, ""): RawAnalyticsOption(key=AttackAnalyticsValue(kind=kind), label=None)} + return cls._options(values) + + @classmethod + def _options(cls, values: list[object]) -> dict[ValueKey, RawAnalyticsOption]: + """ + Match SQLite's registered UnicodeLower function, which uses Python str.lower. + + Do not use ASCII translation or casefold. Keys fold Unicode but display + labels retain the binary-smallest original spelling, matching SQL MIN. + + Returns: + dict[ValueKey, RawAnalyticsOption]: Deduplicated memberships. + + Raises: + AnalyticsDataException: If a member is not text/NULL or its folded key exceeds the contract. + """ + options: dict[ValueKey, RawAnalyticsOption] = {} + for label in values: + if label is not None and not isinstance(label, str): + raise AnalyticsDataException("Stored category/converter metadata contains a non-string value.") + try: + key = ( + AttackAnalyticsValue(kind=AttackAnalyticsValueKind.MISSING) + if label is None + else AttackAnalyticsValue(value=label.lower()) + ) + except ValidationError as error: + raise AnalyticsDataException( + f"Stored attack metadata exceeds the {AttackAnalyticsValue.MAX_VALUE_LENGTH:,}-character " + "analytics key limit. Choose another dimension or inspect the saved result." + ) from error + cls._add_option( + axis=options, key=(key.kind.value, key.value or ""), option=RawAnalyticsOption(key=key, label=label) + ) + return options + + @classmethod + def _add_option( + cls, *, axis: dict[ValueKey, RawAnalyticsOption], key: ValueKey, option: RawAnalyticsOption + ) -> None: + axis[key] = cls._minimum_option(current=axis.get(key), option=option) + + @staticmethod + def _minimum_option(*, current: RawAnalyticsOption | None, option: RawAnalyticsOption) -> RawAnalyticsOption: + if current is None or (option.label is not None and (current.label is None or option.label < current.label)): + return option + return current + + @classmethod + async def _groups_async( + cls, + *, + report: RawAnalyticsReport, + query: AttackAnalyticsQuery, + profiles: list[_Profile], + axis: dict[ValueKey, RawAnalyticsOption], + control: QueryControl, + ) -> None: + counts: dict[ValueKey, dict[str, int]] = defaultdict(lambda: defaultdict(int)) + for index, profile in enumerate(profiles): + if index % cls._CHECK_INTERVAL == 0: + await cls._checkpoint_async(control) + for key in profile.values[0]: + counts[key][profile.outcome] += profile.weight + ordered = sorted(counts, key=lambda key: (-sum(counts[key].values()), key)) + end = query.group_offset + query.group_limit + report.has_more_groups = len(ordered) > end + report.groups = [ + RawAnalyticsGroup(option=axis[key], counts=dict(counts[key])) for key in ordered[query.group_offset : end] + ] + + @classmethod + async def _matrix_async( + cls, + *, + report: RawAnalyticsReport, + query: AttackAnalyticsQuery, + profiles: list[_Profile], + axes: list[dict[ValueKey, RawAnalyticsOption]], + control: QueryControl, + ) -> None: + """Bound axes before expanding pairs; preserve cell-local labels as well as axis labels.""" + selected = [set(sorted(axis)[: query.axis_limit]) for axis in axes] + report.axes_truncated = any(len(axis) > query.axis_limit for axis in axes) + report.rows = [axes[0][key] for key in sorted(selected[0])] + report.columns = [axes[1][key] for key in sorted(selected[1])] + cells: dict[tuple[ValueKey, ValueKey], RawAnalyticsGroup] = {} + for index, profile in enumerate(profiles): + if index % cls._CHECK_INTERVAL == 0: + await cls._checkpoint_async(control) + for row in profile.values[0].keys() & selected[0]: + for column in profile.values[1].keys() & selected[1]: + pair = row, column + if pair not in cells: + cells[pair] = RawAnalyticsGroup( + option=profile.values[0][row], column=profile.values[1][column], counts={} + ) + cell = cells[pair] + cell.counts[profile.outcome] = cell.counts.get(profile.outcome, 0) + profile.weight + cell.option = cls._minimum_option(current=cell.option, option=profile.values[0][row]) + cell.column = cls._minimum_option(current=cell.column, option=profile.values[1][column]) + report.cells = [cells[pair] for pair in sorted(cells)] + + @staticmethod + async def _checkpoint_async(control: QueryControl) -> None: + await asyncio.sleep(0) + control.check() diff --git a/pyrit/analytics/attack_result_analytics.py b/pyrit/analytics/attack_result_analytics.py new file mode 100644 index 0000000000..5ebfa62a76 --- /dev/null +++ b/pyrit/analytics/attack_result_analytics.py @@ -0,0 +1,366 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Async SDK interpretation of persisted AttackResult outcomes.""" + +from __future__ import annotations + +import threading +from typing import TYPE_CHECKING, ClassVar, Self, TypeVar +from weakref import WeakKeyDictionary + +from pydantic import BaseModel + +from pyrit.analytics._execution import AnalyticsExecution +from pyrit.analytics._profile_aggregation import ProfileAggregation +from pyrit.analytics.result_analysis import _compute_stats +from pyrit.exceptions.analytics_exception import AnalyticsBusyException, AnalyticsDataException +from pyrit.memory import CentralMemory +from pyrit.memory.attack_analytics import AttackAnalyticsReader, RawAnalyticsOption, RawAnalyticsReport +from pyrit.models import ( + AttackAnalyticsCell, + AttackAnalyticsDimension, + AttackAnalyticsDimensionName, + AttackAnalyticsFacetQuery, + AttackAnalyticsFacets, + AttackAnalyticsFilter, + AttackAnalyticsFilters, + AttackAnalyticsGroup, + AttackAnalyticsOption, + AttackAnalyticsQuery, + AttackAnalyticsReport, + AttackAnalyticsResults, + AttackAnalyticsResultsQuery, + AttackAnalyticsStatistics, + AttackAnalyticsValue, + AttackAnalyticsValueKind, + AttackOutcome, +) + +if TYPE_CHECKING: + from types import TracebackType + + from pyrit.memory import MemoryInterface + from pyrit.memory.query_control import QueryControl + +QueryT = TypeVar("QueryT", bound=BaseModel) + + +class AttackResultAnalytics: + """ + Analyze every distinct saved result ID, not conversations or scenario execution units. + + Facades sharing a memory instance share one loop-bound execution budget. + Reuse a long-lived SDK owner, then await ``close_async`` on that same loop + before disposing or replacing memory. Closing drains and terminates that + shared controller, including work submitted through other facades. + No authorization, in-flight reuse, or persistent result caching is performed. + """ + + _EXECUTIONS: ClassVar[WeakKeyDictionary[MemoryInterface, AnalyticsExecution]] = WeakKeyDictionary() + _EXECUTION_LOCK: ClassVar[threading.Lock] = threading.Lock() + _MULTIVALUED: ClassVar[frozenset[AttackAnalyticsDimensionName]] = frozenset( + {AttackAnalyticsDimensionName.TARGETED_HARM_CATEGORY, AttackAnalyticsDimensionName.CONVERTER_TYPE} + ) + + def __init__(self, *, memory: MemoryInterface | None = None) -> None: + """ + Retain an initialized backend without querying it or changing its schema/settings. + + Args: + memory (MemoryInterface | None): Backend to read, or the configured + CentralMemory instance. Its resource lifecycle remains caller-owned. + """ + self._memory = memory if memory is not None else CentralMemory.get_memory_instance() + self._reader = AttackAnalyticsReader(memory=self._memory) + self._controller: AnalyticsExecution | None = None + self._closed = False + + async def __aenter__(self) -> Self: + """ + Acquire the shared execution lifetime without opening a database session. + + Returns: + Self: The loop-bound SDK facade. + """ + self._execution()._check_loop() + return self + + async def __aexit__( + self, exc_type: type[BaseException] | None, exc: BaseException | None, traceback: TracebackType | None + ) -> None: + """Drain the shared analytics lifetime before leaving the context.""" + await self.close_async() + + async def query_async(self, *, query: AttackAnalyticsQuery | None = None) -> AttackAnalyticsReport: + """ + Compute a coherent report and its first lightweight result page. + + Args: + query (AttackAnalyticsQuery | None): Cohort, dimensions, and output limits. + None selects all saved result IDs grouped by operation. A deep, + revalidated snapshot is taken before admission or queueing. + + Returns: + AttackAnalyticsReport: Caller-owned counts, rates, annotations and exact + additional drill-down predicates. No snapshot is retained for later calls. + + Raises: + AnalyticsBusyException: If admission is full/expired or this lifetime is closing. + AnalyticsTimeoutException: If the shared database/SDK execution budget expires. + AnalyticsDataException: If stored data cannot be represented faithfully. + RuntimeError: If used from a different event loop. + """ + request = self._copy(query if query is not None else AttackAnalyticsQuery()) + return await self._execution().run_async( + report=True, task=lambda control: self._report_async(query=request, control=control) + ) + + async def results_async(self, *, query: AttackAnalyticsResultsQuery | None = None) -> AttackAnalyticsResults: + """ + Fetch a fresh metadata page without recalculating any report. + + Args: + query (AttackAnalyticsResultsQuery | None): Cohort and opaque cursor. + None requests the first page of all saved results. + + Returns: + AttackAnalyticsResults: Caller-owned result projections with their own freshness timestamp. + + Raises: + ValueError: If the cursor is malformed or belongs to different filters. + AnalyticsBusyException: If quick-query admission fails or shutdown has begun. + AnalyticsTimeoutException: If the operation exceeds its execution budget. + """ + request = self._copy(query if query is not None else AttackAnalyticsResultsQuery()) + return await self._execution().run_async( + report=False, task=lambda control: self._reader.results_async(query=request, control=control) + ) + + async def facets_async(self, *, query: AttackAnalyticsFacetQuery) -> AttackAnalyticsFacets: + """ + Look up one facet page without calculating a report or fetching other facets. + + Args: + query (AttackAnalyticsFacetQuery): Dimension, search, page, and cohort. + Only this exact dimension's predicates are omitted when finding alternatives. + + Returns: + AttackAnalyticsFacets: Typed options and the next offset, if any. + """ + request = self._copy(query) + return await self._execution().run_async( + report=False, task=lambda control: self._facets_async(query=request, control=control) + ) + + async def close_async(self) -> None: + """ + Drain the backend's shared controller; never dispose caller-owned memory. + + Close on the owning loop, after all users of this memory's analytics + lifetime have finished. Cancellation of this await is delayed until actual + operation/session cleanup completes. New controllers remain forbidden + during draining. Used facades cannot reopen; construct a new SDK instance + after closing to begin another lifetime. Closing an unused facade is a no-op + for other facades and still permanently closes that unused facade. + """ + if self._controller is None: + self._closed = True + return + execution = self._controller + try: + await execution.close_async() + finally: + if execution.is_closed: + self._closed = True + with self._EXECUTION_LOCK: + if self._EXECUTIONS.get(self._memory) is execution: + del self._EXECUTIONS[self._memory] + + def _execution(self) -> AnalyticsExecution: + if self._closed: + raise AnalyticsBusyException + if self._controller is None: + with self._EXECUTION_LOCK: + execution = self._EXECUTIONS.get(self._memory) + if execution is None: + execution = AnalyticsExecution() + self._EXECUTIONS[self._memory] = execution + self._controller = execution + return self._controller + + async def _report_async(self, *, query: AttackAnalyticsQuery, control: QueryControl) -> AttackAnalyticsReport: + raw = await self._reader.report_async(query=query, control=control, use_compact_profiles=True) + summary = self._statistics(raw.counts) + await ProfileAggregation.populate_async(report=raw, query=query, control=control) + report = AttackAnalyticsReport( + filters=query.filters, + group_by=query.group_by, + compare_by=query.compare_by, + summary=summary, + outcome_filter_applied=bool(query.filters.outcomes), + groups_overlap=query.group_by.name in self._MULTIVALUED + or (query.compare_by is not None and query.compare_by.name in self._MULTIVALUED), + drilldown_unavailable_reason=self._drilldown_unavailable_reason(query), + groups=[ + AttackAnalyticsGroup( + **self._option(raw=group.option, dimension=query.group_by).model_dump(), + statistics=self._statistics(group.counts), + drilldown_filters=[self._drilldown(dimension=query.group_by, key=group.option.key)], + ) + for group in raw.groups + ], + has_more_groups=raw.has_more_groups, + next_group_offset=query.group_offset + len(raw.groups) if raw.has_more_groups else None, + rows=[self._option(raw=option, dimension=query.group_by) for option in raw.rows], + columns=[self._option(raw=option, dimension=query.compare_by) for option in raw.columns] + if query.compare_by is not None + else [], + cells=self._cells(query=query, raw=raw), + axes_truncated=raw.axes_truncated, + results=raw.results, + computed_at=raw.results.computed_at, + warnings=raw.warnings, + ) + control.check() + return report + + async def _facets_async(self, *, query: AttackAnalyticsFacetQuery, control: QueryControl) -> AttackAnalyticsFacets: + raw = await self._reader.facets_async(query=query, control=control) + return AttackAnalyticsFacets( + items=[self._option(raw=option, dimension=query.dimension) for option in raw.items], + has_more=raw.has_more, + next_offset=query.offset + len(raw.items) if raw.has_more else None, + computed_at=raw.computed_at, + ) + + @classmethod + def _cells(cls, *, query: AttackAnalyticsQuery, raw: RawAnalyticsReport) -> list[AttackAnalyticsCell]: + """ + Fill the bounded matrix, including empty cells with unavailable rather than zero ASR. + + Returns: + list[AttackAnalyticsCell]: One cell per visible axis pair, with two additional predicates. + """ + if query.compare_by is None: + return [] + counts = { + (cls._value_key(cell.option.key), cls._value_key(cell.column.key)): cell.counts + for cell in raw.cells + if cell.column is not None + } + return [ + AttackAnalyticsCell( + row=row.key, + column=column.key, + statistics=cls._statistics(counts.get((cls._value_key(row.key), cls._value_key(column.key)), {})), + drilldown_filters=[ + cls._drilldown(dimension=query.group_by, key=row.key), + cls._drilldown(dimension=query.compare_by, key=column.key), + ], + ) + for row in raw.rows + for column in raw.columns + ] + + @staticmethod + def _statistics(counts: dict[str, int]) -> AttackAnalyticsStatistics: + """ + Reuse raw-outcome ASR policy; total shares include errors and undetermined results. + + Returns: + AttackAnalyticsStatistics: The existing decided-result rate plus whole-cohort shares. + + Raises: + AnalyticsDataException: If saved outcomes or their counts are invalid. + """ + if set(counts) - {outcome.value for outcome in AttackOutcome}: + raise AnalyticsDataException("Stored results contain an unsupported attack outcome.") + if any(type(count) is not int or count < 0 for count in counts.values()): + raise AnalyticsDataException("Stored results contain invalid outcome counts.") + stats = _compute_stats( + successes=counts.get(AttackOutcome.SUCCESS.value, 0), + failures=counts.get(AttackOutcome.FAILURE.value, 0), + undetermined=counts.get(AttackOutcome.UNDETERMINED.value, 0), + errors=counts.get(AttackOutcome.ERROR.value, 0), + ) + total = stats.total_decided + stats.undetermined + stats.errors + return AttackAnalyticsStatistics( + success_rate=stats.success_rate, + total_decided=stats.total_decided, + successes=stats.successes, + failures=stats.failures, + undetermined=stats.undetermined, + errors=stats.errors, + total_results=total, + decided_share=stats.total_decided / total if total else None, + outcome_shares={ + outcome: counts.get(outcome.value, 0) / total if total else 0.0 for outcome in AttackOutcome + }, + ) + + @staticmethod + def _option(*, raw: RawAnalyticsOption, dimension: AttackAnalyticsDimension) -> AttackAnalyticsOption: + """ + Keep stored typed identity independent of absence labels or distinguishing display suffixes. + + Returns: + AttackAnalyticsOption: The unchanged key and its SDK display label. + """ + if raw.key.kind is AttackAnalyticsValueKind.MISSING: + label = "Not recorded" + elif raw.key.kind is AttackAnalyticsValueKind.NO_CONVERTERS: + label = "No converters" + else: + label = raw.label if raw.label is not None else raw.key.value or "" + if not label.strip(): + label = "(Blank)" + elif ( + dimension.name in {AttackAnalyticsDimensionName.OBJECTIVE_TARGET, AttackAnalyticsDimensionName.SCENARIO} + and raw.key.value is not None + and label != raw.key.value + ): + label = f"{label} ({raw.key.value[-8:]})" + return AttackAnalyticsOption(key=raw.key, label=label) + + @staticmethod + def _drilldown_unavailable_reason(query: AttackAnalyticsQuery) -> str | None: + """ + Allow a final legal click to reach the budget; only disable the following click. + + Returns: + str | None: An explanation when this chart's additional predicates no longer fit. + """ + additional = 2 if query.compare_by is not None else 1 + if len(query.filters.dimensions) + additional > AttackAnalyticsFilters.MAX_PREDICATES: + return ( + "Remove a dimension filter before drilling down further " + f"(maximum {AttackAnalyticsFilters.MAX_PREDICATES} predicates)." + ) + if ( + sum(len(predicate.values) for predicate in query.filters.dimensions) + additional + > AttackAnalyticsFilters.MAX_VALUES + ): + return ( + "Select fewer dimension values before drilling down further " + f"(maximum {AttackAnalyticsFilters.MAX_VALUES} values)." + ) + return None + + @staticmethod + def _drilldown(*, dimension: AttackAnalyticsDimension, key: AttackAnalyticsValue) -> AttackAnalyticsFilter: + return AttackAnalyticsFilter(dimension=dimension, values=[key]) + + @staticmethod + def _value_key(value: AttackAnalyticsValue) -> tuple[AttackAnalyticsValueKind, str | None]: + return value.kind, value.value + + @staticmethod + def _copy(query: QueryT) -> QueryT: + """ + Snapshot and revalidate nested mutable queries before consuming any admission capacity. + + Returns: + QueryT: Independent normalized data, including UTC bounds and nested filter values. + """ + return type(query).model_validate(query.model_dump()) diff --git a/tests/unit/analytics/test_attack_result_analytics.py b/tests/unit/analytics/test_attack_result_analytics.py new file mode 100644 index 0000000000..a9e1981cf3 --- /dev/null +++ b/tests/unit/analytics/test_attack_result_analytics.py @@ -0,0 +1,621 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +from __future__ import annotations + +import asyncio +from datetime import UTC, datetime, timedelta, timezone +from typing import TYPE_CHECKING +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from pydantic import ValidationError + +from pyrit.analytics import AttackResultAnalytics +from pyrit.common.task_utils import gather_with_cleanup_async +from pyrit.exceptions.analytics_exception import AnalyticsBusyException, AnalyticsDataException +from pyrit.memory import CentralMemory, MemoryInterface, SQLiteMemory +from pyrit.memory.analytics_identity_v1 import ObjectiveTargetAnalyticsIdentityV1 +from pyrit.memory.attack_analytics import RawAnalyticsOption +from pyrit.models import ( + AtomicAttackIdentifier, + AttackAnalyticsDimension, + AttackAnalyticsFacetQuery, + AttackAnalyticsFilter, + AttackAnalyticsFilters, + AttackAnalyticsQuery, + AttackAnalyticsResultsQuery, + AttackAnalyticsValue, + AttackIdentifier, + AttackOutcome, + TargetIdentifier, +) +from unit.memory.test_attack_analytics import make_result, predicate + +if TYPE_CHECKING: + from collections.abc import AsyncGenerator + + from sqlalchemy.ext.asyncio import AsyncSession + + from pyrit.memory.query_control import QueryControl + + +@pytest.fixture +async def analytics(sqlite_instance: SQLiteMemory) -> AsyncGenerator[AttackResultAnalytics, None]: + async with AttackResultAnalytics(memory=sqlite_instance) as analytics: + yield analytics + + +@pytest.fixture +async def mixed_results(sqlite_instance: SQLiteMemory) -> None: + outcomes = [AttackOutcome.SUCCESS] * 4 + [AttackOutcome.FAILURE] * 2 + outcomes += [AttackOutcome.UNDETERMINED] * 3 + [AttackOutcome.ERROR] + await sqlite_instance.add_attack_results_to_memory_async( + attack_results=[make_result(index=index, outcome=outcome) for index, outcome in enumerate(outcomes, 1)] + ) + + +@pytest.mark.usefixtures("mixed_results") +async def test_all_saved_ids_and_outcomes_use_raw_decided_denominator(analytics: AttackResultAnalytics) -> None: + report = await analytics.query_async() + assert report.summary.total_results == 10 + assert report.summary.total_decided == 6 + assert report.summary.success_rate == pytest.approx(4 / 6) + assert report.summary.decided_share == 0.6 + assert report.summary.outcome_shares == { + AttackOutcome.SUCCESS: 0.4, + AttackOutcome.FAILURE: 0.2, + AttackOutcome.UNDETERMINED: 0.3, + AttackOutcome.ERROR: 0.1, + } + assert not report.outcome_filter_applied + assert not report.groups_overlap + assert report.groups[0].statistics == report.summary + assert len(report.results.items) == 10 + assert report.computed_at == report.results.computed_at + + +async def test_empty_report_has_unavailable_rates(analytics: AttackResultAnalytics) -> None: + report = await analytics.query_async() + assert report.summary.total_results == report.summary.total_decided == 0 + assert report.summary.success_rate is None + assert report.summary.decided_share is None + assert set(report.summary.outcome_shares.values()) == {0.0} + assert report.groups == report.cells == report.results.items == [] + + +@pytest.mark.usefixtures("mixed_results") +@pytest.mark.parametrize( + ("outcome", "count", "rate"), + [ + (AttackOutcome.SUCCESS, 4, 1.0), + (AttackOutcome.FAILURE, 2, 0.0), + (AttackOutcome.ERROR, 1, None), + (AttackOutcome.UNDETERMINED, 3, None), + ], +) +async def test_outcome_filter_annotates_and_restricts_entire_report( + analytics: AttackResultAnalytics, outcome: AttackOutcome, count: int, rate: float | None +) -> None: + report = await analytics.query_async(query=AttackAnalyticsQuery(filters=AttackAnalyticsFilters(outcomes=[outcome]))) + assert report.outcome_filter_applied + assert report.summary.total_results == count + assert report.summary.success_rate == rate + assert report.groups[0].statistics == report.summary + assert all(row.outcome == outcome for row in report.results.items) + + +@pytest.mark.usefixtures("mixed_results") +async def test_all_outcomes_normalize_to_unrestricted(analytics: AttackResultAnalytics) -> None: + report = await analytics.query_async( + query=AttackAnalyticsQuery(filters=AttackAnalyticsFilters(outcomes=list(AttackOutcome))) + ) + assert not report.outcome_filter_applied + assert report.filters.outcomes == [] + assert report.summary.total_results == 10 + + +async def test_overall_rate_is_not_an_average_of_groups( + analytics: AttackResultAnalytics, sqlite_instance: SQLiteMemory +) -> None: + await sqlite_instance.add_attack_results_to_memory_async( + attack_results=[ + make_result(operation="one"), + *[make_result(index=index, operation="two", outcome=AttackOutcome.FAILURE) for index in (2, 3, 4)], + ] + ) + with patch.object( + sqlite_instance, "get_attack_results_async", side_effect=AssertionError("Must not hydrate saved results") + ): + report = await analytics.query_async() + assert report.summary.success_rate == 0.25 + assert {group.statistics.success_rate for group in report.groups} == {0.0, 1.0} + assert len(report.results.items) == 4 + + +async def test_drilldown_appends_to_existing_any_and_response_all_predicates( + analytics: AttackResultAnalytics, sqlite_instance: SQLiteMemory +) -> None: + await sqlite_instance.add_attack_results_to_memory_async( + attack_results=[ + make_result(converters=["Alpha", "Gamma"], response_converters=["X", "Y"]), + make_result(index=2, converters=["Gamma"], response_converters=["X", "Y"]), + make_result(index=3, converters=["Beta"], response_converters=["X", "Y"]), + make_result(index=4, converters=["Alpha", "Gamma"], response_converters=["X"]), + ] + ) + filters = AttackAnalyticsFilters.model_validate( + { + "dimensions": [ + predicate(name="converter_type", values=["Alpha", "Beta"]), + predicate( + name="converter_type", + values=["X", "Y"], + dimension_options={"converter_direction": "response"}, + match_mode="all", + ), + ] + } + ) + report = await analytics.query_async( + query=AttackAnalyticsQuery(filters=filters, group_by=AttackAnalyticsDimension(name="converter_type")) + ) + assert report.summary.total_results == 2 + assert report.groups_overlap + assert sum(group.statistics.total_results for group in report.groups) == 3 + gamma = next(group for group in report.groups if group.key.value == "gamma") + narrowed = filters.model_copy(update={"dimensions": [*filters.dimensions, *gamma.drilldown_filters]}) + results = await analytics.results_async(query=AttackAnalyticsResultsQuery(filters=narrowed)) + assert len(results.items) == gamma.statistics.total_results == 1 + assert results.items[0].request_converters == ["Alpha", "Gamma"] + assert len(gamma.drilldown_filters) == 1 + assert filters == report.filters + + +@pytest.mark.parametrize("compare", [False, True]) +@pytest.mark.parametrize("budget", ["predicates", "values"]) +@pytest.mark.parametrize("remaining", [0, 1, 2]) +async def test_terminal_filter_budget_preserves_current_report_and_final_legal_click( + analytics: AttackResultAnalytics, sqlite_instance: SQLiteMemory, compare: bool, budget: str, remaining: int +) -> None: + await sqlite_instance.add_attack_results_to_memory_async(attack_results=[make_result(categories=["privacy"])]) + sizes = [1] * (16 - remaining) if budget == "predicates" else [100, 100, 100, 100, 100 - remaining] + filters = AttackAnalyticsFilters.model_validate( + { + "dimensions": [ + { + "dimension": {"name": "label", "label_key": f"absent-{index}"}, + "values": [{"kind": "missing"}, *({"value": f"value-{item}"} for item in range(size - 1))], + } + for index, size in enumerate(sizes) + ] + } + ) + query = AttackAnalyticsQuery( + filters=filters, + compare_by=AttackAnalyticsDimension(name="targeted_harm_category") if compare else None, + ) + report = await analytics.query_async(query=query) + assert report.summary.total_results == len(report.results.items) == 1 + assert report.filters == filters + extra = report.cells[0].drilldown_filters if compare else report.groups[0].drilldown_filters + assert len(extra) == (2 if compare else 1) + narrowed = filters.model_copy(update={"dimensions": [*filters.dimensions, *extra]}) + if remaining < len(extra): + assert f"maximum {16 if budget == 'predicates' else 500}" in report.drilldown_unavailable_reason + with pytest.raises(ValidationError): + await analytics.results_async(query=AttackAnalyticsResultsQuery(filters=narrowed)) + else: + assert report.drilldown_unavailable_reason is None + terminal = await analytics.query_async(query=query.model_copy(update={"filters": narrowed})) + assert terminal.summary.total_results == 1 + assert (terminal.drilldown_unavailable_reason is not None) == (remaining < len(extra) * 2) + assert terminal.filters.dimensions[: len(filters.dimensions)] == filters.dimensions + assert len((await analytics.results_async(query=AttackAnalyticsResultsQuery(filters=narrowed))).items) == 1 + + +async def test_matrix_empty_cells_and_additional_predicates_are_exact( + analytics: AttackResultAnalytics, sqlite_instance: SQLiteMemory +) -> None: + await sqlite_instance.add_attack_results_to_memory_async( + attack_results=[ + make_result(operation="one", categories=["privacy"]), + make_result(index=2, operation="two", categories=["safety"], outcome=AttackOutcome.FAILURE), + ] + ) + report = await analytics.query_async( + query=AttackAnalyticsQuery(compare_by=AttackAnalyticsDimension(name="targeted_harm_category")) + ) + assert len(report.cells) == 4 + assert report.groups_overlap + for cell in report.cells: + selected = report.filters.model_copy( + update={"dimensions": [*report.filters.dimensions, *cell.drilldown_filters]} + ) + page = await analytics.results_async(query=AttackAnalyticsResultsQuery(filters=selected)) + assert len(page.items) == cell.statistics.total_results + if not page.items: + assert cell.statistics.success_rate is None + assert set(cell.statistics.outcome_shares.values()) == {0.0} + + +@pytest.mark.parametrize( + ("kind", "value", "stored_label", "dimension", "label"), + [ + ("missing", None, "ignored", "operation", "Not recorded"), + ("no_converters", None, None, "converter_type", "No converters"), + ("value", "", "", "operation", "(Blank)"), + ("value", " \t", None, "label", "(Blank)"), + ("value", "Unknown", "Unknown", "operation", "Unknown"), + ("value", "Not recorded", "Not recorded", "operation", "Not recorded"), + ("value", "content-independent-eval-hash", "model", "objective_target", "model (val-hash)"), + ("value", "scenario-12345678", "run", "scenario", "run (12345678)"), + ("value", "scenario-12345678", None, "scenario", "scenario-12345678"), + ], +) +def test_display_labels_do_not_change_typed_keys( + kind: str, value: str | None, stored_label: str | None, dimension: str, label: str +) -> None: + raw = RawAnalyticsOption(key=AttackAnalyticsValue(kind=kind, value=value), label=stored_label) + option = AttackResultAnalytics._option( + raw=raw, + dimension=AttackAnalyticsDimension(name=dimension, label_key="custom" if dimension == "label" else None), + ) + assert option.label == label + assert option.key == raw.key + + +@pytest.mark.parametrize("counts", [{"unexpected": 1}, {"success": True}, {"error": -1}, {"failure": 1.5}]) +def test_statistics_reject_invalid_raw_data(counts: dict[str, int]) -> None: + with pytest.raises(AnalyticsDataException): + AttackResultAnalytics._statistics(counts) + + +async def test_pages_and_facets_are_fresh_reads_without_reports( + analytics: AttackResultAnalytics, sqlite_instance: SQLiteMemory +) -> None: + await sqlite_instance.add_attack_results_to_memory_async( + attack_results=[make_result(), make_result(index=2, operation="operation-b")] + ) + report = await analytics.query_async(query=AttackAnalyticsQuery(result_limit=1)) + assert report.results.has_more + with patch.object(analytics._reader, "report_async", new_callable=AsyncMock) as read_report: + page = await analytics.results_async(query=AttackAnalyticsResultsQuery(cursor=report.results.next_cursor)) + facet = await analytics.facets_async( + query=AttackAnalyticsFacetQuery( + dimension=AttackAnalyticsDimension(name="operation"), + filters=AttackAnalyticsFilters.model_validate( + {"dimensions": [predicate(name="operation", values=["operation-a"])]} + ), + limit=1, + ) + ) + continuation = await analytics.facets_async( + query=AttackAnalyticsFacetQuery( + dimension=AttackAnalyticsDimension(name="operation"), offset=facet.next_offset, limit=1 + ) + ) + read_report.assert_not_awaited() + assert len(page.items) == 1 + assert page.items[0].attack_result_id != report.results.items[0].attack_result_id + assert page.computed_at >= report.computed_at + assert facet.has_more and facet.next_offset == 1 + assert not continuation.has_more and continuation.next_offset is None + assert {facet.items[0].key.value, continuation.items[0].key.value} == {"operation-a", "operation-b"} + with pytest.raises(ValueError, match="cursor"): + await analytics.results_async(query=AttackAnalyticsResultsQuery(cursor="invalid")) + + +async def test_reports_neither_share_mutable_results_nor_cache_outcomes( + analytics: AttackResultAnalytics, sqlite_instance: SQLiteMemory +) -> None: + result = make_result(categories=["privacy"]) + await sqlite_instance.add_attack_results_to_memory_async(attack_results=[result]) + query = AttackAnalyticsQuery(group_by=AttackAnalyticsDimension(name="targeted_harm_category")) + first = await analytics.query_async(query=query) + second = await analytics.query_async(query=query) + first.summary.successes = 999 + first.results.items.clear() + first.group_by.name = "operation" + assert second.summary.successes == 1 + assert len(second.results.items) == 1 + assert second.group_by == query.group_by + await sqlite_instance.update_attack_result_by_id_async( + attack_result_id=result.attack_result_id, update_fields={"outcome": AttackOutcome.FAILURE.value} + ) + updated = await analytics.query_async(query=query) + assert updated.summary.total_results == 1 + assert updated.groups[0].statistics.failures == 1 + assert updated.summary.success_rate == 0.0 + + +async def test_native_report_and_quick_calls_use_independent_sessions( + analytics: AttackResultAnalytics, sqlite_instance: SQLiteMemory +) -> None: + await sqlite_instance.add_attack_results_to_memory_async( + attack_results=[make_result(), make_result(index=2, outcome=AttackOutcome.FAILURE)] + ) + sessions: list[AsyncSession] = [] + acquire = sqlite_instance.get_session_async + execution = analytics._execution() + # This exercises ownership under concurrency, not a machine-speed latency target. + for lane in execution._lanes.values(): + lane.timeout = 10 + + async def acquire_async() -> AsyncSession: + session = await acquire() + sessions.append(session) + return session + + async with AttackResultAnalytics(memory=sqlite_instance) as second: + assert second._execution() is execution + with patch.object(sqlite_instance, "get_session_async", side_effect=acquire_async): + results = await gather_with_cleanup_async( + [ + *(analytics.query_async() for _ in range(5)), + second.results_async(), + second.facets_async( + query=AttackAnalyticsFacetQuery(dimension=AttackAnalyticsDimension(name="operation")) + ), + ] + ) + reports, page, facet = results[:5], results[5], results[6] + assert all(report.summary.total_results == 2 for report in reports) + assert all(report.summary.success_rate == 0.5 for report in reports) + assert len(page.items) == 2 + assert len(facet.items) == 1 + reports[0].summary.successes = 999 + assert all(report.summary.successes == 1 for report in reports[1:]) + assert len({id(session) for session in sessions}) == 7 + assert not any(session.in_transaction() for session in sessions) + assert all(lane.active == 0 for lane in execution._lanes.values()) + + +@pytest.mark.parametrize("operation", ["query", "results", "facets"]) +async def test_query_snapshot_precedes_admission( + analytics: AttackResultAnalytics, sqlite_instance: SQLiteMemory, operation: str +) -> None: + await sqlite_instance.add_attack_results_to_memory_async(attack_results=[make_result()]) + execution = analytics._execution() + lane = execution._lanes[operation == "query"] + lane.limit, lane.timeout = 1, 10 + entered, release = asyncio.Event(), asyncio.Event() + + async def occupy_async(control: QueryControl) -> None: + entered.set() + await release.wait() + + occupying = asyncio.create_task(execution.run_async(report=operation == "query", task=occupy_async)) + await entered.wait() + filters = AttackAnalyticsFilters.model_validate({"dimensions": [predicate(name="operator", values=["operator-a"])]}) + filters.updated_after = datetime(2025, 1, 1, tzinfo=timezone(timedelta(hours=2))) + if operation == "query": + request = AttackAnalyticsQuery(filters=filters) + pending = asyncio.create_task(analytics.query_async(query=request)) + elif operation == "results": + request = AttackAnalyticsResultsQuery(filters=filters) + pending = asyncio.create_task(analytics.results_async(query=request)) + else: + request = AttackAnalyticsFacetQuery(filters=filters, dimension=AttackAnalyticsDimension(name="operation")) + pending = asyncio.create_task(analytics.facets_async(query=request)) + try: + await asyncio.sleep(0) + assert len(lane.queued) == 1 + filters.dimensions[0].values[0].value = "mutated" + filters.outcomes.append(AttackOutcome.ERROR) + release.set() + await occupying + result = await pending + if operation == "query": + assert result.summary.total_results == 1 + assert result.filters.dimensions[0].values[0].value == "operator-a" + assert result.filters.updated_after == datetime(2024, 12, 31, 22, tzinfo=UTC) + assert result.filters.updated_after.tzinfo is UTC + else: + assert len(result.items) == 1 + finally: + release.set() + await asyncio.gather(occupying, pending, return_exceptions=True) + + +@pytest.mark.parametrize("operation", ["query", "results", "facets"]) +async def test_mutated_models_are_revalidated_before_execution( + analytics: AttackResultAnalytics, operation: str +) -> None: + filters = AttackAnalyticsFilters( + dimensions=[ + AttackAnalyticsFilter( + dimension=AttackAnalyticsDimension(name="operation"), values=[AttackAnalyticsValue(value="operation-a")] + ) + ] + ) + filters.dimensions[0].values.clear() + with patch.object(analytics._execution(), "run_async", new_callable=AsyncMock) as run: + with pytest.raises(ValidationError): + if operation == "query": + await analytics.query_async(query=AttackAnalyticsQuery(filters=filters)) + elif operation == "results": + await analytics.results_async(query=AttackAnalyticsResultsQuery(filters=filters)) + else: + await analytics.facets_async( + query=AttackAnalyticsFacetQuery( + filters=filters, dimension=AttackAnalyticsDimension(name="operation") + ) + ) + run.assert_not_awaited() + + +async def test_frozen_target_groups_preserve_content_hashes_and_every_result_id( + analytics: AttackResultAnalytics, sqlite_instance: SQLiteMemory +) -> None: + targets = [ + TargetIdentifier( + class_name="MockTarget", + class_module="tests", + model_name=f"deploy-{index}", + underlying_model_name="model-x", + endpoint=f"https://example-{index}.test", + temperature=0.3, + ) + for index in (1, 2) + ] + results = [make_result(), make_result(index=2, outcome=AttackOutcome.FAILURE)] + for result, target in zip(results, targets, strict=True): + result.atomic_attack_identifier = AtomicAttackIdentifier.build( + attack_identifier=AttackIdentifier(class_name="ProbeAttack", class_module="tests", objective_target=target) + ) + await sqlite_instance.add_attack_results_to_memory_async(attack_results=results) + with patch.object(ObjectiveTargetAnalyticsIdentityV1, "hash", side_effect=AssertionError("Must use stored keys")): + report = await analytics.query_async( + query=AttackAnalyticsQuery(group_by=AttackAnalyticsDimension(name="objective_target")) + ) + group = report.groups[0] + page = await analytics.results_async( + query=AttackAnalyticsResultsQuery(filters=AttackAnalyticsFilters(dimensions=group.drilldown_filters)) + ) + assert len(report.groups) == 1 + assert group.key.value == "32e7c2bf2a31f21d91dc8bebca280a5ffecf149df474052c40889f7c77b84e81" + assert group.statistics.total_results == 2 + assert group.statistics.success_rate == 0.5 + assert {row.target_identifier_hash for row in page.items} == {target.hash for target in targets} + assert {row.attack_result_id for row in page.items} == {result.attack_result_id for result in results} + + +async def test_construction_and_unused_close_do_not_touch_database() -> None: + memory = MagicMock(spec=MemoryInterface) + with patch.object(CentralMemory, "get_memory_instance", return_value=memory): + analytics = AttackResultAnalytics() + memory.get_session_async.assert_not_awaited() + await analytics.close_async() + assert memory not in AttackResultAnalytics._EXECUTIONS + with pytest.raises(AnalyticsBusyException): + await analytics.query_async() + + +@pytest.mark.parametrize("cancel_close", [False, True]) +async def test_same_backend_shares_controller_and_drain_blocks_replacement( + analytics: AttackResultAnalytics, sqlite_instance: SQLiteMemory, cancel_close: bool +) -> None: + second = AttackResultAnalytics(memory=sqlite_instance) + execution = analytics._execution() + assert second._execution() is execution + entered, release = asyncio.Event(), asyncio.Event() + + async def blocked_async(control: QueryControl) -> None: + entered.set() + await release.wait() + + work = asyncio.create_task(execution.run_async(report=True, task=blocked_async)) + await entered.wait() + work.cancel() + with pytest.raises(asyncio.CancelledError): + await work + closing = asyncio.create_task(analytics.close_async()) + try: + await asyncio.sleep(0) + newcomer = AttackResultAnalytics(memory=sqlite_instance) + assert newcomer._execution() is execution + assert AttackResultAnalytics._EXECUTIONS[sqlite_instance] is execution + assert not closing.done() + if cancel_close: + closing.cancel() + await asyncio.sleep(0) + assert not closing.done() + with pytest.raises(AnalyticsBusyException): + await newcomer.results_async() + with pytest.raises(AnalyticsBusyException): + await second.query_async() + finally: + release.set() + if cancel_close: + with pytest.raises(asyncio.CancelledError): + await closing + else: + await closing + assert sqlite_instance not in AttackResultAnalytics._EXECUTIONS + with pytest.raises(AnalyticsBusyException): + await second.results_async() + async with AttackResultAnalytics(memory=sqlite_instance) as replacement: + assert replacement._execution() is not execution + await second.close_async() + assert AttackResultAnalytics._EXECUTIONS[sqlite_instance] is replacement._execution() + assert (await replacement.query_async()).summary.total_results == 0 + await gather_with_cleanup_async([second.close_async(), newcomer.close_async()]) + + +async def test_foreign_loop_does_not_create_another_backend_budget_or_session( + analytics: AttackResultAnalytics, sqlite_instance: SQLiteMemory +) -> None: + execution = analytics._execution() + + async def foreign_loop_async() -> None: + other = AttackResultAnalytics(memory=sqlite_instance) + with pytest.raises(RuntimeError, match="owning event loop"): + await other.query_async() + with pytest.raises(RuntimeError, match="owning event loop"): + await other.close_async() + assert other._controller is execution + + with patch.object( + sqlite_instance, "get_session_async", side_effect=AssertionError("Must reject before using memory") + ): + await asyncio.to_thread(asyncio.run, foreign_loop_async()) + assert not execution.is_closed + assert (await analytics.query_async()).summary.total_results == 0 + + +async def test_different_memory_objects_have_independent_lifetimes() -> None: + first_memory, second_memory = MagicMock(spec=MemoryInterface), MagicMock(spec=MemoryInterface) + async with ( + AttackResultAnalytics(memory=first_memory) as first, + AttackResultAnalytics(memory=second_memory) as second, + ): + assert first._execution() is not second._execution() + unused = AttackResultAnalytics(memory=second_memory) + await unused.close_async() + await first.close_async() + assert not second._execution().is_closed + + +async def test_native_session_cleanup_retains_capacity_after_caller_cancellation( + analytics: AttackResultAnalytics, sqlite_instance: SQLiteMemory +) -> None: + await sqlite_instance.add_attack_results_to_memory_async(attack_results=[make_result()]) + execution = analytics._execution() + execution._lanes[True].limit = 1 + entered, release = asyncio.Event(), asyncio.Event() + session = await sqlite_instance.get_session_async() + close = session.close + + async def delayed_close_async() -> None: + entered.set() + await release.wait() + await close() + + with ( + patch.object(sqlite_instance, "get_session_async", new_callable=AsyncMock, return_value=session) as acquire, + patch.object(session, "close", side_effect=delayed_close_async), + ): + caller = asyncio.create_task(analytics.query_async()) + try: + await entered.wait() + caller.cancel() + with pytest.raises(asyncio.CancelledError): + await caller + assert execution._lanes[True].active == 1 + waiting = asyncio.create_task(analytics.query_async()) + await asyncio.sleep(0) + acquire.assert_awaited_once() + closing = asyncio.create_task(analytics.close_async()) + await asyncio.sleep(0) + with pytest.raises(AnalyticsBusyException): + await waiting + assert not closing.done() + finally: + release.set() + await analytics.close_async() + await asyncio.gather(caller, return_exceptions=True) + await closing + assert execution.is_closed + async with AttackResultAnalytics(memory=sqlite_instance) as replacement: + assert (await replacement.query_async()).summary.total_results == 1 diff --git a/tests/unit/analytics/test_execution.py b/tests/unit/analytics/test_execution.py new file mode 100644 index 0000000000..038c656d80 --- /dev/null +++ b/tests/unit/analytics/test_execution.py @@ -0,0 +1,380 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +from __future__ import annotations + +import asyncio +import gc +import threading +import weakref +from typing import TYPE_CHECKING +from unittest.mock import patch + +import pytest + +from pyrit.analytics._execution import AnalyticsExecution +from pyrit.common.task_utils import gather_with_cleanup_async +from pyrit.exceptions.analytics_exception import AnalyticsBusyException, AnalyticsTimeoutException + +if TYPE_CHECKING: + from collections.abc import AsyncGenerator + + from pyrit.memory.query_control import QueryControl + + +class _Blocked: + def __init__(self, *, result: str = "result", error: Exception | None = None) -> None: + self.entered = asyncio.Event() + self.release = asyncio.Event() + self.exited = asyncio.Event() + self.control: QueryControl | None = None + self.result = result + self.error = error + self.cancellations = 0 + + async def run_async(self, control: QueryControl) -> str: + self.control = control + self.entered.set() + try: + await self.release.wait() + if self.error is not None: + raise self.error + return self.result + except asyncio.CancelledError: + self.cancellations += 1 + raise + finally: + self.exited.set() + + +class _Harness: + def __init__(self) -> None: + self.executions: list[AnalyticsExecution] = [] + self.operations: list[_Blocked] = [] + self.callers: list[asyncio.Task[str]] = [] + + def controller(self, **overrides: int | float) -> AnalyticsExecution: + options = { + "report_workers": 1, + "quick_workers": 1, + "max_queue": 1, + "queue_timeout": 10.0, + "report_timeout": 10.0, + "quick_timeout": 10.0, + **overrides, + } + execution = AnalyticsExecution(**options) + self.executions.append(execution) + return execution + + def blocked(self, *, result: str = "result", error: Exception | None = None) -> _Blocked: + operation = _Blocked(result=result, error=error) + self.operations.append(operation) + return operation + + def submit(self, *, execution: AnalyticsExecution, operation: _Blocked, report: bool = True) -> asyncio.Task[str]: + caller = asyncio.create_task(execution.run_async(report=report, task=operation.run_async)) + self.callers.append(caller) + return caller + + +@pytest.fixture +async def harness() -> AsyncGenerator[_Harness, None]: + harness = _Harness() + try: + yield harness + finally: + for operation in harness.operations: + operation.release.set() + await gather_with_cleanup_async(execution.close_async() for execution in harness.executions) + await asyncio.gather(*harness.callers, return_exceptions=True) + + +async def test_defaults_bound_both_lanes() -> None: + execution = AnalyticsExecution() + assert (execution._lanes[True].limit, execution._lanes[False].limit) == (5, 2) + assert (execution._max_queue, execution._queue_timeout) == (10, 1) + assert (execution._lanes[True].timeout, execution._lanes[False].timeout) == (5, 1) + await execution.close_async() + + +@pytest.mark.parametrize("name", ["report_workers", "quick_workers", "max_queue"]) +@pytest.mark.parametrize("value", [0, -1, 1.5, True, float("inf"), float("nan"), "1", None]) +async def test_invalid_counts_are_rejected(*, name: str, value: object) -> None: + with pytest.raises(ValueError, match=name): + AnalyticsExecution(**{name: value}) + + +@pytest.mark.parametrize("name", ["queue_timeout", "report_timeout", "quick_timeout"]) +@pytest.mark.parametrize("value", [0, -1.0, True, float("inf"), float("nan"), "1", None, 10**400]) +async def test_invalid_timeouts_are_rejected(*, name: str, value: object) -> None: + with pytest.raises(ValueError, match=name): + AnalyticsExecution(**{name: value}) + + +@pytest.mark.parametrize("report", [True, False]) +async def test_work_uses_caller_loop_and_returns_independent_results(harness: _Harness, report: bool) -> None: + execution = harness.controller() + loop, thread = asyncio.get_running_loop(), threading.get_ident() + + async def work_async(control: QueryControl) -> list[int]: + assert asyncio.get_running_loop() is loop + assert threading.get_ident() == thread + control.check() + return [1] + + first = await execution.run_async(report=report, task=work_async) + second = await execution.run_async(report=report, task=work_async) + assert first == second + assert first is not second + + +@pytest.mark.parametrize("error", [ValueError("query failed"), AnalyticsBusyException(), AnalyticsTimeoutException()]) +async def test_failures_propagate_and_release_capacity(harness: _Harness, error: Exception) -> None: + execution = harness.controller() + operation = harness.blocked(error=error) + operation.release.set() + with pytest.raises(type(error)): + await execution.run_async(report=True, task=operation.run_async) + assert execution._lanes[True].active == 0 + successful = harness.blocked() + successful.release.set() + assert await execution.run_async(report=True, task=successful.run_async) == "result" + + +async def test_lanes_reserve_independent_running_and_queued_capacity(harness: _Harness) -> None: + execution = harness.controller() + report, quick, queued_report, queued_quick = [harness.blocked() for _ in range(4)] + report_call = harness.submit(execution=execution, operation=report) + quick_call = harness.submit(execution=execution, operation=quick, report=False) + await gather_with_cleanup_async([report.entered.wait(), quick.entered.wait()]) + waiting_report = harness.submit(execution=execution, operation=queued_report) + waiting_quick = harness.submit(execution=execution, operation=queued_quick, report=False) + await asyncio.sleep(0) + for lane in (True, False): + with pytest.raises(AnalyticsBusyException): + await execution.run_async(report=lane, task=harness.blocked().run_async) + quick.release.set() + await queued_quick.entered.wait() + assert not queued_report.entered.is_set() + assert not report_call.done() + queued_quick.release.set() + assert await quick_call == await waiting_quick == "result" + report.release.set() + await queued_report.entered.wait() + queued_report.release.set() + assert await report_call == await waiting_report == "result" + + +async def test_admission_is_fifo_with_new_arrivals(harness: _Harness) -> None: + execution = harness.controller(max_queue=2) + active, first, second, newcomer = [harness.blocked() for _ in range(4)] + harness.submit(execution=execution, operation=active) + await active.entered.wait() + harness.submit(execution=execution, operation=first) + harness.submit(execution=execution, operation=second) + await asyncio.sleep(0) + active.release.set() + await first.entered.wait() + harness.submit(execution=execution, operation=newcomer) + await asyncio.sleep(0) + first.release.set() + await second.entered.wait() + assert not newcomer.entered.is_set() + second.release.set() + await newcomer.entered.wait() + + +@pytest.mark.parametrize("report", [True, False]) +async def test_queue_expiry_never_starts_database_work(harness: _Harness, report: bool) -> None: + execution = harness.controller(queue_timeout=0.01) + active, queued = harness.blocked(), harness.blocked() + harness.submit(execution=execution, operation=active, report=report) + await active.entered.wait() + with pytest.raises(AnalyticsBusyException): + await execution.run_async(report=report, task=queued.run_async) + assert not queued.entered.is_set() + assert not execution._lanes[report].queued + assert execution._lanes[report].active == 1 + + +@pytest.mark.parametrize("report", [True, False]) +async def test_cancelling_queued_call_reclaims_only_its_queue_entry(harness: _Harness, report: bool) -> None: + execution = harness.controller() + active, queued, replacement = [harness.blocked() for _ in range(3)] + harness.submit(execution=execution, operation=active, report=report) + await active.entered.wait() + caller = harness.submit(execution=execution, operation=queued, report=report) + await asyncio.sleep(0) + caller.cancel() + with pytest.raises(asyncio.CancelledError): + await caller + assert not queued.entered.is_set() + assert not active.control.expired + harness.submit(execution=execution, operation=replacement, report=report) + await asyncio.sleep(0) + assert len(execution._lanes[report].queued) == 1 + active.release.set() + await replacement.entered.wait() + + +@pytest.mark.parametrize("cancel_caller", [False, True], ids=["deadline", "cancellation"]) +async def test_response_exit_keeps_slot_until_operation_exit(harness: _Harness, cancel_caller: bool) -> None: + execution = harness.controller(report_timeout=10 if cancel_caller else 0.01) + active, queued = harness.blocked(), harness.blocked() + caller = harness.submit(execution=execution, operation=active) + await active.entered.wait() + if cancel_caller: + caller.cancel() + with pytest.raises(asyncio.CancelledError if cancel_caller else AnalyticsTimeoutException): + await caller + assert active.control.cancel_event.is_set() + assert not active.exited.is_set() + assert active.cancellations == 0 + execution._lanes[True].timeout = 10 + harness.submit(execution=execution, operation=queued) + await asyncio.sleep(0) + assert not queued.entered.is_set() + with pytest.raises(AnalyticsBusyException): + await execution.run_async(report=True, task=harness.blocked().run_async) + active.release.set() + await queued.entered.wait() + assert active.exited.is_set() + + +async def test_expired_control_cannot_return_a_late_success(harness: _Harness) -> None: + execution = harness.controller() + operation = harness.blocked() + caller = harness.submit(execution=execution, operation=operation) + await operation.entered.wait() + operation.control.deadline = 0 + operation.release.set() + with pytest.raises(AnalyticsTimeoutException): + await caller + + +async def test_operation_stays_strongly_owned_after_caller_leaves(harness: _Harness) -> None: + execution = harness.controller() + entered = asyncio.Event() + gate_ref: weakref.ReferenceType[asyncio.Future[None]] | None = None + + async def wait_async(control: QueryControl) -> None: + nonlocal gate_ref + gate = asyncio.get_running_loop().create_future() + gate_ref = weakref.ref(gate) + entered.set() + await gate + + caller = asyncio.create_task(execution.run_async(report=True, task=wait_async)) + await entered.wait() + caller.cancel() + with pytest.raises(asyncio.CancelledError): + await caller + del caller + gc.collect() + assert gate_ref is not None + gate = gate_ref() + assert gate is not None + gate.set_result(None) + await execution.close_async() + assert execution._lanes[True].active == 0 + + +async def test_cancellation_after_admission_grant_returns_reserved_slot(harness: _Harness) -> None: + execution = harness.controller() + lane = execution._lanes[True] + lane.active = 1 + admission = asyncio.create_task(execution._admit_async(lane)) + await asyncio.sleep(0) + execution._release(lane) + admission.cancel() + with pytest.raises(asyncio.CancelledError): + await admission + assert lane.active == 0 + assert not lane.queued + + +@pytest.mark.parametrize("expire_before_grant", [False, True]) +async def test_admission_rechecks_deadline_even_without_timer_callback( + harness: _Harness, expire_before_grant: bool +) -> None: + execution = harness.controller() + lane = execution._lanes[True] + lane.active = 1 + admission = asyncio.create_task(execution._admit_async(lane)) + await asyncio.sleep(0) + deadline = lane.queued[0].deadline + if not expire_before_grant: + execution._release(lane) + with patch("pyrit.analytics._execution.monotonic", return_value=deadline + 1): + if expire_before_grant: + execution._release(lane) + with pytest.raises(AnalyticsBusyException): + await admission + assert lane.active == 0 + assert not lane.queued + + +async def test_unexpected_failure_after_cancellation_is_logged( + harness: _Harness, caplog: pytest.LogCaptureFixture +) -> None: + execution = harness.controller() + operation = harness.blocked(error=ValueError("cleanup failed")) + caller = harness.submit(execution=execution, operation=operation) + await operation.entered.wait() + caller.cancel() + with pytest.raises(asyncio.CancelledError): + await caller + operation.release.set() + await execution.close_async() + assert "Analytics operation failed after its caller left" in caplog.text + assert "cleanup failed" in caplog.text + + +async def test_close_rejects_queued_work_and_drains_through_repeated_cancellation(harness: _Harness) -> None: + execution = harness.controller() + active, queued = harness.blocked(), harness.blocked() + caller = harness.submit(execution=execution, operation=active) + await active.entered.wait() + waiting = harness.submit(execution=execution, operation=queued) + await asyncio.sleep(0) + closing = asyncio.create_task(execution.close_async()) + await asyncio.sleep(0) + with pytest.raises(AnalyticsBusyException): + await waiting + with pytest.raises(AnalyticsBusyException): + await execution.run_async(report=False, task=queued.run_async) + assert active.control.cancel_event.is_set() + assert not queued.entered.is_set() + closing.cancel() + await asyncio.sleep(0) + closing.cancel() + await asyncio.sleep(0) + assert not closing.done() + assert not execution.is_closed + active.release.set() + with pytest.raises(asyncio.CancelledError): + await closing + with pytest.raises(AnalyticsTimeoutException): + await caller + assert execution.is_closed + assert execution._lanes[True].active == 0 + await execution.close_async() + with pytest.raises(AnalyticsBusyException): + await execution.run_async(report=True, task=queued.run_async) + + +async def test_foreign_loop_cannot_use_or_close_live_controller(harness: _Harness) -> None: + execution = harness.controller() + operation = harness.blocked() + operation.release.set() + + async def foreign_loop_async() -> None: + with pytest.raises(RuntimeError, match="owning event loop"): + await execution.run_async(report=True, task=operation.run_async) + with pytest.raises(RuntimeError, match="owning event loop"): + await execution.close_async() + + await asyncio.to_thread(asyncio.run, foreign_loop_async()) + assert not execution.is_closed + assert await execution.run_async(report=True, task=operation.run_async) == "result" diff --git a/tests/unit/analytics/test_profile_aggregation.py b/tests/unit/analytics/test_profile_aggregation.py new file mode 100644 index 0000000000..21da96ad80 --- /dev/null +++ b/tests/unit/analytics/test_profile_aggregation.py @@ -0,0 +1,414 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +from __future__ import annotations + +import asyncio +import json +from copy import deepcopy +from datetime import UTC, datetime +from typing import TYPE_CHECKING, Any +from unittest.mock import patch + +import pytest +from sqlalchemy import text, update + +from pyrit.analytics import AttackResultAnalytics +from pyrit.analytics._profile_aggregation import ProfileAggregation +from pyrit.exceptions.analytics_exception import AnalyticsDataException, AnalyticsTimeoutException +from pyrit.memory.attack_analytics import AttackAnalyticsReader, RawAnalyticsProfile, RawAnalyticsReport +from pyrit.memory.memory_models import AttackResultEntry +from pyrit.models import ( + AttackAnalyticsDimension, + AttackAnalyticsQuery, + AttackAnalyticsResults, + AttackOutcome, +) +from unit.memory.test_attack_analytics import control, make_result + +if TYPE_CHECKING: + from pyrit.memory import SQLiteMemory + + +@pytest.fixture +async def mixed_reader(sqlite_instance: SQLiteMemory) -> AttackAnalyticsReader: + missing = make_result(index=5, outcome=AttackOutcome.ERROR, categories=["Unknown", ""]) + missing.atomic_attack_identifier = None + results = [ + make_result( + index=index, + categories=["Privacy", "privacy", "", "Unknown"], + converters=["Alpha", "ALPHA", ""], + response_converters=["Beta"], + ) + for index in (1, 2) + ] + results.extend( + [ + make_result( + index=3, + outcome=AttackOutcome.FAILURE, + categories=["privacy"], + converters=["Alpha", "Gamma"], + response_converters=["Beta", "BETA"], + ), + make_result(index=4, outcome=AttackOutcome.UNDETERMINED), + missing, + make_result( + index=6, + categories=["\u00c9vasion", "\u00e9vasion", "\u00df", "SS", "\u0130", "i\u0307"], + converters=["\u00dcnicode", "\u00fcnicode"], + response_converters=["Other"], + attack_class_name="\u00c9cho", + ), + make_result( + index=7, + categories=["\u00e9vasion", "privacy"], + converters=["\u00fcnicode", "Beta"], + response_converters=["Gamma"], + attack_class_name="\u00e9cho", + ), + make_result(index=8, categories=["legacy"]), + make_result(index=9, categories=["privacy"], converters=["Beta"]), + ] + ) + await sqlite_instance.add_attack_results_to_memory_async(attack_results=results) + legacy = { + "children": { + "attack": { + "__type__": "LegacyAttack", + "children": { + "request_converters": [ + {"__type__": "\u00c9vasion"}, + {"class_name": "\u00e9vasion", "__type__": "Ignored"}, + {"class_name": None, "__type__": "AlsoIgnored"}, + {}, + ], + "response_converters": [], + }, + } + } + } + async with sqlite_instance._get_async_engine().begin() as connection: + await connection.execute( + update(AttackResultEntry) + .where(AttackResultEntry.id == results[-2].attack_result_id) + .values(atomic_attack_identifier_hash=None, atomic_attack_identifier=legacy) + ) + await connection.execute( + update(AttackResultEntry) + .where(AttackResultEntry.id == results[-1].attack_result_id) + .values(atomic_attack_identifier_hash=None) + ) + return AttackAnalyticsReader(memory=sqlite_instance) + + +def _assert_same_chart(*, actual: RawAnalyticsReport, expected: RawAnalyticsReport) -> None: + assert actual.counts == expected.counts + assert actual.groups == expected.groups + assert actual.rows == expected.rows + assert actual.columns == expected.columns + assert actual.cells == expected.cells + assert actual.has_more_groups == expected.has_more_groups + assert actual.axes_truncated == expected.axes_truncated + assert actual.results.items == expected.results.items + + +@pytest.mark.parametrize( + "options", + [ + {"group_by": {"name": "targeted_harm_category"}, "group_limit": 3, "group_offset": 1}, + {"group_by": {"name": "targeted_harm_category"}, "group_offset": 50}, + {"group_by": {"name": "converter_type"}, "group_limit": 3, "group_offset": 1}, + {"group_by": {"name": "attack_type"}, "group_limit": 1, "group_offset": 1}, + { + "group_by": {"name": "targeted_harm_category"}, + "compare_by": {"name": "converter_type"}, + "axis_limit": 2, + }, + { + "group_by": {"name": "targeted_harm_category"}, + "compare_by": {"name": "attack_type"}, + "axis_limit": 2, + }, + {"group_by": {"name": "attack_type"}, "compare_by": {"name": "converter_type"}}, + { + "group_by": {"name": "converter_type"}, + "compare_by": {"name": "converter_type", "converter_direction": "response"}, + }, + { + "group_by": {"name": "targeted_harm_category"}, + "compare_by": {"name": "converter_type"}, + "filters": { + "dimensions": [ + {"dimension": {"name": "converter_type"}, "values": [{"value": "Alpha"}, {"value": "Other"}]}, + {"dimension": {"name": "converter_type"}, "values": [{"value": "Gamma"}]}, + ] + }, + }, + { + "group_by": {"name": "targeted_harm_category"}, + "compare_by": {"name": "converter_type"}, + "filters": {"outcomes": ["error"]}, + }, + {"group_by": {"name": "converter_type"}, "filters": {"outcomes": ["failure", "undetermined"]}}, + ], +) +async def test_weighted_unicode_and_legacy_profiles_equal_actual_sql( + mixed_reader: AttackAnalyticsReader, options: dict[str, Any] +) -> None: + query = AttackAnalyticsQuery.model_validate(options) + sql = await mixed_reader.report_async(query=query, control=control()) + compact = await mixed_reader.report_async(query=query, control=control(), use_compact_profiles=True) + assert compact.profiles is not None + sources = deepcopy(compact.profiles) + await ProfileAggregation.populate_async(report=compact, query=query, control=control()) + _assert_same_chart(actual=compact, expected=sql) + assert compact.profiles == sources + + +@pytest.mark.parametrize("compare", [None, "converter_type", "attack_type"]) +@pytest.mark.parametrize("cap", ["MAX_COMPACT_PROFILES", "MAX_COMPACT_VALUE_LENGTH", "MAX_COMPACT_TOTAL_LENGTH"]) +@pytest.mark.parametrize("overflow", [False, True], ids=["at_cap", "overflow"]) +async def test_exact_reader_cap_and_overflow_keep_complete_sql_parity( + mixed_reader: AttackAnalyticsReader, compare: str | None, cap: str, overflow: bool +) -> None: + query = AttackAnalyticsQuery( + group_by=AttackAnalyticsDimension(name="targeted_harm_category"), + compare_by=AttackAnalyticsDimension(name=compare) if compare else None, + axis_limit=2, + group_limit=2, + ) + sql = await mixed_reader.report_async(query=query, control=control()) + probe = await mixed_reader.report_async(query=query, control=control(), use_compact_profiles=True) + assert probe.profiles + boundaries = { + "MAX_COMPACT_PROFILES": len(probe.profiles), + "MAX_COMPACT_VALUE_LENGTH": max( + max(len(profile["source0"] or ""), len(profile.get("source1") or "")) for profile in probe.profiles + ), + "MAX_COMPACT_TOTAL_LENGTH": sum( + len(value) for profile in probe.profiles for value in profile.values() if isinstance(value, str) + ), + } + with patch.object(AttackAnalyticsReader, cap, boundaries[cap] - int(overflow)): + result = await mixed_reader.report_async(query=query, control=control(), use_compact_profiles=True) + assert (result.profiles is None) == overflow + await ProfileAggregation.populate_async(report=result, query=query, control=control()) + _assert_same_chart(actual=result, expected=sql) + + +@pytest.mark.parametrize("compare", [None, "converter_type", "attack_type"]) +async def test_sdk_report_matches_sql_fallback( + mixed_reader: AttackAnalyticsReader, sqlite_instance: SQLiteMemory, compare: str | None +) -> None: + query = AttackAnalyticsQuery( + group_by=AttackAnalyticsDimension(name="targeted_harm_category"), + compare_by=AttackAnalyticsDimension(name=compare) if compare else None, + ) + async with AttackResultAnalytics(memory=sqlite_instance) as analytics: + fast = await analytics.query_async(query=query) + with patch.object(AttackAnalyticsReader, "MAX_COMPACT_PROFILES", 0): + sql = await analytics.query_async(query=query) + fields = { + "summary", + "groups", + "rows", + "columns", + "cells", + "groups_overlap", + "axes_truncated", + "has_more_groups", + "next_group_offset", + "outcome_filter_applied", + "drilldown_unavailable_reason", + } + assert fast.model_dump(include=fields) == sql.model_dump(include=fields) + + +@pytest.mark.parametrize("compare", [False, True]) +async def test_empty_profile_cohort_is_not_sql_fallback(sqlite_instance: SQLiteMemory, compare: bool) -> None: + reader = AttackAnalyticsReader(memory=sqlite_instance) + query = AttackAnalyticsQuery( + group_by=AttackAnalyticsDimension(name="converter_type"), + compare_by=AttackAnalyticsDimension(name="attack_type") if compare else None, + ) + raw = await reader.report_async(query=query, control=control(), use_compact_profiles=True) + assert raw.profiles == [] + await ProfileAggregation.populate_async(report=raw, query=query, control=control()) + assert raw.groups == raw.rows == raw.columns == raw.cells == [] + assert not raw.has_more_groups and not raw.axes_truncated + + +async def test_non_categorical_dimension_uses_complete_sql(sqlite_instance: SQLiteMemory) -> None: + await sqlite_instance.add_attack_results_to_memory_async(attack_results=[make_result()]) + raw = await AttackAnalyticsReader(memory=sqlite_instance).report_async( + query=AttackAnalyticsQuery(), control=control(), use_compact_profiles=True + ) + assert raw.profiles is None + original = deepcopy(raw) + await ProfileAggregation.populate_async(report=raw, query=AttackAnalyticsQuery(), control=control()) + assert raw == original + + +@pytest.mark.parametrize("raw", ["null", " \t null\n"]) +async def test_profile_null_classification_matches_sql_not_json_decoder_defaults( + sqlite_instance: SQLiteMemory, raw: str +) -> None: + await sqlite_instance.add_attack_results_to_memory_async(attack_results=[make_result()]) + async with sqlite_instance._get_async_engine().begin() as connection: + await connection.execute(text("UPDATE AttackResultEntries SET targeted_harm_categories = :raw"), {"raw": raw}) + reader = AttackAnalyticsReader(memory=sqlite_instance) + query = AttackAnalyticsQuery(group_by=AttackAnalyticsDimension(name="targeted_harm_category")) + compact = await reader.report_async(query=query, control=control(), use_compact_profiles=True) + if raw == "null": + sql = await reader.report_async(query=query, control=control()) + await ProfileAggregation.populate_async(report=compact, query=query, control=control()) + _assert_same_chart(actual=compact, expected=sql) + else: + with pytest.raises(AnalyticsDataException): + await reader.report_async(query=query, control=control()) + with pytest.raises(AnalyticsDataException): + await ProfileAggregation.populate_async(report=compact, query=query, control=control()) + + +@pytest.mark.parametrize( + ("dimension", "raw", "expected"), + [ + ("converter_type", None, {("missing", ""): None}), + ("converter_type", "null", {("missing", ""): None}), + ("converter_type", " \t[\r\n ] ", {("no_converters", ""): None}), + ("targeted_harm_category", "[]", {("missing", ""): None}), + ("converter_type", "[null, null]", {("missing", ""): None}), + ("converter_type", '["", "Unknown", "unknown"]', {("value", ""): "", ("value", "unknown"): "Unknown"}), + ("attack_type", "", {("value", ""): ""}), + ("attack_type", "null", {("value", "null"): "null"}), + ( + "targeted_harm_category", + '["\u00c9vasion", "\u00e9vasion", "SS", "\u00df", "\u0130", "i\u0307"]', + { + ("value", "\u00e9vasion"): "\u00c9vasion", + ("value", "ss"): "SS", + ("value", "\u00df"): "\u00df", + ("value", "i\u0307"): "i\u0307", + }, + ), + ], +) +def test_profile_keys_use_unicode_lower_not_ascii_or_casefold( + dimension: str, raw: str | None, expected: dict[tuple[str, str], str | None] +) -> None: + values = ProfileAggregation._values(raw=raw, dimension=AttackAnalyticsDimension(name=dimension)) + assert {key: option.label for key, option in values.items()} == expected + + +@pytest.mark.parametrize( + "raw", ['{"class_name": "Alpha"}', '"Alpha"', "true", "[1]", "[true]", "not json", '[{"__type__": "Alpha"}]'] +) +def test_profile_decoder_rejects_noncanonical_arrays(raw: str) -> None: + with pytest.raises(AnalyticsDataException): + ProfileAggregation._values(raw=raw, dimension=AttackAnalyticsDimension(name="converter_type")) + + +def test_profile_decoder_rejects_lowercase_expansion_beyond_key_limit() -> None: + with pytest.raises(AnalyticsDataException, match="4,096-character"): + ProfileAggregation._values(raw="\u0130" * 4096, dimension=AttackAnalyticsDimension(name="attack_type")) + + +def _report(profiles: list[RawAnalyticsProfile]) -> RawAnalyticsReport: + return RawAnalyticsReport( + counts={}, + groups=[], + rows=[], + columns=[], + cells=[], + has_more_groups=False, + axes_truncated=False, + results=AttackAnalyticsResults(items=[], has_more=False, next_cursor=None, computed_at=datetime.now(tz=UTC)), + warnings=[], + profiles=profiles, + ) + + +async def test_decoder_cache_is_call_local_and_never_caches_counts() -> None: + records: list[RawAnalyticsProfile] = [ + {"source0": '["Alpha", "Alpha"]', "weight": 5, "outcome": "success", "oversized": False}, + {"source0": '["Alpha", "Alpha"]', "weight": 2, "outcome": "failure", "oversized": False}, + {"source0": '["ALPHA"]', "weight": 3, "outcome": "success", "oversized": False}, + ] + query = AttackAnalyticsQuery(group_by=AttackAnalyticsDimension(name="converter_type")) + report = _report(records) + with patch.object(ProfileAggregation, "_values", wraps=ProfileAggregation._values) as decode: + await ProfileAggregation.populate_async(report=report, query=query, control=control()) + assert decode.call_count == 2 + assert report.groups[0].counts == {"success": 8, "failure": 2} + assert report.groups[0].option.label == "ALPHA" + records[0]["weight"] = 9 + await ProfileAggregation.populate_async(report=report, query=query, control=control()) + assert decode.call_count == 4 + assert report.groups[0].counts == {"success": 12, "failure": 2} + + +async def test_cancelled_control_stops_before_decoding() -> None: + query_control = control() + query_control.cancel() + with patch.object(ProfileAggregation, "_profile", side_effect=AssertionError("Must not decode")): + with pytest.raises(AnalyticsTimeoutException): + await ProfileAggregation.populate_async( + report=_report([{"source0": "[]", "weight": 1, "outcome": "success", "oversized": False}]), + query=AttackAnalyticsQuery(group_by=AttackAnalyticsDimension(name="converter_type")), + control=query_control, + ) + + +@pytest.mark.parametrize("oversized", [False, True]) +async def test_incomplete_or_oversized_profiles_are_not_partial_charts(oversized: bool) -> None: + with pytest.raises(AnalyticsDataException, match="Oversized|second dimension"): + await ProfileAggregation.populate_async( + report=_report([{"source0": "[]", "weight": 1, "outcome": "success", "oversized": oversized}]), + query=AttackAnalyticsQuery( + group_by=AttackAnalyticsDimension(name="converter_type"), + compare_by=AttackAnalyticsDimension(name="attack_type"), + ), + control=control(), + ) + + +async def test_full_profile_budget_yields_and_keeps_weighted_totals() -> None: + records: list[RawAnalyticsProfile] = [ + { + "source0": json.dumps([f"category-{index % 64}", "shared", "SHARED"]), + "source1": json.dumps([f"converter-{index // 64}", "shared", "shared"]), + "weight": 25, + "outcome": list(AttackOutcome)[index % 4].value, + "oversized": False, + } + for index in range(AttackAnalyticsReader.MAX_COMPACT_PROFILES) + ] + query = AttackAnalyticsQuery( + group_by=AttackAnalyticsDimension(name="targeted_harm_category"), + compare_by=AttackAnalyticsDimension(name="converter_type"), + ) + done = asyncio.Event() + pulses = 0 + + async def heartbeat_async() -> None: + nonlocal pulses + while not done.is_set(): + pulses += 1 + await asyncio.sleep(0) + + heartbeat = asyncio.create_task(heartbeat_async()) + report = _report(records) + report.counts = {outcome.value: 25_600 for outcome in AttackOutcome} + try: + await ProfileAggregation.populate_async(report=report, query=query, control=control()) + finally: + done.set() + await heartbeat + assert pulses > 1 + assert report.counts == {outcome.value: 25_600 for outcome in AttackOutcome} + assert report.axes_truncated + assert len(report.rows) == len(report.columns) == 20 + assert len(report.cells) == 400 + assert all(sum(cell.counts.values()) == 25 for cell in report.cells) diff --git a/tests/unit/common/test_lazy_package_imports.py b/tests/unit/common/test_lazy_package_imports.py index 5c44139d9b..c05c810305 100644 --- a/tests/unit/common/test_lazy_package_imports.py +++ b/tests/unit/common/test_lazy_package_imports.py @@ -22,6 +22,18 @@ ) _LAZY_IMPORT_SPOT_CHECKS = [ + ( + "pyrit.analytics", + "AttackResultAnalytics", + "pyrit.analytics.attack_result_analytics", + "pyrit.analytics.conversation_analytics", + ), + ( + "pyrit.analytics", + "AttackStats", + "pyrit.analytics.result_analysis", + "pyrit.analytics.attack_result_analytics", + ), ( "pyrit.registry", "AttackRegistry", From 8c22c7c1562c5b090a1eb1ad277d0d73d50f8ad6 Mon Sep 17 00:00:00 2001 From: Roman Lutz Date: Thu, 8 Oct 2026 10:32:02 -0700 Subject: [PATCH 2/8] FIX: Harden analytics task admission and cancellation Preserve cancellation when admission races with a slot grant on Python 3.11. Release unscheduled operation capacity and close rejected coroutines, keep shutdown retryable, and log failures that race with caller cancellation. Add deterministic lifecycle regressions and verify raw-result analytics remains distinct from scenario-unit statistics and explicit result-role policy. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- doc/code/analytics/0_attack_results.md | 8 ++ pyrit/analytics/_execution.py | 51 +++++++-- .../analytics/test_attack_result_analytics.py | 34 +++++- tests/unit/analytics/test_execution.py | 101 +++++++++++++++++- 4 files changed, 181 insertions(+), 13 deletions(-) diff --git a/doc/code/analytics/0_attack_results.md b/doc/code/analytics/0_attack_results.md index e229de40e0..bfd277436a 100644 --- a/doc/code/analytics/0_attack_results.md +++ b/doc/code/analytics/0_attack_results.md @@ -177,6 +177,9 @@ raises `AnalyticsTimeoutException`. Native async readers own the `QueryControl` database deadlines, connection acquisition, interruption, and session restoration. The SDK adds no second SQL timeout mechanism, sync DB facade, or event-loop thread. +Cancelling a queued call removes it without starting database work, including +cancellation racing with an admission grant. + Cancelling a caller or reaching its response deadline signals cooperative cancellation. It does **not** release a still-running query's slot. Capacity remains occupied until the operation and its session cleanup actually finish. @@ -184,6 +187,11 @@ Closing rejects queued/new calls, signals active work, and waits for that cleanu even when it outlasts the response deadline. Cancelling `close_async` is propagated only after draining; repeated cancellation cannot open an overlapping controller. +If the event loop cannot schedule an operation, the scheduling error propagates +and its unused capacity is released. If scheduling shutdown fails, admission +stays closed and `close_async()` can be retried. Unexpected operation failures +after a caller leaves are logged, including failures racing with cancellation. + Closing one bound facade closes the shared controller for every facade attached to it. Those facades are terminal. Construct a new SDK owner only after closing completes to begin another lifetime. Closing an unused facade does not affect diff --git a/pyrit/analytics/_execution.py b/pyrit/analytics/_execution.py index cd61c4d7dd..378ef5dcdd 100644 --- a/pyrit/analytics/_execution.py +++ b/pyrit/analytics/_execution.py @@ -18,7 +18,7 @@ from pyrit.memory.query_control import QueryControl if TYPE_CHECKING: - from collections.abc import Awaitable, Callable + from collections.abc import Awaitable, Callable, Coroutine logger = logging.getLogger(__name__) T = TypeVar("T") @@ -141,19 +141,21 @@ async def run_async(self, *, report: bool, task: Callable[[QueryControl], Awaita finished=self._loop.create_future(), ) lane.running.add(work) - operation = self._loop.create_task(self._execute_async(task=task, control=work.control)) + try: + operation = self._create_task(self._execute_async(task=task, control=work.control)) + except BaseException: + self._finish(lane=lane, work=work) + raise work.task = operation operation.add_done_callback(lambda completed: self._complete(lane=lane, work=work, task=completed)) try: done, _ = await asyncio.wait({operation}, timeout=work.control.remaining) if not done: - work.abandoned = True - work.control.cancel() + self._abandon(work=work, task=operation) raise AnalyticsTimeoutException return operation.result() except asyncio.CancelledError: - work.abandoned = True - work.control.cancel() + self._abandon(work=work, task=operation) raise async def close_async(self) -> None: @@ -164,6 +166,8 @@ async def close_async(self) -> None: returning early would allow a replacement controller to overlap database operations still cleaning up. Cancelling this await, even repeatedly, is propagated only after draining. The memory backend itself is not disposed. + If scheduling the drain fails, admission remains closed and a later + ``close_async`` call can retry the drain. """ self._check_loop() if self._close_task is None: @@ -175,7 +179,7 @@ async def close_async(self) -> None: waiter.ready.set_exception(AnalyticsBusyException()) for work in lane.running: work.control.cancel() - self._close_task = self._loop.create_task(self._drain_async()) + self._close_task = self._create_task(self._drain_async()) cancellation: asyncio.CancelledError | None = None while not self._close_task.done(): try: @@ -190,6 +194,13 @@ def _check_loop(self) -> None: if asyncio.get_running_loop() is not self._loop: raise RuntimeError("Analytics must be used and closed on its owning event loop.") + def _create_task(self, coroutine: Coroutine[object, object, T]) -> asyncio.Task[T]: + try: + return self._loop.create_task(coroutine) + except BaseException: + coroutine.close() + raise + async def _admit_async(self, lane: _Lane) -> None: if self._closing: raise AnalyticsBusyException @@ -202,7 +213,8 @@ async def _admit_async(self, lane: _Lane) -> None: lane.queued.append(waiter) try: try: - await asyncio.wait_for(waiter.ready, timeout=self._queue_timeout) + async with asyncio.timeout(self._queue_timeout): + await waiter.ready except TimeoutError as error: raise AnalyticsBusyException from error if self._closing or monotonic() >= waiter.deadline: @@ -229,14 +241,24 @@ def _release(self, lane: _Lane) -> None: def _complete(self, *, lane: _Lane, work: _Operation, task: asyncio.Task[T]) -> None: if not task.cancelled(): - error = task.exception() - if work.abandoned and error is not None and not isinstance(error, AnalyticsTimeoutException): - logger.error("Analytics operation failed after its caller left.", exc_info=error) + task.exception() + self._finish(lane=lane, work=work) + if work.abandoned: + self._log_abandoned_error(task) + + def _finish(self, *, lane: _Lane, work: _Operation) -> None: lane.running.remove(work) work.task = None self._release(lane) work.finished.set_result(None) + def _abandon(self, *, work: _Operation, task: asyncio.Task[T]) -> None: + work.abandoned = True + work.control.cancel() + # Completion can precede cancellation delivery, after its callback consumed the exception. + if work.finished.done(): + self._log_abandoned_error(task) + async def _drain_async(self) -> None: await gather_with_cleanup_async(work.finished for lane in self._lanes.values() for work in lane.running) self._closed = True @@ -247,3 +269,10 @@ async def _execute_async(*, task: Callable[[QueryControl], Awaitable[T]], contro result = await task(control) control.check() return result + + @staticmethod + def _log_abandoned_error(task: asyncio.Task[T]) -> None: + if not task.cancelled(): + error = task.exception() + if error is not None and not isinstance(error, AnalyticsTimeoutException): + logger.error("Analytics operation failed after its caller left.", exc_info=error) diff --git a/tests/unit/analytics/test_attack_result_analytics.py b/tests/unit/analytics/test_attack_result_analytics.py index a9e1981cf3..a9a0eb7e60 100644 --- a/tests/unit/analytics/test_attack_result_analytics.py +++ b/tests/unit/analytics/test_attack_result_analytics.py @@ -11,7 +11,7 @@ import pytest from pydantic import ValidationError -from pyrit.analytics import AttackResultAnalytics +from pyrit.analytics import AttackResultAnalytics, compute_scenario_statistics from pyrit.common.task_utils import gather_with_cleanup_async from pyrit.exceptions.analytics_exception import AnalyticsBusyException, AnalyticsDataException from pyrit.memory import CentralMemory, MemoryInterface, SQLiteMemory @@ -28,9 +28,12 @@ AttackAnalyticsValue, AttackIdentifier, AttackOutcome, + AttackResultMetadata, + AttackResultRole, TargetIdentifier, ) from unit.memory.test_attack_analytics import make_result, predicate +from unit.mocks import make_scenario_result if TYPE_CHECKING: from collections.abc import AsyncGenerator @@ -133,6 +136,35 @@ async def test_overall_rate_is_not_an_average_of_groups( assert len(report.results.items) == 4 +async def test_retry_results_keep_raw_asr_distinct_from_scenario_unit_success( + analytics: AttackResultAnalytics, sqlite_instance: SQLiteMemory +) -> None: + results = [make_result(outcome=AttackOutcome.FAILURE), make_result(index=2)] + for result in results: + result.objective = "One objective, retried" + scenario = make_scenario_result(attack_results={"attack": results}) + statistics = compute_scenario_statistics(scenario) + assert statistics.overall.completed == statistics.overall.succeeded == 1 + assert statistics.overall.success_percentage == 100 + await sqlite_instance.add_attack_results_to_memory_async(attack_results=results) + report = await analytics.query_async() + assert report.summary.total_results == report.summary.total_decided == 2 + assert report.summary.success_rate == 0.5 + assert {row.attack_result_id for row in report.results.items} == {result.attack_result_id for result in results} + + +async def test_result_roles_do_not_silently_remove_saved_ids( + analytics: AttackResultAnalytics, sqlite_instance: SQLiteMemory +) -> None: + results = [make_result(index=index) for index, _ in enumerate(AttackResultRole, 1)] + for result, role in zip(results, AttackResultRole, strict=True): + result.attribution_data = AttackResultMetadata(result_role=role).to_metadata() + await sqlite_instance.add_attack_results_to_memory_async(attack_results=results) + report = await analytics.query_async() + assert report.summary.total_results == len(AttackResultRole) + assert {row.attack_result_id for row in report.results.items} == {result.attack_result_id for result in results} + + async def test_drilldown_appends_to_existing_any_and_response_all_predicates( analytics: AttackResultAnalytics, sqlite_instance: SQLiteMemory ) -> None: diff --git a/tests/unit/analytics/test_execution.py b/tests/unit/analytics/test_execution.py index 038c656d80..87857df375 100644 --- a/tests/unit/analytics/test_execution.py +++ b/tests/unit/analytics/test_execution.py @@ -5,6 +5,7 @@ import asyncio import gc +import inspect import threading import weakref from typing import TYPE_CHECKING @@ -17,7 +18,7 @@ from pyrit.exceptions.analytics_exception import AnalyticsBusyException, AnalyticsTimeoutException if TYPE_CHECKING: - from collections.abc import AsyncGenerator + from collections.abc import AsyncGenerator, Coroutine from pyrit.memory.query_control import QueryControl @@ -142,6 +143,34 @@ async def test_failures_propagate_and_release_capacity(harness: _Harness, error: assert await execution.run_async(report=True, task=successful.run_async) == "result" +@pytest.mark.parametrize("report", [True, False]) +async def test_failed_task_scheduling_releases_slot_and_closes_coroutine(report: bool) -> None: + execution = AnalyticsExecution(report_workers=1, quick_workers=1) + rejected: list[Coroutine[object, object, object]] = [] + operation = _Blocked() + operation.release.set() + + def reject(coroutine: Coroutine[object, object, object]) -> None: + rejected.append(coroutine) + raise RuntimeError("Task scheduling failed") + + try: + with patch.object(asyncio.get_running_loop(), "create_task", side_effect=reject): + with pytest.raises(RuntimeError, match="Task scheduling failed"): + await execution.run_async(report=report, task=operation.run_async) + assert rejected + assert all(inspect.getcoroutinestate(coroutine) == inspect.CORO_CLOSED for coroutine in rejected) + assert not operation.entered.is_set() + assert execution._lanes[report].active == 0 + assert not execution._lanes[report].running + assert await execution.run_async(report=report, task=operation.run_async) == "result" + finally: + for coroutine in rejected: + coroutine.close() + if not execution._lanes[report].running: + await execution.close_async() + + async def test_lanes_reserve_independent_running_and_queued_capacity(harness: _Harness) -> None: execution = harness.controller() report, quick, queued_report, queued_quick = [harness.blocked() for _ in range(4)] @@ -294,6 +323,28 @@ async def test_cancellation_after_admission_grant_returns_reserved_slot(harness: assert not lane.queued +@pytest.mark.parametrize("report", [True, False]) +async def test_cancellation_racing_with_slot_grant_never_starts_queued_work(harness: _Harness, report: bool) -> None: + execution = harness.controller() + active, queued = harness.blocked(), harness.blocked() + queued.release.set() + active_caller = harness.submit(execution=execution, operation=active, report=report) + await active.entered.wait() + queued_caller = harness.submit(execution=execution, operation=queued, report=report) + await asyncio.sleep(0) + lane = execution._lanes[report] + assert len(lane.queued) == 1 + work = next(iter(lane.running)) + work.task.add_done_callback(lambda _: queued_caller.cancel()) + active.release.set() + assert await active_caller == "result" + with pytest.raises(asyncio.CancelledError): + await queued_caller + assert not queued.entered.is_set() + assert lane.active == 0 + assert not lane.queued + + @pytest.mark.parametrize("expire_before_grant", [False, True]) async def test_admission_rechecks_deadline_even_without_timer_callback( harness: _Harness, expire_before_grant: bool @@ -331,6 +382,54 @@ async def test_unexpected_failure_after_cancellation_is_logged( assert "cleanup failed" in caplog.text +async def test_failure_is_logged_when_completion_precedes_caller_cancellation( + harness: _Harness, caplog: pytest.LogCaptureFixture +) -> None: + execution = harness.controller() + operation = harness.blocked(error=ValueError("query failed at cancellation")) + caller = harness.submit(execution=execution, operation=operation) + await operation.entered.wait() + work = next(iter(execution._lanes[True].running)) + work.task.add_done_callback(lambda _: caller.cancel()) + operation.release.set() + with pytest.raises(asyncio.CancelledError): + await caller + await execution.close_async() + assert caplog.text.count("Analytics operation failed after its caller left") == 1 + assert "query failed at cancellation" in caplog.text + + +async def test_failed_close_scheduling_can_be_retried_without_reopening_admission(harness: _Harness) -> None: + execution = harness.controller() + operation = harness.blocked() + caller = harness.submit(execution=execution, operation=operation) + await operation.entered.wait() + rejected: list[Coroutine[object, object, object]] = [] + + def reject(coroutine: Coroutine[object, object, object]) -> None: + rejected.append(coroutine) + raise RuntimeError("Shutdown scheduling failed") + + try: + with patch.object(asyncio.get_running_loop(), "create_task", side_effect=reject): + with pytest.raises(RuntimeError, match="Shutdown scheduling failed"): + await execution.close_async() + assert rejected + assert all(inspect.getcoroutinestate(coroutine) == inspect.CORO_CLOSED for coroutine in rejected) + assert not execution.is_closed + assert operation.control.cancel_event.is_set() + assert not operation.exited.is_set() + with pytest.raises(AnalyticsBusyException): + await execution.run_async(report=True, task=operation.run_async) + finally: + for coroutine in rejected: + coroutine.close() + operation.release.set() + await execution.close_async() + with pytest.raises(AnalyticsTimeoutException): + await caller + + async def test_close_rejects_queued_work_and_drains_through_repeated_cancellation(harness: _Harness) -> None: execution = harness.controller() active, queued = harness.blocked(), harness.blocked() From 2d969828c8b7772cbbd0c935c5f62ac646523dfd Mon Sep 17 00:00:00 2001 From: Roman Lutz Date: Thu, 8 Oct 2026 16:52:28 -0700 Subject: [PATCH 3/8] FEAT: Share outcome statistics across attack and scenario analytics Expose decided-only and all-outcome success rates through one calculator and shared model. Keep raw saved-result selection distinct from scenario latest-unit selection while reusing count validation, rates, shares, and percentage formatting. Preserve legacy defaults and carry both rates through scenario progress and JSON projections. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- doc/code/analytics/0_attack_results.md | 76 +++++++++-- doc/code/framework.md | 5 +- pyrit/analytics/__init__.py | 3 + pyrit/analytics/attack_result_analytics.py | 35 +---- pyrit/analytics/outcome_statistics.py | 122 ++++++++++++++++++ pyrit/analytics/result_analysis.py | 15 ++- pyrit/analytics/scenario_statistics.py | 46 +++---- .../backend/services/scenario_run_service.py | 3 +- pyrit/models/__init__.py | 2 + pyrit/models/analytics.py | 20 ++- pyrit/models/scenario_progress.py | 11 +- pyrit/output/_derivation.py | 6 +- pyrit/output/scenario_result/json.py | 3 + .../analytics/test_attack_result_analytics.py | 34 +++++ .../unit/analytics/test_outcome_statistics.py | 99 ++++++++++++++ .../analytics/test_scenario_statistics.py | 77 +++++++++++ .../test_scenario_statistics_parity.py | 8 +- .../unit/common/test_lazy_package_imports.py | 16 ++- tests/unit/models/test_analytics.py | 28 ++++ .../unit/output/scenario_result/test_json.py | 11 +- tests/unit/output/test_derivation.py | 28 +++- 21 files changed, 570 insertions(+), 78 deletions(-) create mode 100644 pyrit/analytics/outcome_statistics.py create mode 100644 tests/unit/analytics/test_outcome_statistics.py diff --git a/doc/code/analytics/0_attack_results.md b/doc/code/analytics/0_attack_results.md index bfd277436a..ffcf7d5861 100644 --- a/doc/code/analytics/0_attack_results.md +++ b/doc/code/analytics/0_attack_results.md @@ -13,18 +13,27 @@ a conversation, including failed attempts followed by successful retries, remain separate. Updating an existing ID changes its current outcome, not the result count. Selecting a scenario is a cohort filter, not a switch to scenario-unit statistics. -ASR uses `successes / (successes + failures)`, reusing the existing `AttackStats` -outcome policy. With no successes or failures, ASR is `None`. Errors and -undetermined outcomes remain visible but do not enter that denominator. +The shared `compute_outcome_statistics` calculator returns both denominator +policies together: + +| Field | Denominator | Meaning | +|---|---|---| +| `success_rate` | Successes + failures | Success among decided outcomes, the existing ASR default. | +| `success_rate_all` | All four outcomes | Success among every selected result, including errors and undetermined outcomes. | + +An empty denominator produces `None`. For example, one success and one error +produce `success_rate=1.0` and `success_rate_all=0.5`. An error-only population has +`success_rate=None` and `success_rate_all=0.0`. Neither rate requires a second query. `total_results` includes all four outcomes; `outcome_shares` and `decided_share` use that whole selected cohort. Empty outcome shares are zero, and an empty cohort's decided share is `None`. All proportions are between 0 and 1. -Overall ASR is calculated from overall counts, never by averaging subgroup rates. +Both success rates are calculated from counts, never by averaging subgroup rates. -This differs from latest-execution-unit scenario success statistics. A failed +Population selection is separate from denominator selection. A failed attempt followed by a successful retry yields 50% raw-result ASR even if the -scenario's latest-unit success rate is 100%. No result-role filtering or inference -from conversation presence is applied. +scenario's latest-unit success rate is 100%. Both analyses have both denominator +policies; the difference in this example is which attempts count. No result-role +filtering or inference from conversation presence is applied. Outcome restrictions apply to the entire report and its result page. `outcome_filter_applied` asks a renderer to annotate the rate as **ASR*** and @@ -65,7 +74,7 @@ async with AttackResultAnalytics() as analytics: compare_by=AttackAnalyticsDimension(name="attack_type"), ) ) - print(report.summary.total_results, report.summary.success_rate) + print(report.summary.total_results, report.summary.success_rate, report.summary.success_rate_all) if report.drilldown_unavailable_reason: print(report.drilldown_unavailable_reason) @@ -89,6 +98,57 @@ async with AttackResultAnalytics() as analytics: ) ``` +### The same statistics for scenario units + +`compute_scenario_statistics` selects each scenario execution unit's latest +attempt, then calls the same outcome calculator. Its overall, atomic-attack, and +display-group counts carry an `outcomes` object of the same `OutcomeStatistics` +type as attack-report summaries, groups, and cells. + +```python +from pyrit.analytics import compute_scenario_statistics + +statistics = compute_scenario_statistics(scenario_result) +outcomes = statistics.overall.outcomes +assert outcomes is not None # Always populated by compute_scenario_statistics. +print(outcomes.success_rate, outcomes.success_rate_all) +``` + +`ScenarioProgressCounts.success_percentage` retains its existing all-completed-unit +denominator and truncated 0-100 representation for compatibility. Both rates under +`outcomes` are unrounded 0-1 proportions. Historical `errors` and `retries` remain +separate: an error followed by a successful retry contributes one historical error +but no latest-unit error. Never derive the decided denominator by subtracting +historical errors from completed units. + +Scenario progress projections and JSON reports preserve the shared `outcomes` +object; existing displayed percentages do not change. Count-only scenario payloads +from older callers deserialize with `outcomes=None` because their latest-outcome +breakdown cannot be reconstructed. Combining a nonempty such payload logs a warning +and leaves the combined breakdown unavailable rather than guessing failures. + +For already counted populations, call `compute_outcome_statistics` directly. +`combine_outcome_statistics` sums counts from **disjoint** populations and +recalculates both rates. It also accepts existing `AttackStats` returned by +`analyze_results` or technique analytics; those APIs retain their existing result +shape and decided-only default but use the same shared calculation internally. +`OutcomeStatistics` is an alias for the existing `AttackAnalyticsStatistics` class, +not a second representation. + +```python +from pyrit.analytics import combine_outcome_statistics, compute_outcome_statistics + +first = compute_outcome_statistics({"success": 1, "error": 1}) +second = compute_outcome_statistics({"success": 2, "failure": 1}) +combined = combine_outcome_statistics([first, second]) +print(combined.success_rate, combined.success_rate_all) # 0.75, 0.6 +``` + +Do not combine overlapping converter/harm groups to reconstruct a report's +summary; use its already-computed summary instead. + +### Filters and drill-downs + Omit `compare_by` for a one-dimensional breakdown. Supported dimensions include operation, operator, targeted harm category, attack type, converter type, objective target, model, scenario, and custom labels. Custom labels require a diff --git a/doc/code/framework.md b/doc/code/framework.md index 5fe994527c..ba2fdb3fa9 100644 --- a/doc/code/framework.md +++ b/doc/code/framework.md @@ -334,10 +334,11 @@ The below talks about responsibilities of most modules in the PyRIT library - This is where cross-run analysis belongs: e.g. "which attack performed best for this objective?", "how often did a technique succeed?", or "which responses match known content?". - **Does not own**: live, in-attack decisions — any decision made *during* an attack is a scorer's job. Analytics only operates on stored results, after the fact. - Today it includes `ConversationAnalytics` (inspecting conversation history), `analyze_results` / `AttackStats` (aggregating outcomes across techniques), and text-matching strategies (`ExactTextMatching`, `ApproximateTextMatching`). -- `compute_scenario_statistics` calculates scenario success statistics. It owns execution-unit identity (atomic attack, technique configuration, and seed group), latest-attempt selection, counts, denominators, and rounding. SDK callers, the GUI backend's run detail and progress views, and the console, JSON, and HTML reports all present its results (`ScenarioExecutionStatistics`, `ScenarioExecutionUnit`, and `ScenarioProgressCounts` in `pyrit.models`) instead of calculating their own. The one exception is the GUI run-history list, which aggregates the same statistics in SQL (`MemoryInterface._build_scenario_history_aggregate_statement`) so it can page over many runs; `tests/unit/analytics/test_scenario_statistics_parity.py` keeps the two implementations in agreement. +- `compute_scenario_statistics` owns execution-unit identity (atomic attack, technique configuration, and seed group), latest-attempt selection, and historical retry/error counts. It passes selected outcomes to the shared outcome calculator. SDK callers, the GUI backend's run detail and progress views, and the console, JSON, and HTML reports present its results (`ScenarioExecutionStatistics`, `ScenarioExecutionUnit`, and `ScenarioProgressCounts` in `pyrit.models`). The GUI run-history list aggregates counts in SQL (`MemoryInterface._build_scenario_history_aggregate_statement`) so it can page over many runs, then uses the shared percentage helper; `tests/unit/analytics/test_scenario_statistics_parity.py` keeps the paths in agreement. - Scenario attempts are ordered by timestamp, then their canonical lowercase UUID string. SQL Server uses this string order rather than its native UUID order for history ranking and progress pagination. Explicit seed attribution wins; an objective alone matches a planned seed group only when that match is unique. Legacy runs with identifier-only seed identities are recounted with the shared analytics. - Shared analytics contracts (filters, dimensions, typed values, reports, facets, result pages, and `AttackStats`) live in `pyrit.models.analytics`. They validate data without querying memory or calculating statistics. `AttackResultSelection` defines selection modes without changing existing callers. -- [`AttackResultAnalytics`](./analytics/0_attack_results.md) provides async saved-result reports, lightweight result pages, and facet lookups. It reuses raw-outcome `AttackStats` policy, supplies display labels and exact additional drill-down predicates, and counts every saved result ID rather than latest scenario execution units. Memory owns cohort SQL and consistent projections; the SDK owns their interpretation. +- `compute_outcome_statistics` owns outcome-count validation, both success-rate denominators, totals, and shares. `success_rate` divides by successes plus failures; `success_rate_all` divides by all outcomes. `combine_outcome_statistics` combines disjoint counts, never averages rates. Attack and scenario analytics share the same `OutcomeStatistics` model and calculations after independently selecting their populations. Legacy `AttackStats` and scenario percentage fields retain their existing shapes/defaults. +- [`AttackResultAnalytics`](./analytics/0_attack_results.md) provides async saved-result reports, lightweight result pages, and facet lookups. It supplies both shared success rates, display labels and exact additional drill-down predicates, and counts every saved result ID rather than latest scenario execution units. Memory owns cohort SQL and consistent projections; the SDK owns their interpretation. - Analytics facades sharing a memory instance share one loop-bound, bounded report/quick-query controller. Close that shared SDK lifetime before replacing or disposing memory; cancellation does not free a running query's slot before its actual session cleanup. SDK lifecycle is independent of backend runtime integration. - Filter-bound cursor and label-normalization helpers live in `pyrit.common.pagination`. The backend pagination module retains compatibility exports, including History's invalid-cursor first-page fallback. diff --git a/pyrit/analytics/__init__.py b/pyrit/analytics/__init__.py index fd45a9944e..4f0ced28a1 100644 --- a/pyrit/analytics/__init__.py +++ b/pyrit/analytics/__init__.py @@ -11,6 +11,7 @@ if TYPE_CHECKING: from pyrit.analytics.attack_result_analytics import AttackResultAnalytics from pyrit.analytics.conversation_analytics import ConversationAnalytics + from pyrit.analytics.outcome_statistics import combine_outcome_statistics, compute_outcome_statistics from pyrit.analytics.result_analysis import ( AttackStats, analyze_results, @@ -25,6 +26,8 @@ "ApproximateTextMatching": "pyrit.analytics.text_matching", "AttackResultAnalytics": "pyrit.analytics.attack_result_analytics", "AttackStats": "pyrit.analytics.result_analysis", + "combine_outcome_statistics": "pyrit.analytics.outcome_statistics", + "compute_outcome_statistics": "pyrit.analytics.outcome_statistics", "compute_scenario_statistics": "pyrit.analytics.scenario_statistics", "ConversationAnalytics": "pyrit.analytics.conversation_analytics", "ExactTextMatching": "pyrit.analytics.text_matching", diff --git a/pyrit/analytics/attack_result_analytics.py b/pyrit/analytics/attack_result_analytics.py index 5ebfa62a76..22e7e53bb3 100644 --- a/pyrit/analytics/attack_result_analytics.py +++ b/pyrit/analytics/attack_result_analytics.py @@ -13,7 +13,7 @@ from pyrit.analytics._execution import AnalyticsExecution from pyrit.analytics._profile_aggregation import ProfileAggregation -from pyrit.analytics.result_analysis import _compute_stats +from pyrit.analytics.outcome_statistics import compute_outcome_statistics from pyrit.exceptions.analytics_exception import AnalyticsBusyException, AnalyticsDataException from pyrit.memory import CentralMemory from pyrit.memory.attack_analytics import AttackAnalyticsReader, RawAnalyticsOption, RawAnalyticsReport @@ -34,7 +34,6 @@ AttackAnalyticsStatistics, AttackAnalyticsValue, AttackAnalyticsValueKind, - AttackOutcome, ) if TYPE_CHECKING: @@ -266,38 +265,18 @@ def _cells(cls, *, query: AttackAnalyticsQuery, raw: RawAnalyticsReport) -> list @staticmethod def _statistics(counts: dict[str, int]) -> AttackAnalyticsStatistics: """ - Reuse raw-outcome ASR policy; total shares include errors and undetermined results. + Apply shared outcome statistics, retaining the SDK's stored-data error contract. Returns: - AttackAnalyticsStatistics: The existing decided-result rate plus whole-cohort shares. + AttackAnalyticsStatistics: Both success rates and whole-cohort outcome shares. Raises: AnalyticsDataException: If saved outcomes or their counts are invalid. """ - if set(counts) - {outcome.value for outcome in AttackOutcome}: - raise AnalyticsDataException("Stored results contain an unsupported attack outcome.") - if any(type(count) is not int or count < 0 for count in counts.values()): - raise AnalyticsDataException("Stored results contain invalid outcome counts.") - stats = _compute_stats( - successes=counts.get(AttackOutcome.SUCCESS.value, 0), - failures=counts.get(AttackOutcome.FAILURE.value, 0), - undetermined=counts.get(AttackOutcome.UNDETERMINED.value, 0), - errors=counts.get(AttackOutcome.ERROR.value, 0), - ) - total = stats.total_decided + stats.undetermined + stats.errors - return AttackAnalyticsStatistics( - success_rate=stats.success_rate, - total_decided=stats.total_decided, - successes=stats.successes, - failures=stats.failures, - undetermined=stats.undetermined, - errors=stats.errors, - total_results=total, - decided_share=stats.total_decided / total if total else None, - outcome_shares={ - outcome: counts.get(outcome.value, 0) / total if total else 0.0 for outcome in AttackOutcome - }, - ) + try: + return compute_outcome_statistics(counts) + except ValueError as error: + raise AnalyticsDataException(str(error)) from error @staticmethod def _option(*, raw: RawAnalyticsOption, dimension: AttackAnalyticsDimension) -> AttackAnalyticsOption: diff --git a/pyrit/analytics/outcome_statistics.py b/pyrit/analytics/outcome_statistics.py new file mode 100644 index 0000000000..1642044fcf --- /dev/null +++ b/pyrit/analytics/outcome_statistics.py @@ -0,0 +1,122 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Outcome-count calculations shared by attack and scenario analytics.""" + +from __future__ import annotations + +from collections import Counter +from typing import TYPE_CHECKING + +from pyrit.models import AttackOutcome, OutcomeStatistics + +if TYPE_CHECKING: + from collections.abc import Iterable, Mapping + + from pyrit.models import AttackStats + + +def compute_outcome_statistics(counts: Mapping[str, int] | Mapping[AttackOutcome, int]) -> OutcomeStatistics: + """ + Calculate both success rates for an already selected population. + + This function neither queries memory nor chooses result IDs, retries, roles, + or scenario units. ``success_rate`` uses success + failure; ``success_rate_all`` + uses every outcome. Both are proportions, not rounded percentages. + + Args: + counts (Mapping[str, int] | Mapping[AttackOutcome, int]): Nonnegative saved-outcome + counts. Omitted outcomes mean zero observations. + + Returns: + OutcomeStatistics: Counts, both denominator policies, and whole-population shares. + Empty denominators produce None; empty outcome shares are zero. + + Raises: + ValueError: If an outcome is unsupported or a count is not a nonnegative integer. + """ + values = _validated_counts(counts) + successes = values.get(AttackOutcome.SUCCESS, 0) + failures = values.get(AttackOutcome.FAILURE, 0) + undetermined = values.get(AttackOutcome.UNDETERMINED, 0) + errors = values.get(AttackOutcome.ERROR, 0) + decided = successes + failures + total = decided + undetermined + errors + return OutcomeStatistics( + successes=successes, + failures=failures, + undetermined=undetermined, + errors=errors, + total_decided=decided, + total_results=total, + success_rate=_rate(successes=successes, total=decided), + success_rate_all=_rate(successes=successes, total=total), + decided_share=_rate(successes=decided, total=total), + outcome_shares={ + outcome: _rate(successes=values.get(outcome, 0), total=total) or 0.0 for outcome in AttackOutcome + }, + ) + + +def combine_outcome_statistics(statistics: Iterable[AttackStats]) -> OutcomeStatistics: + """ + Sum disjoint populations' outcome counts and recalculate both rates. + + Never average subgroup percentages. The caller must ensure its populations + are disjoint; overlapping harm/converter groups cannot reconstruct a cohort. + Existing ``AttackStats`` objects are accepted without changing their legacy shape. + + Args: + statistics (Iterable[AttackStats]): Count-bearing statistics for disjoint populations. + + Returns: + OutcomeStatistics: Statistics from the combined counts, including an empty population. + + Raises: + ValueError: If any input count is invalid, even if another input would offset it. + """ + counts: Counter[AttackOutcome] = Counter() + for item in statistics: + counts.update( + _validated_counts( + { + AttackOutcome.SUCCESS: item.successes, + AttackOutcome.FAILURE: item.failures, + AttackOutcome.UNDETERMINED: item.undetermined, + AttackOutcome.ERROR: item.errors, + } + ) + ) + return compute_outcome_statistics(counts) + + +def success_percentage(*, succeeded: int, completed: int) -> int | None: + """ + Format a chosen denominator's success rate as a legacy integer percentage. + + Args: + succeeded (int): Successful observations. + completed (int): Observations in the chosen denominator. + + Returns: + int | None: The shared rate truncated to an integer percentage, or None if empty. + """ + rate = _rate(successes=succeeded, total=completed) + return int(rate * 100) if rate is not None else None + + +def _validated_counts(counts: Mapping[str, int] | Mapping[AttackOutcome, int]) -> dict[AttackOutcome, int]: + values: dict[AttackOutcome, int] = {} + for name, count in counts.items(): + try: + outcome = AttackOutcome(name) + except ValueError as error: + raise ValueError("Stored results contain an unsupported attack outcome.") from error + if type(count) is not int or count < 0: + raise ValueError("Stored results contain invalid outcome counts.") + values[outcome] = count + return values + + +def _rate(*, successes: int, total: int) -> float | None: + return successes / total if total else None diff --git a/pyrit/analytics/result_analysis.py b/pyrit/analytics/result_analysis.py index 4106fbadb0..574160abb7 100644 --- a/pyrit/analytics/result_analysis.py +++ b/pyrit/analytics/result_analysis.py @@ -5,6 +5,7 @@ from collections.abc import Sequence from typing import TYPE_CHECKING +from pyrit.analytics.outcome_statistics import compute_outcome_statistics from pyrit.common.deprecation import print_deprecation_message from pyrit.models import ( AttackOutcome, @@ -22,11 +23,17 @@ def _compute_stats(successes: int, failures: int, undetermined: int, errors: int) -> AttackStats: - total_decided = successes + failures - success_rate = successes / total_decided if total_decided > 0 else None + statistics = compute_outcome_statistics( + { + AttackOutcome.SUCCESS: successes, + AttackOutcome.FAILURE: failures, + AttackOutcome.UNDETERMINED: undetermined, + AttackOutcome.ERROR: errors, + } + ) return AttackStats( - success_rate=success_rate, - total_decided=total_decided, + success_rate=statistics.success_rate, + total_decided=statistics.total_decided, successes=successes, failures=failures, undetermined=undetermined, diff --git a/pyrit/analytics/scenario_statistics.py b/pyrit/analytics/scenario_statistics.py index 2b483ca6f1..3c2f4a3195 100644 --- a/pyrit/analytics/scenario_statistics.py +++ b/pyrit/analytics/scenario_statistics.py @@ -11,21 +11,27 @@ - execution-unit identity: an atomic group (atomic attack name plus technique configuration) and a logical seed group, resolved against the saved run plan when one exists; - attempt selection: each unit counts once, by its latest attempt (timestamp, then attempt ID); -- counts, denominators, and rounding: the success percentage is succeeded units over completed units, - truncated to an integer. +- unit outcome counts, passed to the shared outcome calculator for both denominator policies. Historical attempt, error, and retry counts are reported separately from the effective-unit counts. +The legacy success percentage remains succeeded units over all completed units. """ from __future__ import annotations import logging +from collections import Counter from dataclasses import dataclass from datetime import UTC, datetime from typing import TYPE_CHECKING, Protocol from pydantic import ValidationError +from pyrit.analytics.outcome_statistics import ( + combine_outcome_statistics, + compute_outcome_statistics, + success_percentage, +) from pyrit.common.utils import to_sha256 from pyrit.models import ( SCENARIO_RUN_PLAN_METADATA_KEY, @@ -255,17 +261,6 @@ def retry_pressure(*, attempts_per_unit: Iterable[int], persisted_retries: Itera return within_attempts + repeated_units -def success_percentage(*, succeeded: int, completed: int) -> int | None: - """ - Return the success percentage for effective execution units. - - Returns: - int | None: ``succeeded / completed`` as a truncated integer percentage, or None with no - completed units. - """ - return int((succeeded / completed) * 100) if completed else None - - def count_execution_units( *, units: Iterable[ScenarioExecutionUnit], @@ -279,30 +274,30 @@ def count_execution_units( oldest first; the last attempt decides the unit's outcome. Returns: - ScenarioProgressCounts: Completed and succeeded units plus historical errors and retries. + ScenarioProgressCounts: Shared statistics for latest unit outcomes plus historical errors and retries. """ - completed = 0 - succeeded = 0 + counts: Counter[AttackOutcome] = Counter() errors = 0 retries = 0 for unit in units: attempts = attempts_by_unit.get(unit, ()) if not attempts: continue - completed += 1 - succeeded += int(attempts[-1].outcome == AttackOutcome.SUCCESS) + counts[attempts[-1].outcome] += 1 errors += sum(int(attempt.outcome == AttackOutcome.ERROR) for attempt in attempts) retries += retry_pressure( attempts_per_unit=[len(attempts)], persisted_retries=[attempt.total_retries for attempt in attempts], ) + outcomes = compute_outcome_statistics(counts) return ScenarioProgressCounts( - completed=completed, + completed=outcomes.total_results, planned=planned, - succeeded=succeeded, - success_percentage=success_percentage(succeeded=succeeded, completed=completed), + succeeded=outcomes.successes, + success_percentage=success_percentage(succeeded=outcomes.successes, completed=outcomes.total_results), errors=errors, retries=retries, + outcomes=outcomes, ) @@ -312,12 +307,18 @@ def combine_execution_counts(counts: Iterable[ScenarioProgressCounts]) -> Scenar Returns: ScenarioProgressCounts: The summed counts with the success percentage recomputed. ``planned`` is - None unless every input has one. + None unless every input has one. ``outcomes`` is None if a nonempty legacy input + lacks its breakdown; historical errors cannot reconstruct latest outcomes. """ counts = list(counts) completed = sum(item.completed for item in counts) succeeded = sum(item.succeeded for item in counts) planned = [item.planned for item in counts] + if any(item.completed and item.outcomes is None for item in counts): + logger.warning("Cannot combine outcome statistics from legacy scenario counts without an outcome breakdown.") + outcomes = None + else: + outcomes = combine_outcome_statistics(item.outcomes for item in counts if item.outcomes is not None) return ScenarioProgressCounts( completed=completed, planned=sum(value for value in planned if value is not None) if all(v is not None for v in planned) else None, @@ -325,6 +326,7 @@ def combine_execution_counts(counts: Iterable[ScenarioProgressCounts]) -> Scenar success_percentage=success_percentage(succeeded=succeeded, completed=completed), errors=sum(item.errors for item in counts), retries=sum(item.retries for item in counts), + outcomes=outcomes, ) diff --git a/pyrit/backend/services/scenario_run_service.py b/pyrit/backend/services/scenario_run_service.py index aa0339e32d..74b548076e 100644 --- a/pyrit/backend/services/scenario_run_service.py +++ b/pyrit/backend/services/scenario_run_service.py @@ -25,6 +25,7 @@ from pydantic import TypeAdapter, ValidationError +from pyrit.analytics.outcome_statistics import success_percentage from pyrit.analytics.scenario_statistics import compute_scenario_statistics from pyrit.backend.models.common import PaginationInfo, filter_sensitive_fields from pyrit.backend.models.scenarios import ScenarioRunListResponse @@ -1695,7 +1696,7 @@ def _build_history_summary( techniques_used=techniques, total_attacks=planned_total if atomic_groups is not None or planned_total else None, completed_attacks=completed, - objective_achieved_rate=int((successful / completed) * 100) if completed else 0, + objective_achieved_rate=success_percentage(succeeded=successful, completed=completed) or 0, total_retries=aggregate.total_retries, labels=record.labels, completed_at=record.completed_at if terminal else None, diff --git a/pyrit/models/__init__.py b/pyrit/models/__init__.py index 5dc925a098..0564d680da 100644 --- a/pyrit/models/__init__.py +++ b/pyrit/models/__init__.py @@ -44,6 +44,7 @@ AttackAnalyticsValueKind, AttackResultSelection, AttackStats, + OutcomeStatistics, ) from pyrit.models.catalog import ( ScenarioDatasetSizeCap, @@ -292,6 +293,7 @@ "AttackAnalyticsValueKind": "pyrit.models.analytics", "AttackResultSelection": "pyrit.models.analytics", "AttackStats": "pyrit.models.analytics", + "OutcomeStatistics": "pyrit.models.analytics", "Acquisition": "pyrit.models.score", "AnswerMatches": "pyrit.models.score", "ALLOWED_CHAT_MESSAGE_ROLES": "pyrit.models.messages.chat_message", diff --git a/pyrit/models/analytics.py b/pyrit/models/analytics.py index 53cda714a1..1670e6376a 100644 --- a/pyrit/models/analytics.py +++ b/pyrit/models/analytics.py @@ -5,7 +5,7 @@ from __future__ import annotations -from dataclasses import dataclass +from dataclasses import dataclass, field from datetime import UTC, datetime from enum import Enum from typing import ClassVar, Self @@ -30,17 +30,27 @@ class AttackStats: @dataclass class AttackAnalyticsStatistics(AttackStats): """ - Outcome statistics for one cohort, group, or heatmap cell. + Shared outcome statistics for saved results or selected scenario execution units. ``success_rate`` uses successes / decided results; errors and undetermined - outcomes do not enter that denominator. ``decided_share`` and ``outcome_shares`` - instead use all results. Rates with no applicable denominator are ``None``; - shares for an empty cohort are zero. All proportions are in the range 0 through 1. + outcomes do not enter that denominator. ``success_rate_all`` instead includes + every outcome in its denominator, as do ``decided_share`` and ``outcome_shares``. + Rates with no applicable denominator are ``None``; shares for an empty cohort + are zero. All proportions are in the range 0 through 1. + + Analytics calculates these values after selecting the counted population. + ``total_results`` can therefore describe distinct saved IDs or latest scenario + units. The model does not select attempts or infer a counting policy. """ total_results: int decided_share: float | None outcome_shares: dict[AttackOutcome, float] + success_rate_all: float | None = field(default=None, kw_only=True) + + +# Preserve the existing class/constructor identity while sharing it beyond attack reports. +OutcomeStatistics = AttackAnalyticsStatistics class AttackResultSelection(str, Enum): diff --git a/pyrit/models/scenario_progress.py b/pyrit/models/scenario_progress.py index 6da0733b4d..8bad575b10 100644 --- a/pyrit/models/scenario_progress.py +++ b/pyrit/models/scenario_progress.py @@ -9,6 +9,7 @@ from pydantic import AwareDatetime, BaseModel, ConfigDict, Field, model_validator +from pyrit.models.analytics import OutcomeStatistics from pyrit.models.catalog.scenario import ScenarioOverloadSummary, ScenarioTargetSummary # noqa: TC001 from pyrit.models.identifiers.atomic_attack_identifier import AtomicAttackIdentifier from pyrit.models.results.attack_result import AttackOutcome, AttackResultRole @@ -183,7 +184,14 @@ class ScenarioProgressResult(BaseModel): class ScenarioProgressCounts(BaseModel): - """Canonical progress counts for a set of scenario execution units.""" + """ + Canonical progress counts for a set of scenario execution units. + + ``outcomes`` describes only the selected latest attempts and supplies both + decided-only and all-outcome success rates. ``errors`` and ``retries`` retain + their historical-attempt meaning and must not be used as outcome denominators. + None preserves older count-only payloads whose outcome breakdown is unknown. + """ completed: int = Field(..., ge=0) planned: int | None = Field(default=None, ge=0) @@ -191,6 +199,7 @@ class ScenarioProgressCounts(BaseModel): success_percentage: int | None = Field(default=None, ge=0, le=100) errors: int = Field(..., ge=0) retries: int = Field(..., ge=0) + outcomes: OutcomeStatistics | None = None class ScenarioExecutionUnit(BaseModel): diff --git a/pyrit/output/_derivation.py b/pyrit/output/_derivation.py index bd03cb5321..5de0b7a5bc 100644 --- a/pyrit/output/_derivation.py +++ b/pyrit/output/_derivation.py @@ -18,7 +18,7 @@ from pyrit.analytics.scenario_statistics import combine_execution_counts, compute_scenario_statistics if TYPE_CHECKING: - from pyrit.models import AttackResult, ComponentIdentifier, MessagePiece, ScenarioResult, Score + from pyrit.models import AttackResult, ComponentIdentifier, MessagePiece, OutcomeStatistics, ScenarioResult, Score class TargetInfo(NamedTuple): @@ -61,6 +61,7 @@ class GroupStatistics(NamedTuple): objective_executions: int attempts: int success_rate: int + outcomes: OutcomeStatistics | None = None class ScenarioOverview(NamedTuple): @@ -70,6 +71,7 @@ class ScenarioOverview(NamedTuple): attempts: int success_rate: int groups: list[GroupStatistics] + outcomes: OutcomeStatistics | None = None def scenario_overview(result: ScenarioResult) -> ScenarioOverview: @@ -100,6 +102,7 @@ def scenario_overview(result: ScenarioResult) -> ScenarioOverview: objective_executions=counts.completed, attempts=len(group_results), success_rate=counts.success_percentage or 0, + outcomes=counts.outcomes, ) ) return ScenarioOverview( @@ -107,6 +110,7 @@ def scenario_overview(result: ScenarioResult) -> ScenarioOverview: attempts=statistics.attempts, success_rate=statistics.overall.success_percentage or 0, groups=groups, + outcomes=statistics.overall.outcomes, ) diff --git a/pyrit/output/scenario_result/json.py b/pyrit/output/scenario_result/json.py index 019450c1ac..948d82f5e4 100644 --- a/pyrit/output/scenario_result/json.py +++ b/pyrit/output/scenario_result/json.py @@ -2,6 +2,7 @@ # Licensed under the MIT license. import json +from dataclasses import asdict from typing import TYPE_CHECKING, Any from pyrit.models import AttackResult, ScenarioResult @@ -129,6 +130,7 @@ def _build_overview(self, result: ScenarioResult) -> dict[str, Any]: "num_objective_executions": group.objective_executions, "num_attempts": group.attempts, "success_rate": group.success_rate, + "outcomes": asdict(group.outcomes) if group.outcomes is not None else None, } for group in overview.groups ] @@ -161,6 +163,7 @@ def _build_overview(self, result: ScenarioResult) -> dict[str, Any]: "total_objective_executions": overview.objective_executions, "total_attempts": overview.attempts, "overall_success_rate": overview.success_rate, + "outcomes": asdict(overview.outcomes) if overview.outcomes is not None else None, "unique_objectives": len(result.get_objectives()), }, "groups": groups, diff --git a/tests/unit/analytics/test_attack_result_analytics.py b/tests/unit/analytics/test_attack_result_analytics.py index a9a0eb7e60..c18dd630b1 100644 --- a/tests/unit/analytics/test_attack_result_analytics.py +++ b/tests/unit/analytics/test_attack_result_analytics.py @@ -64,6 +64,7 @@ async def test_all_saved_ids_and_outcomes_use_raw_decided_denominator(analytics: assert report.summary.total_results == 10 assert report.summary.total_decided == 6 assert report.summary.success_rate == pytest.approx(4 / 6) + assert report.summary.success_rate_all == 0.4 assert report.summary.decided_share == 0.6 assert report.summary.outcome_shares == { AttackOutcome.SUCCESS: 0.4, @@ -82,6 +83,7 @@ async def test_empty_report_has_unavailable_rates(analytics: AttackResultAnalyti report = await analytics.query_async() assert report.summary.total_results == report.summary.total_decided == 0 assert report.summary.success_rate is None + assert report.summary.success_rate_all is None assert report.summary.decided_share is None assert set(report.summary.outcome_shares.values()) == {0.0} assert report.groups == report.cells == report.results.items == [] @@ -104,6 +106,7 @@ async def test_outcome_filter_annotates_and_restricts_entire_report( assert report.outcome_filter_applied assert report.summary.total_results == count assert report.summary.success_rate == rate + assert report.summary.success_rate_all == (1.0 if outcome is AttackOutcome.SUCCESS else 0.0) assert report.groups[0].statistics == report.summary assert all(row.outcome == outcome for row in report.results.items) @@ -146,13 +149,42 @@ async def test_retry_results_keep_raw_asr_distinct_from_scenario_unit_success( statistics = compute_scenario_statistics(scenario) assert statistics.overall.completed == statistics.overall.succeeded == 1 assert statistics.overall.success_percentage == 100 + assert statistics.overall.outcomes.success_rate == statistics.overall.outcomes.success_rate_all == 1.0 await sqlite_instance.add_attack_results_to_memory_async(attack_results=results) report = await analytics.query_async() assert report.summary.total_results == report.summary.total_decided == 2 assert report.summary.success_rate == 0.5 + assert report.summary.success_rate_all == 0.5 assert {row.attack_result_id for row in report.results.items} == {result.attack_result_id for result in results} +@pytest.mark.parametrize( + "outcomes", + [ + [], + [AttackOutcome.SUCCESS, AttackOutcome.ERROR], + [AttackOutcome.FAILURE, AttackOutcome.UNDETERMINED], + list(AttackOutcome), + [AttackOutcome.SUCCESS, AttackOutcome.FAILURE, AttackOutcome.ERROR, AttackOutcome.ERROR], + ], +) +@pytest.mark.parametrize("compare", [False, True]) +async def test_same_population_has_identical_shared_statistics_in_attacks_and_scenarios( + analytics: AttackResultAnalytics, sqlite_instance: SQLiteMemory, outcomes: list[AttackOutcome], compare: bool +) -> None: + results = [make_result(index=index, outcome=outcome) for index, outcome in enumerate(outcomes, 1)] + if results: + await sqlite_instance.add_attack_results_to_memory_async(attack_results=results) + scenario = make_scenario_result(attack_results={"attack": results}) + expected = compute_scenario_statistics(scenario).overall.outcomes + report = await analytics.query_async( + query=AttackAnalyticsQuery(compare_by=AttackAnalyticsDimension(name="converter_type") if compare else None) + ) + assert report.summary == expected + assert all(group.statistics == expected for group in report.groups) + assert all(cell.statistics == expected for cell in report.cells) + + async def test_result_roles_do_not_silently_remove_saved_ids( analytics: AttackResultAnalytics, sqlite_instance: SQLiteMemory ) -> None: @@ -268,6 +300,7 @@ async def test_matrix_empty_cells_and_additional_predicates_are_exact( assert len(page.items) == cell.statistics.total_results if not page.items: assert cell.statistics.success_rate is None + assert cell.statistics.success_rate_all is None assert set(cell.statistics.outcome_shares.values()) == {0.0} @@ -359,6 +392,7 @@ async def test_reports_neither_share_mutable_results_nor_cache_outcomes( assert updated.summary.total_results == 1 assert updated.groups[0].statistics.failures == 1 assert updated.summary.success_rate == 0.0 + assert updated.summary.success_rate_all == 0.0 async def test_native_report_and_quick_calls_use_independent_sessions( diff --git a/tests/unit/analytics/test_outcome_statistics.py b/tests/unit/analytics/test_outcome_statistics.py new file mode 100644 index 0000000000..c6ac98f778 --- /dev/null +++ b/tests/unit/analytics/test_outcome_statistics.py @@ -0,0 +1,99 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +from dataclasses import asdict +from itertools import product + +import pytest + +from pyrit.analytics import combine_outcome_statistics, compute_outcome_statistics +from pyrit.analytics.outcome_statistics import success_percentage +from pyrit.analytics.result_analysis import _compute_stats +from pyrit.models import AttackAnalyticsStatistics, AttackOutcome, AttackStats, OutcomeStatistics + + +@pytest.mark.parametrize(("successes", "failures", "undetermined", "errors"), tuple(product(range(3), repeat=4))) +def test_both_denominators_use_the_same_selected_outcome_counts( + successes: int, failures: int, undetermined: int, errors: int +) -> None: + counts = { + AttackOutcome.SUCCESS: successes, + AttackOutcome.FAILURE: failures, + AttackOutcome.UNDETERMINED: undetermined, + AttackOutcome.ERROR: errors, + } + statistics = compute_outcome_statistics(counts) + decided = successes + failures + total = sum(counts.values()) + assert statistics.total_results == total + assert statistics.total_decided == decided + assert statistics.success_rate == (successes / decided if decided else None) + assert statistics.success_rate_all == (successes / total if total else None) + assert statistics.decided_share == (decided / total if total else None) + assert statistics.outcome_shares == {outcome: count / total if total else 0.0 for outcome, count in counts.items()} + assert isinstance(statistics, OutcomeStatistics) + assert isinstance(statistics, AttackStats) + + +def test_string_and_enum_counts_share_one_model_and_legacy_adapter() -> None: + counts = {"success": 3, "failure": 1, "error": 2} + statistics = compute_outcome_statistics(counts) + assert statistics == compute_outcome_statistics({AttackOutcome(name): value for name, value in counts.items()}) + assert OutcomeStatistics is AttackAnalyticsStatistics + assert statistics.success_rate == 0.75 + assert statistics.success_rate_all == 0.5 + assert statistics.undetermined == 0 + assert asdict(_compute_stats(successes=3, failures=1, errors=2, undetermined=0)) == { + "success_rate": 0.75, + "total_decided": 4, + "successes": 3, + "failures": 1, + "errors": 2, + "undetermined": 0, + } + + +def test_empty_population_rates_are_unavailable_not_zero() -> None: + statistics = compute_outcome_statistics({}) + assert statistics.success_rate is None + assert statistics.success_rate_all is None + assert statistics.decided_share is None + assert statistics.total_results == 0 + assert set(statistics.outcome_shares.values()) == {0.0} + assert combine_outcome_statistics([]) == statistics + + +@pytest.mark.parametrize("counts", [{"unexpected": 1}, {"success": -1}, {"success": True}, {"error": 1.5}]) +def test_invalid_outcome_counts_are_explicit_errors(counts: dict[str, int]) -> None: + with pytest.raises(ValueError, match="unsupported|invalid"): + compute_outcome_statistics(counts) + + +def test_combining_populations_recomputes_rates_instead_of_averaging() -> None: + first = compute_outcome_statistics({"success": 1, "error": 9}) + second = compute_outcome_statistics({"success": 9, "failure": 91}) + combined = combine_outcome_statistics([first, second]) + assert combined.success_rate == pytest.approx(10 / 101) + assert combined.success_rate_all == pytest.approx(10 / 110) + assert combined.success_rate != (first.success_rate + second.success_rate) / 2 + assert combined.success_rate_all != (first.success_rate_all + second.success_rate_all) / 2 + assert combined == compute_outcome_statistics({"success": 10, "failure": 91, "error": 9}) + + +def test_combining_legacy_attack_stats_uses_counts_not_caller_supplied_rates() -> None: + legacy = AttackStats(0.99, 999, 1, 1, 1, 1) + statistics = combine_outcome_statistics([legacy]) + assert statistics.success_rate == 0.5 + assert statistics.success_rate_all == 0.25 + assert statistics.total_decided == 2 + assert statistics.total_results == 4 + + +def test_combining_rejects_invalid_counts_before_they_can_cancel_out() -> None: + with pytest.raises(ValueError, match="invalid"): + combine_outcome_statistics([AttackStats(0, 0, -1, 0, 0, 0), AttackStats(1, 1, 1, 0, 0, 0)]) + + +@pytest.mark.parametrize(("succeeded", "completed", "expected"), [(0, 0, None), (0, 2, 0), (2, 3, 66), (1, 1, 100)]) +def test_legacy_percentages_format_the_shared_rate(succeeded: int, completed: int, expected: int | None) -> None: + assert success_percentage(succeeded=succeeded, completed=completed) == expected diff --git a/tests/unit/analytics/test_scenario_statistics.py b/tests/unit/analytics/test_scenario_statistics.py index a3499ade17..8613b1d759 100644 --- a/tests/unit/analytics/test_scenario_statistics.py +++ b/tests/unit/analytics/test_scenario_statistics.py @@ -4,8 +4,12 @@ import uuid from datetime import UTC, datetime, timedelta +import pytest + +from pyrit.analytics import compute_outcome_statistics from pyrit.analytics.scenario_statistics import ( ScenarioPlanLookup, + combine_execution_counts, compute_scenario_statistics, resolve_execution_unit, ) @@ -14,6 +18,7 @@ SCENARIO_RUN_PLAN_METADATA_KEY, AttackOutcome, AttackResult, + ScenarioProgressCounts, ScenarioRunPlan, ScenarioRunPlanAtomicGroup, ScenarioRunPlanSeedGroup, @@ -79,6 +84,78 @@ def test_empty_result_has_no_success_percentage() -> None: assert statistics.overall.completed == 0 assert statistics.overall.success_percentage is None assert statistics.attempts == 0 + assert statistics.overall.outcomes == compute_outcome_statistics({}) + + +def test_latest_outcomes_share_both_denominators_without_historical_errors() -> None: + result = make_scenario_result( + attack_results={ + "one": [ + _result(objective="A", outcome=AttackOutcome.ERROR), + _result(objective="A", outcome=AttackOutcome.SUCCESS, seconds=1), + _result(objective="B", outcome=AttackOutcome.FAILURE), + ], + "two": [ + _result(objective="C", outcome=AttackOutcome.ERROR), + _result(objective="C", outcome=AttackOutcome.ERROR, seconds=1), + _result(objective="D", outcome=AttackOutcome.UNDETERMINED), + ], + }, + display_group_map={"one": "group", "two": "group"}, + ) + statistics = compute_scenario_statistics(result) + expected = compute_outcome_statistics(dict.fromkeys(AttackOutcome, 1)) + assert statistics.overall.outcomes == expected + assert statistics.display_groups["group"].outcomes == expected + assert combine_execution_counts(statistics.atomic_attacks.values()) == statistics.overall + assert statistics.overall.success_percentage == 25 + assert statistics.overall.outcomes.success_rate == 0.5 + assert statistics.overall.outcomes.success_rate_all == 0.25 + assert statistics.overall.errors == 3 + assert statistics.overall.outcomes.errors == 1 + assert statistics.overall.retries == 2 + + +@pytest.mark.parametrize("outcome", list(AttackOutcome)) +def test_each_latest_outcome_uses_shared_statistics(outcome: AttackOutcome) -> None: + result = make_scenario_result( + attack_results={ + "attack": [ + _result(objective="A", outcome=AttackOutcome.FAILURE), + _result(objective="A", outcome=outcome, seconds=1), + ] + } + ) + counts = compute_scenario_statistics(result).overall + assert counts.outcomes == compute_outcome_statistics({outcome: 1}) + assert counts.completed == 1 + assert counts.success_percentage == (100 if outcome is AttackOutcome.SUCCESS else 0) + + +def test_combining_count_only_legacy_payloads_does_not_infer_failures_from_historical_errors(caplog) -> None: + legacy = ScenarioProgressCounts(completed=2, succeeded=1, errors=5, retries=4, success_percentage=50) + combined = combine_execution_counts([legacy]) + assert combined.outcomes is None + assert combined.success_percentage == 50 + assert combined.errors == 5 + assert "without an outcome breakdown" in caplog.text + + +def test_combining_unequal_scenario_groups_recomputes_both_denominators() -> None: + result = make_scenario_result( + attack_results={ + "one": [_result(objective="A", outcome=AttackOutcome.SUCCESS)], + "two": [ + _result(objective=str(index), outcome=outcome) + for index, outcome in enumerate([AttackOutcome.SUCCESS, AttackOutcome.FAILURE, AttackOutcome.ERROR]) + ], + } + ) + statistics = compute_scenario_statistics(result) + combined = combine_execution_counts(statistics.atomic_attacks.values()) + assert combined.outcomes.success_rate == pytest.approx(2 / 3) + assert combined.outcomes.success_rate_all == 0.5 + assert combined == statistics.overall def test_saved_plan_counts_planned_units_and_reports_unattributed_attempts() -> None: diff --git a/tests/unit/analytics/test_scenario_statistics_parity.py b/tests/unit/analytics/test_scenario_statistics_parity.py index fc8276909b..b2378c44ca 100644 --- a/tests/unit/analytics/test_scenario_statistics_parity.py +++ b/tests/unit/analytics/test_scenario_statistics_parity.py @@ -11,7 +11,7 @@ import json import uuid -from dataclasses import dataclass, field +from dataclasses import asdict, dataclass, field from datetime import UTC, datetime, timedelta import pytest @@ -347,10 +347,12 @@ async def test_sdk_api_and_reports_report_identical_statistics(history_name: str assert list_item.total_retries == detail.total_retries == sdk.overall.retries assert progress.summary.overall.succeeded == sdk.overall.succeeded assert progress.summary.overall.errors == sdk.overall.errors + assert progress.summary.overall.outcomes == sdk.overall.outcomes # Reports report = json.loads(await JsonScenarioResultPrinter().render_async(scenario_result)) assert report["stats"]["overall_success_rate"] == (expected or 0) + assert report["stats"]["outcomes"] == asdict(sdk.overall.outcomes) # Per-group numbers agree between the SDK, the saved-plan progress view, and the reports. Compare # key sets first so a group missing from one view fails instead of reading as 0%. @@ -366,6 +368,8 @@ async def test_sdk_api_and_reports_report_identical_statistics(history_name: str assert report_groups == { name: (completed, rate or 0) for name, (completed, rate) in sdk_groups_with_results.items() } + for group in report["groups"]: + assert group["outcomes"] == asdict(sdk.display_groups[group["name"]].outcomes) if history.plan is not None: progress_groups = { group.display_group: (group.completed, group.success_percentage) @@ -373,6 +377,8 @@ async def test_sdk_api_and_reports_report_identical_statistics(history_name: str } assert set(progress_groups) == set(sdk_groups) assert progress_groups == sdk_groups + for group in progress.summary.display_groups: + assert group.outcomes == sdk.display_groups[group.display_group].outcomes async def test_historical_attempt_counts_stay_separate_from_units(sqlite_instance) -> None: diff --git a/tests/unit/common/test_lazy_package_imports.py b/tests/unit/common/test_lazy_package_imports.py index c05c810305..f8acb26af7 100644 --- a/tests/unit/common/test_lazy_package_imports.py +++ b/tests/unit/common/test_lazy_package_imports.py @@ -22,6 +22,18 @@ ) _LAZY_IMPORT_SPOT_CHECKS = [ + ( + "pyrit.analytics", + "compute_outcome_statistics", + "pyrit.analytics.outcome_statistics", + "pyrit.memory.attack_analytics", + ), + ( + "pyrit.models", + "OutcomeStatistics", + "pyrit.models.analytics", + "pyrit.analytics.outcome_statistics", + ), ( "pyrit.analytics", "AttackResultAnalytics", @@ -411,11 +423,13 @@ def test_analytics_foundations_do_not_load_higher_layers() -> None: name for name, module in pyrit.models._LAZY_EXPORTS.items() if module == "pyrit.models.analytics" ] - assert len(names) == 21 + assert len(names) == 22 + assert "OutcomeStatistics" in names for name in names: exported = getattr(pyrit.models, name) assert exported is getattr(importlib.import_module("pyrit.models.analytics"), name) assert pyrit.models.__dict__[name] is exported + assert pyrit.models.OutcomeStatistics is pyrit.models.AttackAnalyticsStatistics forbidden = ( "pyrit.analytics", "pyrit.backend", "pyrit.memory", "pyrit.executor", diff --git a/tests/unit/models/test_analytics.py b/tests/unit/models/test_analytics.py index 8efec7dff8..ce09faf9a1 100644 --- a/tests/unit/models/test_analytics.py +++ b/tests/unit/models/test_analytics.py @@ -10,6 +10,7 @@ from pydantic import ValidationError from pyrit.analytics import AttackStats as AnalyticsAttackStats +from pyrit.analytics import compute_outcome_statistics from pyrit.analytics.result_analysis import AttackStats as ResultAnalysisAttackStats from pyrit.common.pagination import fingerprint_filters from pyrit.models import ( @@ -28,6 +29,8 @@ AttackOutcome, AttackResultSelection, AttackStats, + OutcomeStatistics, + ScenarioProgressCounts, ) @@ -46,6 +49,31 @@ def test_attack_stats_preserves_existing_import_and_constructor() -> None: assert stats == AttackStats(**expected) +def test_shared_outcome_statistics_preserves_existing_attack_type_identity() -> None: + assert OutcomeStatistics is AttackAnalyticsStatistics + statistics = compute_outcome_statistics({"success": 1, "failure": 1, "error": 2}) + assert statistics.success_rate == 0.5 + assert statistics.success_rate_all == 0.25 + counts = ScenarioProgressCounts( + completed=4, succeeded=1, errors=5, retries=3, success_percentage=25, outcomes=statistics + ) + assert counts.model_dump(mode="json")["outcomes"] == asdict(statistics) + restored = ScenarioProgressCounts.model_validate_json(counts.model_dump_json()) + assert restored == counts + assert isinstance(restored.outcomes, OutcomeStatistics) + assert restored.outcomes.errors == 2 + assert restored.errors == 5 + + +def test_legacy_scenario_counts_do_not_fabricate_an_outcome_breakdown() -> None: + counts = ScenarioProgressCounts.model_validate( + {"completed": 2, "succeeded": 1, "success_percentage": 50, "errors": 3, "retries": 2} + ) + assert counts.outcomes is None + assert counts.errors == 3 + assert counts.completed == 2 + + @pytest.mark.parametrize( "data", [ diff --git a/tests/unit/output/scenario_result/test_json.py b/tests/unit/output/scenario_result/test_json.py index 03ce43285b..a0db103d58 100644 --- a/tests/unit/output/scenario_result/test_json.py +++ b/tests/unit/output/scenario_result/test_json.py @@ -227,5 +227,14 @@ async def test_overview_separates_units_from_attempts(printer): assert payload["stats"]["total_attempts"] == 2 assert payload["stats"]["overall_success_rate"] == 100 assert payload["groups"] == [ - {"name": "technique_a", "num_objective_executions": 1, "num_attempts": 2, "success_rate": 100} + { + "name": "technique_a", + "num_objective_executions": 1, + "num_attempts": 2, + "success_rate": 100, + "outcomes": payload["stats"]["outcomes"], + } ] + assert payload["stats"]["outcomes"]["success_rate"] == 1.0 + assert payload["stats"]["outcomes"]["success_rate_all"] == 1.0 + assert payload["stats"]["outcomes"]["errors"] == 0 diff --git a/tests/unit/output/test_derivation.py b/tests/unit/output/test_derivation.py index 77835dca03..93acc0a4cb 100644 --- a/tests/unit/output/test_derivation.py +++ b/tests/unit/output/test_derivation.py @@ -6,6 +6,7 @@ from unit.mocks import make_scenario_result +from pyrit.analytics import compute_outcome_statistics from pyrit.common.utils import to_sha256 from pyrit.models import ( SCENARIO_RUN_PLAN_METADATA_KEY, @@ -71,7 +72,12 @@ def test_scenario_overview_empty_is_zero(): overview = scenario_overview(result) assert (overview.objective_executions, overview.attempts, overview.success_rate) == (0, 0, 0) - assert overview.groups == [GroupStatistics(name="s1", objective_executions=0, attempts=0, success_rate=0)] + assert overview.groups == [ + GroupStatistics( + name="s1", objective_executions=0, attempts=0, success_rate=0, outcomes=compute_outcome_statistics({}) + ) + ] + assert overview.outcomes == compute_outcome_statistics({}) def test_scenario_overview_folds_atomic_attacks_by_display_group(): @@ -90,7 +96,15 @@ def test_scenario_overview_folds_atomic_attacks_by_display_group(): overview = scenario_overview(result) assert overview.success_rate == 66 - assert overview.groups == [GroupStatistics(name="encoding", objective_executions=3, attempts=3, success_rate=66)] + assert overview.groups == [ + GroupStatistics( + name="encoding", + objective_executions=3, + attempts=3, + success_rate=66, + outcomes=compute_outcome_statistics({"success": 2, "failure": 1}), + ) + ] def test_scenario_overview_uses_display_group_map_even_when_plan_labels_differ(): @@ -126,7 +140,15 @@ def test_scenario_overview_uses_display_group_map_even_when_plan_labels_differ() overview = scenario_overview(result) - assert overview.groups == [GroupStatistics(name="encoding", objective_executions=1, attempts=1, success_rate=100)] + assert overview.groups == [ + GroupStatistics( + name="encoding", + objective_executions=1, + attempts=1, + success_rate=100, + outcomes=compute_outcome_statistics({"success": 1}), + ) + ] # --- attack_score_display --- From 2dc59dadaba4e1b58d43573459593e5a4c194711 Mon Sep 17 00:00:00 2001 From: Roman Lutz Date: Fri, 9 Oct 2026 07:18:45 -0700 Subject: [PATCH 4/8] FIX: Validate outcome statistics and improve maintained analytics APIs Reject contradictory scenario totals and outcome breakdowns, including mutated inputs at aggregation boundaries. Expose an explicit read-only success_rate_decided alias in shared statistics and JSON. Add typed include_outcome_statistics opt-ins to maintained result and async technique analytics while preserving the six-field default and existing deprecation schedules. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- doc/code/analytics/0_attack_results.md | 55 ++++++- doc/code/framework.md | 3 +- pyrit/analytics/__init__.py | 2 + pyrit/analytics/outcome_statistics.py | 4 +- pyrit/analytics/result_analysis.py | 101 +++++++------ pyrit/analytics/scenario_statistics.py | 2 +- pyrit/analytics/technique_analysis.py | 84 +++++++---- pyrit/models/analytics.py | 134 ++++++++++++++++-- pyrit/models/scenario_progress.py | 32 ++++- pyrit/output/scenario_result/json.py | 13 +- .../unit/analytics/test_outcome_statistics.py | 15 ++ tests/unit/analytics/test_result_analysis.py | 39 +++++ .../analytics/test_scenario_statistics.py | 29 ++++ .../test_scenario_statistics_parity.py | 10 +- .../unit/analytics/test_technique_analysis.py | 78 +++++++++- .../unit/common/test_lazy_package_imports.py | 6 + tests/unit/models/test_analytics.py | 95 ++++++++++++- 17 files changed, 600 insertions(+), 102 deletions(-) diff --git a/doc/code/analytics/0_attack_results.md b/doc/code/analytics/0_attack_results.md index ffcf7d5861..d46e2c87a2 100644 --- a/doc/code/analytics/0_attack_results.md +++ b/doc/code/analytics/0_attack_results.md @@ -18,12 +18,15 @@ policies together: | Field | Denominator | Meaning | |---|---|---| -| `success_rate` | Successes + failures | Success among decided outcomes, the existing ASR default. | +| `success_rate_decided` | Successes + failures | Success among decided outcomes, the existing ASR default. | | `success_rate_all` | All four outcomes | Success among every selected result, including errors and undetermined outcomes. | +`success_rate` remains the compatible stored field; `success_rate_decided` is its +read-only alias, not an independently mutable value. Shared-model JSON includes +both names and accepts either when validating a payload; conflicting aliases fail. An empty denominator produces `None`. For example, one success and one error -produce `success_rate=1.0` and `success_rate_all=0.5`. An error-only population has -`success_rate=None` and `success_rate_all=0.0`. Neither rate requires a second query. +produce `success_rate_decided=1.0` and `success_rate_all=0.5`. An error-only population +has `success_rate_decided=None` and `success_rate_all=0.0`. Neither rate requires a second query. `total_results` includes all four outcomes; `outcome_shares` and `decided_share` use that whole selected cohort. Empty outcome shares are zero, and an empty cohort's decided share is `None`. All proportions are between 0 and 1. @@ -74,7 +77,7 @@ async with AttackResultAnalytics() as analytics: compare_by=AttackAnalyticsDimension(name="attack_type"), ) ) - print(report.summary.total_results, report.summary.success_rate, report.summary.success_rate_all) + print(report.summary.total_results, report.summary.success_rate_decided, report.summary.success_rate_all) if report.drilldown_unavailable_reason: print(report.drilldown_unavailable_reason) @@ -111,7 +114,7 @@ from pyrit.analytics import compute_scenario_statistics statistics = compute_scenario_statistics(scenario_result) outcomes = statistics.overall.outcomes assert outcomes is not None # Always populated by compute_scenario_statistics. -print(outcomes.success_rate, outcomes.success_rate_all) +print(outcomes.success_rate_decided, outcomes.success_rate_all) ``` `ScenarioProgressCounts.success_percentage` retains its existing all-completed-unit @@ -127,6 +130,13 @@ from older callers deserialize with `outcomes=None` because their latest-outcome breakdown cannot be reconstructed. Combining a nonempty such payload logs a warning and leaves the combined breakdown unavailable rather than guessing failures. +Shared outcome statistics validate supplied counts, totals, rates, and shares at +construction. Scenario progress totals must agree with the nested outcome counts. +Combining statistics also revalidates existing objects, including nested values +modified after construction, instead of silently repairing contradictions. +Objects remain mutable for compatibility: make a new result through the shared +calculator when changing counts, rather than editing cached totals or rates. + For already counted populations, call `compute_outcome_statistics` directly. `combine_outcome_statistics` sums counts from **disjoint** populations and recalculates both rates. It also accepts existing `AttackStats` returned by @@ -141,12 +151,45 @@ from pyrit.analytics import combine_outcome_statistics, compute_outcome_statisti first = compute_outcome_statistics({"success": 1, "error": 1}) second = compute_outcome_statistics({"success": 2, "failure": 1}) combined = combine_outcome_statistics([first, second]) -print(combined.success_rate, combined.success_rate_all) # 0.75, 0.6 +print(combined.success_rate_decided, combined.success_rate_all) # 0.75, 0.6 ``` Do not combine overlapping converter/harm groups to reconstruct a report's summary; use its already-computed summary instead. +### Maintained APIs and deprecated wrappers + +`analyze_results` and `compute_technique_stats_async` remain supported. Their default +results keep the six-field `AttackStats` dataclass and its existing constructor, +`asdict`, and decided-only rate. Pass `include_outcome_statistics=True` to obtain +the shared `OutcomeStatistics` directly, without an extra conversion or query: + +```python +from pyrit.analytics import analyze_results, compute_technique_stats_async + +analysis = analyze_results(attack_results, include_outcome_statistics=True) +print(analysis["Overall"].success_rate_decided, analysis["Overall"].success_rate_all) + +techniques = await compute_technique_stats_async( + technique_eval_hashes=technique_hashes, + include_outcome_statistics=True, +) +for technique_hash, outcomes in techniques.items(): + print(technique_hash, outcomes.success_rate_decided, outcomes.success_rate_all) +``` + +The opt-in changes returned statistics, not grouping, result selection, or retry +policy. Default `AttackStats` also exposes the read-only `success_rate_decided` +property, without adding dataclass fields. + +| API | Status | +|---|---| +| `analyze_results`, `compute_technique_stats_async`, `get_cached_results_for_technique_async` | Maintained. | +| Synchronous `compute_technique_stats` and `get_cached_results_for_technique` | Deprecated; scheduled for removal in 1.4.0. Use the async equivalents. | +| `ScenarioResult.objective_achieved_rate` | Deprecated; scheduled for removal in 1.4.0. Use `compute_scenario_statistics`. | + +No additional API is deprecated by these changes. + ### Filters and drill-downs Omit `compare_by` for a one-dimensional breakdown. Supported dimensions include diff --git a/doc/code/framework.md b/doc/code/framework.md index ba2fdb3fa9..0022f76cdd 100644 --- a/doc/code/framework.md +++ b/doc/code/framework.md @@ -337,7 +337,8 @@ The below talks about responsibilities of most modules in the PyRIT library - `compute_scenario_statistics` owns execution-unit identity (atomic attack, technique configuration, and seed group), latest-attempt selection, and historical retry/error counts. It passes selected outcomes to the shared outcome calculator. SDK callers, the GUI backend's run detail and progress views, and the console, JSON, and HTML reports present its results (`ScenarioExecutionStatistics`, `ScenarioExecutionUnit`, and `ScenarioProgressCounts` in `pyrit.models`). The GUI run-history list aggregates counts in SQL (`MemoryInterface._build_scenario_history_aggregate_statement`) so it can page over many runs, then uses the shared percentage helper; `tests/unit/analytics/test_scenario_statistics_parity.py` keeps the paths in agreement. - Scenario attempts are ordered by timestamp, then their canonical lowercase UUID string. SQL Server uses this string order rather than its native UUID order for history ranking and progress pagination. Explicit seed attribution wins; an objective alone matches a planned seed group only when that match is unique. Legacy runs with identifier-only seed identities are recounted with the shared analytics. - Shared analytics contracts (filters, dimensions, typed values, reports, facets, result pages, and `AttackStats`) live in `pyrit.models.analytics`. They validate data without querying memory or calculating statistics. `AttackResultSelection` defines selection modes without changing existing callers. -- `compute_outcome_statistics` owns outcome-count validation, both success-rate denominators, totals, and shares. `success_rate` divides by successes plus failures; `success_rate_all` divides by all outcomes. `combine_outcome_statistics` combines disjoint counts, never averages rates. Attack and scenario analytics share the same `OutcomeStatistics` model and calculations after independently selecting their populations. Legacy `AttackStats` and scenario percentage fields retain their existing shapes/defaults. +- `compute_outcome_statistics` owns both success-rate denominators, totals, and shares. `success_rate_decided` (the explicit alias of `success_rate`) divides by successes plus failures; `success_rate_all` divides by all outcomes. `combine_outcome_statistics` revalidates and combines disjoint counts, never averages rates. Attack and scenario analytics share the same `OutcomeStatistics` model and calculations after independently selecting their populations. Models reject contradictory supplied totals/rates; they do not select populations or replace the supplied numbers. +- `analyze_results` and `compute_technique_stats_async` remain maintained APIs. Their `include_outcome_statistics=True` option exposes the shared statistics directly; default `AttackStats` results retain their six-field shape. Existing scenario percentage fields also retain their defaults. Only the already-deprecated sync memory wrappers and `ScenarioResult.objective_achieved_rate` are scheduled for removal in 1.4.0. - [`AttackResultAnalytics`](./analytics/0_attack_results.md) provides async saved-result reports, lightweight result pages, and facet lookups. It supplies both shared success rates, display labels and exact additional drill-down predicates, and counts every saved result ID rather than latest scenario execution units. Memory owns cohort SQL and consistent projections; the SDK owns their interpretation. - Analytics facades sharing a memory instance share one loop-bound, bounded report/quick-query controller. Close that shared SDK lifetime before replacing or disposing memory; cancellation does not free a running query's slot before its actual session cleanup. SDK lifecycle is independent of backend runtime integration. - Filter-bound cursor and label-normalization helpers live in `pyrit.common.pagination`. The backend pagination module retains compatibility exports, including History's invalid-cursor first-page fallback. diff --git a/pyrit/analytics/__init__.py b/pyrit/analytics/__init__.py index 4f0ced28a1..a0413932b3 100644 --- a/pyrit/analytics/__init__.py +++ b/pyrit/analytics/__init__.py @@ -19,6 +19,7 @@ get_cached_results_for_technique_async, ) from pyrit.analytics.scenario_statistics import compute_scenario_statistics + from pyrit.analytics.technique_analysis import compute_technique_stats_async from pyrit.analytics.text_matching import ApproximateTextMatching, ExactTextMatching, TextMatching _LAZY_EXPORTS: dict[str, str | tuple[str, str | None]] = { @@ -29,6 +30,7 @@ "combine_outcome_statistics": "pyrit.analytics.outcome_statistics", "compute_outcome_statistics": "pyrit.analytics.outcome_statistics", "compute_scenario_statistics": "pyrit.analytics.scenario_statistics", + "compute_technique_stats_async": "pyrit.analytics.technique_analysis", "ConversationAnalytics": "pyrit.analytics.conversation_analytics", "ExactTextMatching": "pyrit.analytics.text_matching", "get_cached_results_for_technique": "pyrit.analytics.result_analysis", diff --git a/pyrit/analytics/outcome_statistics.py b/pyrit/analytics/outcome_statistics.py index 1642044fcf..ce45b0f1ee 100644 --- a/pyrit/analytics/outcome_statistics.py +++ b/pyrit/analytics/outcome_statistics.py @@ -21,7 +21,7 @@ def compute_outcome_statistics(counts: Mapping[str, int] | Mapping[AttackOutcome Calculate both success rates for an already selected population. This function neither queries memory nor chooses result IDs, retries, roles, - or scenario units. ``success_rate`` uses success + failure; ``success_rate_all`` + or scenario units. ``success_rate_decided`` (also ``success_rate``) uses success + failure; ``success_rate_all`` uses every outcome. Both are proportions, not rounded percentages. Args: @@ -77,6 +77,8 @@ def combine_outcome_statistics(statistics: Iterable[AttackStats]) -> OutcomeStat """ counts: Counter[AttackOutcome] = Counter() for item in statistics: + if isinstance(item, OutcomeStatistics): + item.validate_consistency() counts.update( _validated_counts( { diff --git a/pyrit/analytics/result_analysis.py b/pyrit/analytics/result_analysis.py index 574160abb7..5356e579a9 100644 --- a/pyrit/analytics/result_analysis.py +++ b/pyrit/analytics/result_analysis.py @@ -3,7 +3,7 @@ from collections import defaultdict from collections.abc import Sequence -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Generic, Literal, TypedDict, TypeVar, overload from pyrit.analytics.outcome_statistics import compute_outcome_statistics from pyrit.common.deprecation import print_deprecation_message @@ -14,12 +14,21 @@ IdentifierFilter, IdentifierType, ObjectiveTargetEvaluationIdentifier, + OutcomeStatistics, ) if TYPE_CHECKING: from pyrit.memory.memory_interface import MemoryInterface _SYNC_API_REMOVAL_VERSION = "1.4.0" +_StatisticsT = TypeVar("_StatisticsT", bound=AttackStats) + + +class _AnalysisResult(TypedDict, Generic[_StatisticsT]): + """The existing result dictionary, with precise types for its two keys.""" + + Overall: _StatisticsT + By_attack_identifier: dict[str, _StatisticsT] def _compute_stats(successes: int, failures: int, undetermined: int, errors: int) -> AttackStats: @@ -31,24 +40,55 @@ def _compute_stats(successes: int, failures: int, undetermined: int, errors: int AttackOutcome.ERROR: errors, } ) + return _as_attack_stats(statistics) + + +def _as_attack_stats(statistics: OutcomeStatistics) -> AttackStats: return AttackStats( success_rate=statistics.success_rate, total_decided=statistics.total_decided, - successes=successes, - failures=failures, - undetermined=undetermined, - errors=errors, + successes=statistics.successes, + failures=statistics.failures, + undetermined=statistics.undetermined, + errors=statistics.errors, ) -def analyze_results(attack_results: list[AttackResult]) -> dict[str, AttackStats | dict[str, AttackStats]]: +@overload +def analyze_results( + attack_results: list[AttackResult], *, include_outcome_statistics: Literal[True] +) -> _AnalysisResult[OutcomeStatistics]: ... + + +@overload +def analyze_results( + attack_results: list[AttackResult], *, include_outcome_statistics: Literal[False] = False +) -> _AnalysisResult[AttackStats]: ... + + +@overload +def analyze_results( + attack_results: list[AttackResult], *, include_outcome_statistics: bool +) -> _AnalysisResult[AttackStats] | _AnalysisResult[OutcomeStatistics]: ... + + +def analyze_results( + attack_results: list[AttackResult], *, include_outcome_statistics: bool = False +) -> _AnalysisResult[AttackStats] | _AnalysisResult[OutcomeStatistics]: """ Analyze a list of AttackResult objects and return overall and grouped statistics. + This API remains supported. Both output shapes use the shared outcome calculator + and count the supplied results without changing their grouping or selecting retries. + + Args: + attack_results (list[AttackResult]): Results to count. + include_outcome_statistics (bool): Return ``OutcomeStatistics`` with both success + rates, totals, and shares. False preserves the six-field ``AttackStats`` shape. + Returns: - A dictionary of AttackStats objects. The overall stats are accessible with the key - "Overall", and the stats of any attack can be retrieved using "By_attack_identifier" - followed by the identifier of the attack. + dict: Overall statistics under "Overall" and per-attack-type statistics under + "By_attack_identifier". The opt-in changes values, not dictionary keys. Raises: ValueError: if attack_results is empty. @@ -64,8 +104,8 @@ def analyze_results(attack_results: list[AttackResult]) -> dict[str, AttackStats if not attack_results: raise ValueError("attack_results cannot be empty") - overall_counts: defaultdict[str, int] = defaultdict(int) - by_type_counts: defaultdict[str, defaultdict[str, int]] = defaultdict(lambda: defaultdict(int)) + overall_counts: defaultdict[AttackOutcome, int] = defaultdict(int) + by_type_counts: defaultdict[str, defaultdict[AttackOutcome, int]] = defaultdict(lambda: defaultdict(int)) for attack in attack_results: if not isinstance(attack, AttackResult): @@ -75,39 +115,16 @@ def analyze_results(attack_results: list[AttackResult]) -> dict[str, AttackStats _strategy_id = attack.get_attack_strategy_identifier() attack_type = _strategy_id.class_name if _strategy_id is not None else "unknown" - if outcome == AttackOutcome.SUCCESS: - overall_counts["successes"] += 1 - by_type_counts[attack_type]["successes"] += 1 - elif outcome == AttackOutcome.FAILURE: - overall_counts["failures"] += 1 - by_type_counts[attack_type]["failures"] += 1 - elif outcome == AttackOutcome.ERROR: - overall_counts["errors"] += 1 - by_type_counts[attack_type]["errors"] += 1 - else: - overall_counts["undetermined"] += 1 - by_type_counts[attack_type]["undetermined"] += 1 - - overall_stats = _compute_stats( - successes=overall_counts["successes"], - failures=overall_counts["failures"], - undetermined=overall_counts["undetermined"], - errors=overall_counts["errors"], - ) - - by_type_stats = { - attack_type: _compute_stats( - successes=counts["successes"], - failures=counts["failures"], - undetermined=counts["undetermined"], - errors=counts["errors"], - ) - for attack_type, counts in by_type_counts.items() - } + overall_counts[outcome] += 1 + by_type_counts[attack_type][outcome] += 1 + overall_stats = compute_outcome_statistics(overall_counts) + by_type_stats = {attack_type: compute_outcome_statistics(counts) for attack_type, counts in by_type_counts.items()} + if include_outcome_statistics: + return {"Overall": overall_stats, "By_attack_identifier": by_type_stats} return { - "Overall": overall_stats, - "By_attack_identifier": by_type_stats, + "Overall": _as_attack_stats(overall_stats), + "By_attack_identifier": {name: _as_attack_stats(statistics) for name, statistics in by_type_stats.items()}, } diff --git a/pyrit/analytics/scenario_statistics.py b/pyrit/analytics/scenario_statistics.py index 3c2f4a3195..65305bab8c 100644 --- a/pyrit/analytics/scenario_statistics.py +++ b/pyrit/analytics/scenario_statistics.py @@ -310,7 +310,7 @@ def combine_execution_counts(counts: Iterable[ScenarioProgressCounts]) -> Scenar None unless every input has one. ``outcomes`` is None if a nonempty legacy input lacks its breakdown; historical errors cannot reconstruct latest outcomes. """ - counts = list(counts) + counts = [ScenarioProgressCounts.model_validate(item) for item in counts] completed = sum(item.completed for item in counts) succeeded = sum(item.succeeded for item in counts) planned = [item.planned for item in counts] diff --git a/pyrit/analytics/technique_analysis.py b/pyrit/analytics/technique_analysis.py index c0ebf7b1ac..5be4d7f02e 100644 --- a/pyrit/analytics/technique_analysis.py +++ b/pyrit/analytics/technique_analysis.py @@ -5,17 +5,20 @@ from __future__ import annotations -from typing import TYPE_CHECKING +from collections import Counter +from typing import TYPE_CHECKING, Literal, overload -from pyrit.analytics.result_analysis import AttackStats, _compute_stats +from pyrit.analytics.outcome_statistics import compute_outcome_statistics +from pyrit.analytics.result_analysis import _as_attack_stats from pyrit.common.deprecation import print_deprecation_message from pyrit.memory import CentralMemory -from pyrit.models import AttackOutcome, AttackResult +from pyrit.models import AttackStats as AttackStats # noqa: TC001 - preserve the existing module export if TYPE_CHECKING: from collections.abc import Sequence from pyrit.memory.memory_interface import MemoryInterface + from pyrit.models import AttackOutcome, AttackResult, OutcomeStatistics def compute_technique_stats( @@ -66,16 +69,51 @@ def compute_technique_stats( targeted_harm_categories=targeted_harm_categories, ) - return _aggregate_technique_stats(results=results, technique_eval_hashes=technique_eval_hashes) + statistics = _aggregate_technique_stats(results=results, technique_eval_hashes=technique_eval_hashes) + return {key: _as_attack_stats(value) for key, value in statistics.items()} +@overload async def compute_technique_stats_async( *, technique_eval_hashes: Sequence[str], scenario_result_id: str | None = None, targeted_harm_categories: Sequence[str] | None = None, memory: MemoryInterface | None = None, -) -> dict[str, AttackStats]: + include_outcome_statistics: Literal[True], +) -> dict[str, OutcomeStatistics]: ... + + +@overload +async def compute_technique_stats_async( + *, + technique_eval_hashes: Sequence[str], + scenario_result_id: str | None = None, + targeted_harm_categories: Sequence[str] | None = None, + memory: MemoryInterface | None = None, + include_outcome_statistics: Literal[False] = False, +) -> dict[str, AttackStats]: ... + + +@overload +async def compute_technique_stats_async( + *, + technique_eval_hashes: Sequence[str], + scenario_result_id: str | None = None, + targeted_harm_categories: Sequence[str] | None = None, + memory: MemoryInterface | None = None, + include_outcome_statistics: bool, +) -> dict[str, AttackStats] | dict[str, OutcomeStatistics]: ... + + +async def compute_technique_stats_async( + *, + technique_eval_hashes: Sequence[str], + scenario_result_id: str | None = None, + targeted_harm_categories: Sequence[str] | None = None, + memory: MemoryInterface | None = None, + include_outcome_statistics: bool = False, +) -> dict[str, AttackStats] | dict[str, OutcomeStatistics]: """ Compute per-technique outcome statistics from persisted attack results. @@ -87,6 +125,9 @@ async def compute_technique_stats_async( for behavioral-equivalence aggregation (seeds excluded, scorer excluded, only behavior-relevant target params included). + This async API remains supported. The output opt-in does not change its memory + query, cohort selection, or behavior used by adaptive technique selectors. + Args: technique_eval_hashes (Sequence[str]): Eval hashes to aggregate. Returned dict is keyed by these. @@ -96,10 +137,12 @@ async def compute_technique_stats_async( whose attack targeted these harm categories. Defaults to ``None``. memory (MemoryInterface | None): Memory backend to query. Defaults to ``CentralMemory.get_memory_instance()``. + include_outcome_statistics (bool): Return shared statistics with both success + rates instead of the compatible six-field ``AttackStats`` shape. Defaults to False. Returns: - dict[str, AttackStats]: Stats per technique eval hash. Hashes with no - historical results are omitted from the result. + dict[str, AttackStats] | dict[str, OutcomeStatistics]: Stats per technique eval + hash, including both rates when requested. Hashes with no results are omitted. """ if not technique_eval_hashes: return {} @@ -112,32 +155,23 @@ async def compute_technique_stats_async( targeted_harm_categories=targeted_harm_categories, ) - return _aggregate_technique_stats(results=results, technique_eval_hashes=technique_eval_hashes) + statistics = _aggregate_technique_stats(results=results, technique_eval_hashes=technique_eval_hashes) + if include_outcome_statistics: + return statistics + return {key: _as_attack_stats(value) for key, value in statistics.items()} def _aggregate_technique_stats( *, results: Sequence[AttackResult], technique_eval_hashes: Sequence[str] -) -> dict[str, AttackStats]: +) -> dict[str, OutcomeStatistics]: requested = set(technique_eval_hashes) - counts: dict[str, tuple[int, int, int, int]] = {} + counts: dict[str, Counter[AttackOutcome]] = {} for result in results: identifier = result.atomic_attack_identifier eval_hash = identifier.eval_hash if identifier is not None else None if eval_hash is None or eval_hash not in requested: continue - successes, failures, undetermined, errors = counts.get(eval_hash, (0, 0, 0, 0)) - if result.outcome == AttackOutcome.SUCCESS: - successes += 1 - elif result.outcome == AttackOutcome.FAILURE: - failures += 1 - elif result.outcome == AttackOutcome.ERROR: - errors += 1 - else: - undetermined += 1 - counts[eval_hash] = (successes, failures, undetermined, errors) - - return { - eval_hash: _compute_stats(successes=successes, failures=failures, undetermined=undetermined, errors=errors) - for eval_hash, (successes, failures, undetermined, errors) in counts.items() - } + counts.setdefault(eval_hash, Counter())[result.outcome] += 1 + + return {eval_hash: compute_outcome_statistics(outcomes) for eval_hash, outcomes in counts.items()} diff --git a/pyrit/models/analytics.py b/pyrit/models/analytics.py index 1670e6376a..55b1cc0c05 100644 --- a/pyrit/models/analytics.py +++ b/pyrit/models/analytics.py @@ -8,16 +8,21 @@ from dataclasses import dataclass, field from datetime import UTC, datetime from enum import Enum -from typing import ClassVar, Self +from math import isclose, isfinite +from typing import Annotated, Any, ClassVar, Self -from pydantic import AwareDatetime, BaseModel, ConfigDict, Field, field_validator, model_validator +from pydantic import AwareDatetime, BaseModel, ConfigDict, Field, computed_field, field_validator, model_validator +from pydantic.dataclasses import dataclass as validated_dataclass from pyrit.models.results.attack_result import AttackOutcome +_OutcomeCount = Annotated[int, Field(strict=True, ge=0)] +_OutcomeRate = Annotated[float, Field(strict=True, ge=0, le=1, allow_inf_nan=False)] + @dataclass class AttackStats: - """Statistics for attack analysis results.""" + """Compatible six-field statistics for attack analysis results.""" success_rate: float | None total_decided: int @@ -26,14 +31,20 @@ class AttackStats: undetermined: int errors: int + @property + def success_rate_decided(self) -> float | None: + """The decided-only rate, exposed without changing the legacy dataclass fields.""" + return self.success_rate -@dataclass + +@validated_dataclass(config=ConfigDict(revalidate_instances="always", extra="forbid")) class AttackAnalyticsStatistics(AttackStats): """ Shared outcome statistics for saved results or selected scenario execution units. - ``success_rate`` uses successes / decided results; errors and undetermined - outcomes do not enter that denominator. ``success_rate_all`` instead includes + ``success_rate_decided`` (also available as ``success_rate``) uses successes / + decided results; errors and undetermined outcomes do not enter that denominator. + ``success_rate_all`` instead includes every outcome in its denominator, as do ``decided_share`` and ``outcome_shares``. Rates with no applicable denominator are ``None``; shares for an empty cohort are zero. All proportions are in the range 0 through 1. @@ -41,12 +52,115 @@ class AttackAnalyticsStatistics(AttackStats): Analytics calculates these values after selecting the counted population. ``total_results`` can therefore describe distinct saved IDs or latest scenario units. The model does not select attempts or infer a counting policy. + Construction and explicit revalidation reject contradictory supplied values. + The nullable all-outcome rate preserves older payloads that omit that field. """ - total_results: int - decided_share: float | None - outcome_shares: dict[AttackOutcome, float] - success_rate_all: float | None = field(default=None, kw_only=True) + success_rate: _OutcomeRate | None + total_decided: _OutcomeCount + successes: _OutcomeCount + failures: _OutcomeCount + undetermined: _OutcomeCount + errors: _OutcomeCount + total_results: _OutcomeCount + decided_share: _OutcomeRate | None + outcome_shares: dict[AttackOutcome, _OutcomeRate] + success_rate_all: _OutcomeRate | None = field(default=None, kw_only=True) + + def __post_init__(self) -> None: + """Validate the supplied statistics without recalculating or replacing fields.""" + self.validate_consistency() + + @computed_field # type: ignore[prop-decorator] + @property + def success_rate_decided(self) -> float | None: + """The explicit decided-only rate; a read-only alias of ``success_rate``.""" + return self.success_rate + + def validate_consistency(self) -> None: + """ + Check supplied counts, totals, rates, and shares without replacing any of them. + + This also validates mutable instances at analytics input boundaries. + Analytics remains responsible for calculating statistics from counts. + + Raises: + ValueError: If a count is invalid or a derived value contradicts its counts. + """ + counts = { + AttackOutcome.SUCCESS: self.successes, + AttackOutcome.FAILURE: self.failures, + AttackOutcome.UNDETERMINED: self.undetermined, + AttackOutcome.ERROR: self.errors, + } + if any( + type(count) is not int or count < 0 for count in (*counts.values(), self.total_decided, self.total_results) + ): + raise ValueError("Outcome counts and totals must be nonnegative integers.") + if self.total_decided != self.successes + self.failures: + raise ValueError("total_decided must equal successes plus failures.") + if self.total_results != sum(counts.values()): + raise ValueError("total_results must equal the sum of all outcome counts.") + self._validate_rate( + name="success_rate", value=self.success_rate, numerator=self.successes, denominator=self.total_decided + ) + self._validate_rate( + name="decided_share", value=self.decided_share, numerator=self.total_decided, denominator=self.total_results + ) + if self.success_rate_all is not None: + self._validate_rate( + name="success_rate_all", + value=self.success_rate_all, + numerator=self.successes, + denominator=self.total_results, + ) + if set(self.outcome_shares) != set(AttackOutcome): + raise ValueError("outcome_shares must contain each supported outcome exactly once.") + for outcome, count in counts.items(): + self._validate_rate( + name=f"outcome_shares.{outcome.value}", + value=self.outcome_shares[outcome], + numerator=count, + denominator=self.total_results, + empty=0.0, + ) + + @model_validator(mode="before") + @classmethod + def _normalize_success_rate_alias(cls, value: Any) -> Any: + """ + Accept either rate name in validated payloads, rejecting disagreeing aliases. + + Returns: + Any: The payload with one stored decided-rate value. + + Raises: + ValueError: If both rate names are supplied with different values. + """ + if isinstance(value, dict) and "success_rate_decided" in value: + value = dict(value) + rate = value.pop("success_rate_decided") + if "success_rate" in value and value["success_rate"] != rate: + raise ValueError("success_rate and success_rate_decided must agree.") + value["success_rate"] = rate + return value + + @staticmethod + def _validate_rate( + *, name: str, value: float | None, numerator: int, denominator: int, empty: float | None = None + ) -> None: + if denominator == 0 and value is None and empty is None: + return + if ( + value is None + or isinstance(value, bool) + or not isinstance(value, (int, float)) + or not isfinite(value) + or not 0 <= value <= 1 + or (denominator == 0 and value != empty) + or (denominator > 0 and not isclose(value, numerator / denominator, rel_tol=1e-12, abs_tol=0.0)) + ): + raise ValueError(f"{name} must agree with its outcome counts.") # Preserve the existing class/constructor identity while sharing it beyond attack reports. diff --git a/pyrit/models/scenario_progress.py b/pyrit/models/scenario_progress.py index 8bad575b10..97e05c0e24 100644 --- a/pyrit/models/scenario_progress.py +++ b/pyrit/models/scenario_progress.py @@ -5,7 +5,7 @@ from datetime import datetime from enum import Enum -from typing import Any, Literal +from typing import Any, Literal, Self from pydantic import AwareDatetime, BaseModel, ConfigDict, Field, model_validator @@ -191,8 +191,12 @@ class ScenarioProgressCounts(BaseModel): decided-only and all-outcome success rates. ``errors`` and ``retries`` retain their historical-attempt meaning and must not be used as outcome denominators. None preserves older count-only payloads whose outcome breakdown is unknown. + Supplied totals must agree with the outcome breakdown; mutable values are + revalidated when they are passed back into analytics. """ + model_config = ConfigDict(revalidate_instances="always") + completed: int = Field(..., ge=0) planned: int | None = Field(default=None, ge=0) succeeded: int = Field(..., ge=0) @@ -201,6 +205,32 @@ class ScenarioProgressCounts(BaseModel): retries: int = Field(..., ge=0) outcomes: OutcomeStatistics | None = None + @model_validator(mode="after") + def _validate_statistics(self) -> Self: + """ + Reject contradictory progress totals instead of presenting two success rates for different counts. + + Returns: + Self: The unchanged consistent counts, including legacy count-only payloads. + + Raises: + ValueError: If successes exceed completed units, a supplied percentage is inconsistent, + or the nested latest-outcome counts disagree with progress totals. + """ + if self.succeeded > self.completed: + raise ValueError("succeeded must not exceed completed.") + if self.success_percentage is not None: + expected = int((self.succeeded / self.completed) * 100) if self.completed else 0 + if self.success_percentage != expected: + raise ValueError("success_percentage must agree with succeeded and completed.") + if self.outcomes is not None: + self.outcomes.validate_consistency() + if self.completed != self.outcomes.total_results: + raise ValueError("completed must equal outcomes.total_results.") + if self.succeeded != self.outcomes.successes: + raise ValueError("succeeded must equal outcomes.successes.") + return self + class ScenarioExecutionUnit(BaseModel): """ diff --git a/pyrit/output/scenario_result/json.py b/pyrit/output/scenario_result/json.py index 948d82f5e4..cbd25f2ae3 100644 --- a/pyrit/output/scenario_result/json.py +++ b/pyrit/output/scenario_result/json.py @@ -2,10 +2,11 @@ # Licensed under the MIT license. import json -from dataclasses import asdict -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, ClassVar -from pyrit.models import AttackResult, ScenarioResult +from pydantic import TypeAdapter + +from pyrit.models import AttackResult, OutcomeStatistics, ScenarioResult from pyrit.output._derivation import ( attack_score_display, resolve_target_info, @@ -34,6 +35,8 @@ class JsonScenarioResultPrinter(ScenarioResultPrinterBase): supplies a memory-backed one). """ + _OUTCOME_ADAPTER: ClassVar[TypeAdapter[OutcomeStatistics | None]] = TypeAdapter(OutcomeStatistics | None) + def __init__( self, *, @@ -130,7 +133,7 @@ def _build_overview(self, result: ScenarioResult) -> dict[str, Any]: "num_objective_executions": group.objective_executions, "num_attempts": group.attempts, "success_rate": group.success_rate, - "outcomes": asdict(group.outcomes) if group.outcomes is not None else None, + "outcomes": self._OUTCOME_ADAPTER.dump_python(group.outcomes, mode="json"), } for group in overview.groups ] @@ -163,7 +166,7 @@ def _build_overview(self, result: ScenarioResult) -> dict[str, Any]: "total_objective_executions": overview.objective_executions, "total_attempts": overview.attempts, "overall_success_rate": overview.success_rate, - "outcomes": asdict(overview.outcomes) if overview.outcomes is not None else None, + "outcomes": self._OUTCOME_ADAPTER.dump_python(overview.outcomes, mode="json"), "unique_objectives": len(result.get_objectives()), }, "groups": groups, diff --git a/tests/unit/analytics/test_outcome_statistics.py b/tests/unit/analytics/test_outcome_statistics.py index c6ac98f778..956af564a4 100644 --- a/tests/unit/analytics/test_outcome_statistics.py +++ b/tests/unit/analytics/test_outcome_statistics.py @@ -63,6 +63,14 @@ def test_empty_population_rates_are_unavailable_not_zero() -> None: assert combine_outcome_statistics([]) == statistics +def test_large_integer_counts_validate_without_float_overflow() -> None: + count = 10**400 + statistics = compute_outcome_statistics({"success": count, "failure": count, "error": 2 * count}) + assert statistics.total_results == 4 * count + assert statistics.success_rate_decided == 0.5 + assert statistics.success_rate_all == 0.25 + + @pytest.mark.parametrize("counts", [{"unexpected": 1}, {"success": -1}, {"success": True}, {"error": 1.5}]) def test_invalid_outcome_counts_are_explicit_errors(counts: dict[str, int]) -> None: with pytest.raises(ValueError, match="unsupported|invalid"): @@ -94,6 +102,13 @@ def test_combining_rejects_invalid_counts_before_they_can_cancel_out() -> None: combine_outcome_statistics([AttackStats(0, 0, -1, 0, 0, 0), AttackStats(1, 1, 1, 0, 0, 0)]) +def test_combining_rejects_mutated_rich_statistics_instead_of_silently_repairing_them() -> None: + statistics = compute_outcome_statistics({"success": 1, "failure": 1}) + statistics.successes = 10 + with pytest.raises(ValueError): + combine_outcome_statistics([statistics]) + + @pytest.mark.parametrize(("succeeded", "completed", "expected"), [(0, 0, None), (0, 2, 0), (2, 3, 66), (1, 1, 100)]) def test_legacy_percentages_format_the_shared_rate(succeeded: int, completed: int, expected: int | None) -> None: assert success_percentage(succeeded=succeeded, completed=completed) == expected diff --git a/tests/unit/analytics/test_result_analysis.py b/tests/unit/analytics/test_result_analysis.py index a699272a15..3fc2ae7d36 100644 --- a/tests/unit/analytics/test_result_analysis.py +++ b/tests/unit/analytics/test_result_analysis.py @@ -1,11 +1,14 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT license. +import warnings +from dataclasses import asdict from datetime import UTC, datetime, timedelta from unittest.mock import AsyncMock, MagicMock import pytest +from pyrit.analytics import compute_outcome_statistics from pyrit.analytics.result_analysis import ( AttackStats, _objective_target_eval_hash_for, @@ -21,6 +24,7 @@ IdentifierFilter, IdentifierType, ObjectiveTargetEvaluationIdentifier, + OutcomeStatistics, ) @@ -56,6 +60,41 @@ def test_analyze_results_raises_on_invalid_object(): analyze_results(["not-an-AttackResult"]) +def test_analyze_results_opt_in_returns_shared_statistics_without_changing_default_shape() -> None: + attacks = [ + make_attack(AttackOutcome.SUCCESS, attack_type="one"), + make_attack(AttackOutcome.ERROR, attack_type="one"), + make_attack(AttackOutcome.FAILURE, attack_type="two"), + make_attack(AttackOutcome.UNDETERMINED, attack_type="two"), + ] + legacy = analyze_results(attacks) + rich = analyze_results(attacks, include_outcome_statistics=True) + assert type(legacy["Overall"]) is AttackStats + assert len(asdict(legacy["Overall"])) == 6 + assert isinstance(rich["Overall"], OutcomeStatistics) + assert rich["Overall"] == compute_outcome_statistics(dict.fromkeys(AttackOutcome, 1)) + assert rich["Overall"].success_rate_decided == 0.5 + assert rich["Overall"].success_rate_all == 0.25 + assert rich["By_attack_identifier"]["one"] == compute_outcome_statistics({"success": 1, "error": 1}) + assert rich["By_attack_identifier"]["two"] == compute_outcome_statistics({"failure": 1, "undetermined": 1}) + assert set(legacy) == set(rich) + + +@pytest.mark.parametrize("include_outcome_statistics", [False, True]) +def test_analyze_results_remains_supported_without_deprecation(include_outcome_statistics: bool) -> None: + with warnings.catch_warnings(): + warnings.simplefilter("error", DeprecationWarning) + analyze_results([make_attack(AttackOutcome.SUCCESS)], include_outcome_statistics=include_outcome_statistics) + + +@pytest.mark.parametrize("include_outcome_statistics", [False, True]) +def test_analyze_results_keeps_input_validation_with_both_output_shapes(include_outcome_statistics: bool) -> None: + with pytest.raises(ValueError, match="empty"): + analyze_results([], include_outcome_statistics=include_outcome_statistics) + with pytest.raises(TypeError, match="AttackResult"): + analyze_results(["invalid"], include_outcome_statistics=include_outcome_statistics) + + @pytest.mark.parametrize( "outcomes, expected_successes, expected_failures, expected_undetermined, expected_errors, expected_rate", [ diff --git a/tests/unit/analytics/test_scenario_statistics.py b/tests/unit/analytics/test_scenario_statistics.py index 8613b1d759..0d22ecb6f2 100644 --- a/tests/unit/analytics/test_scenario_statistics.py +++ b/tests/unit/analytics/test_scenario_statistics.py @@ -158,6 +158,35 @@ def test_combining_unequal_scenario_groups_recomputes_both_denominators() -> Non assert combined == statistics.overall +@pytest.mark.parametrize("field", ["completed", "succeeded", "success_percentage"]) +def test_combining_revalidates_mutated_scenario_totals(field: str) -> None: + counts = ScenarioProgressCounts( + completed=2, + succeeded=1, + errors=1, + retries=0, + success_percentage=50, + outcomes=compute_outcome_statistics({"success": 1, "error": 1}), + ) + setattr(counts, field, 10) + with pytest.raises(ValueError, match="completed|succeeded|success_percentage"): + combine_execution_counts([counts]) + + +@pytest.mark.parametrize("field", ["successes", "total_results", "success_rate_all"]) +def test_combining_revalidates_mutated_nested_outcomes(field: str) -> None: + counts = ScenarioProgressCounts( + completed=2, + succeeded=1, + errors=1, + retries=0, + outcomes=compute_outcome_statistics({"success": 1, "error": 1}), + ) + setattr(counts.outcomes, field, 10) + with pytest.raises(ValueError): + combine_execution_counts([counts]) + + def test_saved_plan_counts_planned_units_and_reports_unattributed_attempts() -> None: result = make_scenario_result( attack_results={ diff --git a/tests/unit/analytics/test_scenario_statistics_parity.py b/tests/unit/analytics/test_scenario_statistics_parity.py index b2378c44ca..0a02bc6da3 100644 --- a/tests/unit/analytics/test_scenario_statistics_parity.py +++ b/tests/unit/analytics/test_scenario_statistics_parity.py @@ -11,10 +11,11 @@ import json import uuid -from dataclasses import asdict, dataclass, field +from dataclasses import dataclass, field from datetime import UTC, datetime, timedelta import pytest +from pydantic import TypeAdapter from pyrit.analytics import compute_scenario_statistics from pyrit.backend.services.scenario_run_service import ScenarioRunService @@ -27,6 +28,7 @@ AttackResult, AttackSeedGroup, ComponentIdentifier, + OutcomeStatistics, ScenarioRunPlan, ScenarioRunPlanAtomicGroup, ScenarioRunPlanSeedGroup, @@ -352,7 +354,7 @@ async def test_sdk_api_and_reports_report_identical_statistics(history_name: str # Reports report = json.loads(await JsonScenarioResultPrinter().render_async(scenario_result)) assert report["stats"]["overall_success_rate"] == (expected or 0) - assert report["stats"]["outcomes"] == asdict(sdk.overall.outcomes) + assert report["stats"]["outcomes"] == TypeAdapter(OutcomeStatistics).dump_python(sdk.overall.outcomes, mode="json") # Per-group numbers agree between the SDK, the saved-plan progress view, and the reports. Compare # key sets first so a group missing from one view fails instead of reading as 0%. @@ -369,7 +371,9 @@ async def test_sdk_api_and_reports_report_identical_statistics(history_name: str name: (completed, rate or 0) for name, (completed, rate) in sdk_groups_with_results.items() } for group in report["groups"]: - assert group["outcomes"] == asdict(sdk.display_groups[group["name"]].outcomes) + assert group["outcomes"] == TypeAdapter(OutcomeStatistics).dump_python( + sdk.display_groups[group["name"]].outcomes, mode="json" + ) if history.plan is not None: progress_groups = { group.display_group: (group.completed, group.success_percentage) diff --git a/tests/unit/analytics/test_technique_analysis.py b/tests/unit/analytics/test_technique_analysis.py index f02b8a9347..6e1ce5efbc 100644 --- a/tests/unit/analytics/test_technique_analysis.py +++ b/tests/unit/analytics/test_technique_analysis.py @@ -1,13 +1,17 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT license. +import warnings +from dataclasses import asdict from unittest.mock import AsyncMock, MagicMock, patch import pytest -from pyrit.analytics.technique_analysis import compute_technique_stats_async -from pyrit.memory import MemoryInterface -from pyrit.models import AttackOutcome +from pyrit.analytics import compute_outcome_statistics +from pyrit.analytics.technique_analysis import compute_technique_stats, compute_technique_stats_async +from pyrit.memory import MemoryInterface, SQLiteMemory +from pyrit.models import AttackOutcome, AttackStats, OutcomeStatistics +from unit.memory.test_attack_analytics import make_result def _make_result(*, eval_hash: str | None, outcome: AttackOutcome) -> MagicMock: @@ -32,6 +36,74 @@ def _patch_memory(): class TestComputeTechniqueStats: + async def test_explicit_outcome_statistics_keep_existing_queries_and_default_shape(self, _patch_memory) -> None: + _patch_memory.get_attack_results_async.return_value = [ + _make_result(eval_hash="a", outcome=AttackOutcome.SUCCESS), + _make_result(eval_hash="a", outcome=AttackOutcome.ERROR), + _make_result(eval_hash="b", outcome=AttackOutcome.UNDETERMINED), + ] + legacy = await compute_technique_stats_async(technique_eval_hashes=["a", "b"]) + rich = await compute_technique_stats_async( + technique_eval_hashes=["a", "b"], + scenario_result_id="run", + targeted_harm_categories=["privacy"], + include_outcome_statistics=True, + ) + assert type(legacy["a"]) is AttackStats + assert len(asdict(legacy["a"])) == 6 + assert isinstance(rich["a"], OutcomeStatistics) + assert rich["a"] == compute_outcome_statistics({"success": 1, "error": 1}) + assert rich["a"].success_rate_decided == 1.0 + assert rich["a"].success_rate_all == 0.5 + assert rich["b"].success_rate_decided is None + assert rich["b"].success_rate_all == 0.0 + assert _patch_memory.get_attack_results_async.call_count == 2 + assert _patch_memory.get_attack_results_async.call_args.kwargs == { + "atomic_attack_eval_hashes": ["a", "b"], + "scenario_result_id": "run", + "targeted_harm_categories": ["privacy"], + } + + @pytest.mark.parametrize("include_outcome_statistics", [False, True]) + async def test_maintained_async_api_does_not_emit_deprecation( + self, _patch_memory, include_outcome_statistics: bool + ) -> None: + _patch_memory.get_attack_results_async.return_value = [ + _make_result(eval_hash="a", outcome=AttackOutcome.SUCCESS) + ] + with warnings.catch_warnings(): + warnings.simplefilter("error", DeprecationWarning) + await compute_technique_stats_async( + technique_eval_hashes=["a"], include_outcome_statistics=include_outcome_statistics + ) + + def test_sync_wrapper_keeps_existing_deprecation_and_result_shape(self, _patch_memory) -> None: + _patch_memory.get_attack_results.return_value = [_make_result(eval_hash="a", outcome=AttackOutcome.SUCCESS)] + with pytest.warns(DeprecationWarning, match=r"compute_technique_stats.*1\.4\.0"): + statistics = compute_technique_stats(technique_eval_hashes=["a"]) + assert type(statistics["a"]) is AttackStats + assert len(asdict(statistics["a"])) == 6 + assert statistics["a"].success_rate_decided == 1.0 + + async def test_rich_statistics_from_real_saved_results(self, sqlite_instance: SQLiteMemory) -> None: + first = make_result() + second = make_result(index=2, outcome=AttackOutcome.ERROR) + first.conversation_id = "first" + second.conversation_id = "second" + await sqlite_instance.add_attack_results_to_memory_async(attack_results=[first, second]) + eval_hash = first.atomic_attack_identifier.eval_hash + statistics = await compute_technique_stats_async( + memory=sqlite_instance, technique_eval_hashes=[eval_hash], include_outcome_statistics=True + ) + assert statistics[eval_hash] == compute_outcome_statistics({"success": 1, "error": 1}) + assert statistics[eval_hash].success_rate_decided == 1.0 + assert statistics[eval_hash].success_rate_all == 0.5 + + @pytest.mark.parametrize("hashes", [[], ["absent"]]) + async def test_rich_statistics_preserve_empty_result_semantics(self, _patch_memory, hashes: list[str]) -> None: + assert await compute_technique_stats_async(technique_eval_hashes=hashes, include_outcome_statistics=True) == {} + assert _patch_memory.get_attack_results_async.call_count == (1 if hashes else 0) + async def test_empty_results_returns_empty(self, _patch_memory): stats = await compute_technique_stats_async(technique_eval_hashes=["a", "b"]) assert stats == {} diff --git a/tests/unit/common/test_lazy_package_imports.py b/tests/unit/common/test_lazy_package_imports.py index f65b97dfda..f745e7b406 100644 --- a/tests/unit/common/test_lazy_package_imports.py +++ b/tests/unit/common/test_lazy_package_imports.py @@ -22,6 +22,12 @@ ) _LAZY_IMPORT_SPOT_CHECKS = [ + ( + "pyrit.analytics", + "compute_technique_stats_async", + "pyrit.analytics.technique_analysis", + "pyrit.analytics.attack_result_analytics", + ), ( "pyrit.analytics", "compute_outcome_statistics", diff --git a/tests/unit/models/test_analytics.py b/tests/unit/models/test_analytics.py index ce09faf9a1..df4bebfccb 100644 --- a/tests/unit/models/test_analytics.py +++ b/tests/unit/models/test_analytics.py @@ -2,12 +2,12 @@ # Licensed under the MIT license. import json -from dataclasses import asdict +from dataclasses import asdict, replace from datetime import UTC, datetime from zoneinfo import ZoneInfo import pytest -from pydantic import ValidationError +from pydantic import TypeAdapter, ValidationError from pyrit.analytics import AttackStats as AnalyticsAttackStats from pyrit.analytics import compute_outcome_statistics @@ -47,6 +47,7 @@ def test_attack_stats_preserves_existing_import_and_constructor() -> None: } assert asdict(stats) == expected assert stats == AttackStats(**expected) + assert stats.success_rate_decided == stats.success_rate def test_shared_outcome_statistics_preserves_existing_attack_type_identity() -> None: @@ -57,7 +58,10 @@ def test_shared_outcome_statistics_preserves_existing_attack_type_identity() -> counts = ScenarioProgressCounts( completed=4, succeeded=1, errors=5, retries=3, success_percentage=25, outcomes=statistics ) - assert counts.model_dump(mode="json")["outcomes"] == asdict(statistics) + assert counts.model_dump(mode="json")["outcomes"] == { + **asdict(statistics), + "success_rate_decided": statistics.success_rate, + } restored = ScenarioProgressCounts.model_validate_json(counts.model_dump_json()) assert restored == counts assert isinstance(restored.outcomes, OutcomeStatistics) @@ -74,6 +78,86 @@ def test_legacy_scenario_counts_do_not_fabricate_an_outcome_breakdown() -> None: assert counts.completed == 2 +@pytest.mark.parametrize( + "changes", + [ + {"completed": 10}, + {"succeeded": 3}, + {"success_percentage": 100}, + ], +) +@pytest.mark.parametrize("input_mode", ["object", "dict", "json"]) +def test_scenario_counts_reject_contradictory_outcomes(changes: dict[str, int], input_mode: str) -> None: + statistics = compute_outcome_statistics({"success": 1, "failure": 1, "error": 2}) + fields = { + "completed": 4, + "succeeded": 1, + "errors": 2, + "retries": 0, + "success_percentage": 25, + "outcomes": statistics if input_mode == "object" else asdict(statistics), + **changes, + } + with pytest.raises(ValidationError, match="completed|succeeded|success_percentage"): + if input_mode == "json": + ScenarioProgressCounts.model_validate_json(json.dumps(fields)) + else: + ScenarioProgressCounts.model_validate(fields) + + +@pytest.mark.parametrize( + "changes", + [ + {"successes": -1}, + {"failures": True}, + {"undetermined": 1.5}, + {"total_decided": 99}, + {"total_results": 99}, + {"success_rate": None}, + {"success_rate": 0.25}, + {"success_rate_all": 0.5}, + {"decided_share": 1.0}, + {"outcome_shares": {AttackOutcome.SUCCESS: 1.0}}, + {"success_rate": float("nan")}, + {"success_rate_all": float("inf")}, + ], +) +def test_shared_statistics_reject_inconsistent_counts_and_rates(changes: dict[str, object]) -> None: + statistics = compute_outcome_statistics({"success": 1, "failure": 1, "error": 2}) + with pytest.raises(ValueError): + replace(statistics, **changes) + with pytest.raises(ValidationError): + TypeAdapter(OutcomeStatistics).validate_python({**asdict(statistics), **changes}) + + +def test_reembedding_mutated_outcomes_revalidates_the_counts() -> None: + statistics = compute_outcome_statistics({"success": 1}) + statistics.successes = 5 + with pytest.raises(ValidationError): + ScenarioProgressCounts(completed=1, succeeded=1, errors=0, retries=0, outcomes=statistics) + + +def test_shared_statistics_serialize_and_accept_the_explicit_decided_alias() -> None: + statistics = compute_outcome_statistics({"success": 1, "failure": 1, "error": 2}) + adapter = TypeAdapter(OutcomeStatistics) + payload = adapter.dump_python(statistics, mode="json") + assert statistics.success_rate_decided == statistics.success_rate == 0.5 + assert payload["success_rate_decided"] == payload["success_rate"] == 0.5 + assert adapter.validate_json(adapter.dump_json(statistics)) == statistics + payload.pop("success_rate") + assert adapter.validate_python(payload) == statistics + assert "success_rate_decided" in adapter.json_schema(mode="serialization")["properties"] + with pytest.raises(ValidationError, match="success_rate"): + adapter.validate_python({**payload, "success_rate": 0.75}) + + +def test_explicit_decided_alias_is_not_independently_mutable() -> None: + statistics = compute_outcome_statistics({"success": 1, "failure": 1}) + with pytest.raises(AttributeError): + statistics.success_rate_decided = 1.0 + assert statistics.success_rate_decided == statistics.success_rate == 0.5 + + @pytest.mark.parametrize( "data", [ @@ -436,7 +520,10 @@ def test_report_round_trips_statistics_and_drilldown_availability(reason: str | computed_at=computed_at, ) assert isinstance(report.summary, AttackStats) - assert report.model_dump(mode="json")["summary"] == asdict(statistics) + assert report.model_dump(mode="json")["summary"] == { + **asdict(statistics), + "success_rate_decided": statistics.success_rate, + } assert report.drilldown_unavailable_reason == reason assert AttackAnalyticsReport.model_validate_json(report.model_dump_json()) == report From 763fec7fa2f559890b3f589b13b1a9c9e1debc0e Mon Sep 17 00:00:00 2001 From: Nimit Jain Date: Sat, 10 Oct 2026 01:24:19 +0530 Subject: [PATCH 5/8] feat: Add role-aware Scenario target-attempt accounting (#3043) --- pyrit/analytics/scenario_statistics.py | 75 +++++++++ .../services/scenario_progress_read_model.py | 9 +- .../backend/services/scenario_run_service.py | 32 +++- pyrit/memory/memory_interface.py | 143 ++++++++++++------ pyrit/models/__init__.py | 4 + pyrit/models/catalog/__init__.py | 4 + pyrit/models/catalog/scenario.py | 18 +++ pyrit/models/scenario_progress.py | 8 +- .../test_scenario_progress_read_model.py | 14 +- .../unit/backend/test_scenario_run_service.py | 36 +++++ ...st_interface_seed_dataset_summary_query.py | 14 ++ 11 files changed, 296 insertions(+), 61 deletions(-) diff --git a/pyrit/analytics/scenario_statistics.py b/pyrit/analytics/scenario_statistics.py index 2b483ca6f1..bb3ad8b055 100644 --- a/pyrit/analytics/scenario_statistics.py +++ b/pyrit/analytics/scenario_statistics.py @@ -32,6 +32,8 @@ AtomicAttackIdentifier, AttackOutcome, AttackResult, + AttackResultRole, + AttackResultMetadata, ComponentIdentifier, ScenarioExecutionStatistics, ScenarioExecutionUnit, @@ -58,6 +60,9 @@ def outcome(self) -> AttackOutcome: ... @property def total_retries(self) -> int: ... + @property + def result_role(self) -> AttackResultRole | None: ... + @dataclass(frozen=True, slots=True) class ScenarioPlanLookup: @@ -134,6 +139,7 @@ class ScenarioAttempt: timestamp: datetime attempt_id: str total_retries: int + result_role: AttackResultRole | None def load_scenario_run_plan(scenario_result: ScenarioResult) -> ScenarioRunPlan | None: @@ -240,6 +246,7 @@ def resolve_attack_result_attempt( timestamp=_timestamp_order_key(attack_result.timestamp), attempt_id=str(attack_result.attack_result_id), total_retries=retries if isinstance(retries, int) else 0, + result_role=AttackResultMetadata.from_metadata(metadata=attribution_data).result_role, ) @@ -285,6 +292,16 @@ def count_execution_units( succeeded = 0 errors = 0 retries = 0 + target_facing_attempts = 0 + target_facing_errors = 0 + target_facing_retries = 0 + orchestration_attempts = 0 + orchestration_errors = 0 + orchestration_retries = 0 + unknown_attempts = 0 + unknown_errors = 0 + unknown_retries = 0 + for unit in units: attempts = attempts_by_unit.get(unit, ()) if not attempts: @@ -296,6 +313,42 @@ def count_execution_units( attempts_per_unit=[len(attempts)], persisted_retries=[attempt.total_retries for attempt in attempts], ) + + for attempt in attempts: + role = attempt.result_role + if role == AttackResultRole.TARGET_FACING: + target_facing_attempts += 1 + target_facing_errors += int(attempt.outcome == AttackOutcome.ERROR) + target_facing_retries += max(attempt.total_retries, 0) + elif role == AttackResultRole.ORCHESTRATION: + orchestration_attempts += 1 + orchestration_errors += int(attempt.outcome == AttackOutcome.ERROR) + orchestration_retries += max(attempt.total_retries, 0) + else: + unknown_attempts += 1 + unknown_errors += int(attempt.outcome == AttackOutcome.ERROR) + unknown_retries += max(attempt.total_retries, 0) + + from pyrit.models import ScenarioProducerCounts, ScenarioProducerCategoryCounts + + producer_counts = ScenarioProducerCounts( + target_facing=ScenarioProducerCategoryCounts( + attempts=target_facing_attempts, + errors=target_facing_errors, + retries=target_facing_retries, + ), + orchestration=ScenarioProducerCategoryCounts( + attempts=orchestration_attempts, + errors=orchestration_errors, + retries=orchestration_retries, + ), + unknown=ScenarioProducerCategoryCounts( + attempts=unknown_attempts, + errors=unknown_errors, + retries=unknown_retries, + ), + ) + return ScenarioProgressCounts( completed=completed, planned=planned, @@ -303,6 +356,7 @@ def count_execution_units( success_percentage=success_percentage(succeeded=succeeded, completed=completed), errors=errors, retries=retries, + producer_counts=producer_counts, ) @@ -318,6 +372,26 @@ def combine_execution_counts(counts: Iterable[ScenarioProgressCounts]) -> Scenar completed = sum(item.completed for item in counts) succeeded = sum(item.succeeded for item in counts) planned = [item.planned for item in counts] + + from pyrit.models import ScenarioProducerCounts, ScenarioProducerCategoryCounts + producer_counts = ScenarioProducerCounts( + target_facing=ScenarioProducerCategoryCounts( + attempts=sum(item.producer_counts.target_facing.attempts for item in counts), + errors=sum(item.producer_counts.target_facing.errors for item in counts), + retries=sum(item.producer_counts.target_facing.retries for item in counts), + ), + orchestration=ScenarioProducerCategoryCounts( + attempts=sum(item.producer_counts.orchestration.attempts for item in counts), + errors=sum(item.producer_counts.orchestration.errors for item in counts), + retries=sum(item.producer_counts.orchestration.retries for item in counts), + ), + unknown=ScenarioProducerCategoryCounts( + attempts=sum(item.producer_counts.unknown.attempts for item in counts), + errors=sum(item.producer_counts.unknown.errors for item in counts), + retries=sum(item.producer_counts.unknown.retries for item in counts), + ), + ) + return ScenarioProgressCounts( completed=completed, planned=sum(value for value in planned if value is not None) if all(v is not None for v in planned) else None, @@ -325,6 +399,7 @@ def combine_execution_counts(counts: Iterable[ScenarioProgressCounts]) -> Scenar success_percentage=success_percentage(succeeded=succeeded, completed=completed), errors=sum(item.errors for item in counts), retries=sum(item.retries for item in counts), + producer_counts=producer_counts, ) diff --git a/pyrit/backend/services/scenario_progress_read_model.py b/pyrit/backend/services/scenario_progress_read_model.py index da2bdc9460..53ba860ca6 100644 --- a/pyrit/backend/services/scenario_progress_read_model.py +++ b/pyrit/backend/services/scenario_progress_read_model.py @@ -267,19 +267,16 @@ def calculate_progress_counts( *, scenario_result: ScenarioResult, plan: ScenarioRunPlan | None, - ) -> tuple[int, int, int, int]: + ) -> ScenarioProgressCounts: """ Calculate planned-unit totals without inflating retries or error attempts. Delegates to ``pyrit.analytics.scenario_statistics`` so run details match the SDK and reports. Returns: - tuple[int, int, int, int]: Total, completed, success-rate percentage, - and successful-unit count. + ScenarioProgressCounts: The calculated progress counts. """ - overall = compute_scenario_statistics(scenario_result, plan=plan, use_saved_plan=False).overall - total = overall.planned if overall.planned is not None else overall.completed - return total, overall.completed, overall.success_percentage or 0, overall.succeeded + return compute_scenario_statistics(scenario_result, plan=plan, use_saved_plan=False).overall @staticmethod def total_retry_pressure(*, attempts_per_unit: Iterable[int], persisted_retries: Iterable[int]) -> int: diff --git a/pyrit/backend/services/scenario_run_service.py b/pyrit/backend/services/scenario_run_service.py index aa0339e32d..00ff46dbbd 100644 --- a/pyrit/backend/services/scenario_run_service.py +++ b/pyrit/backend/services/scenario_run_service.py @@ -1430,12 +1430,14 @@ def _build_run_summary( plan = None # Build result fields from DB (always computed so in-progress runs show progress) - total_attacks, completed_attacks, objective_achieved_rate, successful_attacks = ( - self._progress_read_model.calculate_progress_counts( - scenario_result=scenario_result, - plan=plan, - ) + overall_counts = self._progress_read_model.calculate_progress_counts( + scenario_result=scenario_result, + plan=plan, ) + total_attacks = overall_counts.planned if overall_counts.planned is not None else overall_counts.completed + completed_attacks = overall_counts.completed + objective_achieved_rate = overall_counts.success_percentage or 0 + successful_attacks = overall_counts.succeeded techniques_used = ( list(dict.fromkeys(group.technique_name or group.display_group for group in plan.atomic_groups)) if plan is not None @@ -1529,6 +1531,7 @@ def _build_run_summary( queue_position=queue_position, active_scenario_result_id=active_scenario_result_id, overload_summaries=self._build_overload_summaries(retry_events=overload_events), + producer_counts=overall_counts.producer_counts, ) @staticmethod @@ -1681,6 +1684,24 @@ def _build_history_summary( if atomic_groups is not None else list(aggregate.atomic_attack_names) ) + from pyrit.models import ScenarioProducerCounts, ScenarioProducerCategoryCounts + producer_counts = ScenarioProducerCounts( + target_facing=ScenarioProducerCategoryCounts( + attempts=aggregate.target_facing_attempts, + errors=aggregate.target_facing_error_attempts, + retries=aggregate.target_facing_retries, + ), + orchestration=ScenarioProducerCategoryCounts( + attempts=aggregate.orchestration_attempts, + errors=aggregate.orchestration_error_attempts, + retries=aggregate.orchestration_retries, + ), + unknown=ScenarioProducerCategoryCounts( + attempts=aggregate.unknown_role_attempts, + errors=aggregate.unknown_role_error_attempts, + retries=aggregate.unknown_role_retries, + ), + ) return ScenarioRunListItem( scenario_result_id=record.scenario_result_id, scenario_name=record.scenario_name, @@ -1707,6 +1728,7 @@ def _build_history_summary( successful_attacks=successful, error_attacks=aggregate.error_attempts, attack_details_available=False, + producer_counts=producer_counts, ) @staticmethod diff --git a/pyrit/memory/memory_interface.py b/pyrit/memory/memory_interface.py index 63b0c0ad0b..b78a1a81af 100644 --- a/pyrit/memory/memory_interface.py +++ b/pyrit/memory/memory_interface.py @@ -47,6 +47,7 @@ from pyrit.common.async_compatibility import legacy_sync_override, run_legacy_sync_async from pyrit.common.deprecation import print_deprecation_message from pyrit.common.pagination import DecodedKeysetCursor +from pyrit.memory.analytics_sql import JsonScalar if TYPE_CHECKING: from pyrit.memory.memory_embedding import MemoryEmbedding @@ -89,6 +90,7 @@ AttackIdentifier, AttackOutcome, AttackResult, + AttackResultRole, AttackResultSelection, AttackTechniqueIdentifier, ComponentIdentifier, @@ -263,6 +265,18 @@ class ScenarioHistoryAggregate: # group from those, so callers should count the run with pyrit.analytics.compute_scenario_statistics instead. needs_sdk_statistics: bool = False + target_facing_attempts: int + target_facing_error_attempts: int + target_facing_retries: int + + orchestration_attempts: int + orchestration_error_attempts: int + orchestration_retries: int + + unknown_role_attempts: int + unknown_role_error_attempts: int + unknown_role_retries: int + @classmethod def empty(cls, *, scenario_result_id: str) -> "ScenarioHistoryAggregate": """ @@ -280,6 +294,15 @@ def empty(cls, *, scenario_result_id: str) -> "ScenarioHistoryAggregate": total_retries=0, latest_attempt_timestamp=None, atomic_attack_names=(), + target_facing_attempts=0, + target_facing_error_attempts=0, + target_facing_retries=0, + orchestration_attempts=0, + orchestration_error_attempts=0, + orchestration_retries=0, + unknown_role_attempts=0, + unknown_role_error_attempts=0, + unknown_role_retries=0, ) @@ -4337,26 +4360,26 @@ def _execute_get_seed_dataset_summaries(self) -> Sequence[SeedDatasetSummary]: combined_statement = named_statement.union_all(unnamed_statement) with closing(self._get_session()) as session: - rows = session.execute(combined_statement).all() + rows = session.execute(combined_statement).mappings().all() summaries_by_dataset: dict[str | None, dict[str, Any]] = {} dataset_order: list[str | None] = [] for row in rows: - dataset_name = row.dataset_name + dataset_name = row["dataset_name"] if dataset_name not in summaries_by_dataset: summaries_by_dataset[dataset_name] = { - "seed_pieces": int(row.seed_pieces or 0), - "logical_examples": int(row.logical_examples or 0), - "objectives": int(row.objectives or 0), + "seed_pieces": int(row["seed_pieces"] or 0), + "logical_examples": int(row["logical_examples"] or 0), + "objectives": int(row["objectives"] or 0), "modalities": set(), "harm_categories": set(), "has_unlabeled_harm_categories": False, } dataset_order.append(dataset_name) summary = summaries_by_dataset[dataset_name] - if row.data_type: - summary["modalities"].add(row.data_type) - categories = row.harm_categories or [] + if row["data_type"]: + summary["modalities"].add(row["data_type"]) + categories = row["harm_categories"] or [] if categories: summary["harm_categories"].update(categories) else: @@ -4512,14 +4535,14 @@ def _execute_get_seed_examples( with closing(self._get_session()) as session: total = session.execute(select(func.count()).select_from(grouped)).scalar_one() - rows = session.execute(page).all() + rows = session.execute(page).mappings().all() seeds = self._get_seed_example_seeds( - session=session, scope=scope, example_ids=[row.example_id for row in rows[:limit]] + session=session, scope=scope, example_ids=[row["example_id"] for row in rows[:limit]] ) next_after = None if len(rows) > limit: last = rows[limit - 1] - next_after = DecodedKeysetCursor(timestamp=last.first_added, identifier=str(last.example_id)) + next_after = DecodedKeysetCursor(timestamp=last["first_added"], identifier=str(last["example_id"])) return seeds, total, next_after def _execute_get_seed_example(self, *, dataset_name: str | None, example_id: uuid.UUID) -> list[SeedRecord]: @@ -5011,7 +5034,7 @@ def _apply_attack_fields_in_session( entry.atomic_attack_identifier_hash = identifier.hash value = identifier.model_dump() if field == "attack_metadata": - value = {**(entry.attack_metadata or {}), **value} + value = {**(entry.attack_metadata or {}), **cast(dict[str, Any], value)} setattr(entry, field, value) def _execute_promote_attack_conversation(self, *, attack_result_id: str, conversation_id: str) -> bool: @@ -5971,7 +5994,7 @@ def _execute_get_scenario_run_state_page( raise ValueError("Scenario run state projection limit must be between 1 and 500.") conditions = [ScenarioResultEntry.scenario_run_state.in_([state.value for state in states])] if after_id is not None: - conditions.append(ScenarioResultEntry.id > uuid.UUID(after_id)) + conditions.append(ScenarioResultEntry.id > cast(Any, uuid.UUID(after_id))) statement = ( select(ScenarioResultEntry.id, ScenarioResultEntry.scenario_run_state) .where(and_(*conditions)) @@ -5979,12 +6002,12 @@ def _execute_get_scenario_run_state_page( .limit(limit + 1) ) with closing(self._get_session()) as session: - rows = session.execute(statement).all() + rows = session.execute(statement).mappings().all() return ( [ ScenarioRunStateRecord( - scenario_result_id=str(row.id), - state=ScenarioRunState(row.scenario_run_state), + scenario_result_id=str(row["id"]), + state=ScenarioRunState(row["scenario_run_state"]), ) for row in rows[:limit] ], @@ -6078,7 +6101,7 @@ def _execute_get_scenario_history_aggregates( with closing(self._get_session()) as session: aggregate_rows = session.execute( self._build_scenario_history_aggregate_statement(entry_ids=entry_ids, plan_entry_ids=plan_entry_ids) - ).all() + ).mappings().all() name_rows = session.execute( select(AttackResultEntry.attribution_parent_id, self._get_scenario_attempt_unit_expressions()[0]) .where(AttackResultEntry.attribution_parent_id.in_(entry_ids)) @@ -6104,19 +6127,28 @@ def _execute_get_scenario_history_aggregates( continue names_by_run.setdefault(str(scenario_result_id), []).append(atomic_attack_name) for row in aggregate_rows: - if row.scenario_result_id is None: + if row["scenario_result_id"] is None: continue - run_id = str(row.scenario_result_id) + run_id = str(row["scenario_result_id"]) aggregates[run_id] = ScenarioHistoryAggregate( scenario_result_id=run_id, - unit_count=row.unit_count or 0, - completed_units=row.completed_units or 0, - successful_units=row.successful_units or 0, - error_attempts=row.error_attempts or 0, - total_retries=row.total_retries or 0, - latest_attempt_timestamp=row.latest_attempt_timestamp, + unit_count=row["unit_count"] or 0, + completed_units=row["completed_units"] or 0, + successful_units=row["successful_units"] or 0, + error_attempts=row["error_attempts"] or 0, + total_retries=row["total_retries"] or 0, + latest_attempt_timestamp=row["latest_attempt_timestamp"], atomic_attack_names=tuple(sorted(names_by_run.get(run_id, ()))), needs_sdk_statistics=run_id in sdk_run_ids, + target_facing_attempts=row["target_facing_attempts"] or 0, + target_facing_error_attempts=row["target_facing_error_attempts"] or 0, + target_facing_retries=row["target_facing_retries"] or 0, + orchestration_attempts=row["orchestration_attempts"] or 0, + orchestration_error_attempts=row["orchestration_error_attempts"] or 0, + orchestration_retries=row["orchestration_retries"] or 0, + unknown_role_attempts=row["unknown_role_attempts"] or 0, + unknown_role_error_attempts=row["unknown_role_error_attempts"] or 0, + unknown_role_retries=row["unknown_role_retries"] or 0, ) return aggregates @@ -6154,6 +6186,7 @@ def _build_scenario_history_aggregate_statement( ), else_=0, ).label("total_retries"), + JsonScalar(AttackResultEntry.attribution_data, literal("$.result_role")).label("result_role"), ) .where(AttackResultEntry.attribution_parent_id.in_(entry_ids)) .subquery("history_attempts") @@ -6176,6 +6209,8 @@ def _build_scenario_history_aggregate_statement( func.max(units.c.is_planned).over(partition_by=unit_partition).label("is_planned"), unit_retries.label("unit_retries"), func.sum(case((is_error, 1), else_=0)).over(partition_by=unit_partition).label("unit_errors"), + units.c.result_role.label("result_role"), + units.c.total_retries.label("attempt_retries"), func.row_number() .over( partition_by=unit_partition, @@ -6200,6 +6235,15 @@ def _build_scenario_history_aggregate_statement( func.sum(case((and_(counted, ranked.c.unit_retries > 0), ranked.c.unit_retries), else_=0)).label( "total_retries" ), + func.sum(case((ranked.c.result_role == AttackResultRole.TARGET_FACING.value, 1), else_=0)).label("target_facing_attempts"), + func.sum(case((and_(ranked.c.result_role == AttackResultRole.TARGET_FACING.value, ranked.c.latest_outcome == AttackOutcome.ERROR.value), 1), else_=0)).label("target_facing_error_attempts"), + func.sum(case((ranked.c.result_role == AttackResultRole.TARGET_FACING.value, ranked.c.attempt_retries), else_=0)).label("target_facing_retries"), + func.sum(case((ranked.c.result_role == AttackResultRole.ORCHESTRATION.value, 1), else_=0)).label("orchestration_attempts"), + func.sum(case((and_(ranked.c.result_role == AttackResultRole.ORCHESTRATION.value, ranked.c.latest_outcome == AttackOutcome.ERROR.value), 1), else_=0)).label("orchestration_error_attempts"), + func.sum(case((ranked.c.result_role == AttackResultRole.ORCHESTRATION.value, ranked.c.attempt_retries), else_=0)).label("orchestration_retries"), + func.sum(case((or_(ranked.c.result_role.is_(None), and_(ranked.c.result_role != AttackResultRole.TARGET_FACING.value, ranked.c.result_role != AttackResultRole.ORCHESTRATION.value)), 1), else_=0)).label("unknown_role_attempts"), + func.sum(case((and_(or_(ranked.c.result_role.is_(None), and_(ranked.c.result_role != AttackResultRole.TARGET_FACING.value, ranked.c.result_role != AttackResultRole.ORCHESTRATION.value)), ranked.c.latest_outcome == AttackOutcome.ERROR.value), 1), else_=0)).label("unknown_role_error_attempts"), + func.sum(case((or_(ranked.c.result_role.is_(None), and_(ranked.c.result_role != AttackResultRole.TARGET_FACING.value, ranked.c.result_role != AttackResultRole.ORCHESTRATION.value)), ranked.c.attempt_retries), else_=0)).label("unknown_role_retries"), ) .group_by(ranked.c.scenario_result_id) .order_by(ranked.c.scenario_result_id) @@ -6226,6 +6270,7 @@ def _build_scenario_history_unit_statement(self, *, attempts: Any, plan_entry_id attempts.c.outcome, attempts.c.timestamp, attempts.c.total_retries, + attempts.c.result_role, unplanned_group_id.label("unit_group_id"), attempts.c.seed_group_id.label("unit_seed_id"), literal(1).label("is_planned"), @@ -6305,6 +6350,7 @@ def _build_scenario_history_unit_statement(self, *, attempts: Any, plan_entry_id attempts.c.outcome, attempts.c.timestamp, attempts.c.total_retries, + attempts.c.result_role, attempts.c.atomic_attack_name, unplanned_group_id.label("unplanned_group_id"), attempts.c.seed_group_id, @@ -6330,6 +6376,7 @@ def _build_scenario_history_unit_statement(self, *, attempts: Any, plan_entry_id matched.c.outcome, matched.c.timestamp, matched.c.total_retries, + matched.c.result_role, func.coalesce(matched.c.atomic_group_id, matched.c.unplanned_group_id).label("unit_group_id"), func.coalesce(matched.c.planned_seed_group_id, matched.c.seed_group_id).label("unit_seed_id"), # Runs outside the plan-resolution set keep their raw identity and stay counted. @@ -6443,50 +6490,50 @@ def _execute_get_scenario_attack_result_deltas( .limit(limit + 1) ) with closing(self._get_session()) as session: - rows = session.execute(statement).all() + rows = session.execute(statement).mappings().all() has_more = len(rows) > limit deltas: list[ScenarioAttackResultDelta] = [] for row in rows[:limit]: retry_events = [ RetryEvent.model_validate(event) - for event in (json.loads(row.retry_events_json) if row.retry_events_json else []) + for event in (json.loads(row["retry_events_json"]) if row["retry_events_json"] else []) ] atomic_identifier = ( - AtomicAttackIdentifier.model_validate(row.atomic_attack_identifier) - if row.atomic_attack_identifier + AtomicAttackIdentifier.model_validate(row["atomic_attack_identifier"]) + if row["atomic_attack_identifier"] else None ) score = None - if row.score_id is not None: + if row["score_id"] is not None: scorer_identifier = ( - ComponentIdentifier.model_validate(row.scorer_class_identifier) - if row.scorer_class_identifier + ComponentIdentifier.model_validate(row["scorer_class_identifier"]) + if row["scorer_class_identifier"] else None ) score = ScenarioProgressScore( scorer_name=scorer_identifier.class_name if scorer_identifier else "Unknown", - score_type=row.score_type, - status=ScoreStatus(row.score_status), - score_value=row.score_value, - score_rationale=row.score_rationale, + score_type=row["score_type"], + status=ScoreStatus(row["score_status"]), + score_value=row["score_value"], + score_rationale=row["score_rationale"], ) deltas.append( ScenarioAttackResultDelta( - attack_result_id=str(row.id), - conversation_id=row.conversation_id, - objective=row.objective, - objective_sha256=row.objective_sha256, + attack_result_id=str(row["id"]), + conversation_id=row["conversation_id"], + objective=row["objective"], + objective_sha256=row["objective_sha256"], atomic_attack_identifier=atomic_identifier, - outcome=AttackOutcome(row.outcome), - execution_time_ms=row.execution_time_ms, - timestamp=row.timestamp, + outcome=AttackOutcome(row["outcome"]), + execution_time_ms=row["execution_time_ms"], + timestamp=row["timestamp"], retry_events=retry_events, - total_retries=row.total_retries or 0, - error_type=row.error_type, - error_message=row.error_message, - attribution_data=row.attribution_data or {}, - attack_metadata=row.attack_metadata or {}, + total_retries=row["total_retries"] or 0, + error_type=row["error_type"], + error_message=row["error_message"], + attribution_data=row["attribution_data"] or {}, + attack_metadata=row["attack_metadata"] or {}, score=score, ) ) diff --git a/pyrit/models/__init__.py b/pyrit/models/__init__.py index 3cfc87e1aa..9057ebedfe 100644 --- a/pyrit/models/__init__.py +++ b/pyrit/models/__init__.py @@ -49,6 +49,8 @@ ScenarioDatasetSizeCap, ScenarioDatasetSummary, ScenarioDefaultRunSizeEstimate, + ScenarioProducerCategoryCounts, + ScenarioProducerCounts, ScenarioRunListItem, ScenarioRunSizeComponent, ScenarioRunSizeEstimate, @@ -398,6 +400,8 @@ "ScenarioDatasetSizeCap": "pyrit.models.catalog", "ScenarioDatasetSummary": "pyrit.models.catalog", "ScenarioDefaultRunSizeEstimate": "pyrit.models.catalog", + "ScenarioProducerCategoryCounts": "pyrit.models.catalog", + "ScenarioProducerCounts": "pyrit.models.catalog", "ScenarioRunListItem": "pyrit.models.catalog", "ScenarioRunSizeComponent": "pyrit.models.catalog", "ScenarioRunSizeEstimate": "pyrit.models.catalog", diff --git a/pyrit/models/catalog/__init__.py b/pyrit/models/catalog/__init__.py index 3a0740a050..9eb465b5f8 100644 --- a/pyrit/models/catalog/__init__.py +++ b/pyrit/models/catalog/__init__.py @@ -28,6 +28,8 @@ ScenarioDatasetSizeCap, ScenarioDatasetSummary, ScenarioDefaultRunSizeEstimate, + ScenarioProducerCategoryCounts, + ScenarioProducerCounts, ScenarioRunListItem, ScenarioRunSizeComponent, ScenarioRunSizeEstimate, @@ -51,6 +53,8 @@ "ScenarioDatasetSizeCap": "pyrit.models.catalog.scenario", "ScenarioDatasetSummary": "pyrit.models.catalog.scenario", "ScenarioDefaultRunSizeEstimate": "pyrit.models.catalog.scenario", + "ScenarioProducerCategoryCounts": "pyrit.models.catalog.scenario", + "ScenarioProducerCounts": "pyrit.models.catalog.scenario", "ScenarioRunListItem": "pyrit.models.catalog.scenario", "ScenarioRunSizeComponent": "pyrit.models.catalog.scenario", "ScenarioRunSizeEstimate": "pyrit.models.catalog.scenario", diff --git a/pyrit/models/catalog/scenario.py b/pyrit/models/catalog/scenario.py index 7e9bde5aaa..0324a982d5 100644 --- a/pyrit/models/catalog/scenario.py +++ b/pyrit/models/catalog/scenario.py @@ -159,6 +159,22 @@ class ScenarioDatasetSummary(BaseModel): selection_note: str | None = None +class ScenarioProducerCategoryCounts(BaseModel): + """Raw attempt totals for one specific producer role.""" + + attempts: int = Field(default=0, ge=0) + errors: int = Field(default=0, ge=0) + retries: int = Field(default=0, ge=0) + + +class ScenarioProducerCounts(BaseModel): + """Raw attempts grouped by their recorded producer role.""" + + target_facing: ScenarioProducerCategoryCounts = Field(default_factory=ScenarioProducerCategoryCounts) + orchestration: ScenarioProducerCategoryCounts = Field(default_factory=ScenarioProducerCategoryCounts) + unknown: ScenarioProducerCategoryCounts = Field(default_factory=ScenarioProducerCategoryCounts) + + class ScenarioTechniqueSummary(BaseModel): """One concrete attack technique available to a scenario.""" @@ -569,6 +585,7 @@ class ScenarioRunSummary(BaseModel): default_factory=list, description="Bounded recent HTTP 429 and 5xx retry evidence grouped by component role", ) + producer_counts: ScenarioProducerCounts | None = Field(None, description="Raw producer attempts") class ScenarioRunListItem(BaseModel): @@ -611,6 +628,7 @@ class ScenarioRunListItem(BaseModel): True, description="Whether failed_attacks and attack_retries contain per-attempt details", ) + producer_counts: ScenarioProducerCounts | None = Field(None, description="Raw producer attempts") class ScenarioTargetSummary(BaseModel): diff --git a/pyrit/models/scenario_progress.py b/pyrit/models/scenario_progress.py index 6da0733b4d..b7e3ae2036 100644 --- a/pyrit/models/scenario_progress.py +++ b/pyrit/models/scenario_progress.py @@ -9,7 +9,12 @@ from pydantic import AwareDatetime, BaseModel, ConfigDict, Field, model_validator -from pyrit.models.catalog.scenario import ScenarioOverloadSummary, ScenarioTargetSummary # noqa: TC001 +from pyrit.models.catalog.scenario import ( + ScenarioOverloadSummary, + ScenarioProducerCategoryCounts, + ScenarioProducerCounts, + ScenarioTargetSummary, +) # noqa: TC001 from pyrit.models.identifiers.atomic_attack_identifier import AtomicAttackIdentifier from pyrit.models.results.attack_result import AttackOutcome, AttackResultRole from pyrit.models.results.scenario_result import ScenarioRunState @@ -191,6 +196,7 @@ class ScenarioProgressCounts(BaseModel): success_percentage: int | None = Field(default=None, ge=0, le=100) errors: int = Field(..., ge=0) retries: int = Field(..., ge=0) + producer_counts: ScenarioProducerCounts = Field(default_factory=ScenarioProducerCounts) class ScenarioExecutionUnit(BaseModel): diff --git a/tests/unit/backend/test_scenario_progress_read_model.py b/tests/unit/backend/test_scenario_progress_read_model.py index 264adcee52..f3f404de82 100644 --- a/tests/unit/backend/test_scenario_progress_read_model.py +++ b/tests/unit/backend/test_scenario_progress_read_model.py @@ -715,7 +715,19 @@ async def test_roles_do_not_change_progress_counts_async(self, sqlite_instance: ) assert {result.result_role for result in legacy.results} == {AttackResultRole.UNKNOWN} - assert legacy.summary == snapshot.summary + def _strip(obj): + if isinstance(obj, dict): + obj.pop("producer_counts", None) + for v in obj.values(): + _strip(v) + elif isinstance(obj, list): + for item in obj: + _strip(item) + return obj + + legacy_dump = _strip(legacy.summary.model_dump()) + snapshot_dump = _strip(snapshot.summary.model_dump()) + assert legacy_dump == snapshot_dump assert snapshot.summary.overall.planned == 2 diff --git a/tests/unit/backend/test_scenario_run_service.py b/tests/unit/backend/test_scenario_run_service.py index 35c05dcf3e..313c3d1e17 100644 --- a/tests/unit/backend/test_scenario_run_service.py +++ b/tests/unit/backend/test_scenario_run_service.py @@ -1907,6 +1907,15 @@ async def test_history_uses_plan_and_latest_non_error_attempt_per_unit(self, moc total_retries=3, latest_attempt_timestamp=timestamp, atomic_attack_names=("attack",), + target_facing_attempts=0, + target_facing_error_attempts=0, + target_facing_retries=0, + orchestration_attempts=0, + orchestration_error_attempts=0, + orchestration_retries=0, + unknown_role_attempts=0, + unknown_role_error_attempts=0, + unknown_role_retries=0, ) }, False, @@ -2040,6 +2049,15 @@ async def test_history_scopes_duplicate_objective_hashes_to_atomic_groups(self, total_retries=0, latest_attempt_timestamp=timestamp, atomic_attack_names=("attack-1", "attack-2"), + target_facing_attempts=0, + target_facing_error_attempts=0, + target_facing_retries=0, + orchestration_attempts=0, + orchestration_error_attempts=0, + orchestration_retries=0, + unknown_role_attempts=0, + unknown_role_error_attempts=0, + unknown_role_retries=0, ) }, False, @@ -2173,6 +2191,15 @@ async def test_history_requeries_legacy_aggregates_when_plan_is_rejected(self, m total_retries=4, latest_attempt_timestamp=datetime(2026, 8, 7, tzinfo=UTC), atomic_attack_names=("attack",), + target_facing_attempts=0, + target_facing_error_attempts=0, + target_facing_retries=0, + orchestration_attempts=0, + orchestration_error_attempts=0, + orchestration_retries=0, + unknown_role_attempts=0, + unknown_role_error_attempts=0, + unknown_role_retries=0, ) } ) @@ -3383,6 +3410,15 @@ async def test_history_and_detail_retry_work_match_across_attempt_partitions(moc total_retries=5, latest_attempt_timestamp=timestamp + timedelta(seconds=2), atomic_attack_names=("attack",), + target_facing_attempts=0, + target_facing_error_attempts=0, + target_facing_retries=0, + orchestration_attempts=0, + orchestration_error_attempts=0, + orchestration_retries=0, + unknown_role_attempts=0, + unknown_role_error_attempts=0, + unknown_role_retries=0, ) service = ScenarioRunService() diff --git a/tests/unit/memory/memory_interface/test_interface_seed_dataset_summary_query.py b/tests/unit/memory/memory_interface/test_interface_seed_dataset_summary_query.py index 820494cd22..f22c680c33 100644 --- a/tests/unit/memory/memory_interface/test_interface_seed_dataset_summary_query.py +++ b/tests/unit/memory/memory_interface/test_interface_seed_dataset_summary_query.py @@ -28,6 +28,20 @@ async def test_get_seed_dataset_summaries_avoids_metadata_row_multiplication( def execute(session, statement, *args, **kwargs): result = original_execute(session, statement, *args, **kwargs) + + original_mappings = getattr(result, "mappings", None) + if original_mappings: + def mappings_wrapper(): + m_result = original_mappings() + original_all = m_result.all + def all_rows(): + rows = original_all() + captured_rows.extend(rows) + return rows + m_result.all = all_rows + return m_result + result.mappings = mappings_wrapper + original_all = result.all def all_rows(): From 02f241f1e08660b39fcf4f0a42c1e8ec30dcb59e Mon Sep 17 00:00:00 2001 From: Nimit Jain Date: Sat, 10 Oct 2026 02:36:40 +0530 Subject: [PATCH 6/8] fix: integrate OutcomeStatistics and pass all tests (#3043) --- pyrit/memory/memory_interface.py | 18 +++++++++--------- .../test_scenario_statistics_parity.py | 16 ++++++++++++++++ tests/unit/memory/test_azure_sql_memory.py | 9 +++++++++ 3 files changed, 34 insertions(+), 9 deletions(-) diff --git a/pyrit/memory/memory_interface.py b/pyrit/memory/memory_interface.py index b78a1a81af..c187140a64 100644 --- a/pyrit/memory/memory_interface.py +++ b/pyrit/memory/memory_interface.py @@ -6235,15 +6235,15 @@ def _build_scenario_history_aggregate_statement( func.sum(case((and_(counted, ranked.c.unit_retries > 0), ranked.c.unit_retries), else_=0)).label( "total_retries" ), - func.sum(case((ranked.c.result_role == AttackResultRole.TARGET_FACING.value, 1), else_=0)).label("target_facing_attempts"), - func.sum(case((and_(ranked.c.result_role == AttackResultRole.TARGET_FACING.value, ranked.c.latest_outcome == AttackOutcome.ERROR.value), 1), else_=0)).label("target_facing_error_attempts"), - func.sum(case((ranked.c.result_role == AttackResultRole.TARGET_FACING.value, ranked.c.attempt_retries), else_=0)).label("target_facing_retries"), - func.sum(case((ranked.c.result_role == AttackResultRole.ORCHESTRATION.value, 1), else_=0)).label("orchestration_attempts"), - func.sum(case((and_(ranked.c.result_role == AttackResultRole.ORCHESTRATION.value, ranked.c.latest_outcome == AttackOutcome.ERROR.value), 1), else_=0)).label("orchestration_error_attempts"), - func.sum(case((ranked.c.result_role == AttackResultRole.ORCHESTRATION.value, ranked.c.attempt_retries), else_=0)).label("orchestration_retries"), - func.sum(case((or_(ranked.c.result_role.is_(None), and_(ranked.c.result_role != AttackResultRole.TARGET_FACING.value, ranked.c.result_role != AttackResultRole.ORCHESTRATION.value)), 1), else_=0)).label("unknown_role_attempts"), - func.sum(case((and_(or_(ranked.c.result_role.is_(None), and_(ranked.c.result_role != AttackResultRole.TARGET_FACING.value, ranked.c.result_role != AttackResultRole.ORCHESTRATION.value)), ranked.c.latest_outcome == AttackOutcome.ERROR.value), 1), else_=0)).label("unknown_role_error_attempts"), - func.sum(case((or_(ranked.c.result_role.is_(None), and_(ranked.c.result_role != AttackResultRole.TARGET_FACING.value, ranked.c.result_role != AttackResultRole.ORCHESTRATION.value)), ranked.c.attempt_retries), else_=0)).label("unknown_role_retries"), + func.sum(case((and_(ranked.c.is_planned == 1, ranked.c.result_role == AttackResultRole.TARGET_FACING.value), 1), else_=0)).label("target_facing_attempts"), + func.sum(case((and_(ranked.c.is_planned == 1, ranked.c.result_role == AttackResultRole.TARGET_FACING.value, ranked.c.latest_outcome == AttackOutcome.ERROR.value), 1), else_=0)).label("target_facing_error_attempts"), + func.sum(case((and_(ranked.c.is_planned == 1, ranked.c.result_role == AttackResultRole.TARGET_FACING.value), ranked.c.attempt_retries), else_=0)).label("target_facing_retries"), + func.sum(case((and_(ranked.c.is_planned == 1, ranked.c.result_role == AttackResultRole.ORCHESTRATION.value), 1), else_=0)).label("orchestration_attempts"), + func.sum(case((and_(ranked.c.is_planned == 1, ranked.c.result_role == AttackResultRole.ORCHESTRATION.value, ranked.c.latest_outcome == AttackOutcome.ERROR.value), 1), else_=0)).label("orchestration_error_attempts"), + func.sum(case((and_(ranked.c.is_planned == 1, ranked.c.result_role == AttackResultRole.ORCHESTRATION.value), ranked.c.attempt_retries), else_=0)).label("orchestration_retries"), + func.sum(case((and_(ranked.c.is_planned == 1, or_(ranked.c.result_role.is_(None), and_(ranked.c.result_role != AttackResultRole.TARGET_FACING.value, ranked.c.result_role != AttackResultRole.ORCHESTRATION.value))), 1), else_=0)).label("unknown_role_attempts"), + func.sum(case((and_(ranked.c.is_planned == 1, or_(ranked.c.result_role.is_(None), and_(ranked.c.result_role != AttackResultRole.TARGET_FACING.value, ranked.c.result_role != AttackResultRole.ORCHESTRATION.value)), ranked.c.latest_outcome == AttackOutcome.ERROR.value), 1), else_=0)).label("unknown_role_error_attempts"), + func.sum(case((and_(ranked.c.is_planned == 1, or_(ranked.c.result_role.is_(None), and_(ranked.c.result_role != AttackResultRole.TARGET_FACING.value, ranked.c.result_role != AttackResultRole.ORCHESTRATION.value))), ranked.c.attempt_retries), else_=0)).label("unknown_role_retries"), ) .group_by(ranked.c.scenario_result_id) .order_by(ranked.c.scenario_result_id) diff --git a/tests/unit/analytics/test_scenario_statistics_parity.py b/tests/unit/analytics/test_scenario_statistics_parity.py index 0a02bc6da3..feafc22ff3 100644 --- a/tests/unit/analytics/test_scenario_statistics_parity.py +++ b/tests/unit/analytics/test_scenario_statistics_parity.py @@ -26,6 +26,7 @@ AtomicAttackIdentifier, AttackOutcome, AttackResult, + AttackResultRole, AttackSeedGroup, ComponentIdentifier, OutcomeStatistics, @@ -55,6 +56,7 @@ class _Attempt: attributed_seed_context: str | None = None attack_result_id: str | None = None seconds: int | None = None + result_role: AttackResultRole = AttackResultRole.UNKNOWN @dataclass(frozen=True) @@ -245,6 +247,15 @@ def _plan(*groups: ScenarioRunPlanAtomicGroup, seeds: list[ScenarioRunPlanSeedGr plan=_plan(_group(name="attack", eval_hash="eval", seed_ids=["a"]), seeds=[_seed("a", "A")]), attempts=[], ), + "role_aware_accounting": _History( + plan=_plan(_group(name="attack", eval_hash="eval", seed_ids=["a"]), seeds=[_seed("a", "A")]), + attempts=[ + _Attempt("attack", "A", AttackOutcome.SUCCESS, seed_group_id="a", result_role=AttackResultRole.TARGET_FACING), + _Attempt("attack", "A", AttackOutcome.SUCCESS, seed_group_id="a", result_role=AttackResultRole.TARGET_FACING), + _Attempt("attack", "A", AttackOutcome.ERROR, seed_group_id="a", result_role=AttackResultRole.TARGET_FACING), + _Attempt("attack", "A", AttackOutcome.SUCCESS, seed_group_id="a", result_role=AttackResultRole.ORCHESTRATION), + ], + ), } # Effective-unit success percentages each history must report everywhere (None: no completed unit). @@ -266,6 +277,7 @@ def _plan(*groups: ScenarioRunPlanAtomicGroup, seeds: list[ScenarioRunPlanSeedGr "identifier_only_then_attributed_only": 100, "display_groups": 50, "empty_history": None, + "role_aware_accounting": 100, } @@ -297,6 +309,8 @@ async def _persist(memory: MemoryInterface, history: _History) -> str: if attempt.attributed_seed_context is not None: seed_group = _seed_group(attempt.objective, attempt.attributed_seed_context) attribution_data["seed_group_id"] = seed_group.logical_id + if attempt.result_role != AttackResultRole.UNKNOWN: + attribution_data["result_role"] = attempt.result_role.value atomic_attack_identifier = None if attempt.seed_context is not None: atomic_attack_identifier = AtomicAttackIdentifier.build( @@ -350,6 +364,8 @@ async def test_sdk_api_and_reports_report_identical_statistics(history_name: str assert progress.summary.overall.succeeded == sdk.overall.succeeded assert progress.summary.overall.errors == sdk.overall.errors assert progress.summary.overall.outcomes == sdk.overall.outcomes + assert progress.summary.overall.producer_counts == sdk.overall.producer_counts + assert list_item.producer_counts == sdk.overall.producer_counts # Reports report = json.loads(await JsonScenarioResultPrinter().render_async(scenario_result)) diff --git a/tests/unit/memory/test_azure_sql_memory.py b/tests/unit/memory/test_azure_sql_memory.py index 5b943f8230..90d957e7c1 100644 --- a/tests/unit/memory/test_azure_sql_memory.py +++ b/tests/unit/memory/test_azure_sql_memory.py @@ -923,3 +923,12 @@ def test_init_prod_with_skip_schema_migration_still_checks(): finally: Singleton._instances.clear() Singleton._instances.update(saved) + + +def test_scenario_history_aggregate_uses_sql_server_json_value(memory_interface: AzureSQLMemory) -> None: + statement = memory_interface._build_scenario_history_aggregate_statement( + entry_ids=[uuid.uuid4()], plan_entry_ids=[uuid.uuid4()] + ) + compiled = statement.compile(dialect=memory_interface.engine.dialect) + assert "json_value" in str(compiled).lower() + assert "$.result_role" in compiled.params.values() From 16ccb16ccc3ed38c73f52a6a69e71a28a67b447f Mon Sep 17 00:00:00 2001 From: Nimit Jain Date: Sun, 11 Oct 2026 11:49:29 +0530 Subject: [PATCH 7/8] test: add dedicated test coverage and resolve review comments (#3062) --- pyrit/analytics/__init__.py | 8 +- pyrit/analytics/scenario_statistics.py | 117 ++++++---- .../services/scenario_progress_read_model.py | 15 +- .../backend/services/scenario_run_service.py | 9 + pyrit/memory/azure_sql_memory.py | 19 +- pyrit/memory/memory_interface.py | 144 +++--------- pyrit/models/scenario_progress.py | 1 - pyrit/output/_derivation.py | 14 +- pyrit/output/scenario_result/html.py | 21 +- pyrit/output/scenario_result/json.py | 5 +- .../test_scenario_statistics_parity.py | 210 +++++++++++++++++- .../test_scenario_progress_read_model.py | 3 +- ...st_interface_seed_dataset_summary_query.py | 6 +- tests/unit/memory/test_azure_sql_memory.py | 28 ++- .../unit/output/scenario_result/test_html.py | 28 +++ .../unit/output/scenario_result/test_json.py | 1 + tests/unit/output/test_derivation.py | 14 +- 17 files changed, 471 insertions(+), 172 deletions(-) diff --git a/pyrit/analytics/__init__.py b/pyrit/analytics/__init__.py index a0413932b3..f5113c3f66 100644 --- a/pyrit/analytics/__init__.py +++ b/pyrit/analytics/__init__.py @@ -18,7 +18,11 @@ get_cached_results_for_technique, get_cached_results_for_technique_async, ) - from pyrit.analytics.scenario_statistics import compute_scenario_statistics + from pyrit.analytics.scenario_statistics import ( + compute_producer_counts, + compute_scenario_statistics, + count_execution_units, + ) from pyrit.analytics.technique_analysis import compute_technique_stats_async from pyrit.analytics.text_matching import ApproximateTextMatching, ExactTextMatching, TextMatching @@ -29,8 +33,10 @@ "AttackStats": "pyrit.analytics.result_analysis", "combine_outcome_statistics": "pyrit.analytics.outcome_statistics", "compute_outcome_statistics": "pyrit.analytics.outcome_statistics", + "compute_producer_counts": "pyrit.analytics.scenario_statistics", "compute_scenario_statistics": "pyrit.analytics.scenario_statistics", "compute_technique_stats_async": "pyrit.analytics.technique_analysis", + "count_execution_units": "pyrit.analytics.scenario_statistics", "ConversationAnalytics": "pyrit.analytics.conversation_analytics", "ExactTextMatching": "pyrit.analytics.text_matching", "get_cached_results_for_technique": "pyrit.analytics.result_analysis", diff --git a/pyrit/analytics/scenario_statistics.py b/pyrit/analytics/scenario_statistics.py index ada192aa41..3f0dfecc7a 100644 --- a/pyrit/analytics/scenario_statistics.py +++ b/pyrit/analytics/scenario_statistics.py @@ -38,11 +38,13 @@ AtomicAttackIdentifier, AttackOutcome, AttackResult, - AttackResultRole, AttackResultMetadata, + AttackResultRole, ComponentIdentifier, ScenarioExecutionStatistics, ScenarioExecutionUnit, + ScenarioProducerCategoryCounts, + ScenarioProducerCounts, ScenarioProgressCounts, ScenarioRunPlan, ScenarioRunPlanAtomicGroup, @@ -268,24 +270,16 @@ def retry_pressure(*, attempts_per_unit: Iterable[int], persisted_retries: Itera return within_attempts + repeated_units -def count_execution_units( - *, - units: Iterable[ScenarioExecutionUnit], - attempts_by_unit: Mapping[ScenarioExecutionUnit, Sequence[_CountableAttempt]], - planned: int | None, -) -> ScenarioProgressCounts: +def compute_producer_counts(attempts: Iterable[_CountableAttempt]) -> ScenarioProducerCounts: """ - Count effective execution units from chronologically ordered attempts. + Compute role-aware attempt, error, and retry accounting across attempts. - Each item in ``attempts_by_unit`` must expose ``outcome`` and ``total_retries`` and be ordered - oldest first; the last attempt decides the unit's outcome. + Args: + attempts (Iterable[_CountableAttempt]): Chronological or arbitrary attempt sequence. Returns: - ScenarioProgressCounts: Shared statistics for latest unit outcomes plus historical errors and retries. + ScenarioProducerCounts: Role-specific attempts, errors, and retries. """ - counts: Counter[AttackOutcome] = Counter() - errors = 0 - retries = 0 target_facing_attempts = 0 target_facing_errors = 0 target_facing_retries = 0 @@ -296,34 +290,26 @@ def count_execution_units( unknown_errors = 0 unknown_retries = 0 - for unit in units: - attempts = attempts_by_unit.get(unit, ()) - if not attempts: - continue - counts[attempts[-1].outcome] += 1 - errors += sum(int(attempt.outcome == AttackOutcome.ERROR) for attempt in attempts) - retries += retry_pressure( - attempts_per_unit=[len(attempts)], - persisted_retries=[attempt.total_retries for attempt in attempts], - ) - for attempt in attempts: - role = attempt.result_role - if role == AttackResultRole.TARGET_FACING: - target_facing_attempts += 1 - target_facing_errors += int(attempt.outcome == AttackOutcome.ERROR) - target_facing_retries += max(attempt.total_retries, 0) - elif role == AttackResultRole.ORCHESTRATION: - orchestration_attempts += 1 - orchestration_errors += int(attempt.outcome == AttackOutcome.ERROR) - orchestration_retries += max(attempt.total_retries, 0) - else: - unknown_attempts += 1 - unknown_errors += int(attempt.outcome == AttackOutcome.ERROR) - unknown_retries += max(attempt.total_retries, 0) - - from pyrit.models import ScenarioProducerCounts, ScenarioProducerCategoryCounts - - producer_counts = ScenarioProducerCounts( + for attempt in attempts: + role = getattr(attempt, "result_role", None) + outcome = getattr(attempt, "outcome", None) + retries = max(getattr(attempt, "total_retries", 0), 0) + is_error = int(outcome == AttackOutcome.ERROR) + + if role == AttackResultRole.TARGET_FACING: + target_facing_attempts += 1 + target_facing_errors += is_error + target_facing_retries += retries + elif role == AttackResultRole.ORCHESTRATION: + orchestration_attempts += 1 + orchestration_errors += is_error + orchestration_retries += retries + else: + unknown_attempts += 1 + unknown_errors += is_error + unknown_retries += retries + + return ScenarioProducerCounts( target_facing=ScenarioProducerCategoryCounts( attempts=target_facing_attempts, errors=target_facing_errors, @@ -341,6 +327,41 @@ def count_execution_units( ), ) + +def count_execution_units( + *, + units: Iterable[ScenarioExecutionUnit], + attempts_by_unit: Mapping[ScenarioExecutionUnit, Sequence[_CountableAttempt]], + planned: int | None, + producer_attempts: Iterable[_CountableAttempt] | None = None, +) -> ScenarioProgressCounts: + """ + Count effective execution units from chronologically ordered attempts. + + Each item in ``attempts_by_unit`` must expose ``outcome`` and ``total_retries`` and be ordered + oldest first; the last attempt decides the unit's outcome. + + Returns: + ScenarioProgressCounts: Shared statistics for latest unit outcomes plus historical errors and retries. + """ + counts: Counter[AttackOutcome] = Counter() + errors = 0 + retries = 0 + unit_attempts: list[_CountableAttempt] = [] + + for unit in units: + attempts = attempts_by_unit.get(unit, ()) + if not attempts: + continue + counts[attempts[-1].outcome] += 1 + errors += sum(int(attempt.outcome == AttackOutcome.ERROR) for attempt in attempts) + retries += retry_pressure( + attempts_per_unit=[len(attempts)], + persisted_retries=[attempt.total_retries for attempt in attempts], + ) + unit_attempts.extend(attempts) + + producer_counts = compute_producer_counts(producer_attempts if producer_attempts is not None else unit_attempts) outcomes = compute_outcome_statistics(counts) return ScenarioProgressCounts( completed=outcomes.total_results, @@ -367,7 +388,8 @@ def combine_execution_counts(counts: Iterable[ScenarioProgressCounts]) -> Scenar completed = sum(item.completed for item in counts) succeeded = sum(item.succeeded for item in counts) planned = [item.planned for item in counts] - from pyrit.models import ScenarioProducerCounts, ScenarioProducerCategoryCounts + from pyrit.models import ScenarioProducerCategoryCounts, ScenarioProducerCounts + producer_counts = ScenarioProducerCounts( target_facing=ScenarioProducerCategoryCounts( attempts=sum(item.producer_counts.target_facing.attempts for item in counts), @@ -470,15 +492,20 @@ def compute_scenario_statistics( len(unit_attempts) for unit, unit_attempts in attempts_by_unit.items() if unit not in counted ) - def _count(units: Sequence[ScenarioExecutionUnit]) -> ScenarioProgressCounts: + def _count( + units: Sequence[ScenarioExecutionUnit], + *, + producer_attempts: Iterable[_CountableAttempt] | None = None, + ) -> ScenarioProgressCounts: return count_execution_units( units=units, attempts_by_unit=attempts_by_unit, planned=len(units) if planned is not None else None, + producer_attempts=producer_attempts, ) return ScenarioExecutionStatistics( - overall=_count(counted_units), + overall=_count(counted_units, producer_attempts=attempts), atomic_attacks={name: _count(units) for name, units in units_by_name.items()}, display_groups={name: _count(units) for name, units in units_by_display_group.items()}, attempts=len(attempts), diff --git a/pyrit/backend/services/scenario_progress_read_model.py b/pyrit/backend/services/scenario_progress_read_model.py index 53ba860ca6..63451fc4bb 100644 --- a/pyrit/backend/services/scenario_progress_read_model.py +++ b/pyrit/backend/services/scenario_progress_read_model.py @@ -340,8 +340,18 @@ def _build_progress_summary( ) attempts_by_unit.setdefault(identity, []).append(result) - def aggregate(*, units: Sequence[ResultUnitIdentity], planned: int | None) -> ScenarioProgressCounts: - return count_execution_units(units=units, attempts_by_unit=attempts_by_unit, planned=planned) + def aggregate( + *, + units: Sequence[ResultUnitIdentity], + planned: int | None, + producer_attempts: Sequence[ScenarioProgressResult] | None = None, + ) -> ScenarioProgressCounts: + return count_execution_units( + units=units, + attempts_by_unit=attempts_by_unit, + planned=planned, + producer_attempts=producer_attempts, + ) group_units: dict[str, list[ResultUnitIdentity]] = { group.id: [ @@ -356,6 +366,7 @@ def aggregate(*, units: Sequence[ResultUnitIdentity], planned: int | None) -> Sc overall = aggregate( units=overall_units, planned=len(overall_units) if plan_complete else None, + producer_attempts=results, ) planned_units = set(overall_units) unattributed_attempts = sum( diff --git a/pyrit/backend/services/scenario_run_service.py b/pyrit/backend/services/scenario_run_service.py index 1fe729f635..c443f670d6 100644 --- a/pyrit/backend/services/scenario_run_service.py +++ b/pyrit/backend/services/scenario_run_service.py @@ -1621,6 +1621,15 @@ async def _recount_with_sdk_statistics_async( successful_units=overall.succeeded, error_attempts=overall.errors, total_retries=overall.retries, + target_facing_attempts=overall.producer_counts.target_facing.attempts, + target_facing_error_attempts=overall.producer_counts.target_facing.errors, + target_facing_retries=overall.producer_counts.target_facing.retries, + orchestration_attempts=overall.producer_counts.orchestration.attempts, + orchestration_error_attempts=overall.producer_counts.orchestration.errors, + orchestration_retries=overall.producer_counts.orchestration.retries, + unknown_role_attempts=overall.producer_counts.unknown.attempts, + unknown_role_error_attempts=overall.producer_counts.unknown.errors, + unknown_role_retries=overall.producer_counts.unknown.retries, ) return recounted diff --git a/pyrit/memory/azure_sql_memory.py b/pyrit/memory/azure_sql_memory.py index d5a1d93e12..cba9414187 100644 --- a/pyrit/memory/azure_sql_memory.py +++ b/pyrit/memory/azure_sql_memory.py @@ -20,6 +20,7 @@ event, exists, func, + literal, literal_column, text, ) @@ -44,7 +45,7 @@ ) from pyrit.memory.memory_session import MemorySession from pyrit.memory.storage import AzureBlobStorageIO -from pyrit.models import ConversationStats +from pyrit.models import AttackResultRole, ConversationStats if TYPE_CHECKING: from azure.core.credentials import AccessToken @@ -807,6 +808,22 @@ def _get_scenario_attempt_id_order_expression( # Native SQL Server UUID ordering differs from the SDK's canonical string comparison. return func.lower(sql_cast(attempt_id, String(36))).collate("Latin1_General_100_BIN2") + def _get_scenario_role_match_condition(self, *, role_expression: Any, role: AttackResultRole) -> Any: + """ + Return an exact case-sensitive and trailing-space-sensitive condition for AttackResultRole. + + SQL Server pads strings for '=' comparisons and uses case-insensitive collation by default. + Binary collation ensures exact case matching, and DATALENGTH ensures trailing spaces or + malformed padding are not ignored. + + Returns: + A SQL Server condition strictly matching the specified AttackResultRole. + """ + return and_( + role_expression.collate("Latin1_General_100_BIN2") == role.value, + func.datalength(role_expression) == func.datalength(literal(role.value, type_=Unicode)), + ) + def _get_scenario_attempt_unit_expressions(self) -> tuple[Any, Any, Any, Any]: """Return SQL Server JSON expressions for persisted scenario attempt attribution.""" atomic_name = func.coalesce( diff --git a/pyrit/memory/memory_interface.py b/pyrit/memory/memory_interface.py index 86f5ca3d7b..738ef53e1f 100644 --- a/pyrit/memory/memory_interface.py +++ b/pyrit/memory/memory_interface.py @@ -2171,6 +2171,21 @@ def _get_scenario_attempt_id_order_expression( """Return the scenario attempt ID's canonical string ordering.""" return attempt_id + def _get_scenario_role_match_condition(self, *, role_expression: Any, role: AttackResultRole) -> Any: + """ + Return a backend-specific condition that strictly matches an AttackResultRole. + + Subclasses override this when backend comparison semantics (such as SQL Server's string + padding or case-insensitive default collation) require exact scalar matching. + + Returns: + A backend-specific condition matching the specified AttackResultRole. + """ + return and_( + role_expression == role.value, + func.length(role_expression) == len(role.value), + ) + def _get_scenario_attempt_unit_expressions(self) -> tuple[Any, Any, Any, Any]: """ Return backend-specific JSON expressions for scenario attempt unit attribution. @@ -6226,6 +6241,16 @@ def _build_scenario_history_aggregate_statement( .label("unit_rank"), ).subquery("history_ranked_units") counted = and_(ranked.c.unit_rank == 1, ranked.c.is_planned == 1) + is_target_facing = self._get_scenario_role_match_condition( + role_expression=ranked.c.result_role, role=AttackResultRole.TARGET_FACING + ) + is_orchestration = self._get_scenario_role_match_condition( + role_expression=ranked.c.result_role, role=AttackResultRole.ORCHESTRATION + ) + is_unknown = or_( + ranked.c.result_role.is_(None), + and_(not_(is_target_facing), not_(is_orchestration)), + ) return ( select( ranked.c.scenario_result_id, @@ -6239,131 +6264,30 @@ def _build_scenario_history_aggregate_statement( func.sum(case((and_(counted, ranked.c.unit_retries > 0), ranked.c.unit_retries), else_=0)).label( "total_retries" ), + func.sum(case((is_target_facing, 1), else_=0)).label("target_facing_attempts"), func.sum( case( - ( - and_( - ranked.c.is_planned == 1, ranked.c.result_role == AttackResultRole.TARGET_FACING.value - ), - 1, - ), - else_=0, - ) - ).label("target_facing_attempts"), - func.sum( - case( - ( - and_( - ranked.c.is_planned == 1, - ranked.c.result_role == AttackResultRole.TARGET_FACING.value, - ranked.c.latest_outcome == AttackOutcome.ERROR.value, - ), - 1, - ), + (and_(is_target_facing, ranked.c.latest_outcome == AttackOutcome.ERROR.value), 1), else_=0, ) ).label("target_facing_error_attempts"), + func.sum(case((is_target_facing, ranked.c.attempt_retries), else_=0)).label("target_facing_retries"), + func.sum(case((is_orchestration, 1), else_=0)).label("orchestration_attempts"), func.sum( case( - ( - and_( - ranked.c.is_planned == 1, ranked.c.result_role == AttackResultRole.TARGET_FACING.value - ), - ranked.c.attempt_retries, - ), - else_=0, - ) - ).label("target_facing_retries"), - func.sum( - case( - ( - and_( - ranked.c.is_planned == 1, ranked.c.result_role == AttackResultRole.ORCHESTRATION.value - ), - 1, - ), - else_=0, - ) - ).label("orchestration_attempts"), - func.sum( - case( - ( - and_( - ranked.c.is_planned == 1, - ranked.c.result_role == AttackResultRole.ORCHESTRATION.value, - ranked.c.latest_outcome == AttackOutcome.ERROR.value, - ), - 1, - ), + (and_(is_orchestration, ranked.c.latest_outcome == AttackOutcome.ERROR.value), 1), else_=0, ) ).label("orchestration_error_attempts"), + func.sum(case((is_orchestration, ranked.c.attempt_retries), else_=0)).label("orchestration_retries"), + func.sum(case((is_unknown, 1), else_=0)).label("unknown_role_attempts"), func.sum( case( - ( - and_( - ranked.c.is_planned == 1, ranked.c.result_role == AttackResultRole.ORCHESTRATION.value - ), - ranked.c.attempt_retries, - ), - else_=0, - ) - ).label("orchestration_retries"), - func.sum( - case( - ( - and_( - ranked.c.is_planned == 1, - or_( - ranked.c.result_role.is_(None), - and_( - ranked.c.result_role != AttackResultRole.TARGET_FACING.value, - ranked.c.result_role != AttackResultRole.ORCHESTRATION.value, - ), - ), - ), - 1, - ), - else_=0, - ) - ).label("unknown_role_attempts"), - func.sum( - case( - ( - and_( - ranked.c.is_planned == 1, - or_( - ranked.c.result_role.is_(None), - and_( - ranked.c.result_role != AttackResultRole.TARGET_FACING.value, - ranked.c.result_role != AttackResultRole.ORCHESTRATION.value, - ), - ), - ranked.c.latest_outcome == AttackOutcome.ERROR.value, - ), - 1, - ), + (and_(is_unknown, ranked.c.latest_outcome == AttackOutcome.ERROR.value), 1), else_=0, ) ).label("unknown_role_error_attempts"), - func.sum( - case( - ( - and_( - ranked.c.is_planned == 1, - or_( - ranked.c.result_role.is_(None), - and_( - ranked.c.result_role != AttackResultRole.TARGET_FACING.value, - ranked.c.result_role != AttackResultRole.ORCHESTRATION.value, - ), - ), - ), - ranked.c.attempt_retries, - ), - else_=0, - ) - ).label("unknown_role_retries"), + func.sum(case((is_unknown, ranked.c.attempt_retries), else_=0)).label("unknown_role_retries"), ) .group_by(ranked.c.scenario_result_id) .order_by(ranked.c.scenario_result_id) diff --git a/pyrit/models/scenario_progress.py b/pyrit/models/scenario_progress.py index 386b5d860a..8bda2544da 100644 --- a/pyrit/models/scenario_progress.py +++ b/pyrit/models/scenario_progress.py @@ -12,7 +12,6 @@ from pyrit.models.analytics import OutcomeStatistics from pyrit.models.catalog.scenario import ( ScenarioOverloadSummary, - ScenarioProducerCategoryCounts, ScenarioProducerCounts, ScenarioTargetSummary, ) # noqa: TC001 diff --git a/pyrit/output/_derivation.py b/pyrit/output/_derivation.py index 5de0b7a5bc..af68e543b8 100644 --- a/pyrit/output/_derivation.py +++ b/pyrit/output/_derivation.py @@ -18,7 +18,15 @@ from pyrit.analytics.scenario_statistics import combine_execution_counts, compute_scenario_statistics if TYPE_CHECKING: - from pyrit.models import AttackResult, ComponentIdentifier, MessagePiece, OutcomeStatistics, ScenarioResult, Score + from pyrit.models import ( + AttackResult, + ComponentIdentifier, + MessagePiece, + OutcomeStatistics, + ScenarioProducerCounts, + ScenarioResult, + Score, + ) class TargetInfo(NamedTuple): @@ -62,6 +70,7 @@ class GroupStatistics(NamedTuple): attempts: int success_rate: int outcomes: OutcomeStatistics | None = None + producer_counts: ScenarioProducerCounts | None = None class ScenarioOverview(NamedTuple): @@ -72,6 +81,7 @@ class ScenarioOverview(NamedTuple): success_rate: int groups: list[GroupStatistics] outcomes: OutcomeStatistics | None = None + producer_counts: ScenarioProducerCounts | None = None def scenario_overview(result: ScenarioResult) -> ScenarioOverview: @@ -103,6 +113,7 @@ def scenario_overview(result: ScenarioResult) -> ScenarioOverview: attempts=len(group_results), success_rate=counts.success_percentage or 0, outcomes=counts.outcomes, + producer_counts=counts.producer_counts, ) ) return ScenarioOverview( @@ -111,6 +122,7 @@ def scenario_overview(result: ScenarioResult) -> ScenarioOverview: success_rate=statistics.overall.success_percentage or 0, groups=groups, outcomes=statistics.overall.outcomes, + producer_counts=statistics.overall.producer_counts, ) diff --git a/pyrit/output/scenario_result/html.py b/pyrit/output/scenario_result/html.py index d2804e2d48..d939f36ab5 100644 --- a/pyrit/output/scenario_result/html.py +++ b/pyrit/output/scenario_result/html.py @@ -55,6 +55,16 @@
{{ report.overview.stats.total_attempts }}
Objectives
{{ report.overview.stats.unique_objectives }}
+ {% if report.overview.stats.producer_counts %} +
Target-facing attempts
+
{{ report.overview.stats.producer_counts.target_facing.attempts }}
+
Orchestration attempts
+
{{ report.overview.stats.producer_counts.orchestration.attempts }}
+ {% if report.overview.stats.producer_counts.unknown.attempts %} +
Unknown attempts
+
{{ report.overview.stats.producer_counts.unknown.attempts }}
+ {% endif %} + {% endif %}

Target

@@ -66,10 +76,19 @@

Per-group breakdown

- + + + + + + + + {% for g in report.overview.groups %} + + {% endfor %}
GroupObjective executionsAttemptsSuccess rate
GroupObjective executionsAttemptsTarget-facingOrchestrationSuccess rate
{{ g.name }}{{ g.num_objective_executions }} {{ g.num_attempts }}{{ g.producer_counts.target_facing.attempts if g.producer_counts else "—" }}{{ g.producer_counts.orchestration.attempts if g.producer_counts else "—" }} {{ g.success_rate }}%
diff --git a/pyrit/output/scenario_result/json.py b/pyrit/output/scenario_result/json.py index cbd25f2ae3..137c7d6aeb 100644 --- a/pyrit/output/scenario_result/json.py +++ b/pyrit/output/scenario_result/json.py @@ -6,7 +6,7 @@ from pydantic import TypeAdapter -from pyrit.models import AttackResult, OutcomeStatistics, ScenarioResult +from pyrit.models import AttackResult, OutcomeStatistics, ScenarioProducerCounts, ScenarioResult from pyrit.output._derivation import ( attack_score_display, resolve_target_info, @@ -36,6 +36,7 @@ class JsonScenarioResultPrinter(ScenarioResultPrinterBase): """ _OUTCOME_ADAPTER: ClassVar[TypeAdapter[OutcomeStatistics | None]] = TypeAdapter(OutcomeStatistics | None) + _PRODUCER_ADAPTER: ClassVar[TypeAdapter[ScenarioProducerCounts | None]] = TypeAdapter(ScenarioProducerCounts | None) def __init__( self, @@ -134,6 +135,7 @@ def _build_overview(self, result: ScenarioResult) -> dict[str, Any]: "num_attempts": group.attempts, "success_rate": group.success_rate, "outcomes": self._OUTCOME_ADAPTER.dump_python(group.outcomes, mode="json"), + "producer_counts": self._PRODUCER_ADAPTER.dump_python(group.producer_counts, mode="json"), } for group in overview.groups ] @@ -167,6 +169,7 @@ def _build_overview(self, result: ScenarioResult) -> dict[str, Any]: "total_attempts": overview.attempts, "overall_success_rate": overview.success_rate, "outcomes": self._OUTCOME_ADAPTER.dump_python(overview.outcomes, mode="json"), + "producer_counts": self._PRODUCER_ADAPTER.dump_python(overview.producer_counts, mode="json"), "unique_objectives": len(result.get_objectives()), }, "groups": groups, diff --git a/tests/unit/analytics/test_scenario_statistics_parity.py b/tests/unit/analytics/test_scenario_statistics_parity.py index feafc22ff3..062c24a154 100644 --- a/tests/unit/analytics/test_scenario_statistics_parity.py +++ b/tests/unit/analytics/test_scenario_statistics_parity.py @@ -30,6 +30,7 @@ AttackSeedGroup, ComponentIdentifier, OutcomeStatistics, + ScenarioProducerCounts, ScenarioRunPlan, ScenarioRunPlanAtomicGroup, ScenarioRunPlanSeedGroup, @@ -56,7 +57,7 @@ class _Attempt: attributed_seed_context: str | None = None attack_result_id: str | None = None seconds: int | None = None - result_role: AttackResultRole = AttackResultRole.UNKNOWN + result_role: AttackResultRole | str = AttackResultRole.UNKNOWN @dataclass(frozen=True) @@ -250,10 +251,16 @@ def _plan(*groups: ScenarioRunPlanAtomicGroup, seeds: list[ScenarioRunPlanSeedGr "role_aware_accounting": _History( plan=_plan(_group(name="attack", eval_hash="eval", seed_ids=["a"]), seeds=[_seed("a", "A")]), attempts=[ - _Attempt("attack", "A", AttackOutcome.SUCCESS, seed_group_id="a", result_role=AttackResultRole.TARGET_FACING), - _Attempt("attack", "A", AttackOutcome.SUCCESS, seed_group_id="a", result_role=AttackResultRole.TARGET_FACING), + _Attempt( + "attack", "A", AttackOutcome.SUCCESS, seed_group_id="a", result_role=AttackResultRole.TARGET_FACING + ), + _Attempt( + "attack", "A", AttackOutcome.SUCCESS, seed_group_id="a", result_role=AttackResultRole.TARGET_FACING + ), _Attempt("attack", "A", AttackOutcome.ERROR, seed_group_id="a", result_role=AttackResultRole.TARGET_FACING), - _Attempt("attack", "A", AttackOutcome.SUCCESS, seed_group_id="a", result_role=AttackResultRole.ORCHESTRATION), + _Attempt( + "attack", "A", AttackOutcome.SUCCESS, seed_group_id="a", result_role=AttackResultRole.ORCHESTRATION + ), ], ), } @@ -310,7 +317,11 @@ async def _persist(memory: MemoryInterface, history: _History) -> str: seed_group = _seed_group(attempt.objective, attempt.attributed_seed_context) attribution_data["seed_group_id"] = seed_group.logical_id if attempt.result_role != AttackResultRole.UNKNOWN: - attribution_data["result_role"] = attempt.result_role.value + attribution_data["result_role"] = ( + attempt.result_role.value + if isinstance(attempt.result_role, AttackResultRole) + else str(attempt.result_role) + ) atomic_attack_identifier = None if attempt.seed_context is not None: atomic_attack_identifier = AtomicAttackIdentifier.build( @@ -371,6 +382,9 @@ async def test_sdk_api_and_reports_report_identical_statistics(history_name: str report = json.loads(await JsonScenarioResultPrinter().render_async(scenario_result)) assert report["stats"]["overall_success_rate"] == (expected or 0) assert report["stats"]["outcomes"] == TypeAdapter(OutcomeStatistics).dump_python(sdk.overall.outcomes, mode="json") + assert report["stats"]["producer_counts"] == TypeAdapter(ScenarioProducerCounts).dump_python( + sdk.overall.producer_counts, mode="json" + ) # Per-group numbers agree between the SDK, the saved-plan progress view, and the reports. Compare # key sets first so a group missing from one view fails instead of reading as 0%. @@ -390,6 +404,9 @@ async def test_sdk_api_and_reports_report_identical_statistics(history_name: str assert group["outcomes"] == TypeAdapter(OutcomeStatistics).dump_python( sdk.display_groups[group["name"]].outcomes, mode="json" ) + assert group["producer_counts"] == TypeAdapter(ScenarioProducerCounts).dump_python( + sdk.display_groups[group["name"]].producer_counts, mode="json" + ) if history.plan is not None: progress_groups = { group.display_group: (group.completed, group.success_percentage) @@ -462,3 +479,186 @@ async def test_history_aggregate_flags_runs_with_identifier_only_attempts(sqlite assert aggregates[flagged].needs_sdk_statistics assert not aggregates[attributed].needs_sdk_statistics + + +async def test_unmatched_attempts_visible_in_producer_accounting(sqlite_instance) -> None: + # 4 rows: 1 planned target-facing success, 1 unmatched target-facing error, + # 1 unmatched orchestration row, and 1 unmatched malformed-role row. + plan = _plan( + _group(name="attack", eval_hash="eval", seed_ids=["seed_1"]), + seeds=[_seed("seed_1", "Objective 1")], + ) + history = _History( + plan=plan, + attempts=[ + _Attempt( + atomic_attack_name="attack", + objective="Objective 1", + outcome=AttackOutcome.SUCCESS, + seed_group_id="seed_1", + result_role=AttackResultRole.TARGET_FACING, + ), + _Attempt( + atomic_attack_name="attack", + objective="Unmatched Objective 2", + outcome=AttackOutcome.ERROR, + seed_group_id="unmatched_seed", + result_role=AttackResultRole.TARGET_FACING, + ), + _Attempt( + atomic_attack_name="orchestrator", + objective="Unmatched Objective 3", + outcome=AttackOutcome.SUCCESS, + result_role=AttackResultRole.ORCHESTRATION, + ), + _Attempt( + atomic_attack_name="attack", + objective="Unmatched Objective 4", + outcome=AttackOutcome.ERROR, + # Malformed role that falls back to UNKNOWN + result_role="target_facing ", + ), + ], + ) + scenario_result_id = await _persist(sqlite_instance, history) + + # SDK + [scenario_result] = await sqlite_instance.get_scenario_results_async(scenario_result_ids=[scenario_result_id]) + sdk = compute_scenario_statistics(scenario_result) + + assert sdk.attempts == 4 + assert sdk.unattributed_attempts == 3 + assert sdk.overall.completed == 1 + assert sdk.overall.planned == 1 + assert sdk.overall.succeeded == 1 + + # Overall producer accounting includes all 4 rows + assert sdk.overall.producer_counts.target_facing.attempts == 2 + assert sdk.overall.producer_counts.target_facing.errors == 1 + assert sdk.overall.producer_counts.orchestration.attempts == 1 + assert sdk.overall.producer_counts.orchestration.errors == 0 + assert sdk.overall.producer_counts.unknown.attempts == 1 + assert sdk.overall.producer_counts.unknown.errors == 1 + + # API / Service + service = ScenarioRunService() + runs = await service.list_runs_async() + [list_item] = [item for item in runs.items if item.scenario_result_id == scenario_result_id] + progress = await service.get_run_progress_from_storage_async( + scenario_result_id=scenario_result_id, since=None, limit=500, active_group_ids=[] + ) + + # Progress & History reconcile with SDK + assert progress.summary.unattributed_attempts == 3 + assert progress.summary.overall.completed == 1 + assert progress.summary.overall.planned == 1 + assert progress.summary.overall.producer_counts == sdk.overall.producer_counts + assert list_item.producer_counts == sdk.overall.producer_counts + assert list_item.completed_attacks == 1 + + +async def test_malformed_roles_fall_back_to_unknown_in_producer_accounting(sqlite_instance) -> None: + # Verify that malformed role values (trailing spaces or uppercase) fall back to UNKNOWN + # in both SDK statistics and backend SQL aggregation. + plan = _plan( + _group(name="attack", eval_hash="eval", seed_ids=["seed_1"]), + seeds=[_seed("seed_1", "Objective 1")], + ) + history = _History( + plan=plan, + attempts=[ + _Attempt( + atomic_attack_name="attack", + objective="Objective 1", + outcome=AttackOutcome.SUCCESS, + seed_group_id="seed_1", + result_role="target_facing ", # trailing space + ), + _Attempt( + atomic_attack_name="attack", + objective="Objective 1", + outcome=AttackOutcome.SUCCESS, + seed_group_id="seed_1", + result_role="TARGET_FACING", # uppercase + ), + _Attempt( + atomic_attack_name="attack", + objective="Objective 1", + outcome=AttackOutcome.SUCCESS, + seed_group_id="seed_1", + result_role=AttackResultRole.TARGET_FACING, # exact canonical match + ), + ], + ) + scenario_result_id = await _persist(sqlite_instance, history) + + [scenario_result] = await sqlite_instance.get_scenario_results_async(scenario_result_ids=[scenario_result_id]) + sdk = compute_scenario_statistics(scenario_result) + + assert sdk.overall.producer_counts.target_facing.attempts == 1 + assert sdk.overall.producer_counts.unknown.attempts == 2 + + # Verify backend SQL aggregation also treats both as unknown + aggregates = await sqlite_instance.get_scenario_history_aggregates_async(scenario_result_ids=[scenario_result_id]) + agg = aggregates[scenario_result_id] + assert agg.target_facing_attempts == 1 + assert agg.unknown_role_attempts == 2 + + +async def test_ambiguous_objective_with_identifier_recount_updates_producer_fields(sqlite_instance) -> None: + # Two planned seed groups share an objective. + # One target-facing attempt carries only its atomic identifier's seeds. + # SQL flags needs_sdk_statistics=True because objective matching in SQL is ambiguous. + # The fallback recount restores the unit and MUST copy producer fields. + planned_seed_a = _seed_group("Shared Objective", "context_a") + planned_seed_b = _seed_group("Shared Objective", "context_b") + history = _History( + plan=_plan( + _group(name="attack", eval_hash="eval", seed_ids=[planned_seed_a.logical_id, planned_seed_b.logical_id]), + seeds=[ + ScenarioRunPlanSeedGroup( + id=planned_seed_a.logical_id, + objective="Shared Objective", + objective_sha256=to_sha256("Shared Objective"), + ), + ScenarioRunPlanSeedGroup( + id=planned_seed_b.logical_id, + objective="Shared Objective", + objective_sha256=to_sha256("Shared Objective"), + ), + ], + ), + attempts=[ + _Attempt( + atomic_attack_name="attack", + objective="Shared Objective", + outcome=AttackOutcome.SUCCESS, + seed_context="context_a", + result_role=AttackResultRole.TARGET_FACING, + ) + ], + ) + scenario_result_id = await _persist(sqlite_instance, history) + + # Check SQL aggregate directly: needs_sdk_statistics is True + aggregates = await sqlite_instance.get_scenario_history_aggregates_async(scenario_result_ids=[scenario_result_id]) + assert aggregates[scenario_result_id].needs_sdk_statistics + + # SDK + [scenario_result] = await sqlite_instance.get_scenario_results_async(scenario_result_ids=[scenario_result_id]) + sdk = compute_scenario_statistics(scenario_result) + assert sdk.overall.completed == 1 + assert sdk.overall.producer_counts.target_facing.attempts == 1 + + # API: list_runs_async recounts flagged runs + service = ScenarioRunService() + runs = await service.list_runs_async() + [list_item] = [item for item in runs.items if item.scenario_result_id == scenario_result_id] + progress = await service.get_run_progress_from_storage_async( + scenario_result_id=scenario_result_id, since=None, limit=500, active_group_ids=[] + ) + + assert list_item.completed_attacks == 1 + assert list_item.producer_counts.target_facing.attempts == 1 + assert list_item.producer_counts == sdk.overall.producer_counts + assert progress.summary.overall.producer_counts == sdk.overall.producer_counts diff --git a/tests/unit/backend/test_scenario_progress_read_model.py b/tests/unit/backend/test_scenario_progress_read_model.py index f3f404de82..086b9aa069 100644 --- a/tests/unit/backend/test_scenario_progress_read_model.py +++ b/tests/unit/backend/test_scenario_progress_read_model.py @@ -715,6 +715,7 @@ async def test_roles_do_not_change_progress_counts_async(self, sqlite_instance: ) assert {result.result_role for result in legacy.results} == {AttackResultRole.UNKNOWN} + def _strip(obj): if isinstance(obj, dict): obj.pop("producer_counts", None) @@ -724,7 +725,7 @@ def _strip(obj): for item in obj: _strip(item) return obj - + legacy_dump = _strip(legacy.summary.model_dump()) snapshot_dump = _strip(snapshot.summary.model_dump()) assert legacy_dump == snapshot_dump diff --git a/tests/unit/memory/memory_interface/test_interface_seed_dataset_summary_query.py b/tests/unit/memory/memory_interface/test_interface_seed_dataset_summary_query.py index f22c680c33..88156d6ee5 100644 --- a/tests/unit/memory/memory_interface/test_interface_seed_dataset_summary_query.py +++ b/tests/unit/memory/memory_interface/test_interface_seed_dataset_summary_query.py @@ -28,18 +28,22 @@ async def test_get_seed_dataset_summaries_avoids_metadata_row_multiplication( def execute(session, statement, *args, **kwargs): result = original_execute(session, statement, *args, **kwargs) - + original_mappings = getattr(result, "mappings", None) if original_mappings: + def mappings_wrapper(): m_result = original_mappings() original_all = m_result.all + def all_rows(): rows = original_all() captured_rows.extend(rows) return rows + m_result.all = all_rows return m_result + result.mappings = mappings_wrapper original_all = result.all diff --git a/tests/unit/memory/test_azure_sql_memory.py b/tests/unit/memory/test_azure_sql_memory.py index 90d957e7c1..a240738216 100644 --- a/tests/unit/memory/test_azure_sql_memory.py +++ b/tests/unit/memory/test_azure_sql_memory.py @@ -926,9 +926,35 @@ def test_init_prod_with_skip_schema_migration_still_checks(): def test_scenario_history_aggregate_uses_sql_server_json_value(memory_interface: AzureSQLMemory) -> None: + from sqlalchemy.dialects import mssql + statement = memory_interface._build_scenario_history_aggregate_statement( entry_ids=[uuid.uuid4()], plan_entry_ids=[uuid.uuid4()] ) compiled = statement.compile(dialect=memory_interface.engine.dialect) - assert "json_value" in str(compiled).lower() + compiled_str = str(compiled).lower() + assert "json_value" in compiled_str assert "$.result_role" in compiled.params.values() + assert "latin1_general_100_bin2" in compiled_str + assert "datalength" in compiled_str + + mssql_compiled = str(statement.compile(dialect=mssql.dialect())).lower() + assert "collate latin1_general_100_bin2" in mssql_compiled + assert "datalength" in mssql_compiled + + +def test_scenario_role_match_condition_uses_binary_collation_and_datalength_guard( + memory_interface: AzureSQLMemory, +) -> None: + from sqlalchemy import column + from sqlalchemy.dialects import mssql + + from pyrit.models import AttackResultRole + + col = column("result_role") + condition = memory_interface._get_scenario_role_match_condition( + role_expression=col, role=AttackResultRole.TARGET_FACING + ) + compiled_mssql = str(condition.compile(dialect=mssql.dialect(), compile_kwargs={"literal_binds": True})) + assert "(result_role COLLATE Latin1_General_100_BIN2) = 'target_facing'" in compiled_mssql + assert "datalength(result_role) = datalength(N'target_facing')" in compiled_mssql diff --git a/tests/unit/output/scenario_result/test_html.py b/tests/unit/output/scenario_result/test_html.py index be3e97427d..ae4aecaa9a 100644 --- a/tests/unit/output/scenario_result/test_html.py +++ b/tests/unit/output/scenario_result/test_html.py @@ -241,3 +241,31 @@ async def test_html_report_renders_payload_from_real_builders(patch_central_data assert "chain-of-thought" in html # p.reasoning_summary assert "ContractScorer" in html # score.scorer assert "rationale-text" in html # score.score_rationale + assert "Target-facing attempts" in html + assert "Orchestration attempts" in html + assert "Target-facing" in html + assert "Orchestration" in html + + +async def test_html_report_renders_role_aware_accounting(): + target_facing_attack = AttackResult( + conversation_id=str(uuid.uuid4()), + objective="obj", + outcome=AttackOutcome.SUCCESS, + attribution_data={"result_role": "target_facing"}, + ) + orch_attack = AttackResult( + conversation_id=str(uuid.uuid4()), + objective="obj", + outcome=AttackOutcome.SUCCESS, + attribution_data={"result_role": "orchestration"}, + ) + result = make_scenario_result( + attack_results={"tech": [target_facing_attack, orch_attack]}, + ) + overview = JsonScenarioResultPrinter().build(result, view="overview") + payload = build_scenario_full_payload(result=result, overview=overview, entries=[]) + html = await HtmlScenarioReportPrinter().render_async(payload) + + assert "Target-facing attempts" in html + assert "Orchestration attempts" in html diff --git a/tests/unit/output/scenario_result/test_json.py b/tests/unit/output/scenario_result/test_json.py index a0db103d58..01385cca2a 100644 --- a/tests/unit/output/scenario_result/test_json.py +++ b/tests/unit/output/scenario_result/test_json.py @@ -233,6 +233,7 @@ async def test_overview_separates_units_from_attempts(printer): "num_attempts": 2, "success_rate": 100, "outcomes": payload["stats"]["outcomes"], + "producer_counts": payload["stats"]["producer_counts"], } ] assert payload["stats"]["outcomes"]["success_rate"] == 1.0 diff --git a/tests/unit/output/test_derivation.py b/tests/unit/output/test_derivation.py index 93acc0a4cb..2ff8d3f27b 100644 --- a/tests/unit/output/test_derivation.py +++ b/tests/unit/output/test_derivation.py @@ -14,6 +14,8 @@ AttackResult, ComponentIdentifier, MessagePiece, + ScenarioProducerCategoryCounts, + ScenarioProducerCounts, ScenarioRunPlan, ScenarioRunPlanAtomicGroup, ScenarioRunPlanSeedGroup, @@ -74,10 +76,16 @@ def test_scenario_overview_empty_is_zero(): assert (overview.objective_executions, overview.attempts, overview.success_rate) == (0, 0, 0) assert overview.groups == [ GroupStatistics( - name="s1", objective_executions=0, attempts=0, success_rate=0, outcomes=compute_outcome_statistics({}) + name="s1", + objective_executions=0, + attempts=0, + success_rate=0, + outcomes=compute_outcome_statistics({}), + producer_counts=ScenarioProducerCounts(), ) ] assert overview.outcomes == compute_outcome_statistics({}) + assert overview.producer_counts == ScenarioProducerCounts() def test_scenario_overview_folds_atomic_attacks_by_display_group(): @@ -103,8 +111,10 @@ def test_scenario_overview_folds_atomic_attacks_by_display_group(): attempts=3, success_rate=66, outcomes=compute_outcome_statistics({"success": 2, "failure": 1}), + producer_counts=ScenarioProducerCounts(unknown=ScenarioProducerCategoryCounts(attempts=3)), ) ] + assert overview.producer_counts == ScenarioProducerCounts(unknown=ScenarioProducerCategoryCounts(attempts=3)) def test_scenario_overview_uses_display_group_map_even_when_plan_labels_differ(): @@ -147,8 +157,10 @@ def test_scenario_overview_uses_display_group_map_even_when_plan_labels_differ() attempts=1, success_rate=100, outcomes=compute_outcome_statistics({"success": 1}), + producer_counts=ScenarioProducerCounts(unknown=ScenarioProducerCategoryCounts(attempts=1)), ) ] + assert overview.producer_counts == ScenarioProducerCounts(unknown=ScenarioProducerCategoryCounts(attempts=1)) # --- attack_score_display --- From c2468b5582f25acebfeec42880e670bef85ccfed Mon Sep 17 00:00:00 2001 From: Nimit Jain Date: Sun, 11 Oct 2026 15:47:09 +0530 Subject: [PATCH 8/8] fix: include unmatched producer attempts in per-group breakdowns (#3062) --- pyrit/analytics/scenario_statistics.py | 48 ++++- .../services/scenario_progress_read_model.py | 9 + .../test_scenario_statistics_parity.py | 201 +++++++++++++++++- 3 files changed, 255 insertions(+), 3 deletions(-) diff --git a/pyrit/analytics/scenario_statistics.py b/pyrit/analytics/scenario_statistics.py index 3f0dfecc7a..8f5cd23457 100644 --- a/pyrit/analytics/scenario_statistics.py +++ b/pyrit/analytics/scenario_statistics.py @@ -492,6 +492,38 @@ def compute_scenario_statistics( len(unit_attempts) for unit, unit_attempts in attempts_by_unit.items() if unit not in counted ) + producer_attempts_by_name: dict[str, list[ScenarioAttempt]] = {} + producer_attempts_by_display_group: dict[str, list[ScenarioAttempt]] = {} + + if plan is not None: + planned_groups_by_id = {group.id: group for group in plan.atomic_groups} + for attempt in attempts: + planned_group = planned_groups_by_id.get(attempt.unit.atomic_group_id) + if planned_group is not None: + producer_attempts_by_name.setdefault(planned_group.atomic_attack_name, []).append(attempt) + producer_attempts_by_display_group.setdefault(planned_group.display_group, []).append(attempt) + else: + matching_planned = plan_lookup.groups_by_name.get(attempt.atomic_attack_name, ()) + if len(matching_planned) == 0: + name = attempt.atomic_attack_name + display_group = scenario_result.display_group_map.get(name, name) + producer_attempts_by_name.setdefault(name, []).append(attempt) + producer_attempts_by_display_group.setdefault(display_group, []).append(attempt) + # If len(matching_planned) > 1: genuinely ambiguous attribution between planned groups. + # Keep genuinely ambiguous attribution explicitly unmatched rather than guessing. + else: + for attempt in attempts: + name = attempt.atomic_attack_name + display_group = scenario_result.display_group_map.get(name, name) + producer_attempts_by_name.setdefault(name, []).append(attempt) + producer_attempts_by_display_group.setdefault(display_group, []).append(attempt) + + attack_names = list(units_by_name.keys()) + attack_names.extend(name for name in producer_attempts_by_name if name not in units_by_name) + + display_group_names = list(units_by_display_group.keys()) + display_group_names.extend(dg for dg in producer_attempts_by_display_group if dg not in units_by_display_group) + def _count( units: Sequence[ScenarioExecutionUnit], *, @@ -506,8 +538,20 @@ def _count( return ScenarioExecutionStatistics( overall=_count(counted_units, producer_attempts=attempts), - atomic_attacks={name: _count(units) for name, units in units_by_name.items()}, - display_groups={name: _count(units) for name, units in units_by_display_group.items()}, + atomic_attacks={ + name: _count( + units_by_name.get(name, ()), + producer_attempts=producer_attempts_by_name.get(name, ()), + ) + for name in attack_names + }, + display_groups={ + name: _count( + units_by_display_group.get(name, ()), + producer_attempts=producer_attempts_by_display_group.get(name, ()), + ) + for name in display_group_names + }, attempts=len(attempts), unattributed_attempts=unattributed_attempts, ) diff --git a/pyrit/backend/services/scenario_progress_read_model.py b/pyrit/backend/services/scenario_progress_read_model.py index 63451fc4bb..c2a7ddff31 100644 --- a/pyrit/backend/services/scenario_progress_read_model.py +++ b/pyrit/backend/services/scenario_progress_read_model.py @@ -384,6 +384,10 @@ def aggregate( results=latest_results, ) + results_by_atomic_group_id: dict[str, list[ScenarioProgressResult]] = {} + for result in results: + results_by_atomic_group_id.setdefault(result.atomic_group_id, []).append(result) + active_ids = set(active_group_ids) atomic_groups: list[ScenarioAtomicGroupProgress] = [] for group in plan.atomic_groups: @@ -391,6 +395,7 @@ def aggregate( counts = aggregate( units=units, planned=len(units) if plan_complete else None, + producer_attempts=results_by_atomic_group_id.get(group.id, ()), ) if not terminal and group.id in active_ids: group_status: Literal["RUNNING", "PENDING", "INCOMPLETE", "COMPLETED"] = "RUNNING" @@ -432,6 +437,7 @@ def aggregate( counts = aggregate( units=units, planned=len(units) if plan_complete else None, + producer_attempts=[r for group in groups for r in results_by_atomic_group_id.get(group.id, ())], ) display_groups.append( ScenarioDisplayGroupProgress( @@ -450,6 +456,7 @@ def aggregate( counts = aggregate( units=units, planned=len(units) if plan_complete else None, + producer_attempts=[r for group in groups for r in results_by_atomic_group_id.get(group.id, ())], ) descriptions = list(dict.fromkeys(group.description for group in groups if group.description)) tags = sorted({tag for group in groups for tag in group.tags}) @@ -474,9 +481,11 @@ def aggregate( for group in plan.atomic_groups if seed_id in group.seed_group_ids ] + group_ids = {group.id for group in plan.atomic_groups if seed_id in group.seed_group_ids} counts = aggregate( units=units, planned=len(units) if plan_complete else None, + producer_attempts=[r for r in results if r.seed_group_id == seed_id and r.atomic_group_id in group_ids], ) seed_groups.append( ScenarioSeedGroupProgress( diff --git a/tests/unit/analytics/test_scenario_statistics_parity.py b/tests/unit/analytics/test_scenario_statistics_parity.py index 062c24a154..a93c3762f8 100644 --- a/tests/unit/analytics/test_scenario_statistics_parity.py +++ b/tests/unit/analytics/test_scenario_statistics_parity.py @@ -38,7 +38,12 @@ SeedObjective, SeedPrompt, ) -from pyrit.output.scenario_result.json import JsonScenarioResultPrinter +from pyrit.output._derivation import scenario_overview +from pyrit.output.scenario_result.html import HtmlScenarioReportPrinter +from pyrit.output.scenario_result.json import ( + JsonScenarioResultPrinter, + build_scenario_full_payload, +) from unit.mocks import make_scenario_result _T0 = datetime(2026, 9, 1, tzinfo=UTC) @@ -556,6 +561,200 @@ async def test_unmatched_attempts_visible_in_producer_accounting(sqlite_instance assert list_item.producer_counts == sdk.overall.producer_counts assert list_item.completed_attacks == 1 + # Per-group breakdown includes unmatched attempts in recognized group and wholly unplanned group + attack_atomic = sdk.atomic_attacks["attack"] + assert attack_atomic.completed == 1 + assert attack_atomic.planned == 1 + assert attack_atomic.producer_counts.target_facing.attempts == 2 + assert attack_atomic.producer_counts.target_facing.errors == 1 + assert attack_atomic.producer_counts.unknown.attempts == 1 + assert attack_atomic.producer_counts.unknown.errors == 1 + assert attack_atomic.producer_counts.orchestration.attempts == 0 + + orch_atomic = sdk.atomic_attacks["orchestrator"] + assert orch_atomic.completed == 0 + assert orch_atomic.planned == 0 + assert orch_atomic.producer_counts.orchestration.attempts == 1 + assert orch_atomic.producer_counts.orchestration.errors == 0 + + attack_display = sdk.display_groups["attack"] + assert attack_display.completed == 1 + assert attack_display.planned == 1 + assert attack_display.producer_counts.target_facing.attempts == 2 + assert attack_display.producer_counts.target_facing.errors == 1 + assert attack_display.producer_counts.unknown.attempts == 1 + assert attack_display.producer_counts.unknown.errors == 1 + + orch_display = sdk.display_groups["orchestrator"] + assert orch_display.completed == 0 + assert orch_display.planned == 0 + assert orch_display.producer_counts.orchestration.attempts == 1 + assert orch_display.producer_counts.orchestration.errors == 0 + + # Progress rollups assert group counts directly + [progress_atomic] = progress.summary.atomic_groups + assert progress_atomic.atomic_attack_name == "attack" + assert progress_atomic.completed == 1 + assert progress_atomic.planned == 1 + assert progress_atomic.producer_counts.target_facing.attempts == 2 + assert progress_atomic.producer_counts.target_facing.errors == 1 + assert progress_atomic.producer_counts.unknown.attempts == 1 + assert progress_atomic.producer_counts.unknown.errors == 1 + + [progress_display] = progress.summary.display_groups + assert progress_display.display_group == "attack" + assert progress_display.completed == 1 + assert progress_display.planned == 1 + assert progress_display.producer_counts.target_facing.attempts == 2 + assert progress_display.producer_counts.target_facing.errors == 1 + assert progress_display.producer_counts.unknown.attempts == 1 + assert progress_display.producer_counts.unknown.errors == 1 + + [progress_technique] = progress.summary.techniques + assert progress_technique.display_group == "attack" + assert progress_technique.completed == 1 + assert progress_technique.planned == 1 + assert progress_technique.producer_counts.target_facing.attempts == 2 + assert progress_technique.producer_counts.target_facing.errors == 1 + assert progress_technique.producer_counts.unknown.attempts == 1 + assert progress_technique.producer_counts.unknown.errors == 1 + + # Overview and JSON report assert group counts directly + overview = scenario_overview(scenario_result) + overview_groups = {g.name: g for g in overview.groups} + assert overview_groups["attack"].objective_executions == 1 + assert overview_groups["attack"].attempts == 3 + assert overview_groups["attack"].producer_counts.target_facing.attempts == 2 + assert overview_groups["attack"].producer_counts.target_facing.errors == 1 + assert overview_groups["attack"].producer_counts.unknown.attempts == 1 + assert overview_groups["attack"].producer_counts.unknown.errors == 1 + + assert overview_groups["orchestrator"].objective_executions == 0 + assert overview_groups["orchestrator"].attempts == 1 + assert overview_groups["orchestrator"].producer_counts.orchestration.attempts == 1 + assert overview_groups["orchestrator"].producer_counts.orchestration.errors == 0 + + report = json.loads(await JsonScenarioResultPrinter().render_async(scenario_result)) + report_groups = {g["name"]: g for g in report["groups"]} + assert report_groups["attack"]["num_objective_executions"] == 1 + assert report_groups["attack"]["num_attempts"] == 3 + assert report_groups["attack"]["producer_counts"]["target_facing"]["attempts"] == 2 + assert report_groups["attack"]["producer_counts"]["target_facing"]["errors"] == 1 + assert report_groups["attack"]["producer_counts"]["unknown"]["attempts"] == 1 + assert report_groups["attack"]["producer_counts"]["unknown"]["errors"] == 1 + + assert report_groups["orchestrator"]["num_objective_executions"] == 0 + assert report_groups["orchestrator"]["num_attempts"] == 1 + assert report_groups["orchestrator"]["producer_counts"]["orchestration"]["attempts"] == 1 + assert report_groups["orchestrator"]["producer_counts"]["orchestration"]["errors"] == 0 + + +async def test_unmatched_producer_attempts_in_per_group_breakdown(sqlite_instance) -> None: + # Regression: 1 planned group named "attack", 2 target-facing results with that same group/configuration: + # a planned success and an error whose seed ID is outside the saved plan. + plan = _plan( + _group(name="attack", eval_hash="eval", seed_ids=["seed_1"]), + seeds=[_seed("seed_1", "Objective 1")], + ) + history = _History( + plan=plan, + attempts=[ + _Attempt( + atomic_attack_name="attack", + objective="Objective 1", + outcome=AttackOutcome.SUCCESS, + seed_group_id="seed_1", + result_role=AttackResultRole.TARGET_FACING, + ), + _Attempt( + atomic_attack_name="attack", + objective="Unmatched Objective 2", + outcome=AttackOutcome.ERROR, + seed_group_id="unmatched_seed", + result_role=AttackResultRole.TARGET_FACING, + ), + ], + ) + scenario_result_id = await _persist(sqlite_instance, history) + + [scenario_result] = await sqlite_instance.get_scenario_results_async(scenario_result_ids=[scenario_result_id]) + sdk = compute_scenario_statistics(scenario_result) + + # Overall: 2 attempts, 1 error, 1 completed/planned logical unit + assert sdk.overall.completed == 1 + assert sdk.overall.planned == 1 + assert sdk.overall.succeeded == 1 + assert sdk.overall.producer_counts.target_facing.attempts == 2 + assert sdk.overall.producer_counts.target_facing.errors == 1 + + # SDK atomic attacks and display groups assert group counts directly + assert "attack" in sdk.atomic_attacks + attack_atomic = sdk.atomic_attacks["attack"] + assert attack_atomic.completed == 1 + assert attack_atomic.planned == 1 + assert attack_atomic.succeeded == 1 + assert attack_atomic.producer_counts.target_facing.attempts == 2 + assert attack_atomic.producer_counts.target_facing.errors == 1 + + assert "attack" in sdk.display_groups + attack_display = sdk.display_groups["attack"] + assert attack_display.completed == 1 + assert attack_display.planned == 1 + assert attack_display.succeeded == 1 + assert attack_display.producer_counts.target_facing.attempts == 2 + assert attack_display.producer_counts.target_facing.errors == 1 + + # ScenarioProgressReadModel: atomic, display, and technique rollups assert group counts directly + service = ScenarioRunService() + progress = await service.get_run_progress_from_storage_async( + scenario_result_id=scenario_result_id, since=None, limit=500, active_group_ids=[] + ) + assert progress is not None + + [atomic_group] = progress.summary.atomic_groups + assert atomic_group.atomic_attack_name == "attack" + assert atomic_group.completed == 1 + assert atomic_group.planned == 1 + assert atomic_group.producer_counts.target_facing.attempts == 2 + assert atomic_group.producer_counts.target_facing.errors == 1 + + [display_group] = progress.summary.display_groups + assert display_group.display_group == "attack" + assert display_group.completed == 1 + assert display_group.planned == 1 + assert display_group.producer_counts.target_facing.attempts == 2 + assert display_group.producer_counts.target_facing.errors == 1 + + [technique] = progress.summary.techniques + assert technique.display_group == "attack" + assert technique.completed == 1 + assert technique.planned == 1 + assert technique.producer_counts.target_facing.attempts == 2 + assert technique.producer_counts.target_facing.errors == 1 + + # Report overview and JSON printer assert group counts directly + overview = scenario_overview(scenario_result) + [group_overview] = overview.groups + assert group_overview.name == "attack" + assert group_overview.objective_executions == 1 + assert group_overview.attempts == 2 + assert group_overview.producer_counts.target_facing.attempts == 2 + assert group_overview.producer_counts.target_facing.errors == 1 + + report = json.loads(await JsonScenarioResultPrinter().render_async(scenario_result)) + [group_report] = report["groups"] + assert group_report["name"] == "attack" + assert group_report["num_objective_executions"] == 1 + assert group_report["num_attempts"] == 2 + assert group_report["producer_counts"]["target_facing"]["attempts"] == 2 + assert group_report["producer_counts"]["target_facing"]["errors"] == 1 + + # HTML report displays 2 attempts and 2 target-facing attempts without 2-versus-1 mismatch + json_overview = JsonScenarioResultPrinter().build(scenario_result, view="overview") + payload = build_scenario_full_payload(result=scenario_result, overview=json_overview, entries=[]) + html = await HtmlScenarioReportPrinter().render_async(payload) + assert "attack1\n 2\n 2" in html + async def test_malformed_roles_fall_back_to_unknown_in_producer_accounting(sqlite_instance) -> None: # Verify that malformed role values (trailing spaces or uppercase) fall back to UNKNOWN