From 481ff6472f3652d94469f666cca2c89b0b07ddc2 Mon Sep 17 00:00:00 2001 From: Jackson Severino da Rocha Date: Mon, 28 Sep 2026 11:00:59 -0300 Subject: [PATCH 1/9] test: add seed browsing contract coverage --- ..._dataset_service_seed_browsing_contract.py | 217 +++++++++ tests/unit/backend/test_seed_browsing_api.py | 434 ++++++++++++++++++ .../test_interface_seed_browsing_contract.py | 255 ++++++++++ 3 files changed, 906 insertions(+) create mode 100644 tests/unit/backend/test_dataset_service_seed_browsing_contract.py create mode 100644 tests/unit/backend/test_seed_browsing_api.py create mode 100644 tests/unit/memory/memory_interface/test_interface_seed_browsing_contract.py diff --git a/tests/unit/backend/test_dataset_service_seed_browsing_contract.py b/tests/unit/backend/test_dataset_service_seed_browsing_contract.py new file mode 100644 index 0000000000..3c419d4e55 --- /dev/null +++ b/tests/unit/backend/test_dataset_service_seed_browsing_contract.py @@ -0,0 +1,217 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""RED contract tests for DatasetService seed browsing policy (#2748).""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Any +from unittest.mock import patch +from uuid import uuid4 + +import pytest + +from pyrit.backend.models.common import PaginationInfo +from pyrit.backend.services.dataset_service import DatasetService +from pyrit.models import SeedObjective, SeedPrompt + +if TYPE_CHECKING: + from pyrit.memory import MemoryInterface + + +DATASET = "service-browse-contract" +SELECTION_KEY = f"dataset:named:{DATASET}" + + +def _field(value: Any, name: str) -> Any: + """Read a contract field from either a Pydantic result or a mapping.""" + return value.get(name) if isinstance(value, dict) else getattr(value, name) + + +def _service_method(service: DatasetService, name: str): + method = getattr(service, name, None) + assert method is not None, f"RED: DatasetService.{name} is not implemented" + return method + + +async def _add(memory: MemoryInterface, *seeds: SeedPrompt | SeedObjective) -> None: + await memory.add_seeds_to_memory_async(seeds=list(seeds), added_by="2748-service-test") + + +@pytest.fixture +def dataset_service(sqlite_instance: MemoryInterface) -> DatasetService: + with patch( + "pyrit.backend.services.dataset_service.CentralMemory.get_memory_instance", return_value=sqlite_instance + ): + yield DatasetService() + + +class TestDatasetServiceSeedBrowsingContract: + async def test_list_resolves_selection_key_and_returns_pagination_info( + self, dataset_service: DatasetService, sqlite_instance: MemoryInterface + ): + await _add(sqlite_instance, SeedPrompt(value="prompt", dataset_name=DATASET)) + response = await _service_method(dataset_service, "list_seed_examples_async")( + selection_key=SELECTION_KEY, limit=10 + ) + assert isinstance(_field(response, "pagination"), PaginationInfo) + assert _field(_field(response, "pagination"), "limit") == 10 + assert len(_field(response, "items")) == 1 + + async def test_unnamed_selection_key_is_supported_without_display_name_resolution( + self, dataset_service: DatasetService, sqlite_instance: MemoryInterface + ): + await _add(sqlite_instance, SeedPrompt(value="unnamed")) + response = await _service_method(dataset_service, "list_seed_examples_async")( + selection_key="dataset:unnamed", limit=10 + ) + assert len(_field(response, "items")) == 1 + + async def test_cursor_is_bound_to_selection_and_effective_filters( + self, dataset_service: DatasetService, sqlite_instance: MemoryInterface + ): + await _add(sqlite_instance, *(SeedPrompt(value=str(index), dataset_name=DATASET) for index in range(3))) + first = await _service_method(dataset_service, "list_seed_examples_async")(selection_key=SELECTION_KEY, limit=1) + cursor = _field(_field(first, "pagination"), "next_cursor") + assert cursor + with pytest.raises(ValueError): + await _service_method(dataset_service, "list_seed_examples_async")( + selection_key="dataset:unnamed", limit=1, cursor=cursor + ) + with pytest.raises(ValueError): + await _service_method(dataset_service, "list_seed_examples_async")( + selection_key=SELECTION_KEY, limit=1, cursor=cursor, search="changed" + ) + + async def test_malformed_cursor_and_invalid_selection_are_rejected(self, dataset_service: DatasetService): + with pytest.raises(ValueError): + await _service_method(dataset_service, "list_seed_examples_async")( + selection_key=SELECTION_KEY, limit=10, cursor="malformed" + ) + with pytest.raises(ValueError): + await _service_method(dataset_service, "list_seed_examples_async")( + selection_key="display-name-not-selection-key", limit=10 + ) + + async def test_list_formats_group_preview_types_modalities_counts_and_harm_summary( + self, dataset_service: DatasetService, sqlite_instance: MemoryInterface + ): + group_id = uuid4() + await _add( + sqlite_instance, + SeedPrompt( + value="prompt", + dataset_name=DATASET, + prompt_group_id=group_id, + data_type="text", + harm_categories=["violence"], + ), + SeedObjective(value="objective", dataset_name=DATASET, prompt_group_id=group_id), + SeedPrompt(value="".join("x" for _ in range(101)), dataset_name=DATASET, prompt_group_id=group_id), + ) + response = await _service_method(dataset_service, "list_seed_examples_async")( + selection_key=SELECTION_KEY, limit=10 + ) + item = _field(response, "items")[0] + assert _field(item, "preview_truncated") is True + assert len(_field(item, "preview")) <= 103 + assert _field(item, "piece_count") == 3 + assert _field(item, "objective_count") == 1 + assert "text" in _field(item, "modalities") + assert "violence" in _field(item, "harm_categories") + assert "objective" in _field(item, "seed_types") + + async def test_detail_returns_complete_group_and_persisted_provenance( + self, dataset_service: DatasetService, sqlite_instance: MemoryInterface + ): + group_id = uuid4() + seed = SeedPrompt( + value="full text", + dataset_name=DATASET, + prompt_group_id=group_id, + role="user", + sequence=4, + source="source", + authors=["author"], + groups=["group"], + metadata={"persisted": True}, + parameters=["name"], + ) + await _add( + sqlite_instance, seed, SeedObjective(value="condition", dataset_name=DATASET, prompt_group_id=group_id) + ) + detail = await _service_method(dataset_service, "get_seed_example_async")( + selection_key=SELECTION_KEY, example_id=str(group_id) + ) + members = _field(detail, "members") + assert len(members) == 2 + prompt = next(member for member in members if _field(member, "id") == seed.id) + assert _field(prompt, "prompt_group_id") == group_id + assert _field(prompt, "value") == "full text" + assert _field(prompt, "role") == "user" + assert _field(prompt, "sequence") == 4 + for field in ( + "value_sha256", + "dataset_name", + "source", + "authors", + "groups", + "date_added", + "added_by", + "metadata", + "data_type", + ): + assert field in (prompt if isinstance(prompt, dict) else prompt.model_fields_set | set(prompt.model_fields)) + + async def test_detail_preserves_template_parameters_and_objective_conditions( + self, dataset_service: DatasetService, sqlite_instance: MemoryInterface + ): + group_id = uuid4() + await _add( + sqlite_instance, + SeedPrompt( + value="{{ name }}", + dataset_name=DATASET, + prompt_group_id=group_id, + is_jinja_template=True, + parameters=["name"], + ), + SeedObjective(value="condition", dataset_name=DATASET, prompt_group_id=group_id), + ) + detail = await _service_method(dataset_service, "get_seed_example_async")( + selection_key=SELECTION_KEY, example_id=str(group_id) + ) + assert _field(detail, "members") + assert any(_field(member, "parameters") == ["name"] for member in _field(detail, "members")) + assert any(_field(member, "seed_type") == "objective" for member in _field(detail, "members")) + + async def test_invalid_detail_does_not_generate_group_identity(self, dataset_service: DatasetService): + with pytest.raises(ValueError): + await _service_method(dataset_service, "get_seed_example_async")( + selection_key=SELECTION_KEY, example_id=str(uuid4()) + ) + + async def test_browsing_has_no_template_or_generation_side_effects( + self, dataset_service: DatasetService, sqlite_instance: MemoryInterface + ): + await _add(sqlite_instance, SeedPrompt(value="{{ dangerous }}", dataset_name=DATASET, is_jinja_template=True)) + with patch("pyrit.models.SeedPrompt.render_template_value", side_effect=AssertionError("rendered")) as render: + response = await _service_method(dataset_service, "list_seed_examples_async")( + selection_key=SELECTION_KEY, limit=10 + ) + assert _field(response, "items") + assert render.call_count == 0 + + async def test_media_preview_is_type_label_without_bytes_or_path_leak( + self, dataset_service: DatasetService, sqlite_instance: MemoryInterface, tmp_path + ): + media = tmp_path / "private-image.png" + media.write_bytes(b"local image") + await _add(sqlite_instance, SeedPrompt(value=str(media), dataset_name=DATASET, data_type="image_path")) + response = await _service_method(dataset_service, "list_seed_examples_async")( + selection_key=SELECTION_KEY, limit=10 + ) + item = _field(response, "items")[0] + assert "image" in _field(item, "preview").lower() + assert str(tmp_path) not in _field(item, "preview") + assert "bytes" not in (item if isinstance(item, dict) else item.model_dump()) diff --git a/tests/unit/backend/test_seed_browsing_api.py b/tests/unit/backend/test_seed_browsing_api.py new file mode 100644 index 0000000000..b5a416e799 --- /dev/null +++ b/tests/unit/backend/test_seed_browsing_api.py @@ -0,0 +1,434 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""RED contract tests for the paginated seed browsing API (#2748). + +These tests deliberately target the approved public HTTP contract. The route, +models, and narrow memory read helpers are not present until the feature is +implemented; consequently the current expected failure is a missing-route +response, rather than a test-time production stub. +""" + +from __future__ import annotations + +from datetime import UTC, datetime, timedelta +from typing import TYPE_CHECKING +from unittest.mock import patch +from uuid import UUID, uuid4 + +import pytest +from fastapi.testclient import TestClient + +from pyrit.backend.main import app +from pyrit.models import SeedObjective, SeedPrompt, SeedSimulatedConversation + +if TYPE_CHECKING: + from pyrit.memory import MemoryInterface + + +DATASET = "browse-contract" +NAMED_KEY = f"dataset:named:{DATASET}" +UNNAMED_KEY = "dataset:unnamed" + + +@pytest.fixture +def client(patch_central_database) -> TestClient: + """Use the real SQLite memory fixture behind the API application.""" + return TestClient(app) + + +async def _add(memory: MemoryInterface, *seeds: SeedPrompt | SeedObjective) -> None: + await memory.add_seeds_to_memory_async(seeds=list(seeds), added_by="2748-test") + + +def _list(client: TestClient, selection_key: str = NAMED_KEY, **params: object): + return client.get(f"/api/datasets/{selection_key}/seeds", params=params) + + +def _detail(client: TestClient, example_id: str, selection_key: str = NAMED_KEY): + return client.get(f"/api/datasets/{selection_key}/seeds/{example_id}") + + +def _items(response): + assert response.status_code == 200, response.text + return response.json()["items"] + + +class TestEmptyAndDatasetSelection: + async def test_empty_dataset_is_a_valid_empty_page(self, client, sqlite_instance: MemoryInterface): + response = _list(client) + assert response.status_code == 200 + body = response.json() + assert body["items"] == [] + assert body["pagination"]["has_more"] is False + assert body["pagination"]["next_cursor"] is None + + async def test_empty_filter_result_is_not_an_error(self, client, sqlite_instance: MemoryInterface): + await _add(sqlite_instance, SeedPrompt(value="ordinary", dataset_name=DATASET)) + response = _list(client, search="absent") + assert response.status_code == 200 + assert response.json()["items"] == [] + + async def test_named_selection_key_is_not_display_name(self, client, sqlite_instance: MemoryInterface): + await _add(sqlite_instance, SeedPrompt(value="named", dataset_name=DATASET)) + assert len(_items(_list(client, NAMED_KEY))) == 1 + assert _list(client, DATASET).status_code in {400, 404, 422} + + async def test_unnamed_selection_and_invalid_selection(self, client, sqlite_instance: MemoryInterface): + await _add(sqlite_instance, SeedPrompt(value="unnamed")) + assert len(_items(_list(client, UNNAMED_KEY))) == 1 + assert _list(client, "dataset:named:not-loaded").status_code in {400, 404, 422} + + async def test_named_unnamed_namespace_is_distinct(self, client, sqlite_instance: MemoryInterface): + await _add(sqlite_instance, SeedPrompt(value="literal", dataset_name="__unnamed__"), SeedPrompt(value="none")) + assert len(_items(_list(client, "dataset:named:__unnamed__"))) == 1 + assert len(_items(_list(client, UNNAMED_KEY))) == 1 + + +class TestPaginationAndIdentity: + async def test_one_page_has_existing_pagination_shape(self, client, sqlite_instance: MemoryInterface): + await _add(sqlite_instance, SeedPrompt(value="one", dataset_name=DATASET)) + page = _list(client, limit=1) + assert page.status_code == 200 + assert set(page.json()["pagination"]) >= {"limit", "has_more", "next_cursor", "prev_cursor"} + + async def test_multiple_pages_have_no_duplicates_or_omissions(self, client, sqlite_instance: MemoryInterface): + seeds = [SeedPrompt(value=f"prompt-{i}", dataset_name=DATASET) for i in range(5)] + await _add(sqlite_instance, *seeds) + first = _list(client, limit=2) + first_items = _items(first) + second = _list(client, limit=2, cursor=first.json()["pagination"]["next_cursor"]) + all_items = first_items + _items(second) + while second.json()["pagination"]["has_more"]: + second = _list(client, limit=2, cursor=second.json()["pagination"]["next_cursor"]) + all_items += _items(second) + ids = [item["example_id"] for item in all_items] + assert len(ids) == len(set(ids)) == 5 + + @pytest.mark.parametrize("limit", [0, -1, 101]) + async def test_page_size_is_validated(self, client, sqlite_instance: MemoryInterface, limit: int): + assert _list(client, limit=limit).status_code in {400, 422} + + async def test_group_identity_preserves_ids_and_does_not_hash_merge(self, client, sqlite_instance: MemoryInterface): + group_id = uuid4() + first = SeedPrompt(value="same", dataset_name=DATASET, prompt_group_id=group_id) + second = SeedPrompt(value="same", dataset_name=DATASET, prompt_group_id=group_id) + ungrouped = SeedPrompt(value="same", dataset_name=DATASET) + await _add(sqlite_instance, first, second, ungrouped) + items = _items(_list(client)) + assert len(items) == 2 + assert {str(first.id), str(second.id), str(ungrouped.id)} == set(items[0]["seed_ids"] + items[1]["seed_ids"]) + assert sorted(len(item["seed_ids"]) for item in items) == [1, 2] + assert all(item["example_id"] in {str(group_id), str(ungrouped.id)} for item in items) + assert all("generated" not in item["example_id"] for item in items) + + async def test_order_is_complete_group_date_then_id(self, client, sqlite_instance: MemoryInterface): + tied = datetime(2024, 1, 1, tzinfo=UTC) + old_group = uuid4() + new_group = uuid4() + await _add( + sqlite_instance, + SeedPrompt(value="old", dataset_name=DATASET, prompt_group_id=old_group, date_added=tied), + SeedPrompt( + value="new", dataset_name=DATASET, prompt_group_id=new_group, date_added=tied + timedelta(days=1) + ), + SeedPrompt( + value="late member", + dataset_name=DATASET, + prompt_group_id=old_group, + date_added=tied + timedelta(days=2), + ), + ) + items = _items(_list(client)) + assert [item["example_id"] for item in items] == [str(new_group), str(old_group)] + + async def test_tied_dates_use_descending_logical_id_tie_breaker(self, client, sqlite_instance: MemoryInterface): + date_added = datetime(2024, 1, 1, tzinfo=UTC) + lower = UUID("00000000-0000-0000-0000-000000000001") + higher = UUID("00000000-0000-0000-0000-000000000002") + await _add( + sqlite_instance, + SeedPrompt(value="lower", dataset_name=DATASET, prompt_group_id=lower, date_added=date_added), + SeedPrompt(value="higher", dataset_name=DATASET, prompt_group_id=higher, date_added=date_added), + ) + assert [item["example_id"] for item in _items(_list(client))] == [str(higher), str(lower)] + + async def test_group_is_never_split_and_detail_preserves_role_sequence( + self, client, sqlite_instance: MemoryInterface, tmp_path + ): + group_id = uuid4() + image_path = tmp_path / "group-image.png" + image_path.write_bytes(b"local test image") + await _add( + sqlite_instance, + SeedPrompt( + value=str(image_path), + dataset_name=DATASET, + prompt_group_id=group_id, + data_type="image_path", + sequence=0, + ), + SeedPrompt(value="text", dataset_name=DATASET, prompt_group_id=group_id, data_type="text", sequence=1), + ) + page = _list(client, limit=1) + assert len(_items(page)) == 1 + detail = _detail(client, str(group_id)) + members = detail.json()["members"] if detail.status_code == 200 else [] + assert [member["sequence"] for member in members] == [0, 1] + assert all("role" in member for member in members) + + +class TestFilters: + async def test_modality_is_or_and_matching_member_returns_complete_group( + self, client, sqlite_instance: MemoryInterface, tmp_path + ): + group_id = uuid4() + image_path = tmp_path / "group-image.png" + audio_path = tmp_path / "standalone-audio.wav" + image_path.write_bytes(b"local test image") + audio_path.write_bytes(b"local test audio") + await _add( + sqlite_instance, + SeedPrompt(value=str(image_path), dataset_name=DATASET, prompt_group_id=group_id, data_type="image_path"), + SeedPrompt(value="text", dataset_name=DATASET, prompt_group_id=group_id, data_type="text"), + SeedPrompt(value=str(audio_path), dataset_name=DATASET, data_type="audio_path"), + ) + items = _items(_list(client, modality=["image_path", "audio_path"])) + assert {item["example_id"] for item in items} == {str(group_id)} | { + str(next(seed.id for seed in sqlite_instance.get_seeds(data_types=["audio_path"]))) + } + assert len(_detail(client, str(group_id)).json()["members"]) == 2 + + @pytest.mark.parametrize("category, expected", [("Violence", True), ("vio", False), ("missing", False)]) + async def test_harm_category_is_case_insensitive_whole_value_and_missing_is_unlabeled( + self, client, sqlite_instance: MemoryInterface, category: str, expected: bool + ): + seed = SeedPrompt(value="harm", dataset_name=DATASET, harm_categories=["VIOLENCE"]) + unlabeled = SeedPrompt(value="unlabeled", dataset_name=DATASET, harm_categories=[]) + await _add(sqlite_instance, seed, unlabeled) + items = _items(_list(client, harm_category=category)) + assert (len(items) == 1 and str(seed.id) in items[0]["seed_ids"]) is expected + all_items = _items(_list(client)) + unlabeled_item = next(item for item in all_items if str(unlabeled.id) in item["seed_ids"]) + assert unlabeled_item["has_unlabeled_harm"] is True + + async def test_multiple_harm_categories_are_or_not_existing_get_seeds_all_semantics(self, client, sqlite_instance): + await _add( + sqlite_instance, + SeedPrompt(value="hate", dataset_name=DATASET, harm_categories=["hate"]), + SeedPrompt(value="violence", dataset_name=DATASET, harm_categories=["violence"]), + SeedPrompt(value="both", dataset_name=DATASET, harm_categories=["hate", "violence"]), + ) + items = _items(_list(client, harm_category=["hate", "violence"])) + assert len(items) == 3 + assert {item["seed_ids"][0] for item in items} == { + str(seed.id) for seed in sqlite_instance.get_seeds(dataset_name=DATASET) + } + + async def test_seed_type_filter_is_or_and_returns_complete_group(self, client, sqlite_instance): + group_id = uuid4() + await _add( + sqlite_instance, + SeedPrompt(value="prompt", dataset_name=DATASET, prompt_group_id=group_id), + SeedObjective(value="objective", dataset_name=DATASET, prompt_group_id=group_id), + ) + items = _items(_list(client, seed_type=["objective"])) + assert len(items) == 1 and items[0]["piece_count"] == 2 + + async def test_filters_are_and_across_filters_but_match_at_example_level(self, client, sqlite_instance, tmp_path): + group_id = uuid4() + image_path = tmp_path / "filter-image.png" + image_path.write_bytes(b"local test image") + await _add( + sqlite_instance, + SeedPrompt(value=str(image_path), dataset_name=DATASET, prompt_group_id=group_id, data_type="image_path"), + SeedPrompt(value="violence", dataset_name=DATASET, prompt_group_id=group_id, harm_categories=["violence"]), + SeedPrompt(value=str(image_path), dataset_name=DATASET, data_type="image_path"), + ) + items = _items(_list(client, modality="image_path", harm_category="violence")) + assert len(items) == 1 and items[0]["example_id"] == str(group_id) + + +class TestTextSearchAndSafety: + async def test_text_search_is_literal_case_insensitive_and_text_only(self, client, sqlite_instance, tmp_path): + image_path = tmp_path / "media-path.png" + image_path.write_bytes(b"local test image") + await _add( + sqlite_instance, + SeedPrompt(value="Need 100% literal_value", dataset_name=DATASET), + SeedObjective(value="OBJECTIVE text", dataset_name=DATASET), + SeedPrompt(value=str(image_path), dataset_name=DATASET, data_type="image_path"), + SeedPrompt(value="metadata-only", dataset_name=DATASET, metadata={"secret": "literal_value"}), + ) + assert len(_items(_list(client, search="100% literal_value"))) == 1 + assert len(_items(_list(client, search="literal_value"))) == 1 + assert len(_items(_list(client, search="objective"))) == 1 + assert _items(_list(client, search="image_path")) == [] + + async def test_cursor_is_opaque_bound_to_dataset_and_effective_filters(self, client, sqlite_instance): + await _add(sqlite_instance, *(SeedPrompt(value=str(i), dataset_name=DATASET) for i in range(3))) + cursor = _list(client, limit=1).json()["pagination"]["next_cursor"] + assert cursor and not cursor.startswith("1") + assert _list(client, limit=1, cursor="not-a-cursor").status_code in {400, 422} + assert _list(client, UNNAMED_KEY, limit=1, cursor=cursor).status_code in {400, 422} + assert _list(client, limit=1, search="different", cursor=cursor).status_code in {400, 422} + + async def test_missing_detail_example_is_not_silently_empty(self, client, sqlite_instance): + response = _detail(client, str(uuid4())) + assert response.status_code in {404, 422} + + async def test_template_is_not_rendered_or_loaded(self, client, sqlite_instance): + template = SeedPrompt( + value="{{ dangerous }}", dataset_name=DATASET, is_jinja_template=True, parameters=["dangerous"] + ) + await _add(sqlite_instance, template) + with patch.object(SeedPrompt, "render_template_value", side_effect=AssertionError("rendered")) as render: + with patch("pathlib.Path.read_text", side_effect=AssertionError("loaded")) as load: + response = _list(client, search="dangerous") + assert response.status_code == 200 + assert render.call_count == load.call_count == 0 + item = response.json()["items"][0] + assert item["is_template"] is True + assert item["parameters"] == ["dangerous"] + + async def test_simulated_configuration_is_returned_without_generation_or_target_call(self, client, sqlite_instance): + config = SeedSimulatedConversation( + dataset_name=DATASET, + adversarial_chat_system_prompt=SeedPrompt(value="{{ objective }}"), + simulated_target_system_prompt=SeedPrompt(value="{{ objective }}"), + num_turns=2, + ) + await _add(sqlite_instance, config) + with patch( + "pyrit.executor.attack.multi_turn.simulated_conversation.generate_simulated_conversation_async" + ) as generate: + response = _list(client, search="num_turns") + assert response.status_code == 200 + assert generate.call_count == 0 + assert response.json()["items"][0]["seed_types"] == ["simulated_conversation"] + + +class TestPreviewDetailAndCounts: + async def test_preview_uses_100_character_convention_and_hides_full_content(self, client, sqlite_instance): + short = "x" * 100 + long = "y" * 101 + await _add( + sqlite_instance, SeedPrompt(value=short, dataset_name=DATASET), SeedPrompt(value=long, dataset_name=DATASET) + ) + items = _items(_list(client)) + assert any(item["preview"] == short and item["preview_truncated"] is False for item in items) + long_item = next(item for item in items if item["preview"].startswith("y")) + assert len(long_item["preview"]) <= 103 and long_item["preview_truncated"] is True + assert long not in long_item["preview"] + + async def test_media_preview_is_label_only_and_never_bytes_path_or_credentials( + self, client, sqlite_instance, tmp_path + ): + image_path = tmp_path / "image.png" + image_path.write_bytes(b"local test image") + await _add(sqlite_instance, SeedPrompt(value=str(image_path), dataset_name=DATASET, data_type="image_path")) + item = _items(_list(client))[0] + assert "image" in item["preview"].lower() + assert str(tmp_path) not in item["preview"] and "sig=" not in item["preview"] + assert "bytes" not in item and "content" not in item + + async def test_detail_returns_all_persisted_fields_without_new_ids_or_rendering(self, client, sqlite_instance): + seed_id = uuid4() + group_id = uuid4() + seed = SeedPrompt( + id=seed_id, + value="full text", + dataset_name=DATASET, + prompt_group_id=group_id, + role="user", + sequence=4, + source="source", + authors=["author"], + groups=["group"], + metadata={"persisted": "yes"}, + ) + await _add(sqlite_instance, seed) + response = _detail(client, str(group_id)) + assert response.status_code == 200 + member = response.json()["members"][0] + assert member["id"] == str(seed_id) + assert member["prompt_group_id"] == str(group_id) + assert member["value"] == "full text" + for field in ( + "role", + "sequence", + "value_sha256", + "dataset_name", + "source", + "authors", + "groups", + "date_added", + "added_by", + "metadata", + "data_type", + ): + assert field in member + + async def test_counts_are_logical_examples_and_use_same_predicates(self, client, sqlite_instance, tmp_path): + group_id = uuid4() + image_path = tmp_path / "count-image.png" + image_path.write_bytes(b"local test image") + await _add( + sqlite_instance, + SeedPrompt(value=str(image_path), dataset_name=DATASET, prompt_group_id=group_id, data_type="image_path"), + SeedObjective(value="two", dataset_name=DATASET, prompt_group_id=group_id), + SeedPrompt(value="three", dataset_name=DATASET, data_type="text"), + ) + response = _list(client, modality="image_path") + assert response.status_code == 200 + body = response.json() + assert body["total"] == 1 + assert body["items"][0]["piece_count"] == 2 + assert body["items"][0]["objective_count"] == 1 + + +class TestDatabaseBoundsAndSideEffects: + async def test_page_query_is_bounded_and_does_not_n_plus_one(self, client, sqlite_instance): + await _add(sqlite_instance, *(SeedPrompt(value=f"p{i}", dataset_name=DATASET) for i in range(250))) + statements = [] + from sqlalchemy import event + + def capture(_connection, _cursor, statement, _parameters, _context, _executemany): + statements.append(statement.lower()) + + event.listen(sqlite_instance.engine, "before_cursor_execute", capture) + try: + response = _list(client, limit=2) + finally: + event.remove(sqlite_instance.engine, "before_cursor_execute", capture) + assert response.status_code == 200 + assert len(response.json()["items"]) == 2 + assert len(statements) < 12 + assert not any("select" in statement and "250" in statement for statement in statements) + + async def test_browsing_is_read_only_and_does_not_fetch_provider_or_write(self, client, sqlite_instance): + with ( + patch( + "pyrit.datasets.SeedDatasetProvider.get_all_dataset_names_async", + side_effect=AssertionError("provider fetch"), + ), + patch.object(sqlite_instance, "add_seeds_to_memory_async", side_effect=AssertionError("write")) as write, + ): + response = _list(client) + assert response.status_code == 200 + assert write.call_count == 0 + + +class TestCompatibility: + async def test_existing_dataset_list_route_remains_unchanged(self, client): + response = client.get("/api/datasets") + assert response.status_code == 200 + assert "items" in response.json() + + async def test_existing_get_seeds_harm_semantics_remain_all_categories(self, sqlite_instance: MemoryInterface): + await _add( + sqlite_instance, + SeedPrompt(value="one", harm_categories=["hate"]), + SeedPrompt(value="two", harm_categories=["hate", "violence"]), + ) + assert len(sqlite_instance.get_seeds(harm_categories=["hate", "violence"])) == 1 diff --git a/tests/unit/memory/memory_interface/test_interface_seed_browsing_contract.py b/tests/unit/memory/memory_interface/test_interface_seed_browsing_contract.py new file mode 100644 index 0000000000..3bb73c871c --- /dev/null +++ b/tests/unit/memory/memory_interface/test_interface_seed_browsing_contract.py @@ -0,0 +1,255 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""RED contract tests for the portable #2748 logical-example read helper.""" + +from __future__ import annotations + +from datetime import UTC, datetime, timedelta +from typing import TYPE_CHECKING, Any +from uuid import UUID, uuid4 + +import pytest +from sqlalchemy import event + +from pyrit.models import SeedObjective, SeedPrompt + +if TYPE_CHECKING: + from pyrit.memory import MemoryInterface + + +DATASET = "memory-browse-contract" + + +def _field(value: Any, name: str) -> Any: + """Read a contract field from either a Pydantic result or a mapping.""" + return value.get(name) if isinstance(value, dict) else getattr(value, name) + + +def _page(memory: MemoryInterface, **kwargs: Any) -> Any: + """Call the proposed narrow, database-backed logical-example page helper.""" + helper = getattr(memory, "get_seed_example_page", None) + assert helper is not None, "RED: MemoryInterface.get_seed_example_page is not implemented" + return helper(**kwargs) + + +async def _add(memory: MemoryInterface, *seeds: SeedPrompt | SeedObjective) -> None: + await memory.add_seeds_to_memory_async(seeds=list(seeds), added_by="2748-memory-test") + + +class TestSeedBrowsingMemoryContract: + async def test_logical_identity_uses_group_id_else_seed_id(self, sqlite_instance: MemoryInterface): + group_id = uuid4() + grouped = SeedPrompt(value="grouped", dataset_name=DATASET, prompt_group_id=group_id) + ungrouped = SeedPrompt(value="ungrouped", dataset_name=DATASET) + await _add(sqlite_instance, grouped, ungrouped) + + page = _page(sqlite_instance, dataset_name=DATASET, limit=10) + examples = _field(page, "items") + assert {str(_field(item, "example_id")) for item in examples} == {str(group_id), str(ungrouped.id)} + assert {str(_field(item, "seed_ids")[0]) for item in examples} == {str(grouped.id), str(ungrouped.id)} + + async def test_order_is_earliest_complete_date_then_id(self, sqlite_instance: MemoryInterface): + old_group = UUID("00000000-0000-0000-0000-000000000001") + new_group = UUID("00000000-0000-0000-0000-000000000002") + start = datetime(2024, 1, 1, tzinfo=UTC) + await _add( + sqlite_instance, + SeedPrompt(value="old", dataset_name=DATASET, prompt_group_id=old_group, date_added=start), + SeedPrompt( + value="new", dataset_name=DATASET, prompt_group_id=new_group, date_added=start + timedelta(days=1) + ), + SeedPrompt( + value="late member", + dataset_name=DATASET, + prompt_group_id=old_group, + date_added=start + timedelta(days=3), + ), + ) + ids = [ + _field(item, "example_id") + for item in _field(_page(sqlite_instance, dataset_name=DATASET, limit=10), "items") + ] + assert [str(value) for value in ids] == [str(new_group), str(old_group)] + + async def test_tied_timestamps_use_descending_logical_id(self, sqlite_instance: MemoryInterface): + timestamp = datetime(2024, 1, 1, tzinfo=UTC) + lower = UUID("00000000-0000-0000-0000-000000000001") + higher = UUID("00000000-0000-0000-0000-000000000002") + await _add( + sqlite_instance, + SeedPrompt(value="lower", dataset_name=DATASET, prompt_group_id=lower, date_added=timestamp), + SeedPrompt(value="higher", dataset_name=DATASET, prompt_group_id=higher, date_added=timestamp), + ) + ids = [ + _field(item, "example_id") + for item in _field(_page(sqlite_instance, dataset_name=DATASET, limit=10), "items") + ] + assert [str(value) for value in ids] == [str(higher), str(lower)] + + async def test_cursor_continuation_pages_logical_examples_not_seed_rows(self, sqlite_instance: MemoryInterface): + group_id = uuid4() + await _add( + sqlite_instance, + SeedPrompt(value="first member", dataset_name=DATASET, prompt_group_id=group_id), + SeedPrompt(value="second member", dataset_name=DATASET, prompt_group_id=group_id), + *(SeedPrompt(value=f"single-{index}", dataset_name=DATASET) for index in range(3)), + ) + first = _page(sqlite_instance, dataset_name=DATASET, limit=1) + second = _page(sqlite_instance, dataset_name=DATASET, limit=1, cursor=_field(first, "next_cursor")) + assert len(_field(first, "items")) == len(_field(second, "items")) == 1 + assert _field(first, "items")[0] != _field(second, "items")[0] + assert len(_field(_field(first, "items")[0], "members")) == 2 + + seen = _field(first, "items") + _field(second, "items") + cursor = _field(second, "next_cursor") + while cursor: + page = _page(sqlite_instance, dataset_name=DATASET, limit=1, cursor=cursor) + seen += _field(page, "items") + cursor = _field(page, "next_cursor") + assert len({_field(item, "example_id") for item in seen}) == 4 + + async def test_filters_are_member_or_and_example_and(self, sqlite_instance: MemoryInterface): + group_id = uuid4() + await _add( + sqlite_instance, + SeedPrompt(value="image", dataset_name=DATASET, prompt_group_id=group_id, data_type="text"), + SeedPrompt(value="violence", dataset_name=DATASET, prompt_group_id=group_id, harm_categories=["violence"]), + SeedPrompt(value="only image", dataset_name=DATASET, data_type="text"), + ) + page = _page( + sqlite_instance, + dataset_name=DATASET, + data_types=["text"], + harm_categories=["violence"], + limit=10, + ) + assert len(_field(page, "items")) == 1 + item = _field(page, "items")[0] + expected_ids = {str(seed.id) for seed in sqlite_instance.get_seeds(prompt_group_ids=[group_id])} + assert expected_ids == {str(seed_id) for seed_id in _field(item, "seed_ids")} + assert len(_field(item, "members")) == 2 + + async def test_modality_harm_and_seed_type_values_are_or(self, sqlite_instance: MemoryInterface): + await _add( + sqlite_instance, + SeedPrompt(value="one", dataset_name=DATASET, data_type="text", harm_categories=["hate"]), + SeedPrompt(value="two", dataset_name=DATASET, data_type="url", harm_categories=["violence"]), + SeedObjective(value="three", dataset_name=DATASET), + ) + page = _page( + sqlite_instance, + dataset_name=DATASET, + data_types=["text", "url"], + harm_categories=["hate", "violence"], + seed_types=["prompt", "objective"], + limit=10, + ) + assert len(_field(page, "items")) == 3 + + async def test_harm_matching_is_case_insensitive_whole_value_and_missing_is_unlabeled( + self, sqlite_instance: MemoryInterface + ): + labeled = SeedPrompt(value="labeled", dataset_name=DATASET, harm_categories=["VIOLENCE"]) + unlabeled = SeedPrompt(value="unlabeled", dataset_name=DATASET, harm_categories=[]) + await _add(sqlite_instance, labeled, unlabeled) + exact = _page(sqlite_instance, dataset_name=DATASET, harm_categories=["violence"], limit=10) + substring = _page(sqlite_instance, dataset_name=DATASET, harm_categories=["vio"], limit=10) + assert len(_field(exact, "items")) == 1 + assert str(_field(exact, "items")[0]["seed_ids"][0]) == str(labeled.id) + assert len(_field(substring, "items")) == 0 + all_items = _field(_page(sqlite_instance, dataset_name=DATASET, limit=10), "items") + unlabeled_item = next( + item for item in all_items if str(unlabeled.id) in [str(i) for i in _field(item, "seed_ids")] + ) + assert _field(unlabeled_item, "has_unlabeled_harm") is True + + @pytest.mark.parametrize("text", ["100% literal", "literal_value"]) + async def test_text_search_is_case_insensitive_literal_and_text_only( + self, sqlite_instance: MemoryInterface, text: str, tmp_path + ): + media = tmp_path / "media-path.png" + media.write_bytes(b"local image") + await _add( + sqlite_instance, + SeedPrompt(value="Need 100% literal_value", dataset_name=DATASET), + SeedPrompt(value=str(media), dataset_name=DATASET, data_type="image_path"), + SeedPrompt(value="metadata", dataset_name=DATASET, metadata={"search": "metadata_only"}), + ) + page = _page(sqlite_instance, dataset_name=DATASET, value_search=text.upper(), limit=10) + assert len(_field(page, "items")) == 1 + assert ( + len(_field(_page(sqlite_instance, dataset_name=DATASET, value_search="metadata_only", limit=10), "items")) + == 0 + ) + assert ( + len(_field(_page(sqlite_instance, dataset_name=DATASET, value_search="media-path", limit=10), "items")) == 0 + ) + + async def test_count_uses_same_logical_predicates_as_page(self, sqlite_instance: MemoryInterface): + group_id = uuid4() + await _add( + sqlite_instance, + SeedPrompt(value="matching", dataset_name=DATASET, prompt_group_id=group_id, data_type="text"), + SeedObjective(value="related", dataset_name=DATASET, prompt_group_id=group_id), + SeedPrompt(value="other", dataset_name=DATASET, data_type="text"), + ) + page = _page(sqlite_instance, dataset_name=DATASET, data_types=["text"], limit=10) + assert _field(page, "total") == 2 + grouped = next(item for item in _field(page, "items") if _field(item, "piece_count") == 2) + assert str(_field(grouped, "example_id")) == str(group_id) + + async def test_only_selected_examples_are_expanded_to_all_members(self, sqlite_instance: MemoryInterface): + selected = uuid4() + not_selected = uuid4() + await _add( + sqlite_instance, + SeedPrompt(value="selected", dataset_name=DATASET, prompt_group_id=selected, harm_categories=["hate"]), + SeedPrompt(value="selected related", dataset_name=DATASET, prompt_group_id=selected), + SeedPrompt(value="unselected", dataset_name=DATASET, prompt_group_id=not_selected), + ) + page = _page(sqlite_instance, dataset_name=DATASET, harm_categories=["hate"], limit=10) + assert len(_field(page, "items")) == 1 + assert len(_field(_field(page, "items")[0], "members")) == 2 + + async def test_query_is_database_bounded_and_not_n_plus_one(self, sqlite_instance: MemoryInterface): + await _add(sqlite_instance, *(SeedPrompt(value=f"seed-{index}", dataset_name=DATASET) for index in range(250))) + statements: list[str] = [] + + def capture(_connection, _cursor, statement, _parameters, _context, _executemany): + statements.append(statement.lower()) + + event.listen(sqlite_instance.engine, "before_cursor_execute", capture) + try: + page = _page(sqlite_instance, dataset_name=DATASET, limit=2) + finally: + event.remove(sqlite_instance.engine, "before_cursor_execute", capture) + assert len(_field(page, "items")) == 2 + assert len(statements) < 12 + assert any(" limit " in f" {statement} " for statement in statements) + + async def test_malformed_cursor_and_filter_mismatch_are_rejected(self, sqlite_instance: MemoryInterface): + await _add( + sqlite_instance, + SeedPrompt(value="one", dataset_name=DATASET), + SeedPrompt(value="two", dataset_name=DATASET, harm_categories=["violence"]), + ) + with pytest.raises(ValueError): + _page(sqlite_instance, dataset_name=DATASET, limit=1, cursor="malformed") + first = _page(sqlite_instance, dataset_name=DATASET, limit=1) + with pytest.raises(ValueError): + _page( + sqlite_instance, + dataset_name=DATASET, + harm_categories=["violence"], + limit=1, + cursor=_field(first, "next_cursor"), + ) + + async def test_existing_get_seeds_harm_semantics_remain_all_categories(self, sqlite_instance: MemoryInterface): + await _add( + sqlite_instance, + SeedPrompt(value="one", harm_categories=["hate"]), + SeedPrompt(value="two", harm_categories=["hate", "violence"]), + ) + assert len(sqlite_instance.get_seeds(harm_categories=["hate", "violence"])) == 1 From ca2d5daca3b981192b1436b7afad909233f6d294 Mon Sep 17 00:00:00 2001 From: Jackson Severino da Rocha Date: Mon, 28 Sep 2026 13:41:39 -0300 Subject: [PATCH 2/9] test: harden seed browsing contract coverage --- ..._dataset_service_seed_browsing_contract.py | 72 ++++-- tests/unit/backend/test_seed_browsing_api.py | 67 ++++- .../test_interface_seed_browsing_contract.py | 241 +++++++++++++++--- 3 files changed, 324 insertions(+), 56 deletions(-) diff --git a/tests/unit/backend/test_dataset_service_seed_browsing_contract.py b/tests/unit/backend/test_dataset_service_seed_browsing_contract.py index 3c419d4e55..50a4789ae5 100644 --- a/tests/unit/backend/test_dataset_service_seed_browsing_contract.py +++ b/tests/unit/backend/test_dataset_service_seed_browsing_contract.py @@ -83,6 +83,32 @@ async def test_cursor_is_bound_to_selection_and_effective_filters( selection_key=SELECTION_KEY, limit=1, cursor=cursor, search="changed" ) + @pytest.mark.parametrize( + "changed_filters", + [ + {"search": "changed"}, + {"data_types": ["url"]}, + {"harm_categories": ["violence"]}, + {"seed_types": ["objective"]}, + ], + ) + async def test_cursor_is_bound_to_each_effective_filter( + self, dataset_service: DatasetService, sqlite_instance: MemoryInterface, changed_filters: dict[str, object] + ): + await _add( + sqlite_instance, + SeedPrompt(value="one", dataset_name=DATASET), + SeedPrompt( + value="https://example.com/two", dataset_name=DATASET, data_type="url", harm_categories=["violence"] + ), + ) + first = await _service_method(dataset_service, "list_seed_examples_async")(selection_key=SELECTION_KEY, limit=1) + cursor = _field(_field(first, "pagination"), "next_cursor") + with pytest.raises(ValueError): + await _service_method(dataset_service, "list_seed_examples_async")( + selection_key=SELECTION_KEY, limit=1, cursor=cursor, **changed_filters + ) + async def test_malformed_cursor_and_invalid_selection_are_rejected(self, dataset_service: DatasetService): with pytest.raises(ValueError): await _service_method(dataset_service, "list_seed_examples_async")( @@ -114,7 +140,7 @@ async def test_list_formats_group_preview_types_modalities_counts_and_harm_summa ) item = _field(response, "items")[0] assert _field(item, "preview_truncated") is True - assert len(_field(item, "preview")) <= 103 + assert _field(item, "preview") == ("x" * 100) + "..." assert _field(item, "piece_count") == 3 assert _field(item, "objective_count") == 1 assert "text" in _field(item, "modalities") @@ -138,7 +164,14 @@ async def test_detail_returns_complete_group_and_persisted_provenance( parameters=["name"], ) await _add( - sqlite_instance, seed, SeedObjective(value="condition", dataset_name=DATASET, prompt_group_id=group_id) + sqlite_instance, + seed, + SeedObjective( + value="condition", + dataset_name=DATASET, + prompt_group_id=group_id, + metadata={"conditions": "must hold"}, + ), ) detail = await _service_method(dataset_service, "get_seed_example_async")( selection_key=SELECTION_KEY, example_id=str(group_id) @@ -146,22 +179,26 @@ async def test_detail_returns_complete_group_and_persisted_provenance( members = _field(detail, "members") assert len(members) == 2 prompt = next(member for member in members if _field(member, "id") == seed.id) + objective = next(member for member in members if _field(member, "seed_type") == "objective") assert _field(prompt, "prompt_group_id") == group_id - assert _field(prompt, "value") == "full text" + assert _field(prompt, "name") == seed.name + assert _field(prompt, "value") == seed.value assert _field(prompt, "role") == "user" assert _field(prompt, "sequence") == 4 - for field in ( - "value_sha256", - "dataset_name", - "source", - "authors", - "groups", - "date_added", - "added_by", - "metadata", - "data_type", - ): - assert field in (prompt if isinstance(prompt, dict) else prompt.model_fields_set | set(prompt.model_fields)) + stored_prompt = next( + stored for stored in sqlite_instance.get_seeds(prompt_group_ids=[group_id]) if stored.id == seed.id + ) + assert _field(prompt, "value_sha256") == stored_prompt.value_sha256 + assert _field(prompt, "dataset_name") == DATASET + assert _field(prompt, "source") == "source" + assert _field(prompt, "authors") == ["author"] + assert _field(prompt, "groups") == ["group"] + assert _field(prompt, "date_added") == stored_prompt.date_added + assert _field(prompt, "added_by") == "2748-service-test" + assert _field(prompt, "metadata") == {"persisted": True} + assert _field(prompt, "data_type") == "text" + assert _field(objective, "value") == "condition" + assert _field(objective, "metadata") == {"conditions": "must hold"} async def test_detail_preserves_template_parameters_and_objective_conditions( self, dataset_service: DatasetService, sqlite_instance: MemoryInterface @@ -182,7 +219,10 @@ async def test_detail_preserves_template_parameters_and_objective_conditions( selection_key=SELECTION_KEY, example_id=str(group_id) ) assert _field(detail, "members") - assert any(_field(member, "parameters") == ["name"] for member in _field(detail, "members")) + template = next(member for member in _field(detail, "members") if _field(member, "seed_type") == "prompt") + assert _field(template, "value") == "{{ name }}" + assert _field(template, "is_jinja_template") is True + assert _field(template, "parameters") == ["name"] assert any(_field(member, "seed_type") == "objective" for member in _field(detail, "members")) async def test_invalid_detail_does_not_generate_group_identity(self, dataset_service: DatasetService): diff --git a/tests/unit/backend/test_seed_browsing_api.py b/tests/unit/backend/test_seed_browsing_api.py index b5a416e799..c5a36b1c35 100644 --- a/tests/unit/backend/test_seed_browsing_api.py +++ b/tests/unit/backend/test_seed_browsing_api.py @@ -37,7 +37,7 @@ def client(patch_central_database) -> TestClient: return TestClient(app) -async def _add(memory: MemoryInterface, *seeds: SeedPrompt | SeedObjective) -> None: +async def _add(memory: MemoryInterface, *seeds: SeedPrompt | SeedObjective | SeedSimulatedConversation) -> None: await memory.add_seeds_to_memory_async(seeds=list(seeds), added_by="2748-test") @@ -72,12 +72,12 @@ async def test_empty_filter_result_is_not_an_error(self, client, sqlite_instance async def test_named_selection_key_is_not_display_name(self, client, sqlite_instance: MemoryInterface): await _add(sqlite_instance, SeedPrompt(value="named", dataset_name=DATASET)) assert len(_items(_list(client, NAMED_KEY))) == 1 - assert _list(client, DATASET).status_code in {400, 404, 422} + assert _list(client, DATASET).status_code == 404 async def test_unnamed_selection_and_invalid_selection(self, client, sqlite_instance: MemoryInterface): await _add(sqlite_instance, SeedPrompt(value="unnamed")) assert len(_items(_list(client, UNNAMED_KEY))) == 1 - assert _list(client, "dataset:named:not-loaded").status_code in {400, 404, 422} + assert _list(client, "dataset:named:not-loaded").status_code == 404 async def test_named_unnamed_namespace_is_distinct(self, client, sqlite_instance: MemoryInterface): await _add(sqlite_instance, SeedPrompt(value="literal", dataset_name="__unnamed__"), SeedPrompt(value="none")) @@ -107,7 +107,7 @@ async def test_multiple_pages_have_no_duplicates_or_omissions(self, client, sqli @pytest.mark.parametrize("limit", [0, -1, 101]) async def test_page_size_is_validated(self, client, sqlite_instance: MemoryInterface, limit: int): - assert _list(client, limit=limit).status_code in {400, 422} + assert _list(client, limit=limit).status_code == 422 async def test_group_identity_preserves_ids_and_does_not_hash_merge(self, client, sqlite_instance: MemoryInterface): group_id = uuid4() @@ -269,13 +269,37 @@ async def test_cursor_is_opaque_bound_to_dataset_and_effective_filters(self, cli await _add(sqlite_instance, *(SeedPrompt(value=str(i), dataset_name=DATASET) for i in range(3))) cursor = _list(client, limit=1).json()["pagination"]["next_cursor"] assert cursor and not cursor.startswith("1") - assert _list(client, limit=1, cursor="not-a-cursor").status_code in {400, 422} - assert _list(client, UNNAMED_KEY, limit=1, cursor=cursor).status_code in {400, 422} - assert _list(client, limit=1, search="different", cursor=cursor).status_code in {400, 422} + assert _list(client, limit=1, cursor="not-a-cursor").status_code == 400 + assert _list(client, UNNAMED_KEY, limit=1, cursor=cursor).status_code == 400 + assert _list(client, limit=1, search="different", cursor=cursor).status_code == 400 + + @pytest.mark.parametrize( + "changed_filters", + [ + {"search": "different"}, + {"modality": "url"}, + {"harm_category": "violence"}, + {"seed_type": "objective"}, + ], + ) + async def test_cursor_rejects_each_changed_effective_filter(self, client, sqlite_instance, changed_filters): + await _add( + sqlite_instance, + SeedPrompt(value="one", dataset_name=DATASET), + SeedPrompt( + value="https://example.com/two", dataset_name=DATASET, data_type="url", harm_categories=["violence"] + ), + ) + first = _list(client, limit=1) + cursor = first.json()["pagination"]["next_cursor"] + response = _list(client, limit=1, cursor=cursor, **changed_filters) + assert response.status_code == 400 + assert response.json()["detail"] async def test_missing_detail_example_is_not_silently_empty(self, client, sqlite_instance): response = _detail(client, str(uuid4())) - assert response.status_code in {404, 422} + assert response.status_code == 404 + assert response.json()["detail"] async def test_template_is_not_rendered_or_loaded(self, client, sqlite_instance): template = SeedPrompt( @@ -307,6 +331,10 @@ async def test_simulated_configuration_is_returned_without_generation_or_target_ assert generate.call_count == 0 assert response.json()["items"][0]["seed_types"] == ["simulated_conversation"] + detail = _detail(client, response.json()["items"][0]["example_id"]) + assert detail.status_code == 200 + assert detail.json()["members"][0]["value"] == config.value + class TestPreviewDetailAndCounts: async def test_preview_uses_100_character_convention_and_hides_full_content(self, client, sqlite_instance): @@ -318,9 +346,22 @@ async def test_preview_uses_100_character_convention_and_hides_full_content(self items = _items(_list(client)) assert any(item["preview"] == short and item["preview_truncated"] is False for item in items) long_item = next(item for item in items if item["preview"].startswith("y")) - assert len(long_item["preview"]) <= 103 and long_item["preview_truncated"] is True + assert long_item["preview"] == ("y" * 100) + "..." + assert long_item["preview_truncated"] is True assert long not in long_item["preview"] + async def test_detail_returns_full_content_after_list_preview_truncation(self, client, sqlite_instance): + long_value = "long-value-" + ("x" * 150) + seed = SeedPrompt(value=long_value, dataset_name=DATASET) + await _add(sqlite_instance, seed) + item = _items(_list(client))[0] + assert item["preview_truncated"] is True + detail = _detail(client, item["example_id"]) + assert detail.status_code == 200 + member = detail.json()["members"][0] + assert member["value"] == long_value + assert member["prompt_group_id"] is None + async def test_media_preview_is_label_only_and_never_bytes_path_or_credentials( self, client, sqlite_instance, tmp_path ): @@ -354,6 +395,12 @@ async def test_detail_returns_all_persisted_fields_without_new_ids_or_rendering( assert member["id"] == str(seed_id) assert member["prompt_group_id"] == str(group_id) assert member["value"] == "full text" + assert member["role"] == "user" + assert member["sequence"] == 4 + assert member["source"] == "source" + assert member["authors"] == ["author"] + assert member["groups"] == ["group"] + assert member["metadata"] == {"persisted": "yes"} for field in ( "role", "sequence", @@ -409,7 +456,7 @@ def capture(_connection, _cursor, statement, _parameters, _context, _executemany async def test_browsing_is_read_only_and_does_not_fetch_provider_or_write(self, client, sqlite_instance): with ( patch( - "pyrit.datasets.SeedDatasetProvider.get_all_dataset_names_async", + "pyrit.datasets.SeedDatasetProvider.fetch_datasets_async", side_effect=AssertionError("provider fetch"), ), patch.object(sqlite_instance, "add_seeds_to_memory_async", side_effect=AssertionError("write")) as write, diff --git a/tests/unit/memory/memory_interface/test_interface_seed_browsing_contract.py b/tests/unit/memory/memory_interface/test_interface_seed_browsing_contract.py index 3bb73c871c..bbd1954f3c 100644 --- a/tests/unit/memory/memory_interface/test_interface_seed_browsing_contract.py +++ b/tests/unit/memory/memory_interface/test_interface_seed_browsing_contract.py @@ -12,7 +12,7 @@ import pytest from sqlalchemy import event -from pyrit.models import SeedObjective, SeedPrompt +from pyrit.models import SeedObjective, SeedPrompt, SeedSimulatedConversation if TYPE_CHECKING: from pyrit.memory import MemoryInterface @@ -33,7 +33,7 @@ def _page(memory: MemoryInterface, **kwargs: Any) -> Any: return helper(**kwargs) -async def _add(memory: MemoryInterface, *seeds: SeedPrompt | SeedObjective) -> None: +async def _add(memory: MemoryInterface, *seeds: SeedPrompt | SeedObjective | SeedSimulatedConversation) -> None: await memory.add_seeds_to_memory_async(seeds=list(seeds), added_by="2748-memory-test") @@ -88,26 +88,63 @@ async def test_tied_timestamps_use_descending_logical_id(self, sqlite_instance: assert [str(value) for value in ids] == [str(higher), str(lower)] async def test_cursor_continuation_pages_logical_examples_not_seed_rows(self, sqlite_instance: MemoryInterface): - group_id = uuid4() + group_id = UUID("00000000-0000-0000-0000-000000000003") + second_group_id = UUID("00000000-0000-0000-0000-000000000002") + first_group_id = UUID("00000000-0000-0000-0000-000000000001") + ungrouped_id = UUID("00000000-0000-0000-0000-000000000004") + first_timestamp = datetime(2024, 1, 1, tzinfo=UTC) + tied_timestamp = datetime(2024, 1, 2, tzinfo=UTC) + group_three_first = SeedPrompt( + value="group three first", dataset_name=DATASET, prompt_group_id=group_id, date_added=tied_timestamp + ) + group_three_second = SeedPrompt( + value="group three second", dataset_name=DATASET, prompt_group_id=group_id, date_added=tied_timestamp + ) + group_two = SeedPrompt( + value="group two", dataset_name=DATASET, prompt_group_id=second_group_id, date_added=tied_timestamp + ) + group_one_early = SeedPrompt( + value="group one early", dataset_name=DATASET, prompt_group_id=first_group_id, date_added=first_timestamp + ) + group_one_late = SeedPrompt( + value="group one late", + dataset_name=DATASET, + prompt_group_id=first_group_id, + date_added=datetime(2024, 1, 3, tzinfo=UTC), + ) + ungrouped = SeedPrompt(value="ungrouped", dataset_name=DATASET, id=ungrouped_id, date_added=first_timestamp) await _add( sqlite_instance, - SeedPrompt(value="first member", dataset_name=DATASET, prompt_group_id=group_id), - SeedPrompt(value="second member", dataset_name=DATASET, prompt_group_id=group_id), - *(SeedPrompt(value=f"single-{index}", dataset_name=DATASET) for index in range(3)), + group_three_first, + group_three_second, + group_two, + group_one_early, + group_one_late, + ungrouped, ) - first = _page(sqlite_instance, dataset_name=DATASET, limit=1) - second = _page(sqlite_instance, dataset_name=DATASET, limit=1, cursor=_field(first, "next_cursor")) - assert len(_field(first, "items")) == len(_field(second, "items")) == 1 - assert _field(first, "items")[0] != _field(second, "items")[0] - assert len(_field(_field(first, "items")[0], "members")) == 2 - - seen = _field(first, "items") + _field(second, "items") - cursor = _field(second, "next_cursor") - while cursor: + pages: list[Any] = [] + cursor = None + while True: page = _page(sqlite_instance, dataset_name=DATASET, limit=1, cursor=cursor) - seen += _field(page, "items") + pages.extend(_field(page, "items")) cursor = _field(page, "next_cursor") - assert len({_field(item, "example_id") for item in seen}) == 4 + if cursor is None: + break + + assert [str(_field(item, "example_id")) for item in pages] == [ + str(group_id), + str(second_group_id), + str(ungrouped_id), + str(first_group_id), + ] + assert [len(_field(item, "members")) for item in pages] == [2, 1, 1, 2] + assert [{str(seed_id) for seed_id in _field(item, "seed_ids")} for item in pages] == [ + {str(group_three_first.id), str(group_three_second.id)}, + {str(group_two.id)}, + {str(ungrouped.id)}, + {str(group_one_early.id), str(group_one_late.id)}, + ] + assert len({_field(item, "example_id") for item in pages}) == len(pages) == 4 async def test_filters_are_member_or_and_example_and(self, sqlite_instance: MemoryInterface): group_id = uuid4() @@ -130,22 +167,106 @@ async def test_filters_are_member_or_and_example_and(self, sqlite_instance: Memo assert expected_ids == {str(seed_id) for seed_id in _field(item, "seed_ids")} assert len(_field(item, "members")) == 2 - async def test_modality_harm_and_seed_type_values_are_or(self, sqlite_instance: MemoryInterface): + async def test_modality_harm_and_seed_type_values_are_or(self, sqlite_instance: MemoryInterface, tmp_path): + modality_only = SeedPrompt(value="https://example.com/modality-only", dataset_name=DATASET, data_type="url") + modality_only_second = SeedPrompt( + value=str(tmp_path / "modality-only.png"), dataset_name=DATASET, data_type="image_path" + ) + (tmp_path / "modality-only.png").write_bytes(b"image") + harm_only = SeedPrompt(value="harm only", dataset_name=DATASET, data_type="reasoning", harm_categories=["hate"]) + seed_type_only = SeedSimulatedConversation( + dataset_name=DATASET, + adversarial_chat_system_prompt=SeedPrompt(value="adversarial"), + simulated_target_system_prompt=SeedPrompt(value="target"), + ) + all_filters_group = uuid4() + all_filters_prompt = SeedPrompt( + value="https://example.com/all-filters", + dataset_name=DATASET, + prompt_group_id=all_filters_group, + data_type="url", + harm_categories=["violence"], + ) + all_filters_objective = SeedObjective( + value="all filters objective", dataset_name=DATASET, prompt_group_id=all_filters_group + ) + no_filters = SeedPrompt( + value="no filters", dataset_name=DATASET, data_type="reasoning", harm_categories=["other"] + ) await _add( sqlite_instance, - SeedPrompt(value="one", dataset_name=DATASET, data_type="text", harm_categories=["hate"]), - SeedPrompt(value="two", dataset_name=DATASET, data_type="url", harm_categories=["violence"]), - SeedObjective(value="three", dataset_name=DATASET), + modality_only, + modality_only_second, + harm_only, + seed_type_only, + all_filters_prompt, + all_filters_objective, + no_filters, ) - page = _page( + modality_page = _page(sqlite_instance, dataset_name=DATASET, data_types=["url", "image_path"], limit=10) + harm_page = _page(sqlite_instance, dataset_name=DATASET, harm_categories=["hate", "violence"], limit=10) + seed_type_page = _page( + sqlite_instance, + dataset_name=DATASET, + seed_types=["objective", "simulated_conversation"], + limit=10, + ) + combined_page = _page( sqlite_instance, dataset_name=DATASET, - data_types=["text", "url"], + data_types=["url", "image_path"], harm_categories=["hate", "violence"], - seed_types=["prompt", "objective"], + seed_types=["objective", "simulated_conversation"], limit=10, ) - assert len(_field(page, "items")) == 3 + + assert {str(_field(item, "example_id")) for item in _field(modality_page, "items")} == { + str(modality_only.id), + str(modality_only_second.id), + str(all_filters_group), + } + assert {str(_field(item, "example_id")) for item in _field(harm_page, "items")} == { + str(harm_only.id), + str(all_filters_group), + } + assert {str(_field(item, "example_id")) for item in _field(seed_type_page, "items")} == { + str(seed_type_only.id), + str(all_filters_group), + } + assert [str(_field(item, "example_id")) for item in _field(combined_page, "items")] == [str(all_filters_group)] + assert _field(combined_page, "total") == 1 + + async def test_filtered_order_uses_earliest_member_not_earliest_matching_member( + self, sqlite_instance: MemoryInterface + ): + group_a = uuid4() + group_b = uuid4() + t1 = datetime(2024, 1, 1, tzinfo=UTC) + t2 = datetime(2024, 1, 2, tzinfo=UTC) + t3 = datetime(2024, 1, 3, tzinfo=UTC) + await _add( + sqlite_instance, + SeedPrompt(value="a earliest", dataset_name=DATASET, prompt_group_id=group_a, date_added=t1), + SeedPrompt( + value="https://example.com/a-matching-later", + dataset_name=DATASET, + prompt_group_id=group_a, + date_added=t3, + data_type="url", + ), + SeedPrompt( + value="https://example.com/b-matching", + dataset_name=DATASET, + prompt_group_id=group_b, + date_added=t2, + data_type="url", + ), + ) + page = _page(sqlite_instance, dataset_name=DATASET, data_types=["url"], limit=10) + assert [str(_field(item, "example_id")) for item in _field(page, "items")] == [ + str(group_b), + str(group_a), + ] async def test_harm_matching_is_case_insensitive_whole_value_and_missing_is_unlabeled( self, sqlite_instance: MemoryInterface @@ -186,6 +307,20 @@ async def test_text_search_is_case_insensitive_literal_and_text_only( len(_field(_page(sqlite_instance, dataset_name=DATASET, value_search="media-path", limit=10), "items")) == 0 ) + async def test_percent_search_is_literal(self, sqlite_instance: MemoryInterface): + literal_match = SeedPrompt(value="contains 100%", dataset_name=DATASET) + wildcard_only = SeedPrompt(value="contains 1000", dataset_name=DATASET) + await _add(sqlite_instance, literal_match, wildcard_only) + page = _page(sqlite_instance, dataset_name=DATASET, value_search="100%", limit=10) + assert [str(_field(item, "seed_ids")[0]) for item in _field(page, "items")] == [str(literal_match.id)] + + async def test_underscore_search_is_literal(self, sqlite_instance: MemoryInterface): + literal_match = SeedPrompt(value="contains a_b", dataset_name=DATASET) + wildcard_only = SeedPrompt(value="contains acb", dataset_name=DATASET) + await _add(sqlite_instance, literal_match, wildcard_only) + page = _page(sqlite_instance, dataset_name=DATASET, value_search="a_b", limit=10) + assert [str(_field(item, "seed_ids")[0]) for item in _field(page, "items")] == [str(literal_match.id)] + async def test_count_uses_same_logical_predicates_as_page(self, sqlite_instance: MemoryInterface): group_id = uuid4() await _add( @@ -213,7 +348,15 @@ async def test_only_selected_examples_are_expanded_to_all_members(self, sqlite_i assert len(_field(_field(page, "items")[0], "members")) == 2 async def test_query_is_database_bounded_and_not_n_plus_one(self, sqlite_instance: MemoryInterface): - await _add(sqlite_instance, *(SeedPrompt(value=f"seed-{index}", dataset_name=DATASET) for index in range(250))) + groups = [uuid4() for _ in range(20)] + await _add( + sqlite_instance, + *( + SeedPrompt(value=f"seed-{index}", dataset_name=DATASET, prompt_group_id=group_id) + for index, group_id in enumerate(groups) + ), + *(SeedPrompt(value=f"extra-{index}", dataset_name=DATASET) for index in range(250)), + ) statements: list[str] = [] def capture(_connection, _cursor, statement, _parameters, _context, _executemany): @@ -221,12 +364,17 @@ def capture(_connection, _cursor, statement, _parameters, _context, _executemany event.listen(sqlite_instance.engine, "before_cursor_execute", capture) try: - page = _page(sqlite_instance, dataset_name=DATASET, limit=2) + page = _page(sqlite_instance, dataset_name=DATASET, limit=20) finally: event.remove(sqlite_instance.engine, "before_cursor_execute", capture) - assert len(_field(page, "items")) == 2 - assert len(statements) < 12 - assert any(" limit " in f" {statement} " for statement in statements) + assert len(_field(page, "items")) == 20 + select_statements = [statement for statement in statements if statement.lstrip().startswith("select")] + grouped_page_queries = [ + statement for statement in select_statements if "group by" in statement and "order by" in statement + ] + assert grouped_page_queries + assert all(" limit " in f" {statement} " for statement in grouped_page_queries) + assert len(select_statements) <= 5 async def test_malformed_cursor_and_filter_mismatch_are_rejected(self, sqlite_instance: MemoryInterface): await _add( @@ -246,6 +394,39 @@ async def test_malformed_cursor_and_filter_mismatch_are_rejected(self, sqlite_in cursor=_field(first, "next_cursor"), ) + @pytest.mark.parametrize( + ("changed_filters", "changed_dataset"), + [ + ({"value_search": "changed"}, None), + ({"data_types": ["url"]}, None), + ({"harm_categories": ["violence"]}, None), + ({"seed_types": ["objective"]}, None), + ({}, "another-memory-dataset"), + ], + ) + async def test_cursor_is_bound_to_every_effective_filter( + self, + sqlite_instance: MemoryInterface, + changed_filters: dict[str, object], + changed_dataset: str | None, + ): + await _add( + sqlite_instance, + SeedPrompt(value="one", dataset_name=DATASET), + SeedPrompt( + value="https://example.com/two", dataset_name=DATASET, data_type="url", harm_categories=["violence"] + ), + ) + first = _page(sqlite_instance, dataset_name=DATASET, limit=1) + with pytest.raises(ValueError): + _page( + sqlite_instance, + dataset_name=changed_dataset or DATASET, + limit=1, + cursor=_field(first, "next_cursor"), + **changed_filters, + ) + async def test_existing_get_seeds_harm_semantics_remain_all_categories(self, sqlite_instance: MemoryInterface): await _add( sqlite_instance, From c3193360e719bb2264308f7cf1749318b3663aee Mon Sep 17 00:00:00 2001 From: Jackson Severino da Rocha Date: Mon, 28 Sep 2026 18:26:47 -0300 Subject: [PATCH 3/9] feat: add database-bounded seed example browsing --- pyrit/memory/__init__.py | 8 + pyrit/memory/azure_sql_memory.py | 26 + pyrit/memory/memory_interface.py | 550 +++++++++++++++++- pyrit/memory/sqlite_memory.py | 16 +- ...est_seed_browsing_azure_sql_integration.py | 159 +++++ .../test_interface_seed_browsing_contract.py | 137 ++++- .../test_interface_seed_browsing_sql.py | 108 ++++ 7 files changed, 997 insertions(+), 7 deletions(-) create mode 100644 tests/integration/memory/test_seed_browsing_azure_sql_integration.py create mode 100644 tests/unit/memory/memory_interface/test_interface_seed_browsing_sql.py diff --git a/pyrit/memory/__init__.py b/pyrit/memory/__init__.py index 42aae06309..3376a877fd 100644 --- a/pyrit/memory/__init__.py +++ b/pyrit/memory/__init__.py @@ -23,6 +23,10 @@ ScenarioHistoryKeysetCursor, ScenarioHistoryRunRecord, ScenarioRunStateRecord, + SeedExample, + SeedExampleDatasetScope, + SeedExampleMember, + SeedExamplePage, ) from pyrit.memory.memory_models import AttackResultEntry, EmbeddingDataEntry, PromptMemoryEntry, SeedEntry from pyrit.memory.sqlite_memory import SQLiteMemory @@ -61,6 +65,10 @@ "ErrorDataTypeSerializer": "pyrit.memory.storage", "ImagePathDataTypeSerializer": "pyrit.memory.storage", "MemoryInterface": "pyrit.memory.memory_interface", + "SeedExample": "pyrit.memory.memory_interface", + "SeedExampleDatasetScope": "pyrit.memory.memory_interface", + "SeedExampleMember": "pyrit.memory.memory_interface", + "SeedExamplePage": "pyrit.memory.memory_interface", "MemoryEmbedding": "pyrit.memory.memory_embedding", "ScenarioHistoryKeysetCursor": "pyrit.memory.memory_interface", "ScenarioHistoryRunRecord": "pyrit.memory.memory_interface", diff --git a/pyrit/memory/azure_sql_memory.py b/pyrit/memory/azure_sql_memory.py index 651bfb76aa..d4109dcd78 100644 --- a/pyrit/memory/azure_sql_memory.py +++ b/pyrit/memory/azure_sql_memory.py @@ -14,11 +14,14 @@ Unicode, and_, bindparam, + case, create_engine, event, exists, func, + literal, literal_column, + select, text, ) from sqlalchemy.engine.base import Engine @@ -459,6 +462,29 @@ def _get_condition_json_array_match( combined = joiner.join(conditions) return text(f"""ISJSON("{table_name}".{column_name}) = 1 AND ({combined})""").bindparams(**bindparams_dict) + def _get_seed_harm_category_condition( + self, *, json_column: InstrumentedAttribute[Any], categories: Sequence[str] + ) -> Any: + """ + Build an aliased-column-safe Azure SQL harm-category membership predicate. + + Returns: + Any: A SQLAlchemy predicate matching any requested category. + """ + values = [category.lower() for category in categories] + safe_array = case( + ( + and_( + func.ISJSON(json_column) == literal(1), + func.LEFT(func.LTRIM(json_column), literal(1)) == literal("["), + ), + json_column, + ), + else_=literal("[]"), + ) + elements = func.OPENJSON(safe_array).table_valued("value") + return exists(select(1).select_from(elements).where(func.lower(elements.c.value).in_(values))) + def _get_attack_result_label_condition(self, *, labels: dict[str, str | Sequence[str]]) -> Any: """ Azure SQL implementation for filtering AttackResults by labels. diff --git a/pyrit/memory/memory_interface.py b/pyrit/memory/memory_interface.py index fff51a507b..a010da49d3 100644 --- a/pyrit/memory/memory_interface.py +++ b/pyrit/memory/memory_interface.py @@ -10,7 +10,7 @@ import re import uuid import weakref -from collections.abc import Collection, Iterator, Mapping, MutableSequence, Sequence +from collections.abc import Callable, Collection, Iterator, Mapping, MutableSequence, Sequence from contextlib import closing from dataclasses import dataclass from datetime import UTC, datetime, timedelta @@ -18,14 +18,15 @@ from typing import TYPE_CHECKING, Any, ClassVar, Literal, NamedTuple, TypeVar from urllib.parse import urlparse -from sqlalchemy import MetaData, and_, case, exists, func, literal, not_, or_, select, update +from sqlalchemy import MetaData, String, and_, case, cast, exists, func, literal, not_, or_, select, update from sqlalchemy.engine.base import Engine from sqlalchemy.exc import IntegrityError, SQLAlchemyError -from sqlalchemy.orm import joinedload +from sqlalchemy.orm import aliased, joinedload from sqlalchemy.orm.attributes import InstrumentedAttribute, flag_modified from sqlalchemy.orm.session import Session from pyrit.common.deprecation import print_deprecation_message +from pyrit.common.pagination import decode_keyset_cursor, encode_keyset_cursor, fingerprint_filters if TYPE_CHECKING: from pyrit.memory.memory_embedding import MemoryEmbedding @@ -143,6 +144,475 @@ class _PreparedScorableContent: value_sha256: str +@dataclass(frozen=True, slots=True, kw_only=True) +class SeedExample: + """One logical seed example and its persisted members.""" + + example_id: uuid.UUID + dataset_name: str | None + seed_ids: list[uuid.UUID] + members: list[SeedExampleMember] + piece_count: int + objective_count: int + modalities: list[str] + seed_types: list[str] + harm_categories: list[str] + has_unlabeled_harm: bool + + def __getitem__(self, key: str) -> Any: + """ + Allow response-style field access alongside typed attributes. + + Returns: + Any: The requested field value. + """ + return getattr(self, key) + + +@dataclass(frozen=True, slots=True, kw_only=True) +class SeedExamplePage: + """A bounded page of logical seed examples.""" + + items: list[SeedExample] + total: int + next_cursor: str | None + + +@dataclass(frozen=True, slots=True) +class SeedExampleMember: + """Side-effect-free persisted seed data for browsing.""" + + id: uuid.UUID + prompt_group_id: uuid.UUID | None + seed_type: str + data_type: str + value: str + value_sha256: str | None + role: str | None + sequence: int | None + name: str | None + dataset_name: str | None + harm_categories: list[str] | None + description: str | None + source: str | None + authors: list[str] | None + groups: list[str] | None + date_added: datetime + added_by: str + prompt_metadata: dict[str, Any] | None + parameters: list[str] | None + is_jinja_template: bool | None + + @property + def metadata(self) -> dict[str, Any] | None: + """Persisted metadata under the Seed model's public name.""" + return self.prompt_metadata + + +@dataclass(frozen=True, slots=True) +class SeedExampleDatasetScope: + """Explicit dataset namespace for logical seed-example browsing.""" + + kind: Literal["named", "unnamed"] + name: str | None = None + + @classmethod + def named(cls, name: str) -> SeedExampleDatasetScope: + """ + Create a named dataset scope. + + Args: + name: The non-empty persisted dataset name. + + Returns: + The named dataset scope. + + Raises: + ValueError: If the name is empty. + """ + if not name: + raise ValueError("A named dataset scope requires a non-empty name") + return cls(kind="named", name=name) + + @classmethod + def unnamed(cls) -> SeedExampleDatasetScope: + """ + Create the combined NULL/empty dataset scope. + + Returns: + The unnamed dataset scope. + """ + return cls(kind="unnamed") + + def __post_init__(self) -> None: + """ + Validate the scope's name and kind combination. + + Raises: + ValueError: If the scope kind and name do not agree. + """ + if self.kind not in {"named", "unnamed"}: + raise ValueError(f"Unsupported dataset scope kind: {self.kind}") + if self.kind == "named" and not self.name: + raise ValueError("A named dataset scope requires a non-empty name") + if self.kind == "unnamed" and self.name is not None: + raise ValueError("An unnamed dataset scope cannot have a name") + + +@dataclass(frozen=True, slots=True, kw_only=True) +class _SeedExampleQuery: + """Immutable filters and pagination state for a logical seed example query.""" + + dataset_scope: SeedExampleDatasetScope + limit: int + cursor: Any + fingerprint: str + data_types: tuple[str, ...] + harm_categories: tuple[str, ...] + seed_types: tuple[str, ...] + value_search: str | None + + +def _seed_example_logical_id(member: Any) -> Any: + """Return the logical example key for a seed entry or alias.""" + return func.coalesce(member.prompt_group_id, member.id) + + +def _seed_example_order_key(logical_id: Any) -> Any: + """Return the canonical textual UUID key shared by ordering and seeking.""" + return func.lower(cast(logical_id, String(36))) + + +def _seed_example_dataset_condition(member: Any, *, dataset_name: str | None) -> Any: + """ + Build the dataset scope condition used by seed example queries. + + Returns: + Any: The SQLAlchemy condition for the requested dataset scope. + """ + if dataset_name is None: + return or_(member.dataset_name.is_(None), member.dataset_name == "") + return member.dataset_name == dataset_name + + +def _seed_example_filter_fingerprint( + *, + dataset_name: str | None, + data_types: tuple[str, ...], + harm_categories: tuple[str, ...], + seed_types: tuple[str, ...], + value_search: str | None, +) -> str: + """Return the stable identity of the effective seed example filters.""" + return fingerprint_filters( + filters={ + "dataset_name": dataset_name, + "data_types": data_types, + "harm_categories": harm_categories, + "seed_types": seed_types, + "value_search": value_search or "", + }, + length=64, + ) + + +def _build_seed_example_query( + *, + dataset_scope: SeedExampleDatasetScope, + limit: int, + cursor: str | None, + data_types: Sequence[str] | None, + harm_categories: Sequence[str] | None, + seed_types: Sequence[str] | None, + value_search: str | None, +) -> _SeedExampleQuery: + """ + Normalize seed example filters and decode a filter-bound cursor. + + Returns: + _SeedExampleQuery: The immutable effective query state. + + Raises: + ValueError: If ``limit`` is invalid or the cursor is malformed or mismatched. + """ + if not 1 <= limit <= 100: + raise ValueError("limit must be between 1 and 100") + + normalized_types = tuple(sorted(set(data_types or ()))) + normalized_harms = tuple(sorted(set(harm_categories or ()))) + normalized_seed_types = tuple(sorted(set(seed_types or ()))) + if len(normalized_types) + len(normalized_harms) + len(normalized_seed_types) > 100: + raise ValueError("Too many seed example filter values") + fingerprint = _seed_example_filter_fingerprint( + dataset_name=dataset_scope.name if dataset_scope.kind == "named" else None, + data_types=normalized_types, + harm_categories=normalized_harms, + seed_types=normalized_seed_types, + value_search=value_search, + ) + decoded_cursor = decode_keyset_cursor(cursor=cursor, fingerprint=fingerprint) + if cursor is not None and decoded_cursor is None: + raise ValueError("Invalid or filter-mismatched seed example cursor") + return _SeedExampleQuery( + dataset_scope=dataset_scope, + limit=limit, + cursor=decoded_cursor, + fingerprint=fingerprint, + data_types=normalized_types, + harm_categories=normalized_harms, + seed_types=normalized_seed_types, + value_search=value_search, + ) + + +def _seed_example_member_predicate( + member: Any, + *, + query: _SeedExampleQuery, + filter_name: str, + harm_condition_builder: Callable[[Any, Sequence[str]], Any], +) -> Any: + """ + Build one member-level predicate for a logical example filter. + + Returns: + Any: The SQLAlchemy condition for the requested member filter. + """ + predicates: list[Any] = [] + if filter_name == "data_types" and query.data_types: + predicates.append(member.data_type.in_(query.data_types)) + if filter_name == "seed_types" and query.seed_types: + predicates.append(member.seed_type.in_(query.seed_types)) + if filter_name == "harm_categories" and query.harm_categories: + predicates.append(harm_condition_builder(member.harm_categories, query.harm_categories)) + if filter_name == "value_search" and query.value_search: + escaped = query.value_search.lower().replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") + predicates.extend([member.data_type == "text", func.lower(member.value).like(f"%{escaped}%", escape="\\")]) + return and_(*predicates) if predicates else literal(True) + + +def _build_seed_example_filter_conditions( + *, + query: _SeedExampleQuery, + logical_id: Any, + harm_condition_builder: Callable[[Any, Sequence[str]], Any], +) -> list[Any]: + """ + Build the dataset and independent member-existence filters. + + Returns: + list[Any]: SQLAlchemy conditions for the grouped logical-example query. + """ + dataset_name = query.dataset_scope.name if query.dataset_scope.kind == "named" else None + conditions: list[Any] = [_seed_example_dataset_condition(SeedEntry, dataset_name=dataset_name)] + filter_values = ( + ("data_types", query.data_types), + ("harm_categories", query.harm_categories), + ("seed_types", query.seed_types), + ("value_search", (query.value_search,) if query.value_search else ()), + ) + for filter_name, values in filter_values: + if not values: + continue + member = aliased(SeedEntry) + conditions.append( + exists( + select(1).where( + _seed_example_dataset_condition(member, dataset_name=dataset_name), + _seed_example_logical_id(member) == logical_id, + _seed_example_member_predicate( + member, + query=query, + filter_name=filter_name, + harm_condition_builder=harm_condition_builder, + ), + ) + ) + ) + return conditions + + +def _seed_example_order_by(logical_id: Any) -> list[Any]: + """Return deterministic member ordering within each logical example.""" + return [ + logical_id, + case( + (SeedEntry.seed_type == "objective", 0), + (SeedEntry.seed_type == "simulated_conversation", 1), + else_=2, + ), + case((SeedEntry.sequence.is_(None), 1), else_=0), + SeedEntry.sequence, + SeedEntry.id, + ] + + +def _seed_example_keyset_seek_condition(*, grouped_subquery: Any, cursor: Any) -> Any: + """ + Build the seek predicate for first-added descending, ID descending order. + + Returns: + Any: The SQLAlchemy condition selecting rows after the cursor. + """ + return or_( + grouped_subquery.c.first_added < cursor.timestamp, + and_( + grouped_subquery.c.first_added == cursor.timestamp, + grouped_subquery.c.example_id_key < cursor.identifier, + ), + ) + + +def _query_seed_example_page( + *, + session: Session, + query: _SeedExampleQuery, + harm_condition_builder: Callable[[Any, Sequence[str]], Any], +) -> tuple[list[Any], int]: + """ + Select one logical seed example page and its total logical count. + + Returns: + tuple[list[Any], int]: Over-fetched page rows and the logical total. + """ + logical_id = _seed_example_logical_id(SeedEntry) + logical_id_key = _seed_example_order_key(logical_id) + grouped_subquery = ( + select( + logical_id.label("example_id"), + logical_id_key.label("example_id_key"), + func.min(SeedEntry.date_added).label("first_added"), + ) + .where( + and_( + *_build_seed_example_filter_conditions( + query=query, + logical_id=logical_id, + harm_condition_builder=harm_condition_builder, + ) + ) + ) + .group_by(logical_id, logical_id_key) + .subquery() + ) + group_query = select( + grouped_subquery.c.example_id, + grouped_subquery.c.first_added, + literal(query.dataset_scope.name if query.dataset_scope.kind == "named" else None).label("dataset_name"), + ) + if query.cursor is not None: + group_query = group_query.where( + _seed_example_keyset_seek_condition(grouped_subquery=grouped_subquery, cursor=query.cursor) + ) + page_rows = list( + session.execute( + group_query.order_by(grouped_subquery.c.first_added.desc(), grouped_subquery.c.example_id_key.desc()).limit( + query.limit + 1 + ) + ).all() + ) + total = int(session.execute(select(func.count()).select_from(grouped_subquery)).scalar_one()) + return page_rows, total + + +def _query_seed_example_members( + *, session: Session, query: _SeedExampleQuery, example_ids: Sequence[uuid.UUID] +) -> list[SeedEntry]: + """ + Load all members for the selected logical examples in one query. + + Returns: + list[SeedEntry]: Persisted members ordered within each logical example. + """ + logical_id = _seed_example_logical_id(SeedEntry) + return list( + session.execute( + select(SeedEntry) + .where( + _seed_example_dataset_condition( + SeedEntry, + dataset_name=query.dataset_scope.name if query.dataset_scope.kind == "named" else None, + ), + logical_id.in_(example_ids), + ) + .order_by(*_seed_example_order_by(logical_id)) + ).scalars() + ) + + +def _build_seed_examples(*, rows: Sequence[Any], members: Sequence[SeedEntry]) -> list[SeedExample]: + """ + Materialize logical seed examples from selected rows and persisted members. + + Returns: + list[SeedExample]: The logical examples represented by the selected rows. + """ + selected_ids = [row.example_id for row in rows] + members_by_group: dict[uuid.UUID, list[SeedEntry]] = {example_id: [] for example_id in selected_ids} + for entry in members: + members_by_group.setdefault(entry.prompt_group_id or entry.id, []).append(entry) + + items: list[SeedExample] = [] + for row in rows: + entries = members_by_group[row.example_id] + items.append( + SeedExample( + example_id=row.example_id, + dataset_name=row.dataset_name, + seed_ids=[entry.id for entry in entries], + members=[ + SeedExampleMember( + id=entry.id, + prompt_group_id=entry.prompt_group_id, + seed_type=entry.seed_type, + data_type=entry.data_type, + value=entry.value, + value_sha256=entry.value_sha256, + role=entry.role, + sequence=entry.sequence, + name=entry.name, + dataset_name=entry.dataset_name, + harm_categories=list(entry.harm_categories) if entry.harm_categories is not None else None, + description=entry.description, + source=entry.source, + authors=list(entry.authors) if entry.authors is not None else None, + groups=list(entry.groups) if entry.groups is not None else None, + date_added=entry.date_added, + added_by=entry.added_by, + prompt_metadata=dict(entry.prompt_metadata) if entry.prompt_metadata is not None else None, + parameters=list(entry.parameters) if entry.parameters is not None else None, + is_jinja_template=True if entry.parameters else None, + ) + for entry in entries + ], + piece_count=len(entries), + objective_count=sum(entry.seed_type == "objective" for entry in entries), + modalities=sorted({entry.data_type for entry in entries}), + seed_types=sorted({entry.seed_type for entry in entries}), + harm_categories=sorted({category for entry in entries for category in (entry.harm_categories or [])}), + has_unlabeled_harm=any(not entry.harm_categories for entry in entries), + ) + ) + return items + + +def _build_seed_example_next_cursor(*, rows: Sequence[Any], query: _SeedExampleQuery) -> str | None: + """ + Build the continuation cursor when the page was over-fetched. + + Returns: + str | None: The filter-bound continuation cursor, if another page exists. + """ + if len(rows) <= query.limit: + return None + last = rows[query.limit - 1] + return encode_keyset_cursor( + timestamp=last.first_added, + identifier=str(last.example_id), + fingerprint=query.fingerprint, + ) + + class AttackResultKeysetCursor(NamedTuple): """ Keyset (seek) anchor identifying the last attack result on a page. @@ -682,6 +1152,80 @@ def add_conversation_to_memory(self, *, conversation: Conversation) -> None: """ self._insert_conversation(conversation=conversation) + def get_seed_example_page( + self, + *, + dataset_scope: SeedExampleDatasetScope, + limit: int, + cursor: str | None = None, + data_types: Sequence[str] | None = None, + harm_categories: Sequence[str] | None = None, + seed_types: Sequence[str] | None = None, + value_search: str | None = None, + ) -> SeedExamplePage: + """ + Read a bounded page of logical seed examples and their members. + + Returns: + SeedExamplePage: The selected logical examples and continuation metadata. + + Raises: + ValueError: If ``limit`` is invalid or the cursor is malformed or + bound to different filters. + """ + query = _build_seed_example_query( + dataset_scope=dataset_scope, + limit=limit, + cursor=cursor, + data_types=data_types, + harm_categories=harm_categories, + seed_types=seed_types, + value_search=value_search, + ) + with closing(self.get_session()) as session: + page_rows, total = _query_seed_example_page( + session=session, + query=query, + harm_condition_builder=self._seed_example_harm_condition, + ) + selected_rows = page_rows[:limit] + if not selected_rows: + return SeedExamplePage(items=[], total=total, next_cursor=None) + + selected_ids = [row.example_id for row in selected_rows] + members = _query_seed_example_members(session=session, query=query, example_ids=selected_ids) + + return SeedExamplePage( + items=_build_seed_examples(rows=selected_rows, members=members), + total=total, + next_cursor=_build_seed_example_next_cursor(rows=page_rows, query=query), + ) + + def _seed_example_harm_condition(self, column: Any, categories: Sequence[str]) -> Any: + """ + Build the backend-native any-category predicate for seed harm labels. + + Returns: + Any: The backend-specific SQLAlchemy predicate. + """ + return self._get_seed_harm_category_condition(json_column=column, categories=categories) + + @abc.abstractmethod + def _get_seed_harm_category_condition( + self, *, json_column: InstrumentedAttribute[Any], categories: Sequence[str] + ) -> Any: + """ + Build an aliased-column-safe, case-insensitive any-category predicate. + + Args: + json_column: The possibly aliased harm-category JSON column. + categories: Categories whose presence should match with OR semantics. + + Returns: + Any: A backend-specific SQLAlchemy predicate. + """ + ... + def add_message_pieces_to_memory(self, *, message_pieces: Sequence[MessagePiece]) -> None: """ Insert a list of message pieces into the memory storage. diff --git a/pyrit/memory/sqlite_memory.py b/pyrit/memory/sqlite_memory.py index 04421c09f3..1c9a7512b5 100644 --- a/pyrit/memory/sqlite_memory.py +++ b/pyrit/memory/sqlite_memory.py @@ -11,7 +11,7 @@ from pathlib import Path from typing import Any, Literal -from sqlalchemy import and_, case, create_engine, exists, func, or_, select, text +from sqlalchemy import and_, case, create_engine, exists, func, literal, or_, select, text from sqlalchemy.engine.base import Engine from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.orm import InstrumentedAttribute, sessionmaker @@ -280,6 +280,20 @@ def _get_condition_json_array_match( combined = joiner.join(conditions) return text(f"({combined})").bindparams(**bindparams_dict) + def _get_seed_harm_category_condition( + self, *, json_column: InstrumentedAttribute[Any], categories: Sequence[str] + ) -> Any: + """ + Build an aliased-column-safe SQLite harm-category membership predicate. + + Returns: + Any: A SQLAlchemy predicate matching any requested category. + """ + values = [category.lower() for category in categories] + array = func.json_extract(json_column, literal("$")) + elements = func.json_each(array).table_valued("value") + return exists(select(1).select_from(elements).where(func.lower(elements.c.value).in_(values))) + def get_all_table_models(self) -> list[type[Base]]: """ Return a list of all table models used in the database by inspecting the Base registry. diff --git a/tests/integration/memory/test_seed_browsing_azure_sql_integration.py b/tests/integration/memory/test_seed_browsing_azure_sql_integration.py new file mode 100644 index 0000000000..7021a18233 --- /dev/null +++ b/tests/integration/memory/test_seed_browsing_azure_sql_integration.py @@ -0,0 +1,159 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Azure SQL execution contract for the #2748 seed browsing seam.""" + +from __future__ import annotations + +from contextlib import closing +from datetime import UTC, datetime, timedelta +from uuid import uuid4 + +import pytest + +from pyrit.memory import AzureSQLMemory, SeedExampleDatasetScope +from pyrit.memory.memory_models import SeedEntry +from pyrit.models import SeedObjective, SeedPrompt + + +@pytest.mark.run_only_if_all_tests +async def test_seed_browsing_contract_on_azure_sql(azuresql_instance: AzureSQLMemory): + test_id = str(uuid4()) + dataset = f"2748-azure-{test_id}" + shared_group = uuid4() + newer_group = uuid4() + older_group = uuid4() + base_time = datetime(2024, 1, 1, tzinfo=UTC) + seeds = [ + SeedPrompt( + value="https://example.com/violence", + dataset_name=dataset, + prompt_group_id=shared_group, + data_type="url", + harm_categories=["violence", "hate", 'special_%_"_é'], + date_added=base_time, + added_by=test_id, + ), + SeedObjective( + value="objective", + dataset_name=dataset, + prompt_group_id=shared_group, + date_added=base_time + timedelta(seconds=1), + added_by=test_id, + ), + SeedPrompt( + value="literal 100% a_b", + dataset_name=dataset, + prompt_group_id=newer_group, + date_added=base_time + timedelta(days=1), + harm_categories=["nonviolence", "violence-extra"], + added_by=test_id, + ), + SeedPrompt( + value="older", + dataset_name=dataset, + prompt_group_id=older_group, + date_added=base_time, + harm_categories=[], + added_by=test_id, + ), + SeedPrompt( + value="unnamed null", + prompt_group_id=shared_group, + harm_categories=None, + date_added=base_time, + added_by=test_id, + ), + SeedPrompt( + value="unnamed empty", + dataset_name="", + prompt_group_id=shared_group, + harm_categories=[], + date_added=base_time + timedelta(seconds=1), + added_by=test_id, + ), + ] + await azuresql_instance.add_seeds_to_memory_async(seeds=seeds, added_by=test_id) + + try: + named = SeedExampleDatasetScope.named(dataset) + unnamed = SeedExampleDatasetScope.unnamed() + + named_page = azuresql_instance.get_seed_example_page(dataset_scope=named, limit=100) + named_ids = {item.example_id for item in named_page.items} + assert named_ids == {shared_group, newer_group, older_group} + assert [item.example_id for item in named_page.items] == [ + newer_group, + *sorted((shared_group, older_group), reverse=True), + ] + shared = next(item for item in named_page.items if item.example_id == shared_group) + assert {member.id for member in shared.members} == {seeds[0].id, seeds[1].id} + assert shared.objective_count == 1 + + assert { + item.example_id for item in azuresql_instance.get_seed_example_page(dataset_scope=unnamed, limit=100).items + } == {shared_group} + + assert { + item.example_id + for item in azuresql_instance.get_seed_example_page( + dataset_scope=named, data_types=["url"], limit=100 + ).items + } == {shared_group} + assert { + item.example_id + for item in azuresql_instance.get_seed_example_page( + dataset_scope=named, seed_types=["objective"], limit=100 + ).items + } == {shared_group} + assert { + item.example_id + for item in azuresql_instance.get_seed_example_page( + dataset_scope=named, harm_categories=["VIOLENCE"], limit=100 + ).items + } == {shared_group} + assert { + item.example_id + for item in azuresql_instance.get_seed_example_page( + dataset_scope=named, harm_categories=["missing", "HATE"], limit=100 + ).items + } == {shared_group} + assert { + item.example_id + for item in azuresql_instance.get_seed_example_page( + dataset_scope=named, harm_categories=['special_%_"_É'], limit=100 + ).items + } == {shared_group} + assert { + item.example_id + for item in azuresql_instance.get_seed_example_page( + dataset_scope=named, harm_categories=["nonviolence"], limit=100 + ).items + } == {newer_group} + + assert { + item.example_id + for item in azuresql_instance.get_seed_example_page( + dataset_scope=named, value_search="100%", limit=100 + ).items + } == {newer_group} + assert { + item.example_id + for item in azuresql_instance.get_seed_example_page( + dataset_scope=named, value_search="a_b", limit=100 + ).items + } == {newer_group} + + unlabeled = next(item for item in named_page.items if item.example_id == older_group) + assert unlabeled.has_unlabeled_harm is True + + first = azuresql_instance.get_seed_example_page(dataset_scope=named, limit=2) + assert first.next_cursor is not None + second = azuresql_instance.get_seed_example_page(dataset_scope=named, limit=100, cursor=first.next_cursor) + paged_ids = [item.example_id for item in first.items + second.items] + assert len(paged_ids) == len(set(paged_ids)) == 3 + assert set(paged_ids) == named_ids + finally: + with closing(azuresql_instance.get_session()) as session: + session.query(SeedEntry).filter(SeedEntry.added_by == test_id).delete(synchronize_session=False) + session.commit() diff --git a/tests/unit/memory/memory_interface/test_interface_seed_browsing_contract.py b/tests/unit/memory/memory_interface/test_interface_seed_browsing_contract.py index bbd1954f3c..2d34fefc29 100644 --- a/tests/unit/memory/memory_interface/test_interface_seed_browsing_contract.py +++ b/tests/unit/memory/memory_interface/test_interface_seed_browsing_contract.py @@ -5,13 +5,19 @@ from __future__ import annotations +import base64 +import json from datetime import UTC, datetime, timedelta +from pathlib import Path from typing import TYPE_CHECKING, Any +from unittest.mock import patch from uuid import UUID, uuid4 import pytest from sqlalchemy import event +from pyrit.memory import SeedExampleDatasetScope +from pyrit.memory.memory_models import SeedEntry from pyrit.models import SeedObjective, SeedPrompt, SeedSimulatedConversation if TYPE_CHECKING: @@ -30,6 +36,10 @@ def _page(memory: MemoryInterface, **kwargs: Any) -> Any: """Call the proposed narrow, database-backed logical-example page helper.""" helper = getattr(memory, "get_seed_example_page", None) assert helper is not None, "RED: MemoryInterface.get_seed_example_page is not implemented" + dataset_name = kwargs.pop("dataset_name", None) + kwargs["dataset_scope"] = ( + SeedExampleDatasetScope.named(dataset_name) if dataset_name is not None else SeedExampleDatasetScope.unnamed() + ) return helper(**kwargs) @@ -38,6 +48,91 @@ async def _add(memory: MemoryInterface, *seeds: SeedPrompt | SeedObjective | See class TestSeedBrowsingMemoryContract: + def test_scope_and_page_limits_are_explicit(self): + with pytest.raises(ValueError): + SeedExampleDatasetScope(kind="invalid") # type: ignore[arg-type] + + with pytest.raises(ValueError): + SeedExampleDatasetScope.named("") + + async def test_named_and_unnamed_scopes_are_distinct(self, sqlite_instance: MemoryInterface): + shared_group = uuid4() + unnamed_null = SeedPrompt(value="unnamed-null", prompt_group_id=shared_group) + unnamed_empty = SeedPrompt(value="unnamed-empty", dataset_name="", prompt_group_id=shared_group) + named_seed = SeedPrompt(value="named", dataset_name=DATASET, prompt_group_id=shared_group) + await _add( + sqlite_instance, + unnamed_null, + unnamed_empty, + named_seed, + ) + + unnamed = _page(sqlite_instance, limit=10) + named = _page(sqlite_instance, dataset_name=DATASET, limit=10) + assert len(_field(unnamed, "items")) == 1 + assert {str(member.id) for member in _field(unnamed, "items")[0].members} == { + str(unnamed_null.id), + str(unnamed_empty.id), + } + assert len(_field(named, "items")) == 1 + assert {str(member.id) for member in _field(named, "items")[0].members} == { + str(named_seed.id), + } + + with pytest.raises(ValueError): + _page(sqlite_instance, dataset_name=DATASET, limit=101) + + async def test_browsing_uses_persisted_rows_without_get_seed_or_filesystem_access( + self, sqlite_instance: MemoryInterface + ): + legacy_path = "C:/not-present/legacy-prompt.yaml" + legacy_value = json.dumps( + { + "num_turns": 2, + "sequence": 0, + "adversarial_chat_system_prompt_path": legacy_path, + "simulated_target_system_prompt_path": legacy_path, + } + ) + entry = SeedEntry(entry=SeedPrompt(value="stored", dataset_name=DATASET)) + entry.value = legacy_value + entry.value_sha256 = "legacy-hash" + entry.prompt_metadata = {"persisted": "yes"} + entry.seed_type = "simulated_conversation" + entry.date_added = datetime(2024, 1, 1, tzinfo=UTC) + entry.added_by = "legacy-test" + entry_id = entry.id + with sqlite_instance.get_session() as session: + session.add(entry) + session.commit() + + with ( + patch.object(SeedEntry, "get_seed", side_effect=AssertionError("get_seed called")), + patch.object(Path, "is_file", side_effect=AssertionError("filesystem touched")), + ): + page = _page(sqlite_instance, dataset_name=DATASET, limit=10) + + member = _field(page, "items")[0].members[0] + assert member.id == entry_id + assert member.value == legacy_value + assert member.value_sha256 == "legacy-hash" + assert member.metadata == {"persisted": "yes"} + assert member.seed_type == "simulated_conversation" + + async def test_template_persisted_value_and_parameters_are_not_rendered(self, sqlite_instance: MemoryInterface): + template = SeedPrompt( + value="{{ name }}", + dataset_name=DATASET, + parameters=["name"], + is_jinja_template=True, + ) + await _add(sqlite_instance, template) + with patch.object(SeedEntry, "get_seed", side_effect=AssertionError("get_seed called")): + member = _field(_page(sqlite_instance, dataset_name=DATASET, limit=10), "items")[0].members[0] + assert member.value == "{{ name }}" + assert member.parameters == ["name"] + assert member.is_jinja_template is True + async def test_logical_identity_uses_group_id_else_seed_id(self, sqlite_instance: MemoryInterface): group_id = uuid4() grouped = SeedPrompt(value="grouped", dataset_name=DATASET, prompt_group_id=group_id) @@ -272,12 +367,41 @@ async def test_harm_matching_is_case_insensitive_whole_value_and_missing_is_unla self, sqlite_instance: MemoryInterface ): labeled = SeedPrompt(value="labeled", dataset_name=DATASET, harm_categories=["VIOLENCE"]) + multi_labeled = SeedPrompt( + value="multi-labeled", dataset_name=DATASET, harm_categories=["hate", 'special_%_"_é'] + ) + substring_only = SeedPrompt(value="substring-only", dataset_name=DATASET, harm_categories=["nonviolence"]) + suffix_only = SeedPrompt(value="suffix-only", dataset_name=DATASET, harm_categories=["violence-extra"]) unlabeled = SeedPrompt(value="unlabeled", dataset_name=DATASET, harm_categories=[]) - await _add(sqlite_instance, labeled, unlabeled) + null_labeled = SeedPrompt(value="null-labeled", dataset_name=DATASET, harm_categories=None) + grouped = uuid4() + grouped_unmatched_member = SeedPrompt( + value="grouped-unmatched", dataset_name=DATASET, prompt_group_id=grouped, harm_categories=[] + ) + grouped_matching_member = SeedPrompt( + value="grouped-matching", dataset_name=DATASET, prompt_group_id=grouped, harm_categories=["violence"] + ) + await _add( + sqlite_instance, + labeled, + multi_labeled, + substring_only, + suffix_only, + unlabeled, + null_labeled, + grouped_unmatched_member, + grouped_matching_member, + ) exact = _page(sqlite_instance, dataset_name=DATASET, harm_categories=["violence"], limit=10) substring = _page(sqlite_instance, dataset_name=DATASET, harm_categories=["vio"], limit=10) - assert len(_field(exact, "items")) == 1 - assert str(_field(exact, "items")[0]["seed_ids"][0]) == str(labeled.id) + assert {str(item.example_id) for item in _field(exact, "items")} == { + str(labeled.id), + str(grouped), + } + multi = _page(sqlite_instance, dataset_name=DATASET, harm_categories=["missing", "HATE"], limit=10) + assert {str(item.example_id) for item in _field(multi, "items")} == {str(multi_labeled.id)} + special = _page(sqlite_instance, dataset_name=DATASET, harm_categories=['special_%_"_É'], limit=10) + assert {str(item.example_id) for item in _field(special, "items")} == {str(multi_labeled.id)} assert len(_field(substring, "items")) == 0 all_items = _field(_page(sqlite_instance, dataset_name=DATASET, limit=10), "items") unlabeled_item = next( @@ -393,6 +517,13 @@ async def test_malformed_cursor_and_filter_mismatch_are_rejected(self, sqlite_in limit=1, cursor=_field(first, "next_cursor"), ) + payload = json.loads( + base64.urlsafe_b64decode(_field(first, "next_cursor") + "=" * (-len(_field(first, "next_cursor")) % 4)) + ) + payload["i"] = "not-a-uuid" + malformed_identifier = base64.urlsafe_b64encode(json.dumps(payload).encode()).decode().rstrip("=") + with pytest.raises(ValueError): + _page(sqlite_instance, dataset_name=DATASET, limit=1, cursor=malformed_identifier) @pytest.mark.parametrize( ("changed_filters", "changed_dataset"), diff --git a/tests/unit/memory/memory_interface/test_interface_seed_browsing_sql.py b/tests/unit/memory/memory_interface/test_interface_seed_browsing_sql.py new file mode 100644 index 0000000000..418b3fe77f --- /dev/null +++ b/tests/unit/memory/memory_interface/test_interface_seed_browsing_sql.py @@ -0,0 +1,108 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""SQL dialect compilation coverage for the #2748 seed browsing seam.""" + +from __future__ import annotations + +from datetime import UTC, datetime +from uuid import uuid4 + +from sqlalchemy.dialects import mssql, sqlite + +from pyrit.common.pagination import encode_keyset_cursor +from pyrit.memory.azure_sql_memory import AzureSQLMemory +from pyrit.memory.memory_interface import ( + SeedExampleDatasetScope, + _build_seed_example_query, + _query_seed_example_members, + _query_seed_example_page, +) +from pyrit.memory.sqlite_memory import SQLiteMemory + + +class _Result: + def all(self): + return [] + + def scalar_one(self): + return 0 + + def scalars(self): + return [] + + +class _StatementCapture: + def __init__(self): + self.statements = [] + + def execute(self, statement): + self.statements.append(statement) + return _Result() + + +def _query(*, scope: SeedExampleDatasetScope, cursor: str | None = None): + return _build_seed_example_query( + dataset_scope=scope, + limit=100, + cursor=cursor, + data_types=["url", "image_path"], + harm_categories=["violence", "self_harm_%"], + seed_types=["prompt", "objective"], + value_search=r"literal_%\value", + ) + + +def _capture_statements(*, scope: SeedExampleDatasetScope, memory_type): + initial = _query(scope=scope) + cursor = encode_keyset_cursor( + timestamp=datetime(2024, 1, 2, tzinfo=UTC), + identifier=str(uuid4()), + fingerprint=initial.fingerprint, + ) + query = _query(scope=scope, cursor=cursor) + capture = _StatementCapture() + _query_seed_example_page( + session=capture, + query=query, + harm_condition_builder=object.__new__(memory_type)._seed_example_harm_condition, + ) + _query_seed_example_members( + session=capture, + query=query, + example_ids=[uuid4() for _ in range(100)], + ) + return capture.statements + + +def test_seed_browsing_statements_compile_for_sql_server_and_sqlite(): + for scope in (SeedExampleDatasetScope.named("dataset"), SeedExampleDatasetScope.unnamed()): + sql_server_statements = _capture_statements(scope=scope, memory_type=AzureSQLMemory) + sqlite_statements = _capture_statements(scope=scope, memory_type=SQLiteMemory) + sql_server_compiled = [ + statement.compile(dialect=mssql.dialect(), compile_kwargs={"render_postcompile": True}) + for statement in sql_server_statements + ] + sql_server_sql = [str(statement).lower() for statement in sql_server_compiled] + sqlite_sql = [str(statement.compile(dialect=sqlite.dialect())).lower() for statement in sqlite_statements] + + assert all(sql for sql in sql_server_sql) + assert all(sql for sql in sqlite_sql) + assert all(len(statement.params) < 2100 for statement in sql_server_compiled) + assert any("openjson" in sql for sql in sql_server_sql) + assert any("json_each" in sql for sql in sqlite_sql) + assert any("group by" in sql and "min" in sql for sql in sql_server_sql) + assert any("group by" in sql and "min" in sql for sql in sqlite_sql) + assert any("order by" in sql and "offset" in sql or "top" in sql for sql in sql_server_sql) + assert any("order by" in sql and "limit" in sql for sql in sqlite_sql) + assert any("sequence" in sql and "case" in sql for sql in sql_server_sql) + page_sql = sql_server_sql[0] + assert "lower(cast(coalesce" in page_sql + assert "example_id_key <" in page_sql + assert "example_id_key desc" in page_sql + assert "openjson(case when" in page_sql + assert "isjson([seedpromptentries_" in page_sql + assert "left(ltrim([seedpromptentries_" in page_sql + assert "else :" in page_sql + assert "openjson([seedpromptentries_" not in page_sql + assert "openjson(json_query(" not in page_sql From 25ef5934829b866f6dfc6156883c19c024a8449b Mon Sep 17 00:00:00 2001 From: Jackson Severino da Rocha Date: Mon, 28 Sep 2026 18:43:31 -0300 Subject: [PATCH 4/9] feat: add seed example browsing service --- pyrit/backend/models/datasets.py | 70 +++++++++++ pyrit/backend/services/dataset_service.py | 141 +++++++++++++++++++++- 2 files changed, 210 insertions(+), 1 deletion(-) diff --git a/pyrit/backend/models/datasets.py b/pyrit/backend/models/datasets.py index 9f8814534e..02f123b899 100644 --- a/pyrit/backend/models/datasets.py +++ b/pyrit/backend/models/datasets.py @@ -9,8 +9,14 @@ listing available datasets. """ +from datetime import datetime +from typing import Any +from uuid import UUID + from pydantic import BaseModel, Field +from pyrit.backend.models.common import PaginationInfo + class DatasetInfo(BaseModel): """Metadata about a single available dataset.""" @@ -41,3 +47,67 @@ class DatasetListResponse(BaseModel): """Response for listing available datasets.""" items: list[DatasetInfo] = Field(..., description="List of available datasets") + + +class SeedExampleMemberView(BaseModel): + """Persisted member of a logical seed example.""" + + id: UUID + prompt_group_id: UUID | None = None + seed_type: str + data_type: str + value: str + value_sha256: str | None = None + role: str | None = None + sequence: int | None = None + name: str | None = None + dataset_name: str | None = None + harm_categories: list[str] | None = None + description: str | None = None + source: str | None = None + authors: list[str] | None = None + groups: list[str] | None = None + date_added: datetime + added_by: str + metadata: dict[str, Any] | None = None + parameters: list[str] | None = None + is_jinja_template: bool | None = None + + +class SeedExampleSummary(BaseModel): + """List representation of one complete logical seed example.""" + + example_id: UUID + dataset_name: str | None = None + name: str | None = None + preview: str + preview_truncated: bool + seed_ids: list[UUID] + modalities: list[str] + seed_types: list[str] + piece_count: int + objective_count: int + harm_categories: list[str] + has_unlabeled_harm: bool + + +class SeedExampleListResponse(BaseModel): + """Paginated logical seed examples.""" + + items: list[SeedExampleSummary] + pagination: PaginationInfo + + +class SeedExampleDetailResponse(BaseModel): + """Complete persisted logical seed example.""" + + example_id: UUID + dataset_name: str | None = None + seed_ids: list[UUID] + piece_count: int + objective_count: int + modalities: list[str] + seed_types: list[str] + harm_categories: list[str] + has_unlabeled_harm: bool + members: list[SeedExampleMemberView] diff --git a/pyrit/backend/services/dataset_service.py b/pyrit/backend/services/dataset_service.py index 4d14b1699f..caef27d564 100644 --- a/pyrit/backend/services/dataset_service.py +++ b/pyrit/backend/services/dataset_service.py @@ -10,13 +10,20 @@ import logging from collections.abc import Sequence from functools import lru_cache +from re import match +from urllib.parse import urlparse +from pyrit.backend.models.common import PaginationInfo from pyrit.backend.models.datasets import ( DatasetInfo, DatasetListResponse, + SeedExampleDetailResponse, + SeedExampleListResponse, + SeedExampleMemberView, + SeedExampleSummary, ) from pyrit.datasets import SeedDatasetProvider -from pyrit.memory import CentralMemory +from pyrit.memory import CentralMemory, SeedExampleDatasetScope from pyrit.models import SeedDatasetSummary logger = logging.getLogger(__name__) @@ -94,6 +101,138 @@ 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[str] | None = None, + harm_categories: Sequence[str] | None = None, + seed_types: Sequence[str] | None = None, + ) -> SeedExampleListResponse: + """ + List logical seed examples using Memory's database-backed page query. + + Returns: + SeedExampleListResponse: The selected logical examples and pagination metadata. + """ + scope = self._selection_scope(selection_key) + page = self._memory.get_seed_example_page( + dataset_scope=scope, + limit=limit, + cursor=cursor, + data_types=data_types, + harm_categories=harm_categories, + seed_types=seed_types, + value_search=search, + ) + items = [self._summary(item) for item in page.items] + return SeedExampleListResponse( + items=items, + pagination=PaginationInfo( + limit=limit, + has_more=page.next_cursor is not None, + next_cursor=page.next_cursor, + prev_cursor=cursor, + ), + ) + + async def get_seed_example_async(self, *, selection_key: str, example_id: str) -> SeedExampleDetailResponse: + """Return one complete logical seed example without materializing seed models.""" + scope = self._selection_scope(selection_key) + page = self._memory.get_seed_example_page(dataset_scope=scope, limit=100) + item = next((candidate for candidate in page.items if str(candidate.example_id) == example_id), None) + if item is None: + raise ValueError(f"Seed example not found: {example_id}") + return SeedExampleDetailResponse( + example_id=item.example_id, + dataset_name=item.dataset_name, + seed_ids=item.seed_ids, + piece_count=item.piece_count, + objective_count=item.objective_count, + modalities=item.modalities, + seed_types=item.seed_types, + harm_categories=item.harm_categories, + has_unlabeled_harm=item.has_unlabeled_harm, + members=[self._member(member) for member in item.members], + ) + + @staticmethod + def _selection_scope(selection_key: str) -> SeedExampleDatasetScope: + """ + Resolve the stable dataset selection namespace. + + Returns: + SeedExampleDatasetScope: The named or unnamed memory query scope. + """ + if selection_key == "dataset:unnamed": + return SeedExampleDatasetScope.unnamed() + prefix = "dataset:named:" + if selection_key.startswith(prefix) and selection_key[len(prefix) :]: + return SeedExampleDatasetScope.named(selection_key[len(prefix) :]) + raise ValueError(f"Invalid dataset selection key: {selection_key}") + + @classmethod + def _summary(cls, item: object) -> SeedExampleSummary: + members = item.members # type: ignore[attr-defined] + preview, truncated = cls._preview(members) + return SeedExampleSummary( + example_id=item.example_id, # type: ignore[attr-defined] + dataset_name=item.dataset_name, # type: ignore[attr-defined] + name=next((member.name for member in members if member.name), None), + preview=preview, + preview_truncated=truncated, + seed_ids=item.seed_ids, # type: ignore[attr-defined] + modalities=item.modalities, # type: ignore[attr-defined] + seed_types=item.seed_types, # type: ignore[attr-defined] + piece_count=item.piece_count, # type: ignore[attr-defined] + objective_count=item.objective_count, # type: ignore[attr-defined] + harm_categories=item.harm_categories, # type: ignore[attr-defined] + has_unlabeled_harm=item.has_unlabeled_harm, # type: ignore[attr-defined] + ) + + @staticmethod + def _preview(members: Sequence[object]) -> tuple[str, bool]: + safe_text = [ + member.value + for member in members + if member.data_type == "text" + and not urlparse(member.value).scheme + and not match(r"^(?:[A-Za-z]:[\\/]|/)", member.value) + ] + if safe_text: + value = max(safe_text, key=len) + return (value[:100] + "...", len(value) > 100) if len(value) > 100 else (value, False) + data_type = members[0].data_type if members else "seed" + return f"{data_type} seed", False + + @staticmethod + def _member(member: object) -> SeedExampleMemberView: + return SeedExampleMemberView( + id=member.id, + prompt_group_id=member.prompt_group_id, + seed_type=member.seed_type, + data_type=member.data_type, + value=member.value, + value_sha256=member.value_sha256, + role=member.role, + sequence=member.sequence, + name=member.name, + dataset_name=member.dataset_name, + harm_categories=member.harm_categories, + description=member.description, + source=member.source, + authors=member.authors, + groups=member.groups, + date_added=member.date_added, + added_by=member.added_by, + metadata=member.metadata, + parameters=member.parameters, + is_jinja_template=member.is_jinja_template, + ) + @staticmethod def _merge_unnamed_summaries(summaries: Sequence[SeedDatasetSummary]) -> SeedDatasetSummary | None: """ From 487f53dfc46229fd3fccb170c2b6c9c5f2521538 Mon Sep 17 00:00:00 2001 From: Jackson Severino da Rocha Date: Mon, 28 Sep 2026 20:52:50 -0300 Subject: [PATCH 5/9] fix: harden seed browsing contract --- pyrit/backend/services/dataset_service.py | 13 ++- ...f2a4c6e8b0d2_persist_seed_template_flag.py | 26 ++++++ pyrit/memory/memory_interface.py | 42 ++++++++- pyrit/memory/memory_models.py | 15 ++++ pyrit/memory/sqlite_memory.py | 11 ++- ..._dataset_service_seed_browsing_contract.py | 40 +++++++++ .../test_interface_seed_browsing_contract.py | 89 ++++++++++++++++++- tests/unit/memory/test_memory_models.py | 26 ++++++ tests/unit/memory/test_migration.py | 26 ++++++ 9 files changed, 280 insertions(+), 8 deletions(-) create mode 100644 pyrit/memory/alembic/versions/f2a4c6e8b0d2_persist_seed_template_flag.py diff --git a/pyrit/backend/services/dataset_service.py b/pyrit/backend/services/dataset_service.py index caef27d564..83d67b5245 100644 --- a/pyrit/backend/services/dataset_service.py +++ b/pyrit/backend/services/dataset_service.py @@ -8,10 +8,12 @@ """ import logging +import ntpath +import posixpath from collections.abc import Sequence from functools import lru_cache -from re import match from urllib.parse import urlparse +from uuid import UUID from pyrit.backend.models.common import PaginationInfo from pyrit.backend.models.datasets import ( @@ -142,8 +144,11 @@ async def list_seed_examples_async( async def get_seed_example_async(self, *, selection_key: str, example_id: str) -> SeedExampleDetailResponse: """Return one complete logical seed example without materializing seed models.""" scope = self._selection_scope(selection_key) - page = self._memory.get_seed_example_page(dataset_scope=scope, limit=100) - item = next((candidate for candidate in page.items if str(candidate.example_id) == example_id), None) + try: + logical_id = UUID(example_id) + except ValueError as exc: + raise ValueError(f"Seed example not found: {example_id}") from exc + item = self._memory.get_seed_example(dataset_scope=scope, example_id=logical_id) if item is None: raise ValueError(f"Seed example not found: {example_id}") return SeedExampleDetailResponse( @@ -200,7 +205,7 @@ def _preview(members: Sequence[object]) -> tuple[str, bool]: for member in members if member.data_type == "text" and not urlparse(member.value).scheme - and not match(r"^(?:[A-Za-z]:[\\/]|/)", member.value) + and not (ntpath.isabs(member.value) or posixpath.isabs(member.value)) ] if safe_text: value = max(safe_text, key=len) diff --git a/pyrit/memory/alembic/versions/f2a4c6e8b0d2_persist_seed_template_flag.py b/pyrit/memory/alembic/versions/f2a4c6e8b0d2_persist_seed_template_flag.py new file mode 100644 index 0000000000..356e84f7c4 --- /dev/null +++ b/pyrit/memory/alembic/versions/f2a4c6e8b0d2_persist_seed_template_flag.py @@ -0,0 +1,26 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Persist the seed template marker for side-effect-free browsing.""" + +from collections.abc import Sequence + +import sqlalchemy as sa +from alembic import op + +revision: str = "f2a4c6e8b0d2" +down_revision: str | None = "7a9c1e3f5b2d" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + + +def upgrade() -> None: + """Add the nullable persisted template marker.""" + with op.batch_alter_table("SeedPromptEntries") as batch_op: + batch_op.add_column(sa.Column("is_jinja_template", sa.Boolean(), nullable=True)) + + +def downgrade() -> None: + """Remove the persisted template marker.""" + with op.batch_alter_table("SeedPromptEntries") as batch_op: + batch_op.drop_column("is_jinja_template") diff --git a/pyrit/memory/memory_interface.py b/pyrit/memory/memory_interface.py index a010da49d3..c9c454fc28 100644 --- a/pyrit/memory/memory_interface.py +++ b/pyrit/memory/memory_interface.py @@ -581,7 +581,7 @@ def _build_seed_examples(*, rows: Sequence[Any], members: Sequence[SeedEntry]) - added_by=entry.added_by, prompt_metadata=dict(entry.prompt_metadata) if entry.prompt_metadata is not None else None, parameters=list(entry.parameters) if entry.parameters is not None else None, - is_jinja_template=True if entry.parameters else None, + is_jinja_template=entry.is_jinja_template, ) for entry in entries ], @@ -1201,6 +1201,46 @@ def get_seed_example_page( next_cursor=_build_seed_example_next_cursor(rows=page_rows, query=query), ) + def get_seed_example( + self, *, dataset_scope: SeedExampleDatasetScope, example_id: uuid.UUID + ) -> SeedExample | None: + """ + Read one logical seed example by its persisted identity. + + Returns: + SeedExample | None: The complete logical example, or None when it is absent + from the requested dataset scope. + """ + logical_id = _seed_example_logical_id(SeedEntry) + dataset_name = dataset_scope.name if dataset_scope.kind == "named" else None + with closing(self.get_session()) as session: + row = session.execute( + select( + logical_id.label("example_id"), + literal(dataset_name).label("dataset_name"), + ) + .where( + _seed_example_dataset_condition(SeedEntry, dataset_name=dataset_name), + logical_id == example_id, + ) + .group_by(logical_id) + ).first() + if row is None: + return None + + members = list( + session.execute( + select(SeedEntry) + .where( + _seed_example_dataset_condition(SeedEntry, dataset_name=dataset_name), + logical_id == row.example_id, + ) + .order_by(*_seed_example_order_by(logical_id)) + ).scalars() + ) + + return _build_seed_examples(rows=[row], members=members)[0] + def _seed_example_harm_condition(self, column: Any, categories: Sequence[str]) -> Any: """ Build the backend-native any-category predicate for seed harm labels. diff --git a/pyrit/memory/memory_models.py b/pyrit/memory/memory_models.py index 65bd82c20f..2d67fa133a 100644 --- a/pyrit/memory/memory_models.py +++ b/pyrit/memory/memory_models.py @@ -1488,6 +1488,7 @@ class SeedEntry(Base): added_by = mapped_column(String, nullable=False) prompt_metadata: Mapped[dict[str, str | int] | None] = mapped_column(JSON, nullable=True) parameters: Mapped[list[str] | None] = mapped_column(JSON, nullable=True) + is_jinja_template: Mapped[bool | None] = mapped_column(Boolean, nullable=True) prompt_group_id: Mapped[uuid.UUID | None] = mapped_column(CustomUUID, nullable=True) sequence: Mapped[int | None] = mapped_column(INTEGER, nullable=True) role: Mapped[ChatMessageRole | None] = mapped_column(String, nullable=True) @@ -1522,6 +1523,7 @@ def __init__(self, *, entry: Seed) -> None: self.date_added = entry.date_added self.added_by = entry.added_by self.prompt_metadata = self._pack_seed_metadata(entry) + self.is_jinja_template = entry.is_jinja_template self.prompt_group_id = entry.prompt_group_id self.seed_type = seed_type @@ -1626,6 +1628,7 @@ def get_seed(self) -> Seed: names a prompt file that is not present on this machine. """ cleaned_metadata, decoded_schema = self._unpack_seed_metadata(self.prompt_metadata) + domain_template_flag = self._domain_template_flag() if self.seed_type == "objective": return SeedObjective( id=self.id, @@ -1642,6 +1645,7 @@ def get_seed(self) -> Seed: added_by=self.added_by, metadata=cleaned_metadata, prompt_group_id=self.prompt_group_id, + is_jinja_template=domain_template_flag, ) if self.seed_type == "simulated_conversation": # Reconstruct SeedSimulatedConversation from JSON value. Records written before the @@ -1677,6 +1681,7 @@ def get_seed(self) -> Seed: added_by=self.added_by, metadata=cleaned_metadata, prompt_group_id=self.prompt_group_id, + is_jinja_template=domain_template_flag, num_turns=config.get("num_turns", 3), sequence=config.get("sequence", 0), pyrit_version=config.get("pyrit_version"), @@ -1707,11 +1712,21 @@ def get_seed(self) -> Seed: metadata=cleaned_metadata, response_json_schema=decoded_schema, parameters=self.parameters, + is_jinja_template=domain_template_flag, prompt_group_id=self.prompt_group_id, sequence=self.sequence or 0, role=self.role, ) + def _domain_template_flag(self) -> bool: + """ + Normalize an unknown historical template flag to the domain default. + + Returns: + bool: The persisted flag, or the domain's false default for historical NULL values. + """ + return self.is_jinja_template if self.is_jinja_template is not None else False + class AttackResultEntry(Base): """ diff --git a/pyrit/memory/sqlite_memory.py b/pyrit/memory/sqlite_memory.py index 1c9a7512b5..5419927097 100644 --- a/pyrit/memory/sqlite_memory.py +++ b/pyrit/memory/sqlite_memory.py @@ -290,8 +290,15 @@ def _get_seed_harm_category_condition( Any: A SQLAlchemy predicate matching any requested category. """ values = [category.lower() for category in categories] - array = func.json_extract(json_column, literal("$")) - elements = func.json_each(array).table_valued("value") + safe_json = case( + (func.json_valid(json_column) == 1, json_column), + else_=literal("[]"), + ) + safe_array = case( + (func.json_type(safe_json, literal("$")) == "array", safe_json), + else_=literal("[]"), + ) + elements = func.json_each(safe_array).table_valued("value") return exists(select(1).select_from(elements).where(func.lower(elements.c.value).in_(values))) def get_all_table_models(self) -> list[type[Base]]: diff --git a/tests/unit/backend/test_dataset_service_seed_browsing_contract.py b/tests/unit/backend/test_dataset_service_seed_browsing_contract.py index 50a4789ae5..586a9a3f36 100644 --- a/tests/unit/backend/test_dataset_service_seed_browsing_contract.py +++ b/tests/unit/backend/test_dataset_service_seed_browsing_contract.py @@ -5,6 +5,7 @@ from __future__ import annotations +from datetime import UTC, datetime from typing import TYPE_CHECKING, Any from unittest.mock import patch from uuid import uuid4 @@ -231,6 +232,27 @@ async def test_invalid_detail_does_not_generate_group_identity(self, dataset_ser selection_key=SELECTION_KEY, example_id=str(uuid4()) ) + async def test_detail_retrieves_logical_example_beyond_first_list_page( + self, dataset_service: DatasetService, sqlite_instance: MemoryInterface + ): + group_id = uuid4() + target = SeedPrompt( + value="target", + dataset_name=DATASET, + prompt_group_id=group_id, + date_added=datetime(2020, 1, 1, tzinfo=UTC), + ) + await _add(sqlite_instance, target) + await _add( + sqlite_instance, + *(SeedPrompt(value=f"later-{index}", dataset_name=DATASET) for index in range(100)), + ) + + detail = await _service_method(dataset_service, "get_seed_example_async")( + selection_key=SELECTION_KEY, example_id=str(group_id) + ) + assert [member.id for member in _field(detail, "members")] == [target.id] + async def test_browsing_has_no_template_or_generation_side_effects( self, dataset_service: DatasetService, sqlite_instance: MemoryInterface ): @@ -255,3 +277,21 @@ async def test_media_preview_is_type_label_without_bytes_or_path_leak( assert "image" in _field(item, "preview").lower() assert str(tmp_path) not in _field(item, "preview") assert "bytes" not in (item if isinstance(item, dict) else item.model_dump()) + + @pytest.mark.parametrize( + ("value", "expected_preview"), + [ + ("/private/secret.txt", "text seed"), + (r"C:\\private\\secret.txt", "text seed"), + (r"\\\\server\\share\\secret.txt", "text seed"), + ("ordinary safe text", "ordinary safe text"), + ], + ) + async def test_text_preview_rejects_absolute_paths_including_unc( + self, dataset_service: DatasetService, sqlite_instance: MemoryInterface, value: str, expected_preview: str + ): + await _add(sqlite_instance, SeedPrompt(value=value, dataset_name=DATASET, data_type="text")) + response = await _service_method(dataset_service, "list_seed_examples_async")( + selection_key=SELECTION_KEY, limit=10 + ) + assert _field(_field(response, "items")[0], "preview") == expected_preview diff --git a/tests/unit/memory/memory_interface/test_interface_seed_browsing_contract.py b/tests/unit/memory/memory_interface/test_interface_seed_browsing_contract.py index 2d34fefc29..048c2d82de 100644 --- a/tests/unit/memory/memory_interface/test_interface_seed_browsing_contract.py +++ b/tests/unit/memory/memory_interface/test_interface_seed_browsing_contract.py @@ -14,7 +14,7 @@ from uuid import UUID, uuid4 import pytest -from sqlalchemy import event +from sqlalchemy import event, text from pyrit.memory import SeedExampleDatasetScope from pyrit.memory.memory_models import SeedEntry @@ -43,6 +43,17 @@ def _page(memory: MemoryInterface, **kwargs: Any) -> Any: return helper(**kwargs) +def _example(memory: MemoryInterface, **kwargs: Any) -> Any: + """Call the proposed bounded logical-example point lookup helper.""" + helper = getattr(memory, "get_seed_example", None) + assert helper is not None, "RED: MemoryInterface.get_seed_example is not implemented" + dataset_name = kwargs.pop("dataset_name", None) + kwargs["dataset_scope"] = ( + SeedExampleDatasetScope.named(dataset_name) if dataset_name is not None else SeedExampleDatasetScope.unnamed() + ) + return helper(**kwargs) + + async def _add(memory: MemoryInterface, *seeds: SeedPrompt | SeedObjective | SeedSimulatedConversation) -> None: await memory.add_seeds_to_memory_async(seeds=list(seeds), added_by="2748-memory-test") @@ -133,6 +144,82 @@ async def test_template_persisted_value_and_parameters_are_not_rendered(self, sq assert member.parameters == ["name"] assert member.is_jinja_template is True + @pytest.mark.parametrize( + ("is_template", "parameters"), + [(True, []), (False, ["name"]), (False, [])], + ) + async def test_template_flag_is_persisted_independently_of_parameters( + self, sqlite_instance: MemoryInterface, is_template: bool, parameters: list[str] + ): + seed = SeedPrompt( + value="{{ name }}" if is_template else "ordinary", + dataset_name=DATASET, + is_jinja_template=is_template, + parameters=parameters, + ) + await _add(sqlite_instance, seed) + + member = _field(_page(sqlite_instance, dataset_name=DATASET, limit=10), "items")[0].members[0] + assert member.is_jinja_template is is_template + + async def test_historical_null_template_flag_stays_null_in_browsing_projection( + self, sqlite_instance: MemoryInterface + ): + entry = SeedEntry( + entry=SeedPrompt(value="historical", dataset_name=DATASET, parameters=["name"], added_by="legacy-test") + ) + entry.is_jinja_template = None + with sqlite_instance.get_session() as session: + session.add(entry) + session.commit() + + member = _field(_page(sqlite_instance, dataset_name=DATASET, limit=10), "items")[0].members[0] + assert member.parameters == ["name"] + assert member.is_jinja_template is None + + async def test_point_lookup_is_bounded_and_scope_isolated(self, sqlite_instance: MemoryInterface): + target_group = uuid4() + target = SeedPrompt( + value="target", + dataset_name=DATASET, + prompt_group_id=target_group, + date_added=datetime(2020, 1, 1, tzinfo=UTC), + ) + other_dataset_member = SeedPrompt( + value="wrong dataset", + dataset_name="another-dataset", + prompt_group_id=target_group, + ) + unnamed_member = SeedPrompt(value="unnamed", prompt_group_id=target_group) + await _add(sqlite_instance, target, other_dataset_member, unnamed_member) + await _add( + sqlite_instance, + *(SeedPrompt(value=f"later-{index}", dataset_name=DATASET) for index in range(100)), + ) + + result = _example(sqlite_instance, dataset_name=DATASET, example_id=target_group) + assert result is not None + assert {member.id for member in result.members} == {target.id} + + assert _example(sqlite_instance, example_id=target_group) is not None + assert _example(sqlite_instance, dataset_name="another-dataset", example_id=target_group) is not None + assert _example(sqlite_instance, example_id=uuid4()) is None + + async def test_sqlite_harm_json_invalid_and_non_arrays_are_unlabeled(self, sqlite_instance: MemoryInterface): + values = [None, "null", "[]", '"violence"', '{"category":"violence"}', "not-json", '["violence"]'] + seeds = [SeedPrompt(value=f"harm-{index}", dataset_name=DATASET) for index in range(len(values))] + await _add(sqlite_instance, *seeds) + with sqlite_instance.get_session() as session: + for seed, value in zip(seeds, values, strict=True): + session.execute( + text('UPDATE "SeedPromptEntries" SET harm_categories = :value WHERE id = :id'), + {"value": value, "id": str(seed.id)}, + ) + session.commit() + + page = _page(sqlite_instance, dataset_name=DATASET, harm_categories=["violence"], limit=10) + assert [member.id for member in _field(page, "items")[0].members] == [seeds[-1].id] + async def test_logical_identity_uses_group_id_else_seed_id(self, sqlite_instance: MemoryInterface): group_id = uuid4() grouped = SeedPrompt(value="grouped", dataset_name=DATASET, prompt_group_id=group_id) diff --git a/tests/unit/memory/test_memory_models.py b/tests/unit/memory/test_memory_models.py index e4e345ed69..715c88006b 100644 --- a/tests/unit/memory/test_memory_models.py +++ b/tests/unit/memory/test_memory_models.py @@ -552,6 +552,32 @@ def test_roundtrip_seed_objective(self): assert isinstance(recovered, SeedObjective) assert recovered.value == "objective text" + @pytest.mark.parametrize("seed_kind", ["prompt", "objective", "simulated_conversation"]) + def test_get_seed_normalizes_historical_null_template_flag(self, seed_kind: str): + if seed_kind == "prompt": + seed = _make_seed_prompt() + elif seed_kind == "objective": + seed = SeedObjective(value="objective text", dataset_name="ds", added_by="tester") + else: + seed = SeedSimulatedConversation( + adversarial_chat_system_prompt=SeedPrompt(value="adversarial"), + pyrit_version="1.0.0", + ) + + entry = SeedEntry(entry=seed) + entry.is_jinja_template = None + + recovered = entry.get_seed() + assert recovered.is_jinja_template is False + + @pytest.mark.parametrize("is_template", [True, False]) + def test_get_seed_preserves_persisted_template_flag(self, is_template: bool): + entry = SeedEntry(entry=_make_seed_prompt(is_jinja_template=is_template)) + entry.is_jinja_template = is_template + + recovered = entry.get_seed() + assert recovered.is_jinja_template is is_template + def test_seed_prompt_preserves_parameters(self): seed = _make_seed_prompt(parameters=["param1", "param2"]) entry = SeedEntry(entry=seed) diff --git a/tests/unit/memory/test_migration.py b/tests/unit/memory/test_migration.py index 6587bf7b61..48b3ea1f5a 100644 --- a/tests/unit/memory/test_migration.py +++ b/tests/unit/memory/test_migration.py @@ -174,6 +174,32 @@ def test_run_schema_migrations_applies_head_revision(): engine.dispose() +def test_seed_template_flag_migration_lifecycle(): + """The seed template marker is added and removed through the normal Alembic lifecycle.""" + with tempfile.TemporaryDirectory() as temp_dir: + db_path = os.path.join(temp_dir, "seed-template-flag.db") + engine = create_engine(f"sqlite:///{db_path}") + try: + with engine.begin() as connection: + config = _config_for(connection) + command.upgrade(config, "7a9c1e3f5b2d") + assert "is_jinja_template" not in { + column["name"] for column in inspect(connection).get_columns("SeedPromptEntries") + } + + command.upgrade(config, "head") + assert "is_jinja_template" in { + column["name"] for column in inspect(connection).get_columns("SeedPromptEntries") + } + + command.downgrade(config, "7a9c1e3f5b2d") + assert "is_jinja_template" not in { + column["name"] for column in inspect(connection).get_columns("SeedPromptEntries") + } + finally: + engine.dispose() + + def test_scenario_progress_migration_adds_composite_index(): """The migration head contains the parent/timestamp/id keyset index.""" with tempfile.TemporaryDirectory() as temp_dir: From d449524a363187962c4c28cdaf747129b77e4492 Mon Sep 17 00:00:00 2001 From: Jackson Severino da Rocha Date: Sun, 4 Oct 2026 17:42:28 -0300 Subject: [PATCH 6/9] feat: add paginated seed browsing API --- pyrit/backend/models/datasets.py | 3 + pyrit/backend/routes/datasets.py | 68 +++++++++++++++++++- pyrit/backend/services/dataset_service.py | 59 ++++++++++++++++- tests/unit/backend/test_seed_browsing_api.py | 35 +++++++++- 4 files changed, 157 insertions(+), 8 deletions(-) diff --git a/pyrit/backend/models/datasets.py b/pyrit/backend/models/datasets.py index 02f123b899..24a60c3365 100644 --- a/pyrit/backend/models/datasets.py +++ b/pyrit/backend/models/datasets.py @@ -82,6 +82,8 @@ class SeedExampleSummary(BaseModel): name: str | None = None preview: str preview_truncated: bool + is_template: bool | None = None + parameters: list[str] | None = None seed_ids: list[UUID] modalities: list[str] seed_types: list[str] @@ -96,6 +98,7 @@ class SeedExampleListResponse(BaseModel): items: list[SeedExampleSummary] pagination: PaginationInfo + total: int class SeedExampleDetailResponse(BaseModel): diff --git a/pyrit/backend/routes/datasets.py b/pyrit/backend/routes/datasets.py index 39109fbb58..7c2e0c8ed6 100644 --- a/pyrit/backend/routes/datasets.py +++ b/pyrit/backend/routes/datasets.py @@ -8,13 +8,19 @@ discovered from registered ``SeedDatasetProvider`` subclasses. """ -from fastapi import APIRouter +from fastapi import APIRouter, HTTPException, Query, status from pyrit.backend.models.common import ProblemDetail from pyrit.backend.models.datasets import ( DatasetListResponse, + SeedExampleDetailResponse, + SeedExampleListResponse, +) +from pyrit.backend.services.dataset_service import ( + DatasetNotFoundError, + InvalidDatasetSelectionError, + get_dataset_service, ) -from pyrit.backend.services.dataset_service import get_dataset_service router = APIRouter(prefix="/datasets", tags=["datasets"]) @@ -38,3 +44,61 @@ 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( + "/{selection_key}/seeds", + response_model=SeedExampleListResponse, + responses={ + 400: {"model": ProblemDetail, "description": "Invalid cursor or selection"}, + 404: {"model": ProblemDetail, "description": "Dataset not found"}, + 422: {"model": ProblemDetail, "description": "Invalid query parameters"}, + }, +) +async def list_seed_examples( + selection_key: str, + limit: int = Query(20, ge=1, le=100, description="Maximum examples per page"), + cursor: str | None = Query(None, description="Opaque page cursor"), + search: str | None = Query(None, description="Literal text search"), + modality: list[str] | None = Query(None), + harm_category: list[str] | None = Query(None), + seed_type: list[str] | None = Query(None), +) -> SeedExampleListResponse: + """ + List a page of logical seed examples. + + Returns: + SeedExampleListResponse: The selected examples and pagination metadata. + """ + service = get_dataset_service() + try: + return await 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, + ) + except (DatasetNotFoundError, InvalidDatasetSelectionError) as exc: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)) from exc + except ValueError as exc: + raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)) from exc + + +@router.get( + "/{selection_key}/seeds/{example_id}", + response_model=SeedExampleDetailResponse, + responses={ + 404: {"model": ProblemDetail, "description": "Dataset or example not found"}, + 400: {"model": ProblemDetail, "description": "Invalid example identifier"}, + }, +) +async def get_seed_example(selection_key: str, example_id: str) -> SeedExampleDetailResponse: + """Return the complete persisted members of one logical seed example.""" + service = get_dataset_service() + try: + return await service.get_seed_example_async(selection_key=selection_key, example_id=example_id) + except ValueError as exc: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)) from exc diff --git a/pyrit/backend/services/dataset_service.py b/pyrit/backend/services/dataset_service.py index 83d67b5245..c713ccd049 100644 --- a/pyrit/backend/services/dataset_service.py +++ b/pyrit/backend/services/dataset_service.py @@ -31,6 +31,14 @@ logger = logging.getLogger(__name__) +class DatasetNotFoundError(ValueError): + """Raised when a syntactically valid named selection is not loaded in Memory.""" + + +class InvalidDatasetSelectionError(ValueError): + """Raised when a dataset selection key has an invalid format.""" + + class DatasetService: """Service for listing seed datasets.""" @@ -120,7 +128,7 @@ async def list_seed_examples_async( Returns: SeedExampleListResponse: The selected logical examples and pagination metadata. """ - scope = self._selection_scope(selection_key) + scope = await self._resolve_browsing_selection_async(selection_key=selection_key) page = self._memory.get_seed_example_page( dataset_scope=scope, limit=limit, @@ -139,11 +147,12 @@ async def list_seed_examples_async( next_cursor=page.next_cursor, prev_cursor=cursor, ), + total=page.total, ) async def get_seed_example_async(self, *, selection_key: str, example_id: str) -> SeedExampleDetailResponse: """Return one complete logical seed example without materializing seed models.""" - scope = self._selection_scope(selection_key) + scope = await self._resolve_browsing_selection_async(selection_key=selection_key) try: logical_id = UUID(example_id) except ValueError as exc: @@ -164,6 +173,24 @@ async def get_seed_example_async(self, *, selection_key: str, example_id: str) - members=[self._member(member) for member in item.members], ) + async def _resolve_browsing_selection_async(self, *, selection_key: str) -> SeedExampleDatasetScope: + """ + Resolve a browse selection from persisted dataset identities only. + + Returns: + SeedExampleDatasetScope: The validated named or unnamed scope. + + Raises: + DatasetNotFoundError: If a named dataset is not represented in Memory. + InvalidDatasetSelectionError: If the selection key is malformed. + """ + scope = self._selection_scope(selection_key) + if scope.kind == "named": + summaries = self._memory.get_seed_dataset_summaries() + if not any(summary.dataset_name == scope.name for summary in summaries): + raise DatasetNotFoundError(f"Dataset not found: {selection_key}") + return scope + @staticmethod def _selection_scope(selection_key: str) -> SeedExampleDatasetScope: """ @@ -177,7 +204,7 @@ def _selection_scope(selection_key: str) -> SeedExampleDatasetScope: prefix = "dataset:named:" if selection_key.startswith(prefix) and selection_key[len(prefix) :]: return SeedExampleDatasetScope.named(selection_key[len(prefix) :]) - raise ValueError(f"Invalid dataset selection key: {selection_key}") + raise InvalidDatasetSelectionError(f"Invalid dataset selection key: {selection_key}") @classmethod def _summary(cls, item: object) -> SeedExampleSummary: @@ -189,6 +216,8 @@ def _summary(cls, item: object) -> SeedExampleSummary: name=next((member.name for member in members if member.name), None), preview=preview, preview_truncated=truncated, + is_template=cls._template_status(members), + parameters=cls._template_parameters(members), seed_ids=item.seed_ids, # type: ignore[attr-defined] modalities=item.modalities, # type: ignore[attr-defined] seed_types=item.seed_types, # type: ignore[attr-defined] @@ -198,6 +227,30 @@ def _summary(cls, item: object) -> SeedExampleSummary: has_unlabeled_harm=item.has_unlabeled_harm, # type: ignore[attr-defined] ) + @staticmethod + def _template_status(members: Sequence[object]) -> bool | None: + """ + Aggregate nullable template status without treating unknown as false. + + Returns: + bool | None: True if any member is a template, None if the status is unknown, + otherwise False. + """ + statuses = [member.is_jinja_template for member in members] + if any(status is True for status in statuses): + return True + if any(status is None for status in statuses): + return None + return False + + @staticmethod + def _template_parameters(members: Sequence[object]) -> list[str] | None: + """Return persisted parameters for known template members only.""" + for member in members: + if member.is_jinja_template is True: + return member.parameters + return None + @staticmethod def _preview(members: Sequence[object]) -> tuple[str, bool]: safe_text = [ diff --git a/tests/unit/backend/test_seed_browsing_api.py b/tests/unit/backend/test_seed_browsing_api.py index c5a36b1c35..8771353337 100644 --- a/tests/unit/backend/test_seed_browsing_api.py +++ b/tests/unit/backend/test_seed_browsing_api.py @@ -20,6 +20,8 @@ from fastapi.testclient import TestClient from pyrit.backend.main import app +from pyrit.backend.services.dataset_service import get_dataset_service +from pyrit.datasets import SeedDatasetProvider from pyrit.models import SeedObjective, SeedPrompt, SeedSimulatedConversation if TYPE_CHECKING: @@ -34,7 +36,14 @@ @pytest.fixture def client(patch_central_database) -> TestClient: """Use the real SQLite memory fixture behind the API application.""" - return TestClient(app) + # DatasetService is globally cached in production. Each test gets a fresh + # fixture-owned SQLiteMemory, so do not retain a service bound to a previous + # test's in-memory engine after that fixture disposes it. + get_dataset_service.cache_clear() + try: + yield TestClient(app) + finally: + get_dataset_service.cache_clear() async def _add(memory: MemoryInterface, *seeds: SeedPrompt | SeedObjective | SeedSimulatedConversation) -> None: @@ -56,7 +65,7 @@ def _items(response): class TestEmptyAndDatasetSelection: async def test_empty_dataset_is_a_valid_empty_page(self, client, sqlite_instance: MemoryInterface): - response = _list(client) + response = _list(client, UNNAMED_KEY) assert response.status_code == 200 body = response.json() assert body["items"] == [] @@ -250,6 +259,26 @@ async def test_filters_are_and_across_filters_but_match_at_example_level(self, c class TestTextSearchAndSafety: + async def test_browsing_selection_validation_does_not_discover_providers_or_read_files( + self, client, sqlite_instance + ): + seed = SeedPrompt(value="stored", dataset_name=DATASET) + await _add(sqlite_instance, seed) + with ( + patch.object( + SeedDatasetProvider, + "get_all_dataset_names_async", + side_effect=AssertionError("provider metadata discovery"), + ), + patch.object(SeedDatasetProvider, "_parse_metadata_async", side_effect=AssertionError("metadata parse")), + patch("pathlib.Path.read_text", side_effect=AssertionError("provider file read")), + ): + listed = _list(client) + assert listed.status_code == 200 + example_id = listed.json()["items"][0]["example_id"] + detail = _detail(client, example_id) + assert detail.status_code == 200 + async def test_text_search_is_literal_case_insensitive_and_text_only(self, client, sqlite_instance, tmp_path): image_path = tmp_path / "media-path.png" image_path.write_bytes(b"local test image") @@ -461,7 +490,7 @@ async def test_browsing_is_read_only_and_does_not_fetch_provider_or_write(self, ), patch.object(sqlite_instance, "add_seeds_to_memory_async", side_effect=AssertionError("write")) as write, ): - response = _list(client) + response = _list(client, UNNAMED_KEY) assert response.status_code == 200 assert write.call_count == 0 From 343b7a701b5fed1460e5967e31b2a344dc7fe743 Mon Sep 17 00:00:00 2001 From: Richard Lundeen Date: Thu, 8 Oct 2026 12:35:12 -0700 Subject: [PATCH 7/9] Simplify seed browsing API to return domain seeds - Return domain Seed objects from memory; detail members are SeedUnion. - Remove the is_jinja_template column, its migration, and the legacy routes. - Remove the is_template and is_configuration summary labels. - Exclude simulated-conversation JSON from text search. - Replace the seed browsing tests with focused memory, route, and integration tests. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- doc/code/datasets/0_dataset.md | 33 + pyrit/backend/models/datasets.py | 71 +- pyrit/backend/routes/datasets.py | 108 +-- pyrit/backend/services/dataset_service.py | 257 +++---- pyrit/memory/__init__.py | 8 - ...f2a4c6e8b0d2_persist_seed_template_flag.py | 26 - pyrit/memory/azure_sql_memory.py | 26 - pyrit/memory/memory_interface.py | 700 ++++++------------ pyrit/memory/memory_models.py | 19 +- pyrit/memory/sqlite_memory.py | 23 +- ...est_seed_browsing_azure_sql_integration.py | 159 ---- ...est_seed_examples_azure_sql_integration.py | 65 ++ ..._dataset_service_seed_browsing_contract.py | 297 -------- tests/unit/backend/test_seed_browsing_api.py | 510 ------------- .../unit/backend/test_seed_example_routes.py | 172 +++++ .../test_interface_seed_browsing_contract.py | 654 ---------------- .../test_interface_seed_browsing_sql.py | 108 --- .../test_interface_seed_examples.py | 230 ++++++ tests/unit/memory/test_memory_models.py | 26 - tests/unit/memory/test_migration.py | 30 +- 20 files changed, 921 insertions(+), 2601 deletions(-) delete mode 100644 pyrit/memory/alembic/versions/f2a4c6e8b0d2_persist_seed_template_flag.py delete mode 100644 tests/integration/memory/test_seed_browsing_azure_sql_integration.py create mode 100644 tests/integration/memory/test_seed_examples_azure_sql_integration.py delete mode 100644 tests/unit/backend/test_dataset_service_seed_browsing_contract.py delete mode 100644 tests/unit/backend/test_seed_browsing_api.py create mode 100644 tests/unit/backend/test_seed_example_routes.py delete mode 100644 tests/unit/memory/memory_interface/test_interface_seed_browsing_contract.py delete mode 100644 tests/unit/memory/memory_interface/test_interface_seed_browsing_sql.py create mode 100644 tests/unit/memory/memory_interface/test_interface_seed_examples.py diff --git a/doc/code/datasets/0_dataset.md b/doc/code/datasets/0_dataset.md index 008ed96dce..4f4bb26aff 100644 --- a/doc/code/datasets/0_dataset.md +++ b/doc/code/datasets/0_dataset.md @@ -93,3 +93,36 @@ 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 files, render templates, or generate conversations. A simulated-conversation seed +saved by an older PyRIT version stores prompt file paths, so the browser loads those files. +If a seed cannot be read, for example because a file is missing, the browser skips it and +logs a warning. An example with no readable seeds is not shown, but `total` counts it. + +- `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 seed objects (`SeedPrompt`, `SeedObjective`, or `SeedSimulatedConversation`). + +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 example ID. 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, and +other types show a type label. The browser does not render templates or run +simulated conversations. diff --git a/pyrit/backend/models/datasets.py b/pyrit/backend/models/datasets.py index 24a60c3365..e411ea0f62 100644 --- a/pyrit/backend/models/datasets.py +++ b/pyrit/backend/models/datasets.py @@ -6,16 +6,15 @@ 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 datetime import datetime -from typing import Any from uuid import UUID from pydantic import BaseModel, Field from pyrit.backend.models.common import PaginationInfo +from pyrit.models import PromptDataType, SeedType, SeedUnion class DatasetInfo(BaseModel): @@ -49,68 +48,30 @@ class DatasetListResponse(BaseModel): items: list[DatasetInfo] = Field(..., description="List of available datasets") -class SeedExampleMemberView(BaseModel): - """Persisted member of a logical seed example.""" - - id: UUID - prompt_group_id: UUID | None = None - seed_type: str - data_type: str - value: str - value_sha256: str | None = None - role: str | None = None - sequence: int | None = None - name: str | None = None - dataset_name: str | None = None - harm_categories: list[str] | None = None - description: str | None = None - source: str | None = None - authors: list[str] | None = None - groups: list[str] | None = None - date_added: datetime - added_by: str - metadata: dict[str, Any] | None = None - parameters: list[str] | None = None - is_jinja_template: bool | None = None - - class SeedExampleSummary(BaseModel): - """List representation of one complete logical seed example.""" - - example_id: UUID - dataset_name: str | None = None - name: str | None = None - preview: str - preview_truncated: bool - is_template: bool | None = None - parameters: list[str] | None = None - seed_ids: list[UUID] - modalities: list[str] - seed_types: list[str] + """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 + has_unlabeled_harm: bool = Field(..., description="Whether any member has no harm category") class SeedExampleListResponse(BaseModel): - """Paginated logical seed examples.""" + """One page of logical seed examples.""" items: list[SeedExampleSummary] pagination: PaginationInfo - total: int + total: int = Field(..., description="Number of logical examples that match the filters") -class SeedExampleDetailResponse(BaseModel): - """Complete persisted logical seed example.""" +class SeedExampleDetailResponse(SeedExampleSummary): + """One logical seed example with all of its stored seeds.""" - example_id: UUID - dataset_name: str | None = None - seed_ids: list[UUID] - piece_count: int - objective_count: int - modalities: list[str] - seed_types: list[str] - harm_categories: list[str] - has_unlabeled_harm: bool - members: list[SeedExampleMemberView] + members: list[SeedUnion] = Field(..., description="Stored seeds, objectives first, then by sequence") diff --git a/pyrit/backend/routes/datasets.py b/pyrit/backend/routes/datasets.py index 7c2e0c8ed6..7dc342be94 100644 --- a/pyrit/backend/routes/datasets.py +++ b/pyrit/backend/routes/datasets.py @@ -4,26 +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 uuid import UUID + from fastapi import APIRouter, HTTPException, Query, status -from pyrit.backend.models.common import ProblemDetail +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 ( - DatasetNotFoundError, - InvalidDatasetSelectionError, - get_dataset_service, -) +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( "", @@ -47,58 +48,73 @@ async def list_datasets(loaded_only: bool = False) -> DatasetListResponse: # py @router.get( - "/{selection_key}/seeds", + "/seeds", response_model=SeedExampleListResponse, responses={ - 400: {"model": ProblemDetail, "description": "Invalid cursor or selection"}, - 404: {"model": ProblemDetail, "description": "Dataset not found"}, - 422: {"model": ProblemDetail, "description": "Invalid query parameters"}, + 400: {"model": ProblemDetail, "description": "Invalid selection key or cursor"}, }, ) -async def list_seed_examples( - selection_key: str, +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: str | None = Query(None, description="Opaque page cursor"), - search: str | None = Query(None, description="Literal text search"), - modality: list[str] | None = Query(None), - harm_category: list[str] | None = Query(None), - seed_type: list[str] | None = Query(None), + 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 a page of logical seed examples. + 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 selected examples and pagination metadata. + SeedExampleListResponse: The page, its pagination data, and the number of matching examples. """ - service = get_dataset_service() - try: - return await 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, - ) - except (DatasetNotFoundError, InvalidDatasetSelectionError) as exc: - raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)) from exc - except ValueError as exc: - raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)) from exc + 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( - "/{selection_key}/seeds/{example_id}", + "/seeds/{example_id}", response_model=SeedExampleDetailResponse, responses={ - 404: {"model": ProblemDetail, "description": "Dataset or example not found"}, - 400: {"model": ProblemDetail, "description": "Invalid example identifier"}, + 400: {"model": ProblemDetail, "description": "Invalid selection key"}, + 404: {"model": ProblemDetail, "description": "Seed example not found in the dataset"}, }, ) -async def get_seed_example(selection_key: str, example_id: str) -> SeedExampleDetailResponse: - """Return the complete persisted members of one logical seed example.""" - service = get_dataset_service() - try: - return await service.get_seed_example_async(selection_key=selection_key, example_id=example_id) - except ValueError as exc: - raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)) from exc +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 9f884d179d..b7fa79f877 100644 --- a/pyrit/backend/services/dataset_service.py +++ b/pyrit/backend/services/dataset_service.py @@ -2,43 +2,42 @@ # 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 from collections.abc import Sequence from functools import lru_cache -from urllib.parse import urlparse 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, - SeedExampleMemberView, SeedExampleSummary, ) +from pyrit.common.pagination import decode_keyset_cursor, encode_keyset_cursor, fingerprint_filters from pyrit.datasets import SeedDatasetProvider -from pyrit.memory import CentralMemory, SeedExampleDatasetScope -from pyrit.models import SeedDatasetSummary +from pyrit.memory import CentralMemory +from pyrit.models import ( + MEDIA_PATH_DATA_TYPES, + ConversationStats, + PromptDataType, + SeedDatasetSummary, + SeedObjective, + SeedSimulatedConversation, + SeedType, + SeedUnion, +) logger = logging.getLogger(__name__) -class DatasetNotFoundError(ValueError): - """Raised when a syntactically valid named selection is not loaded in Memory.""" - - -class InvalidDatasetSelectionError(ValueError): - """Raised when a dataset selection key has an invalid format.""" - - class DatasetService: """Service for listing seed datasets.""" @@ -118,178 +117,146 @@ async def list_seed_examples_async( limit: int = 20, cursor: str | None = None, search: str | None = None, - data_types: Sequence[str] | None = None, + data_types: Sequence[PromptDataType] | None = None, harm_categories: Sequence[str] | None = None, - seed_types: Sequence[str] | None = None, + seed_types: Sequence[SeedType] | None = None, ) -> SeedExampleListResponse: """ - List logical seed examples using Memory's database-backed page query. + 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 selected logical examples and pagination metadata. + SeedExampleListResponse: The page, its pagination data, and the number of matching examples. + + Raises: + ValueError: If the selection key or the cursor is not valid. """ - scope = await self._resolve_browsing_selection_async(selection_key=selection_key) - page = self._memory.get_seed_example_page( - dataset_scope=scope, + 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, - cursor=cursor, + after=after, data_types=data_types, harm_categories=harm_categories, seed_types=seed_types, value_search=search, ) - items = [self._summary(item) for item in page.items] + next_cursor = ( + encode_keyset_cursor( + timestamp=next_after.timestamp, identifier=next_after.identifier, fingerprint=fingerprint + ) + if next_after + else None + ) return SeedExampleListResponse( - items=items, + items=[self._summarize(example_id=example_id, seeds=seeds) for example_id, seeds in examples.items()], pagination=PaginationInfo( - limit=limit, - has_more=page.next_cursor is not None, - next_cursor=page.next_cursor, - prev_cursor=cursor, + limit=limit, has_more=next_after is not None, next_cursor=next_cursor, prev_cursor=cursor ), - total=page.total, - ) - - async def get_seed_example_async(self, *, selection_key: str, example_id: str) -> SeedExampleDetailResponse: - """Return one complete logical seed example without materializing seed models.""" - scope = await self._resolve_browsing_selection_async(selection_key=selection_key) - try: - logical_id = UUID(example_id) - except ValueError as exc: - raise ValueError(f"Seed example not found: {example_id}") from exc - item = self._memory.get_seed_example(dataset_scope=scope, example_id=logical_id) - if item is None: - raise ValueError(f"Seed example not found: {example_id}") - return SeedExampleDetailResponse( - example_id=item.example_id, - dataset_name=item.dataset_name, - seed_ids=item.seed_ids, - piece_count=item.piece_count, - objective_count=item.objective_count, - modalities=item.modalities, - seed_types=item.seed_types, - harm_categories=item.harm_categories, - has_unlabeled_harm=item.has_unlabeled_harm, - members=[self._member(member) for member in item.members], + total=total, ) - async def _resolve_browsing_selection_async(self, *, selection_key: str) -> SeedExampleDatasetScope: + async def get_seed_example_async(self, *, selection_key: str, example_id: UUID) -> SeedExampleDetailResponse | None: """ - Resolve a browse selection from persisted dataset identities only. + 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: - SeedExampleDatasetScope: The validated named or unnamed scope. + SeedExampleDetailResponse | None: The example, or None if the dataset does not contain it. Raises: - DatasetNotFoundError: If a named dataset is not represented in Memory. - InvalidDatasetSelectionError: If the selection key is malformed. + ValueError: If the selection key is not valid. """ - scope = self._selection_scope(selection_key) - if scope.kind == "named": - summaries = self._memory.get_seed_dataset_summaries() - if not any(summary.dataset_name == scope.name for summary in summaries): - raise DatasetNotFoundError(f"Dataset not found: {selection_key}") - return scope + 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 _selection_scope(selection_key: str) -> SeedExampleDatasetScope: + def _parse_selection_key(selection_key: str) -> str | None: """ - Resolve the stable dataset selection namespace. + Get the dataset name of a selection key. Returns: - SeedExampleDatasetScope: The named or unnamed memory query scope. + 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 SeedExampleDatasetScope.unnamed() - prefix = "dataset:named:" - if selection_key.startswith(prefix) and selection_key[len(prefix) :]: - return SeedExampleDatasetScope.named(selection_key[len(prefix) :]) - raise InvalidDatasetSelectionError(f"Invalid dataset selection key: {selection_key}") + 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 _summary(cls, item: object) -> SeedExampleSummary: - members = item.members # type: ignore[attr-defined] - preview, truncated = cls._preview(members) + def _summarize(cls, *, example_id: UUID, seeds: Sequence[SeedUnion]) -> SeedExampleSummary: + preview, truncated = cls._preview(seeds) return SeedExampleSummary( - example_id=item.example_id, # type: ignore[attr-defined] - dataset_name=item.dataset_name, # type: ignore[attr-defined] - name=next((member.name for member in members if member.name), None), + example_id=example_id, + name=next((seed.name for seed in seeds if seed.name), None), preview=preview, preview_truncated=truncated, - is_template=cls._template_status(members), - parameters=cls._template_parameters(members), - seed_ids=item.seed_ids, # type: ignore[attr-defined] - modalities=item.modalities, # type: ignore[attr-defined] - seed_types=item.seed_types, # type: ignore[attr-defined] - piece_count=item.piece_count, # type: ignore[attr-defined] - objective_count=item.objective_count, # type: ignore[attr-defined] - harm_categories=item.harm_categories, # type: ignore[attr-defined] - has_unlabeled_harm=item.has_unlabeled_harm, # type: ignore[attr-defined] + modalities=sorted({seed.data_type for seed in seeds if seed.data_type}), + seed_types=sorted({seed.seed_type for seed in seeds}), + piece_count=len(seeds), + objective_count=sum(isinstance(seed, SeedObjective) 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 _template_status(members: Sequence[object]) -> bool | None: + def _preview(seeds: Sequence[SeedUnion]) -> tuple[str, bool]: """ - Aggregate nullable template status without treating unknown as false. + 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, so the preview does not expose paths or URL credentials. Returns: - bool | None: True if any member is a template, None if the status is unknown, - otherwise False. + tuple[str, bool]: The preview and whether the text was shortened. """ - statuses = [member.is_jinja_template for member in members] - if any(status is True for status in statuses): - return True - if any(status is None for status in statuses): - return None - return False - - @staticmethod - def _template_parameters(members: Sequence[object]) -> list[str] | None: - """Return persisted parameters for known template members only.""" - for member in members: - if member.is_jinja_template is True: - return member.parameters - return None - - @staticmethod - def _preview(members: Sequence[object]) -> tuple[str, bool]: - safe_text = [ - member.value - for member in members - if member.data_type == "text" - and not urlparse(member.value).scheme - and not (ntpath.isabs(member.value) or posixpath.isabs(member.value)) - ] - if safe_text: - value = max(safe_text, key=len) - return (value[:100] + "...", len(value) > 100) if len(value) > 100 else (value, False) - data_type = members[0].data_type if members else "seed" - return f"{data_type} seed", False - - @staticmethod - def _member(member: object) -> SeedExampleMemberView: - return SeedExampleMemberView( - id=member.id, - prompt_group_id=member.prompt_group_id, - seed_type=member.seed_type, - data_type=member.data_type, - value=member.value, - value_sha256=member.value_sha256, - role=member.role, - sequence=member.sequence, - name=member.name, - dataset_name=member.dataset_name, - harm_categories=member.harm_categories, - description=member.description, - source=member.source, - authors=member.authors, - groups=member.groups, - date_added=member.date_added, - added_by=member.added_by, - metadata=member.metadata, - parameters=member.parameters, - is_jinja_template=member.is_jinja_template, - ) + shown = [seed for seed in seeds if not isinstance(seed, SeedSimulatedConversation)] + 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 + 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/__init__.py b/pyrit/memory/__init__.py index 3376a877fd..42aae06309 100644 --- a/pyrit/memory/__init__.py +++ b/pyrit/memory/__init__.py @@ -23,10 +23,6 @@ ScenarioHistoryKeysetCursor, ScenarioHistoryRunRecord, ScenarioRunStateRecord, - SeedExample, - SeedExampleDatasetScope, - SeedExampleMember, - SeedExamplePage, ) from pyrit.memory.memory_models import AttackResultEntry, EmbeddingDataEntry, PromptMemoryEntry, SeedEntry from pyrit.memory.sqlite_memory import SQLiteMemory @@ -65,10 +61,6 @@ "ErrorDataTypeSerializer": "pyrit.memory.storage", "ImagePathDataTypeSerializer": "pyrit.memory.storage", "MemoryInterface": "pyrit.memory.memory_interface", - "SeedExample": "pyrit.memory.memory_interface", - "SeedExampleDatasetScope": "pyrit.memory.memory_interface", - "SeedExampleMember": "pyrit.memory.memory_interface", - "SeedExamplePage": "pyrit.memory.memory_interface", "MemoryEmbedding": "pyrit.memory.memory_embedding", "ScenarioHistoryKeysetCursor": "pyrit.memory.memory_interface", "ScenarioHistoryRunRecord": "pyrit.memory.memory_interface", diff --git a/pyrit/memory/alembic/versions/f2a4c6e8b0d2_persist_seed_template_flag.py b/pyrit/memory/alembic/versions/f2a4c6e8b0d2_persist_seed_template_flag.py deleted file mode 100644 index 356e84f7c4..0000000000 --- a/pyrit/memory/alembic/versions/f2a4c6e8b0d2_persist_seed_template_flag.py +++ /dev/null @@ -1,26 +0,0 @@ -# Copyright (c) Microsoft Corporation. -# Licensed under the MIT license. - -"""Persist the seed template marker for side-effect-free browsing.""" - -from collections.abc import Sequence - -import sqlalchemy as sa -from alembic import op - -revision: str = "f2a4c6e8b0d2" -down_revision: str | None = "7a9c1e3f5b2d" -branch_labels: str | Sequence[str] | None = None -depends_on: str | Sequence[str] | None = None - - -def upgrade() -> None: - """Add the nullable persisted template marker.""" - with op.batch_alter_table("SeedPromptEntries") as batch_op: - batch_op.add_column(sa.Column("is_jinja_template", sa.Boolean(), nullable=True)) - - -def downgrade() -> None: - """Remove the persisted template marker.""" - with op.batch_alter_table("SeedPromptEntries") as batch_op: - batch_op.drop_column("is_jinja_template") diff --git a/pyrit/memory/azure_sql_memory.py b/pyrit/memory/azure_sql_memory.py index 96736d8e2f..c5afd08b9c 100644 --- a/pyrit/memory/azure_sql_memory.py +++ b/pyrit/memory/azure_sql_memory.py @@ -15,14 +15,11 @@ Unicode, and_, bindparam, - case, create_engine, event, exists, func, - literal, literal_column, - select, text, ) from sqlalchemy.engine import make_url @@ -530,29 +527,6 @@ def _get_condition_json_array_match( combined = joiner.join(conditions) return text(f"""ISJSON("{table_name}".{column_name}) = 1 AND ({combined})""").bindparams(**bindparams_dict) - def _get_seed_harm_category_condition( - self, *, json_column: InstrumentedAttribute[Any], categories: Sequence[str] - ) -> Any: - """ - Build an aliased-column-safe Azure SQL harm-category membership predicate. - - Returns: - Any: A SQLAlchemy predicate matching any requested category. - """ - values = [category.lower() for category in categories] - safe_array = case( - ( - and_( - func.ISJSON(json_column) == literal(1), - func.LEFT(func.LTRIM(json_column), literal(1)) == literal("["), - ), - json_column, - ), - else_=literal("[]"), - ) - elements = func.OPENJSON(safe_array).table_valued("value") - return exists(select(1).select_from(elements).where(func.lower(elements.c.value).in_(values))) - def _get_attack_result_label_condition(self, *, labels: dict[str, str | Sequence[str]]) -> Any: """ Azure SQL implementation for filtering AttackResults by labels. diff --git a/pyrit/memory/memory_interface.py b/pyrit/memory/memory_interface.py index 4c53d03034..92057d8368 100644 --- a/pyrit/memory/memory_interface.py +++ b/pyrit/memory/memory_interface.py @@ -20,17 +20,17 @@ from typing import TYPE_CHECKING, Any, ClassVar, Literal, NamedTuple, ParamSpec, TypeVar, cast from urllib.parse import urlparse -from sqlalchemy import MetaData, String, and_, case, cast, exists, false, func, literal, not_, or_, select, update +from sqlalchemy import MetaData, and_, case, exists, false, func, literal, not_, or_, select, update from sqlalchemy.engine.base import Engine from sqlalchemy.exc import IntegrityError, SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession -from sqlalchemy.orm import aliased, joinedload +from sqlalchemy.orm import joinedload from sqlalchemy.orm.attributes import InstrumentedAttribute, flag_modified from sqlalchemy.orm.session import Session 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 decode_keyset_cursor, encode_keyset_cursor, fingerprint_filters +from pyrit.common.pagination import DecodedKeysetCursor if TYPE_CHECKING: from pyrit.memory.memory_embedding import MemoryEmbedding @@ -88,6 +88,7 @@ MessagePiece, MessageScorable, Observation, + PromptDataType, RetryEvent, ScenarioAttackResultDelta, ScenarioIdentifier, @@ -106,6 +107,7 @@ SeedOrigin, SeedPrompt, SeedType, + SeedUnion, TargetIdentifier, group_conversation_message_pieces_by_sequence, sort_message_pieces, @@ -162,475 +164,6 @@ class _PreparedScorableContent: value_sha256: str -@dataclass(frozen=True, slots=True, kw_only=True) -class SeedExample: - """One logical seed example and its persisted members.""" - - example_id: uuid.UUID - dataset_name: str | None - seed_ids: list[uuid.UUID] - members: list[SeedExampleMember] - piece_count: int - objective_count: int - modalities: list[str] - seed_types: list[str] - harm_categories: list[str] - has_unlabeled_harm: bool - - def __getitem__(self, key: str) -> Any: - """ - Allow response-style field access alongside typed attributes. - - Returns: - Any: The requested field value. - """ - return getattr(self, key) - - -@dataclass(frozen=True, slots=True, kw_only=True) -class SeedExamplePage: - """A bounded page of logical seed examples.""" - - items: list[SeedExample] - total: int - next_cursor: str | None - - -@dataclass(frozen=True, slots=True) -class SeedExampleMember: - """Side-effect-free persisted seed data for browsing.""" - - id: uuid.UUID - prompt_group_id: uuid.UUID | None - seed_type: str - data_type: str - value: str - value_sha256: str | None - role: str | None - sequence: int | None - name: str | None - dataset_name: str | None - harm_categories: list[str] | None - description: str | None - source: str | None - authors: list[str] | None - groups: list[str] | None - date_added: datetime - added_by: str - prompt_metadata: dict[str, Any] | None - parameters: list[str] | None - is_jinja_template: bool | None - - @property - def metadata(self) -> dict[str, Any] | None: - """Persisted metadata under the Seed model's public name.""" - return self.prompt_metadata - - -@dataclass(frozen=True, slots=True) -class SeedExampleDatasetScope: - """Explicit dataset namespace for logical seed-example browsing.""" - - kind: Literal["named", "unnamed"] - name: str | None = None - - @classmethod - def named(cls, name: str) -> SeedExampleDatasetScope: - """ - Create a named dataset scope. - - Args: - name: The non-empty persisted dataset name. - - Returns: - The named dataset scope. - - Raises: - ValueError: If the name is empty. - """ - if not name: - raise ValueError("A named dataset scope requires a non-empty name") - return cls(kind="named", name=name) - - @classmethod - def unnamed(cls) -> SeedExampleDatasetScope: - """ - Create the combined NULL/empty dataset scope. - - Returns: - The unnamed dataset scope. - """ - return cls(kind="unnamed") - - def __post_init__(self) -> None: - """ - Validate the scope's name and kind combination. - - Raises: - ValueError: If the scope kind and name do not agree. - """ - if self.kind not in {"named", "unnamed"}: - raise ValueError(f"Unsupported dataset scope kind: {self.kind}") - if self.kind == "named" and not self.name: - raise ValueError("A named dataset scope requires a non-empty name") - if self.kind == "unnamed" and self.name is not None: - raise ValueError("An unnamed dataset scope cannot have a name") - - -@dataclass(frozen=True, slots=True, kw_only=True) -class _SeedExampleQuery: - """Immutable filters and pagination state for a logical seed example query.""" - - dataset_scope: SeedExampleDatasetScope - limit: int - cursor: Any - fingerprint: str - data_types: tuple[str, ...] - harm_categories: tuple[str, ...] - seed_types: tuple[str, ...] - value_search: str | None - - -def _seed_example_logical_id(member: Any) -> Any: - """Return the logical example key for a seed entry or alias.""" - return func.coalesce(member.prompt_group_id, member.id) - - -def _seed_example_order_key(logical_id: Any) -> Any: - """Return the canonical textual UUID key shared by ordering and seeking.""" - return func.lower(cast(logical_id, String(36))) - - -def _seed_example_dataset_condition(member: Any, *, dataset_name: str | None) -> Any: - """ - Build the dataset scope condition used by seed example queries. - - Returns: - Any: The SQLAlchemy condition for the requested dataset scope. - """ - if dataset_name is None: - return or_(member.dataset_name.is_(None), member.dataset_name == "") - return member.dataset_name == dataset_name - - -def _seed_example_filter_fingerprint( - *, - dataset_name: str | None, - data_types: tuple[str, ...], - harm_categories: tuple[str, ...], - seed_types: tuple[str, ...], - value_search: str | None, -) -> str: - """Return the stable identity of the effective seed example filters.""" - return fingerprint_filters( - filters={ - "dataset_name": dataset_name, - "data_types": data_types, - "harm_categories": harm_categories, - "seed_types": seed_types, - "value_search": value_search or "", - }, - length=64, - ) - - -def _build_seed_example_query( - *, - dataset_scope: SeedExampleDatasetScope, - limit: int, - cursor: str | None, - data_types: Sequence[str] | None, - harm_categories: Sequence[str] | None, - seed_types: Sequence[str] | None, - value_search: str | None, -) -> _SeedExampleQuery: - """ - Normalize seed example filters and decode a filter-bound cursor. - - Returns: - _SeedExampleQuery: The immutable effective query state. - - Raises: - ValueError: If ``limit`` is invalid or the cursor is malformed or mismatched. - """ - if not 1 <= limit <= 100: - raise ValueError("limit must be between 1 and 100") - - normalized_types = tuple(sorted(set(data_types or ()))) - normalized_harms = tuple(sorted(set(harm_categories or ()))) - normalized_seed_types = tuple(sorted(set(seed_types or ()))) - if len(normalized_types) + len(normalized_harms) + len(normalized_seed_types) > 100: - raise ValueError("Too many seed example filter values") - fingerprint = _seed_example_filter_fingerprint( - dataset_name=dataset_scope.name if dataset_scope.kind == "named" else None, - data_types=normalized_types, - harm_categories=normalized_harms, - seed_types=normalized_seed_types, - value_search=value_search, - ) - decoded_cursor = decode_keyset_cursor(cursor=cursor, fingerprint=fingerprint) - if cursor is not None and decoded_cursor is None: - raise ValueError("Invalid or filter-mismatched seed example cursor") - return _SeedExampleQuery( - dataset_scope=dataset_scope, - limit=limit, - cursor=decoded_cursor, - fingerprint=fingerprint, - data_types=normalized_types, - harm_categories=normalized_harms, - seed_types=normalized_seed_types, - value_search=value_search, - ) - - -def _seed_example_member_predicate( - member: Any, - *, - query: _SeedExampleQuery, - filter_name: str, - harm_condition_builder: Callable[[Any, Sequence[str]], Any], -) -> Any: - """ - Build one member-level predicate for a logical example filter. - - Returns: - Any: The SQLAlchemy condition for the requested member filter. - """ - predicates: list[Any] = [] - if filter_name == "data_types" and query.data_types: - predicates.append(member.data_type.in_(query.data_types)) - if filter_name == "seed_types" and query.seed_types: - predicates.append(member.seed_type.in_(query.seed_types)) - if filter_name == "harm_categories" and query.harm_categories: - predicates.append(harm_condition_builder(member.harm_categories, query.harm_categories)) - if filter_name == "value_search" and query.value_search: - escaped = query.value_search.lower().replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") - predicates.extend([member.data_type == "text", func.lower(member.value).like(f"%{escaped}%", escape="\\")]) - return and_(*predicates) if predicates else literal(True) - - -def _build_seed_example_filter_conditions( - *, - query: _SeedExampleQuery, - logical_id: Any, - harm_condition_builder: Callable[[Any, Sequence[str]], Any], -) -> list[Any]: - """ - Build the dataset and independent member-existence filters. - - Returns: - list[Any]: SQLAlchemy conditions for the grouped logical-example query. - """ - dataset_name = query.dataset_scope.name if query.dataset_scope.kind == "named" else None - conditions: list[Any] = [_seed_example_dataset_condition(SeedEntry, dataset_name=dataset_name)] - filter_values = ( - ("data_types", query.data_types), - ("harm_categories", query.harm_categories), - ("seed_types", query.seed_types), - ("value_search", (query.value_search,) if query.value_search else ()), - ) - for filter_name, values in filter_values: - if not values: - continue - member = aliased(SeedEntry) - conditions.append( - exists( - select(1).where( - _seed_example_dataset_condition(member, dataset_name=dataset_name), - _seed_example_logical_id(member) == logical_id, - _seed_example_member_predicate( - member, - query=query, - filter_name=filter_name, - harm_condition_builder=harm_condition_builder, - ), - ) - ) - ) - return conditions - - -def _seed_example_order_by(logical_id: Any) -> list[Any]: - """Return deterministic member ordering within each logical example.""" - return [ - logical_id, - case( - (SeedEntry.seed_type == "objective", 0), - (SeedEntry.seed_type == "simulated_conversation", 1), - else_=2, - ), - case((SeedEntry.sequence.is_(None), 1), else_=0), - SeedEntry.sequence, - SeedEntry.id, - ] - - -def _seed_example_keyset_seek_condition(*, grouped_subquery: Any, cursor: Any) -> Any: - """ - Build the seek predicate for first-added descending, ID descending order. - - Returns: - Any: The SQLAlchemy condition selecting rows after the cursor. - """ - return or_( - grouped_subquery.c.first_added < cursor.timestamp, - and_( - grouped_subquery.c.first_added == cursor.timestamp, - grouped_subquery.c.example_id_key < cursor.identifier, - ), - ) - - -def _query_seed_example_page( - *, - session: Session, - query: _SeedExampleQuery, - harm_condition_builder: Callable[[Any, Sequence[str]], Any], -) -> tuple[list[Any], int]: - """ - Select one logical seed example page and its total logical count. - - Returns: - tuple[list[Any], int]: Over-fetched page rows and the logical total. - """ - logical_id = _seed_example_logical_id(SeedEntry) - logical_id_key = _seed_example_order_key(logical_id) - grouped_subquery = ( - select( - logical_id.label("example_id"), - logical_id_key.label("example_id_key"), - func.min(SeedEntry.date_added).label("first_added"), - ) - .where( - and_( - *_build_seed_example_filter_conditions( - query=query, - logical_id=logical_id, - harm_condition_builder=harm_condition_builder, - ) - ) - ) - .group_by(logical_id, logical_id_key) - .subquery() - ) - group_query = select( - grouped_subquery.c.example_id, - grouped_subquery.c.first_added, - literal(query.dataset_scope.name if query.dataset_scope.kind == "named" else None).label("dataset_name"), - ) - if query.cursor is not None: - group_query = group_query.where( - _seed_example_keyset_seek_condition(grouped_subquery=grouped_subquery, cursor=query.cursor) - ) - page_rows = list( - session.execute( - group_query.order_by(grouped_subquery.c.first_added.desc(), grouped_subquery.c.example_id_key.desc()).limit( - query.limit + 1 - ) - ).all() - ) - total = int(session.execute(select(func.count()).select_from(grouped_subquery)).scalar_one()) - return page_rows, total - - -def _query_seed_example_members( - *, session: Session, query: _SeedExampleQuery, example_ids: Sequence[uuid.UUID] -) -> list[SeedEntry]: - """ - Load all members for the selected logical examples in one query. - - Returns: - list[SeedEntry]: Persisted members ordered within each logical example. - """ - logical_id = _seed_example_logical_id(SeedEntry) - return list( - session.execute( - select(SeedEntry) - .where( - _seed_example_dataset_condition( - SeedEntry, - dataset_name=query.dataset_scope.name if query.dataset_scope.kind == "named" else None, - ), - logical_id.in_(example_ids), - ) - .order_by(*_seed_example_order_by(logical_id)) - ).scalars() - ) - - -def _build_seed_examples(*, rows: Sequence[Any], members: Sequence[SeedEntry]) -> list[SeedExample]: - """ - Materialize logical seed examples from selected rows and persisted members. - - Returns: - list[SeedExample]: The logical examples represented by the selected rows. - """ - selected_ids = [row.example_id for row in rows] - members_by_group: dict[uuid.UUID, list[SeedEntry]] = {example_id: [] for example_id in selected_ids} - for entry in members: - members_by_group.setdefault(entry.prompt_group_id or entry.id, []).append(entry) - - items: list[SeedExample] = [] - for row in rows: - entries = members_by_group[row.example_id] - items.append( - SeedExample( - example_id=row.example_id, - dataset_name=row.dataset_name, - seed_ids=[entry.id for entry in entries], - members=[ - SeedExampleMember( - id=entry.id, - prompt_group_id=entry.prompt_group_id, - seed_type=entry.seed_type, - data_type=entry.data_type, - value=entry.value, - value_sha256=entry.value_sha256, - role=entry.role, - sequence=entry.sequence, - name=entry.name, - dataset_name=entry.dataset_name, - harm_categories=list(entry.harm_categories) if entry.harm_categories is not None else None, - description=entry.description, - source=entry.source, - authors=list(entry.authors) if entry.authors is not None else None, - groups=list(entry.groups) if entry.groups is not None else None, - date_added=entry.date_added, - added_by=entry.added_by, - prompt_metadata=dict(entry.prompt_metadata) if entry.prompt_metadata is not None else None, - parameters=list(entry.parameters) if entry.parameters is not None else None, - is_jinja_template=entry.is_jinja_template, - ) - for entry in entries - ], - piece_count=len(entries), - objective_count=sum(entry.seed_type == "objective" for entry in entries), - modalities=sorted({entry.data_type for entry in entries}), - seed_types=sorted({entry.seed_type for entry in entries}), - harm_categories=sorted({category for entry in entries for category in (entry.harm_categories or [])}), - has_unlabeled_harm=any(not entry.harm_categories for entry in entries), - ) - ) - return items - - -def _build_seed_example_next_cursor(*, rows: Sequence[Any], query: _SeedExampleQuery) -> str | None: - """ - Build the continuation cursor when the page was over-fetched. - - Returns: - str | None: The filter-bound continuation cursor, if another page exists. - """ - if len(rows) <= query.limit: - return None - last = rows[query.limit - 1] - return encode_keyset_cursor( - timestamp=last.first_added, - identifier=str(last.example_id), - fingerprint=query.fingerprint, - ) - - class AttackResultKeysetCursor(NamedTuple): """ Keyset (seek) anchor identifying the last attack result on a page. @@ -1376,7 +909,6 @@ def _execute_add_conversation_to_memory(self, *, conversation: Conversation) -> """ self._insert_conversation(conversation=conversation) - def _execute_add_message_pieces_to_memory(self, *, message_pieces: Sequence[MessagePiece]) -> None: """ Insert a list of message pieces into the memory storage. @@ -4771,6 +4303,162 @@ 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[SeedUnion]]: + """ + Read the stored seeds of the given logical examples. + + A seed that cannot be rebuilt is skipped with a warning, so one bad row does not stop the read. + For example, a simulated conversation saved with prompt file paths fails when a file is missing. + + Returns: + dict[uuid.UUID, list[SeedUnion]]: The seeds of each example that has readable seeds, 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, SeedEntry.id) + ).all() + seeds: dict[uuid.UUID, list[SeedUnion]] = {example_id: [] for example_id in example_ids} + for entry in entries: + try: + seeds[entry.prompt_group_id or entry.id].append(entry.get_seed()) + except ValueError as e: + logger.warning(f"Skipping stored seed {entry.id} because it cannot be read: {e}") + 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[SeedUnion]], int, DecodedKeysetCursor | None]: + """ + Read one keyset page of complete logical seed examples. + + Returns: + tuple[dict[uuid.UUID, list[SeedUnion]], 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) + 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"), func.min(SeedEntry.date_added).label("first_added")) + .where(scope, *filters) + .group_by(logical_id) + .subquery() + ) + page = select(grouped.c.example_id, grouped.c.first_added) + if after is not None: + anchor_id = 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 < anchor_id), + ) + ) + page = page.order_by(grouped.c.first_added.desc(), grouped.c.example_id.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[SeedUnion]: + """ + Read one complete logical seed example. + + Returns: + list[SeedUnion]: The seeds of the example. The list is empty if the dataset does not contain it. + """ + 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. @@ -8895,6 +8583,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[SeedUnion]], 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. Values inside one filter use OR, different filters use AND, and any seed + of an example can satisfy a filter. This method does not render templates or load media. + A simulated conversation saved with prompt file paths loads those files. If it cannot be + read, it is skipped with a warning. + + 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[SeedUnion]], int, DecodedKeysetCursor | None]: The seeds 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[SeedUnion]: + """ + 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[SeedUnion]: The seeds of the example, 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 3fae1d9754..27f3755720 100644 --- a/pyrit/memory/memory_models.py +++ b/pyrit/memory/memory_models.py @@ -1508,7 +1508,6 @@ class SeedEntry(Base): added_by = mapped_column(String, nullable=False) prompt_metadata: Mapped[dict[str, str | int] | None] = mapped_column(JSON, nullable=True) parameters: Mapped[list[str] | None] = mapped_column(JSON, nullable=True) - is_jinja_template: Mapped[bool | None] = mapped_column(Boolean, nullable=True) prompt_group_id: Mapped[uuid.UUID | None] = mapped_column(CustomUUID, nullable=True) sequence: Mapped[int | None] = mapped_column(INTEGER, nullable=True) role: Mapped[ChatMessageRole | None] = mapped_column(String, nullable=True) @@ -1547,7 +1546,6 @@ def __init__(self, *, entry: Seed) -> None: self.date_added = entry.date_added self.added_by = entry.added_by self.prompt_metadata = self._pack_seed_metadata(entry) - self.is_jinja_template = entry.is_jinja_template self.prompt_group_id = entry.prompt_group_id self.seed_type = seed_type self.origin = entry.origin.value @@ -1646,12 +1644,12 @@ def _unpack_seed_metadata( decoded = None return cleaned, decoded - def get_seed(self) -> Seed: + 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, @@ -1661,7 +1659,6 @@ def get_seed(self) -> Seed: if self.seed_type != "objective" and self.conditions not in (None, []): raise ValueError("Only objective seeds can have persisted conditions.") cleaned_metadata, decoded_schema = self._unpack_seed_metadata(self.prompt_metadata) - domain_template_flag = self._domain_template_flag() if self.seed_type == "objective": return SeedObjective( id=self.id, @@ -1679,7 +1676,6 @@ def get_seed(self) -> Seed: added_by=self.added_by, metadata=cleaned_metadata, prompt_group_id=self.prompt_group_id, - is_jinja_template=domain_template_flag, conditions=self.conditions if self.conditions is not None else (), ) if self.seed_type == "simulated_conversation": @@ -1717,7 +1713,6 @@ def get_seed(self) -> Seed: added_by=self.added_by, metadata=cleaned_metadata, prompt_group_id=self.prompt_group_id, - is_jinja_template=domain_template_flag, num_turns=config.get("num_turns", 3), sequence=config.get("sequence", 0), pyrit_version=config.get("pyrit_version"), @@ -1749,21 +1744,11 @@ def get_seed(self) -> Seed: metadata=cleaned_metadata, response_json_schema=decoded_schema, parameters=self.parameters, - is_jinja_template=domain_template_flag, prompt_group_id=self.prompt_group_id, sequence=self.sequence or 0, role=self.role, ) - def _domain_template_flag(self) -> bool: - """ - Normalize an unknown historical template flag to the domain default. - - Returns: - bool: The persisted flag, or the domain's false default for historical NULL values. - """ - return self.is_jinja_template if self.is_jinja_template is not None else False - class AttackResultEntry(Base): """ diff --git a/pyrit/memory/sqlite_memory.py b/pyrit/memory/sqlite_memory.py index 7aaf83ead6..3f17bc9c62 100644 --- a/pyrit/memory/sqlite_memory.py +++ b/pyrit/memory/sqlite_memory.py @@ -16,7 +16,7 @@ from types import TracebackType from typing import TYPE_CHECKING, Any, Literal -from sqlalchemy import and_, case, create_engine, event, exists, func, literal, or_, select, text +from sqlalchemy import and_, case, create_engine, event, exists, func, or_, select, text from sqlalchemy.engine import AdaptedConnection, ExceptionContext from sqlalchemy.engine.base import Engine from sqlalchemy.exc import SQLAlchemyError @@ -491,27 +491,6 @@ def _get_condition_json_array_match( combined = joiner.join(conditions) return text(f"({combined})").bindparams(**bindparams_dict) - def _get_seed_harm_category_condition( - self, *, json_column: InstrumentedAttribute[Any], categories: Sequence[str] - ) -> Any: - """ - Build an aliased-column-safe SQLite harm-category membership predicate. - - Returns: - Any: A SQLAlchemy predicate matching any requested category. - """ - values = [category.lower() for category in categories] - safe_json = case( - (func.json_valid(json_column) == 1, json_column), - else_=literal("[]"), - ) - safe_array = case( - (func.json_type(safe_json, literal("$")) == "array", safe_json), - else_=literal("[]"), - ) - elements = func.json_each(safe_array).table_valued("value") - return exists(select(1).select_from(elements).where(func.lower(elements.c.value).in_(values))) - def get_all_table_models(self) -> list[type[Base]]: """ Return a list of all table models used in the database by inspecting the Base registry. diff --git a/tests/integration/memory/test_seed_browsing_azure_sql_integration.py b/tests/integration/memory/test_seed_browsing_azure_sql_integration.py deleted file mode 100644 index 7021a18233..0000000000 --- a/tests/integration/memory/test_seed_browsing_azure_sql_integration.py +++ /dev/null @@ -1,159 +0,0 @@ -# Copyright (c) Microsoft Corporation. -# Licensed under the MIT license. - -"""Azure SQL execution contract for the #2748 seed browsing seam.""" - -from __future__ import annotations - -from contextlib import closing -from datetime import UTC, datetime, timedelta -from uuid import uuid4 - -import pytest - -from pyrit.memory import AzureSQLMemory, SeedExampleDatasetScope -from pyrit.memory.memory_models import SeedEntry -from pyrit.models import SeedObjective, SeedPrompt - - -@pytest.mark.run_only_if_all_tests -async def test_seed_browsing_contract_on_azure_sql(azuresql_instance: AzureSQLMemory): - test_id = str(uuid4()) - dataset = f"2748-azure-{test_id}" - shared_group = uuid4() - newer_group = uuid4() - older_group = uuid4() - base_time = datetime(2024, 1, 1, tzinfo=UTC) - seeds = [ - SeedPrompt( - value="https://example.com/violence", - dataset_name=dataset, - prompt_group_id=shared_group, - data_type="url", - harm_categories=["violence", "hate", 'special_%_"_é'], - date_added=base_time, - added_by=test_id, - ), - SeedObjective( - value="objective", - dataset_name=dataset, - prompt_group_id=shared_group, - date_added=base_time + timedelta(seconds=1), - added_by=test_id, - ), - SeedPrompt( - value="literal 100% a_b", - dataset_name=dataset, - prompt_group_id=newer_group, - date_added=base_time + timedelta(days=1), - harm_categories=["nonviolence", "violence-extra"], - added_by=test_id, - ), - SeedPrompt( - value="older", - dataset_name=dataset, - prompt_group_id=older_group, - date_added=base_time, - harm_categories=[], - added_by=test_id, - ), - SeedPrompt( - value="unnamed null", - prompt_group_id=shared_group, - harm_categories=None, - date_added=base_time, - added_by=test_id, - ), - SeedPrompt( - value="unnamed empty", - dataset_name="", - prompt_group_id=shared_group, - harm_categories=[], - date_added=base_time + timedelta(seconds=1), - added_by=test_id, - ), - ] - await azuresql_instance.add_seeds_to_memory_async(seeds=seeds, added_by=test_id) - - try: - named = SeedExampleDatasetScope.named(dataset) - unnamed = SeedExampleDatasetScope.unnamed() - - named_page = azuresql_instance.get_seed_example_page(dataset_scope=named, limit=100) - named_ids = {item.example_id for item in named_page.items} - assert named_ids == {shared_group, newer_group, older_group} - assert [item.example_id for item in named_page.items] == [ - newer_group, - *sorted((shared_group, older_group), reverse=True), - ] - shared = next(item for item in named_page.items if item.example_id == shared_group) - assert {member.id for member in shared.members} == {seeds[0].id, seeds[1].id} - assert shared.objective_count == 1 - - assert { - item.example_id for item in azuresql_instance.get_seed_example_page(dataset_scope=unnamed, limit=100).items - } == {shared_group} - - assert { - item.example_id - for item in azuresql_instance.get_seed_example_page( - dataset_scope=named, data_types=["url"], limit=100 - ).items - } == {shared_group} - assert { - item.example_id - for item in azuresql_instance.get_seed_example_page( - dataset_scope=named, seed_types=["objective"], limit=100 - ).items - } == {shared_group} - assert { - item.example_id - for item in azuresql_instance.get_seed_example_page( - dataset_scope=named, harm_categories=["VIOLENCE"], limit=100 - ).items - } == {shared_group} - assert { - item.example_id - for item in azuresql_instance.get_seed_example_page( - dataset_scope=named, harm_categories=["missing", "HATE"], limit=100 - ).items - } == {shared_group} - assert { - item.example_id - for item in azuresql_instance.get_seed_example_page( - dataset_scope=named, harm_categories=['special_%_"_É'], limit=100 - ).items - } == {shared_group} - assert { - item.example_id - for item in azuresql_instance.get_seed_example_page( - dataset_scope=named, harm_categories=["nonviolence"], limit=100 - ).items - } == {newer_group} - - assert { - item.example_id - for item in azuresql_instance.get_seed_example_page( - dataset_scope=named, value_search="100%", limit=100 - ).items - } == {newer_group} - assert { - item.example_id - for item in azuresql_instance.get_seed_example_page( - dataset_scope=named, value_search="a_b", limit=100 - ).items - } == {newer_group} - - unlabeled = next(item for item in named_page.items if item.example_id == older_group) - assert unlabeled.has_unlabeled_harm is True - - first = azuresql_instance.get_seed_example_page(dataset_scope=named, limit=2) - assert first.next_cursor is not None - second = azuresql_instance.get_seed_example_page(dataset_scope=named, limit=100, cursor=first.next_cursor) - paged_ids = [item.example_id for item in first.items + second.items] - assert len(paged_ids) == len(set(paged_ids)) == 3 - assert set(paged_ids) == named_ids - finally: - with closing(azuresql_instance.get_session()) as session: - session.query(SeedEntry).filter(SeedEntry.added_by == test_id).delete(synchronize_session=False) - session.commit() 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..6c50f1529a --- /dev/null +++ b/tests/integration/memory/test_seed_examples_azure_sql_integration.py @@ -0,0 +1,65 @@ +# 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 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, newer, older = uuid4(), uuid4(), 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, *sorted([group, older], reverse=True)] + 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=1) + rest, _, last = await azuresql_instance.get_seed_examples_async(dataset_name=dataset, limit=100, after=after) + assert after is not None + 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_dataset_service_seed_browsing_contract.py b/tests/unit/backend/test_dataset_service_seed_browsing_contract.py deleted file mode 100644 index 586a9a3f36..0000000000 --- a/tests/unit/backend/test_dataset_service_seed_browsing_contract.py +++ /dev/null @@ -1,297 +0,0 @@ -# Copyright (c) Microsoft Corporation. -# Licensed under the MIT license. - -"""RED contract tests for DatasetService seed browsing policy (#2748).""" - -from __future__ import annotations - -from datetime import UTC, datetime -from typing import TYPE_CHECKING, Any -from unittest.mock import patch -from uuid import uuid4 - -import pytest - -from pyrit.backend.models.common import PaginationInfo -from pyrit.backend.services.dataset_service import DatasetService -from pyrit.models import SeedObjective, SeedPrompt - -if TYPE_CHECKING: - from pyrit.memory import MemoryInterface - - -DATASET = "service-browse-contract" -SELECTION_KEY = f"dataset:named:{DATASET}" - - -def _field(value: Any, name: str) -> Any: - """Read a contract field from either a Pydantic result or a mapping.""" - return value.get(name) if isinstance(value, dict) else getattr(value, name) - - -def _service_method(service: DatasetService, name: str): - method = getattr(service, name, None) - assert method is not None, f"RED: DatasetService.{name} is not implemented" - return method - - -async def _add(memory: MemoryInterface, *seeds: SeedPrompt | SeedObjective) -> None: - await memory.add_seeds_to_memory_async(seeds=list(seeds), added_by="2748-service-test") - - -@pytest.fixture -def dataset_service(sqlite_instance: MemoryInterface) -> DatasetService: - with patch( - "pyrit.backend.services.dataset_service.CentralMemory.get_memory_instance", return_value=sqlite_instance - ): - yield DatasetService() - - -class TestDatasetServiceSeedBrowsingContract: - async def test_list_resolves_selection_key_and_returns_pagination_info( - self, dataset_service: DatasetService, sqlite_instance: MemoryInterface - ): - await _add(sqlite_instance, SeedPrompt(value="prompt", dataset_name=DATASET)) - response = await _service_method(dataset_service, "list_seed_examples_async")( - selection_key=SELECTION_KEY, limit=10 - ) - assert isinstance(_field(response, "pagination"), PaginationInfo) - assert _field(_field(response, "pagination"), "limit") == 10 - assert len(_field(response, "items")) == 1 - - async def test_unnamed_selection_key_is_supported_without_display_name_resolution( - self, dataset_service: DatasetService, sqlite_instance: MemoryInterface - ): - await _add(sqlite_instance, SeedPrompt(value="unnamed")) - response = await _service_method(dataset_service, "list_seed_examples_async")( - selection_key="dataset:unnamed", limit=10 - ) - assert len(_field(response, "items")) == 1 - - async def test_cursor_is_bound_to_selection_and_effective_filters( - self, dataset_service: DatasetService, sqlite_instance: MemoryInterface - ): - await _add(sqlite_instance, *(SeedPrompt(value=str(index), dataset_name=DATASET) for index in range(3))) - first = await _service_method(dataset_service, "list_seed_examples_async")(selection_key=SELECTION_KEY, limit=1) - cursor = _field(_field(first, "pagination"), "next_cursor") - assert cursor - with pytest.raises(ValueError): - await _service_method(dataset_service, "list_seed_examples_async")( - selection_key="dataset:unnamed", limit=1, cursor=cursor - ) - with pytest.raises(ValueError): - await _service_method(dataset_service, "list_seed_examples_async")( - selection_key=SELECTION_KEY, limit=1, cursor=cursor, search="changed" - ) - - @pytest.mark.parametrize( - "changed_filters", - [ - {"search": "changed"}, - {"data_types": ["url"]}, - {"harm_categories": ["violence"]}, - {"seed_types": ["objective"]}, - ], - ) - async def test_cursor_is_bound_to_each_effective_filter( - self, dataset_service: DatasetService, sqlite_instance: MemoryInterface, changed_filters: dict[str, object] - ): - await _add( - sqlite_instance, - SeedPrompt(value="one", dataset_name=DATASET), - SeedPrompt( - value="https://example.com/two", dataset_name=DATASET, data_type="url", harm_categories=["violence"] - ), - ) - first = await _service_method(dataset_service, "list_seed_examples_async")(selection_key=SELECTION_KEY, limit=1) - cursor = _field(_field(first, "pagination"), "next_cursor") - with pytest.raises(ValueError): - await _service_method(dataset_service, "list_seed_examples_async")( - selection_key=SELECTION_KEY, limit=1, cursor=cursor, **changed_filters - ) - - async def test_malformed_cursor_and_invalid_selection_are_rejected(self, dataset_service: DatasetService): - with pytest.raises(ValueError): - await _service_method(dataset_service, "list_seed_examples_async")( - selection_key=SELECTION_KEY, limit=10, cursor="malformed" - ) - with pytest.raises(ValueError): - await _service_method(dataset_service, "list_seed_examples_async")( - selection_key="display-name-not-selection-key", limit=10 - ) - - async def test_list_formats_group_preview_types_modalities_counts_and_harm_summary( - self, dataset_service: DatasetService, sqlite_instance: MemoryInterface - ): - group_id = uuid4() - await _add( - sqlite_instance, - SeedPrompt( - value="prompt", - dataset_name=DATASET, - prompt_group_id=group_id, - data_type="text", - harm_categories=["violence"], - ), - SeedObjective(value="objective", dataset_name=DATASET, prompt_group_id=group_id), - SeedPrompt(value="".join("x" for _ in range(101)), dataset_name=DATASET, prompt_group_id=group_id), - ) - response = await _service_method(dataset_service, "list_seed_examples_async")( - selection_key=SELECTION_KEY, limit=10 - ) - item = _field(response, "items")[0] - assert _field(item, "preview_truncated") is True - assert _field(item, "preview") == ("x" * 100) + "..." - assert _field(item, "piece_count") == 3 - assert _field(item, "objective_count") == 1 - assert "text" in _field(item, "modalities") - assert "violence" in _field(item, "harm_categories") - assert "objective" in _field(item, "seed_types") - - async def test_detail_returns_complete_group_and_persisted_provenance( - self, dataset_service: DatasetService, sqlite_instance: MemoryInterface - ): - group_id = uuid4() - seed = SeedPrompt( - value="full text", - dataset_name=DATASET, - prompt_group_id=group_id, - role="user", - sequence=4, - source="source", - authors=["author"], - groups=["group"], - metadata={"persisted": True}, - parameters=["name"], - ) - await _add( - sqlite_instance, - seed, - SeedObjective( - value="condition", - dataset_name=DATASET, - prompt_group_id=group_id, - metadata={"conditions": "must hold"}, - ), - ) - detail = await _service_method(dataset_service, "get_seed_example_async")( - selection_key=SELECTION_KEY, example_id=str(group_id) - ) - members = _field(detail, "members") - assert len(members) == 2 - prompt = next(member for member in members if _field(member, "id") == seed.id) - objective = next(member for member in members if _field(member, "seed_type") == "objective") - assert _field(prompt, "prompt_group_id") == group_id - assert _field(prompt, "name") == seed.name - assert _field(prompt, "value") == seed.value - assert _field(prompt, "role") == "user" - assert _field(prompt, "sequence") == 4 - stored_prompt = next( - stored for stored in sqlite_instance.get_seeds(prompt_group_ids=[group_id]) if stored.id == seed.id - ) - assert _field(prompt, "value_sha256") == stored_prompt.value_sha256 - assert _field(prompt, "dataset_name") == DATASET - assert _field(prompt, "source") == "source" - assert _field(prompt, "authors") == ["author"] - assert _field(prompt, "groups") == ["group"] - assert _field(prompt, "date_added") == stored_prompt.date_added - assert _field(prompt, "added_by") == "2748-service-test" - assert _field(prompt, "metadata") == {"persisted": True} - assert _field(prompt, "data_type") == "text" - assert _field(objective, "value") == "condition" - assert _field(objective, "metadata") == {"conditions": "must hold"} - - async def test_detail_preserves_template_parameters_and_objective_conditions( - self, dataset_service: DatasetService, sqlite_instance: MemoryInterface - ): - group_id = uuid4() - await _add( - sqlite_instance, - SeedPrompt( - value="{{ name }}", - dataset_name=DATASET, - prompt_group_id=group_id, - is_jinja_template=True, - parameters=["name"], - ), - SeedObjective(value="condition", dataset_name=DATASET, prompt_group_id=group_id), - ) - detail = await _service_method(dataset_service, "get_seed_example_async")( - selection_key=SELECTION_KEY, example_id=str(group_id) - ) - assert _field(detail, "members") - template = next(member for member in _field(detail, "members") if _field(member, "seed_type") == "prompt") - assert _field(template, "value") == "{{ name }}" - assert _field(template, "is_jinja_template") is True - assert _field(template, "parameters") == ["name"] - assert any(_field(member, "seed_type") == "objective" for member in _field(detail, "members")) - - async def test_invalid_detail_does_not_generate_group_identity(self, dataset_service: DatasetService): - with pytest.raises(ValueError): - await _service_method(dataset_service, "get_seed_example_async")( - selection_key=SELECTION_KEY, example_id=str(uuid4()) - ) - - async def test_detail_retrieves_logical_example_beyond_first_list_page( - self, dataset_service: DatasetService, sqlite_instance: MemoryInterface - ): - group_id = uuid4() - target = SeedPrompt( - value="target", - dataset_name=DATASET, - prompt_group_id=group_id, - date_added=datetime(2020, 1, 1, tzinfo=UTC), - ) - await _add(sqlite_instance, target) - await _add( - sqlite_instance, - *(SeedPrompt(value=f"later-{index}", dataset_name=DATASET) for index in range(100)), - ) - - detail = await _service_method(dataset_service, "get_seed_example_async")( - selection_key=SELECTION_KEY, example_id=str(group_id) - ) - assert [member.id for member in _field(detail, "members")] == [target.id] - - async def test_browsing_has_no_template_or_generation_side_effects( - self, dataset_service: DatasetService, sqlite_instance: MemoryInterface - ): - await _add(sqlite_instance, SeedPrompt(value="{{ dangerous }}", dataset_name=DATASET, is_jinja_template=True)) - with patch("pyrit.models.SeedPrompt.render_template_value", side_effect=AssertionError("rendered")) as render: - response = await _service_method(dataset_service, "list_seed_examples_async")( - selection_key=SELECTION_KEY, limit=10 - ) - assert _field(response, "items") - assert render.call_count == 0 - - async def test_media_preview_is_type_label_without_bytes_or_path_leak( - self, dataset_service: DatasetService, sqlite_instance: MemoryInterface, tmp_path - ): - media = tmp_path / "private-image.png" - media.write_bytes(b"local image") - await _add(sqlite_instance, SeedPrompt(value=str(media), dataset_name=DATASET, data_type="image_path")) - response = await _service_method(dataset_service, "list_seed_examples_async")( - selection_key=SELECTION_KEY, limit=10 - ) - item = _field(response, "items")[0] - assert "image" in _field(item, "preview").lower() - assert str(tmp_path) not in _field(item, "preview") - assert "bytes" not in (item if isinstance(item, dict) else item.model_dump()) - - @pytest.mark.parametrize( - ("value", "expected_preview"), - [ - ("/private/secret.txt", "text seed"), - (r"C:\\private\\secret.txt", "text seed"), - (r"\\\\server\\share\\secret.txt", "text seed"), - ("ordinary safe text", "ordinary safe text"), - ], - ) - async def test_text_preview_rejects_absolute_paths_including_unc( - self, dataset_service: DatasetService, sqlite_instance: MemoryInterface, value: str, expected_preview: str - ): - await _add(sqlite_instance, SeedPrompt(value=value, dataset_name=DATASET, data_type="text")) - response = await _service_method(dataset_service, "list_seed_examples_async")( - selection_key=SELECTION_KEY, limit=10 - ) - assert _field(_field(response, "items")[0], "preview") == expected_preview diff --git a/tests/unit/backend/test_seed_browsing_api.py b/tests/unit/backend/test_seed_browsing_api.py deleted file mode 100644 index 8771353337..0000000000 --- a/tests/unit/backend/test_seed_browsing_api.py +++ /dev/null @@ -1,510 +0,0 @@ -# Copyright (c) Microsoft Corporation. -# Licensed under the MIT license. - -"""RED contract tests for the paginated seed browsing API (#2748). - -These tests deliberately target the approved public HTTP contract. The route, -models, and narrow memory read helpers are not present until the feature is -implemented; consequently the current expected failure is a missing-route -response, rather than a test-time production stub. -""" - -from __future__ import annotations - -from datetime import UTC, datetime, timedelta -from typing import TYPE_CHECKING -from unittest.mock import patch -from uuid import UUID, uuid4 - -import pytest -from fastapi.testclient import TestClient - -from pyrit.backend.main import app -from pyrit.backend.services.dataset_service import get_dataset_service -from pyrit.datasets import SeedDatasetProvider -from pyrit.models import SeedObjective, SeedPrompt, SeedSimulatedConversation - -if TYPE_CHECKING: - from pyrit.memory import MemoryInterface - - -DATASET = "browse-contract" -NAMED_KEY = f"dataset:named:{DATASET}" -UNNAMED_KEY = "dataset:unnamed" - - -@pytest.fixture -def client(patch_central_database) -> TestClient: - """Use the real SQLite memory fixture behind the API application.""" - # DatasetService is globally cached in production. Each test gets a fresh - # fixture-owned SQLiteMemory, so do not retain a service bound to a previous - # test's in-memory engine after that fixture disposes it. - get_dataset_service.cache_clear() - try: - yield TestClient(app) - finally: - get_dataset_service.cache_clear() - - -async def _add(memory: MemoryInterface, *seeds: SeedPrompt | SeedObjective | SeedSimulatedConversation) -> None: - await memory.add_seeds_to_memory_async(seeds=list(seeds), added_by="2748-test") - - -def _list(client: TestClient, selection_key: str = NAMED_KEY, **params: object): - return client.get(f"/api/datasets/{selection_key}/seeds", params=params) - - -def _detail(client: TestClient, example_id: str, selection_key: str = NAMED_KEY): - return client.get(f"/api/datasets/{selection_key}/seeds/{example_id}") - - -def _items(response): - assert response.status_code == 200, response.text - return response.json()["items"] - - -class TestEmptyAndDatasetSelection: - async def test_empty_dataset_is_a_valid_empty_page(self, client, sqlite_instance: MemoryInterface): - response = _list(client, UNNAMED_KEY) - assert response.status_code == 200 - body = response.json() - assert body["items"] == [] - assert body["pagination"]["has_more"] is False - assert body["pagination"]["next_cursor"] is None - - async def test_empty_filter_result_is_not_an_error(self, client, sqlite_instance: MemoryInterface): - await _add(sqlite_instance, SeedPrompt(value="ordinary", dataset_name=DATASET)) - response = _list(client, search="absent") - assert response.status_code == 200 - assert response.json()["items"] == [] - - async def test_named_selection_key_is_not_display_name(self, client, sqlite_instance: MemoryInterface): - await _add(sqlite_instance, SeedPrompt(value="named", dataset_name=DATASET)) - assert len(_items(_list(client, NAMED_KEY))) == 1 - assert _list(client, DATASET).status_code == 404 - - async def test_unnamed_selection_and_invalid_selection(self, client, sqlite_instance: MemoryInterface): - await _add(sqlite_instance, SeedPrompt(value="unnamed")) - assert len(_items(_list(client, UNNAMED_KEY))) == 1 - assert _list(client, "dataset:named:not-loaded").status_code == 404 - - async def test_named_unnamed_namespace_is_distinct(self, client, sqlite_instance: MemoryInterface): - await _add(sqlite_instance, SeedPrompt(value="literal", dataset_name="__unnamed__"), SeedPrompt(value="none")) - assert len(_items(_list(client, "dataset:named:__unnamed__"))) == 1 - assert len(_items(_list(client, UNNAMED_KEY))) == 1 - - -class TestPaginationAndIdentity: - async def test_one_page_has_existing_pagination_shape(self, client, sqlite_instance: MemoryInterface): - await _add(sqlite_instance, SeedPrompt(value="one", dataset_name=DATASET)) - page = _list(client, limit=1) - assert page.status_code == 200 - assert set(page.json()["pagination"]) >= {"limit", "has_more", "next_cursor", "prev_cursor"} - - async def test_multiple_pages_have_no_duplicates_or_omissions(self, client, sqlite_instance: MemoryInterface): - seeds = [SeedPrompt(value=f"prompt-{i}", dataset_name=DATASET) for i in range(5)] - await _add(sqlite_instance, *seeds) - first = _list(client, limit=2) - first_items = _items(first) - second = _list(client, limit=2, cursor=first.json()["pagination"]["next_cursor"]) - all_items = first_items + _items(second) - while second.json()["pagination"]["has_more"]: - second = _list(client, limit=2, cursor=second.json()["pagination"]["next_cursor"]) - all_items += _items(second) - ids = [item["example_id"] for item in all_items] - assert len(ids) == len(set(ids)) == 5 - - @pytest.mark.parametrize("limit", [0, -1, 101]) - async def test_page_size_is_validated(self, client, sqlite_instance: MemoryInterface, limit: int): - assert _list(client, limit=limit).status_code == 422 - - async def test_group_identity_preserves_ids_and_does_not_hash_merge(self, client, sqlite_instance: MemoryInterface): - group_id = uuid4() - first = SeedPrompt(value="same", dataset_name=DATASET, prompt_group_id=group_id) - second = SeedPrompt(value="same", dataset_name=DATASET, prompt_group_id=group_id) - ungrouped = SeedPrompt(value="same", dataset_name=DATASET) - await _add(sqlite_instance, first, second, ungrouped) - items = _items(_list(client)) - assert len(items) == 2 - assert {str(first.id), str(second.id), str(ungrouped.id)} == set(items[0]["seed_ids"] + items[1]["seed_ids"]) - assert sorted(len(item["seed_ids"]) for item in items) == [1, 2] - assert all(item["example_id"] in {str(group_id), str(ungrouped.id)} for item in items) - assert all("generated" not in item["example_id"] for item in items) - - async def test_order_is_complete_group_date_then_id(self, client, sqlite_instance: MemoryInterface): - tied = datetime(2024, 1, 1, tzinfo=UTC) - old_group = uuid4() - new_group = uuid4() - await _add( - sqlite_instance, - SeedPrompt(value="old", dataset_name=DATASET, prompt_group_id=old_group, date_added=tied), - SeedPrompt( - value="new", dataset_name=DATASET, prompt_group_id=new_group, date_added=tied + timedelta(days=1) - ), - SeedPrompt( - value="late member", - dataset_name=DATASET, - prompt_group_id=old_group, - date_added=tied + timedelta(days=2), - ), - ) - items = _items(_list(client)) - assert [item["example_id"] for item in items] == [str(new_group), str(old_group)] - - async def test_tied_dates_use_descending_logical_id_tie_breaker(self, client, sqlite_instance: MemoryInterface): - date_added = datetime(2024, 1, 1, tzinfo=UTC) - lower = UUID("00000000-0000-0000-0000-000000000001") - higher = UUID("00000000-0000-0000-0000-000000000002") - await _add( - sqlite_instance, - SeedPrompt(value="lower", dataset_name=DATASET, prompt_group_id=lower, date_added=date_added), - SeedPrompt(value="higher", dataset_name=DATASET, prompt_group_id=higher, date_added=date_added), - ) - assert [item["example_id"] for item in _items(_list(client))] == [str(higher), str(lower)] - - async def test_group_is_never_split_and_detail_preserves_role_sequence( - self, client, sqlite_instance: MemoryInterface, tmp_path - ): - group_id = uuid4() - image_path = tmp_path / "group-image.png" - image_path.write_bytes(b"local test image") - await _add( - sqlite_instance, - SeedPrompt( - value=str(image_path), - dataset_name=DATASET, - prompt_group_id=group_id, - data_type="image_path", - sequence=0, - ), - SeedPrompt(value="text", dataset_name=DATASET, prompt_group_id=group_id, data_type="text", sequence=1), - ) - page = _list(client, limit=1) - assert len(_items(page)) == 1 - detail = _detail(client, str(group_id)) - members = detail.json()["members"] if detail.status_code == 200 else [] - assert [member["sequence"] for member in members] == [0, 1] - assert all("role" in member for member in members) - - -class TestFilters: - async def test_modality_is_or_and_matching_member_returns_complete_group( - self, client, sqlite_instance: MemoryInterface, tmp_path - ): - group_id = uuid4() - image_path = tmp_path / "group-image.png" - audio_path = tmp_path / "standalone-audio.wav" - image_path.write_bytes(b"local test image") - audio_path.write_bytes(b"local test audio") - await _add( - sqlite_instance, - SeedPrompt(value=str(image_path), dataset_name=DATASET, prompt_group_id=group_id, data_type="image_path"), - SeedPrompt(value="text", dataset_name=DATASET, prompt_group_id=group_id, data_type="text"), - SeedPrompt(value=str(audio_path), dataset_name=DATASET, data_type="audio_path"), - ) - items = _items(_list(client, modality=["image_path", "audio_path"])) - assert {item["example_id"] for item in items} == {str(group_id)} | { - str(next(seed.id for seed in sqlite_instance.get_seeds(data_types=["audio_path"]))) - } - assert len(_detail(client, str(group_id)).json()["members"]) == 2 - - @pytest.mark.parametrize("category, expected", [("Violence", True), ("vio", False), ("missing", False)]) - async def test_harm_category_is_case_insensitive_whole_value_and_missing_is_unlabeled( - self, client, sqlite_instance: MemoryInterface, category: str, expected: bool - ): - seed = SeedPrompt(value="harm", dataset_name=DATASET, harm_categories=["VIOLENCE"]) - unlabeled = SeedPrompt(value="unlabeled", dataset_name=DATASET, harm_categories=[]) - await _add(sqlite_instance, seed, unlabeled) - items = _items(_list(client, harm_category=category)) - assert (len(items) == 1 and str(seed.id) in items[0]["seed_ids"]) is expected - all_items = _items(_list(client)) - unlabeled_item = next(item for item in all_items if str(unlabeled.id) in item["seed_ids"]) - assert unlabeled_item["has_unlabeled_harm"] is True - - async def test_multiple_harm_categories_are_or_not_existing_get_seeds_all_semantics(self, client, sqlite_instance): - await _add( - sqlite_instance, - SeedPrompt(value="hate", dataset_name=DATASET, harm_categories=["hate"]), - SeedPrompt(value="violence", dataset_name=DATASET, harm_categories=["violence"]), - SeedPrompt(value="both", dataset_name=DATASET, harm_categories=["hate", "violence"]), - ) - items = _items(_list(client, harm_category=["hate", "violence"])) - assert len(items) == 3 - assert {item["seed_ids"][0] for item in items} == { - str(seed.id) for seed in sqlite_instance.get_seeds(dataset_name=DATASET) - } - - async def test_seed_type_filter_is_or_and_returns_complete_group(self, client, sqlite_instance): - group_id = uuid4() - await _add( - sqlite_instance, - SeedPrompt(value="prompt", dataset_name=DATASET, prompt_group_id=group_id), - SeedObjective(value="objective", dataset_name=DATASET, prompt_group_id=group_id), - ) - items = _items(_list(client, seed_type=["objective"])) - assert len(items) == 1 and items[0]["piece_count"] == 2 - - async def test_filters_are_and_across_filters_but_match_at_example_level(self, client, sqlite_instance, tmp_path): - group_id = uuid4() - image_path = tmp_path / "filter-image.png" - image_path.write_bytes(b"local test image") - await _add( - sqlite_instance, - SeedPrompt(value=str(image_path), dataset_name=DATASET, prompt_group_id=group_id, data_type="image_path"), - SeedPrompt(value="violence", dataset_name=DATASET, prompt_group_id=group_id, harm_categories=["violence"]), - SeedPrompt(value=str(image_path), dataset_name=DATASET, data_type="image_path"), - ) - items = _items(_list(client, modality="image_path", harm_category="violence")) - assert len(items) == 1 and items[0]["example_id"] == str(group_id) - - -class TestTextSearchAndSafety: - async def test_browsing_selection_validation_does_not_discover_providers_or_read_files( - self, client, sqlite_instance - ): - seed = SeedPrompt(value="stored", dataset_name=DATASET) - await _add(sqlite_instance, seed) - with ( - patch.object( - SeedDatasetProvider, - "get_all_dataset_names_async", - side_effect=AssertionError("provider metadata discovery"), - ), - patch.object(SeedDatasetProvider, "_parse_metadata_async", side_effect=AssertionError("metadata parse")), - patch("pathlib.Path.read_text", side_effect=AssertionError("provider file read")), - ): - listed = _list(client) - assert listed.status_code == 200 - example_id = listed.json()["items"][0]["example_id"] - detail = _detail(client, example_id) - assert detail.status_code == 200 - - async def test_text_search_is_literal_case_insensitive_and_text_only(self, client, sqlite_instance, tmp_path): - image_path = tmp_path / "media-path.png" - image_path.write_bytes(b"local test image") - await _add( - sqlite_instance, - SeedPrompt(value="Need 100% literal_value", dataset_name=DATASET), - SeedObjective(value="OBJECTIVE text", dataset_name=DATASET), - SeedPrompt(value=str(image_path), dataset_name=DATASET, data_type="image_path"), - SeedPrompt(value="metadata-only", dataset_name=DATASET, metadata={"secret": "literal_value"}), - ) - assert len(_items(_list(client, search="100% literal_value"))) == 1 - assert len(_items(_list(client, search="literal_value"))) == 1 - assert len(_items(_list(client, search="objective"))) == 1 - assert _items(_list(client, search="image_path")) == [] - - async def test_cursor_is_opaque_bound_to_dataset_and_effective_filters(self, client, sqlite_instance): - await _add(sqlite_instance, *(SeedPrompt(value=str(i), dataset_name=DATASET) for i in range(3))) - cursor = _list(client, limit=1).json()["pagination"]["next_cursor"] - assert cursor and not cursor.startswith("1") - assert _list(client, limit=1, cursor="not-a-cursor").status_code == 400 - assert _list(client, UNNAMED_KEY, limit=1, cursor=cursor).status_code == 400 - assert _list(client, limit=1, search="different", cursor=cursor).status_code == 400 - - @pytest.mark.parametrize( - "changed_filters", - [ - {"search": "different"}, - {"modality": "url"}, - {"harm_category": "violence"}, - {"seed_type": "objective"}, - ], - ) - async def test_cursor_rejects_each_changed_effective_filter(self, client, sqlite_instance, changed_filters): - await _add( - sqlite_instance, - SeedPrompt(value="one", dataset_name=DATASET), - SeedPrompt( - value="https://example.com/two", dataset_name=DATASET, data_type="url", harm_categories=["violence"] - ), - ) - first = _list(client, limit=1) - cursor = first.json()["pagination"]["next_cursor"] - response = _list(client, limit=1, cursor=cursor, **changed_filters) - assert response.status_code == 400 - assert response.json()["detail"] - - async def test_missing_detail_example_is_not_silently_empty(self, client, sqlite_instance): - response = _detail(client, str(uuid4())) - assert response.status_code == 404 - assert response.json()["detail"] - - async def test_template_is_not_rendered_or_loaded(self, client, sqlite_instance): - template = SeedPrompt( - value="{{ dangerous }}", dataset_name=DATASET, is_jinja_template=True, parameters=["dangerous"] - ) - await _add(sqlite_instance, template) - with patch.object(SeedPrompt, "render_template_value", side_effect=AssertionError("rendered")) as render: - with patch("pathlib.Path.read_text", side_effect=AssertionError("loaded")) as load: - response = _list(client, search="dangerous") - assert response.status_code == 200 - assert render.call_count == load.call_count == 0 - item = response.json()["items"][0] - assert item["is_template"] is True - assert item["parameters"] == ["dangerous"] - - async def test_simulated_configuration_is_returned_without_generation_or_target_call(self, client, sqlite_instance): - config = SeedSimulatedConversation( - dataset_name=DATASET, - adversarial_chat_system_prompt=SeedPrompt(value="{{ objective }}"), - simulated_target_system_prompt=SeedPrompt(value="{{ objective }}"), - num_turns=2, - ) - await _add(sqlite_instance, config) - with patch( - "pyrit.executor.attack.multi_turn.simulated_conversation.generate_simulated_conversation_async" - ) as generate: - response = _list(client, search="num_turns") - assert response.status_code == 200 - assert generate.call_count == 0 - assert response.json()["items"][0]["seed_types"] == ["simulated_conversation"] - - detail = _detail(client, response.json()["items"][0]["example_id"]) - assert detail.status_code == 200 - assert detail.json()["members"][0]["value"] == config.value - - -class TestPreviewDetailAndCounts: - async def test_preview_uses_100_character_convention_and_hides_full_content(self, client, sqlite_instance): - short = "x" * 100 - long = "y" * 101 - await _add( - sqlite_instance, SeedPrompt(value=short, dataset_name=DATASET), SeedPrompt(value=long, dataset_name=DATASET) - ) - items = _items(_list(client)) - assert any(item["preview"] == short and item["preview_truncated"] is False for item in items) - long_item = next(item for item in items if item["preview"].startswith("y")) - assert long_item["preview"] == ("y" * 100) + "..." - assert long_item["preview_truncated"] is True - assert long not in long_item["preview"] - - async def test_detail_returns_full_content_after_list_preview_truncation(self, client, sqlite_instance): - long_value = "long-value-" + ("x" * 150) - seed = SeedPrompt(value=long_value, dataset_name=DATASET) - await _add(sqlite_instance, seed) - item = _items(_list(client))[0] - assert item["preview_truncated"] is True - detail = _detail(client, item["example_id"]) - assert detail.status_code == 200 - member = detail.json()["members"][0] - assert member["value"] == long_value - assert member["prompt_group_id"] is None - - async def test_media_preview_is_label_only_and_never_bytes_path_or_credentials( - self, client, sqlite_instance, tmp_path - ): - image_path = tmp_path / "image.png" - image_path.write_bytes(b"local test image") - await _add(sqlite_instance, SeedPrompt(value=str(image_path), dataset_name=DATASET, data_type="image_path")) - item = _items(_list(client))[0] - assert "image" in item["preview"].lower() - assert str(tmp_path) not in item["preview"] and "sig=" not in item["preview"] - assert "bytes" not in item and "content" not in item - - async def test_detail_returns_all_persisted_fields_without_new_ids_or_rendering(self, client, sqlite_instance): - seed_id = uuid4() - group_id = uuid4() - seed = SeedPrompt( - id=seed_id, - value="full text", - dataset_name=DATASET, - prompt_group_id=group_id, - role="user", - sequence=4, - source="source", - authors=["author"], - groups=["group"], - metadata={"persisted": "yes"}, - ) - await _add(sqlite_instance, seed) - response = _detail(client, str(group_id)) - assert response.status_code == 200 - member = response.json()["members"][0] - assert member["id"] == str(seed_id) - assert member["prompt_group_id"] == str(group_id) - assert member["value"] == "full text" - assert member["role"] == "user" - assert member["sequence"] == 4 - assert member["source"] == "source" - assert member["authors"] == ["author"] - assert member["groups"] == ["group"] - assert member["metadata"] == {"persisted": "yes"} - for field in ( - "role", - "sequence", - "value_sha256", - "dataset_name", - "source", - "authors", - "groups", - "date_added", - "added_by", - "metadata", - "data_type", - ): - assert field in member - - async def test_counts_are_logical_examples_and_use_same_predicates(self, client, sqlite_instance, tmp_path): - group_id = uuid4() - image_path = tmp_path / "count-image.png" - image_path.write_bytes(b"local test image") - await _add( - sqlite_instance, - SeedPrompt(value=str(image_path), dataset_name=DATASET, prompt_group_id=group_id, data_type="image_path"), - SeedObjective(value="two", dataset_name=DATASET, prompt_group_id=group_id), - SeedPrompt(value="three", dataset_name=DATASET, data_type="text"), - ) - response = _list(client, modality="image_path") - assert response.status_code == 200 - body = response.json() - assert body["total"] == 1 - assert body["items"][0]["piece_count"] == 2 - assert body["items"][0]["objective_count"] == 1 - - -class TestDatabaseBoundsAndSideEffects: - async def test_page_query_is_bounded_and_does_not_n_plus_one(self, client, sqlite_instance): - await _add(sqlite_instance, *(SeedPrompt(value=f"p{i}", dataset_name=DATASET) for i in range(250))) - statements = [] - from sqlalchemy import event - - def capture(_connection, _cursor, statement, _parameters, _context, _executemany): - statements.append(statement.lower()) - - event.listen(sqlite_instance.engine, "before_cursor_execute", capture) - try: - response = _list(client, limit=2) - finally: - event.remove(sqlite_instance.engine, "before_cursor_execute", capture) - assert response.status_code == 200 - assert len(response.json()["items"]) == 2 - assert len(statements) < 12 - assert not any("select" in statement and "250" in statement for statement in statements) - - async def test_browsing_is_read_only_and_does_not_fetch_provider_or_write(self, client, sqlite_instance): - with ( - patch( - "pyrit.datasets.SeedDatasetProvider.fetch_datasets_async", - side_effect=AssertionError("provider fetch"), - ), - patch.object(sqlite_instance, "add_seeds_to_memory_async", side_effect=AssertionError("write")) as write, - ): - response = _list(client, UNNAMED_KEY) - assert response.status_code == 200 - assert write.call_count == 0 - - -class TestCompatibility: - async def test_existing_dataset_list_route_remains_unchanged(self, client): - response = client.get("/api/datasets") - assert response.status_code == 200 - assert "items" in response.json() - - async def test_existing_get_seeds_harm_semantics_remain_all_categories(self, sqlite_instance: MemoryInterface): - await _add( - sqlite_instance, - SeedPrompt(value="one", harm_categories=["hate"]), - SeedPrompt(value="two", harm_categories=["hate", "violence"]), - ) - assert len(sqlite_instance.get_seeds(harm_categories=["hate", "violence"])) == 1 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..4dd90a256c --- /dev/null +++ b/tests/unit/backend/test_seed_example_routes.py @@ -0,0 +1,172 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +from collections.abc import AsyncIterator +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.memory import MemoryInterface +from pyrit.memory.memory_models import SeedEntry +from pyrit.models import AnswerMatches, MatchesObjective, 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]" + + +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_members_as_domain_seeds(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)]["num_turns"] == 2 + assert members[str(configuration.id)]["adversarial_chat_system_prompt"]["value"] == "adversarial" + assert (body["preview"], body["objective_count"], body["has_unlabeled_harm"]) == ("objective", 1, True) + assert unnamed.status_code == 404 diff --git a/tests/unit/memory/memory_interface/test_interface_seed_browsing_contract.py b/tests/unit/memory/memory_interface/test_interface_seed_browsing_contract.py deleted file mode 100644 index 048c2d82de..0000000000 --- a/tests/unit/memory/memory_interface/test_interface_seed_browsing_contract.py +++ /dev/null @@ -1,654 +0,0 @@ -# Copyright (c) Microsoft Corporation. -# Licensed under the MIT license. - -"""RED contract tests for the portable #2748 logical-example read helper.""" - -from __future__ import annotations - -import base64 -import json -from datetime import UTC, datetime, timedelta -from pathlib import Path -from typing import TYPE_CHECKING, Any -from unittest.mock import patch -from uuid import UUID, uuid4 - -import pytest -from sqlalchemy import event, text - -from pyrit.memory import SeedExampleDatasetScope -from pyrit.memory.memory_models import SeedEntry -from pyrit.models import SeedObjective, SeedPrompt, SeedSimulatedConversation - -if TYPE_CHECKING: - from pyrit.memory import MemoryInterface - - -DATASET = "memory-browse-contract" - - -def _field(value: Any, name: str) -> Any: - """Read a contract field from either a Pydantic result or a mapping.""" - return value.get(name) if isinstance(value, dict) else getattr(value, name) - - -def _page(memory: MemoryInterface, **kwargs: Any) -> Any: - """Call the proposed narrow, database-backed logical-example page helper.""" - helper = getattr(memory, "get_seed_example_page", None) - assert helper is not None, "RED: MemoryInterface.get_seed_example_page is not implemented" - dataset_name = kwargs.pop("dataset_name", None) - kwargs["dataset_scope"] = ( - SeedExampleDatasetScope.named(dataset_name) if dataset_name is not None else SeedExampleDatasetScope.unnamed() - ) - return helper(**kwargs) - - -def _example(memory: MemoryInterface, **kwargs: Any) -> Any: - """Call the proposed bounded logical-example point lookup helper.""" - helper = getattr(memory, "get_seed_example", None) - assert helper is not None, "RED: MemoryInterface.get_seed_example is not implemented" - dataset_name = kwargs.pop("dataset_name", None) - kwargs["dataset_scope"] = ( - SeedExampleDatasetScope.named(dataset_name) if dataset_name is not None else SeedExampleDatasetScope.unnamed() - ) - return helper(**kwargs) - - -async def _add(memory: MemoryInterface, *seeds: SeedPrompt | SeedObjective | SeedSimulatedConversation) -> None: - await memory.add_seeds_to_memory_async(seeds=list(seeds), added_by="2748-memory-test") - - -class TestSeedBrowsingMemoryContract: - def test_scope_and_page_limits_are_explicit(self): - with pytest.raises(ValueError): - SeedExampleDatasetScope(kind="invalid") # type: ignore[arg-type] - - with pytest.raises(ValueError): - SeedExampleDatasetScope.named("") - - async def test_named_and_unnamed_scopes_are_distinct(self, sqlite_instance: MemoryInterface): - shared_group = uuid4() - unnamed_null = SeedPrompt(value="unnamed-null", prompt_group_id=shared_group) - unnamed_empty = SeedPrompt(value="unnamed-empty", dataset_name="", prompt_group_id=shared_group) - named_seed = SeedPrompt(value="named", dataset_name=DATASET, prompt_group_id=shared_group) - await _add( - sqlite_instance, - unnamed_null, - unnamed_empty, - named_seed, - ) - - unnamed = _page(sqlite_instance, limit=10) - named = _page(sqlite_instance, dataset_name=DATASET, limit=10) - assert len(_field(unnamed, "items")) == 1 - assert {str(member.id) for member in _field(unnamed, "items")[0].members} == { - str(unnamed_null.id), - str(unnamed_empty.id), - } - assert len(_field(named, "items")) == 1 - assert {str(member.id) for member in _field(named, "items")[0].members} == { - str(named_seed.id), - } - - with pytest.raises(ValueError): - _page(sqlite_instance, dataset_name=DATASET, limit=101) - - async def test_browsing_uses_persisted_rows_without_get_seed_or_filesystem_access( - self, sqlite_instance: MemoryInterface - ): - legacy_path = "C:/not-present/legacy-prompt.yaml" - legacy_value = json.dumps( - { - "num_turns": 2, - "sequence": 0, - "adversarial_chat_system_prompt_path": legacy_path, - "simulated_target_system_prompt_path": legacy_path, - } - ) - entry = SeedEntry(entry=SeedPrompt(value="stored", dataset_name=DATASET)) - entry.value = legacy_value - entry.value_sha256 = "legacy-hash" - entry.prompt_metadata = {"persisted": "yes"} - entry.seed_type = "simulated_conversation" - entry.date_added = datetime(2024, 1, 1, tzinfo=UTC) - entry.added_by = "legacy-test" - entry_id = entry.id - with sqlite_instance.get_session() as session: - session.add(entry) - session.commit() - - with ( - patch.object(SeedEntry, "get_seed", side_effect=AssertionError("get_seed called")), - patch.object(Path, "is_file", side_effect=AssertionError("filesystem touched")), - ): - page = _page(sqlite_instance, dataset_name=DATASET, limit=10) - - member = _field(page, "items")[0].members[0] - assert member.id == entry_id - assert member.value == legacy_value - assert member.value_sha256 == "legacy-hash" - assert member.metadata == {"persisted": "yes"} - assert member.seed_type == "simulated_conversation" - - async def test_template_persisted_value_and_parameters_are_not_rendered(self, sqlite_instance: MemoryInterface): - template = SeedPrompt( - value="{{ name }}", - dataset_name=DATASET, - parameters=["name"], - is_jinja_template=True, - ) - await _add(sqlite_instance, template) - with patch.object(SeedEntry, "get_seed", side_effect=AssertionError("get_seed called")): - member = _field(_page(sqlite_instance, dataset_name=DATASET, limit=10), "items")[0].members[0] - assert member.value == "{{ name }}" - assert member.parameters == ["name"] - assert member.is_jinja_template is True - - @pytest.mark.parametrize( - ("is_template", "parameters"), - [(True, []), (False, ["name"]), (False, [])], - ) - async def test_template_flag_is_persisted_independently_of_parameters( - self, sqlite_instance: MemoryInterface, is_template: bool, parameters: list[str] - ): - seed = SeedPrompt( - value="{{ name }}" if is_template else "ordinary", - dataset_name=DATASET, - is_jinja_template=is_template, - parameters=parameters, - ) - await _add(sqlite_instance, seed) - - member = _field(_page(sqlite_instance, dataset_name=DATASET, limit=10), "items")[0].members[0] - assert member.is_jinja_template is is_template - - async def test_historical_null_template_flag_stays_null_in_browsing_projection( - self, sqlite_instance: MemoryInterface - ): - entry = SeedEntry( - entry=SeedPrompt(value="historical", dataset_name=DATASET, parameters=["name"], added_by="legacy-test") - ) - entry.is_jinja_template = None - with sqlite_instance.get_session() as session: - session.add(entry) - session.commit() - - member = _field(_page(sqlite_instance, dataset_name=DATASET, limit=10), "items")[0].members[0] - assert member.parameters == ["name"] - assert member.is_jinja_template is None - - async def test_point_lookup_is_bounded_and_scope_isolated(self, sqlite_instance: MemoryInterface): - target_group = uuid4() - target = SeedPrompt( - value="target", - dataset_name=DATASET, - prompt_group_id=target_group, - date_added=datetime(2020, 1, 1, tzinfo=UTC), - ) - other_dataset_member = SeedPrompt( - value="wrong dataset", - dataset_name="another-dataset", - prompt_group_id=target_group, - ) - unnamed_member = SeedPrompt(value="unnamed", prompt_group_id=target_group) - await _add(sqlite_instance, target, other_dataset_member, unnamed_member) - await _add( - sqlite_instance, - *(SeedPrompt(value=f"later-{index}", dataset_name=DATASET) for index in range(100)), - ) - - result = _example(sqlite_instance, dataset_name=DATASET, example_id=target_group) - assert result is not None - assert {member.id for member in result.members} == {target.id} - - assert _example(sqlite_instance, example_id=target_group) is not None - assert _example(sqlite_instance, dataset_name="another-dataset", example_id=target_group) is not None - assert _example(sqlite_instance, example_id=uuid4()) is None - - async def test_sqlite_harm_json_invalid_and_non_arrays_are_unlabeled(self, sqlite_instance: MemoryInterface): - values = [None, "null", "[]", '"violence"', '{"category":"violence"}', "not-json", '["violence"]'] - seeds = [SeedPrompt(value=f"harm-{index}", dataset_name=DATASET) for index in range(len(values))] - await _add(sqlite_instance, *seeds) - with sqlite_instance.get_session() as session: - for seed, value in zip(seeds, values, strict=True): - session.execute( - text('UPDATE "SeedPromptEntries" SET harm_categories = :value WHERE id = :id'), - {"value": value, "id": str(seed.id)}, - ) - session.commit() - - page = _page(sqlite_instance, dataset_name=DATASET, harm_categories=["violence"], limit=10) - assert [member.id for member in _field(page, "items")[0].members] == [seeds[-1].id] - - async def test_logical_identity_uses_group_id_else_seed_id(self, sqlite_instance: MemoryInterface): - group_id = uuid4() - grouped = SeedPrompt(value="grouped", dataset_name=DATASET, prompt_group_id=group_id) - ungrouped = SeedPrompt(value="ungrouped", dataset_name=DATASET) - await _add(sqlite_instance, grouped, ungrouped) - - page = _page(sqlite_instance, dataset_name=DATASET, limit=10) - examples = _field(page, "items") - assert {str(_field(item, "example_id")) for item in examples} == {str(group_id), str(ungrouped.id)} - assert {str(_field(item, "seed_ids")[0]) for item in examples} == {str(grouped.id), str(ungrouped.id)} - - async def test_order_is_earliest_complete_date_then_id(self, sqlite_instance: MemoryInterface): - old_group = UUID("00000000-0000-0000-0000-000000000001") - new_group = UUID("00000000-0000-0000-0000-000000000002") - start = datetime(2024, 1, 1, tzinfo=UTC) - await _add( - sqlite_instance, - SeedPrompt(value="old", dataset_name=DATASET, prompt_group_id=old_group, date_added=start), - SeedPrompt( - value="new", dataset_name=DATASET, prompt_group_id=new_group, date_added=start + timedelta(days=1) - ), - SeedPrompt( - value="late member", - dataset_name=DATASET, - prompt_group_id=old_group, - date_added=start + timedelta(days=3), - ), - ) - ids = [ - _field(item, "example_id") - for item in _field(_page(sqlite_instance, dataset_name=DATASET, limit=10), "items") - ] - assert [str(value) for value in ids] == [str(new_group), str(old_group)] - - async def test_tied_timestamps_use_descending_logical_id(self, sqlite_instance: MemoryInterface): - timestamp = datetime(2024, 1, 1, tzinfo=UTC) - lower = UUID("00000000-0000-0000-0000-000000000001") - higher = UUID("00000000-0000-0000-0000-000000000002") - await _add( - sqlite_instance, - SeedPrompt(value="lower", dataset_name=DATASET, prompt_group_id=lower, date_added=timestamp), - SeedPrompt(value="higher", dataset_name=DATASET, prompt_group_id=higher, date_added=timestamp), - ) - ids = [ - _field(item, "example_id") - for item in _field(_page(sqlite_instance, dataset_name=DATASET, limit=10), "items") - ] - assert [str(value) for value in ids] == [str(higher), str(lower)] - - async def test_cursor_continuation_pages_logical_examples_not_seed_rows(self, sqlite_instance: MemoryInterface): - group_id = UUID("00000000-0000-0000-0000-000000000003") - second_group_id = UUID("00000000-0000-0000-0000-000000000002") - first_group_id = UUID("00000000-0000-0000-0000-000000000001") - ungrouped_id = UUID("00000000-0000-0000-0000-000000000004") - first_timestamp = datetime(2024, 1, 1, tzinfo=UTC) - tied_timestamp = datetime(2024, 1, 2, tzinfo=UTC) - group_three_first = SeedPrompt( - value="group three first", dataset_name=DATASET, prompt_group_id=group_id, date_added=tied_timestamp - ) - group_three_second = SeedPrompt( - value="group three second", dataset_name=DATASET, prompt_group_id=group_id, date_added=tied_timestamp - ) - group_two = SeedPrompt( - value="group two", dataset_name=DATASET, prompt_group_id=second_group_id, date_added=tied_timestamp - ) - group_one_early = SeedPrompt( - value="group one early", dataset_name=DATASET, prompt_group_id=first_group_id, date_added=first_timestamp - ) - group_one_late = SeedPrompt( - value="group one late", - dataset_name=DATASET, - prompt_group_id=first_group_id, - date_added=datetime(2024, 1, 3, tzinfo=UTC), - ) - ungrouped = SeedPrompt(value="ungrouped", dataset_name=DATASET, id=ungrouped_id, date_added=first_timestamp) - await _add( - sqlite_instance, - group_three_first, - group_three_second, - group_two, - group_one_early, - group_one_late, - ungrouped, - ) - pages: list[Any] = [] - cursor = None - while True: - page = _page(sqlite_instance, dataset_name=DATASET, limit=1, cursor=cursor) - pages.extend(_field(page, "items")) - cursor = _field(page, "next_cursor") - if cursor is None: - break - - assert [str(_field(item, "example_id")) for item in pages] == [ - str(group_id), - str(second_group_id), - str(ungrouped_id), - str(first_group_id), - ] - assert [len(_field(item, "members")) for item in pages] == [2, 1, 1, 2] - assert [{str(seed_id) for seed_id in _field(item, "seed_ids")} for item in pages] == [ - {str(group_three_first.id), str(group_three_second.id)}, - {str(group_two.id)}, - {str(ungrouped.id)}, - {str(group_one_early.id), str(group_one_late.id)}, - ] - assert len({_field(item, "example_id") for item in pages}) == len(pages) == 4 - - async def test_filters_are_member_or_and_example_and(self, sqlite_instance: MemoryInterface): - group_id = uuid4() - await _add( - sqlite_instance, - SeedPrompt(value="image", dataset_name=DATASET, prompt_group_id=group_id, data_type="text"), - SeedPrompt(value="violence", dataset_name=DATASET, prompt_group_id=group_id, harm_categories=["violence"]), - SeedPrompt(value="only image", dataset_name=DATASET, data_type="text"), - ) - page = _page( - sqlite_instance, - dataset_name=DATASET, - data_types=["text"], - harm_categories=["violence"], - limit=10, - ) - assert len(_field(page, "items")) == 1 - item = _field(page, "items")[0] - expected_ids = {str(seed.id) for seed in sqlite_instance.get_seeds(prompt_group_ids=[group_id])} - assert expected_ids == {str(seed_id) for seed_id in _field(item, "seed_ids")} - assert len(_field(item, "members")) == 2 - - async def test_modality_harm_and_seed_type_values_are_or(self, sqlite_instance: MemoryInterface, tmp_path): - modality_only = SeedPrompt(value="https://example.com/modality-only", dataset_name=DATASET, data_type="url") - modality_only_second = SeedPrompt( - value=str(tmp_path / "modality-only.png"), dataset_name=DATASET, data_type="image_path" - ) - (tmp_path / "modality-only.png").write_bytes(b"image") - harm_only = SeedPrompt(value="harm only", dataset_name=DATASET, data_type="reasoning", harm_categories=["hate"]) - seed_type_only = SeedSimulatedConversation( - dataset_name=DATASET, - adversarial_chat_system_prompt=SeedPrompt(value="adversarial"), - simulated_target_system_prompt=SeedPrompt(value="target"), - ) - all_filters_group = uuid4() - all_filters_prompt = SeedPrompt( - value="https://example.com/all-filters", - dataset_name=DATASET, - prompt_group_id=all_filters_group, - data_type="url", - harm_categories=["violence"], - ) - all_filters_objective = SeedObjective( - value="all filters objective", dataset_name=DATASET, prompt_group_id=all_filters_group - ) - no_filters = SeedPrompt( - value="no filters", dataset_name=DATASET, data_type="reasoning", harm_categories=["other"] - ) - await _add( - sqlite_instance, - modality_only, - modality_only_second, - harm_only, - seed_type_only, - all_filters_prompt, - all_filters_objective, - no_filters, - ) - modality_page = _page(sqlite_instance, dataset_name=DATASET, data_types=["url", "image_path"], limit=10) - harm_page = _page(sqlite_instance, dataset_name=DATASET, harm_categories=["hate", "violence"], limit=10) - seed_type_page = _page( - sqlite_instance, - dataset_name=DATASET, - seed_types=["objective", "simulated_conversation"], - limit=10, - ) - combined_page = _page( - sqlite_instance, - dataset_name=DATASET, - data_types=["url", "image_path"], - harm_categories=["hate", "violence"], - seed_types=["objective", "simulated_conversation"], - limit=10, - ) - - assert {str(_field(item, "example_id")) for item in _field(modality_page, "items")} == { - str(modality_only.id), - str(modality_only_second.id), - str(all_filters_group), - } - assert {str(_field(item, "example_id")) for item in _field(harm_page, "items")} == { - str(harm_only.id), - str(all_filters_group), - } - assert {str(_field(item, "example_id")) for item in _field(seed_type_page, "items")} == { - str(seed_type_only.id), - str(all_filters_group), - } - assert [str(_field(item, "example_id")) for item in _field(combined_page, "items")] == [str(all_filters_group)] - assert _field(combined_page, "total") == 1 - - async def test_filtered_order_uses_earliest_member_not_earliest_matching_member( - self, sqlite_instance: MemoryInterface - ): - group_a = uuid4() - group_b = uuid4() - t1 = datetime(2024, 1, 1, tzinfo=UTC) - t2 = datetime(2024, 1, 2, tzinfo=UTC) - t3 = datetime(2024, 1, 3, tzinfo=UTC) - await _add( - sqlite_instance, - SeedPrompt(value="a earliest", dataset_name=DATASET, prompt_group_id=group_a, date_added=t1), - SeedPrompt( - value="https://example.com/a-matching-later", - dataset_name=DATASET, - prompt_group_id=group_a, - date_added=t3, - data_type="url", - ), - SeedPrompt( - value="https://example.com/b-matching", - dataset_name=DATASET, - prompt_group_id=group_b, - date_added=t2, - data_type="url", - ), - ) - page = _page(sqlite_instance, dataset_name=DATASET, data_types=["url"], limit=10) - assert [str(_field(item, "example_id")) for item in _field(page, "items")] == [ - str(group_b), - str(group_a), - ] - - async def test_harm_matching_is_case_insensitive_whole_value_and_missing_is_unlabeled( - self, sqlite_instance: MemoryInterface - ): - labeled = SeedPrompt(value="labeled", dataset_name=DATASET, harm_categories=["VIOLENCE"]) - multi_labeled = SeedPrompt( - value="multi-labeled", dataset_name=DATASET, harm_categories=["hate", 'special_%_"_é'] - ) - substring_only = SeedPrompt(value="substring-only", dataset_name=DATASET, harm_categories=["nonviolence"]) - suffix_only = SeedPrompt(value="suffix-only", dataset_name=DATASET, harm_categories=["violence-extra"]) - unlabeled = SeedPrompt(value="unlabeled", dataset_name=DATASET, harm_categories=[]) - null_labeled = SeedPrompt(value="null-labeled", dataset_name=DATASET, harm_categories=None) - grouped = uuid4() - grouped_unmatched_member = SeedPrompt( - value="grouped-unmatched", dataset_name=DATASET, prompt_group_id=grouped, harm_categories=[] - ) - grouped_matching_member = SeedPrompt( - value="grouped-matching", dataset_name=DATASET, prompt_group_id=grouped, harm_categories=["violence"] - ) - await _add( - sqlite_instance, - labeled, - multi_labeled, - substring_only, - suffix_only, - unlabeled, - null_labeled, - grouped_unmatched_member, - grouped_matching_member, - ) - exact = _page(sqlite_instance, dataset_name=DATASET, harm_categories=["violence"], limit=10) - substring = _page(sqlite_instance, dataset_name=DATASET, harm_categories=["vio"], limit=10) - assert {str(item.example_id) for item in _field(exact, "items")} == { - str(labeled.id), - str(grouped), - } - multi = _page(sqlite_instance, dataset_name=DATASET, harm_categories=["missing", "HATE"], limit=10) - assert {str(item.example_id) for item in _field(multi, "items")} == {str(multi_labeled.id)} - special = _page(sqlite_instance, dataset_name=DATASET, harm_categories=['special_%_"_É'], limit=10) - assert {str(item.example_id) for item in _field(special, "items")} == {str(multi_labeled.id)} - assert len(_field(substring, "items")) == 0 - all_items = _field(_page(sqlite_instance, dataset_name=DATASET, limit=10), "items") - unlabeled_item = next( - item for item in all_items if str(unlabeled.id) in [str(i) for i in _field(item, "seed_ids")] - ) - assert _field(unlabeled_item, "has_unlabeled_harm") is True - - @pytest.mark.parametrize("text", ["100% literal", "literal_value"]) - async def test_text_search_is_case_insensitive_literal_and_text_only( - self, sqlite_instance: MemoryInterface, text: str, tmp_path - ): - media = tmp_path / "media-path.png" - media.write_bytes(b"local image") - await _add( - sqlite_instance, - SeedPrompt(value="Need 100% literal_value", dataset_name=DATASET), - SeedPrompt(value=str(media), dataset_name=DATASET, data_type="image_path"), - SeedPrompt(value="metadata", dataset_name=DATASET, metadata={"search": "metadata_only"}), - ) - page = _page(sqlite_instance, dataset_name=DATASET, value_search=text.upper(), limit=10) - assert len(_field(page, "items")) == 1 - assert ( - len(_field(_page(sqlite_instance, dataset_name=DATASET, value_search="metadata_only", limit=10), "items")) - == 0 - ) - assert ( - len(_field(_page(sqlite_instance, dataset_name=DATASET, value_search="media-path", limit=10), "items")) == 0 - ) - - async def test_percent_search_is_literal(self, sqlite_instance: MemoryInterface): - literal_match = SeedPrompt(value="contains 100%", dataset_name=DATASET) - wildcard_only = SeedPrompt(value="contains 1000", dataset_name=DATASET) - await _add(sqlite_instance, literal_match, wildcard_only) - page = _page(sqlite_instance, dataset_name=DATASET, value_search="100%", limit=10) - assert [str(_field(item, "seed_ids")[0]) for item in _field(page, "items")] == [str(literal_match.id)] - - async def test_underscore_search_is_literal(self, sqlite_instance: MemoryInterface): - literal_match = SeedPrompt(value="contains a_b", dataset_name=DATASET) - wildcard_only = SeedPrompt(value="contains acb", dataset_name=DATASET) - await _add(sqlite_instance, literal_match, wildcard_only) - page = _page(sqlite_instance, dataset_name=DATASET, value_search="a_b", limit=10) - assert [str(_field(item, "seed_ids")[0]) for item in _field(page, "items")] == [str(literal_match.id)] - - async def test_count_uses_same_logical_predicates_as_page(self, sqlite_instance: MemoryInterface): - group_id = uuid4() - await _add( - sqlite_instance, - SeedPrompt(value="matching", dataset_name=DATASET, prompt_group_id=group_id, data_type="text"), - SeedObjective(value="related", dataset_name=DATASET, prompt_group_id=group_id), - SeedPrompt(value="other", dataset_name=DATASET, data_type="text"), - ) - page = _page(sqlite_instance, dataset_name=DATASET, data_types=["text"], limit=10) - assert _field(page, "total") == 2 - grouped = next(item for item in _field(page, "items") if _field(item, "piece_count") == 2) - assert str(_field(grouped, "example_id")) == str(group_id) - - async def test_only_selected_examples_are_expanded_to_all_members(self, sqlite_instance: MemoryInterface): - selected = uuid4() - not_selected = uuid4() - await _add( - sqlite_instance, - SeedPrompt(value="selected", dataset_name=DATASET, prompt_group_id=selected, harm_categories=["hate"]), - SeedPrompt(value="selected related", dataset_name=DATASET, prompt_group_id=selected), - SeedPrompt(value="unselected", dataset_name=DATASET, prompt_group_id=not_selected), - ) - page = _page(sqlite_instance, dataset_name=DATASET, harm_categories=["hate"], limit=10) - assert len(_field(page, "items")) == 1 - assert len(_field(_field(page, "items")[0], "members")) == 2 - - async def test_query_is_database_bounded_and_not_n_plus_one(self, sqlite_instance: MemoryInterface): - groups = [uuid4() for _ in range(20)] - await _add( - sqlite_instance, - *( - SeedPrompt(value=f"seed-{index}", dataset_name=DATASET, prompt_group_id=group_id) - for index, group_id in enumerate(groups) - ), - *(SeedPrompt(value=f"extra-{index}", dataset_name=DATASET) for index in range(250)), - ) - statements: list[str] = [] - - def capture(_connection, _cursor, statement, _parameters, _context, _executemany): - statements.append(statement.lower()) - - event.listen(sqlite_instance.engine, "before_cursor_execute", capture) - try: - page = _page(sqlite_instance, dataset_name=DATASET, limit=20) - finally: - event.remove(sqlite_instance.engine, "before_cursor_execute", capture) - assert len(_field(page, "items")) == 20 - select_statements = [statement for statement in statements if statement.lstrip().startswith("select")] - grouped_page_queries = [ - statement for statement in select_statements if "group by" in statement and "order by" in statement - ] - assert grouped_page_queries - assert all(" limit " in f" {statement} " for statement in grouped_page_queries) - assert len(select_statements) <= 5 - - async def test_malformed_cursor_and_filter_mismatch_are_rejected(self, sqlite_instance: MemoryInterface): - await _add( - sqlite_instance, - SeedPrompt(value="one", dataset_name=DATASET), - SeedPrompt(value="two", dataset_name=DATASET, harm_categories=["violence"]), - ) - with pytest.raises(ValueError): - _page(sqlite_instance, dataset_name=DATASET, limit=1, cursor="malformed") - first = _page(sqlite_instance, dataset_name=DATASET, limit=1) - with pytest.raises(ValueError): - _page( - sqlite_instance, - dataset_name=DATASET, - harm_categories=["violence"], - limit=1, - cursor=_field(first, "next_cursor"), - ) - payload = json.loads( - base64.urlsafe_b64decode(_field(first, "next_cursor") + "=" * (-len(_field(first, "next_cursor")) % 4)) - ) - payload["i"] = "not-a-uuid" - malformed_identifier = base64.urlsafe_b64encode(json.dumps(payload).encode()).decode().rstrip("=") - with pytest.raises(ValueError): - _page(sqlite_instance, dataset_name=DATASET, limit=1, cursor=malformed_identifier) - - @pytest.mark.parametrize( - ("changed_filters", "changed_dataset"), - [ - ({"value_search": "changed"}, None), - ({"data_types": ["url"]}, None), - ({"harm_categories": ["violence"]}, None), - ({"seed_types": ["objective"]}, None), - ({}, "another-memory-dataset"), - ], - ) - async def test_cursor_is_bound_to_every_effective_filter( - self, - sqlite_instance: MemoryInterface, - changed_filters: dict[str, object], - changed_dataset: str | None, - ): - await _add( - sqlite_instance, - SeedPrompt(value="one", dataset_name=DATASET), - SeedPrompt( - value="https://example.com/two", dataset_name=DATASET, data_type="url", harm_categories=["violence"] - ), - ) - first = _page(sqlite_instance, dataset_name=DATASET, limit=1) - with pytest.raises(ValueError): - _page( - sqlite_instance, - dataset_name=changed_dataset or DATASET, - limit=1, - cursor=_field(first, "next_cursor"), - **changed_filters, - ) - - async def test_existing_get_seeds_harm_semantics_remain_all_categories(self, sqlite_instance: MemoryInterface): - await _add( - sqlite_instance, - SeedPrompt(value="one", harm_categories=["hate"]), - SeedPrompt(value="two", harm_categories=["hate", "violence"]), - ) - assert len(sqlite_instance.get_seeds(harm_categories=["hate", "violence"])) == 1 diff --git a/tests/unit/memory/memory_interface/test_interface_seed_browsing_sql.py b/tests/unit/memory/memory_interface/test_interface_seed_browsing_sql.py deleted file mode 100644 index 418b3fe77f..0000000000 --- a/tests/unit/memory/memory_interface/test_interface_seed_browsing_sql.py +++ /dev/null @@ -1,108 +0,0 @@ -# Copyright (c) Microsoft Corporation. -# Licensed under the MIT license. - -"""SQL dialect compilation coverage for the #2748 seed browsing seam.""" - -from __future__ import annotations - -from datetime import UTC, datetime -from uuid import uuid4 - -from sqlalchemy.dialects import mssql, sqlite - -from pyrit.common.pagination import encode_keyset_cursor -from pyrit.memory.azure_sql_memory import AzureSQLMemory -from pyrit.memory.memory_interface import ( - SeedExampleDatasetScope, - _build_seed_example_query, - _query_seed_example_members, - _query_seed_example_page, -) -from pyrit.memory.sqlite_memory import SQLiteMemory - - -class _Result: - def all(self): - return [] - - def scalar_one(self): - return 0 - - def scalars(self): - return [] - - -class _StatementCapture: - def __init__(self): - self.statements = [] - - def execute(self, statement): - self.statements.append(statement) - return _Result() - - -def _query(*, scope: SeedExampleDatasetScope, cursor: str | None = None): - return _build_seed_example_query( - dataset_scope=scope, - limit=100, - cursor=cursor, - data_types=["url", "image_path"], - harm_categories=["violence", "self_harm_%"], - seed_types=["prompt", "objective"], - value_search=r"literal_%\value", - ) - - -def _capture_statements(*, scope: SeedExampleDatasetScope, memory_type): - initial = _query(scope=scope) - cursor = encode_keyset_cursor( - timestamp=datetime(2024, 1, 2, tzinfo=UTC), - identifier=str(uuid4()), - fingerprint=initial.fingerprint, - ) - query = _query(scope=scope, cursor=cursor) - capture = _StatementCapture() - _query_seed_example_page( - session=capture, - query=query, - harm_condition_builder=object.__new__(memory_type)._seed_example_harm_condition, - ) - _query_seed_example_members( - session=capture, - query=query, - example_ids=[uuid4() for _ in range(100)], - ) - return capture.statements - - -def test_seed_browsing_statements_compile_for_sql_server_and_sqlite(): - for scope in (SeedExampleDatasetScope.named("dataset"), SeedExampleDatasetScope.unnamed()): - sql_server_statements = _capture_statements(scope=scope, memory_type=AzureSQLMemory) - sqlite_statements = _capture_statements(scope=scope, memory_type=SQLiteMemory) - sql_server_compiled = [ - statement.compile(dialect=mssql.dialect(), compile_kwargs={"render_postcompile": True}) - for statement in sql_server_statements - ] - sql_server_sql = [str(statement).lower() for statement in sql_server_compiled] - sqlite_sql = [str(statement.compile(dialect=sqlite.dialect())).lower() for statement in sqlite_statements] - - assert all(sql for sql in sql_server_sql) - assert all(sql for sql in sqlite_sql) - assert all(len(statement.params) < 2100 for statement in sql_server_compiled) - assert any("openjson" in sql for sql in sql_server_sql) - assert any("json_each" in sql for sql in sqlite_sql) - assert any("group by" in sql and "min" in sql for sql in sql_server_sql) - assert any("group by" in sql and "min" in sql for sql in sqlite_sql) - assert any("order by" in sql and "offset" in sql or "top" in sql for sql in sql_server_sql) - assert any("order by" in sql and "limit" in sql for sql in sqlite_sql) - assert any("sequence" in sql and "case" in sql for sql in sql_server_sql) - page_sql = sql_server_sql[0] - assert "lower(cast(coalesce" in page_sql - assert "example_id_key <" in page_sql - assert "example_id_key desc" in page_sql - assert "openjson(case when" in page_sql - assert "isjson([seedpromptentries_" in page_sql - assert "left(ltrim([seedpromptentries_" in page_sql - assert "else :" in page_sql - assert "openjson([seedpromptentries_" not in page_sql - assert "openjson(json_query(" not in page_sql 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..ceaccfcad4 --- /dev/null +++ b/tests/unit/memory/memory_interface/test_interface_seed_examples.py @@ -0,0 +1,230 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +import json +import logging +from datetime import UTC, datetime, timedelta +from pathlib import Path +from uuid import UUID, uuid4 + +import pytest + +from pyrit.memory import MemoryInterface +from pyrit.memory.memory_models import SeedEntry +from pyrit.models import Seed, SeedObjective, SeedPrompt, 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, tied_high = sorted([uuid4(), uuid4()]) + 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_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], SeedObjective) + + +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(*, missing_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(missing_file)}, + sort_keys=True, + separators=(",", ":"), + ) + return entry + + +async def test_get_seed_examples_skips_seeds_that_cannot_be_read( + sqlite_instance: MemoryInterface, tmp_path: Path, caplog: pytest.LogCaptureFixture +): + mixed, broken = uuid4(), uuid4() + prompt = SeedPrompt(value="readable", dataset_name=DATASET, prompt_group_id=mixed) + await _add(sqlite_instance, prompt) + mixed_entry = _legacy_conversation_entry(missing_file=tmp_path / "gone.yaml", prompt_group_id=mixed) + broken_entry = _legacy_conversation_entry(missing_file=tmp_path / "gone.yaml", prompt_group_id=broken) + skipped_ids = [str(mixed_entry.id), str(broken_entry.id)] + await _store(sqlite_instance, mixed_entry) + await _store(sqlite_instance, broken_entry) + + with ( + caplog.at_level(logging.WARNING, logger="pyrit.memory.memory_interface"), + pytest.warns(DeprecationWarning, match="adversarial_chat_system_prompt_path"), + ): + examples, total, _ = await sqlite_instance.get_seed_examples_async(dataset_name=DATASET, limit=10) + + assert [seed.id for seed in examples[mixed]] == [prompt.id] + assert broken not in examples + assert total == 2 + messages = [record.getMessage() for record in caplog.records] + assert all(any(seed_id in message for message in messages) for seed_id in skipped_ids) + + +async def test_get_seed_examples_returns_domain_simulated_conversation(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], SeedSimulatedConversation) + assert seeds[0].adversarial_chat_system_prompt.value == "adversarial" + + +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 f6b68fac5c..297f611f89 100644 --- a/tests/unit/memory/test_memory_models.py +++ b/tests/unit/memory/test_memory_models.py @@ -545,32 +545,6 @@ def test_roundtrip_seed_objective(self): assert isinstance(recovered, SeedObjective) assert recovered.value == "objective text" - @pytest.mark.parametrize("seed_kind", ["prompt", "objective", "simulated_conversation"]) - def test_get_seed_normalizes_historical_null_template_flag(self, seed_kind: str): - if seed_kind == "prompt": - seed = _make_seed_prompt() - elif seed_kind == "objective": - seed = SeedObjective(value="objective text", dataset_name="ds", added_by="tester") - else: - seed = SeedSimulatedConversation( - adversarial_chat_system_prompt=SeedPrompt(value="adversarial"), - pyrit_version="1.0.0", - ) - - entry = SeedEntry(entry=seed) - entry.is_jinja_template = None - - recovered = entry.get_seed() - assert recovered.is_jinja_template is False - - @pytest.mark.parametrize("is_template", [True, False]) - def test_get_seed_preserves_persisted_template_flag(self, is_template: bool): - entry = SeedEntry(entry=_make_seed_prompt(is_jinja_template=is_template)) - entry.is_jinja_template = is_template - - recovered = entry.get_seed() - assert recovered.is_jinja_template is is_template - def test_seed_prompt_preserves_parameters(self): seed = _make_seed_prompt(parameters=["param1", "param2"]) entry = SeedEntry(entry=seed) diff --git a/tests/unit/memory/test_migration.py b/tests/unit/memory/test_migration.py index 25fe49cdb4..a81a631d09 100644 --- a/tests/unit/memory/test_migration.py +++ b/tests/unit/memory/test_migration.py @@ -178,32 +178,6 @@ def test_run_schema_migrations_applies_head_revision(): engine.dispose() -def test_seed_template_flag_migration_lifecycle(): - """The seed template marker is added and removed through the normal Alembic lifecycle.""" - with tempfile.TemporaryDirectory() as temp_dir: - db_path = os.path.join(temp_dir, "seed-template-flag.db") - engine = create_engine(f"sqlite:///{db_path}") - try: - with engine.begin() as connection: - config = _config_for(connection) - command.upgrade(config, "7a9c1e3f5b2d") - assert "is_jinja_template" not in { - column["name"] for column in inspect(connection).get_columns("SeedPromptEntries") - } - - command.upgrade(config, "head") - assert "is_jinja_template" in { - column["name"] for column in inspect(connection).get_columns("SeedPromptEntries") - } - - command.downgrade(config, "7a9c1e3f5b2d") - assert "is_jinja_template" not in { - column["name"] for column in inspect(connection).get_columns("SeedPromptEntries") - } - finally: - engine.dispose() - - @pytest.mark.parametrize("starting_revision", ["9b2d4f6a8c0e", "fcecd0617e61"]) def test_seed_conditions_and_follow_up_template_migrations_merge(starting_revision: str) -> None: engine = create_engine("sqlite:///:memory:") @@ -218,9 +192,7 @@ def test_seed_conditions_and_follow_up_template_migrations_merge(starting_revisi with engine.connect() as connection: version = connection.execute(text("SELECT version_num FROM pyrit_memory_alembic_version")).scalar_one() assert version == _get_alembic_head_revision(config=config) - assert "conditions" in { - column["name"] for column in inspect(connection).get_columns("SeedPromptEntries") - } + assert "conditions" in {column["name"] for column in inspect(connection).get_columns("SeedPromptEntries")} assert "adversarial_prompt_template" in { column["name"] for column in inspect(connection).get_columns("AttackIdentifiers") } From fe7b9112bd1eeb9da92dc183ea7901cc361ef79e Mon Sep 17 00:00:00 2001 From: Roman Lutz Date: Thu, 8 Oct 2026 13:21:06 -0700 Subject: [PATCH 8/9] Fix seed browsing fidelity, previews, and UUID ordering Return stored seed records without loading legacy template files or dropping group members. Mask standalone text references in previews and use canonical UUID ordering for cross-database paging. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- doc/code/datasets/0_dataset.md | 22 +-- pyrit/backend/models/datasets.py | 6 +- pyrit/backend/services/dataset_service.py | 24 ++-- pyrit/memory/memory_interface.py | 66 ++++----- pyrit/memory/memory_models.py | 34 +++++ pyrit/models/__init__.py | 2 + pyrit/models/seeds/__init__.py | 2 + pyrit/models/seeds/seed_record.py | 42 ++++++ ...est_seed_examples_azure_sql_integration.py | 13 +- .../unit/backend/test_seed_example_routes.py | 85 +++++++++++- .../unit/common/test_lazy_package_imports.py | 6 + .../test_interface_seed_examples.py | 128 ++++++++++++++---- tests/unit/memory/test_memory_models.py | 27 ++++ 13 files changed, 376 insertions(+), 81 deletions(-) create mode 100644 pyrit/models/seeds/seed_record.py diff --git a/doc/code/datasets/0_dataset.md b/doc/code/datasets/0_dataset.md index 4f4bb26aff..b22c702627 100644 --- a/doc/code/datasets/0_dataset.md +++ b/doc/code/datasets/0_dataset.md @@ -97,14 +97,17 @@ preserve the recorded origin. ## Browse stored seeds The seed browser reads stored seeds from memory only. It does not load providers, open -media files, render templates, or generate conversations. A simulated-conversation seed -saved by an older PyRIT version stores prompt file paths, so the browser loads those files. -If a seed cannot be read, for example because a file is missing, the browser skips it and -logs a warning. An example with no readable seeds is not shown, but `total` counts it. +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 seed objects (`SeedPrompt`, `SeedObjective`, or `SeedSimulatedConversation`). + 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 @@ -119,10 +122,13 @@ 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 example ID. A +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, and -other types show a type label. The browser does not render templates or run +and `preview_truncated` when it is shortened. Media members show only the file name. +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/models/datasets.py b/pyrit/backend/models/datasets.py index e411ea0f62..f3566c340a 100644 --- a/pyrit/backend/models/datasets.py +++ b/pyrit/backend/models/datasets.py @@ -14,7 +14,7 @@ from pydantic import BaseModel, Field from pyrit.backend.models.common import PaginationInfo -from pyrit.models import PromptDataType, SeedType, SeedUnion +from pyrit.models import PromptDataType, SeedRecord, SeedType class DatasetInfo(BaseModel): @@ -74,4 +74,6 @@ class SeedExampleListResponse(BaseModel): class SeedExampleDetailResponse(SeedExampleSummary): """One logical seed example with all of its stored seeds.""" - members: list[SeedUnion] = Field(..., description="Stored seeds, objectives first, then by sequence") + members: list[SeedRecord] = Field( + ..., description="Stored seed records without reconstruction, objectives first, then by sequence" + ) diff --git a/pyrit/backend/services/dataset_service.py b/pyrit/backend/services/dataset_service.py index b7fa79f877..84920ffb85 100644 --- a/pyrit/backend/services/dataset_service.py +++ b/pyrit/backend/services/dataset_service.py @@ -8,6 +8,9 @@ """ import logging +import ntpath +import posixpath +import re from collections.abc import Sequence from functools import lru_cache from uuid import UUID @@ -29,10 +32,8 @@ ConversationStats, PromptDataType, SeedDatasetSummary, - SeedObjective, - SeedSimulatedConversation, + SeedRecord, SeedType, - SeedUnion, ) logger = logging.getLogger(__name__) @@ -222,36 +223,41 @@ def _parse_selection_key(selection_key: str) -> str | None: return name @classmethod - def _summarize(cls, *, example_id: UUID, seeds: Sequence[SeedUnion]) -> SeedExampleSummary: + 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 if seed.data_type}), + 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(isinstance(seed, SeedObjective) for seed in 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[SeedUnion]) -> tuple[str, bool]: + 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, so the preview does not expose paths or URL credentials. + 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 not isinstance(seed, SeedSimulatedConversation)] + 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) diff --git a/pyrit/memory/memory_interface.py b/pyrit/memory/memory_interface.py index 92057d8368..597c161b69 100644 --- a/pyrit/memory/memory_interface.py +++ b/pyrit/memory/memory_interface.py @@ -20,7 +20,8 @@ from typing import TYPE_CHECKING, Any, ClassVar, Literal, NamedTuple, ParamSpec, TypeVar, cast from urllib.parse import urlparse -from sqlalchemy import MetaData, and_, case, exists, false, func, literal, not_, or_, select, update +from sqlalchemy import MetaData, String, and_, case, exists, false, func, literal, not_, or_, select, 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 @@ -106,8 +107,8 @@ SeedObjective, SeedOrigin, SeedPrompt, + SeedRecord, SeedType, - SeedUnion, TargetIdentifier, group_conversation_message_pieces_by_sequence, sort_message_pieces, @@ -4362,29 +4363,27 @@ def _seed_example_filters( @staticmethod def _get_seed_example_seeds( *, session: Session, scope: "ColumnElement[bool]", example_ids: Sequence[uuid.UUID] - ) -> dict[uuid.UUID, list[SeedUnion]]: + ) -> dict[uuid.UUID, list[SeedRecord]]: """ - Read the stored seeds of the given logical examples. - - A seed that cannot be rebuilt is skipped with a warning, so one bad row does not stop the read. - For example, a simulated conversation saved with prompt file paths fails when a file is missing. + Read all stored members without reconstructing seeds or resolving configuration paths. Returns: - dict[uuid.UUID, list[SeedUnion]]: The seeds of each example that has readable seeds, in the + 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, SeedEntry.id) + .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[SeedUnion]] = {example_id: [] for example_id in example_ids} + seeds: dict[uuid.UUID, list[SeedRecord]] = {example_id: [] for example_id in example_ids} for entry in entries: - try: - seeds[entry.prompt_group_id or entry.id].append(entry.get_seed()) - except ValueError as e: - logger.warning(f"Skipping stored seed {entry.id} because it cannot be read: {e}") + 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( @@ -4397,16 +4396,17 @@ def _execute_get_seed_examples( harm_categories: Sequence[str] | None, seed_types: Sequence[SeedType] | None, value_search: str | None, - ) -> tuple[dict[uuid.UUID, list[SeedUnion]], int, DecodedKeysetCursor | 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[SeedUnion]], int, DecodedKeysetCursor | None]: The seeds of each + 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, @@ -4416,21 +4416,25 @@ def _execute_get_seed_examples( value_search=value_search, ) grouped = ( - select(logical_id.label("example_id"), func.min(SeedEntry.date_added).label("first_added")) + 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) + .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 = uuid.UUID(after.identifier) + 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 < anchor_id), + 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.desc()).limit(limit + 1) + 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() @@ -4444,12 +4448,12 @@ def _execute_get_seed_examples( 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[SeedUnion]: + def _execute_get_seed_example(self, *, dataset_name: str | None, example_id: uuid.UUID) -> list[SeedRecord]: """ Read one complete logical seed example. Returns: - list[SeedUnion]: The seeds of the example. The list is empty if the dataset does not contain it. + 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( @@ -8593,16 +8597,16 @@ async def get_seed_examples_async( harm_categories: Sequence[str] | None = None, seed_types: Sequence[SeedType] | None = None, value_search: str | None = None, - ) -> tuple[dict[uuid.UUID, list[SeedUnion]], int, DecodedKeysetCursor | 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. Values inside one filter use OR, different filters use AND, and any seed - of an example can satisfy a filter. This method does not render templates or load media. - A simulated conversation saved with prompt file paths loads those files. If it cannot be - read, it is skipped with a warning. + 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. @@ -8616,7 +8620,7 @@ async def get_seed_examples_async( text, ignoring case. Simulated-conversation configurations are not searched. Returns: - tuple[dict[uuid.UUID, list[SeedUnion]], int, DecodedKeysetCursor | None]: The seeds of each + 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. """ @@ -8631,7 +8635,7 @@ async def get_seed_examples_async( value_search=value_search, ) - async def get_seed_example_async(self, *, dataset_name: str | None, example_id: uuid.UUID) -> list[SeedUnion]: + 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. @@ -8642,7 +8646,7 @@ async def get_seed_example_async(self, *, dataset_name: str | None, example_id: example_id: The ``prompt_group_id`` of the example, or the seed ID of a seed without a group. Returns: - list[SeedUnion]: The seeds of the example, objectives first. The list is empty if the + 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( diff --git a/pyrit/memory/memory_models.py b/pyrit/memory/memory_models.py index 27f3755720..73046186cd 100644 --- a/pyrit/memory/memory_models.py +++ b/pyrit/memory/memory_models.py @@ -70,6 +70,7 @@ SeedObjective, SeedOrigin, SeedPrompt, + SeedRecord, SeedSimulatedConversation, SeedType, TargetIdentifier, @@ -1644,6 +1645,39 @@ def _unpack_seed_metadata( decoded = None return cleaned, decoded + 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. diff --git a/pyrit/models/__init__.py b/pyrit/models/__init__.py index a1754fc22e..e1a99d25d0 100644 --- a/pyrit/models/__init__.py +++ b/pyrit/models/__init__.py @@ -214,6 +214,7 @@ SeedObjective, SeedOrigin, SeedPrompt, + SeedRecord, SeedSimulatedConversation, SeedUnion, SimulatedTargetSystemPromptPaths, @@ -418,6 +419,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 index 6c50f1529a..b13f2e7392 100644 --- a/tests/integration/memory/test_seed_examples_azure_sql_integration.py +++ b/tests/integration/memory/test_seed_examples_azure_sql_integration.py @@ -4,7 +4,7 @@ """Azure SQL execution of the seed example read queries.""" from datetime import UTC, datetime, timedelta -from uuid import uuid4 +from uuid import UUID, uuid4 import pytest from sqlalchemy import delete @@ -18,7 +18,9 @@ async def test_seed_examples_on_azure_sql(azuresql_instance: AzureSQLMemory): test_id = str(uuid4()) dataset = f"2748-azure-{test_id}" - group, newer, older = uuid4(), uuid4(), uuid4() + 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), @@ -44,16 +46,19 @@ async def ids(**filters) -> list: return list(examples) try: - assert await ids() == [newer, *sorted([group, older], reverse=True)] + 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=1) + 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() diff --git a/tests/unit/backend/test_seed_example_routes.py b/tests/unit/backend/test_seed_example_routes.py index 4dd90a256c..927e5ea08f 100644 --- a/tests/unit/backend/test_seed_example_routes.py +++ b/tests/unit/backend/test_seed_example_routes.py @@ -1,7 +1,10 @@ # 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 @@ -9,6 +12,7 @@ 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, Seed, SeedObjective, SeedPrompt, SeedSimulatedConversation @@ -104,6 +108,45 @@ async def test_list_seed_examples_builds_safe_previews(client: AsyncClient, sqli assert items[str(configuration.id)]["preview"] == "[Simulated conversation configuration]" +@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) @@ -148,7 +191,7 @@ async def test_list_seed_examples_returns_empty_page(client: AsyncClient, sqlite assert (body["pagination"]["has_more"], body["pagination"]["next_cursor"]) == (False, None) -async def test_get_seed_example_returns_all_members_as_domain_seeds(client: AsyncClient, sqlite_instance): +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" @@ -166,7 +209,43 @@ async def test_get_seed_example_returns_all_members_as_domain_seeds(client: Asyn 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)]["num_turns"] == 2 - assert members[str(configuration.id)]["adversarial_chat_system_prompt"]["value"] == "adversarial" + 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 8caed73637..fb448c8c8a 100644 --- a/tests/unit/common/test_lazy_package_imports.py +++ b/tests/unit/common/test_lazy_package_imports.py @@ -34,6 +34,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 index ceaccfcad4..cadce48a55 100644 --- a/tests/unit/memory/memory_interface/test_interface_seed_examples.py +++ b/tests/unit/memory/memory_interface/test_interface_seed_examples.py @@ -2,16 +2,22 @@ # Licensed under the MIT license. import json -import logging 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, SeedSimulatedConversation +from pyrit.models import Seed, SeedObjective, SeedPrompt, SeedRecord, SeedSimulatedConversation DATASET = "browse" T0 = datetime(2024, 1, 1, tzinfo=UTC) @@ -33,7 +39,8 @@ async def _ids(memory: MemoryInterface, dataset_name: str | None = DATASET, **fi async def test_get_seed_examples_orders_by_first_date_then_id_and_seeks(sqlite_instance: MemoryInterface): - tied_low, tied_high = sorted([uuid4(), uuid4()]) + tied_low = UUID("00000000-0000-0000-0000-000000000001") + tied_high = UUID("ffffffff-ffff-ffff-ffff-000000000000") newest, oldest = uuid4(), uuid4() await _add( sqlite_instance, @@ -56,6 +63,45 @@ async def test_get_seed_examples_orders_by_first_date_then_id_and_seeks(sqlite_i 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, ): @@ -83,7 +129,8 @@ async def test_get_seed_examples_returns_complete_groups_objective_first(sqlite_ assert total == 1 assert [seed.id for seed in examples[group]] == [objective.id, first.id, second.id] - assert isinstance(examples[group][0], SeedObjective) + 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): @@ -158,7 +205,7 @@ async def test_get_seed_examples_search_ignores_simulated_conversation_json(sqli assert set(await _ids(sqlite_instance, seed_types=["simulated_conversation"])) == {standalone.id, group} -def _legacy_conversation_entry(*, missing_file: Path, prompt_group_id: UUID) -> SeedEntry: +def _legacy_conversation_entry(*, prompt_file: Path, prompt_group_id: UUID) -> SeedEntry: entry = SeedEntry( entry=SeedSimulatedConversation( num_turns=2, @@ -169,39 +216,52 @@ def _legacy_conversation_entry(*, missing_file: Path, prompt_group_id: UUID) -> ) ) entry.value = json.dumps( - {"num_turns": 2, "sequence": 0, "adversarial_chat_system_prompt_path": str(missing_file)}, + {"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 -async def test_get_seed_examples_skips_seeds_that_cannot_be_read( - sqlite_instance: MemoryInterface, tmp_path: Path, caplog: pytest.LogCaptureFixture -): - mixed, broken = uuid4(), uuid4() +@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(missing_file=tmp_path / "gone.yaml", prompt_group_id=mixed) - broken_entry = _legacy_conversation_entry(missing_file=tmp_path / "gone.yaml", prompt_group_id=broken) - skipped_ids = [str(mixed_entry.id), str(broken_entry.id)] + 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, broken_entry) + await _store(sqlite_instance, standalone_entry) with ( - caplog.at_level(logging.WARNING, logger="pyrit.memory.memory_interface"), - pytest.warns(DeprecationWarning, match="adversarial_chat_system_prompt_path"), + 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) + 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] - assert broken not in examples + 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 - messages = [record.getMessage() for record in caplog.records] - assert all(any(seed_id in message for message in messages) for seed_id in skipped_ids) + 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_returns_domain_simulated_conversation(sqlite_instance: MemoryInterface): +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"]), @@ -213,8 +273,28 @@ async def test_get_seed_examples_returns_domain_simulated_conversation(sqlite_in seeds = await sqlite_instance.get_seed_example_async(dataset_name=DATASET, example_id=conversation.id) assert len(seeds) == 1 - assert isinstance(seeds[0], SeedSimulatedConversation) - assert seeds[0].adversarial_chat_system_prompt.value == "adversarial" + 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): 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): From 0410ac39e990c465551d37206f8aceb1cc31603d Mon Sep 17 00:00:00 2001 From: Roman Lutz Date: Thu, 8 Oct 2026 16:22:11 -0700 Subject: [PATCH 9/9] Fix credential exposure in normalized media previews Recognize HTTP(S) and data URI schemes case-insensitively after leading whitespace. Derive media labels only from URL paths while preserving stored values and local filenames. Cover the shared formatter and dataset list/detail routes. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- doc/code/datasets/0_dataset.md | 2 + pyrit/backend/mappers/_preview.py | 13 ++++--- tests/unit/backend/test_preview.py | 35 ++++++++++++++++++ .../unit/backend/test_seed_example_routes.py | 37 ++++++++++++++++++- 4 files changed, 81 insertions(+), 6 deletions(-) diff --git a/doc/code/datasets/0_dataset.md b/doc/code/datasets/0_dataset.md index b22c702627..2384e9f5e4 100644 --- a/doc/code/datasets/0_dataset.md +++ b/doc/code/datasets/0_dataset.md @@ -128,6 +128,8 @@ cursor is valid only for the same `selection_key` and filters. Other cursors ret 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 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/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 index 927e5ea08f..0f39b7a9e8 100644 --- a/tests/unit/backend/test_seed_example_routes.py +++ b/tests/unit/backend/test_seed_example_routes.py @@ -15,7 +15,15 @@ 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, Seed, SeedObjective, SeedPrompt, SeedSimulatedConversation +from pyrit.models import ( + AnswerMatches, + MatchesObjective, + PromptDataType, + Seed, + SeedObjective, + SeedPrompt, + SeedSimulatedConversation, +) URL = "/api/datasets/seeds" DATASET = "browse" @@ -108,6 +116,33 @@ async def test_list_seed_examples_builds_safe_previews(client: AsyncClient, sqli 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", [