Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 10 additions & 6 deletions pyrit/memory/azure_sql_memory.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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:
Expand Down
4 changes: 2 additions & 2 deletions pyrit/memory/sqlite_memory.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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_(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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)."""

Expand Down
38 changes: 31 additions & 7 deletions tests/unit/memory/test_azure_sql_memory.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -421,17 +445,18 @@ 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):
"""Sequence values produce one placeholder per element."""
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",
}


Expand All @@ -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):
Expand Down