_decision_engine.py 9.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245
  1. """Pure Planner-decision rules for the explicit orchestration aggregate."""
  2. from __future__ import annotations
  3. from dataclasses import asdict
  4. from typing import Any, Callable, Dict, Optional
  5. from ._task_graph import ACTIVE_TASK_STATUSES, TERMINAL_TASK_STATUSES, TaskGraph
  6. from .errors import TaskConflict
  7. from .models import (
  8. DecisionAction,
  9. PlannerDecision,
  10. TaskLedger,
  11. TaskRecord,
  12. TaskStatus,
  13. ValidationReport,
  14. ValidationRunStatus,
  15. ValidationVerdict,
  16. json_values,
  17. new_id,
  18. utc_now,
  19. )
  20. from .state_machine import transition
  21. class DecisionEngine:
  22. """Apply a Planner decision to a caller-owned in-memory ledger."""
  23. def __init__(
  24. self,
  25. task_graph: TaskGraph,
  26. max_repair_continuations: Callable[[], int],
  27. ) -> None:
  28. self._task_graph = task_graph
  29. self._max_repair_continuations = max_repair_continuations
  30. def apply(
  31. self,
  32. ledger: TaskLedger,
  33. *,
  34. task_id: str,
  35. validation_id: Optional[str],
  36. action: DecisionAction,
  37. payload: Dict[str, Any],
  38. ) -> Dict[str, Any]:
  39. task = self._task(ledger, task_id)
  40. before = task.status
  41. is_root = task.task_id == ledger.root_task_id
  42. validation = ledger.validations.get(validation_id) if validation_id else None
  43. attempt_id = (
  44. validation.attempt_id
  45. if validation
  46. else task.attempt_ids[-1]
  47. if task.attempt_ids
  48. else None
  49. )
  50. reason = str(payload.get("reason", "")).strip()
  51. if not reason and action != DecisionAction.UNBLOCK:
  52. raise ValueError("Planner decision reason is required")
  53. if action == DecisionAction.ACCEPT:
  54. self.guard_accept(task, validation, ledger)
  55. target = TaskStatus.COMPLETED
  56. elif action == DecisionAction.REPAIR:
  57. self.guard_replan(task, validation, ledger)
  58. count_key = str(task.current_spec_version)
  59. used = task.repair_count_by_version.get(count_key, 0)
  60. if used >= self._max_repair_continuations():
  61. raise TaskConflict(
  62. "Repair continuation limit reached; retry, revise, split, or block"
  63. )
  64. task.repair_count_by_version[count_key] = used + 1
  65. target = TaskStatus.PENDING
  66. elif action == DecisionAction.RETRY:
  67. self._require_replan_state(task, "Retry")
  68. target = TaskStatus.PENDING
  69. elif action == DecisionAction.REVISE:
  70. self._require_replan_state(task, "Revise")
  71. objective = payload.get("objective", task.current_spec.objective)
  72. version = task.current_spec_version + 1
  73. draft = {
  74. "objective": objective,
  75. "acceptance_criteria": payload.get(
  76. "acceptance_criteria",
  77. [
  78. json_values(asdict(item))
  79. for item in task.current_spec.acceptance_criteria
  80. ],
  81. ),
  82. "context_refs": payload.get(
  83. "context_refs",
  84. task.current_spec.context_refs,
  85. ),
  86. }
  87. task.specs.append(TaskGraph.task_spec_from_draft(draft, version=version))
  88. task.current_spec_version = version
  89. target = TaskStatus.PENDING
  90. elif action == DecisionAction.SPLIT:
  91. self._require_replan_state(task, "Split")
  92. drafts = payload.get("tasks") or []
  93. if not drafts:
  94. raise ValueError("Split requires payload.tasks")
  95. child_ids = self._task_graph.create_child_records(ledger, task, drafts)
  96. payload["child_task_ids"] = child_ids
  97. target = TaskStatus.WAITING_CHILDREN
  98. elif action == DecisionAction.BLOCK:
  99. if task.status in TERMINAL_TASK_STATUSES:
  100. raise TaskConflict("Terminal task cannot be blocked")
  101. if any(
  102. child.status in ACTIVE_TASK_STATUSES
  103. for child in TaskGraph.nonterminal_descendants(ledger, task)
  104. ):
  105. raise TaskConflict("Task with active descendants cannot be blocked")
  106. task.blocked_reason = reason
  107. target = TaskStatus.BLOCKED
  108. elif action == DecisionAction.UNBLOCK:
  109. if task.status != TaskStatus.BLOCKED:
  110. raise TaskConflict("Only a blocked task can be unblocked")
  111. task.blocked_reason = None
  112. target = TaskStatus.NEEDS_REPLAN
  113. elif action == DecisionAction.CANCEL:
  114. self._guard_replace_or_cancel(ledger, task, is_root, "cancelled")
  115. target = TaskStatus.CANCELLED
  116. elif action == DecisionAction.SUPERSEDE:
  117. self._guard_replace_or_cancel(ledger, task, is_root, "superseded")
  118. replacement = payload.get("replacement") or {}
  119. replacement_ids = self._task_graph.create_sibling_records(
  120. ledger,
  121. task,
  122. [replacement],
  123. )
  124. task.superseded_by = replacement_ids[0]
  125. payload["replacement_task_id"] = replacement_ids[0]
  126. target = TaskStatus.SUPERSEDED
  127. elif action == DecisionAction.REVALIDATE:
  128. raise ValueError("Use revalidate_attempt() / validate_attempt tool")
  129. else:
  130. raise ValueError(f"Unsupported decision action: {action.value}")
  131. transition(before, target)
  132. task.status = target
  133. task.updated_at = utc_now()
  134. decision = PlannerDecision(
  135. decision_id=new_id(),
  136. task_id=task_id,
  137. action=action,
  138. reason=reason,
  139. from_status=before,
  140. to_status=target,
  141. attempt_id=attempt_id,
  142. validation_id=validation_id,
  143. payload=payload,
  144. )
  145. ledger.decisions[decision.decision_id] = decision
  146. task.decision_ids.append(decision.decision_id)
  147. TaskGraph.update_parent_after_child(ledger, task)
  148. return {
  149. "task_id": task_id,
  150. "decision_id": decision.decision_id,
  151. "action": action.value,
  152. "status": target.value,
  153. "payload": payload,
  154. }
  155. @staticmethod
  156. def guard_accept(
  157. task: TaskRecord,
  158. validation: Optional[ValidationReport],
  159. ledger: TaskLedger,
  160. ) -> None:
  161. if task.status != TaskStatus.AWAITING_DECISION:
  162. raise TaskConflict("Accept requires awaiting_decision")
  163. if not validation or validation.status != ValidationRunStatus.COMPLETED:
  164. raise TaskConflict("Accept requires a completed validation")
  165. if validation.spec_version != task.current_spec_version:
  166. raise TaskConflict("Validation belongs to an obsolete TaskSpec version")
  167. if validation.verdict != ValidationVerdict.PASSED:
  168. raise TaskConflict("Only a passed validation can be accepted")
  169. if not task.attempt_ids or validation.attempt_id != task.attempt_ids[-1]:
  170. raise TaskConflict("Validation does not belong to the current attempt")
  171. attempt = ledger.attempts[validation.attempt_id]
  172. if validation.snapshot_id != attempt.snapshot_id:
  173. raise TaskConflict("Validation does not belong to the current snapshot")
  174. if TaskGraph.nonterminal_descendants(ledger, task):
  175. raise TaskConflict(
  176. "Task cannot be accepted while descendants are non-terminal"
  177. )
  178. @staticmethod
  179. def guard_replan(
  180. task: TaskRecord,
  181. validation: Optional[ValidationReport],
  182. ledger: TaskLedger,
  183. ) -> None:
  184. if task.status != TaskStatus.AWAITING_DECISION:
  185. raise TaskConflict("Repair requires awaiting_decision")
  186. if not validation or validation.status != ValidationRunStatus.COMPLETED:
  187. raise TaskConflict("Repair requires a completed validation")
  188. if validation.spec_version != task.current_spec_version:
  189. raise TaskConflict(
  190. "Repair validation belongs to an obsolete TaskSpec version"
  191. )
  192. if not task.attempt_ids or validation.attempt_id != task.attempt_ids[-1]:
  193. raise TaskConflict(
  194. "Repair validation does not belong to the current attempt"
  195. )
  196. attempt = ledger.attempts[validation.attempt_id]
  197. if validation.snapshot_id != attempt.snapshot_id:
  198. raise TaskConflict(
  199. "Repair validation does not belong to the current snapshot"
  200. )
  201. if validation.verdict == ValidationVerdict.PASSED:
  202. raise TaskConflict("Passed work should be accepted, not repaired")
  203. @staticmethod
  204. def _task(ledger: TaskLedger, task_id: str) -> TaskRecord:
  205. try:
  206. return ledger.tasks[task_id]
  207. except KeyError as exc:
  208. raise ValueError(f"Task not found: {task_id}") from exc
  209. @staticmethod
  210. def _require_replan_state(task: TaskRecord, action: str) -> None:
  211. if task.status not in (
  212. TaskStatus.AWAITING_DECISION,
  213. TaskStatus.NEEDS_REPLAN,
  214. ):
  215. raise TaskConflict(f"{action} requires awaiting_decision or needs_replan")
  216. @staticmethod
  217. def _guard_replace_or_cancel(
  218. ledger: TaskLedger,
  219. task: TaskRecord,
  220. is_root: bool,
  221. action: str,
  222. ) -> None:
  223. if is_root:
  224. raise TaskConflict(f"The root task cannot be {action}")
  225. if task.status in TERMINAL_TASK_STATUSES:
  226. raise TaskConflict("Task is already terminal")
  227. if TaskGraph.nonterminal_descendants(ledger, task):
  228. raise TaskConflict(f"Task with non-terminal descendants cannot be {action}")
  229. __all__ = ["DecisionEngine"]