diff --git a/doc/code/datasets/0_dataset.md b/doc/code/datasets/0_dataset.md index 008ed96dce..2384e9f5e4 100644 --- a/doc/code/datasets/0_dataset.md +++ b/doc/code/datasets/0_dataset.md @@ -93,3 +93,44 @@ an explicit origin other than `local`. Remote dataset providers assign use `GENERATED`. Origin does not describe upstream authorship. Use `origin=SeedOrigin.USER` for explicit user entries. Unspecified and legacy origins remain `UNKNOWN`, while edits preserve the recorded origin. + +## Browse stored seeds + +The seed browser reads stored seeds from memory only. It does not load providers, open +media or template files, render templates, or generate conversations. Members are +`SeedRecord` projections, not reconstructed execution-ready seeds. Simulated-conversation +configurations remain unchanged in `value`, including legacy file references and +configurations that cannot be executed. Missing files do not remove members or examples. +Stored IDs, hashes, nullable roles and sequences, parameters, objective conditions, and +provenance are retained. Existing `get_seeds_async()` reconstruction is unchanged. + +- `GET /api/datasets/seeds?selection_key=` lists one page of logical examples. +- `GET /api/datasets/seeds/{example_id}?selection_key=` returns all members of one example + as stored records, identified by `seed_type`. Configuration fields such as `num_turns` + remain in the stored JSON `value`; browsing does not resolve them into live seed objects. + +Get the `selection_key` from `GET /api/datasets`. The unnamed key `dataset:unnamed` includes +NULL and empty dataset names. The example ID is the `prompt_group_id`, or the seed ID when +the seed has no group. Only members in the selected dataset are returned. + +The list accepts `limit` (1 to 100), `cursor`, `search`, and repeated `modality`, +`seed_type`, and `harm_category` parameters. Values of one parameter use OR. Different +parameters use AND, and different members of an example can match different parameters. +Harm categories match complete values without case sensitivity. `search` finds literal +text in the values of text prompts and objectives; `%`, `_`, and `[` are not patterns. +SQLite ignores case for ASCII characters only. `search` does not look in +simulated-conversation configurations, because their stored value is JSON. Use +`seed_type=simulated_conversation` to find them. + +Examples sort by the earliest member `date_added`, newest first, then by canonical textual +example ID, descending. SQLite and Azure SQL use the same UUID order. A +cursor is valid only for the same `selection_key` and filters. Other cursors return 400. + +Each list item has a preview of the first text member: at most 100 characters, with `...` +and `preview_truncated` when it is shortened. Media members show only the file name. +HTTP(S) media URLs are recognized regardless of scheme case or leading whitespace; +their authority, query, and fragment are excluded from the preview. +Standalone absolute paths and URLs stored as text show `[Text reference]` rather than +paths or credentials; detail retains the full stored value. Other types show a type label. +The browser does not render templates or run +simulated conversations. diff --git a/pyrit/backend/mappers/_preview.py b/pyrit/backend/mappers/_preview.py index 83a94d648e..afbaf0dc99 100644 --- a/pyrit/backend/mappers/_preview.py +++ b/pyrit/backend/mappers/_preview.py @@ -42,11 +42,12 @@ def _derive_basename(value: str) -> str | None: The basename (filename portion) of *value*, or ``None`` if one can't be derived (e.g. data URI, empty value). """ - if not value or value.startswith("data:"): + url_value = value.lstrip() + if not url_value or url_value.lower().startswith("data:"): return None - if value.startswith(("http://", "https://")): - # Strip query string (e.g. SAS tokens) before taking the basename. - parsed = urlparse(value) + if url_value.lower().startswith(("http://", "https://")): + # Use only the URL path, excluding authority, query, and fragment credentials. + parsed = urlparse(url_value) name = PureWindowsPath(parsed.path).name return name or None # Local path — PureWindowsPath treats both ``/`` and ``\`` as separators, @@ -66,7 +67,9 @@ def format_last_message_preview( Media-path data types are rendered as ``[Image: ]`` (and variants) so the absolute filesystem path of memory artifacts is never - exposed through API responses or UI previews. Error values are replaced + exposed through API responses or UI previews. HTTP(S) URL recognition + ignores scheme case and leading whitespace, and only the path contributes + to the filename. Error values are replaced with a generic status so persisted exception tracebacks are not exposed. Text-like data types pass through with truncation and an ellipsis suffix when they exceed *max_len*. diff --git a/pyrit/backend/models/datasets.py b/pyrit/backend/models/datasets.py index 9f8814534e..f3566c340a 100644 --- a/pyrit/backend/models/datasets.py +++ b/pyrit/backend/models/datasets.py @@ -6,11 +6,16 @@ Datasets are seed prompt/objective collections provided by ``SeedDatasetProvider`` subclasses. These models describe the wire format for -listing available datasets. +listing available datasets and browsing their stored seed examples. """ +from uuid import UUID + from pydantic import BaseModel, Field +from pyrit.backend.models.common import PaginationInfo +from pyrit.models import PromptDataType, SeedRecord, SeedType + class DatasetInfo(BaseModel): """Metadata about a single available dataset.""" @@ -41,3 +46,34 @@ class DatasetListResponse(BaseModel): """Response for listing available datasets.""" items: list[DatasetInfo] = Field(..., description="List of available datasets") + + +class SeedExampleSummary(BaseModel): + """One logical seed example: the seeds that share a group ID, or one seed without a group.""" + + example_id: UUID = Field(..., description="The prompt_group_id, or the seed ID of a seed without a group") + name: str | None = Field(None, description="The first member name, if any") + preview: str = Field(..., description="Text preview of at most 100 characters, or a type label") + preview_truncated: bool = Field(..., description="Whether the preview text was shortened") + modalities: list[PromptDataType] + seed_types: list[SeedType] + piece_count: int + objective_count: int + harm_categories: list[str] + has_unlabeled_harm: bool = Field(..., description="Whether any member has no harm category") + + +class SeedExampleListResponse(BaseModel): + """One page of logical seed examples.""" + + items: list[SeedExampleSummary] + pagination: PaginationInfo + total: int = Field(..., description="Number of logical examples that match the filters") + + +class SeedExampleDetailResponse(SeedExampleSummary): + """One logical seed example with all of its stored seeds.""" + + members: list[SeedRecord] = Field( + ..., description="Stored seed records without reconstruction, objectives first, then by sequence" + ) diff --git a/pyrit/backend/routes/datasets.py b/pyrit/backend/routes/datasets.py index 39109fbb58..7dc342be94 100644 --- a/pyrit/backend/routes/datasets.py +++ b/pyrit/backend/routes/datasets.py @@ -4,20 +4,27 @@ """ Dataset API routes. -Provides an endpoint for listing available seed datasets. Datasets are -discovered from registered ``SeedDatasetProvider`` subclasses. +Lists available seed datasets and browses the stored seed examples of one dataset. Datasets are +discovered from registered ``SeedDatasetProvider`` subclasses and from memory. """ -from fastapi import APIRouter +from uuid import UUID -from pyrit.backend.models.common import ProblemDetail +from fastapi import APIRouter, HTTPException, Query, status + +from pyrit.backend.models.common import MAX_ITEMS, CursorStr, ProblemDetail from pyrit.backend.models.datasets import ( DatasetListResponse, + SeedExampleDetailResponse, + SeedExampleListResponse, ) from pyrit.backend.services.dataset_service import get_dataset_service +from pyrit.models import PromptDataType, SeedType router = APIRouter(prefix="/datasets", tags=["datasets"]) +_SELECTION_KEY_DESCRIPTION = "The selection_key of a dataset from GET /datasets" + @router.get( "", @@ -38,3 +45,76 @@ async def list_datasets(loaded_only: bool = False) -> DatasetListResponse: # py """ service = get_dataset_service() return await service.list_datasets_async(loaded_only=loaded_only) + + +@router.get( + "/seeds", + response_model=SeedExampleListResponse, + responses={ + 400: {"model": ProblemDetail, "description": "Invalid selection key or cursor"}, + }, +) +async def list_seed_examples( # pyrit-async-suffix-exempt + selection_key: str = Query(..., description=_SELECTION_KEY_DESCRIPTION), + limit: int = Query(20, ge=1, le=100, description="Maximum examples per page"), + cursor: CursorStr | None = Query( + None, + description="The next_cursor of the previous page. A cursor is valid only with the same " + "selection_key and filters.", + ), + search: str | None = Query( + None, + max_length=1000, + description="Case-insensitive literal text to find in text prompt and objective values", + ), + modality: list[PromptDataType] | None = Query(None, max_length=MAX_ITEMS, description="Data types, OR-matched"), + harm_category: list[str] | None = Query( + None, max_length=MAX_ITEMS, description="Whole harm categories, case-insensitive, OR-matched" + ), + seed_type: list[SeedType] | None = Query(None, max_length=MAX_ITEMS, description="Seed types, OR-matched"), +) -> SeedExampleListResponse: + """ + List one page of the stored logical seed examples of a dataset. + + Examples are ordered newest first. Different filters are AND-matched, and different members + of an example can match different filters. + + Returns: + SeedExampleListResponse: The page, its pagination data, and the number of matching examples. + """ + return await get_dataset_service().list_seed_examples_async( + selection_key=selection_key, + limit=limit, + cursor=cursor, + search=search, + data_types=modality, + harm_categories=harm_category, + seed_types=seed_type, + ) + + +@router.get( + "/seeds/{example_id}", + response_model=SeedExampleDetailResponse, + responses={ + 400: {"model": ProblemDetail, "description": "Invalid selection key"}, + 404: {"model": ProblemDetail, "description": "Seed example not found in the dataset"}, + }, +) +async def get_seed_example( # pyrit-async-suffix-exempt + example_id: UUID, + selection_key: str = Query(..., description=_SELECTION_KEY_DESCRIPTION), +) -> SeedExampleDetailResponse: + """ + Get one stored logical seed example with all of its members. + + Returns: + SeedExampleDetailResponse: The example and its members. + + Raises: + HTTPException: If the dataset does not contain the example. + """ + example = await get_dataset_service().get_seed_example_async(selection_key=selection_key, example_id=example_id) + if example is None: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=f"Seed example '{example_id}' not found") + return example diff --git a/pyrit/backend/services/dataset_service.py b/pyrit/backend/services/dataset_service.py index ce83c58290..84920ffb85 100644 --- a/pyrit/backend/services/dataset_service.py +++ b/pyrit/backend/services/dataset_service.py @@ -2,22 +2,39 @@ # Licensed under the MIT license. """ -Dataset service for listing seed datasets. +Dataset service for listing seed datasets and browsing their stored seed examples. Wraps ``SeedDatasetProvider`` discovery and memory to list available datasets. """ import logging +import ntpath +import posixpath +import re from collections.abc import Sequence from functools import lru_cache +from uuid import UUID +from pyrit.backend.mappers import format_last_message_preview +from pyrit.backend.models.common import PaginationInfo from pyrit.backend.models.datasets import ( DatasetInfo, DatasetListResponse, + SeedExampleDetailResponse, + SeedExampleListResponse, + SeedExampleSummary, ) +from pyrit.common.pagination import decode_keyset_cursor, encode_keyset_cursor, fingerprint_filters from pyrit.datasets import SeedDatasetProvider from pyrit.memory import CentralMemory -from pyrit.models import SeedDatasetSummary +from pyrit.models import ( + MEDIA_PATH_DATA_TYPES, + ConversationStats, + PromptDataType, + SeedDatasetSummary, + SeedRecord, + SeedType, +) logger = logging.getLogger(__name__) @@ -94,6 +111,159 @@ async def list_datasets_async(self, *, loaded_only: bool = False) -> DatasetList return DatasetListResponse(items=items) + async def list_seed_examples_async( + self, + *, + selection_key: str, + limit: int = 20, + cursor: str | None = None, + search: str | None = None, + data_types: Sequence[PromptDataType] | None = None, + harm_categories: Sequence[str] | None = None, + seed_types: Sequence[SeedType] | None = None, + ) -> SeedExampleListResponse: + """ + List one page of the stored logical seed examples of a dataset. + + The cursor is bound to the dataset and the filters. A cursor that is not valid for the + request causes an error. It does not restart at the first page. + + Args: + selection_key (str): The ``selection_key`` of a dataset from ``list_datasets_async``. + limit (int): The maximum number of examples to return. + cursor (str | None): The ``next_cursor`` of the previous page. + search (str | None): Literal text that the value of a text prompt or objective must contain, + ignoring case. Simulated-conversation configurations are not searched. + data_types (Sequence[PromptDataType] | None): Match examples with a member of any of these data types. + harm_categories (Sequence[str] | None): Match examples with a member in any of these harm categories. + seed_types (Sequence[SeedType] | None): Match examples with a member of any of these seed types. + + Returns: + SeedExampleListResponse: The page, its pagination data, and the number of matching examples. + + Raises: + ValueError: If the selection key or the cursor is not valid. + """ + dataset_name = self._parse_selection_key(selection_key) + fingerprint = fingerprint_filters( + filters={ + "selection_key": selection_key, + "data_types": data_types or None, + "harm_categories": [category.lower() for category in harm_categories] if harm_categories else None, + "seed_types": seed_types or None, + "search": search or None, + } + ) + after = decode_keyset_cursor(cursor=cursor, fingerprint=fingerprint) + if cursor and after is None: + raise ValueError("The cursor is not valid for this dataset and these filters") + + examples, total, next_after = await self._memory.get_seed_examples_async( + dataset_name=dataset_name, + limit=limit, + after=after, + data_types=data_types, + harm_categories=harm_categories, + seed_types=seed_types, + value_search=search, + ) + next_cursor = ( + encode_keyset_cursor( + timestamp=next_after.timestamp, identifier=next_after.identifier, fingerprint=fingerprint + ) + if next_after + else None + ) + return SeedExampleListResponse( + items=[self._summarize(example_id=example_id, seeds=seeds) for example_id, seeds in examples.items()], + pagination=PaginationInfo( + limit=limit, has_more=next_after is not None, next_cursor=next_cursor, prev_cursor=cursor + ), + total=total, + ) + + async def get_seed_example_async(self, *, selection_key: str, example_id: UUID) -> SeedExampleDetailResponse | None: + """ + Get one stored logical seed example of a dataset with all of its members. + + Args: + selection_key (str): The ``selection_key`` of a dataset from ``list_datasets_async``. + example_id (UUID): The ``example_id`` from the list response. + + Returns: + SeedExampleDetailResponse | None: The example, or None if the dataset does not contain it. + + Raises: + ValueError: If the selection key is not valid. + """ + seeds = await self._memory.get_seed_example_async( + dataset_name=self._parse_selection_key(selection_key), example_id=example_id + ) + if not seeds: + return None + summary = self._summarize(example_id=example_id, seeds=seeds) + return SeedExampleDetailResponse(**summary.model_dump(), members=seeds) + + @staticmethod + def _parse_selection_key(selection_key: str) -> str | None: + """ + Get the dataset name of a selection key. + + Returns: + str | None: The dataset name, or None for the unnamed population. + + Raises: + ValueError: If the selection key is not valid. + """ + if selection_key == "dataset:unnamed": + return None + name = selection_key.removeprefix("dataset:named:") + if not name or name == selection_key: + raise ValueError(f"Invalid dataset selection key: {selection_key}") + return name + + @classmethod + def _summarize(cls, *, example_id: UUID, seeds: Sequence[SeedRecord]) -> SeedExampleSummary: + preview, truncated = cls._preview(seeds) + return SeedExampleSummary( + example_id=example_id, + name=next((seed.name for seed in seeds if seed.name), None), + preview=preview, + preview_truncated=truncated, + modalities=sorted({seed.data_type for seed in seeds}), + seed_types=sorted({seed.seed_type for seed in seeds}), + piece_count=len(seeds), + objective_count=sum(seed.seed_type == "objective" for seed in seeds), + harm_categories=sorted({category for seed in seeds for category in seed.harm_categories or []}), + has_unlabeled_harm=any(not seed.harm_categories for seed in seeds), + ) + + @staticmethod + def _preview(seeds: Sequence[SeedRecord]) -> tuple[str, bool]: + """ + Build the list preview from the first text seed, or else from a type label. + + A simulated-conversation configuration is not a prompt, so it gets only a label. + Media seeds show only a file name. Standalone paths and URLs stored as text get a label. + + Returns: + tuple[str, bool]: The preview and whether the text was shortened. + """ + shown = [seed for seed in seeds if seed.seed_type != "simulated_conversation"] + seed = next((seed for seed in shown if seed.data_type == "text"), shown[0] if shown else None) + if seed is None: + return "[Simulated conversation configuration]", False + value = seed.value.lstrip() + if seed.data_type == "text" and ( + ntpath.isabs(value) or posixpath.isabs(value) or re.match(r"^[a-zA-Z][a-zA-Z0-9+.-]*://", value) + ): + return "[Text reference]", False + preview = None + if seed.data_type == "text" or seed.data_type in MEDIA_PATH_DATA_TYPES: + preview = format_last_message_preview(value=seed.value, data_type=seed.data_type) + truncated = seed.data_type == "text" and len(seed.value) > ConversationStats.PREVIEW_MAX_LEN + return preview or f"[{seed.data_type}]", truncated + @staticmethod def _merge_unnamed_summaries(summaries: Sequence[SeedDatasetSummary]) -> SeedDatasetSummary | None: """ diff --git a/pyrit/memory/memory_interface.py b/pyrit/memory/memory_interface.py index 4492507194..63b0c0ad0b 100644 --- a/pyrit/memory/memory_interface.py +++ b/pyrit/memory/memory_interface.py @@ -22,6 +22,7 @@ from sqlalchemy import ( MetaData, + String, Unicode, and_, case, @@ -35,6 +36,7 @@ type_coerce, update, ) +from sqlalchemy import cast as sql_cast from sqlalchemy.engine.base import Engine from sqlalchemy.exc import IntegrityError, SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession @@ -44,6 +46,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 if TYPE_CHECKING: from pyrit.memory.memory_embedding import MemoryEmbedding @@ -102,6 +105,7 @@ MessagePiece, MessageScorable, Observation, + PromptDataType, RetryEvent, ScenarioAttackResultDelta, ScenarioIdentifier, @@ -119,6 +123,7 @@ SeedObjective, SeedOrigin, SeedPrompt, + SeedRecord, SeedType, TargetIdentifier, group_conversation_message_pieces_by_sequence, @@ -4373,6 +4378,165 @@ def _execute_get_seed_dataset_summaries(self) -> Sequence[SeedDatasetSummary]: logger.exception(f"Failed to retrieve dataset summaries with error {e}") raise + @staticmethod + def _seed_example_scope(*, dataset_name: str | None) -> "ColumnElement[bool]": + if dataset_name: + return SeedEntry.dataset_name == dataset_name + return or_(SeedEntry.dataset_name.is_(None), SeedEntry.dataset_name == "") + + def _seed_example_filters( + self, + *, + scope: "ColumnElement[bool]", + data_types: Sequence[PromptDataType] | None, + harm_categories: Sequence[str] | None, + seed_types: Sequence[SeedType] | None, + value_search: str | None, + ) -> list[Any]: + """ + Build one logical-example membership condition for each active filter. + + Each filter matches when any member of the example matches it, so different members can + satisfy different filters. The IN subqueries use the unaliased table because the + dialect JSON array match emits SQL text that references the table name. + + Returns: + list[Any]: SQLAlchemy conditions to combine with AND. + """ + member_conditions: list[Any] = [] + if data_types: + member_conditions.append(SeedEntry.data_type.in_(list(data_types))) + if harm_categories: + member_conditions.append( + self._get_condition_json_array_match( + json_column=SeedEntry.harm_categories, + property_path="$", + array_to_match=list(harm_categories), + match_mode="any", + ) + ) + if seed_types: + member_conditions.append(SeedEntry.seed_type.in_(list(seed_types))) + if value_search: + # A simulated-conversation value is JSON, so its keys would match common words such as "prompt". + pattern = "%" + re.sub(r"([\\%_\[])", r"\\\1", value_search) + "%" + member_conditions.append( + and_( + SeedEntry.data_type == "text", + SeedEntry.seed_type != "simulated_conversation", + SeedEntry.value.ilike(pattern, escape="\\"), + ) + ) + + logical_id = func.coalesce(SeedEntry.prompt_group_id, SeedEntry.id) + return [ + logical_id.in_(select(logical_id).where(scope, condition).correlate(None)) + for condition in member_conditions + ] + + @staticmethod + def _get_seed_example_seeds( + *, session: Session, scope: "ColumnElement[bool]", example_ids: Sequence[uuid.UUID] + ) -> dict[uuid.UUID, list[SeedRecord]]: + """ + Read all stored members without reconstructing seeds or resolving configuration paths. + + Returns: + dict[uuid.UUID, list[SeedRecord]]: The stored members of each example, in the + order of ``example_ids``. Objectives come first, then seeds by sequence and ID. + """ + logical_id = func.coalesce(SeedEntry.prompt_group_id, SeedEntry.id) + entries = session.scalars( + select(SeedEntry) + .where(scope, logical_id.in_(example_ids)) + .order_by( + case((SeedEntry.seed_type == "objective", 0), else_=1), + SeedEntry.sequence, + func.lower(sql_cast(SeedEntry.id, String(36))), + ) + ).all() + seeds: dict[uuid.UUID, list[SeedRecord]] = {example_id: [] for example_id in example_ids} + for entry in entries: + seeds[entry.prompt_group_id or entry.id].append(entry.get_seed_record()) + return {example_id: example_seeds for example_id, example_seeds in seeds.items() if example_seeds} + + def _execute_get_seed_examples( + self, + *, + dataset_name: str | None, + limit: int, + after: DecodedKeysetCursor | None, + data_types: Sequence[PromptDataType] | None, + harm_categories: Sequence[str] | None, + seed_types: Sequence[SeedType] | None, + value_search: str | None, + ) -> tuple[dict[uuid.UUID, list[SeedRecord]], int, DecodedKeysetCursor | None]: + """ + Read one keyset page of complete logical seed examples. + + Returns: + tuple[dict[uuid.UUID, list[SeedRecord]], int, DecodedKeysetCursor | None]: The seeds of each + example in page order, the number of examples that match the filters, and the sort key of + the last example when more examples follow. + """ + logical_id = func.coalesce(SeedEntry.prompt_group_id, SeedEntry.id) + logical_id_key = func.lower(sql_cast(logical_id, String(36))) + scope = self._seed_example_scope(dataset_name=dataset_name) + filters = self._seed_example_filters( + scope=scope, + data_types=data_types, + harm_categories=harm_categories, + seed_types=seed_types, + value_search=value_search, + ) + grouped = ( + select( + logical_id.label("example_id"), + logical_id_key.label("example_id_key"), + func.min(SeedEntry.date_added).label("first_added"), + ) + .where(scope, *filters) + .group_by(logical_id, logical_id_key) + .subquery() + ) + page = select(grouped.c.example_id, grouped.c.first_added) + if after is not None: + anchor_id = str(uuid.UUID(after.identifier)) + page = page.where( + or_( + grouped.c.first_added < after.timestamp, + and_(grouped.c.first_added == after.timestamp, grouped.c.example_id_key < anchor_id), + ) + ) + page = page.order_by(grouped.c.first_added.desc(), grouped.c.example_id_key.desc()).limit(limit + 1) + + with closing(self._get_session()) as session: + total = session.execute(select(func.count()).select_from(grouped)).scalar_one() + rows = session.execute(page).all() + seeds = self._get_seed_example_seeds( + 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)) + return seeds, total, next_after + + def _execute_get_seed_example(self, *, dataset_name: str | None, example_id: uuid.UUID) -> list[SeedRecord]: + """ + Read one complete logical seed example. + + Returns: + list[SeedRecord]: The stored members. The list is empty if the dataset does not contain the example. + """ + with closing(self._get_session()) as session: + seeds = self._get_seed_example_seeds( + session=session, + scope=self._seed_example_scope(dataset_name=dataset_name), + example_ids=[example_id], + ) + return seeds.get(example_id, []) + def _execute_get_seed_dataset_names(self) -> Sequence[str]: """ Return a list of all seed dataset names in the memory storage. @@ -8574,6 +8738,72 @@ async def get_seed_dataset_summaries_async(self) -> Sequence[SeedDatasetSummary] """ return await self._run_database_operation_async(self._execute_get_seed_dataset_summaries) + async def get_seed_examples_async( + self, + *, + dataset_name: str | None, + limit: int, + after: DecodedKeysetCursor | None = None, + data_types: Sequence[PromptDataType] | None = None, + harm_categories: Sequence[str] | None = None, + seed_types: Sequence[SeedType] | None = None, + value_search: str | None = None, + ) -> tuple[dict[uuid.UUID, list[SeedRecord]], int, DecodedKeysetCursor | None]: + """ + Read one page of complete logical seed examples from one dataset. + + A logical example is all seeds in the dataset that share a ``prompt_group_id``, or one seed + without a group. Examples are ordered by their earliest ``date_added``, then by example ID, + both descending, using the canonical textual UUID order on every backend. Values inside + one filter use OR, different filters use AND, and any member can satisfy a filter. + Members are stored-record projections: configurations remain raw text, and no templates + are rendered or referenced files loaded. + + Args: + dataset_name: The dataset name. None or an empty string selects seeds without a dataset name. + limit: The maximum number of examples to return. + after: The sort key of the last example on the previous page. + data_types: Match seeds with any of these data types. + harm_categories: Match seeds with any of these harm categories, as whole values that + ignore case. + seed_types: Match seeds with any of these seed types. + value_search: Match text prompts and objectives whose stored value contains this literal + text, ignoring case. Simulated-conversation configurations are not searched. + + Returns: + tuple[dict[uuid.UUID, list[SeedRecord]], int, DecodedKeysetCursor | None]: The stored members of each + example keyed by example ID in page order, objectives first; the number of examples that + match the filters; and the sort key of the last example when more examples follow. + """ + return await self._run_database_operation_async( + self._execute_get_seed_examples, + dataset_name=dataset_name, + limit=limit, + after=after, + data_types=data_types, + harm_categories=harm_categories, + seed_types=seed_types, + value_search=value_search, + ) + + async def get_seed_example_async(self, *, dataset_name: str | None, example_id: uuid.UUID) -> list[SeedRecord]: + """ + Read one complete logical seed example from one dataset. + + Seeds are read as in ``get_seed_examples_async``. + + Args: + dataset_name: The dataset name. None or an empty string selects seeds without a dataset name. + example_id: The ``prompt_group_id`` of the example, or the seed ID of a seed without a group. + + Returns: + list[SeedRecord]: The stored members, objectives first. The list is empty if the + dataset does not contain it. + """ + return await self._run_database_operation_async( + self._execute_get_seed_example, dataset_name=dataset_name, example_id=example_id + ) + def get_seed_dataset_names(self) -> Sequence[str]: """ Use ``get_seed_dataset_names_async``. diff --git a/pyrit/memory/memory_models.py b/pyrit/memory/memory_models.py index afc5c07601..4369bff228 100644 --- a/pyrit/memory/memory_models.py +++ b/pyrit/memory/memory_models.py @@ -74,6 +74,7 @@ SeedObjective, SeedOrigin, SeedPrompt, + SeedRecord, SeedSimulatedConversation, SeedType, TargetIdentifier, @@ -1648,12 +1649,45 @@ def _unpack_seed_metadata( decoded = None return cleaned, decoded - def get_seed(self) -> Seed: + def get_seed_record(self) -> SeedRecord: + """ + Project stored fields without reconstructing an executable seed. + + Returns: + SeedRecord: Stored content, identifiers, and metadata, including raw configuration text. + """ + metadata, response_json_schema = self._unpack_seed_metadata(self.prompt_metadata) + return SeedRecord( + id=self.id, + seed_type=self.seed_type, + value=self.value, + value_sha256=self.value_sha256, + data_type=self.data_type, + name=self.name, + dataset_name=self.dataset_name, + origin=SeedOrigin(self.origin), + harm_categories=self.harm_categories, + description=self.description, + authors=self.authors, + groups=self.groups, + source=self.source, + date_added=self.date_added, + added_by=self.added_by, + metadata=metadata, + prompt_group_id=self.prompt_group_id, + sequence=self.sequence, + role=self.role, + parameters=self.parameters, + conditions=self.conditions, + response_json_schema=response_json_schema, + ) + + def get_seed(self) -> SeedPrompt | SeedObjective | SeedSimulatedConversation: """ Convert this database entry back into a Seed object. Returns: - Seed: The reconstructed seed object (SeedPrompt, SeedObjective, or SeedSimulatedConversation) + SeedPrompt | SeedObjective | SeedSimulatedConversation: The reconstructed seed object. Raises: ValueError: If persisted conditions are invalid or attached to a non-objective seed, diff --git a/pyrit/models/__init__.py b/pyrit/models/__init__.py index 5dc925a098..3cfc87e1aa 100644 --- a/pyrit/models/__init__.py +++ b/pyrit/models/__init__.py @@ -225,6 +225,7 @@ SeedObjective, SeedOrigin, SeedPrompt, + SeedRecord, SeedSimulatedConversation, SeedUnion, SimulatedTargetSystemPromptPaths, @@ -441,6 +442,7 @@ "SeedObjective": "pyrit.models.seeds", "SeedOrigin": "pyrit.models.seeds", "SeedPrompt": "pyrit.models.seeds", + "SeedRecord": "pyrit.models.seeds", "SeedDataset": "pyrit.models.seeds", "SeedDatasetSummary": "pyrit.models.seeds", "SeedGroup": "pyrit.models.seeds", diff --git a/pyrit/models/seeds/__init__.py b/pyrit/models/seeds/__init__.py index ce8f58460b..4b355409d6 100644 --- a/pyrit/models/seeds/__init__.py +++ b/pyrit/models/seeds/__init__.py @@ -32,6 +32,7 @@ from pyrit.models.seeds.seed_objective import SeedObjective from pyrit.models.seeds.seed_origin import SeedOrigin from pyrit.models.seeds.seed_prompt import SeedPrompt + from pyrit.models.seeds.seed_record import SeedRecord from pyrit.models.seeds.seed_simulated_conversation import ( NextMessageSystemPromptPaths, SeedSimulatedConversation, @@ -66,6 +67,7 @@ "SeedObjective": "pyrit.models.seeds.seed_objective", "SeedOrigin": "pyrit.models.seeds.seed_origin", "SeedPrompt": "pyrit.models.seeds.seed_prompt", + "SeedRecord": "pyrit.models.seeds.seed_record", "SeedSimulatedConversation": "pyrit.models.seeds.seed_simulated_conversation", "SeedUnion": "pyrit.models.seeds.seed_group", "SimulatedTargetSystemPromptPaths": "pyrit.models.seeds.seed_simulated_conversation", diff --git a/pyrit/models/seeds/seed_record.py b/pyrit/models/seeds/seed_record.py new file mode 100644 index 0000000000..d081f20c69 --- /dev/null +++ b/pyrit/models/seeds/seed_record.py @@ -0,0 +1,42 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +from datetime import datetime +from typing import Any +from uuid import UUID + +from pydantic import BaseModel, ConfigDict + +from pyrit.models.literals import ChatMessageRole, PromptDataType, SeedType +from pyrit.models.score.condition import ConditionTuple +from pyrit.models.seeds.seed_origin import SeedOrigin +from pyrit.models.target.json_schema_definition import JsonSchemaDefinition + + +class SeedRecord(BaseModel): + """Stored seed data for inspection, without rendering, file loading, or generated defaults.""" + + model_config = ConfigDict(frozen=True, extra="forbid") + + id: UUID + seed_type: SeedType + value: str + value_sha256: str | None + data_type: PromptDataType + name: str | None + dataset_name: str | None + origin: SeedOrigin + harm_categories: list[str] | None + description: str | None + authors: list[str] | None + groups: list[str] | None + source: str | None + date_added: datetime + added_by: str + metadata: dict[str, Any] | None + prompt_group_id: UUID | None + sequence: int | None + role: ChatMessageRole | None + parameters: list[str] | None + conditions: ConditionTuple | None + response_json_schema: JsonSchemaDefinition | None diff --git a/tests/integration/memory/test_seed_examples_azure_sql_integration.py b/tests/integration/memory/test_seed_examples_azure_sql_integration.py new file mode 100644 index 0000000000..b13f2e7392 --- /dev/null +++ b/tests/integration/memory/test_seed_examples_azure_sql_integration.py @@ -0,0 +1,70 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Azure SQL execution of the seed example read queries.""" + +from datetime import UTC, datetime, timedelta +from uuid import UUID, uuid4 + +import pytest +from sqlalchemy import delete + +from pyrit.memory import AzureSQLMemory +from pyrit.memory.memory_models import SeedEntry +from pyrit.models import SeedObjective, SeedPrompt + + +@pytest.mark.run_only_if_all_tests +async def test_seed_examples_on_azure_sql(azuresql_instance: AzureSQLMemory): + test_id = str(uuid4()) + dataset = f"2748-azure-{test_id}" + group = UUID("ffffffff-ffff-ffff-ffff-000000000000") + older = UUID("00000000-0000-0000-0000-000000000001") + newer = uuid4() + base_time = datetime(2024, 1, 1, tzinfo=UTC) + seeds = [ + SeedPrompt(value="hello", dataset_name=dataset, prompt_group_id=group, date_added=base_time), + SeedObjective( + value="objective", + dataset_name=dataset, + prompt_group_id=group, + harm_categories=["Violence"], + date_added=base_time, + ), + SeedPrompt( + value="literal 100% a_b [ab] \\", + dataset_name=dataset, + prompt_group_id=newer, + date_added=base_time + timedelta(1), + ), + SeedPrompt(value="older", dataset_name=dataset, prompt_group_id=older, date_added=base_time), + ] + await azuresql_instance.add_seeds_to_memory_async(seeds=seeds, added_by=test_id) + + async def ids(**filters) -> list: + examples, _, _ = await azuresql_instance.get_seed_examples_async(dataset_name=dataset, limit=100, **filters) + return list(examples) + + try: + assert await ids() == [newer, group, older] + assert await ids(harm_categories=["missing", "VIOLENCE"], value_search="HELLO") == [group] + assert await ids(seed_types=["objective"], data_types=["text"]) == [group] + for search in ["100%", "a_b", "[ab]", "\\"]: + assert await ids(value_search=search) == [newer] + assert await ids(value_search="0_") == [] + + first, _, after = await azuresql_instance.get_seed_examples_async(dataset_name=dataset, limit=2) + rest, _, last = await azuresql_instance.get_seed_examples_async(dataset_name=dataset, limit=100, after=after) + assert after is not None + assert after.identifier == str(group) + assert list(first) == [newer, group] + assert list(rest) == [older] + assert last is None + assert [*first, *rest] == await ids() + + detail = await azuresql_instance.get_seed_example_async(dataset_name=dataset, example_id=group) + assert [seed.id for seed in detail] == [seeds[1].id, seeds[0].id] + finally: + async with await azuresql_instance.get_session_async() as session: + await session.execute(delete(SeedEntry).where(SeedEntry.added_by == test_id)) + await session.commit() diff --git a/tests/unit/backend/test_preview.py b/tests/unit/backend/test_preview.py index 5b7b071d37..fc593e4fe0 100644 --- a/tests/unit/backend/test_preview.py +++ b/tests/unit/backend/test_preview.py @@ -82,6 +82,36 @@ def test_media_azure_blob_url_strips_query_and_keeps_filename(self) -> None: assert "sig=" not in (result or "") assert "blob.core.windows.net" not in (result or "") + @pytest.mark.parametrize( + ("data_type", "label"), + [ + ("image_path", "Image"), + ("audio_path", "Audio"), + ("video_path", "Video"), + ("binary_path", "File"), + ], + ) + @pytest.mark.parametrize("prefix", ["HTTPS://", "hTtP://", " \thttps://", "\r\n HTTPS://"]) + def test_media_url_normalization_hides_credentials(self, *, data_type: str, label: str, prefix: str) -> None: + url = f"{prefix}reader:password@example.test/private/file.png?sig=secret#token=secret" + + result = format_last_message_preview(value=url, data_type=data_type) + + assert result == f"[{label}: file.png]" + + @pytest.mark.parametrize("path", ["", "/"]) + def test_media_url_without_filename_does_not_expose_authority(self, *, path: str) -> None: + url = f" \tHTTPS://reader:password@example.test{path}?sig=secret#token=secret" + + result = format_last_message_preview(value=url, data_type="image_path") + + assert result == "[Image]" + + def test_media_local_filename_keeps_leading_space_and_hash(self) -> None: + result = format_last_message_preview(value=" image#1.png", data_type="image_path") + + assert result == "[Image: image#1.png]" + def test_media_empty_value_falls_back_to_label_only(self) -> None: result = format_last_message_preview(value="", data_type="image_path", max_len=100) assert result == "[Image]" @@ -99,6 +129,11 @@ def test_media_data_uri_falls_back_to_label_only(self) -> None: ) assert result == "[Image]" + def test_media_uppercase_data_uri_with_whitespace_falls_back_to_label_only(self) -> None: + result = format_last_message_preview(value=" \tDATA:image/png;base64,iVBORw0KGgo=", data_type="image_path") + + assert result == "[Image]" + def test_media_long_path_basename_not_truncated(self) -> None: # Even with a 100-char text limit, the basename label should not be # truncated. Memory layer fetches up to PREVIEW_FETCH_MAX_LEN chars so diff --git a/tests/unit/backend/test_seed_example_routes.py b/tests/unit/backend/test_seed_example_routes.py new file mode 100644 index 0000000000..0f39b7a9e8 --- /dev/null +++ b/tests/unit/backend/test_seed_example_routes.py @@ -0,0 +1,286 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +import json +from collections.abc import AsyncIterator +from pathlib import Path +from unittest.mock import patch +from uuid import uuid4 + +import pytest +from httpx import ASGITransport, AsyncClient + +from pyrit.backend.main import app +from pyrit.backend.services.dataset_service import get_dataset_service +from pyrit.common.utils import to_sha256 +from pyrit.memory import MemoryInterface +from pyrit.memory.memory_models import SeedEntry +from pyrit.models import ( + AnswerMatches, + MatchesObjective, + PromptDataType, + Seed, + SeedObjective, + SeedPrompt, + SeedSimulatedConversation, +) + +URL = "/api/datasets/seeds" +DATASET = "browse" +NAMED = f"dataset:named:{DATASET}" +MISSING_ID = "00000000-0000-4000-8000-000000000000" + + +@pytest.fixture +async def client(patch_central_database, compatibility_headers: dict[str, str]) -> AsyncIterator[AsyncClient]: + get_dataset_service.cache_clear() + transport = ASGITransport(app=app) + async with AsyncClient(transport=transport, base_url="http://test", headers=compatibility_headers) as client: + yield client + get_dataset_service.cache_clear() + + +async def _store(memory: MemoryInterface, *seeds: Seed) -> None: + async with await memory.get_session_async() as session: + session.add_all(SeedEntry(entry=seed) for seed in seeds) + await session.commit() + + +def _prompt(value: str, **kwargs) -> SeedPrompt: + return SeedPrompt(value=value, dataset_name=DATASET, added_by="test", **kwargs) + + +def _conversation(**kwargs) -> SeedSimulatedConversation: + return SeedSimulatedConversation( + num_turns=2, + adversarial_chat_system_prompt=SeedPrompt(value="adversarial", parameters=["objective"]), + simulated_target_system_prompt=SeedPrompt(value="target", parameters=["objective", "num_turns"]), + dataset_name=DATASET, + added_by="test", + **kwargs, + ) + + +async def test_list_seed_examples_pages_with_a_filter_bound_cursor(client: AsyncClient, sqlite_instance): + await _store(sqlite_instance, *(_prompt(f"prompt {index}") for index in range(3)), _prompt("other")) + params = {"selection_key": NAMED, "limit": 2, "search": "PROMPT"} + + first = (await client.get(URL, params=params)).json() + cursor = first["pagination"]["next_cursor"] + second = (await client.get(URL, params={**params, "cursor": cursor})).json() + mismatched = await client.get(URL, params={"selection_key": NAMED, "limit": 2, "cursor": cursor}) + + assert len({item["example_id"] for item in first["items"] + second["items"]}) == 3 + assert first["total"] == second["total"] == 3 + assert (first["pagination"]["has_more"], second["pagination"]["has_more"]) == (True, False) + assert second["pagination"]["prev_cursor"] == cursor + assert mismatched.status_code == 400 + + +@pytest.mark.parametrize( + ("path", "params", "expected_status"), + [ + ("", {"selection_key": NAMED, "cursor": "not-a-cursor"}, 400), + ("", {"selection_key": DATASET}, 400), + ("", {"selection_key": "dataset:named:"}, 400), + ("", {"selection_key": NAMED, "limit": 0}, 422), + (f"/{MISSING_ID}", {"selection_key": "dataset:other"}, 400), + (f"/{MISSING_ID}", {"selection_key": NAMED}, 404), + ("/not-a-uuid", {"selection_key": NAMED}, 422), + ], +) +async def test_seed_example_routes_reject_bad_requests( + client: AsyncClient, path: str, params: dict[str, str | int], expected_status: int +): + response = await client.get(URL + path, params=params) + + assert response.status_code == expected_status + + +async def test_list_seed_examples_builds_safe_previews(client: AsyncClient, sqlite_instance): + long_text = _prompt("x" * 150) + image = _prompt("https://account.blob.core.windows.net/c/cat.png?sig=secret", data_type="image_path") + url = _prompt("https://example.com/?token=secret", data_type="url") + configuration = _conversation() + await _store(sqlite_instance, long_text, image, url, configuration) + + response = await client.get(URL, params={"selection_key": NAMED}) + + items = {item["example_id"]: item for item in response.json()["items"]} + assert (items[str(long_text.id)]["preview"], items[str(long_text.id)]["preview_truncated"]) == ( + "x" * 100 + "...", + True, + ) + assert items[str(image.id)]["preview"] == "[Image: cat.png]" + assert items[str(url.id)]["preview"] == "[url]" + assert items[str(configuration.id)]["preview"] == "[Simulated conversation configuration]" + + +@pytest.mark.parametrize( + ("data_type", "label"), + [("image_path", "Image"), ("audio_path", "Audio"), ("video_path", "Video"), ("binary_path", "File")], +) +@pytest.mark.parametrize("prefix", ["HTTPS://", "hTtP://", " \thttps://"]) +async def test_seed_example_media_preview_normalizes_url_prefix_async( + *, + client: AsyncClient, + sqlite_instance: MemoryInterface, + data_type: PromptDataType, + label: str, + prefix: str, +) -> None: + url = f"{prefix}reader:password@example.test/private/file.png?sig=secret#token=secret" + seed = _prompt(url, data_type=data_type) + await _store(sqlite_instance, seed) + + listed = await client.get(URL, params={"selection_key": NAMED}) + detail = await client.get(f"{URL}/{seed.id}", params={"selection_key": NAMED}) + + assert listed.status_code == detail.status_code == 200 + assert listed.json()["items"][0]["preview"] == f"[{label}: file.png]" + assert listed.json()["items"][0]["preview_truncated"] is False + assert detail.json()["preview"] == f"[{label}: file.png]" + assert detail.json()["members"][0]["value"] == url + + +@pytest.mark.parametrize( + "value", + [ + "/private/seed.txt", + r"C:\private\seed.txt", + r"\\server\private\seed.txt", + "https://storage.example.test/seed.txt?sig=secret", + "HTTPS://user:password@example.test/seed.txt?token=secret", + " \thttps://example.test/seed.txt?token=secret", + ], +) +async def test_seed_example_preview_hides_text_references_async( + client: AsyncClient, sqlite_instance: MemoryInterface, value: str +) -> None: + seed = _prompt(value, data_type="text") + await _store(sqlite_instance, seed) + + listed = await client.get(URL, params={"selection_key": NAMED}) + detail = await client.get(f"{URL}/{seed.id}", params={"selection_key": NAMED}) + + assert listed.status_code == detail.status_code == 200 + item = listed.json()["items"][0] + assert item["preview"] == "[Text reference]" + assert item["preview_truncated"] is False + assert detail.json()["members"][0]["value"] == value + + +@pytest.mark.parametrize("value", ["User: describe the image", "Discuss /private/example without opening it"]) +async def test_seed_example_preview_preserves_ordinary_prose_async( + client: AsyncClient, sqlite_instance: MemoryInterface, value: str +) -> None: + await _store(sqlite_instance, _prompt(value)) + + response = await client.get(URL, params={"selection_key": NAMED}) + + assert response.status_code == 200 + assert response.json()["items"][0]["preview"] == value + + +async def test_get_seed_example_returns_full_text_of_truncated_preview(client: AsyncClient, sqlite_instance): + seed = _prompt("x" * 150) + await _store(sqlite_instance, seed) + + listed = (await client.get(URL, params={"selection_key": NAMED})).json() + detail = (await client.get(f"{URL}/{seed.id}", params={"selection_key": NAMED})).json() + + assert listed["items"][0]["preview_truncated"] is True + assert detail["preview_truncated"] is True + assert detail["members"][0]["value"] == "x" * 150 + + +async def test_get_seed_example_returns_objective_conditions(client: AsyncClient, sqlite_instance): + objective = SeedObjective( + value="question", + dataset_name=DATASET, + added_by="test", + conditions=(AnswerMatches(correct_answer="Paris", correct_answer_label="A"), MatchesObjective()), + ) + await _store(sqlite_instance, objective) + + response = await client.get(f"{URL}/{objective.id}", params={"selection_key": NAMED}) + + conditions = response.json()["members"][0]["conditions"] + assert [condition["condition_type"] for condition in conditions] == ["answer_matches", "matches_objective"] + assert conditions[0]["correct_answer"] == "Paris" + assert conditions[0]["correct_answer_label"] == "A" + + +@pytest.mark.parametrize( + "params", + [{"selection_key": "dataset:named:empty"}, {"selection_key": NAMED, "search": "no match"}], +) +async def test_list_seed_examples_returns_empty_page(client: AsyncClient, sqlite_instance, params: dict[str, str]): + await _store(sqlite_instance, _prompt("prompt")) + + response = await client.get(URL, params=params) + + assert response.status_code == 200 + body = response.json() + assert (body["items"], body["total"]) == ([], 0) + assert (body["pagination"]["has_more"], body["pagination"]["next_cursor"]) == (False, None) + + +async def test_get_seed_example_returns_all_stored_members(client: AsyncClient, sqlite_instance): + group = uuid4() + objective = SeedObjective( + value="objective", dataset_name=DATASET, prompt_group_id=group, harm_categories=["violence"], added_by="test" + ) + prompt = _prompt("prompt", prompt_group_id=group, metadata={"source_id": 7}, sequence=1) + configuration = _conversation(prompt_group_id=group) + await _store(sqlite_instance, prompt, configuration, objective) + + response = await client.get(f"{URL}/{group}", params={"selection_key": NAMED}) + unnamed = await client.get(f"{URL}/{group}", params={"selection_key": "dataset:unnamed"}) + + body = response.json() + members = {member["id"]: member for member in body["members"]} + assert body["members"][0]["id"] == str(objective.id) + assert set(members) == {str(objective.id), str(prompt.id), str(configuration.id)} + assert members[str(prompt.id)]["metadata"] == {"source_id": 7} + assert members[str(configuration.id)]["seed_type"] == "simulated_conversation" + assert members[str(configuration.id)]["value"] == configuration.value + stored_configuration = json.loads(members[str(configuration.id)]["value"]) + assert stored_configuration["num_turns"] == 2 + assert stored_configuration["adversarial_chat_system_prompt"]["value"] == "adversarial" + assert (body["preview"], body["objective_count"], body["has_unlabeled_harm"]) == ("objective", 1, True) + assert unnamed.status_code == 404 + + +async def test_seed_example_routes_preserve_legacy_configuration_without_reconstruction_async( + client: AsyncClient, sqlite_instance: MemoryInterface, tmp_path: Path +) -> None: + group = uuid4() + prompt = _prompt("related prompt", prompt_group_id=group) + configuration = SeedEntry(entry=_conversation(prompt_group_id=group)) + configuration.value = json.dumps( + {"num_turns": 2, "adversarial_chat_system_prompt_path": str(tmp_path / "missing.yaml")} + ) + configuration.value_sha256 = to_sha256(configuration.value) + expected_id, expected_value, expected_hash = configuration.id, configuration.value, configuration.value_sha256 + await _store(sqlite_instance, prompt) + async with await sqlite_instance.get_session_async() as session: + session.add(configuration) + await session.commit() + + with ( + patch.object(SeedEntry, "get_seed", side_effect=AssertionError("seed reconstruction")), + patch.object(Path, "read_text", side_effect=AssertionError("file read")), + patch.object(SeedPrompt, "render_template_value_silent", side_effect=AssertionError("template rendering")), + ): + listed = await client.get(URL, params={"selection_key": NAMED, "seed_type": "simulated_conversation"}) + detail = await client.get(f"{URL}/{group}", params={"selection_key": NAMED}) + + assert listed.status_code == detail.status_code == 200 + assert listed.json()["total"] == 1 + assert listed.json()["items"][0]["piece_count"] == 2 + assert listed.json()["items"][0]["seed_types"] == ["prompt", "simulated_conversation"] + members = {member["id"]: member for member in detail.json()["members"]} + assert set(members) == {str(prompt.id), str(expected_id)} + assert members[str(expected_id)]["value"] == expected_value + assert members[str(expected_id)]["value_sha256"] == expected_hash diff --git a/tests/unit/common/test_lazy_package_imports.py b/tests/unit/common/test_lazy_package_imports.py index 5c44139d9b..5a4cfa9d8c 100644 --- a/tests/unit/common/test_lazy_package_imports.py +++ b/tests/unit/common/test_lazy_package_imports.py @@ -46,6 +46,12 @@ "pyrit.models.seeds.seed_dataset_summary", "pyrit.memory", ), + ( + "pyrit.models", + "SeedRecord", + "pyrit.models.seeds.seed_record", + "pyrit.memory", + ), ( "pyrit.models.catalog", "RegisteredInitializer", diff --git a/tests/unit/memory/memory_interface/test_interface_seed_examples.py b/tests/unit/memory/memory_interface/test_interface_seed_examples.py new file mode 100644 index 0000000000..cadce48a55 --- /dev/null +++ b/tests/unit/memory/memory_interface/test_interface_seed_examples.py @@ -0,0 +1,310 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +import json +from datetime import UTC, datetime, timedelta +from pathlib import Path +from typing import Any +from unittest.mock import patch +from uuid import UUID, uuid4 + +import pytest +from sqlalchemy import event +from sqlalchemy.dialects import mssql, sqlite +from sqlalchemy.sql import Select + +from pyrit.common.pagination import DecodedKeysetCursor +from pyrit.common.utils import to_sha256 +from pyrit.memory import MemoryInterface +from pyrit.memory.memory_models import SeedEntry +from pyrit.models import Seed, SeedObjective, SeedPrompt, SeedRecord, SeedSimulatedConversation + +DATASET = "browse" +T0 = datetime(2024, 1, 1, tzinfo=UTC) + + +async def _add(memory: MemoryInterface, *seeds: Seed) -> None: + await memory.add_seeds_to_memory_async(seeds=list(seeds), added_by="test") + + +async def _store(memory: MemoryInterface, entry: SeedEntry) -> None: + async with await memory.get_session_async() as session: + session.add(entry) + await session.commit() + + +async def _ids(memory: MemoryInterface, dataset_name: str | None = DATASET, **filters) -> list[UUID]: + examples, _, _ = await memory.get_seed_examples_async(dataset_name=dataset_name, limit=100, **filters) + return list(examples) + + +async def test_get_seed_examples_orders_by_first_date_then_id_and_seeks(sqlite_instance: MemoryInterface): + tied_low = UUID("00000000-0000-0000-0000-000000000001") + tied_high = UUID("ffffffff-ffff-ffff-ffff-000000000000") + newest, oldest = uuid4(), uuid4() + await _add( + sqlite_instance, + SeedPrompt(value="newest late", dataset_name=DATASET, prompt_group_id=newest, date_added=T0 + timedelta(3)), + SeedPrompt(value="newest", dataset_name=DATASET, prompt_group_id=newest, date_added=T0 + timedelta(2)), + SeedPrompt(value="tied low", dataset_name=DATASET, prompt_group_id=tied_low, date_added=T0 + timedelta(1)), + SeedPrompt(value="tied high", dataset_name=DATASET, prompt_group_id=tied_high, date_added=T0 + timedelta(1)), + SeedPrompt(value="oldest", dataset_name=DATASET, id=oldest, date_added=T0), + ) + + first, first_total, after = await sqlite_instance.get_seed_examples_async(dataset_name=DATASET, limit=2) + second, second_total, last = await sqlite_instance.get_seed_examples_async( + dataset_name=DATASET, limit=2, after=after + ) + + assert [*first, *second] == [newest, tied_high, tied_low, oldest] + assert after is not None + assert (after.timestamp, after.identifier) == (T0 + timedelta(1), str(tied_high)) + assert last is None + assert first_total == second_total == 4 + + +async def test_get_seed_examples_uses_textual_uuid_keys_on_both_dialects_async( + sqlite_instance: MemoryInterface, +) -> None: + seed = SeedPrompt(value="prompt", dataset_name=DATASET, date_added=T0) + await _add(sqlite_instance, seed) + statements: list[Select[Any]] = [] + + def capture( + _connection: Any, + statement: Any, + _multiparams: Any, + _params: Any, + _execution_options: Any, + ) -> None: + if isinstance(statement, Select): + statements.append(statement) + + engine = sqlite_instance._get_async_engine().sync_engine + event.listen(engine, "before_execute", capture) + try: + await sqlite_instance.get_seed_examples_async( + dataset_name=DATASET, + limit=1, + after=DecodedKeysetCursor(timestamp=T0, identifier="ffffffff-ffff-ffff-ffff-ffffffffffff"), + ) + finally: + event.remove(engine, "before_execute", capture) + + assert len(statements) == 3 + for dialect in (mssql.dialect(), sqlite.dialect()): + page = str(statements[1].compile(dialect=dialect, compile_kwargs={"literal_binds": True})).lower() + members = str(statements[2].compile(dialect=dialect, compile_kwargs={"literal_binds": True})).lower() + assert "lower(cast(coalesce(" in page + assert "as varchar(36)" in page + assert "example_id_key < 'ffffffff-ffff-ffff-ffff-ffffffffffff'" in page + assert "example_id_key desc" in page + assert "order by" in members and "lower(cast(" in members + + +async def test_get_seed_examples_orders_filtered_examples_by_first_date_of_all_members( + sqlite_instance: MemoryInterface, +): + group, single = uuid4(), uuid4() + await _add( + sqlite_instance, + SeedPrompt(value="old", dataset_name=DATASET, prompt_group_id=group, date_added=T0), + SeedPrompt(value="needle", dataset_name=DATASET, prompt_group_id=group, date_added=T0 + timedelta(3)), + SeedPrompt(value="needle", dataset_name=DATASET, id=single, date_added=T0 + timedelta(1)), + ) + + assert await _ids(sqlite_instance, value_search="needle") == [single, group] + + +async def test_get_seed_examples_returns_complete_groups_objective_first(sqlite_instance: MemoryInterface): + group = uuid4() + objective = SeedObjective(value="objective", dataset_name=DATASET, prompt_group_id=group) + second = SeedPrompt(value="second", dataset_name=DATASET, prompt_group_id=group, sequence=1) + first = SeedPrompt(value="first", dataset_name=DATASET, prompt_group_id=group, sequence=0) + await _add(sqlite_instance, second, first, objective) + + examples, total, _ = await sqlite_instance.get_seed_examples_async( + dataset_name=DATASET, limit=10, value_search="second" + ) + + assert total == 1 + assert [seed.id for seed in examples[group]] == [objective.id, first.id, second.id] + assert isinstance(examples[group][0], SeedRecord) + assert examples[group][0].seed_type == "objective" + + +async def test_get_seed_examples_filters_match_across_members(sqlite_instance: MemoryInterface): + both, harm_only = uuid4(), uuid4() + await _add( + sqlite_instance, + SeedObjective(value="objective", dataset_name=DATASET, prompt_group_id=both, harm_categories=["Violence"]), + SeedPrompt(value="say hello", dataset_name=DATASET, prompt_group_id=both), + SeedObjective(value="objective", dataset_name=DATASET, prompt_group_id=harm_only, harm_categories=["Violence"]), + SeedPrompt(value="goodbye", dataset_name=DATASET, prompt_group_id=harm_only), + ) + + assert set(await _ids(sqlite_instance, harm_categories=["missing", "VIOLENCE"])) == {both, harm_only} + assert await _ids(sqlite_instance, harm_categories=["violence"], value_search="HELLO") == [both] + assert await _ids(sqlite_instance, seed_types=["objective"], data_types=["image_path"]) == [] + + +async def test_get_seed_examples_keeps_named_and_unnamed_scopes_apart(sqlite_instance: MemoryInterface): + group = uuid4() + named = SeedPrompt(value="named", dataset_name=DATASET, prompt_group_id=group) + null_name = SeedPrompt(value="null name", prompt_group_id=group) + empty_name = SeedPrompt(value="empty name", dataset_name="", prompt_group_id=group) + await _add(sqlite_instance, named, null_name, empty_name) + + named_examples, _, _ = await sqlite_instance.get_seed_examples_async(dataset_name=DATASET, limit=10) + unnamed_examples, _, _ = await sqlite_instance.get_seed_examples_async(dataset_name=None, limit=10) + + assert [seed.id for seed in named_examples[group]] == [named.id] + assert {seed.id for seed in unnamed_examples[group]} == {null_name.id, empty_name.id} + assert await _ids(sqlite_instance, dataset_name="other") == [] + + +@pytest.mark.parametrize( + ("search", "expected"), + [("%", ["50% off"]), ("_", ["snake_case"]), ("\\", ["a\\b"]), ("[x]", ["[x]"]), ("0_", []), ("x%", [])], +) +async def test_get_seed_examples_search_is_literal(sqlite_instance: MemoryInterface, search: str, expected: list[str]): + await _add( + sqlite_instance, + *(SeedPrompt(value=value, dataset_name=DATASET) for value in ["50% off", "snake_case", "a\\b", "[x]"]), + ) + + examples, _, _ = await sqlite_instance.get_seed_examples_async(dataset_name=DATASET, limit=10, value_search=search) + + assert [seeds[0].value for seeds in examples.values()] == expected + + +async def test_get_seed_examples_search_ignores_media_values(sqlite_instance: MemoryInterface): + media = SeedPrompt(value="/data/needle.png", data_type="image_path", dataset_name=DATASET, added_by="test") + await _store(sqlite_instance, SeedEntry(entry=media)) + + assert await _ids(sqlite_instance, value_search="needle") == [] + + +async def test_get_seed_examples_search_ignores_simulated_conversation_json(sqlite_instance: MemoryInterface): + group = uuid4() + standalone = SeedSimulatedConversation( + num_turns=2, adversarial_chat_system_prompt=SeedPrompt(value="needle"), dataset_name=DATASET + ) + grouped = SeedSimulatedConversation( + num_turns=2, + adversarial_chat_system_prompt=SeedPrompt(value="needle"), + dataset_name=DATASET, + prompt_group_id=group, + ) + objective = SeedObjective(value="objective", dataset_name=DATASET, prompt_group_id=group) + await _add(sqlite_instance, standalone, grouped, objective) + + assert await _ids(sqlite_instance, value_search="needle") == [] + assert await _ids(sqlite_instance, value_search="num_turns") == [] + assert await _ids(sqlite_instance, value_search="objective") == [group] + assert set(await _ids(sqlite_instance, seed_types=["simulated_conversation"])) == {standalone.id, group} + + +def _legacy_conversation_entry(*, prompt_file: Path, prompt_group_id: UUID) -> SeedEntry: + entry = SeedEntry( + entry=SeedSimulatedConversation( + num_turns=2, + adversarial_chat_system_prompt=SeedPrompt(value="placeholder"), + dataset_name=DATASET, + prompt_group_id=prompt_group_id, + added_by="test", + ) + ) + entry.value = json.dumps( + {"num_turns": 2, "sequence": 0, "adversarial_chat_system_prompt_path": str(prompt_file)}, + sort_keys=True, + separators=(",", ":"), + ) + entry.value_sha256 = to_sha256(entry.value) + return entry + + +@pytest.mark.parametrize("file_exists", [True, False]) +async def test_get_seed_examples_preserves_legacy_members_without_loading_files_async( + sqlite_instance: MemoryInterface, tmp_path: Path, file_exists: bool +) -> None: + mixed, standalone = uuid4(), uuid4() + prompt_file = tmp_path / "legacy.yaml" + if file_exists: + prompt_file.write_text('value: "{{ 1 + 1 }}"\ndata_type: text', encoding="utf-8") + prompt = SeedPrompt(value="readable", dataset_name=DATASET, prompt_group_id=mixed) + await _add(sqlite_instance, prompt) + mixed_entry = _legacy_conversation_entry(prompt_file=prompt_file, prompt_group_id=mixed) + standalone_entry = _legacy_conversation_entry(prompt_file=prompt_file, prompt_group_id=standalone) + mixed_id, standalone_id = mixed_entry.id, standalone_entry.id + expected = {entry.id: (entry.value, entry.value_sha256) for entry in (mixed_entry, standalone_entry)} + await _store(sqlite_instance, mixed_entry) + await _store(sqlite_instance, standalone_entry) + + with ( + patch.object(SeedEntry, "get_seed", side_effect=AssertionError("seed reconstruction")), + patch.object(Path, "read_text", side_effect=AssertionError("file read")), + patch.object(SeedPrompt, "render_template_value_silent", side_effect=AssertionError("template rendering")), + ): + examples, total, _ = await sqlite_instance.get_seed_examples_async( + dataset_name=DATASET, limit=10, seed_types=["simulated_conversation"] + ) + detail = await sqlite_instance.get_seed_example_async(dataset_name=DATASET, example_id=mixed) + + assert {seed.id for seed in examples[mixed]} == {prompt.id, mixed_id} + assert [seed.id for seed in examples[standalone]] == [standalone_id] + assert detail == examples[mixed] + assert total == 2 + for members in examples.values(): + for record in members: + if record.seed_type == "simulated_conversation": + assert (record.value, record.value_sha256) == expected[record.id] + + +async def test_get_seed_examples_preserves_simulated_configuration(sqlite_instance: MemoryInterface): + conversation = SeedSimulatedConversation( + num_turns=2, + adversarial_chat_system_prompt=SeedPrompt(value="adversarial", parameters=["objective"]), + simulated_target_system_prompt=SeedPrompt(value="target", parameters=["objective", "num_turns"]), + dataset_name=DATASET, + ) + await _add(sqlite_instance, conversation) + + seeds = await sqlite_instance.get_seed_example_async(dataset_name=DATASET, example_id=conversation.id) + + assert len(seeds) == 1 + assert isinstance(seeds[0], SeedRecord) + assert seeds[0].seed_type == "simulated_conversation" + assert seeds[0].value == conversation.value + assert seeds[0].value_sha256 == conversation.value_sha256 + + +async def test_get_seed_examples_keeps_unparseable_configuration_inspectable_async( + sqlite_instance: MemoryInterface, tmp_path: Path +) -> None: + group = uuid4() + entry = _legacy_conversation_entry(prompt_file=tmp_path / "missing.yaml", prompt_group_id=group) + entry.value = "unparseable stored configuration" + expected_value, expected_hash = entry.value, entry.value_sha256 + await _store(sqlite_instance, entry) + + examples, total, _ = await sqlite_instance.get_seed_examples_async(dataset_name=DATASET, limit=10) + detail = await sqlite_instance.get_seed_example_async(dataset_name=DATASET, example_id=group) + + assert total == 1 + assert detail == examples[group] + assert detail[0].value == expected_value + assert detail[0].value_sha256 == expected_hash + + +async def test_get_seed_example_returns_empty_outside_its_dataset(sqlite_instance: MemoryInterface): + group = uuid4() + objective = SeedObjective(value="objective", dataset_name=DATASET, prompt_group_id=group, date_added=T0) + await _add(sqlite_instance, objective, SeedPrompt(value="prompt", dataset_name=DATASET, prompt_group_id=group)) + + seeds = await sqlite_instance.get_seed_example_async(dataset_name=DATASET, example_id=group) + + assert seeds[0].id == objective.id + assert len(seeds) == 2 + assert await sqlite_instance.get_seed_example_async(dataset_name="other", example_id=group) == [] + assert await sqlite_instance.get_seed_example_async(dataset_name=None, example_id=group) == [] diff --git a/tests/unit/memory/test_memory_models.py b/tests/unit/memory/test_memory_models.py index 297f611f89..cbc062939c 100644 --- a/tests/unit/memory/test_memory_models.py +++ b/tests/unit/memory/test_memory_models.py @@ -58,6 +58,7 @@ SeedIdentifier, SeedObjective, SeedPrompt, + SeedRecord, SeedSimulatedConversation, TargetIdentifier, ) @@ -550,6 +551,32 @@ def test_seed_prompt_preserves_parameters(self): entry = SeedEntry(entry=seed) assert entry.parameters == ["param1", "param2"] + def test_get_seed_record_preserves_stored_fields_without_generated_defaults(self) -> None: + seed = _make_seed_prompt( + value="{{ unchanged }}", + value_sha256="stored-hash", + parameters=["unchanged"], + metadata={"source_id": 7}, + response_json_schema={"type": "object"}, + ) + entry = SeedEntry(entry=seed) + entry.sequence = None + + record = entry.get_seed_record() + + assert isinstance(record, SeedRecord) + assert record.id == seed.id + assert record.value == "{{ unchanged }}" + assert record.value_sha256 == "stored-hash" + assert record.parameters == ["unchanged"] + assert record.sequence is None + assert record.prompt_group_id is None + assert record.metadata == {"source_id": 7} + assert record.response_json_schema == {"type": "object"} + assert record.date_added == seed.date_added + assert record.added_by == seed.added_by + assert "is_jinja_template" not in record.model_dump() + # ---- response_json_schema persistence --------------------------------- def test_roundtrip_seed_prompt_preserves_inline_response_json_schema(self):