Skip to content

Commit 7d80d95

Browse files
u7k4rs6romanlutzhannahwestra25
authored
FIX Seed removal list filters match whole elements instead of JSON substrings (#2811)
Co-authored-by: Roman Lutz <romanlutz13@gmail.com> Co-authored-by: hannahwestra25 <hannahwestra@microsoft.com>
1 parent d4d58c8 commit 7d80d95

4 files changed

Lines changed: 109 additions & 23 deletions

File tree

‎doc/code/memory/8_seed_database.ipynb‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -286,7 +286,7 @@
286286
"\n",
287287
"For the most part these are user errors, but when you want to remove whole groups rather than individual seeds, use `remove_seed_groups_from_memory`. It applies the same filters, but removes every seed that shares a `prompt_group_id` with any match, so groups are never left partial. Note that it only affects seeds that belong to a group: a matching seed added individually (with no `prompt_group_id`) is skipped, so use `remove_seeds_from_memory` for those.\n",
288288
"\n",
289-
"> **Note on deleting by `value`.** For the remove methods, the `value` filter defaults to full-string equality (`exact=True`), so `remove_seeds_from_memory(value=\"the\")` deletes only seeds whose value is exactly `\"the\"` — not everything containing it. This differs from `get_seeds`, which always matches `value` by substring. Pass `exact=False` to opt into substring deletion when you really want it. As a general rule, preview with the same filters via `get_seeds(...)` first and prefer a specific filter (such as `dataset_name` or `value_sha256`) for deletion.\n",
289+
"> **Note on deleting by `value`.** For the remove methods, the `value` filter defaults to full-string equality (`exact=True`), so `remove_seeds_from_memory(value=\"the\")` deletes only seeds whose value is exactly `\"the\"` — not everything containing it. This differs from `get_seeds`, which always matches `value` by substring. The same applies to the list filters: `harm_categories`, `authors`, `groups` and `parameters` must match whole list elements (case-insensitive), so `remove_seeds_from_memory(harm_categories=[\"hate\"])` does not also remove seeds tagged `\"hate_speech\"`. Pass `exact=False` to opt into substring deletion when you really want it. As a general rule, preview with the same filters via `get_seeds(...)` first and prefer a specific filter (such as `dataset_name` or `value_sha256`) for deletion. Because `get_seeds` always matches by substring, the preview is a superset of what the remove methods delete: some previewed seeds will survive the deletion, which is expected.\n",
290290
"\n",
291291
"> **Note on file-backed seeds.** For `image_path`, `audio_path`, and `video_path` seeds, removal deletes only the database record; the serialized file on disk is left in place. Delete those files separately if they are no longer needed."
292292
]

‎doc/code/memory/8_seed_database.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -145,6 +145,6 @@ def print_group(seed_group):
145145
#
146146
# For the most part these are user errors, but when you want to remove whole groups rather than individual seeds, use `remove_seed_groups_from_memory`. It applies the same filters, but removes every seed that shares a `prompt_group_id` with any match, so groups are never left partial. Note that it only affects seeds that belong to a group: a matching seed added individually (with no `prompt_group_id`) is skipped, so use `remove_seeds_from_memory` for those.
147147
#
148-
# > **Note on deleting by `value`.** For the remove methods, the `value` filter defaults to full-string equality (`exact=True`), so `remove_seeds_from_memory(value="the")` deletes only seeds whose value is exactly `"the"` — not everything containing it. This differs from `get_seeds`, which always matches `value` by substring. Pass `exact=False` to opt into substring deletion when you really want it. As a general rule, preview with the same filters via `get_seeds(...)` first and prefer a specific filter (such as `dataset_name` or `value_sha256`) for deletion.
148+
# > **Note on deleting by `value`.** For the remove methods, the `value` filter defaults to full-string equality (`exact=True`), so `remove_seeds_from_memory(value="the")` deletes only seeds whose value is exactly `"the"` — not everything containing it. This differs from `get_seeds`, which always matches `value` by substring. The same applies to the list filters: `harm_categories`, `authors`, `groups` and `parameters` must match whole list elements (case-insensitive), so `remove_seeds_from_memory(harm_categories=["hate"])` does not also remove seeds tagged `"hate_speech"`. Pass `exact=False` to opt into substring deletion when you really want it. As a general rule, preview with the same filters via `get_seeds(...)` first and prefer a specific filter (such as `dataset_name` or `value_sha256`) for deletion. Because `get_seeds` always matches by substring, the preview is a superset of what the remove methods delete: some previewed seeds will survive the deletion, which is expected.
149149
#
150150
# > **Note on file-backed seeds.** For `image_path`, `audio_path`, and `video_path` seeds, removal deletes only the database record; the serialized file on disk is left in place. Delete those files separately if they are no longer needed.

‎pyrit/memory/memory_interface.py‎

Lines changed: 41 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -3047,8 +3047,10 @@ def _build_seed_filter_conditions(
30473047
Args:
30483048
value (str): The value to match. By default this matches by substring; pass exact=True to
30493049
require full-string equality instead. If None, all values are returned.
3050-
exact (bool): When True, ``value`` is matched by full-string equality rather than substring.
3051-
Has no effect unless ``value`` is provided. Defaults to False (substring matching).
3050+
exact (bool): When True, ``value`` is matched by full-string equality rather than substring,
3051+
and ``harm_categories``, ``authors``, ``groups`` and ``parameters`` must match whole list
3052+
elements (case-insensitive) rather than substrings of the stored list. Defaults to False
3053+
(substring matching).
30523054
value_sha256 (Sequence[str] | None): A list of SHA256 hashes of values to match.
30533055
If None, all values are returned.
30543056
dataset_name (str): The dataset name to match exactly. If None, all dataset names are considered.
@@ -3104,12 +3106,14 @@ def _build_seed_filter_conditions(
31043106
elif seed_type is not None:
31053107
conditions.append(SeedEntry.seed_type == seed_type)
31063108

3107-
self._add_list_conditions(field=SeedEntry.harm_categories, values=harm_categories, conditions=conditions)
3108-
self._add_list_conditions(field=SeedEntry.authors, values=authors, conditions=conditions)
3109-
self._add_list_conditions(field=SeedEntry.groups, values=groups, conditions=conditions)
3109+
self._add_list_conditions(
3110+
field=SeedEntry.harm_categories, values=harm_categories, conditions=conditions, exact=exact
3111+
)
3112+
self._add_list_conditions(field=SeedEntry.authors, values=authors, conditions=conditions, exact=exact)
3113+
self._add_list_conditions(field=SeedEntry.groups, values=groups, conditions=conditions, exact=exact)
31103114

31113115
if parameters:
3112-
self._add_list_conditions(field=SeedEntry.parameters, values=parameters, conditions=conditions)
3116+
self._add_list_conditions(field=SeedEntry.parameters, values=parameters, conditions=conditions, exact=exact)
31133117

31143118
if metadata:
31153119
conditions.append(self._get_seed_metadata_conditions(metadata=metadata))
@@ -3229,10 +3233,11 @@ def remove_seeds_from_memory(
32293233
value (str): The value to match. For the remove methods this defaults to full-string equality
32303234
(exact=True) so a short or common value does not delete far more seeds than intended; pass
32313235
exact=False to match by substring instead. If None, all values are considered.
3232-
exact (bool): When True, ``value`` is matched by full-string equality rather than substring.
3233-
Has no effect unless ``value`` is provided. Defaults to True for the remove methods (the
3234-
safer choice for deletion). Note this differs from get_seeds, which always matches ``value``
3235-
by substring.
3236+
exact (bool): When True, ``value`` is matched by full-string equality rather than substring, and
3237+
``harm_categories``, ``authors``, ``groups`` and ``parameters`` must match whole list elements
3238+
(case-insensitive), so ``harm_categories=["hate"]`` does not also remove seeds tagged
3239+
``"hate_speech"``. Defaults to True for the remove methods (the safer choice for deletion).
3240+
Note this differs from get_seeds, which always matches these filters by substring.
32363241
value_sha256 (Sequence[str] | None): A list of SHA256 hashes of values to match.
32373242
If None, all values are considered.
32383243
dataset_name (str): The dataset name to match exactly. If None, all dataset names are considered.
@@ -3246,9 +3251,9 @@ def remove_seeds_from_memory(
32463251
all harm categories are considered.
32473252
Specifying multiple harm categories matches only prompts that are marked with all harm categories.
32483253
added_by (str): The user who added the prompts.
3249-
authors (Sequence[str]): A list of authors to filter by.
3250-
Note that this filters by substring, so a query for "Adam Jones" may not return results if the record
3251-
is "A. Jones", "Jones, Adam", etc. If None, all authors are considered.
3254+
authors (Sequence[str]): A list of authors to filter by. With exact=True (the default) each author
3255+
must match a stored author exactly (case-insensitive); with exact=False this filters by substring.
3256+
If None, all authors are considered.
32523257
groups (Sequence[str]): A list of groups to filter by. If None, all groups are considered.
32533258
source (str): The source to filter by. If None, all sources are considered.
32543259
seed_type (SeedType): The type of seed to filter by ("prompt", "objective", or
@@ -3344,10 +3349,11 @@ def remove_seed_groups_from_memory(
33443349
value (str): The value to match. For the remove methods this defaults to full-string equality
33453350
(exact=True) so a short or common value does not delete far more seeds than intended; pass
33463351
exact=False to match by substring instead. If None, all values are considered.
3347-
exact (bool): When True, ``value`` is matched by full-string equality rather than substring.
3348-
Has no effect unless ``value`` is provided. Defaults to True for the remove methods (the
3349-
safer choice for deletion). Note this differs from get_seeds, which always matches ``value``
3350-
by substring.
3352+
exact (bool): When True, ``value`` is matched by full-string equality rather than substring, and
3353+
``harm_categories``, ``authors``, ``groups`` and ``parameters`` must match whole list elements
3354+
(case-insensitive), so ``harm_categories=["hate"]`` does not also remove seeds tagged
3355+
``"hate_speech"``. Defaults to True for the remove methods (the safer choice for deletion).
3356+
Note this differs from get_seeds, which always matches these filters by substring.
33513357
value_sha256 (Sequence[str] | None): A list of SHA256 hashes of values to match.
33523358
If None, all values are considered.
33533359
dataset_name (str): The dataset name to match exactly. If None, all dataset names are considered.
@@ -3361,9 +3367,9 @@ def remove_seed_groups_from_memory(
33613367
all harm categories are considered.
33623368
Specifying multiple harm categories matches only prompts that are marked with all harm categories.
33633369
added_by (str): The user who added the prompts.
3364-
authors (Sequence[str]): A list of authors to filter by.
3365-
Note that this filters by substring, so a query for "Adam Jones" may not return results if the record
3366-
is "A. Jones", "Jones, Adam", etc. If None, all authors are considered.
3370+
authors (Sequence[str]): A list of authors to filter by. With exact=True (the default) each author
3371+
must match a stored author exactly (case-insensitive); with exact=False this filters by substring.
3372+
If None, all authors are considered.
33673373
groups (Sequence[str]): A list of groups to filter by. If None, all groups are considered.
33683374
source (str): The source to filter by. If None, all sources are considered.
33693375
seed_type (SeedType): The type of seed to filter by ("prompt", "objective", or
@@ -3431,11 +3437,25 @@ def remove_seed_groups_from_memory(
34313437

34323438
def _add_list_conditions(
34333439
self,
3440+
*,
34343441
field: InstrumentedAttribute[Any],
34353442
conditions: "list[ColumnElement[bool]]",
34363443
values: Sequence[str] | None = None,
3444+
exact: bool = False,
34373445
) -> None:
3438-
if values:
3446+
if not values:
3447+
return
3448+
if exact:
3449+
# Match whole list elements (case-insensitive) so "hate" does not match "hate_speech" or "whatever".
3450+
conditions.append(
3451+
self._get_condition_json_array_match(
3452+
json_column=field,
3453+
property_path="$",
3454+
array_to_match=list(values),
3455+
match_mode="all",
3456+
)
3457+
)
3458+
else:
34393459
conditions.extend(field.contains(value) for value in values)
34403460

34413461
async def _serialize_seed_value_async(self, prompt: Seed) -> str:

‎tests/unit/memory/memory_interface/test_interface_remove_seeds.py‎

Lines changed: 66 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -249,6 +249,55 @@ async def test_remove_seeds_by_value_exact_is_narrow(sqlite_instance: MemoryInte
249249
assert remaining[0].value == "the lazy dog"
250250

251251

252+
async def test_remove_seeds_by_harm_categories_matches_whole_elements(sqlite_instance: MemoryInterface):
253+
seed_prompts = [
254+
SeedPrompt(value="exact", harm_categories=["hate"], data_type="text"),
255+
SeedPrompt(value="longer_label", harm_categories=["hate_speech"], data_type="text"),
256+
SeedPrompt(value="inner_substring", harm_categories=["whatever"], data_type="text"),
257+
SeedPrompt(value="wildcard", harm_categories=["selfXharm"], data_type="text"),
258+
]
259+
await sqlite_instance.add_seeds_to_memory_async(seeds=seed_prompts, added_by="test")
260+
261+
assert sqlite_instance.remove_seeds_from_memory(harm_categories=["hate"]) == 1
262+
assert sqlite_instance.remove_seeds_from_memory(harm_categories=["self_harm"]) == 0
263+
264+
remaining = sorted(seed.value for seed in sqlite_instance.get_seeds())
265+
assert remaining == ["inner_substring", "longer_label", "wildcard"]
266+
267+
268+
async def test_remove_seeds_by_list_filters_is_case_insensitive(sqlite_instance: MemoryInterface):
269+
seed_prompts = [
270+
SeedPrompt(value="p1", harm_categories=["Violence"], data_type="text"),
271+
SeedPrompt(value="p2", harm_categories=["fraud"], data_type="text"),
272+
]
273+
await sqlite_instance.add_seeds_to_memory_async(seeds=seed_prompts, added_by="test")
274+
275+
assert sqlite_instance.remove_seeds_from_memory(harm_categories=["violence"]) == 1
276+
277+
278+
async def test_remove_seeds_by_groups_authors_parameters_match_whole_elements(sqlite_instance: MemoryInterface):
279+
seed_prompts = [
280+
SeedPrompt(value="p1", groups=["team"], authors=["Ann"], parameters=["goal"], data_type="text"),
281+
SeedPrompt(value="p2", groups=["team_b"], authors=["Annabel"], parameters=["goal_2"], data_type="text"),
282+
]
283+
await sqlite_instance.add_seeds_to_memory_async(seeds=seed_prompts, added_by="test")
284+
285+
assert sqlite_instance.remove_seeds_from_memory(groups=["team"], authors=["Ann"], parameters=["goal"]) == 1
286+
287+
remaining = sqlite_instance.get_seeds()
288+
assert [seed.value for seed in remaining] == ["p2"]
289+
290+
291+
async def test_remove_seeds_by_harm_categories_substring_when_not_exact(sqlite_instance: MemoryInterface):
292+
seed_prompts = [
293+
SeedPrompt(value="p1", harm_categories=["hate"], data_type="text"),
294+
SeedPrompt(value="p2", harm_categories=["hate_speech"], data_type="text"),
295+
]
296+
await sqlite_instance.add_seeds_to_memory_async(seeds=seed_prompts, added_by="test")
297+
298+
assert sqlite_instance.remove_seeds_from_memory(harm_categories=["hate"], exact=False) == 2
299+
300+
252301
async def test_remove_seeds_multi_filter_narrowing(sqlite_instance: MemoryInterface):
253302
seed_prompts = [
254303
SeedPrompt(value="prompt1", dataset_name="ds1", added_by="user1", data_type="text"),
@@ -325,6 +374,23 @@ async def test_remove_seed_groups_removes_entire_group(sqlite_instance: MemoryIn
325374
assert remaining[0].value == "keep"
326375

327376

377+
async def test_remove_seed_groups_by_harm_categories_matches_whole_elements(sqlite_instance: MemoryInterface):
378+
hate_group = SeedGroup(
379+
seeds=[SeedPrompt(value="hate", harm_categories=["hate"], data_type="text", sequence=0, role="user")]
380+
)
381+
speech_group = SeedGroup(
382+
seeds=[
383+
SeedPrompt(value="hate_speech", harm_categories=["hate_speech"], data_type="text", sequence=0, role="user")
384+
]
385+
)
386+
await sqlite_instance.add_seed_groups_to_memory_async(prompt_groups=[hate_group, speech_group], added_by="test")
387+
388+
assert sqlite_instance.remove_seed_groups_from_memory(harm_categories=["hate"]) == 1
389+
390+
remaining = sqlite_instance.get_seeds()
391+
assert [seed.value for seed in remaining] == ["hate_speech"]
392+
393+
328394
async def test_remove_seed_groups_spanning_multiple_datasets(sqlite_instance: MemoryInterface):
329395
group = SeedGroup(
330396
seeds=[

0 commit comments

Comments
 (0)