diff --git a/graph/langgraph_runtime.py b/graph/langgraph_runtime.py index e6b9cc0..eb5f33c 100644 --- a/graph/langgraph_runtime.py +++ b/graph/langgraph_runtime.py @@ -3,7 +3,8 @@ from __future__ import annotations from typing import Any, TypedDict from uuid import UUID -from graph.models import GraphRun +from control_plane.events.bus import EventBus +from graph.models import GraphEdgeTraversal, GraphRun, GraphRunStatus from graph.native_runtime import NativeGraphRuntime from graph.registry import NodeHandlerRegistry from graph.runtime import GraphRuntime @@ -17,11 +18,16 @@ class LangGraphState(TypedDict, total=False): edge_result: str +class GraphInterrupted(RuntimeError): + pass + + class LangGraphRuntime(GraphRuntime): """LangGraph adapter behind Artifex's runtime-neutral graph boundary.""" - def __init__(self, registry: NodeHandlerRegistry | None = None) -> None: + def __init__(self, registry: NodeHandlerRegistry | None = None, bus: EventBus | None = None) -> None: self.registry = registry + self.bus = bus or EventBus() async def start(self, project_id: UUID | None = None, **kwargs: Any) -> str: graph_run = kwargs.get("graph_run") @@ -45,7 +51,7 @@ class LangGraphRuntime(GraphRuntime): async def signal(self, run_id: str, event: dict[str, Any]) -> None: return None - def run_until_terminal_or_paused(self, graph_run: GraphRun) -> GraphRun: + def run_until_terminal_or_paused(self, graph_run: GraphRun, *, interrupt_after: str | None = None) -> GraphRun: if self.registry is None: raise RuntimeError("LangGraphRuntime requires a NodeHandlerRegistry") try: @@ -58,7 +64,7 @@ class LangGraphRuntime(GraphRuntime): workflow = StateGraph(LangGraphState) for node_id in spec.nodes: - workflow.add_node(node_id, self._node_runner(native, graph_run, node_id)) + workflow.add_node(node_id, self._node_runner(native, graph_run, node_id, interrupt_after)) workflow.set_entry_point(graph_run.current_node or spec.entry) for node_id in spec.nodes: if node_id in spec.terminal_nodes: @@ -68,28 +74,65 @@ class LangGraphRuntime(GraphRuntime): if not edges: workflow.add_edge(node_id, END) continue - workflow.add_conditional_edges( - node_id, - lambda state: str(state.get("edge_result", "success")), - {edge.condition or "success": edge.target for edge in edges}, - ) + workflow.add_conditional_edges(node_id, self._route(graph_run, spec, node_id), {edge.condition or "success": edge.target for edge in edges}) compiled = workflow.compile() - compiled.invoke({"graph_run_id": graph_run.id, "current_node": graph_run.current_node or spec.entry}) + try: + compiled.invoke({"graph_run_id": graph_run.id, "current_node": graph_run.current_node or spec.entry}) + except GraphInterrupted: + pass graph_run.refresh_from_db() return graph_run - def _node_runner(self, native: NativeGraphRuntime, graph_run: GraphRun, node_id: str): + def _node_runner(self, native: NativeGraphRuntime, graph_run: GraphRun, node_id: str, interrupt_after: str | None): def run_node(state: LangGraphState) -> LangGraphState: if graph_run.current_node != node_id: graph_run.current_node = node_id graph_run.save(update_fields=["current_node", "updated_at"]) native.run_until_terminal_or_paused(graph_run, interrupt_after=node_id) graph_run.refresh_from_db() + if interrupt_after == node_id: + raise GraphInterrupted(node_id) return { "graph_run_id": graph_run.id, "current_node": graph_run.current_node, "status": graph_run.status, - "edge_result": str(graph_run.metadata.get("last_edge_result", "success")), + "edge_result": str(graph_run.metadata.get("interrupted_edge_result", graph_run.metadata.get("last_edge_result", "success"))), } return run_node + + def _route(self, graph_run: GraphRun, spec: ExecutionGraphSpec, node_id: str): + def route(state: LangGraphState) -> str: + graph_run.refresh_from_db() + edge_result = str(state.get("edge_result") or graph_run.metadata.get("interrupted_edge_result") or "success") + target = None + for edge in spec.edges: + if edge.source == node_id and edge.condition == edge_result: + target = edge.target + break + if target is None: + for edge in spec.edges: + if edge.source == node_id and not edge.condition: + target = edge.target + edge_result = edge.condition or "success" + break + if target is None: + graph_run.status = GraphRunStatus.FAILED + graph_run.failure_reason = f"No edge from {node_id} for {edge_result}" + graph_run.save(update_fields=["status", "failure_reason", "updated_at"]) + return edge_result + GraphEdgeTraversal.objects.create(graph_run=graph_run, source_node=node_id, target_node=target, condition=edge_result, result=edge_result) + metadata = dict(graph_run.metadata) + metadata["last_node_id"] = node_id + metadata["last_edge_result"] = edge_result + metadata["last_output"] = metadata.get("interrupted_output", {}) + metadata.pop("interrupted_after_node", None) + metadata.pop("interrupted_edge_result", None) + metadata.pop("interrupted_output", None) + graph_run.current_node = target + graph_run.metadata = metadata + graph_run.save(update_fields=["current_node", "metadata", "updated_at"]) + self.bus.publish("GRAPH_EDGE_TRAVERSED", project=graph_run.project, task=graph_run.task, payload={"graph_run_id": graph_run.id, "source": node_id, "target": target, "condition": edge_result}) + return edge_result + + return route diff --git a/graph/native_runtime.py b/graph/native_runtime.py index ef7d79c..332d9dd 100644 --- a/graph/native_runtime.py +++ b/graph/native_runtime.py @@ -101,7 +101,11 @@ class NativeGraphRuntime(GraphRuntime): else: result = self._execute_node(context, node_run) if interrupt_after == node_id: - graph_run.metadata = {**dict(graph_run.metadata), "interrupted_after_node": node_id} + metadata = self._merge_summary(graph_run.metadata, result) + metadata["interrupted_after_node"] = node_id + metadata["interrupted_edge_result"] = result.edge_result + metadata["interrupted_output"] = self._bounded(result.output_metadata or {}) + graph_run.metadata = metadata graph_run.save(update_fields=["metadata", "updated_at"]) graph_run.refresh_from_db() return graph_run diff --git a/runtime_loop/autonomous_task_loop.py b/runtime_loop/autonomous_task_loop.py index c5f3fbc..2d4dfd3 100644 --- a/runtime_loop/autonomous_task_loop.py +++ b/runtime_loop/autonomous_task_loop.py @@ -11,8 +11,8 @@ from control_plane.events.models import EventType from control_plane.projects.models import CommitRecord, Task, TaskAttempt, TaskStatus, Worktree from control_plane.verification.models import VerificationResult from graph.bootstrap import champion_task_execution_graph_v1 +from graph.langgraph_runtime import LangGraphRuntime from graph.models import GraphRun -from graph.native_runtime import NativeGraphRuntime from graph.scheduler import TaskScheduler from graph.task_nodes import TaskExecutionServices, task_execution_registry from knowledge.context_builder import WorkerContextBuilder @@ -59,7 +59,7 @@ class AutonomousTaskLoop: current_node=graph_version.graph_spec["entry"], ) services = TaskExecutionServices(self.router, bus=self.bus, test_command=test_command) - NativeGraphRuntime(task_execution_registry(services), bus=self.bus).run_until_terminal_or_paused(graph_run) + LangGraphRuntime(task_execution_registry(services), bus=self.bus).run_until_terminal_or_paused(graph_run) return coder_version = self._champion(AgentRole.CODER) diff --git a/tests/test_task_execution_native_graph.py b/tests/test_task_execution_native_graph.py index 839c4b1..6e55e7c 100644 --- a/tests/test_task_execution_native_graph.py +++ b/tests/test_task_execution_native_graph.py @@ -9,6 +9,7 @@ from control_plane.projects.models import CommitRecord, Project, ProjectPlan, Mi from control_plane.verification.models import Review, TestRun, Verification, VerificationResult from graph.models import ExecutionGraphDefinition, ExecutionGraphVersion, ExecutionGraphVersionStatus, GraphEdgeTraversal, GraphRun, GraphRunStatus from graph.native_runtime import NativeGraphRuntime +from graph.langgraph_runtime import LangGraphRuntime from graph.task_execution import task_execution_graph_v1 from graph.task_nodes import TaskExecutionServices, task_execution_registry from model_router.router import ModelRouter @@ -58,6 +59,13 @@ def run_graph(task: Task, *, interrupt_after: str | None = None) -> GraphRun: return runtime.run_until_terminal_or_paused(graph_run_for_task(task), interrupt_after=interrupt_after) +def run_langgraph(task: Task, *, interrupt_after: str | None = None) -> GraphRun: + SeedAgentsCommand().handle() + services = TaskExecutionServices(ModelRouter({"qwen": DeterministicCodingProvider()}), test_command=["python", "manage.py", "test"]) + runtime = LangGraphRuntime(task_execution_registry(services)) + return runtime.run_until_terminal_or_paused(graph_run_for_task(task), interrupt_after=interrupt_after) + + def test_native_task_execution_graph_success_matches_loop_semantics(tmp_path: Path) -> None: repo = create_disposable_django_repo(tmp_path) task = create_task(repo, "Add a /health endpoint returning JSON {\"status\": \"ok\"} and add tests.", ["/health returns JSON ok", "tests pass"]) @@ -110,3 +118,58 @@ def test_native_task_execution_graph_resume_after_coder_prevents_duplicate_commi assert task.status == TaskStatus.COMPLETE assert CommitRecord.objects.filter(task=task).count() == 1 assert graph_run.node_runs.filter(node_id="coder").count() == 1 + + +def test_langgraph_task_execution_graph_success_matches_native_domain_outcome(tmp_path: Path) -> None: + repo = create_disposable_django_repo(tmp_path) + task = create_task(repo, "Add a /health endpoint returning JSON {\"status\": \"ok\"} and add tests.", ["/health returns JSON ok", "tests pass"]) + + graph_run = run_langgraph(task) + + task.refresh_from_db() + assert graph_run.status == GraphRunStatus.COMPLETE + assert graph_run.execution_graph_version.graph.name == "task_execution" + assert task.status == TaskStatus.COMPLETE + commit = CommitRecord.objects.get(task=task) + assert commit.graph_run == graph_run + assert TestRun.objects.get(task=task).status == "PASS" + assert Review.objects.get(task=task).status == "PASS" + assert Verification.objects.get(task=task).result == VerificationResult.PASS + assert GraphEdgeTraversal.objects.filter(graph_run=graph_run, source_node="review", target_node="judge", condition="PASS").exists() + + +def test_langgraph_task_execution_graph_retry_exhaustion_matches_native_domain_outcome(tmp_path: Path) -> None: + repo = create_disposable_django_repo(tmp_path) + task = create_task(repo, "FORCE_BAD_IMPLEMENTATION Add a /health endpoint returning JSON {\"status\": \"ok\"} and add tests.", ["/health returns JSON ok", "tests pass"], max_retries=1) + + graph_run = run_langgraph(task) + + task.refresh_from_db() + assert graph_run.status == GraphRunStatus.FAILED + assert task.status == TaskStatus.FAILED + assert task.retry_count == 2 + assert CommitRecord.objects.filter(task=task).count() == 0 + assert Event.objects.filter(task=task, event_type="TASK_RETRY_EXHAUSTED").exists() + + +def test_langgraph_resume_after_coder_tests_and_judge_does_not_duplicate_commit(tmp_path: Path) -> None: + for node_id in ["coder", "run_tests", "judge"]: + root = tmp_path / node_id + root.mkdir() + repo = create_disposable_django_repo(root) + task = create_task(repo, "Add a /health endpoint returning JSON {\"status\": \"ok\"} and add tests.", ["/health returns JSON ok", "tests pass"]) + SeedAgentsCommand().handle() + graph_run = graph_run_for_task(task) + services = TaskExecutionServices(ModelRouter({"qwen": DeterministicCodingProvider()}), test_command=["python", "manage.py", "test"]) + runtime = LangGraphRuntime(task_execution_registry(services)) + + runtime.run_until_terminal_or_paused(graph_run, interrupt_after=node_id) + graph_run.refresh_from_db() + assert graph_run.current_node == node_id + + runtime.run_until_terminal_or_paused(graph_run) + runtime.run_until_terminal_or_paused(graph_run) + + task.refresh_from_db() + assert task.status == TaskStatus.COMPLETE + assert CommitRecord.objects.filter(task=task).count() == 1