106 lines
5.8 KiB
Python
106 lines
5.8 KiB
Python
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
|