from __future__ import annotations from agents.scenario_lab import ScenarioLabService from control_plane.projects.models import ScenarioSuite from graph.lifecycle import ApprovalNode from graph.native_runtime import GraphExecutionContext from graph.registry import NodeHandlerRegistry, NodeResult from graph.spec import ExecutionGraphSpec, GraphEdgeSpec, GraphNodeSpec def scenario_lab_graph_v1() -> ExecutionGraphSpec: nodes = ["prepare", "gather_context", "generate_or_select_scenarios", "validate", "await_approval", "freeze_suite", "execute_scenarios", "collect_results", "classify_findings", "route_findings", "summarize", "complete"] spec = ExecutionGraphSpec( name="scenario_lab", version=1, graph_type="SCENARIO_LAB", entry="prepare", nodes={node: GraphNodeSpec(node, node if node == "complete" else f"scenario_{node}") for node in nodes}, edges=[GraphEdgeSpec("prepare", "gather_context", "success"), GraphEdgeSpec("gather_context", "generate_or_select_scenarios", "success"), GraphEdgeSpec("generate_or_select_scenarios", "validate", "success"), GraphEdgeSpec("validate", "await_approval", "approval_required"), GraphEdgeSpec("validate", "freeze_suite", "success"), GraphEdgeSpec("await_approval", "freeze_suite", "approved"), GraphEdgeSpec("await_approval", "complete", "rejected"), GraphEdgeSpec("freeze_suite", "execute_scenarios", "success"), GraphEdgeSpec("execute_scenarios", "collect_results", "success"), GraphEdgeSpec("collect_results", "classify_findings", "success"), GraphEdgeSpec("classify_findings", "route_findings", "success"), GraphEdgeSpec("route_findings", "summarize", "success"), GraphEdgeSpec("summarize", "complete", "success")], terminal_nodes=["complete"], metadata={"description": "Scenario Lab workflow: safe scenario generation/validation/execution, finding classification and routing."}, ) spec.validate() return spec class ScenarioNode: idempotent = True replay_safe = True destructive = False def __init__(self, service: ScenarioLabService, node_type: str) -> None: self.service = service self.node_type = node_type def suite(self, context: GraphExecutionContext) -> ScenarioSuite: return ScenarioSuite.objects.get(id=context.graph_run.metadata["scenario_suite_id"]) class ScenarioSimpleNode(ScenarioNode): def run(self, context: GraphExecutionContext) -> NodeResult: return NodeResult("COMPLETE", "success") class ScenarioGatherContextNode(ScenarioNode): def run(self, context: GraphExecutionContext) -> NodeResult: suite = self.suite(context) metadata = dict(context.graph_run.metadata) metadata["scenario_context"] = self.service.project_context(suite.project) context.graph_run.metadata = metadata context.graph_run.save(update_fields=["metadata", "updated_at"]) return NodeResult("COMPLETE", "success") class ScenarioGenerateNode(ScenarioNode): def run(self, context: GraphExecutionContext) -> NodeResult: suite = self.suite(context) generated = [] if suite.scenarios.exists() else self.service.generate_scenarios(suite) return NodeResult("COMPLETE", "success", {"generated_count": len(generated), "scenario_count": suite.scenarios.count()}) class ScenarioValidateNode(ScenarioNode): def run(self, context: GraphExecutionContext) -> NodeResult: suite = self.suite(context) valid = self.service.validate_suite(suite) approval_required = any(s.severity in ["HIGH", "CRITICAL"] for s in valid) return NodeResult("COMPLETE", "approval_required" if approval_required else "success", {"valid_count": len(valid), "approval_required": approval_required}) class ScenarioFreezeNode(ScenarioNode): def run(self, context: GraphExecutionContext) -> NodeResult: suite = self.service.freeze_suite(self.suite(context)) return NodeResult("COMPLETE", "success", {"suite_id": str(suite.id), "status": suite.status}) class ScenarioExecuteNode(ScenarioNode): destructive = True def run(self, context: GraphExecutionContext) -> NodeResult: suite = self.suite(context) runs = self.service.execute_suite(suite, graph_run=context.graph_run) return NodeResult("COMPLETE", "success", {"scenario_run_ids": [str(run.id) for run in runs]}) class ScenarioCollectNode(ScenarioNode): def run(self, context: GraphExecutionContext) -> NodeResult: suite = self.suite(context) return NodeResult("COMPLETE", "success", {"results": list(suite.project.scenario_runs.filter(scenario__suite=suite).values("scenario__title", "result", "failure_evidence"))}) class ScenarioRouteNode(ScenarioNode): def run(self, context: GraphExecutionContext) -> NodeResult: routed = self.service.route_findings(self.suite(context)) return NodeResult("COMPLETE", "success", {"routed_count": len(routed)}) class ScenarioSummarizeNode(ScenarioNode): def run(self, context: GraphExecutionContext) -> NodeResult: return NodeResult("COMPLETE", "success", self.service.summarize_suite(self.suite(context))) def scenario_lab_registry(service: ScenarioLabService) -> NodeHandlerRegistry: registry = NodeHandlerRegistry() for handler in [ScenarioSimpleNode(service, "scenario_prepare"), ScenarioGatherContextNode(service, "scenario_gather_context"), ScenarioGenerateNode(service, "scenario_generate_or_select_scenarios"), ScenarioValidateNode(service, "scenario_validate"), ApprovalNode("scenario_await_approval"), ScenarioFreezeNode(service, "scenario_freeze_suite"), ScenarioExecuteNode(service, "scenario_execute_scenarios"), ScenarioCollectNode(service, "scenario_collect_results"), ScenarioSimpleNode(service, "scenario_classify_findings"), ScenarioRouteNode(service, "scenario_route_findings"), ScenarioSummarizeNode(service, "scenario_summarize")]: registry.register(handler) return registry