Artifex/graph/task_nodes.py

485 lines
23 KiB
Python
Raw Normal View History

2026-08-15 17:05:01 +07:00
from __future__ import annotations
from pathlib import Path
from agents.coder import Coder
from agents.judge import Judge
2026-08-15 17:53:52 +07:00
from agents.progeny import ProgenyService
2026-08-15 17:05:01 +07:00
from agents.reviewer import Reviewer
from control_plane.agents.models import AgentRole, AgentVersion
from control_plane.events.bus import EventBus
from control_plane.events.models import EventType
from control_plane.projects.models import CommitRecord, Task, TaskAttempt, TaskStatus, Worktree
from control_plane.verification.models import TestRun, VerificationResult
from graph.native_runtime import GraphExecutionContext
from graph.registry import NodeHandlerRegistry, NodeResult
from knowledge.context_builder import WorkerContextBuilder
from model_router.router import ModelRouter
from tools.capabilities import Capability
from tools.runtime import WorktreeTools
from tools.test_runner import DeterministicTestRunner
from workspace.worktrees import WorktreeManager
class TaskExecutionServices:
def __init__(
self,
router: ModelRouter,
*,
bus: EventBus | None = None,
test_command: list[str] | None = None,
2026-08-15 18:14:21 +07:00
agent_overrides: dict[str, AgentVersion] | None = None,
2026-08-15 17:05:01 +07:00
) -> None:
self.router = router
self.bus = bus or EventBus()
self.context_builder = WorkerContextBuilder()
self.worktrees = WorktreeManager()
self.coder = Coder(router)
self.reviewer = Reviewer()
self.judge = Judge()
2026-08-15 17:53:52 +07:00
self.progeny = ProgenyService(self.bus)
2026-08-15 17:05:01 +07:00
self.tests = DeterministicTestRunner()
self.test_command = test_command or ["python", "-m", "pytest"]
2026-08-15 18:14:21 +07:00
self.agent_overrides = agent_overrides or {}
2026-08-15 17:05:01 +07:00
def champion(self, role: AgentRole) -> AgentVersion:
2026-08-15 18:14:21 +07:00
if role in self.agent_overrides:
return self.agent_overrides[role]
if role.value in self.agent_overrides:
return self.agent_overrides[role.value]
2026-08-15 17:05:01 +07:00
return AgentVersion.objects.select_related("agent").get(agent__role=role, promotion_status="CHAMPION")
def tools(self, worktree: Worktree) -> WorktreeTools:
return WorktreeTools(
Path(worktree.worktree_path),
{
Capability.READ_REPOSITORY,
Capability.INVESTIGATE_WORKTREE,
Capability.WRITE_WORKTREE,
Capability.RUN_TESTS,
Capability.COMMIT_CHANGES,
},
)
def worktree(self, task: Task) -> Worktree:
try:
return task.worktree
except Worktree.DoesNotExist:
if not task.project.repository_path:
raise RuntimeError("Task project has no repository_path")
return self.worktrees.create_for_task(task, Path(task.project.repository_path))
class TaskNode:
idempotent = True
replay_safe = "checkpointed"
destructive = False
def __init__(self, services: TaskExecutionServices, node_type: str) -> None:
self.services = services
self.node_type = node_type
def task(self, context: GraphExecutionContext) -> Task:
if context.graph_run.task_id is None:
raise RuntimeError("Task execution graph run requires task")
return Task.objects.select_related("project", "milestone", "feature").get(id=context.graph_run.task_id)
def metadata(self, context: GraphExecutionContext) -> dict[str, object]:
return dict(context.graph_run.metadata)
def save_metadata(self, context: GraphExecutionContext, metadata: dict[str, object]) -> None:
context.graph_run.metadata = metadata
context.graph_run.save(update_fields=["metadata", "updated_at"])
2026-08-15 17:53:52 +07:00
def current_node_run(self, context: GraphExecutionContext):
return context.graph_run.node_runs.filter(node_id=context.graph_run.current_node).order_by("-visit_index").first()
def _clear_current_failure(self, metadata: dict[str, object]) -> None:
metadata.pop("current_failure_reason", None)
metadata.pop("current_failure_findings", None)
metadata.pop("current_failure_node_id", None)
metadata.pop("last_failure_reason", None)
metadata.pop("last_failure_findings", None)
2026-08-15 17:05:01 +07:00
class ClaimTaskNode(TaskNode):
def __init__(self, services: TaskExecutionServices) -> None:
super().__init__(services, "claim_task")
self.idempotent = True
def run(self, context: GraphExecutionContext) -> NodeResult:
task = self.task(context)
return NodeResult("COMPLETE", "success", {"task_id": str(task.id), "task_status": task.status})
class PrepareWorktreeNode(TaskNode):
def __init__(self, services: TaskExecutionServices) -> None:
super().__init__(services, "prepare_worktree")
def run(self, context: GraphExecutionContext) -> NodeResult:
task = self.task(context)
worktree = self.services.worktree(task)
metadata = self.metadata(context)
metadata["worktree_id"] = str(worktree.id)
metadata["worktree_path"] = worktree.worktree_path
self.save_metadata(context, metadata)
return NodeResult("COMPLETE", "success", {"worktree_id": str(worktree.id), "worktree_path": worktree.worktree_path})
class BuildContextNode(TaskNode):
def __init__(self, services: TaskExecutionServices) -> None:
super().__init__(services, "build_context")
def run(self, context: GraphExecutionContext) -> NodeResult:
task = self.task(context)
worktree = self.services.worktree(task)
metadata = self.metadata(context)
attempt_id = metadata.get("current_attempt_id")
if attempt_id:
attempt = TaskAttempt.objects.get(id=attempt_id)
return NodeResult("COMPLETE", "success", {"attempt_id": str(attempt.id), "attempt_number": attempt.attempt_number})
coder_version = self.services.champion(AgentRole.CODER)
attempt = TaskAttempt.objects.create(task=task, attempt_number=task.retry_count + 1, coder=coder_version, status="RUNNING")
task_context = self.services.context_builder.build_for_task(task, Path(worktree.worktree_path))
task_context["previous_attempts"] = list(
task.attempts.exclude(id=attempt.id).order_by("attempt_number").values("attempt_number", "status", "coder_result", "review_findings", "judge_findings")
)
attempt.context_snapshot = task_context
attempt.save(update_fields=["context_snapshot", "updated_at"])
metadata["current_attempt_id"] = str(attempt.id)
self.save_metadata(context, metadata)
context.graph_run.task_attempt = attempt
context.graph_run.save(update_fields=["task_attempt", "updated_at"])
2026-08-15 17:05:01 +07:00
return NodeResult("COMPLETE", "success", {"attempt_id": str(attempt.id), "attempt_number": attempt.attempt_number})
class CoderNode(TaskNode):
def __init__(self, services: TaskExecutionServices) -> None:
super().__init__(services, "coder")
self.idempotent = False
def run(self, context: GraphExecutionContext) -> NodeResult:
task = self.task(context)
worktree = self.services.worktree(task)
metadata = self.metadata(context)
attempt = TaskAttempt.objects.get(id=metadata["current_attempt_id"])
coder_version = attempt.coder
2026-08-15 17:53:52 +07:00
try:
result = self.services.coder.execute(attempt.context_snapshot, self.services.tools(worktree), project=task.project, agent_version=coder_version)
except Exception as exc:
self.services.progeny.create_provider_signal(
task.project,
task,
task.milestone,
coder_version,
self._provider_failure_category(str(exc)),
f"Coder provider/runtime failure: {exc}",
{"error": str(exc)},
graph_run=context.graph_run,
graph_node_run=self.current_node_run(context),
)
raise
2026-08-15 17:05:01 +07:00
attempt.coder_result = {"status": result.status, "summary": result.summary, "changed_files": result.changed_files, "metadata": result.metadata}
attempt.save(update_fields=["coder_result", "updated_at"])
if result.status != "COMPLETE":
metadata["current_failure_reason"] = "coder_failed"
metadata["current_failure_findings"] = [result.summary]
metadata["current_failure_node_id"] = "coder"
2026-08-15 17:53:52 +07:00
if "malformed json" in result.summary.lower() or "json" in result.summary.lower():
self.services.progeny.create_model_output_signal(
task.project,
task,
task.milestone,
coder_version,
result.summary,
str(result.metadata),
graph_run=context.graph_run,
graph_node_run=self.current_node_run(context),
)
else:
self._clear_current_failure(metadata)
2026-08-15 17:05:01 +07:00
self.save_metadata(context, metadata)
return NodeResult("COMPLETE", "success" if result.status == "COMPLETE" else "failure", {"coder_status": result.status, "summary": result.summary}, result.metadata.get("telemetry", {}) if isinstance(result.metadata, dict) else {})
2026-08-15 17:53:52 +07:00
def _provider_failure_category(self, error: str) -> str:
lowered = error.lower()
if "timeout" in lowered or "timed out" in lowered:
return "PROVIDER_TIMEOUT"
if "unavailable" in lowered or "connection" in lowered or "configured" in lowered:
return "PROVIDER_UNAVAILABLE"
return "PROVIDER_RUNTIME_ERROR"
2026-08-15 17:05:01 +07:00
class RunTestsNode(TaskNode):
def __init__(self, services: TaskExecutionServices) -> None:
super().__init__(services, "run_tests")
def run(self, context: GraphExecutionContext) -> NodeResult:
task = self.task(context)
worktree = self.services.worktree(task)
metadata = self.metadata(context)
attempt = TaskAttempt.objects.get(id=metadata["current_attempt_id"])
test_run = self.services.tests.run(task.project, task, Path(worktree.worktree_path), self.services.test_command)
metadata["test_run_id"] = str(test_run.id)
self.save_metadata(context, metadata)
if test_run.status != "PASS":
self._attach_test_failure_evidence(attempt, test_run)
self.services.bus.publish(EventType.TEST_FAILED, project=task.project, task=task, payload={"test_run_id": str(test_run.id)})
return NodeResult("COMPLETE", "complete", {"test_run_id": str(test_run.id), "test_status": test_run.status}, {"test_status": test_run.status})
def _attach_test_failure_evidence(self, attempt: TaskAttempt, test_run: TestRun) -> None:
content = test_run.output_artifact.content if test_run.output_artifact else {}
stdout = str(content.get("stdout", "")) if isinstance(content, dict) else ""
stderr = str(content.get("stderr", "")) if isinstance(content, dict) else ""
coder_result = dict(attempt.coder_result or {})
attempt_metadata = dict(coder_result.get("metadata", {})) if isinstance(coder_result.get("metadata", {}), dict) else {}
attempt_metadata["test_failure_evidence"] = {"test_run_id": str(test_run.id), "status": test_run.status, "stdout_excerpt": stdout[-12000:], "stderr_excerpt": stderr[-4000:]}
coder_result["metadata"] = attempt_metadata
attempt.coder_result = coder_result
attempt.save(update_fields=["coder_result", "updated_at"])
class ReviewNode(TaskNode):
def __init__(self, services: TaskExecutionServices) -> None:
super().__init__(services, "review")
def run(self, context: GraphExecutionContext) -> NodeResult:
task = self.task(context)
metadata = self.metadata(context)
attempt = TaskAttempt.objects.get(id=metadata["current_attempt_id"])
reviewer_version = self.services.champion(AgentRole.REVIEWER)
worktree = self.services.worktree(task)
tools = self.services.tools(worktree)
tools.git(["add", "-N", "."])
diff = tools.diff()
test_run = TestRun.objects.get(id=metadata["test_run_id"])
review = self.services.reviewer.review(task, reviewer_version, diff, test_run.status)
attempt.review_findings = review.findings
attempt.save(update_fields=["review_findings", "updated_at"])
metadata["review_id"] = str(review.id)
if review.status != "PASS":
metadata["current_failure_reason"] = "review_failed"
metadata["current_failure_findings"] = review.findings
metadata["current_failure_node_id"] = "review"
2026-08-15 17:53:52 +07:00
self.services.progeny.create_reviewer_signal(
task.project,
task,
task.milestone,
reviewer_version,
review.status,
review.findings,
review.summary,
graph_run=context.graph_run,
graph_node_run=self.current_node_run(context),
metadata={"affected_agent_version_id": str(attempt.coder_id), "reviewer_version_id": str(reviewer_version.id), "test_status": test_run.status},
)
else:
self._clear_current_failure(metadata)
2026-08-15 17:05:01 +07:00
self.save_metadata(context, metadata)
if review.status != "PASS":
self.services.bus.publish(EventType.REVIEW_FAILED, project=task.project, task=task, payload={"review_id": str(review.id), "findings": review.findings})
return NodeResult("COMPLETE", review.status, {"review_id": str(review.id), "review_status": review.status, "findings": review.findings}, {"review_status": review.status})
class JudgeNode(TaskNode):
def __init__(self, services: TaskExecutionServices) -> None:
super().__init__(services, "judge")
def run(self, context: GraphExecutionContext) -> NodeResult:
task = self.task(context)
metadata = self.metadata(context)
attempt = TaskAttempt.objects.get(id=metadata["current_attempt_id"])
judge_version = self.services.champion(AgentRole.PROJECT_JUDGE)
test_run = TestRun.objects.get(id=metadata["test_run_id"])
worktree = self.services.worktree(task)
tools = self.services.tools(worktree)
tools.git(["add", "-N", "."])
diff = tools.diff()
verification = self.services.judge.judge(task.project, task, judge_version, diff, test_run.status)
attempt.judge_findings = verification.evidence
attempt.save(update_fields=["judge_findings", "updated_at"])
metadata["verification_id"] = str(verification.id)
if verification.result != VerificationResult.PASS:
metadata["current_failure_reason"] = "judge_failed"
metadata["current_failure_findings"] = verification.evidence
metadata["current_failure_node_id"] = "judge"
2026-08-15 17:53:52 +07:00
self.services.progeny.create_judge_signal(
task.project,
task,
task.milestone,
judge_version,
verification.result,
verification.evidence,
verification.summary,
graph_run=context.graph_run,
graph_node_run=self.current_node_run(context),
metadata={"affected_agent_version_id": str(attempt.coder_id), "judge_version_id": str(judge_version.id), "test_status": test_run.status},
)
else:
self._clear_current_failure(metadata)
2026-08-15 17:05:01 +07:00
self.save_metadata(context, metadata)
edge = "PASS" if verification.result == VerificationResult.PASS else "FAIL"
return NodeResult("COMPLETE", edge, {"verification_id": str(verification.id), "result": verification.result, "evidence": verification.evidence}, {"judge_result": verification.result})
class RetryOrFailNode(TaskNode):
def __init__(self, services: TaskExecutionServices) -> None:
super().__init__(services, "retry_or_fail")
self.idempotent = False
def run(self, context: GraphExecutionContext) -> NodeResult:
task = self.task(context)
metadata = self.metadata(context)
attempt = TaskAttempt.objects.get(id=metadata["current_attempt_id"])
effective_max_retries = min(task.max_retries, 2)
attempt.status = "REWORK_REQUIRED" if task.retry_count < effective_max_retries else "FAILED"
attempt.save(update_fields=["status", "updated_at"])
task.retry_count += 1
reason = str(metadata.get("current_failure_reason", "task_failed"))
findings = metadata.get("current_failure_findings", [])
self._append_historical_failure(metadata, attempt, reason, findings)
2026-08-15 17:05:01 +07:00
classification = self._classify_failure(reason, findings)
metadata.pop("current_attempt_id", None)
if task.retry_count > effective_max_retries:
metadata["final_failure_reason"] = reason
self._clear_current_failure(metadata)
2026-08-15 17:05:01 +07:00
self.save_metadata(context, metadata)
if task.retry_count <= effective_max_retries:
task.status = TaskStatus.RUNNING
task.save(update_fields=["retry_count", "status", "updated_at"])
self.services.bus.publish(EventType.TASK_FAILED, project=task.project, task=task, payload={"reason": reason, "will_retry": True, "classification": classification, "findings": findings})
return NodeResult("COMPLETE", "retry_available", {"retry_count": task.retry_count, "classification": classification})
task.status = TaskStatus.FAILED
task.save(update_fields=["retry_count", "status", "updated_at"])
2026-08-15 17:53:52 +07:00
self.services.progeny.create_retry_exhausted_signal(
task.project,
task,
task.milestone,
attempt.coder,
task.retry_count,
graph_run=context.graph_run,
graph_node_run=self.current_node_run(context),
evidence={"reason": reason, "classification": classification, "findings": findings},
)
2026-08-15 17:05:01 +07:00
self.services.bus.publish("TASK_RETRY_EXHAUSTED", project=task.project, task=task, payload={"reason": reason, "classification": classification, "findings": findings})
self.services.bus.publish(EventType.TASK_FAILED, project=task.project, task=task, payload={"reason": reason, "will_retry": False, "classification": classification, "findings": findings})
return NodeResult("COMPLETE", "retry_exhausted", {"retry_count": task.retry_count, "classification": classification})
def _classify_failure(self, reason: str, findings: object) -> str:
text = f"{reason} {findings}".lower()
if "unsupported operation" in text or "missing capability" in text:
return "missing_capability"
if "context" in text or "migration" in text:
return "context_problem"
if "timeout" in text or "provider" in text:
return "environment_problem"
if "malformed json" in text or "model" in text:
return "model_problem"
if "ambiguous" in text:
return "intent_ambiguity"
if reason == "review_failed" or reason == "judge_failed":
return "replan"
return "split_task"
def _append_historical_failure(self, metadata: dict[str, object], attempt: TaskAttempt, reason: str, findings: object) -> None:
failures = list(metadata.get("historical_failures", []))
failures.append(
{
"node_id": metadata.get("current_failure_node_id", "unknown"),
"visit_index": len(failures) + 1,
"attempt_id": str(attempt.id),
"attempt_number": attempt.attempt_number,
"reason": reason,
"evidence": findings,
}
)
metadata["historical_failures"] = failures[-20:]
2026-08-15 17:05:01 +07:00
class CommitNode(TaskNode):
destructive = True
def __init__(self, services: TaskExecutionServices) -> None:
super().__init__(services, "commit")
self.idempotent = False
def run(self, context: GraphExecutionContext) -> NodeResult:
task = self.task(context)
existing = CommitRecord.objects.filter(task=task).first()
if existing is not None:
return NodeResult("COMPLETE", "success", {"commit_id": str(existing.id), "sha": existing.sha, "deduplicated": True})
metadata = self.metadata(context)
attempt = TaskAttempt.objects.get(id=metadata["current_attempt_id"])
worktree = self.services.worktree(task)
test_run = TestRun.objects.get(id=metadata["test_run_id"])
sha = self.services.tools(worktree).commit_all(f"Artifex task: {task.goal[:80]}")
commit = CommitRecord.objects.create(
project=task.project,
task=task,
worktree=worktree,
coder=attempt.coder,
reviewer=self.services.champion(AgentRole.REVIEWER),
judge=self.services.champion(AgentRole.PROJECT_JUDGE),
test_run=test_run,
review_id=metadata.get("review_id"),
verification_id=metadata.get("verification_id"),
graph_run=context.graph_run,
2026-08-15 17:05:01 +07:00
sha=sha,
branch_name=worktree.branch_name,
message=f"Artifex task: {task.goal[:80]}",
)
attempt.status = "COMPLETE"
attempt.save(update_fields=["status", "updated_at"])
task.status = TaskStatus.COMPLETE
task.save(update_fields=["status", "updated_at"])
self.services.bus.publish(EventType.COMMIT_CREATED, project=task.project, task=task, payload={"commit_id": str(commit.id), "sha": sha})
self.services.bus.publish(EventType.TASK_COMPLETED, project=task.project, task=task, payload={"task_id": str(task.id)})
metadata["commit_id"] = str(commit.id)
metadata["commit_sha"] = sha
self.save_metadata(context, metadata)
return NodeResult("COMPLETE", "success", {"commit_id": str(commit.id), "sha": sha})
2026-08-15 18:14:21 +07:00
class StaticAnalysisNode(TaskNode):
def __init__(self, services: TaskExecutionServices) -> None:
super().__init__(services, "static_analysis")
def run(self, context: GraphExecutionContext) -> NodeResult:
task = self.task(context)
metadata = self.metadata(context)
metadata["static_analysis"] = {"status": "PASS", "checks": ["deterministic_fixture"]}
self.save_metadata(context, metadata)
return NodeResult("COMPLETE", "PASS", {"static_analysis_status": "PASS"}, {"static_analysis_checks": 1})
2026-08-15 17:05:01 +07:00
class CleanupNode(TaskNode):
def __init__(self, services: TaskExecutionServices) -> None:
super().__init__(services, "cleanup")
def run(self, context: GraphExecutionContext) -> NodeResult:
task = self.task(context)
if task.status == TaskStatus.COMPLETE:
self.services.worktrees.validate_clean_worktree(task.worktree)
self.services.worktrees.cleanup(task.worktree)
return NodeResult("COMPLETE", "success", {"cleaned": True})
return NodeResult("COMPLETE", "failed", {"cleaned": False})
def task_execution_registry(services: TaskExecutionServices) -> NodeHandlerRegistry:
registry = NodeHandlerRegistry()
for handler in [
ClaimTaskNode(services),
PrepareWorktreeNode(services),
BuildContextNode(services),
CoderNode(services),
RunTestsNode(services),
ReviewNode(services),
JudgeNode(services),
CommitNode(services),
2026-08-15 18:14:21 +07:00
StaticAnalysisNode(services),
2026-08-15 17:05:01 +07:00
RetryOrFailNode(services),
CleanupNode(services),
]:
registry.register(handler)
return registry