Skip to content

Commit b76704a

Browse files
authored
MAINT: Adding simulated assistant role (microsoft#1292)
1 parent 3813661 commit b76704a

25 files changed

Lines changed: 387 additions & 64 deletions

pyrit/executor/attack/component/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@
77
ConversationManager,
88
ConversationState,
99
format_conversation_context,
10+
mark_messages_as_simulated,
1011
)
1112
from pyrit.executor.attack.component.objective_evaluator import ObjectiveEvaluator
1213
from pyrit.executor.attack.component.simulated_conversation import (
@@ -19,6 +20,7 @@
1920
"ConversationManager",
2021
"ConversationState",
2122
"format_conversation_context",
23+
"mark_messages_as_simulated",
2224
"ObjectiveEvaluator",
2325
"generate_simulated_conversation_async",
2426
"SimulatedConversationResult",

pyrit/executor/attack/component/conversation_manager.py

Lines changed: 42 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@
44
import logging
55
import uuid
66
from dataclasses import dataclass, field
7-
from typing import Dict, List, Optional
7+
from typing import Dict, List, Optional, Sequence
88

99
from pyrit.memory import CentralMemory
1010
from pyrit.models import ChatMessageRole, Message, MessagePiece, Score
@@ -18,6 +18,29 @@
1818
logger = logging.getLogger(__name__)
1919

2020

21+
def mark_messages_as_simulated(messages: Sequence[Message]) -> List[Message]:
22+
"""
23+
Mark assistant messages as simulated_assistant for traceability.
24+
25+
This function converts all assistant roles to simulated_assistant in the
26+
provided messages. This is useful when loading conversations from YAML files
27+
or other sources where the responses are not from actual targets.
28+
29+
Args:
30+
messages (Sequence[Message]): The messages to mark as simulated.
31+
32+
Returns:
33+
List[Message]: The same messages with assistant roles converted to simulated_assistant.
34+
Modifies the messages in place and also returns them for convenience.
35+
"""
36+
result = list(messages)
37+
for message in result:
38+
for piece in message.message_pieces:
39+
if piece._role == "assistant":
40+
piece._role = "simulated_assistant"
41+
return result
42+
43+
2144
def format_conversation_context(messages: List[Message]) -> str:
2245
"""
2346
Format a list of messages into a context string for adversarial chat system prompts.
@@ -55,17 +78,22 @@ def format_conversation_context(messages: List[Message]) -> str:
5578
piece = message.get_piece()
5679

5780
# Skip system messages - they're handled separately
58-
if piece.role == "system":
81+
if piece.api_role == "system":
5982
continue
6083

6184
# Start a new turn when we see a user message
62-
if piece.role == "user":
85+
if piece.api_role == "user":
6386
turn_number += 1
6487
context_parts.append(f"Turn {turn_number}:")
6588

6689
# Format the piece content
6790
content = _format_piece_content(piece)
68-
role_label = "User" if piece.role == "user" else "Assistant"
91+
if piece.api_role == "user":
92+
role_label = "User"
93+
elif piece.is_simulated:
94+
role_label = "Assistant (simulated)"
95+
else:
96+
role_label = "Assistant"
6997
context_parts.append(f"{role_label}: {content}")
7098

7199
return "\n".join(context_parts)
@@ -175,7 +203,7 @@ def get_last_message(
175203
if role:
176204
for m in reversed(conversation):
177205
piece = m.get_piece()
178-
if piece.role == role:
206+
if piece.api_role == role:
179207
return piece
180208
return None
181209

@@ -288,7 +316,7 @@ async def update_conversation_state_async(
288316
# Determine if we should exclude the last message (if it's a user message in multi-turn context)
289317
last_message = valid_requests[-1].message_pieces[0]
290318
is_multi_turn = max_turns is not None
291-
should_exclude_last = is_multi_turn and last_message.role == "user"
319+
should_exclude_last = is_multi_turn and last_message.api_role == "user"
292320

293321
# Process all messages except potentially the last one
294322
for i, request in enumerate(valid_requests):
@@ -351,9 +379,9 @@ async def _apply_role_specific_converters_async(
351379
for piece in request.message_pieces:
352380
applicable_converters: Optional[List[PromptConverterConfiguration]] = None
353381

354-
if piece.role == "user" and request_converters:
382+
if piece.api_role == "user" and request_converters:
355383
applicable_converters = request_converters
356-
elif piece.role == "assistant" and response_converters:
384+
elif piece.api_role == "assistant" and response_converters:
357385
applicable_converters = response_converters
358386
# System messages get no converters (applicable_converters remains None)
359387

@@ -429,7 +457,7 @@ def _process_piece(
429457
is_multi_turn = max_turns is not None
430458

431459
# Only assistant messages count as turns
432-
if piece.role == "assistant" and is_multi_turn:
460+
if piece.api_role == "assistant" and is_multi_turn:
433461
conversation_state.turn_count += 1
434462

435463
if conversation_state.turn_count > max_turns:
@@ -466,11 +494,11 @@ async def _populate_conversation_state_async(
466494
return # Nothing to extract from empty history
467495

468496
# Extract the last user message and assistant message scores from the last message
469-
if last_message.role == "user":
497+
if last_message.api_role == "user":
470498
conversation_state.last_user_message = last_message.converted_value
471499
logger.debug(f"Extracted last user message: {conversation_state.last_user_message[:50]}...")
472500

473-
elif last_message.role == "assistant":
501+
elif last_message.api_role == "assistant":
474502
# Get scores for the last assistant message based off of the original id
475503
conversation_state.last_assistant_message_scores = list(
476504
self._memory.get_prompt_scores(prompt_ids=[str(last_message.original_prompt_id)])
@@ -482,7 +510,7 @@ async def _populate_conversation_state_async(
482510
return
483511

484512
# Check assumption that there will be a user message preceding the assistant message
485-
if len(prepended_conversation) > 1 and prepended_conversation[-2].get_piece().role == "user":
513+
if len(prepended_conversation) > 1 and prepended_conversation[-2].get_piece().api_role == "user":
486514
conversation_state.last_user_message = prepended_conversation[-2].get_value()
487515
logger.debug(f"Extracted preceding user message: {conversation_state.last_user_message[:50]}...")
488516
else:
@@ -533,11 +561,11 @@ async def prepend_to_adversarial_chat_async(
533561
for message in prepended_conversation:
534562
for piece in message.message_pieces:
535563
# Skip system messages - adversarial chat has its own system prompt
536-
if piece.role == "system":
564+
if piece.api_role == "system":
537565
continue
538566

539567
# Create a new piece with swapped role for adversarial chat
540-
swapped_role = role_swap.get(piece.role, piece.role)
568+
swapped_role = role_swap.get(piece.api_role, piece.api_role)
541569

542570
adversarial_piece = MessagePiece(
543571
id=uuid.uuid4(),

pyrit/executor/attack/component/simulated_conversation.py

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -63,7 +63,7 @@ def _effective_turn_index(self) -> int:
6363
# Calculate total complete turns (user+assistant pairs)
6464
total_turns = len(self.conversation) // 2
6565
# Account for trailing user message (incomplete turn)
66-
if len(self.conversation) % 2 == 1 and self.conversation[-1].role == "user":
66+
if len(self.conversation) % 2 == 1 and self.conversation[-1].api_role == "user":
6767
total_turns += 1
6868

6969
if self.turn_index is None:
@@ -108,7 +108,7 @@ def next_message(self) -> Optional[Message]:
108108
return None
109109
# User message for turn N is at index (N-1) * 2
110110
user_idx = (turn - 1) * 2
111-
if user_idx < len(self.conversation) and self.conversation[user_idx].role == "user":
111+
if user_idx < len(self.conversation) and self.conversation[user_idx].api_role == "user":
112112
return self.conversation[user_idx].duplicate_message()
113113
return None
114114

@@ -229,9 +229,14 @@ async def generate_simulated_conversation_async(
229229

230230
# Filter out system messages - prepended_conversation should only have user/assistant turns
231231
# System prompts are set separately on each target during attack execution
232+
# Also mark assistant messages as simulated for traceability
232233
filtered_messages: List[Message] = []
233234
for message in raw_messages:
234-
if message.role != "system":
235+
if message.api_role != "system":
236+
# Mark assistant responses as simulated since this is a simulated conversation
237+
if message.api_role == "assistant":
238+
for piece in message.message_pieces:
239+
piece._role = "simulated_assistant"
235240
filtered_messages.append(message)
236241

237242
# Get the score from the result (there should be one score for the last turn)

pyrit/executor/attack/multi_turn/tree_of_attacks.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -777,7 +777,7 @@ async def _generate_subsequent_turn_prompt_async(self, objective: str) -> str:
777777
target_messages = self._memory.get_conversation(conversation_id=self.objective_target_conversation_id)
778778

779779
# Extract the last assistant response
780-
assistant_responses = [r for r in target_messages if r.get_piece().role == "assistant"]
780+
assistant_responses = [r for r in target_messages if r.get_piece().api_role == "assistant"]
781781
if not assistant_responses:
782782
logger.error(f"No assistant responses found in the conversation {self.objective_target_conversation_id}.")
783783
raise RuntimeError("Cannot proceed without an assistant response.")

pyrit/executor/attack/printer/console_printer.py

Lines changed: 7 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -152,14 +152,14 @@ async def print_messages_async(
152152
turn_number = 0
153153
for message in messages:
154154
# Increment turn number once per message with role="user"
155-
if message.role == "user":
155+
if message.api_role == "user":
156156
turn_number += 1
157157
# User message header
158158
print()
159159
self._print_colored("─" * self._width, Fore.BLUE)
160160
self._print_colored(f"🔹 Turn {turn_number} - USER", Style.BRIGHT, Fore.BLUE)
161161
self._print_colored("─" * self._width, Fore.BLUE)
162-
elif message.role == "system":
162+
elif message.api_role == "system":
163163
# System message header (not counted as a turn)
164164
print()
165165
self._print_colored("─" * self._width, Fore.MAGENTA)
@@ -169,7 +169,8 @@ async def print_messages_async(
169169
# Assistant or other role message header
170170
print()
171171
self._print_colored("─" * self._width, Fore.YELLOW)
172-
self._print_colored(f"🔸 {message.role.upper()}", Style.BRIGHT, Fore.YELLOW)
172+
role_label = "ASSISTANT (SIMULATED)" if message.is_simulated else message.api_role.upper()
173+
self._print_colored(f"🔸 {role_label}", Style.BRIGHT, Fore.YELLOW)
173174
self._print_colored("─" * self._width, Fore.YELLOW)
174175

175176
# Now print all pieces in this message
@@ -179,15 +180,15 @@ async def print_messages_async(
179180
continue
180181

181182
# Handle converted values for user messages
182-
if piece.role == "user" and piece.converted_value != piece.original_value:
183+
if piece.api_role == "user" and piece.converted_value != piece.original_value:
183184
self._print_colored(f"{self._indent} Original:", Fore.CYAN)
184185
self._print_wrapped_text(piece.original_value, Fore.WHITE)
185186
print()
186187
self._print_colored(f"{self._indent} Converted:", Fore.CYAN)
187188
self._print_wrapped_text(piece.converted_value, Fore.WHITE)
188-
elif piece.role == "user":
189+
elif piece.api_role == "user":
189190
self._print_wrapped_text(piece.converted_value, Fore.BLUE)
190-
elif piece.role == "system":
191+
elif piece.api_role == "system":
191192
self._print_wrapped_text(piece.converted_value, Fore.MAGENTA)
192193
else:
193194
self._print_wrapped_text(piece.converted_value, Fore.YELLOW)

pyrit/executor/attack/printer/markdown_printer.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -213,7 +213,7 @@ async def _get_conversation_markdown_async(
213213
if not message.message_pieces:
214214
continue
215215

216-
message_role = message.get_piece().role
216+
message_role = message.get_piece().api_role
217217

218218
if message_role == "system":
219219
markdown_lines.extend(self._format_system_message(message))
@@ -284,7 +284,8 @@ async def _format_assistant_message_async(self, *, message: Message) -> List[str
284284
List[str]: List of markdown strings representing the response message.
285285
"""
286286
lines = []
287-
role_name = message.message_pieces[0].role.capitalize()
287+
piece = message.message_pieces[0]
288+
role_name = "Assistant (Simulated)" if piece.is_simulated else piece.api_role.capitalize()
288289

289290
lines.append(f"\n#### {role_name}\n")
290291

pyrit/executor/promptgen/fuzzer/fuzzer.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -452,11 +452,11 @@ def _print_conversations(self, result: FuzzerResult) -> None:
452452
continue
453453

454454
for message in target_messages:
455-
if message.role == "user":
455+
if message.api_role == "user":
456456
self._print_colored(f"{self._indent * 2} USER:", Style.BRIGHT, Fore.BLUE)
457457
self._print_wrapped_text(message.converted_value, Fore.BLUE)
458458
else:
459-
self._print_colored(f"{self._indent * 2} {message.role.upper()}:", Style.BRIGHT, Fore.YELLOW)
459+
self._print_colored(f"{self._indent * 2} {message.api_role.upper()}:", Style.BRIGHT, Fore.YELLOW)
460460
self._print_wrapped_text(message.converted_value, Fore.YELLOW)
461461

462462
# Print scores if available

pyrit/memory/memory_interface.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -471,7 +471,7 @@ def get_request_from_response(self, *, response: Message) -> Message:
471471
Raises:
472472
ValueError: If the response is not from an assistant role or has no preceding request.
473473
"""
474-
if response.role != "assistant":
474+
if response.api_role != "assistant":
475475
raise ValueError("The provided request is not a response (role must be 'assistant').")
476476
if response.sequence < 1:
477477
raise ValueError("The provided request does not have a preceding request (sequence < 1).")
@@ -628,7 +628,7 @@ def duplicate_conversation_excluding_last_turn(self, *, conversation_id: str) ->
628628

629629
length_of_sequence_to_remove = 0
630630

631-
if last_message.role == "system" or last_message.role == "user":
631+
if last_message.api_role == "system" or last_message.api_role == "user":
632632
length_of_sequence_to_remove = 1
633633
else:
634634
length_of_sequence_to_remove = 2
@@ -778,7 +778,7 @@ def get_chat_messages_with_conversation_id(self, *, conversation_id: str) -> Seq
778778
Sequence[ChatMessage]: The list of chat messages.
779779
"""
780780
memory_entries = self.get_message_pieces(conversation_id=conversation_id)
781-
return [ChatMessage(role=me.role, content=me.converted_value) for me in memory_entries] # type: ignore
781+
return [ChatMessage(role=me.api_role, content=me.converted_value) for me in memory_entries] # type: ignore
782782

783783
def get_seeds(
784784
self,

pyrit/memory/memory_models.py

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -149,7 +149,9 @@ class PromptMemoryEntry(Base):
149149
__tablename__ = "PromptMemoryEntries"
150150
__table_args__ = {"extend_existing": True}
151151
id = mapped_column(CustomUUID, nullable=False, primary_key=True)
152-
role: Mapped[Literal["system", "user", "assistant", "tool", "developer"]] = mapped_column(String, nullable=False)
152+
role: Mapped[Literal["system", "user", "assistant", "simulated_assistant", "tool", "developer"]] = mapped_column(
153+
String, nullable=False
154+
)
153155
conversation_id = mapped_column(String, nullable=False)
154156
sequence = mapped_column(INTEGER, nullable=False)
155157
timestamp = mapped_column(DateTime, nullable=False)
@@ -192,7 +194,7 @@ def __init__(self, *, entry: MessagePiece):
192194
entry (MessagePiece): The message piece to convert into a database entry.
193195
"""
194196
self.id = entry.id
195-
self.role = entry.role
197+
self.role = entry._role
196198
self.conversation_id = entry.conversation_id
197199
self.sequence = entry.sequence
198200
self.timestamp = entry.timestamp

pyrit/models/chat_message.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@
77

88
from pyrit.models.literals import ChatMessageRole
99

10-
ALLOWED_CHAT_MESSAGE_ROLES = ["system", "user", "assistant", "tool", "developer"]
10+
ALLOWED_CHAT_MESSAGE_ROLES = ["system", "user", "assistant", "simulated_assistant", "tool", "developer"]
1111

1212

1313
class ToolCall(BaseModel):

0 commit comments

Comments
 (0)