484 lines
23 KiB
Python
484 lines
23 KiB
Python
from __future__ import annotations
|
|
|
|
from pathlib import Path
|
|
|
|
from agents.coder import Coder
|
|
from agents.judge import Judge
|
|
from agents.progeny import ProgenyService
|
|
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,
|
|
agent_overrides: dict[str, AgentVersion] | None = None,
|
|
) -> 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()
|
|
self.progeny = ProgenyService(self.bus)
|
|
self.tests = DeterministicTestRunner()
|
|
self.test_command = test_command or ["python", "-m", "pytest"]
|
|
self.agent_overrides = agent_overrides or {}
|
|
|
|
def champion(self, role: AgentRole) -> AgentVersion:
|
|
if role in self.agent_overrides:
|
|
return self.agent_overrides[role]
|
|
if role.value in self.agent_overrides:
|
|
return self.agent_overrides[role.value]
|
|
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"])
|
|
|
|
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)
|
|
|
|
|
|
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"])
|
|
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
|
|
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
|
|
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"
|
|
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)
|
|
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 {})
|
|
|
|
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"
|
|
|
|
|
|
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"
|
|
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)
|
|
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"
|
|
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)
|
|
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)
|
|
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)
|
|
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"])
|
|
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},
|
|
)
|
|
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:]
|
|
|
|
|
|
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,
|
|
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})
|
|
|
|
|
|
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})
|
|
|
|
|
|
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),
|
|
StaticAnalysisNode(services),
|
|
RetryOrFailNode(services),
|
|
CleanupNode(services),
|
|
]:
|
|
registry.register(handler)
|
|
return registry
|