From 4284e13bfd66b1a22fc4d50a25b68cc69e3fe516 Mon Sep 17 00:00:00 2001 From: JR Date: Fri, 9 Oct 2026 22:36:01 +0200 Subject: [PATCH 1/2] FIX: match label keys containing dots or hyphens in attack result and message filters --- pyrit/memory/azure_sql_memory.py | 16 +++++--- pyrit/memory/sqlite_memory.py | 4 +- .../test_interface_attack_results.py | 34 +++++++++++++++++ tests/unit/memory/test_azure_sql_memory.py | 38 +++++++++++++++---- 4 files changed, 77 insertions(+), 15 deletions(-) diff --git a/pyrit/memory/azure_sql_memory.py b/pyrit/memory/azure_sql_memory.py index d5a1d93e12..422713b641 100644 --- a/pyrit/memory/azure_sql_memory.py +++ b/pyrit/memory/azure_sql_memory.py @@ -348,9 +348,11 @@ def _get_message_pieces_memory_label_conditions(self, *, memory_labels: dict[str are_label_parts: list[str] = [] are_bindparams: dict[str, str] = {} - for key, value in memory_labels.items(): - are_param = f"are_ml_{key}" - are_label_parts.append(f"JSON_VALUE(\"AttackResultEntries\".labels, '$.{key}') = :{are_param}") + for key_index, (key, value) in enumerate(memory_labels.items()): + path_param = f"are_ml_path_{key_index}" + are_param = f"are_ml_{key_index}" + are_label_parts.append(f'JSON_VALUE("AttackResultEntries".labels, :{path_param}) = :{are_param}') + are_bindparams[path_param] = f'$."{key}"' are_bindparams[are_param] = str(value) combined_are = " AND ".join(are_label_parts) @@ -545,17 +547,19 @@ def _get_attack_result_label_condition(self, *, labels: dict[str, str | Sequence are_label_conditions: list[str] = [] are_bindparams: dict[str, str] = {} - for key, raw_value in labels.items(): + for key_index, (key, raw_value) in enumerate(labels.items()): values = [raw_value] if isinstance(raw_value, str) else list(raw_value) if not values: continue + path_param = f"are_label_path_{key_index}" + are_bindparams[path_param] = f'$."{key}"' are_placeholders = [] for idx, v in enumerate(values): - are_param = f"are_label_{key}_{idx}" + are_param = f"are_label_{key_index}_{idx}" are_placeholders.append(f":{are_param}") are_bindparams[are_param] = str(v) are_in = ", ".join(are_placeholders) - are_label_conditions.append(f"JSON_VALUE(\"AttackResultEntries\".labels, '$.{key}') IN ({are_in})") + are_label_conditions.append(f'JSON_VALUE("AttackResultEntries".labels, :{path_param}) IN ({are_in})') are_parts: list[Any] = [AttackResultEntry.labels.isnot(None)] if are_label_conditions: diff --git a/pyrit/memory/sqlite_memory.py b/pyrit/memory/sqlite_memory.py index a94febd211..d5219c9bc0 100644 --- a/pyrit/memory/sqlite_memory.py +++ b/pyrit/memory/sqlite_memory.py @@ -378,7 +378,7 @@ def _get_message_pieces_memory_label_conditions(self, *, memory_labels: dict[str """ per_key_are_conditions = [] for key, value in memory_labels.items(): - are_col = func.json_extract(AttackResultEntry.labels, f"$.{key}") + are_col = func.json_extract(AttackResultEntry.labels, f'$."{key}"') per_key_are_conditions.append(are_col == str(value)) return [ exists().where( @@ -626,7 +626,7 @@ def _get_attack_result_label_condition(self, *, labels: dict[str, str | Sequence values = [raw_value] if isinstance(raw_value, str) else list(raw_value) if not values: continue - are_col = func.json_extract(AttackResultEntry.labels, f"$.{key}") + are_col = func.json_extract(AttackResultEntry.labels, f'$."{key}"') per_key_are_conditions.append(are_col.in_(values)) return and_( diff --git a/tests/unit/memory/memory_interface/test_interface_attack_results.py b/tests/unit/memory/memory_interface/test_interface_attack_results.py index c87a202036..1f6f017449 100644 --- a/tests/unit/memory/memory_interface/test_interface_attack_results.py +++ b/tests/unit/memory/memory_interface/test_interface_attack_results.py @@ -1162,6 +1162,40 @@ async def test_get_attack_results_rejects_invalid_label_keys(sqlite_instance: Me (await sqlite_instance.get_attack_results_async(labels={bad_key: "value"})) +@pytest.mark.parametrize("key", ["team.name", "run-id"]) +async def test_get_attack_results_by_labels_key_with_dot_or_hyphen(sqlite_instance: MemoryInterface, key: str): + """A label key the allowlist accepts is matched as one key, not as a nested JSON path.""" + await sqlite_instance.add_attack_results_to_memory_async( + attack_results=[ + create_attack_result("conv_1", 1, labels={key: "safety"}), + create_attack_result("conv_2", 2, labels={"team": "safety"}), + ] + ) + + results = await sqlite_instance.get_attack_results_async(labels={key: "safety"}) + assert [r.conversation_id for r in results] == ["conv_1"] + + results = await sqlite_instance.get_attack_results_async(labels={key: ["other", "safety"]}) + assert [r.conversation_id for r in results] == ["conv_1"] + + +async def test_get_message_pieces_by_label_key_with_dot(sqlite_instance: MemoryInterface): + """Message pieces are found through a dotted label key on their attack result.""" + for conversation_id in ("conv_1", "conv_2"): + await sqlite_instance.add_message_pieces_to_memory_async( + message_pieces=[MessagePiece(role="user", original_value="hello", conversation_id=conversation_id)] + ) + await sqlite_instance.add_attack_results_to_memory_async( + attack_results=[ + create_attack_result("conv_1", 1, labels={"team.name": "safety"}), + create_attack_result("conv_2", 2, labels={"team": "safety"}), + ] + ) + + pieces = await sqlite_instance.get_message_pieces_async(labels={"team.name": "safety"}) + assert [p.conversation_id for p in pieces] == ["conv_1"] + + async def test_get_attack_results_by_labels_multiple(sqlite_instance: MemoryInterface): """Test filtering attack results by multiple labels (AND logic).""" diff --git a/tests/unit/memory/test_azure_sql_memory.py b/tests/unit/memory/test_azure_sql_memory.py index 5b943f8230..74ce80ef8d 100644 --- a/tests/unit/memory/test_azure_sql_memory.py +++ b/tests/unit/memory/test_azure_sql_memory.py @@ -308,7 +308,31 @@ def test_get_message_pieces_memory_label_conditions_bind_params(uninitialized_me memory_labels={"operation": "test_op"} ) params = conditions[0].compile().params - assert params == {"are_ml_operation": "test_op"} + assert params == {"are_ml_path_0": '$."operation"', "are_ml_0": "test_op"} + + +def test_label_conditions_bind_whole_keys_with_dot_or_hyphen(uninitialized_memory_interface: AzureSQLMemory): + """Allowlisted keys with ``.`` or ``-`` build a query that looks up each whole key.""" + attack_condition = uninitialized_memory_interface._get_attack_result_label_condition( + labels={"team.name": ["red", "blue"], "run-id": "r1"} + ) + assert attack_condition.compile().params == { + "are_label_path_0": '$."team.name"', + "are_label_0_0": "red", + "are_label_0_1": "blue", + "are_label_path_1": '$."run-id"', + "are_label_1_0": "r1", + } + + piece_condition = uninitialized_memory_interface._get_message_pieces_memory_label_conditions( + memory_labels={"team.name": "red", "run-id": "r1"} + )[0] + assert piece_condition.compile().params == { + "are_ml_path_0": '$."team.name"', + "are_ml_0": "red", + "are_ml_path_1": '$."run-id"', + "are_ml_1": "r1", + } async def test_update_entries_async(memory_interface: AzureSQLMemory): @@ -421,7 +445,7 @@ def test_get_attack_result_label_condition_with_string_value(memory_interface: A """String values produce a single-placeholder IN clause with the stringified value.""" condition = memory_interface._get_attack_result_label_condition(labels={"operator": "roakey"}) params = condition.compile().params - assert params == {"are_label_operator_0": "roakey"} + assert params == {"are_label_path_0": '$."operator"', "are_label_0_0": "roakey"} def test_get_attack_result_label_condition_with_sequence_value(memory_interface: AzureSQLMemory): @@ -429,9 +453,10 @@ def test_get_attack_result_label_condition_with_sequence_value(memory_interface: condition = memory_interface._get_attack_result_label_condition(labels={"operation": ["op_a", "op_b", "op_c"]}) params = condition.compile().params assert params == { - "are_label_operation_0": "op_a", - "are_label_operation_1": "op_b", - "are_label_operation_2": "op_c", + "are_label_path_0": '$."operation"', + "are_label_0_0": "op_a", + "are_label_0_1": "op_b", + "are_label_0_2": "op_c", } @@ -440,8 +465,7 @@ def test_get_attack_result_label_condition_skips_empty_sequence(memory_interface condition = memory_interface._get_attack_result_label_condition(labels={"operator": "roakey", "operation": []}) params = condition.compile().params # operator gets bind params; operation (empty) does not. - assert params == {"are_label_operator_0": "roakey"} - assert not any("label_operation_" in k for k in params) + assert params == {"are_label_path_0": '$."operator"', "are_label_0_0": "roakey"} def test_get_attack_result_label_condition_empty_labels_dict(memory_interface: AzureSQLMemory): From b12812af1c324813cf08dc9e36637df2de7554e8 Mon Sep 17 00:00:00 2001 From: Roman Lutz Date: Sat, 10 Oct 2026 21:39:52 -0700 Subject: [PATCH 2/2] Fix literal SQLite message label keys Match bound JSON object keys without JSON-path interpretation and cover quoted and backslashed labels. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- pyrit/memory/sqlite_memory.py | 15 +++-- .../test_interface_attack_results.py | 56 +++++++++++++++++-- 2 files changed, 60 insertions(+), 11 deletions(-) diff --git a/pyrit/memory/sqlite_memory.py b/pyrit/memory/sqlite_memory.py index d5219c9bc0..01db1fe859 100644 --- a/pyrit/memory/sqlite_memory.py +++ b/pyrit/memory/sqlite_memory.py @@ -368,18 +368,23 @@ def _create_engine(self, *, has_echo: bool) -> Engine: def _get_message_pieces_memory_label_conditions(self, *, memory_labels: dict[str, str]) -> list[Any]: """ Generate SQLAlchemy filter conditions for filtering conversation pieces by memory labels. - For SQLite, we use JSON_EXTRACT function to handle JSON fields. + Match literal object keys through json_each(), without interpreting keys as JSON paths. - Matches if labels are on the PromptMemoryEntry itself OR on any - AttackResultEntry that shares the same conversation_id. + All labels must match one AttackResultEntry that shares the same conversation_id. Returns: list: A list of SQLAlchemy conditions. """ per_key_are_conditions = [] + label_entries = func.json_each(AttackResultEntry.labels).table_valued("key", "value") for key, value in memory_labels.items(): - are_col = func.json_extract(AttackResultEntry.labels, f'$."{key}"') - per_key_are_conditions.append(are_col == str(value)) + per_key_are_conditions.append( + select(1) + .select_from(label_entries) + .where(label_entries.c.key == key, label_entries.c.value == str(value)) + .correlate(AttackResultEntry) + .exists() + ) return [ exists().where( and_( diff --git a/tests/unit/memory/memory_interface/test_interface_attack_results.py b/tests/unit/memory/memory_interface/test_interface_attack_results.py index 1f6f017449..b3bfb0ded0 100644 --- a/tests/unit/memory/memory_interface/test_interface_attack_results.py +++ b/tests/unit/memory/memory_interface/test_interface_attack_results.py @@ -1179,21 +1179,65 @@ async def test_get_attack_results_by_labels_key_with_dot_or_hyphen(sqlite_instan assert [r.conversation_id for r in results] == ["conv_1"] -async def test_get_message_pieces_by_label_key_with_dot(sqlite_instance: MemoryInterface): - """Message pieces are found through a dotted label key on their attack result.""" - for conversation_id in ("conv_1", "conv_2"): +@pytest.mark.parametrize( + "key, value", + [ + ("team.name", "safety"), + ("run-id", "safety"), + ('team"name', "safety"), + (r"team\name", "safety"), + ('team\\"name', "safety"), + ("team[0]", "safety"), + ("team name", "safety"), + ("caf\u00e9", "safety"), + ("", "safety"), + ("team.name", ""), + ('team"name', 'quoted"\\value'), + (r"team\name", "caf\u00e9"), + ], +) +async def test_get_message_pieces_by_literal_label_key_async( + sqlite_instance: MemoryInterface, key: str, value: str +) -> None: + """Message labels match complete string keys and values without JSON-path interpretation.""" + for conversation_id in ("conv_1", "conv_2", "conv_3"): await sqlite_instance.add_message_pieces_to_memory_async( message_pieces=[MessagePiece(role="user", original_value="hello", conversation_id=conversation_id)] ) await sqlite_instance.add_attack_results_to_memory_async( attack_results=[ - create_attack_result("conv_1", 1, labels={"team.name": "safety"}), - create_attack_result("conv_2", 2, labels={"team": "safety"}), + create_attack_result("conv_1", 1, labels={key: value}), + create_attack_result("conv_2", 2, labels={"team": value}), + create_attack_result("conv_3", 3, labels={key: f"{value}-other"}), + ] + ) + + pieces = await sqlite_instance.get_message_pieces_async(labels={key: value}) + assert [p.conversation_id for p in pieces] == ["conv_1"] + + +async def test_get_message_pieces_literal_labels_match_one_attack_without_duplicates_async( + sqlite_instance: MemoryInterface, +) -> None: + """All labels must match one attack result, without multiplying its conversation's pieces.""" + labels = {'team"name': "safety", r"run\id": "r1"} + for conversation_id in ("conv_1", "conv_2", "conv_3"): + await sqlite_instance.add_message_pieces_to_memory_async( + message_pieces=[MessagePiece(role="user", original_value="hello", conversation_id=conversation_id)] + ) + await sqlite_instance.add_attack_results_to_memory_async( + attack_results=[ + create_attack_result("conv_1", 1, labels=labels), + create_attack_result("conv_1", 2, labels=labels), + create_attack_result("conv_2", 3, labels={'team"name': "safety"}), + create_attack_result("conv_2", 4, labels={r"run\id": "r1"}), + create_attack_result("conv_3", 5, labels={**labels, r"run\id": "other"}), ] ) - pieces = await sqlite_instance.get_message_pieces_async(labels={"team.name": "safety"}) + pieces = await sqlite_instance.get_message_pieces_async(labels=labels, role="user") assert [p.conversation_id for p in pieces] == ["conv_1"] + assert await sqlite_instance.get_message_pieces_async(labels=labels, conversation_id="conv_2") == [] async def test_get_attack_results_by_labels_multiple(sqlite_instance: MemoryInterface):