Skip to content

Commit 79810bc

Browse files
rlundeen2Copilot
andcommitted
MAINT: tighten model from_dict edge cases (labels, pyrit_version, completion_time)
Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
1 parent 11b30d5 commit 79810bc

5 files changed

Lines changed: 51 additions & 4 deletions

File tree

pyrit/models/attack_result.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -240,7 +240,7 @@ def to_dict(self) -> dict[str, Any]:
240240
"execution_time_ms": self.execution_time_ms,
241241
"outcome": self.outcome.value,
242242
"outcome_reason": self.outcome_reason,
243-
"timestamp": self.timestamp.isoformat() if self.timestamp else None,
243+
"timestamp": self.timestamp.isoformat(),
244244
"related_conversations": sorted(
245245
[ref.to_dict() for ref in self.related_conversations],
246246
key=lambda r: r["conversation_id"],

pyrit/models/message_piece.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -370,7 +370,7 @@ def from_dict(cls, data: dict[str, Any]) -> MessagePiece:
370370
conversation_id=data.get("conversation_id"),
371371
sequence=data.get("sequence", -1),
372372
timestamp=(datetime.fromisoformat(str(data["timestamp"])) if data.get("timestamp") else None),
373-
labels=data.get("labels"),
373+
labels=data.get("labels") or None,
374374
targeted_harm_categories=data.get("targeted_harm_categories"),
375375
prompt_metadata=data.get("prompt_metadata"),
376376
converter_identifiers=(

pyrit/models/scenario_result.py

Lines changed: 7 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -79,7 +79,7 @@ def from_dict(cls, data: dict[str, Any]) -> ScenarioIdentifier:
7979
description=data.get("description", ""),
8080
scenario_version=data.get("scenario_version", 1),
8181
init_data=data.get("init_data"),
82-
pyrit_version=data.get("pyrit_version"),
82+
pyrit_version=data.get("pyrit_version") or "unknown",
8383
)
8484

8585

@@ -339,7 +339,7 @@ def from_dict(cls, data: dict[str, Any]) -> ScenarioResult:
339339
"""
340340
from pyrit.identifiers.component_identifier import ComponentIdentifier
341341

342-
return cls(
342+
result = cls(
343343
id=uuid.UUID(data["id"]) if data.get("id") else None,
344344
scenario_identifier=ScenarioIdentifier.from_dict(data["scenario_identifier"]),
345345
objective_target_identifier=(
@@ -366,3 +366,8 @@ def from_dict(cls, data: dict[str, Any]) -> ScenarioResult:
366366
error_message=data.get("error_message"),
367367
error_type=data.get("error_type"),
368368
)
369+
# Preserve missing completion_time: __init__ defaults it to now(), but a
370+
# still-running scenario shouldn't be marked as completed-at-load-time.
371+
if not data.get("completion_time"):
372+
result.completion_time = None # type: ignore[ty:invalid-assignment]
373+
return result

tests/unit/models/test_message_piece.py

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1244,3 +1244,13 @@ def test_to_dict_from_dict_roundtrip_after_set_piece_not_in_database():
12441244
roundtripped = MessagePiece.from_dict(serialized)
12451245
assert isinstance(roundtripped.id, uuid.UUID)
12461246
assert isinstance(roundtripped.original_prompt_id, uuid.UUID)
1247+
1248+
1249+
def test_from_dict_does_not_emit_deprecation_warning_for_default_labels():
1250+
"""from_dict on a piece with no labels should not trigger the labels deprecation warning."""
1251+
piece = MessagePiece(role="user", original_value="hello", conversation_id="conv-1")
1252+
serialized = piece.to_dict()
1253+
1254+
with warnings.catch_warnings():
1255+
warnings.simplefilter("error", DeprecationWarning)
1256+
MessagePiece.from_dict(serialized)

tests/unit/models/test_scenario_result.py

Lines changed: 32 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -268,3 +268,35 @@ def test_scenario_result_to_dict_from_dict_roundtrip():
268268
)
269269
roundtripped = ScenarioResult.from_dict(original.to_dict())
270270
assert original.to_dict() == roundtripped.to_dict()
271+
272+
273+
def test_scenario_identifier_from_dict_missing_pyrit_version_yields_unknown():
274+
"""from_dict should preserve 'unknown' rather than fabricating the current version when missing."""
275+
data = {
276+
"name": "Legacy",
277+
"description": "loaded from older payload",
278+
"scenario_version": 1,
279+
"init_data": None,
280+
# pyrit_version intentionally absent
281+
}
282+
identifier = ScenarioIdentifier.from_dict(data)
283+
assert identifier.pyrit_version == "unknown"
284+
285+
286+
def test_scenario_result_from_dict_preserves_missing_completion_time():
287+
"""An in-progress scenario serialized without completion_time should round-trip with completion_time=None."""
288+
scenario_id = ScenarioIdentifier(name="Test", scenario_version=1, pyrit_version="0.14.0")
289+
target_id = ComponentIdentifier(class_name="OpenAIChatTarget", class_module="pyrit.prompt_target")
290+
291+
original = ScenarioResult(
292+
scenario_identifier=scenario_id,
293+
objective_target_identifier=target_id,
294+
objective_scorer_identifier=None,
295+
attack_results={},
296+
scenario_run_state="IN_PROGRESS",
297+
)
298+
original.completion_time = None # type: ignore[ty:invalid-assignment]
299+
300+
roundtripped = ScenarioResult.from_dict(original.to_dict())
301+
assert roundtripped.completion_time is None
302+
assert roundtripped.scenario_run_state == "IN_PROGRESS"

0 commit comments

Comments
 (0)