From 9338b24cc1523a2b5742310d7eb39efa0f75f30e Mon Sep 17 00:00:00 2001 From: Copilot <223556219+Copilot@users.noreply.github.com> Date: Tue, 4 Aug 2026 01:36:19 -0700 Subject: [PATCH] FIX: Correlate multi-turn attack errors Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Copilot-Session: 0d1f97bd-49a5-4a6d-b814-fecef1e2a8cb --- pyrit/executor/attack/core/attack_strategy.py | 7 ++++-- .../attack/core/test_attack_strategy.py | 22 +++++++++++++++++++ 2 files changed, 27 insertions(+), 2 deletions(-) diff --git a/pyrit/executor/attack/core/attack_strategy.py b/pyrit/executor/attack/core/attack_strategy.py index 6bab4492b7..11c58b60ad 100644 --- a/pyrit/executor/attack/core/attack_strategy.py +++ b/pyrit/executor/attack/core/attack_strategy.py @@ -355,8 +355,11 @@ async def _on_error_async( collector = get_retry_collector() retry_events = collector.events if collector else [] - # Build a conversation_id — use context's if available, otherwise generate one - conversation_id = getattr(context, "conversation_id", None) or str(uuid.uuid4()) + # Multi-turn contexts keep the active ID on their conversation session. + conversation_id = getattr(context, "conversation_id", None) + if not conversation_id: + conversation_id = getattr(getattr(context, "session", None), "conversation_id", None) + conversation_id = conversation_id or str(uuid.uuid4()) error_result = AttackResult( conversation_id=conversation_id, diff --git a/tests/unit/executor/attack/core/test_attack_strategy.py b/tests/unit/executor/attack/core/test_attack_strategy.py index f4f8701245..433f6f3df9 100644 --- a/tests/unit/executor/attack/core/test_attack_strategy.py +++ b/tests/unit/executor/attack/core/test_attack_strategy.py @@ -15,6 +15,7 @@ AttackStrategy, _DefaultAttackStrategyEventHandler, ) +from pyrit.executor.attack.multi_turn.multi_turn_attack_strategy import ConversationSession, MultiTurnAttackContext from pyrit.executor.core import StrategyEvent, StrategyEventData from pyrit.memory.central_memory import CentralMemory from pyrit.models import ( @@ -631,6 +632,27 @@ async def test_on_error_persists_result_to_memory(self, sample_attack_context, m assert stored_result.error_type == "ValueError" assert stored_result.execution_time_ms == 500 + async def test_on_error_uses_multi_turn_session_conversation_id(self, mock_memory): + """Test that multi-turn failures remain correlated with their active conversation.""" + context = MultiTurnAttackContext( + params=AttackParameters(objective="Test harmful objective"), + session=ConversationSession(conversation_id="active-conversation-id"), + ) + + with patch("pyrit.memory.central_memory.CentralMemory.get_memory_instance", return_value=mock_memory): + handler = _DefaultAttackStrategyEventHandler() + event_data = StrategyEventData( + event=StrategyEvent.ON_ERROR, + strategy_name="TestStrategy", + strategy_id="test-id", + context=context, + error=TimeoutError("target timed out"), + ) + await handler.on_event_async(event_data) + + stored_result = mock_memory.add_attack_results_to_memory.call_args.kwargs["attack_results"][0] + assert stored_result.conversation_id == "active-conversation-id" + async def test_on_error_skips_when_no_error_or_context(self, mock_memory): """Test that error handler returns early when error or context is None""" with patch("pyrit.memory.central_memory.CentralMemory.get_memory_instance", return_value=mock_memory):