Skip to content

Commit a6dbaa3

Browse files
rlundeen2Copilot
andcommitted
MAINT: address PR feedback on model to_dict/from_dict
Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
1 parent 9f39ff7 commit a6dbaa3

7 files changed

Lines changed: 36 additions & 23 deletions

File tree

pyrit/models/attack_result.py

Lines changed: 2 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -242,11 +242,8 @@ def to_dict(self) -> dict[str, Any]:
242242
"outcome_reason": self.outcome_reason,
243243
"timestamp": self.timestamp.isoformat() if self.timestamp else None,
244244
"related_conversations": sorted(
245-
[
246-
ref.to_dict() if isinstance(ref, ConversationReference) else ref
247-
for ref in self.related_conversations
248-
],
249-
key=lambda r: r["conversation_id"] if isinstance(r, dict) else "",
245+
[ref.to_dict() for ref in self.related_conversations],
246+
key=lambda r: r["conversation_id"],
250247
),
251248
"metadata": self.metadata,
252249
"labels": self.labels,

pyrit/models/message.py

Lines changed: 6 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -287,18 +287,16 @@ def to_dict(self) -> dict[str, object]:
287287
"""
288288
Convert the message to a dictionary representation including all piece details.
289289
290-
Serializes each piece individually via MessagePiece.to_dict(). This is the format
291-
expected by from_dict().
290+
Serializes each piece individually via MessagePiece.to_dict(). All message-level
291+
attributes (role, conversation_id, sequence, is_simulated) are derived from the
292+
pieces themselves, so only 'pieces' is included. This is the format expected by
293+
from_dict().
292294
293295
Returns:
294-
dict[str, object]: Dictionary with 'role', 'is_simulated', 'conversation_id',
295-
'sequence', and 'pieces' (list of MessagePiece.to_dict() dicts).
296+
dict[str, object]: Dictionary with a single 'pieces' key containing a list
297+
of MessagePiece.to_dict() dicts.
296298
"""
297299
return {
298-
"role": self.api_role,
299-
"is_simulated": self.is_simulated,
300-
"conversation_id": self.conversation_id,
301-
"sequence": self.sequence,
302300
"pieces": [piece.to_dict() for piece in self.message_pieces],
303301
}
304302

pyrit/models/message_piece.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -314,7 +314,7 @@ def to_dict(self) -> dict[str, object]:
314314
315315
"""
316316
return {
317-
"id": str(self.id),
317+
"id": str(self.id) if self.id is not None else None,
318318
"role": self._role,
319319
"conversation_id": self.conversation_id,
320320
"sequence": self.sequence,
@@ -336,7 +336,7 @@ def to_dict(self) -> dict[str, object]:
336336
"converted_value_sha256": self.converted_value_sha256,
337337
"response_error": self.response_error,
338338
"originator": self.originator,
339-
"original_prompt_id": str(self.original_prompt_id),
339+
"original_prompt_id": str(self.original_prompt_id) if self.original_prompt_id is not None else None,
340340
"scores": [score.to_dict() for score in self.scores],
341341
}
342342

pyrit/models/score.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -210,10 +210,10 @@ def from_dict(cls, data: dict[str, Any]) -> Score:
210210
return cls(
211211
id=data.get("id"),
212212
score_value=data["score_value"],
213-
score_value_description=data.get("score_value_description", ""),
213+
score_value_description=data["score_value_description"],
214214
score_type=data["score_type"],
215215
score_category=data.get("score_category"),
216-
score_rationale=data.get("score_rationale", ""),
216+
score_rationale=data["score_rationale"],
217217
score_metadata=data.get("score_metadata"),
218218
scorer_class_identifier=ComponentIdentifier.from_dict(data["scorer_class_identifier"]),
219219
message_piece_id=data["message_piece_id"],

tests/unit/message_normalizer/test_generic_system_squash_normalizer.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -61,9 +61,9 @@ async def test_generic_squash_normalize_to_dicts_async():
6161
assert isinstance(result, list)
6262
assert len(result) == 1
6363
assert isinstance(result[0], dict)
64-
assert result[0]["role"] == "user"
65-
assert "pieces" in result[0]
64+
assert set(result[0].keys()) == {"pieces"}
6665
assert len(result[0]["pieces"]) == 1
66+
assert result[0]["pieces"][0]["role"] == "user"
6767
assert "### Instructions ###" in result[0]["pieces"][0]["converted_value"]
6868
assert "System message" in result[0]["pieces"][0]["converted_value"]
6969
assert "User message" in result[0]["pieces"][0]["converted_value"]

tests/unit/models/test_message.py

Lines changed: 2 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -226,11 +226,9 @@ def test_message_to_dict() -> None:
226226
message = Message.from_prompt(prompt="Hello world", role="user")
227227
result = message.to_dict()
228228

229-
assert result["role"] == "user"
230-
assert result["is_simulated"] is False
231-
assert "conversation_id" in result
232-
assert "sequence" in result
229+
assert set(result.keys()) == {"pieces"}
233230
assert len(result["pieces"]) == 1
231+
assert result["pieces"][0]["role"] == "user"
234232
assert result["pieces"][0]["converted_value"] == "Hello world"
235233
assert result["pieces"][0]["converted_value_data_type"] == "text"
236234

tests/unit/models/test_message_piece.py

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1224,3 +1224,23 @@ def test_to_dict_from_dict_roundtrip():
12241224
)
12251225
roundtripped = MessagePiece.from_dict(original.to_dict())
12261226
assert original.to_dict() == roundtripped.to_dict()
1227+
1228+
1229+
def test_to_dict_from_dict_roundtrip_after_set_piece_not_in_database():
1230+
"""Pieces marked not-in-database (id=None) must serialize and deserialize cleanly without ValueError."""
1231+
piece = MessagePiece(
1232+
role="user",
1233+
original_value="Hello world",
1234+
conversation_id="conv-not-in-db",
1235+
)
1236+
piece.set_piece_not_in_database()
1237+
piece.original_prompt_id = None # type: ignore[assignment]
1238+
1239+
serialized = piece.to_dict()
1240+
assert serialized["id"] is None
1241+
assert serialized["original_prompt_id"] is None
1242+
1243+
# Must not raise ValueError on the literal string "None" or similar corruption.
1244+
roundtripped = MessagePiece.from_dict(serialized)
1245+
assert isinstance(roundtripped.id, uuid.UUID)
1246+
assert isinstance(roundtripped.original_prompt_id, uuid.UUID)

0 commit comments

Comments
 (0)