diff --git a/doc/code/analytics/0_attack_results.md b/doc/code/analytics/0_attack_results.md new file mode 100644 index 0000000000..d46e2c87a2 --- /dev/null +++ b/doc/code/analytics/0_attack_results.md @@ -0,0 +1,333 @@ +# 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. + +The shared `compute_outcome_statistics` calculator returns both denominator +policies together: + +| Field | Denominator | Meaning | +|---|---|---| +| `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_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. +Both success rates are calculated from counts, never by averaging subgroup rates. + +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%. 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 +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_decided, report.summary.success_rate_all) + + 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", + ) + ) +``` + +### 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_decided, 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. + +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 +`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_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 +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 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. +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. + +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 +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 0e067b4ab2..cf0b171dd7 100644 --- a/doc/code/framework.md +++ b/doc/code/framework.md @@ -333,9 +333,13 @@ 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. +- `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. ## Auth @@ -361,6 +365,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 6a8754afc7..aa28dca9fb 100644 --- a/doc/myst.yml +++ b/doc/myst.yml @@ -189,6 +189,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 57c78d1a01..f5113c3f66 100644 --- a/pyrit/analytics/__init__.py +++ b/pyrit/analytics/__init__.py @@ -9,21 +9,34 @@ 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.outcome_statistics import combine_outcome_statistics, compute_outcome_statistics from pyrit.analytics.result_analysis import ( AttackStats, analyze_results, 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 _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", + "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/_execution.py b/pyrit/analytics/_execution.py new file mode 100644 index 0000000000..378ef5dcdd --- /dev/null +++ b/pyrit/analytics/_execution.py @@ -0,0 +1,278 @@ +# 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, Coroutine + +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) + 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: + self._abandon(work=work, task=operation) + raise AnalyticsTimeoutException + return operation.result() + except asyncio.CancelledError: + self._abandon(work=work, task=operation) + 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. + 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: + 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._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.") + + 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 + 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: + 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: + 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(): + 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 + + @staticmethod + async def _execute_async(*, task: Callable[[QueryControl], Awaitable[T]], control: QueryControl) -> T: + control.check() + 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/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..22e7e53bb3 --- /dev/null +++ b/pyrit/analytics/attack_result_analytics.py @@ -0,0 +1,345 @@ +# 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.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 +from pyrit.models import ( + AttackAnalyticsCell, + AttackAnalyticsDimension, + AttackAnalyticsDimensionName, + AttackAnalyticsFacetQuery, + AttackAnalyticsFacets, + AttackAnalyticsFilter, + AttackAnalyticsFilters, + AttackAnalyticsGroup, + AttackAnalyticsOption, + AttackAnalyticsQuery, + AttackAnalyticsReport, + AttackAnalyticsResults, + AttackAnalyticsResultsQuery, + AttackAnalyticsStatistics, + AttackAnalyticsValue, + AttackAnalyticsValueKind, +) + +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: + """ + Apply shared outcome statistics, retaining the SDK's stored-data error contract. + + Returns: + AttackAnalyticsStatistics: Both success rates and whole-cohort outcome shares. + + Raises: + AnalyticsDataException: If saved outcomes or their counts are invalid. + """ + try: + return compute_outcome_statistics(counts) + except ValueError as error: + raise AnalyticsDataException(str(error)) from error + + @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/pyrit/analytics/outcome_statistics.py b/pyrit/analytics/outcome_statistics.py new file mode 100644 index 0000000000..ce45b0f1ee --- /dev/null +++ b/pyrit/analytics/outcome_statistics.py @@ -0,0 +1,124 @@ +# 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_decided`` (also ``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: + if isinstance(item, OutcomeStatistics): + item.validate_consistency() + 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..5356e579a9 100644 --- a/pyrit/analytics/result_analysis.py +++ b/pyrit/analytics/result_analysis.py @@ -3,8 +3,9 @@ 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 from pyrit.models import ( AttackOutcome, @@ -13,35 +14,81 @@ 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: - 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 _as_attack_stats(statistics) + + +def _as_attack_stats(statistics: OutcomeStatistics) -> AttackStats: return AttackStats( - success_rate=success_rate, - total_decided=total_decided, - successes=successes, - failures=failures, - undetermined=undetermined, - errors=errors, + success_rate=statistics.success_rate, + total_decided=statistics.total_decided, + 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. @@ -57,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): @@ -68,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 2b483ca6f1..3f0dfecc7a 100644 --- a/pyrit/analytics/scenario_statistics.py +++ b/pyrit/analytics/scenario_statistics.py @@ -11,30 +11,40 @@ - 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, AtomicAttackIdentifier, AttackOutcome, AttackResult, + AttackResultMetadata, + AttackResultRole, ComponentIdentifier, ScenarioExecutionStatistics, ScenarioExecutionUnit, + ScenarioProducerCategoryCounts, + ScenarioProducerCounts, ScenarioProgressCounts, ScenarioRunPlan, ScenarioRunPlanAtomicGroup, @@ -58,6 +68,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 +147,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 +254,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, ) @@ -255,15 +270,62 @@ 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: +def compute_producer_counts(attempts: Iterable[_CountableAttempt]) -> ScenarioProducerCounts: """ - Return the success percentage for effective execution units. + Compute role-aware attempt, error, and retry accounting across attempts. + + Args: + attempts (Iterable[_CountableAttempt]): Chronological or arbitrary attempt sequence. Returns: - int | None: ``succeeded / completed`` as a truncated integer percentage, or None with no - completed units. + ScenarioProducerCounts: Role-specific attempts, errors, and retries. """ - return int((succeeded / completed) * 100) if completed else None + 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 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, + 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, + ), + ) def count_execution_units( @@ -271,6 +333,7 @@ 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. @@ -279,30 +342,36 @@ 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 + unit_attempts: list[_CountableAttempt] = [] + 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], ) + 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=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, + producer_counts=producer_counts, + outcomes=outcomes, ) @@ -312,12 +381,39 @@ 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) + 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] + from pyrit.models import ScenarioProducerCategoryCounts, ScenarioProducerCounts + + 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), + ), + ) + + 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 +421,8 @@ 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, + outcomes=outcomes, ) @@ -394,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/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/backend/services/scenario_progress_read_model.py b/pyrit/backend/services/scenario_progress_read_model.py index da2bdc9460..63451fc4bb 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: @@ -343,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: [ @@ -359,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 89e151be6f..c443f670d6 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 @@ -1436,12 +1437,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 @@ -1535,6 +1538,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 @@ -1617,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 @@ -1687,6 +1700,25 @@ def _build_history_summary( if atomic_groups is not None else list(aggregate.atomic_attack_names) ) + from pyrit.models import ScenarioProducerCategoryCounts, ScenarioProducerCounts + + 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, @@ -1701,7 +1733,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, @@ -1713,6 +1745,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/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 ca0997e73d..738ef53e1f 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, ) @@ -2148,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. @@ -4337,26 +4375,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 +4550,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 +5049,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 +6009,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 +6017,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] ], @@ -6076,9 +6114,13 @@ def _execute_get_scenario_history_aggregates( if scenario_result_id in 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() + aggregate_rows = ( + session.execute( + self._build_scenario_history_aggregate_statement(entry_ids=entry_ids, plan_entry_ids=plan_entry_ids) + ) + .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 +6146,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 +6205,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 +6228,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, @@ -6187,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, @@ -6200,6 +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_(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_(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_(is_unknown, ranked.c.latest_outcome == AttackOutcome.ERROR.value), 1), + else_=0, + ) + ).label("unknown_role_error_attempts"), + 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) @@ -6226,6 +6314,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 +6394,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 +6420,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 +6534,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..2552659de8 100644 --- a/pyrit/models/__init__.py +++ b/pyrit/models/__init__.py @@ -44,11 +44,14 @@ AttackAnalyticsValueKind, AttackResultSelection, AttackStats, + OutcomeStatistics, ) from pyrit.models.catalog import ( ScenarioDatasetSizeCap, ScenarioDatasetSummary, ScenarioDefaultRunSizeEstimate, + ScenarioProducerCategoryCounts, + ScenarioProducerCounts, ScenarioRunListItem, ScenarioRunSizeComponent, ScenarioRunSizeEstimate, @@ -293,6 +296,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", @@ -398,6 +402,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/analytics.py b/pyrit/models/analytics.py index 53cda714a1..55b1cc0c05 100644 --- a/pyrit/models/analytics.py +++ b/pyrit/models/analytics.py @@ -5,19 +5,24 @@ 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 +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,21 +31,140 @@ 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): """ - Outcome statistics for one cohort, group, or heatmap cell. - - ``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. + Shared outcome statistics for saved results or selected scenario execution units. + + ``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. + + 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: _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. +OutcomeStatistics = AttackAnalyticsStatistics class AttackResultSelection(str, Enum): 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 9a9f269a25..5fed818bc4 100644 --- a/pyrit/models/catalog/scenario.py +++ b/pyrit/models/catalog/scenario.py @@ -166,6 +166,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.""" @@ -586,6 +602,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): @@ -628,6 +645,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..8bda2544da 100644 --- a/pyrit/models/scenario_progress.py +++ b/pyrit/models/scenario_progress.py @@ -5,11 +5,16 @@ 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 -from pyrit.models.catalog.scenario import ScenarioOverloadSummary, ScenarioTargetSummary # noqa: TC001 +from pyrit.models.analytics import OutcomeStatistics +from pyrit.models.catalog.scenario import ( + ScenarioOverloadSummary, + 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 @@ -183,7 +188,18 @@ 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. + 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) @@ -191,6 +207,34 @@ 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) + 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/_derivation.py b/pyrit/output/_derivation.py index bd03cb5321..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, ScenarioResult, Score + from pyrit.models import ( + AttackResult, + ComponentIdentifier, + MessagePiece, + OutcomeStatistics, + ScenarioProducerCounts, + ScenarioResult, + Score, + ) class TargetInfo(NamedTuple): @@ -61,6 +69,8 @@ class GroupStatistics(NamedTuple): objective_executions: int attempts: int success_rate: int + outcomes: OutcomeStatistics | None = None + producer_counts: ScenarioProducerCounts | None = None class ScenarioOverview(NamedTuple): @@ -70,6 +80,8 @@ class ScenarioOverview(NamedTuple): attempts: int success_rate: int groups: list[GroupStatistics] + outcomes: OutcomeStatistics | None = None + producer_counts: ScenarioProducerCounts | None = None def scenario_overview(result: ScenarioResult) -> ScenarioOverview: @@ -100,6 +112,8 @@ def scenario_overview(result: ScenarioResult) -> ScenarioOverview: objective_executions=counts.completed, attempts=len(group_results), success_rate=counts.success_percentage or 0, + outcomes=counts.outcomes, + producer_counts=counts.producer_counts, ) ) return ScenarioOverview( @@ -107,6 +121,8 @@ def scenario_overview(result: ScenarioResult) -> ScenarioOverview: attempts=statistics.attempts, 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 @@
| Group | Objective executions | Attempts | Success rate | ||
|---|---|---|---|---|---|
| Group | +Objective executions | +Attempts | +Target-facing | +Orchestration | +Success 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 }}% |