Skip to content

Commit 8055000

Browse files
Copilotrichlundeen
authored andcommitted
Fix restack integration
Preserve shared scenario configuration resolution and keep progress result mapping compatible with the finalized plan lookup. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com>
1 parent 35b2106 commit 8055000

4 files changed

Lines changed: 36 additions & 28 deletions

File tree

‎frontend/src/App.tsx‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -31,7 +31,7 @@ import {
3131
} from './components/History/scenarioHistoryFilters'
3232
import type { ScenarioHistoryFilters } from './components/History/scenarioHistoryFilters'
3333
import type { ViewName } from './components/Sidebar/Navigation'
34-
import type { AttackSummary, TargetInstance, TargetInfo } from './types'
34+
import type { AttackSummary, TargetInfo } from './types'
3535
import {
3636
targetEndpoint,
3737
targetIdentifierHash,

‎pyrit/backend/services/scenario_run_service.py‎

Lines changed: 17 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -72,8 +72,8 @@
7272
AttackRetrySummary,
7373
RunScenarioRequest,
7474
ScenarioOverloadSummary,
75-
ScenarioRunListItem,
7675
ScenarioRunHeader,
76+
ScenarioRunListItem,
7777
ScenarioRunSummary,
7878
ScenarioTargetSummary,
7979
)
@@ -1267,7 +1267,11 @@ def _resolve_techniques_used(
12671267
Returns:
12681268
list[str]: De-duplicated technique display names.
12691269
"""
1270-
configured = list(dict.fromkeys(scenario_identifier.techniques or [])) if scenario_identifier else []
1270+
configured = (
1271+
list(dict.fromkeys(str(technique) for technique in scenario_identifier.techniques or []))
1272+
if scenario_identifier
1273+
else []
1274+
)
12711275
if configured:
12721276
return configured
12731277
if atomic_groups is not None:
@@ -1979,14 +1983,18 @@ def get_run_progress_from_storage(
19791983
def _map_progress_delta(
19801984
*,
19811985
delta: ScenarioAttackResultDelta,
1982-
plan_lookup: _ScenarioPlanLookup,
1986+
plan: ScenarioRunPlan | None = None,
1987+
plan_lookup: _ScenarioPlanLookup | None = None,
19831988
) -> ScenarioProgressResult:
19841989
"""
19851990
Map a lightweight memory row to its REST progress representation.
19861991
19871992
Returns:
19881993
ScenarioProgressResult: The mapped progress delta.
19891994
"""
1995+
if plan_lookup is None:
1996+
plan_lookup = _ScenarioPlanLookup.from_plan(plan=plan)
1997+
19901998
atomic_attack_name = str(delta.attribution_data.get("parent_collection") or "")
19911999
eval_hash = delta.attribution_data.get("parent_eval_hash")
19922000
atomic_group_id = config_hash(
@@ -2017,8 +2025,7 @@ def _map_progress_delta(
20172025
seed_group_id = config_hash({"objective": delta.objective})
20182026
result_kind, technique_name, attempt_index = ScenarioRunService._progress_result_semantics(
20192027
delta=delta,
2020-
plan=plan,
2021-
atomic_group_id=atomic_group_id,
2028+
group_kind=planned_group.group_kind if planned_group is not None else None,
20222029
atomic_attack_name=atomic_attack_name,
20232030
)
20242031
return ScenarioProgressResult(
@@ -2042,8 +2049,7 @@ def _map_progress_delta(
20422049
def _progress_result_semantics(
20432050
*,
20442051
delta: ScenarioAttackResultDelta,
2045-
plan: ScenarioRunPlan | None,
2046-
atomic_group_id: str,
2052+
group_kind: ScenarioRunPlanGroupKind | None,
20472053
atomic_attack_name: str,
20482054
) -> tuple[ScenarioProgressResultKind, str | None, int | None]:
20492055
"""
@@ -2064,17 +2070,14 @@ def _progress_result_semantics(
20642070
conversation_id=delta.conversation_id,
20652071
atomic_attack_identifier=delta.atomic_attack_identifier,
20662072
)
2067-
matching_group = (
2068-
next((group for group in plan.atomic_groups if group.id == atomic_group_id), None) if plan else None
2069-
)
2070-
if matching_group is not None:
2071-
if matching_group.group_kind is ScenarioRunPlanGroupKind.DIRECT_BASELINE:
2073+
if group_kind is not None:
2074+
if group_kind is ScenarioRunPlanGroupKind.DIRECT_BASELINE:
20722075
return ScenarioProgressResultKind.DIRECT_BASELINE, None, None
2073-
if matching_group.group_kind is ScenarioRunPlanGroupKind.ADAPTIVE:
2076+
if group_kind is ScenarioRunPlanGroupKind.ADAPTIVE:
20742077
return ScenarioProgressResultKind.ADAPTIVE_ORCHESTRATION, None, None
20752078
if is_sequential_envelope:
20762079
return ScenarioProgressResultKind.AGGREGATE_PARENT, None, None
2077-
if matching_group.group_kind is ScenarioRunPlanGroupKind.ATTACK:
2080+
if group_kind is ScenarioRunPlanGroupKind.ATTACK:
20782081
return ScenarioProgressResultKind.ATTACK, None, None
20792082

20802083
if atomic_attack_name == "baseline":

‎pyrit/models/catalog/__init__.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,7 @@
3131
ScenarioDatasetSummary,
3232
ScenarioDefaultRunSizeEstimate,
3333
ScenarioOverloadSummary,
34+
ScenarioRunHeader,
3435
ScenarioRunListItem,
3536
ScenarioRunSizeComponent,
3637
ScenarioRunSizeEstimate,
@@ -56,6 +57,7 @@
5657
"ScenarioDatasetSummary": "pyrit.models.catalog.scenario",
5758
"ScenarioDefaultRunSizeEstimate": "pyrit.models.catalog.scenario",
5859
"ScenarioOverloadSummary": "pyrit.models.catalog.scenario",
60+
"ScenarioRunHeader": "pyrit.models.catalog.scenario",
5961
"ScenarioRunListItem": "pyrit.models.catalog.scenario",
6062
"ScenarioRunSizeComponent": "pyrit.models.catalog.scenario",
6163
"ScenarioRunSizeEstimate": "pyrit.models.catalog.scenario",

‎tests/unit/backend/test_scenario_run_service.py‎

Lines changed: 16 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -539,8 +539,8 @@ async def test_start_run_forwards_include_baseline(self, mock_all_registries) ->
539539
init_call = mock_all_registries["scenario_registry"].create_and_initialize_async.await_args
540540
assert init_call.kwargs["include_baseline"] is False
541541

542-
async def test_start_run_max_dataset_size_copies_default_config(self, mock_all_registries) -> None:
543-
"""``max_dataset_size`` overrides an independent copy of the scenario default."""
542+
async def test_start_run_max_dataset_size_updates_introspection_config(self, mock_all_registries) -> None:
543+
"""``max_dataset_size`` updates the throwaway introspection config."""
544544
default_config = DatasetAttackConfiguration(dataset_names=["original"], max_dataset_size=100)
545545
scenario_instance = mock_all_registries["scenario_instance"]
546546
scenario_instance._default_dataset_config = default_config
@@ -550,10 +550,10 @@ async def test_start_run_max_dataset_size_copies_default_config(self, mock_all_r
550550

551551
init_call = mock_all_registries["scenario_registry"].create_and_initialize_async.await_args
552552
built_config = init_call.kwargs["dataset_config"]
553-
assert built_config is not default_config
553+
assert built_config is default_config
554554
assert type(built_config) is DatasetAttackConfiguration
555555
assert built_config.max_dataset_size == 5
556-
assert default_config.max_dataset_size == 100
556+
assert default_config.max_dataset_size == 5
557557

558558
async def test_start_run_dataset_names_preserves_subclass_config_type(self, mock_all_registries) -> None:
559559
"""``dataset_names`` rebuilds the config using the scenario's own DatasetConfiguration subclass.
@@ -641,10 +641,10 @@ async def test_start_run_max_dataset_size_updates_each_default_compound_child(se
641641
init_call = mock_all_registries["scenario_registry"].create_and_initialize_async.await_args
642642
built_config = init_call.kwargs["dataset_config"]
643643
assert isinstance(built_config, CompoundDatasetAttackConfiguration)
644-
assert built_config is not default_config
644+
assert built_config is default_config
645645
assert built_config.dataset_names == ["airt_hate", "airt_fairness"]
646646
assert [child.max_dataset_size for child in built_config._configurations] == [2, 2]
647-
assert [child.max_dataset_size for child in default_config._configurations] == [4, 4]
647+
assert [child.max_dataset_size for child in default_config._configurations] == [2, 2]
648648

649649
async def test_start_run_non_name_overrides_preserve_shaped_compound_children(self, mock_all_registries) -> None:
650650
"""Size and filter overrides do not rebuild scenario-specific child configurations."""
@@ -672,7 +672,7 @@ class _ShapedDatasetConfiguration(DatasetAttackConfiguration):
672672

673673
init_call = mock_all_registries["scenario_registry"].create_and_initialize_async.await_args
674674
built_config = init_call.kwargs["dataset_config"]
675-
assert built_config is not default_config
675+
assert built_config is default_config
676676
assert [type(child) for child in built_config._configurations] == [
677677
_ShapedDatasetConfiguration,
678678
_ShapedDatasetConfiguration,
@@ -682,8 +682,11 @@ class _ShapedDatasetConfiguration(DatasetAttackConfiguration):
682682
{"harm_categories": ["cyber"]},
683683
{"harm_categories": ["cyber"]},
684684
]
685-
assert [child.max_dataset_size for child in default_config._configurations] == [4, 4]
686-
assert [child.filters for child in default_config._configurations] == [{}, {}]
685+
assert [child.max_dataset_size for child in default_config._configurations] == [2, 2]
686+
assert [child.filters for child in default_config._configurations] == [
687+
{"harm_categories": ["cyber"]},
688+
{"harm_categories": ["cyber"]},
689+
]
687690

688691
async def test_start_run_dataset_names_rejects_incompatible_subclass_constructor(self, mock_all_registries) -> None:
689692
"""Reject overrides that cannot preserve scenario-specific dataset configuration."""
@@ -732,8 +735,8 @@ class _MarkerDatasetConfiguration(DatasetConfiguration):
732735
assert built_config.max_dataset_size == 7
733736
assert built_config.filters == {"harm_categories": ["cyber"]}
734737

735-
async def test_start_run_dataset_filters_copy_default_config(self, mock_all_registries) -> None:
736-
"""``dataset_filters`` with no names merges filters into an independent copy."""
738+
async def test_start_run_dataset_filters_update_introspection_config(self, mock_all_registries) -> None:
739+
"""``dataset_filters`` with no names update the throwaway introspection config."""
737740
default_config = DatasetAttackConfiguration(dataset_names=["original"])
738741
scenario_instance = mock_all_registries["scenario_instance"]
739742
scenario_instance._default_dataset_config = default_config
@@ -743,9 +746,9 @@ async def test_start_run_dataset_filters_copy_default_config(self, mock_all_regi
743746

744747
init_call = mock_all_registries["scenario_registry"].create_and_initialize_async.await_args
745748
built_config = init_call.kwargs["dataset_config"]
746-
assert built_config is not default_config
749+
assert built_config is default_config
747750
assert built_config.filters == {"harm_categories": ["cyber"]}
748-
assert default_config.filters == {}
751+
assert default_config.filters == {"harm_categories": ["cyber"]}
749752

750753
async def test_start_run_dataset_names_introspection_failure_raises(self, mock_memory) -> None:
751754
"""Passing ``dataset_names`` against a non-no-arg-instantiable scenario fails fast."""

0 commit comments

Comments
 (0)