Skip to content
Draft
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
17 changes: 17 additions & 0 deletions src/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -520,6 +520,23 @@ class AMReXAgentConfig(BaseModel):
description="Auto-approve pre-confirmation gates without prompting."
)

gate_strategy: Literal["auto", "terminal", "selective"] = Field(
default="auto",
description="Decision-level gating strategy for GateManager."
)
gate_points: List[str] = Field(
default_factory=list,
description="Decision points to gate (e.g., solver, baseline, modifications, execution)."
)
router_gate_strategy: Literal["off", "terminal", "selective"] = Field(
default="off",
description="Router-level gating strategy for interactivity gates."
)
router_gate_points: List[str] = Field(
default_factory=list,
description="Router steps to gate (e.g., architect, reviewer, input_writer, runner, analysis)."
)

run_mode: Literal["dry", "stage", "submit", "full"] = Field(
default="full",
description="Run execution strategy: dry (scripts only), stage (stage inputs only), "
Expand Down
121 changes: 116 additions & 5 deletions src/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,11 +30,15 @@
reviewer_node,
runner_node,
visualization_node, # Phase 4
router_gate_node,
)
from src.router_func import (
route_after_analysis, # Phase 4
route_after_architect,
route_after_input_writer,
route_after_reviewer,
route_after_runner, # Phase 4
route_after_router_gate,
)


Expand Down Expand Up @@ -437,6 +441,28 @@ def parse_arguments(args: list[str] | None = None) -> argparse.Namespace:
action='store_true',
help='Pause for a pre-confirmation gate before validation'
)
parser.add_argument(
'--gate-strategy',
choices=['auto', 'terminal', 'selective'],
dest='gate_strategy',
help='Decision gate strategy: auto, terminal, selective'
)
parser.add_argument(
'--gate-points',
dest='gate_points',
help='Comma-separated decision points to gate (solver,baseline,modifications,execution,input_writer,analysis,visualization)'
)
parser.add_argument(
'--router-gate-strategy',
choices=['off', 'terminal', 'selective'],
dest='router_gate_strategy',
help='Router gate strategy: off, terminal, selective'
)
parser.add_argument(
'--router-gate-points',
dest='router_gate_points',
help='Comma-separated router steps to gate (architect,reviewer,input_writer,runner,analysis,visualization)'
)
parser.add_argument(
'--llm-gate-strategy',
choices=[
Expand Down Expand Up @@ -559,6 +585,48 @@ def load_prompt_content(args: argparse.Namespace) -> str:
return path.read_text().strip()


def apply_gate_cli_settings(config, parsed_args) -> None:
if getattr(parsed_args, "preconfirm", False):
logger.warning(
"Preconfirm gating is deprecated; prefer router gating for new workflows."
)
config.preconfirm_gate = True

gate_strategy = getattr(parsed_args, "gate_strategy", None)
gate_points_raw = getattr(parsed_args, "gate_points", None)
gate_points = []
if gate_points_raw:
gate_points = [p.strip() for p in gate_points_raw.split(",") if p.strip()]

if gate_strategy:
config.gate_strategy = gate_strategy
elif gate_points:
config.gate_strategy = "selective"
else:
config.gate_strategy = getattr(config, "gate_strategy", "auto") or "auto"

config.gate_points = gate_points

router_gate_strategy = getattr(parsed_args, "router_gate_strategy", None)
router_gate_points_raw = getattr(parsed_args, "router_gate_points", None)
router_gate_points = []
if router_gate_points_raw:
router_gate_points = [
p.strip() for p in router_gate_points_raw.split(",") if p.strip()
]

if router_gate_strategy:
config.router_gate_strategy = router_gate_strategy
elif router_gate_points:
config.router_gate_strategy = "selective"
else:
config.router_gate_strategy = (
getattr(config, "router_gate_strategy", "off") or "off"
)

config.router_gate_points = router_gate_points


def _warn_if_schema_missing(config: AMReXAgentConfig, baseline_override: str | None) -> None:
"""Warn if schema is missing for the baseline override solver."""
if not baseline_override:
Expand Down Expand Up @@ -647,8 +715,7 @@ def main(args: list[str] | None = None) -> None:
if parsed_args.environment:
config.environment = parsed_args.environment

if parsed_args.preconfirm:
config.preconfirm_gate = True
apply_gate_cli_settings(config, parsed_args)

if parsed_args.llm_gate_strategy:
config.llm_gate_strategy = parsed_args.llm_gate_strategy
Expand Down Expand Up @@ -694,14 +761,23 @@ def main(args: list[str] | None = None) -> None:
if 'run_directory' in result:
run_dir = Path(result['run_directory'])
workflow_path = run_dir / "workflow_history.json"
gate_history_path = run_dir / "gate_history.json"
else:
base_dir = Path(parsed_args.output_dir) if parsed_args.output_dir else Path("output")
base_dir.mkdir(parents=True, exist_ok=True)
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
workflow_path = base_dir / f"workflow_history_{timestamp}.json"
gate_history_path = base_dir / f"gate_history_{timestamp}.json"
with open(workflow_path, 'w') as f:
json.dump(result.get('workflow_history', []), f, indent=2, default=str)
logger.info(f"Workflow history saved to {workflow_path}")
from src.utils.gate import build_gate_history_from_workflow_history
gate_history = build_gate_history_from_workflow_history(
result.get('workflow_history', [])
)
with open(gate_history_path, 'w') as f:
json.dump({"gates": gate_history}, f, indent=2, default=str)
logger.info(f"Gate history saved to {gate_history_path}")
except Exception as e:
logger.warning(f"Failed to save workflow history: {e}")

Expand Down Expand Up @@ -887,6 +963,7 @@ def create_amrex_agent_graph(checkpointer: Any = None) -> StateGraph:
workflow.add_node("runner", runner_node)
workflow.add_node("analysis", analysis_node)
workflow.add_node("visualization", visualization_node)
workflow.add_node("router_gate", router_gate_node)

# Add edges
workflow.add_edge(START, "architect")
Expand All @@ -898,8 +975,16 @@ def create_amrex_agent_graph(checkpointer: Any = None) -> StateGraph:
# 1. Entry point
workflow.add_edge(START, "architect")

# 2. Architect always sends plan to reviewer
workflow.add_edge("architect", "reviewer")
# 2. Architect routing (includes router-level gating)
workflow.add_conditional_edges(
"architect",
route_after_architect,
{
"reviewer": "reviewer",
"router_gate": "router_gate",
END: END,
}
)

# 3. Reviewer conditional routing (reflexion loop)
workflow.add_conditional_edges(
Expand All @@ -908,19 +993,29 @@ def create_amrex_agent_graph(checkpointer: Any = None) -> StateGraph:
{
"input_writer": "input_writer", # Proceed (validation passed)
"architect": "architect", # Retry (validation failed, attempts remain)
"router_gate": "router_gate",
END: END # Fail (max retries or critical error)
}
)

# 4. Linear execution path
workflow.add_edge("input_writer", "runner")
workflow.add_conditional_edges(
"input_writer",
route_after_input_writer,
{
"runner": "runner",
"router_gate": "router_gate",
END: END,
}
)

# 4b. Conditional routing from runner (check for failures)
workflow.add_conditional_edges(
"runner",
route_after_runner,
{
"analysis": "analysis", # Proceed to analysis on success
"router_gate": "router_gate",
END: END # Stop execution if runner fails
}
)
Expand All @@ -932,13 +1027,29 @@ def create_amrex_agent_graph(checkpointer: Any = None) -> StateGraph:
{
"visualization": "visualization", # Success (simulation completed)
"reviewer": "reviewer", # Failure (post-execution diagnosis)
"router_gate": "router_gate",
END: END # Terminal analysis failure
}
)

# 6. Terminal node
workflow.add_edge("visualization", END)

# 7. Router-level gate resume path
workflow.add_conditional_edges(
"router_gate",
route_after_router_gate,
{
"architect": "architect",
"reviewer": "reviewer",
"input_writer": "input_writer",
"runner": "runner",
"analysis": "analysis",
"visualization": "visualization",
END: END,
}
)

# ========================================
# COMPILATION
# ========================================
Expand Down
5 changes: 5 additions & 0 deletions src/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -94,6 +94,11 @@ class GraphState(TypedDict, total=False):
# Phase 2: Workflow History (rich event log for debugging/demo)
workflow_history: List[Dict[str, Any]] # [{iteration, node, action, timestamp, metadata}, ...]

# Router-level gate payloads (interactivity gates)
router_gate: Optional[Dict[str, Any]]
gate_proposal: Optional[Dict[str, Any]]
gate_resolution: Optional[Dict[str, Any]]

# ========================================
# History & Metadata
# ========================================
Expand Down
3 changes: 3 additions & 0 deletions src/models/graph_state_canonical.py
Original file line number Diff line number Diff line change
Expand Up @@ -161,6 +161,9 @@ class GraphState(TypedDict, total=False):
errors_fixed: List[str] # Errors successfully resolved
error_logs: List[str] # Runtime error messages from execution/analysis
preconfirm_action: Optional[str] # "proceed" | "cancel"
router_gate: Optional[Dict[str, Any]] # Router-level gate state
gate_proposal: Optional[Dict[str, Any]] # Router gate proposal payload
gate_resolution: Optional[Dict[str, Any]] # Router gate resolution payload

# ========================================
# OBSERVABILITY (Immutable Audit Log)
Expand Down
6 changes: 6 additions & 0 deletions src/nodes/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,11 +48,17 @@
except ImportError:
visualization_node = None

try:
from .router_gate_node import router_gate_node
except ImportError:
router_gate_node = None

__all__ = [
'architect_node',
'input_writer_node',
'runner_node',
'reviewer_node',
'analysis_node',
'visualization_node',
'router_gate_node',
]
37 changes: 32 additions & 5 deletions src/nodes/analysis_node.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@

from src.models import GraphState
from src.services.analysis import AnalysisService
from src.utils.gate import run_preconfirm_gate
from src.utils.gate import GateManager, run_preconfirm_gate

logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -91,6 +91,7 @@ def analysis_node(state: GraphState) -> dict[str, Any]:

config = state["config"]
iteration = state.get("iteration", 0)
workflow_history = state.get("workflow_history", [])

run_mode = getattr(config, "run_mode", None)
if run_mode is None or run_mode == "full":
Expand All @@ -101,7 +102,6 @@ def analysis_node(state: GraphState) -> dict[str, Any]:

if run_mode in {"dry", "stage", "submit"}:
logger.info("[INFO] Run mode %s - skipping analysis", run_mode)
workflow_history = state.get("workflow_history", [])
history_entry = {
"node": "analysis",
"timestamp": datetime.utcnow().isoformat() + "Z",
Expand All @@ -125,6 +125,35 @@ def analysis_node(state: GraphState) -> dict[str, Any]:
}
}

gate_manager = GateManager(
strategy=getattr(config, "gate_strategy", "auto") or "auto",
gate_points=getattr(config, "gate_points", []) or [],
)
allowed_gate_points = set(getattr(config, "gate_points", []) or [])
if (not allowed_gate_points or "analysis" in allowed_gate_points) and gate_manager.should_gate("analysis"):
decision = gate_manager.present_gate(
gate_point="analysis",
selected=state.get("run_directory") or "analysis",
confidence=1.0,
reasoning="Analyze simulation output for errors and metrics.",
evidence={"job_status": state.get("job_status", "unknown")},
alternatives=[],
)
decision_entry = {
"node": "preconfirm_gate",
"timestamp": datetime.utcnow().isoformat() + "Z",
"action": decision.user_action,
"iteration": iteration,
"details": {
"gate_node": "analysis",
"selection": {"value": decision.selected_option},
"reason": "decision_gate",
},
}
if decision.user_modification:
decision_entry["details"]["user_modification"] = decision.user_modification
workflow_history = workflow_history + [decision_entry]

gate_entry = None
run_dir = get_run_directory(state)
auto_approve = getattr(config, "preconfirm_gate_auto_approve", False) is True
Expand Down Expand Up @@ -152,7 +181,7 @@ def analysis_node(state: GraphState) -> dict[str, Any]:
"status": "skipped",
"message": "User canceled analysis at pre-confirm gate.",
},
"workflow_history": state.get("workflow_history", []) + ([gate_entry] if gate_entry else []),
"workflow_history": workflow_history + ([gate_entry] if gate_entry else []),
}

# Get run_directory from canonical path (workflow_history) with fallback to state
Expand All @@ -165,7 +194,6 @@ def analysis_node(state: GraphState) -> dict[str, Any]:
# WORKFLOW HISTORY ENTRY (SKIPPED)
# ========================================

workflow_history = state.get("workflow_history", [])
history_entry = {
"node": "analysis",
"timestamp": datetime.utcnow().isoformat() + "Z",
Expand Down Expand Up @@ -236,7 +264,6 @@ def analysis_node(state: GraphState) -> dict[str, Any]:
# ========================================
# Store complete analysis results in workflow_history.details

workflow_history = state.get("workflow_history", [])
if gate_entry:
workflow_history = workflow_history + [gate_entry]
status = report.get('status', 'unknown')
Expand Down
Loading
Loading