Artifex/graph/scenario_lab.py
2026-08-15 19:46:07 +07:00

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