Skip to content

Commit ac714b1

Browse files
[FIX] add tests and fix bugs (microsoft#1152)
1 parent bee8e61 commit ac714b1

7 files changed

Lines changed: 51 additions & 8 deletions

File tree

.pre-commit-config.yaml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -55,7 +55,7 @@ repos:
5555
- id: flake8
5656
additional_dependencies: ['flake8-copyright']
5757
types: [python]
58-
exclude: (doc/|.github/|pyrit/prompt_converter/morse_converter.py|tests/unit/converter/test_prompt_converter.py|pyrit/prompt_converter/emoji_converter.py|tests/unit/models/test_seed_prompt.py|tests/unit/converter/test_unicode_confusable_converter.py|tests/unit/converter/test_first_letter_converter.py|tests/unit/converter/test_base2048_converter.py|tests/unit/converter/test_ecoji_converter.py|tests/unit/converter/test_bin_ascii_converter.py|pyrit/scenarios/printer/console_printer.py)
58+
exclude: (doc/|.github/|pyrit/prompt_converter/morse_converter.py|tests/unit/converter/test_prompt_converter.py|pyrit/prompt_converter/emoji_converter.py|tests/unit/models/test_seed.py|tests/unit/converter/test_unicode_confusable_converter.py|tests/unit/converter/test_first_letter_converter.py|tests/unit/converter/test_base2048_converter.py|tests/unit/converter/test_ecoji_converter.py|tests/unit/converter/test_bin_ascii_converter.py|pyrit/scenarios/printer/console_printer.py)
5959

6060
- repo: local
6161
hooks:

doc/code/memory/11_harm_categories.ipynb

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -29,10 +29,10 @@
2929
"source": [
3030
"import pathlib\n",
3131
"\n",
32-
"from pyrit.common.initialization import initialize_pyrit\n",
3332
"from pyrit.common.path import DATASETS_PATH\n",
3433
"from pyrit.memory.central_memory import CentralMemory\n",
3534
"from pyrit.models import SeedDataset\n",
35+
"from pyrit.setup.initialization import initialize_pyrit\n",
3636
"\n",
3737
"initialize_pyrit(memory_db_type=\"InMemory\")\n",
3838
"\n",

doc/code/memory/11_harm_categories.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -21,10 +21,10 @@
2121
# %%
2222
import pathlib
2323

24-
from pyrit.common.initialization import initialize_pyrit
2524
from pyrit.common.path import DATASETS_PATH
2625
from pyrit.memory.central_memory import CentralMemory
2726
from pyrit.models import SeedDataset
27+
from pyrit.setup.initialization import initialize_pyrit
2828

2929
initialize_pyrit(memory_db_type="InMemory")
3030

pyrit/memory/memory_interface.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -856,6 +856,7 @@ async def add_seed_groups_to_memory(
856856

857857
all_prompts.extend(prompt_group.prompts)
858858
if prompt_group.objective:
859+
prompt_group.objective.prompt_group_id = prompt_group_id
859860
all_prompts.append(prompt_group.objective)
860861
await self.add_seeds_to_memory_async(prompts=all_prompts, added_by=added_by)
861862

pyrit/models/seed_dataset.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -97,7 +97,7 @@ def __init__(
9797
value=p["value"],
9898
data_type="text",
9999
value_sha256=p.get("value_sha256"),
100-
id=p.get("id"),
100+
id=uuid.uuid4(),
101101
name=p.get("name"),
102102
dataset_name=p.get("dataset_name"),
103103
harm_categories=p.get("harm_categories", []),

tests/unit/memory/memory_interface/memory_interface.py

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -39,7 +39,6 @@
3939
Score,
4040
SeedDataset,
4141
SeedGroup,
42-
SeedPrompt,
4342
StorageIO,
4443
data_serializer_factory,
4544
group_conversation_message_pieces_by_sequence,
@@ -816,7 +815,7 @@ async def add_seed_groups_to_memory(
816815
raise ValueError("At least one prompt group must be provided.")
817816
# Validates the prompt group IDs and sets them if possible before leveraging
818817
# the add_seed_prompts_to_memory method.
819-
all_prompts: MutableSequence[SeedPrompt] = []
818+
all_prompts: MutableSequence[Seed] = []
820819
for prompt_group in prompt_groups:
821820
if not prompt_group.prompts:
822821
raise ValueError("Prompt group must have at least one prompt.")
@@ -832,6 +831,9 @@ async def add_seed_groups_to_memory(
832831
prompt_group_id = group_id_set.pop() or uuid.uuid4()
833832
for prompt in prompt_group.prompts:
834833
prompt.prompt_group_id = prompt_group_id
834+
if prompt_group.objective:
835+
prompt_group.objective.prompt_group_id = prompt_group_id
836+
all_prompts.append(prompt_group.objective)
835837
all_prompts.extend(prompt_group.prompts)
836838
await self.add_seeds_to_memory_async(prompts=all_prompts, added_by=added_by)
837839

Lines changed: 42 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -13,8 +13,7 @@
1313
from scipy.io import wavfile
1414

1515
from pyrit.common.path import DATASETS_PATH
16-
from pyrit.models import SeedDataset, SeedGroup, SeedPrompt
17-
from pyrit.models.seed_objective import SeedObjective
16+
from pyrit.models import SeedDataset, SeedGroup, SeedObjective, SeedPrompt
1817

1918

2019
@pytest.fixture
@@ -37,6 +36,30 @@ def seed_prompt_fixture():
3736
)
3837

3938

39+
@pytest.fixture
40+
def seed_objective_fixture():
41+
return SeedObjective(
42+
value="Test objective",
43+
data_type="text",
44+
name="Test Name",
45+
dataset_name="Test Dataset",
46+
harm_categories=["category1", "category2"],
47+
description="Test Description",
48+
authors=["Author1"],
49+
groups=["Group1"],
50+
source="Test Source",
51+
added_by="Tester",
52+
metadata={"key": "value"},
53+
prompt_group_id=uuid.uuid4(),
54+
)
55+
56+
57+
def test_seed_objective_initialization(seed_objective_fixture):
58+
assert isinstance(seed_objective_fixture.id, uuid.UUID)
59+
assert seed_objective_fixture.value == "Test objective"
60+
assert seed_objective_fixture.data_type == "text"
61+
62+
4063
def test_seed_prompt_initialization(seed_prompt_fixture):
4164
assert isinstance(seed_prompt_fixture.id, uuid.UUID)
4265
assert seed_prompt_fixture.value == "Test prompt"
@@ -132,6 +155,23 @@ def test_seed_dataset_initialization(seed_prompt_fixture):
132155
assert dataset.prompts[0].value == "Test prompt"
133156

134157

158+
def test_seed_dataset_initialization_with_objective():
159+
seed_objective = SeedObjective(value="Test objective")
160+
dataset = SeedDataset(prompts=[seed_objective])
161+
assert len(dataset.objectives) == 1
162+
assert dataset.objectives[0].value == "Test objective"
163+
164+
165+
def test_seed_dataset_initialization_with_yaml_objective():
166+
dataset = SeedDataset.from_yaml_file(
167+
pathlib.Path(DATASETS_PATH) / "seed_prompts" / "illegal-multimodal-objective-group.prompt"
168+
)
169+
assert len(dataset.objectives) == 1
170+
assert dataset.objectives[0].value == "Use this image as inspiration"
171+
172+
assert len(dataset.prompts) == 3
173+
174+
135175
def test_seed_dataset_get_values():
136176
dataset = SeedDataset.from_yaml_file(pathlib.Path(DATASETS_PATH) / "seed_prompts" / "illegal.prompt")
137177
values = dataset.get_values()

0 commit comments

Comments
 (0)