diff --git a/pyrit/memory/memory_interface.py b/pyrit/memory/memory_interface.py index 098d7d262b..4bbef89456 100644 --- a/pyrit/memory/memory_interface.py +++ b/pyrit/memory/memory_interface.py @@ -7,9 +7,11 @@ import re import uuid import weakref -from collections.abc import Iterator, MutableSequence, Sequence +from collections.abc import Iterator, Mapping, MutableSequence, Sequence from contextlib import closing +from dataclasses import dataclass from datetime import datetime, timezone +from types import MappingProxyType from typing import TYPE_CHECKING, Any, ClassVar, Literal, NamedTuple, TypeVar from sqlalchemy import MetaData, and_, func, not_, or_, select @@ -119,6 +121,56 @@ def from_attack_result(cls, result: AttackResult) -> "AttackResultsKeysetCursor" ) +@dataclass(frozen=True, slots=True, kw_only=True) +class _AttackResultQuery: + """ + Immutable filters and pagination settings for an attack-result query. + + Sequence and mapping inputs are defensively copied into immutable containers so a + query cannot change while its database conditions are being assembled or executed. + """ + + _SEQUENCE_FIELDS: ClassVar[tuple[str, ...]] = ( + "attack_result_ids", + "objective_sha256", + "attack_classes", + "atomic_attack_eval_hashes", + "converter_classes", + "targeted_harm_categories", + "identifier_filters", + ) + + attack_result_ids: Sequence[str] | None = None + conversation_id: str | None = None + objective: str | None = None + objective_sha256: Sequence[str] | None = None + outcome: str | None = None + attack_classes: Sequence[str] | None = None + atomic_attack_eval_hashes: Sequence[str] | None = None + converter_classes: Sequence[str] | None = None + converter_classes_match: Literal["all", "any"] = "all" + has_converters: bool | None = None + labels: Mapping[str, str | Sequence[str]] | None = None + targeted_harm_categories: Sequence[str] | None = None + identifier_filters: Sequence[IdentifierFilter] | None = None + scenario_result_id: str | None = None + min_turns: int | None = None + max_turns: int | None = None + limit: int | None = None + after: AttackResultsKeysetCursor | None = None + + def __post_init__(self) -> None: + """Snapshot mutable sequence and mapping inputs.""" + for field_name in self._SEQUENCE_FIELDS: + value = getattr(self, field_name) + if value is not None: + object.__setattr__(self, field_name, tuple(value)) + + if self.labels is not None: + labels = {key: value if isinstance(value, str) else tuple(value) for key, value in self.labels.items()} + object.__setattr__(self, "labels", MappingProxyType(labels)) + + class MemoryInterface(abc.ABC): """ Abstract interface for conversation memory storage systems. @@ -2806,160 +2858,268 @@ def get_attack_results( ValueError: If ``limit`` or ``after`` is combined with ``attack_result_ids`` or ``objective_sha256`` (id-batched lookups do not support SQL pagination). """ - # Handle empty list cases - if attack_result_ids is not None and len(attack_result_ids) == 0: - return [] - if objective_sha256 is not None and len(objective_sha256) == 0: + query = _AttackResultQuery( + attack_result_ids=attack_result_ids, + conversation_id=conversation_id, + objective=objective, + objective_sha256=objective_sha256, + outcome=outcome, + attack_classes=attack_classes, + atomic_attack_eval_hashes=atomic_attack_eval_hashes, + converter_classes=converter_classes, + converter_classes_match=converter_classes_match, + has_converters=has_converters, + labels=labels, + targeted_harm_categories=targeted_harm_categories, + identifier_filters=identifier_filters, + scenario_result_id=scenario_result_id, + min_turns=min_turns, + max_turns=max_turns, + limit=limit, + after=after, + ) + return self._query_attack_results(query=query) + + def _query_attack_results(self, *, query: _AttackResultQuery) -> Sequence[AttackResult]: + """ + Retrieve attack results matching an immutable query. + + Args: + query (_AttackResultQuery): Filters and pagination settings to apply. + + Returns: + Sequence[AttackResult]: Attack results matching the query. + + Raises: + ValueError: If the query contains invalid label keys or combines pagination + with an ID-batched lookup. + """ + if self._attack_result_query_has_empty_lookup(query=query): return [] - # Build non-list conditions - conditions: list[ColumnElement[bool]] = [] - if conversation_id: - conditions.append(AttackResultEntry.conversation_id == conversation_id) - if objective: - conditions.append(AttackResultEntry.objective.contains(objective)) - if outcome: - conditions.append(AttackResultEntry.outcome == outcome) - if scenario_result_id: - conditions.append(AttackResultEntry.attribution_parent_id == uuid.UUID(scenario_result_id)) - - if attack_classes: - # Case-insensitive to mirror converter_classes; forgives casing drift in - # REST/CLI callers. PyRIT class names are PascalCase with no case-variant - # collisions so this never changes match results for well-formed inputs. + conditions = self._build_attack_result_conditions(query=query) + paginating = query.limit is not None or query.after is not None + self._validate_attack_result_query_pagination(query=query, paginating=paginating) + try: + if paginating: + return self._query_paginated_attack_results( + conditions=conditions, + min_turns=query.min_turns, + max_turns=query.max_turns, + limit=query.limit, + after=query.after, + ) + + entries = self._query_with_list_params( + AttackResultEntry, conditions=conditions, list_params=self._build_attack_result_list_params(query=query) + ) + results = self._dedup_attack_entries(entries) + return self._filter_attack_results_by_turns( + results, + min_turns=query.min_turns, + max_turns=query.max_turns, + ) + except Exception as e: + logger.exception(f"Failed to retrieve attack results with error {e}") + raise + + @staticmethod + def _attack_result_query_has_empty_lookup(*, query: _AttackResultQuery) -> bool: + """Return whether an explicitly empty ID lookup must produce no results.""" + return ( + query.attack_result_ids is not None + and len(query.attack_result_ids) == 0 + or query.objective_sha256 is not None + and len(query.objective_sha256) == 0 + ) + + def _build_attack_result_conditions(self, *, query: _AttackResultQuery) -> list[Any]: + """ + Build backend-neutral and backend-specific SQL conditions for a query. + + Returns: + list[Any]: SQLAlchemy conditions for the query. + """ + conditions = self._build_attack_result_scalar_conditions(query=query) + conditions.extend(self._build_attack_result_identifier_conditions(query=query)) + conditions.extend(self._build_attack_result_converter_conditions(query=query)) + conditions.extend(self._build_attack_result_label_conditions(query=query)) + conditions.extend(self._build_attack_result_category_conditions(query=query)) + conditions.extend(self._build_attack_result_generic_identifier_conditions(query=query)) + return conditions + + @staticmethod + def _build_attack_result_scalar_conditions(*, query: _AttackResultQuery) -> list[Any]: + """ + Build conditions for scalar attack-result columns. + + Returns: + list[Any]: SQLAlchemy conditions for populated scalar filters. + """ + conditions: list[Any] = [] + if query.conversation_id: + conditions.append(AttackResultEntry.conversation_id == query.conversation_id) + if query.objective: + conditions.append(AttackResultEntry.objective.contains(query.objective)) + if query.outcome: + conditions.append(AttackResultEntry.outcome == query.outcome) + if query.scenario_result_id: + conditions.append(AttackResultEntry.attribution_parent_id == uuid.UUID(query.scenario_result_id)) + return conditions + + def _build_attack_result_identifier_conditions(self, *, query: _AttackResultQuery) -> list[Any]: + """ + Build conditions for attack identifier JSON properties. + + Returns: + list[Any]: SQLAlchemy conditions for identifier filters. + """ + conditions: list[Any] = [] + if query.attack_classes: conditions.append( or_( *[ self._get_condition_json_property_match( json_column=AttackResultEntry.atomic_attack_identifier, property_path="$.children.attack_technique.children.attack.class_name", - value=ac, + value=attack_class, ) - for ac in attack_classes + for attack_class in query.attack_classes ] ) ) - - if atomic_attack_eval_hashes: - # Single JSON path query on the auto-stamped eval_hash. OR-combined across - # supplied hashes so callers can fetch history for multiple technique - # configurations in one round trip. + if query.atomic_attack_eval_hashes: conditions.append( or_( *[ self._get_condition_json_property_match( json_column=AttackResultEntry.atomic_attack_identifier, property_path="$.eval_hash", - value=h, + value=eval_hash, case_sensitive=True, ) - for h in atomic_attack_eval_hashes + for eval_hash in query.atomic_attack_eval_hashes ] ) ) + return conditions - if converter_classes is not None: - # Non-empty sequence: filter to attacks that used ALL (or ANY, depending on - # converter_classes_match) of the listed converters. - # Empty sequence: filter to attacks that used NO converters. - # None: no filter. + def _build_attack_result_converter_conditions(self, *, query: _AttackResultQuery) -> list[Any]: + """ + Build conditions for converter class names and converter presence. + + Returns: + list[Any]: SQLAlchemy conditions for converter filters. + """ + conditions: list[Any] = [] + if query.converter_classes is not None: conditions.append( self._get_condition_json_array_match( json_column=AttackResultEntry.atomic_attack_identifier, property_path="$.children.attack_technique.children.attack.children.request_converters", array_element_path="$.class_name", - array_to_match=converter_classes, - match_mode=converter_classes_match, + array_to_match=query.converter_classes, + match_mode=query.converter_classes_match, ) ) - - # Skip when has_converters=True and converter_classes is already non-empty: - # the "ALL listed converters" constraint strictly implies "at least one - # converter", so adding this predicate is redundant work. - if has_converters is not None and not (has_converters is True and converter_classes): - # Reuse the array-empty match (array_to_match=[]) as the "no converters" - # condition; invert it for "has at least one converter". + if query.has_converters is not None and not (query.has_converters is True and query.converter_classes): empty_condition = self._get_condition_json_array_match( json_column=AttackResultEntry.atomic_attack_identifier, property_path="$.children.attack_technique.children.attack.children.request_converters", array_element_path="$.class_name", array_to_match=[], ) - conditions.append(not_(empty_condition) if has_converters else empty_condition) + conditions.append(not_(empty_condition) if query.has_converters else empty_condition) + return conditions - if labels: - # Strip keys whose value is an empty sequence — an empty sequence means - # "no OR-candidates", and per the docstring applies no filter for that - # key. Without this, the per-backend helpers would still emit a base - # EXISTS(... labels IS NOT NULL) predicate that is strictly more - # restrictive than "no filter". - effective_labels = {k: v for k, v in labels.items() if not (isinstance(v, (list, tuple)) and len(v) == 0)} - # Validate label keys against an allowlist: backend helpers - # interpolate keys into JSON path expressions (e.g. ``$.key``), - # so a key with quotes or SQL punctuation could otherwise break - # out and inject SQL. - invalid_keys = [k for k in effective_labels if not self._LABEL_KEY_PATTERN.match(k)] - if invalid_keys: - raise ValueError( - f"Invalid label key(s) {invalid_keys!r}: keys must match {self._LABEL_KEY_PATTERN.pattern}." - ) - if effective_labels: - # Use database-specific JSON query method - conditions.append(self._get_attack_result_label_condition(labels=effective_labels)) + def _build_attack_result_label_conditions(self, *, query: _AttackResultQuery) -> list[Any]: + """ + Build a validated backend-specific attack-label condition. - if targeted_harm_categories: - # Match attacks whose targeted_harm_categories array contains ANY of the - # requested categories. - conditions.append( - self._get_condition_json_array_match( - json_column=AttackResultEntry.targeted_harm_categories, - property_path="$", - array_to_match=list(targeted_harm_categories), - match_mode="any", - ) + Returns: + list[Any]: The backend-specific label condition, if labels are effective. + + Raises: + ValueError: If a label key falls outside the safe allowlist. + """ + if not query.labels: + return [] + effective_labels = { + key: value + for key, value in query.labels.items() + if not (isinstance(value, (list, tuple)) and len(value) == 0) + } + invalid_keys = [key for key in effective_labels if not self._LABEL_KEY_PATTERN.match(key)] + if invalid_keys: + raise ValueError( + f"Invalid label key(s) {invalid_keys!r}: keys must match {self._LABEL_KEY_PATTERN.pattern}." ) + if not effective_labels: + return [] + return [self._get_attack_result_label_condition(labels=effective_labels)] - if identifier_filters: - conditions.extend( - self._build_identifier_filter_conditions( - identifier_filters=identifier_filters, - identifier_column_map={IdentifierType.ATTACK: AttackResultEntry.atomic_attack_identifier}, - caller="get_attack_results", - ) + def _build_attack_result_category_conditions(self, *, query: _AttackResultQuery) -> list[Any]: + """ + Build the targeted-harm-category condition. + + Returns: + list[Any]: The backend-specific category condition, if categories are supplied. + """ + if not query.targeted_harm_categories: + return [] + return [ + self._get_condition_json_array_match( + json_column=AttackResultEntry.targeted_harm_categories, + property_path="$", + array_to_match=query.targeted_harm_categories, + match_mode="any", ) + ] - paginating = limit is not None or after is not None - if paginating and (attack_result_ids or objective_sha256): + def _build_attack_result_generic_identifier_conditions(self, *, query: _AttackResultQuery) -> list[Any]: + """ + Build generic attack identifier conditions. + + Returns: + list[Any]: SQLAlchemy conditions for generic identifier filters. + """ + if not query.identifier_filters: + return [] + return self._build_identifier_filter_conditions( + identifier_filters=query.identifier_filters, + identifier_column_map={IdentifierType.ATTACK: AttackResultEntry.atomic_attack_identifier}, + caller="get_attack_results", + ) + + @staticmethod + def _validate_attack_result_query_pagination(*, query: _AttackResultQuery, paginating: bool) -> None: + """ + Reject pagination combined with unsupported ID-batched lookups. + + Raises: + ValueError: If pagination is combined with an ID-batched lookup. + """ + if paginating and (query.attack_result_ids or query.objective_sha256): raise ValueError( "limit/keyset pagination cannot be combined with attack_result_ids or objective_sha256 lookups." ) - try: - if paginating: - return self._query_paginated_attack_results( - conditions=conditions, - min_turns=min_turns, - max_turns=max_turns, - limit=limit, - after=after, - ) - - list_params: list[tuple[InstrumentedAttribute[Any], Sequence[Any], str]] = [] - if attack_result_ids: - list_params.append((AttackResultEntry.id, list(attack_result_ids), "id")) - if objective_sha256: - list_params.append((AttackResultEntry.objective_sha256, list(objective_sha256), "objective_sha256")) + @staticmethod + def _build_attack_result_list_params( + *, query: _AttackResultQuery + ) -> list[tuple[InstrumentedAttribute[Any], Sequence[Any], str]]: + """ + Build batched list lookup descriptors for the unpaginated query path. - entries = self._query_with_list_params( - AttackResultEntry, - conditions=conditions, - list_params=list_params, - ) - results = self._dedup_attack_entries(entries) - return self._filter_attack_results_by_turns(results, min_turns=min_turns, max_turns=max_turns) - except Exception as e: - logger.exception(f"Failed to retrieve attack results with error {e}") - raise + Returns: + list[tuple[InstrumentedAttribute[Any], Sequence[Any], str]]: Batched lookup descriptors. + """ + list_params: list[tuple[InstrumentedAttribute[Any], Sequence[Any], str]] = [] + if query.attack_result_ids: + list_params.append((AttackResultEntry.id, query.attack_result_ids, "id")) + if query.objective_sha256: + list_params.append((AttackResultEntry.objective_sha256, query.objective_sha256, "objective_sha256")) + return list_params def _query_paginated_attack_results( self, diff --git a/tests/unit/memory/memory_interface/test_interface_attack_results.py b/tests/unit/memory/memory_interface/test_interface_attack_results.py index 83b2e0c6a9..e532a1fe54 100644 --- a/tests/unit/memory/memory_interface/test_interface_attack_results.py +++ b/tests/unit/memory/memory_interface/test_interface_attack_results.py @@ -3,13 +3,16 @@ import uuid +from dataclasses import FrozenInstanceError from datetime import datetime, timedelta, timezone from typing import TYPE_CHECKING +from unittest.mock import patch import pytest from pyrit.common.utils import to_sha256 from pyrit.memory import AttackResultsKeysetCursor, MemoryInterface +from pyrit.memory.memory_interface import _AttackResultQuery from pyrit.memory.memory_models import AttackResultEntry from pyrit.models import ( AtomicAttackIdentifier, @@ -100,6 +103,102 @@ def _drain_keyset(memory: MemoryInterface, *, page_size: int, **filters) -> list return drained +def test_attack_result_query_snapshots_mutable_inputs(): + """The internal query remains stable when caller-owned containers change.""" + attack_classes = ["CrescendoAttack"] + labels = {"operator": ["alice"]} + query = _AttackResultQuery(attack_classes=attack_classes, labels=labels) + + attack_classes.append("ManualAttack") + labels["operator"].append("bob") + + assert query.attack_classes == ("CrescendoAttack",) + assert query.labels == {"operator": ("alice",)} + field_name = "limit" + with pytest.raises(FrozenInstanceError): + setattr(query, field_name, 10) + + +def test_attack_result_query_requires_keyword_arguments(): + """The internal query does not expose field ordering as a positional API.""" + with pytest.raises(TypeError): + _AttackResultQuery(["id"]) # type: ignore[misc] + + +def test_get_attack_results_forwards_all_parameters_to_query(sqlite_instance: MemoryInterface): + """The compatibility API maps every parameter onto the internal query.""" + cursor = AttackResultsKeysetCursor(timestamp=_BASE_TS, attack_result_id=str(uuid.uuid4())) + identifier_filter = IdentifierFilter( + identifier_type=IdentifierType.ATTACK, + property_path="$.hash", + value="hash", + ) + + with patch.object(sqlite_instance, "_query_attack_results", return_value=[]) as query_mock: + sqlite_instance.get_attack_results( + attack_result_ids=["id"], + conversation_id="conversation", + objective="objective", + objective_sha256=["sha"], + outcome="success", + attack_classes=["Attack"], + atomic_attack_eval_hashes=["eval"], + converter_classes=["Converter"], + converter_classes_match="any", + has_converters=True, + labels={"operator": ["alice"]}, + targeted_harm_categories=["violence"], + identifier_filters=[identifier_filter], + scenario_result_id=str(uuid.uuid4()), + min_turns=1, + max_turns=5, + limit=10, + after=cursor, + ) + + query = query_mock.call_args.kwargs["query"] + assert isinstance(query, _AttackResultQuery) + assert query.attack_result_ids == ("id",) + assert query.conversation_id == "conversation" + assert query.objective == "objective" + assert query.objective_sha256 == ("sha",) + assert query.outcome == "success" + assert query.attack_classes == ("Attack",) + assert query.atomic_attack_eval_hashes == ("eval",) + assert query.converter_classes == ("Converter",) + assert query.converter_classes_match == "any" + assert query.has_converters is True + assert query.labels == {"operator": ("alice",)} + assert query.targeted_harm_categories == ("violence",) + assert query.identifier_filters == (identifier_filter,) + assert query.scenario_result_id is not None + assert query.min_turns == 1 + assert query.max_turns == 5 + assert query.limit == 10 + assert query.after == cursor + + +def test_query_attack_results_matches_compatibility_api(sqlite_instance: MemoryInterface): + """The internal query executor and compatibility API return the same page.""" + sqlite_instance.add_attack_results_to_memory( + attack_results=[ + _make_attack_result("conv-1", executed_turns=2, outcome=AttackOutcome.SUCCESS, ts_offset=1), + _make_attack_result("conv-2", executed_turns=5, outcome=AttackOutcome.SUCCESS, ts_offset=2), + _make_attack_result("conv-3", executed_turns=7, outcome=AttackOutcome.FAILURE, ts_offset=3), + ] + ) + + query = _AttackResultQuery(outcome=AttackOutcome.SUCCESS.value, min_turns=3, limit=1) + direct = sqlite_instance._query_attack_results(query=query) + compatibility = sqlite_instance.get_attack_results( + outcome=AttackOutcome.SUCCESS.value, + min_turns=3, + limit=1, + ) + + assert [result.attack_result_id for result in direct] == [result.attack_result_id for result in compatibility] + + def test_add_attack_results_to_memory(sqlite_instance: MemoryInterface): """Test adding attack results to memory.""" # Create sample attack results