Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
90 changes: 84 additions & 6 deletions pyrit/executor/attack/core/attack_executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -195,8 +195,10 @@ async def execute_attack_from_seed_groups_async(
field_overrides: Optional per-seed-group field overrides. If provided,
must match the length of seed_groups. Each dict is passed to
from_seed_group() as overrides.
return_partial_on_failure: If True, returns partial results when some
objectives fail. If False (default), raises the first exception.
return_partial_on_failure: If True, executes successfully constructed
objectives and returns parameter-build or execution failures alongside
completed results. If False (default), a parameter-build failure
suppresses execution and raises the first exception by input order.
attribution: Optional ``AttackResultAttribution`` stamped onto every
per-task ``AttackContext`` so the persisted ``AttackResultEntry``
row carries ``attribution_parent_id`` + ``attribution_data``.
Expand Down Expand Up @@ -239,13 +241,39 @@ async def build_params_async(i: int, sg: AttackSeedGroup) -> AttackParameters:
**combined_overrides,
)

params_list = list(await asyncio.gather(*[build_params_async(i, sg) for i, sg in enumerate(seed_groups)]))
build_results = list(
await asyncio.gather(
*[build_params_async(i, sg) for i, sg in enumerate(seed_groups)],
return_exceptions=True,
)
)
params_list: list[AttackParameters] = []
successful_input_indices: list[int] = []
build_failures: list[tuple[int, str, Exception]] = []
for index, (seed_group, build_result) in enumerate(zip(seed_groups, build_results, strict=True)):
if isinstance(build_result, Exception):
assert seed_group.objective is not None
build_failures.append((index, seed_group.objective.value, build_result))
elif isinstance(build_result, BaseException):
raise build_result
else:
params_list.append(build_result)
successful_input_indices.append(index)

if build_failures and not return_partial_on_failure:
raise build_failures[0][2]

return await self._execute_with_params_list_async(
execution_result = await self._execute_with_params_list_async(
attack=attack,
params_list=params_list,
return_partial_on_failure=return_partial_on_failure,
attribution=attribution,
input_indices=successful_input_indices,
)
return self._merge_parameter_build_failures(
build_failures=build_failures,
successful_input_indices=successful_input_indices,
execution_result=execution_result,
)

async def execute_attack_async(
Expand Down Expand Up @@ -323,6 +351,7 @@ async def _execute_with_params_list_async(
params_list: Sequence[AttackParameters],
return_partial_on_failure: bool = False,
attribution: AttackResultAttribution | None = None,
input_indices: Sequence[int] | None = None,
) -> AttackExecutorResult[AttackStrategyResultT]:
"""
Execute attacks in parallel with a list of pre-built parameters.
Expand All @@ -337,6 +366,8 @@ async def _execute_with_params_list_async(
attribution: Optional ``AttackResultAttribution`` stamped onto every
per-task ``AttackContext`` so the persistence path can record
orchestrator linkage.
input_indices: Original input positions for ``params_list``. Defaults
to sequential positions when parameters were constructed directly.

Returns:
AttackExecutorResult with completed results and any incomplete objectives.
Expand All @@ -357,6 +388,7 @@ async def run_one_async(index: int, params: AttackParameters) -> AttackStrategyR
objectives=[p.objective for p in params_list],
results_or_exceptions=list(results_or_exceptions),
return_partial_on_failure=return_partial_on_failure,
input_indices=input_indices,
)

def _process_execution_results(
Expand All @@ -365,6 +397,7 @@ def _process_execution_results(
objectives: Sequence[str],
results_or_exceptions: list[Any],
return_partial_on_failure: bool,
input_indices: Sequence[int] | None = None,
) -> AttackExecutorResult[AttackStrategyResultT]:
"""
Process results from parallel execution into an AttackExecutorResult.
Expand All @@ -373,23 +406,30 @@ def _process_execution_results(
objectives: The objectives that were executed.
results_or_exceptions: Results or exceptions from asyncio.gather.
return_partial_on_failure: Whether to return partial results on failure.
input_indices: Original input positions corresponding to ``objectives``.

Returns:
AttackExecutorResult with completed and incomplete results.

Raises:
BaseException: If return_partial_on_failure=False and any failed.
ValueError: If input_indices length doesn't match objectives length.
"""
completed: list[AttackStrategyResultT] = []
incomplete: list[tuple[str, BaseException]] = []
completed_indices: list[int] = []
source_indices = list(input_indices) if input_indices is not None else list(range(len(objectives)))
if len(source_indices) != len(objectives):
raise ValueError("input_indices length must match objectives length")

self._raise_first_fatal_exception(results_or_exceptions=results_or_exceptions)

for i, (objective, result) in enumerate(zip(objectives, results_or_exceptions, strict=False)):
for i, (objective, result) in enumerate(zip(objectives, results_or_exceptions, strict=True)):
if isinstance(result, BaseException):
incomplete.append((objective, result))
else:
completed.append(result)
completed_indices.append(i)
completed_indices.append(source_indices[i])

executor_result: AttackExecutorResult[AttackStrategyResultT] = AttackExecutorResult(
completed_results=completed,
Expand All @@ -401,3 +441,41 @@ def _process_execution_results(
executor_result.raise_if_incomplete()

return executor_result

@staticmethod
def _raise_first_fatal_exception(*, results_or_exceptions: Sequence[Any]) -> None:
"""Propagate cancellation and other non-recoverable base exceptions."""
for result in results_or_exceptions:
if isinstance(result, BaseException) and not isinstance(result, Exception):
raise result

@staticmethod
def _merge_parameter_build_failures(
*,
build_failures: Sequence[tuple[int, str, Exception]],
successful_input_indices: Sequence[int],
execution_result: AttackExecutorResult[AttackStrategyResultT],
) -> AttackExecutorResult[AttackStrategyResultT]:
"""
Merge parameter-build and execution failures by original input position.

Returns:
The combined executor result.
"""
if not build_failures:
return execution_result

incomplete_by_index: dict[int, tuple[str, BaseException]] = {
index: (objective, error) for index, objective, error in build_failures
}
completed_indices = set(execution_result.input_indices)
execution_failures = iter(execution_result.incomplete_objectives)
for input_index in successful_input_indices:
if input_index not in completed_indices:
incomplete_by_index[input_index] = next(execution_failures)

return AttackExecutorResult(
completed_results=execution_result.completed_results,
incomplete_objectives=[incomplete_by_index[index] for index in sorted(incomplete_by_index)],
input_indices=execution_result.input_indices,
)
Loading