Skip to content

Commit 4daaf4a

Browse files
rlundeen2Copilot
andcommitted
Add to_dict/from_dict roundtrip serialization to model classes
Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
1 parent 361c2a4 commit 4daaf4a

12 files changed

Lines changed: 610 additions & 20 deletions

pyrit/models/attack_result.py

Lines changed: 84 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -224,6 +224,90 @@ def __str__(self) -> str:
224224
"""
225225
return f"AttackResult: {self.conversation_id}: {self.outcome.value}: {self.objective[:50]}..."
226226

227+
def to_dict(self) -> dict[str, Any]:
228+
"""
229+
Serialize this attack result to a JSON-compatible dictionary.
230+
231+
Returns:
232+
dict[str, Any]: Serialized payload suitable for REST APIs or persistence.
233+
"""
234+
from pyrit.models.conversation_reference import ConversationReference
235+
236+
return {
237+
"conversation_id": self.conversation_id,
238+
"objective": self.objective,
239+
"attack_result_id": self.attack_result_id,
240+
"atomic_attack_identifier": (
241+
self.atomic_attack_identifier.to_dict() if self.atomic_attack_identifier else None
242+
),
243+
"last_response": self.last_response.to_dict() if self.last_response else None,
244+
"last_score": self.last_score.to_dict() if self.last_score else None,
245+
"executed_turns": self.executed_turns,
246+
"execution_time_ms": self.execution_time_ms,
247+
"outcome": self.outcome.value,
248+
"outcome_reason": self.outcome_reason,
249+
"timestamp": self.timestamp.isoformat() if self.timestamp else None,
250+
"related_conversations": sorted(
251+
[
252+
ref.to_dict() if isinstance(ref, ConversationReference) else ref
253+
for ref in self.related_conversations
254+
],
255+
key=lambda r: r["conversation_id"] if isinstance(r, dict) else "",
256+
),
257+
"metadata": self.metadata,
258+
"labels": self.labels,
259+
"error_message": self.error_message,
260+
"error_type": self.error_type,
261+
"error_traceback": self.error_traceback,
262+
"retry_events": [e.to_dict() for e in self.retry_events],
263+
"total_retries": self.total_retries,
264+
}
265+
266+
@classmethod
267+
def from_dict(cls, data: dict[str, Any]) -> AttackResult:
268+
"""
269+
Reconstruct an AttackResult from a dictionary.
270+
271+
Args:
272+
data (dict[str, Any]): Dictionary as produced by to_dict().
273+
274+
Returns:
275+
AttackResult: Reconstructed instance.
276+
"""
277+
from pyrit.identifiers.component_identifier import ComponentIdentifier
278+
from pyrit.models.conversation_reference import ConversationReference
279+
from pyrit.models.message_piece import MessagePiece
280+
from pyrit.models.retry_event import RetryEvent
281+
from pyrit.models.score import Score
282+
283+
return cls(
284+
conversation_id=data["conversation_id"],
285+
objective=data["objective"],
286+
attack_result_id=data.get("attack_result_id", str(uuid.uuid4())),
287+
atomic_attack_identifier=(
288+
ComponentIdentifier.from_dict(data["atomic_attack_identifier"])
289+
if data.get("atomic_attack_identifier")
290+
else None
291+
),
292+
last_response=(MessagePiece.from_dict(data["last_response"]) if data.get("last_response") else None),
293+
last_score=Score.from_dict(data["last_score"]) if data.get("last_score") else None,
294+
executed_turns=data.get("executed_turns", 0),
295+
execution_time_ms=data.get("execution_time_ms", 0),
296+
outcome=AttackOutcome(data.get("outcome", "undetermined")),
297+
outcome_reason=data.get("outcome_reason"),
298+
timestamp=(
299+
datetime.fromisoformat(data["timestamp"]) if data.get("timestamp") else datetime.now(timezone.utc)
300+
),
301+
related_conversations={ConversationReference.from_dict(r) for r in data.get("related_conversations", [])},
302+
metadata=data.get("metadata", {}),
303+
labels=data.get("labels", {}),
304+
error_message=data.get("error_message"),
305+
error_type=data.get("error_type"),
306+
error_traceback=data.get("error_traceback"),
307+
retry_events=[RetryEvent.from_dict(e) for e in data.get("retry_events", [])],
308+
total_retries=data.get("total_retries", 0),
309+
)
310+
227311

228312
def _add_attack_identifier_compat(cls: type) -> type:
229313
"""

pyrit/models/conversation_reference.py

Lines changed: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -36,6 +36,36 @@ def __hash__(self) -> int:
3636
"""
3737
return hash(self.conversation_id)
3838

39+
def to_dict(self) -> dict[str, str | None]:
40+
"""
41+
Serialize to a JSON-compatible dictionary.
42+
43+
Returns:
44+
dict[str, str | None]: Dictionary with conversation_id, conversation_type, and description.
45+
"""
46+
return {
47+
"conversation_id": self.conversation_id,
48+
"conversation_type": self.conversation_type.value,
49+
"description": self.description,
50+
}
51+
52+
@classmethod
53+
def from_dict(cls, data: dict[str, str | None]) -> ConversationReference:
54+
"""
55+
Reconstruct a ConversationReference from a dictionary.
56+
57+
Args:
58+
data (dict[str, str | None]): Dictionary as produced by to_dict().
59+
60+
Returns:
61+
ConversationReference: Reconstructed instance.
62+
"""
63+
return cls(
64+
conversation_id=str(data["conversation_id"]),
65+
conversation_type=ConversationType(data["conversation_type"]),
66+
description=data.get("description"),
67+
)
68+
3969
def __eq__(self, other: object) -> bool:
4070
"""
4171
Compare two references by conversation ID.

pyrit/models/message.py

Lines changed: 27 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,7 @@
66
import copy
77
import uuid
88
from datetime import datetime, timezone
9-
from typing import TYPE_CHECKING, Optional, Union
9+
from typing import TYPE_CHECKING, Any, Optional, Union
1010

1111
from pyrit.common.utils import combine_dict
1212
from pyrit.models.message_piece import MessagePiece
@@ -285,28 +285,41 @@ def __str__(self) -> str:
285285

286286
def to_dict(self) -> dict[str, object]:
287287
"""
288-
Convert the message to a dictionary representation.
288+
Convert the message to a dictionary representation including all piece details.
289289
290-
Returns:
291-
dict: A dictionary with 'role', 'converted_value', 'conversation_id', 'sequence',
292-
and 'converted_value_data_type' keys.
290+
Serializes each piece individually via MessagePiece.to_dict(). This is the format
291+
expected by from_dict().
293292
293+
Returns:
294+
dict[str, object]: Dictionary with 'role', 'is_simulated', 'conversation_id',
295+
'sequence', and 'pieces' (list of MessagePiece.to_dict() dicts).
294296
"""
295-
if len(self.message_pieces) == 1:
296-
converted_value: str | list[str] = self.message_pieces[0].converted_value
297-
converted_value_data_type: str | list[str] = self.message_pieces[0].converted_value_data_type
298-
else:
299-
converted_value = [piece.converted_value for piece in self.message_pieces]
300-
converted_value_data_type = [piece.converted_value_data_type for piece in self.message_pieces]
301-
302297
return {
303298
"role": self.api_role,
304-
"converted_value": converted_value,
299+
"is_simulated": self.is_simulated,
305300
"conversation_id": self.conversation_id,
306301
"sequence": self.sequence,
307-
"converted_value_data_type": converted_value_data_type,
302+
"pieces": [piece.to_dict() for piece in self.message_pieces],
308303
}
309304

305+
@classmethod
306+
def from_dict(cls, data: dict[str, Any]) -> Message:
307+
"""
308+
Reconstruct a Message from a dictionary.
309+
310+
Expects the format produced by to_dict(), which includes a 'pieces' key
311+
containing a list of MessagePiece dictionaries.
312+
313+
Args:
314+
data (dict[str, Any]): Dictionary as produced by to_dict().
315+
316+
Returns:
317+
Message: Reconstructed instance.
318+
"""
319+
pieces_data = data.get("pieces", [])
320+
message_pieces = [MessagePiece.from_dict(p) for p in pieces_data]
321+
return cls(message_pieces, skip_validation=True)
322+
310323
@staticmethod
311324
def get_all_values(messages: Sequence[Message]) -> list[str]:
312325
"""

pyrit/models/message_piece.py

Lines changed: 52 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55

66
import uuid
77
from datetime import datetime, timezone
8-
from typing import TYPE_CHECKING, Literal, Optional, Union, get_args
8+
from typing import TYPE_CHECKING, Any, Literal, Optional, Union, get_args
99
from uuid import uuid4
1010

1111
from pyrit.common.deprecation import print_deprecation_message
@@ -354,6 +354,57 @@ def __str__(self) -> str:
354354

355355
__repr__ = __str__
356356

357+
@classmethod
358+
def from_dict(cls, data: dict[str, Any]) -> MessagePiece:
359+
"""
360+
Reconstruct a MessagePiece from a dictionary.
361+
362+
Args:
363+
data (dict[str, Any]): Dictionary as produced by to_dict().
364+
365+
Returns:
366+
MessagePiece: Reconstructed instance.
367+
"""
368+
from pyrit.identifiers.component_identifier import ComponentIdentifier
369+
from pyrit.models.score import Score
370+
371+
return cls(
372+
id=data.get("id"),
373+
role=data.get("role", "user"),
374+
conversation_id=data.get("conversation_id"),
375+
sequence=data.get("sequence", -1),
376+
timestamp=(datetime.fromisoformat(str(data["timestamp"])) if data.get("timestamp") else None),
377+
labels=data.get("labels"),
378+
targeted_harm_categories=data.get("targeted_harm_categories"),
379+
prompt_metadata=data.get("prompt_metadata"),
380+
converter_identifiers=(
381+
[ComponentIdentifier.from_dict(c) for c in data["converter_identifiers"]]
382+
if data.get("converter_identifiers")
383+
else None
384+
),
385+
prompt_target_identifier=(
386+
ComponentIdentifier.from_dict(data["prompt_target_identifier"])
387+
if data.get("prompt_target_identifier")
388+
else None
389+
),
390+
attack_identifier=(
391+
ComponentIdentifier.from_dict(data["attack_identifier"]) if data.get("attack_identifier") else None
392+
),
393+
scorer_identifier=(
394+
ComponentIdentifier.from_dict(data["scorer_identifier"]) if data.get("scorer_identifier") else None
395+
),
396+
original_value_data_type=data.get("original_value_data_type", "text"),
397+
original_value=data.get("original_value", ""),
398+
original_value_sha256=data.get("original_value_sha256"),
399+
converted_value_data_type=data.get("converted_value_data_type"),
400+
converted_value=data.get("converted_value"),
401+
converted_value_sha256=data.get("converted_value_sha256"),
402+
response_error=data.get("response_error", "none"),
403+
originator=data.get("originator", "undefined"),
404+
original_prompt_id=(uuid.UUID(str(data["original_prompt_id"])) if data.get("original_prompt_id") else None),
405+
scores=([Score.from_dict(s) for s in data["scores"]] if data.get("scores") else None),
406+
)
407+
357408
def __eq__(self, other: object) -> bool:
358409
"""
359410
Compare this message piece with another for semantic equality.

0 commit comments

Comments
 (0)