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):