| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245 |
- """Pure Planner-decision rules for the explicit orchestration aggregate."""
- from __future__ import annotations
- from dataclasses import asdict
- from typing import Any, Callable, Dict, Optional
- from ._task_graph import ACTIVE_TASK_STATUSES, TERMINAL_TASK_STATUSES, TaskGraph
- from .errors import TaskConflict
- from .models import (
- DecisionAction,
- PlannerDecision,
- TaskLedger,
- TaskRecord,
- TaskStatus,
- ValidationReport,
- ValidationRunStatus,
- ValidationVerdict,
- json_values,
- new_id,
- utc_now,
- )
- from .state_machine import transition
- class DecisionEngine:
- """Apply a Planner decision to a caller-owned in-memory ledger."""
- def __init__(
- self,
- task_graph: TaskGraph,
- max_repair_continuations: Callable[[], int],
- ) -> None:
- self._task_graph = task_graph
- self._max_repair_continuations = max_repair_continuations
- def apply(
- self,
- ledger: TaskLedger,
- *,
- task_id: str,
- validation_id: Optional[str],
- action: DecisionAction,
- payload: Dict[str, Any],
- ) -> Dict[str, Any]:
- task = self._task(ledger, task_id)
- before = task.status
- is_root = task.task_id == ledger.root_task_id
- validation = ledger.validations.get(validation_id) if validation_id else None
- attempt_id = (
- validation.attempt_id
- if validation
- else task.attempt_ids[-1]
- if task.attempt_ids
- else None
- )
- reason = str(payload.get("reason", "")).strip()
- if not reason and action != DecisionAction.UNBLOCK:
- raise ValueError("Planner decision reason is required")
- if action == DecisionAction.ACCEPT:
- self.guard_accept(task, validation, ledger)
- target = TaskStatus.COMPLETED
- elif action == DecisionAction.REPAIR:
- self.guard_replan(task, validation, ledger)
- count_key = str(task.current_spec_version)
- used = task.repair_count_by_version.get(count_key, 0)
- if used >= self._max_repair_continuations():
- raise TaskConflict(
- "Repair continuation limit reached; retry, revise, split, or block"
- )
- task.repair_count_by_version[count_key] = used + 1
- target = TaskStatus.PENDING
- elif action == DecisionAction.RETRY:
- self._require_replan_state(task, "Retry")
- target = TaskStatus.PENDING
- elif action == DecisionAction.REVISE:
- self._require_replan_state(task, "Revise")
- objective = payload.get("objective", task.current_spec.objective)
- version = task.current_spec_version + 1
- draft = {
- "objective": objective,
- "acceptance_criteria": payload.get(
- "acceptance_criteria",
- [
- json_values(asdict(item))
- for item in task.current_spec.acceptance_criteria
- ],
- ),
- "context_refs": payload.get(
- "context_refs",
- task.current_spec.context_refs,
- ),
- }
- task.specs.append(TaskGraph.task_spec_from_draft(draft, version=version))
- task.current_spec_version = version
- target = TaskStatus.PENDING
- elif action == DecisionAction.SPLIT:
- self._require_replan_state(task, "Split")
- drafts = payload.get("tasks") or []
- if not drafts:
- raise ValueError("Split requires payload.tasks")
- child_ids = self._task_graph.create_child_records(ledger, task, drafts)
- payload["child_task_ids"] = child_ids
- target = TaskStatus.WAITING_CHILDREN
- elif action == DecisionAction.BLOCK:
- if task.status in TERMINAL_TASK_STATUSES:
- raise TaskConflict("Terminal task cannot be blocked")
- if any(
- child.status in ACTIVE_TASK_STATUSES
- for child in TaskGraph.nonterminal_descendants(ledger, task)
- ):
- raise TaskConflict("Task with active descendants cannot be blocked")
- task.blocked_reason = reason
- target = TaskStatus.BLOCKED
- elif action == DecisionAction.UNBLOCK:
- if task.status != TaskStatus.BLOCKED:
- raise TaskConflict("Only a blocked task can be unblocked")
- task.blocked_reason = None
- target = TaskStatus.NEEDS_REPLAN
- elif action == DecisionAction.CANCEL:
- self._guard_replace_or_cancel(ledger, task, is_root, "cancelled")
- target = TaskStatus.CANCELLED
- elif action == DecisionAction.SUPERSEDE:
- self._guard_replace_or_cancel(ledger, task, is_root, "superseded")
- replacement = payload.get("replacement") or {}
- replacement_ids = self._task_graph.create_sibling_records(
- ledger,
- task,
- [replacement],
- )
- task.superseded_by = replacement_ids[0]
- payload["replacement_task_id"] = replacement_ids[0]
- target = TaskStatus.SUPERSEDED
- elif action == DecisionAction.REVALIDATE:
- raise ValueError("Use revalidate_attempt() / validate_attempt tool")
- else:
- raise ValueError(f"Unsupported decision action: {action.value}")
- transition(before, target)
- task.status = target
- task.updated_at = utc_now()
- decision = PlannerDecision(
- decision_id=new_id(),
- task_id=task_id,
- action=action,
- reason=reason,
- from_status=before,
- to_status=target,
- attempt_id=attempt_id,
- validation_id=validation_id,
- payload=payload,
- )
- ledger.decisions[decision.decision_id] = decision
- task.decision_ids.append(decision.decision_id)
- TaskGraph.update_parent_after_child(ledger, task)
- return {
- "task_id": task_id,
- "decision_id": decision.decision_id,
- "action": action.value,
- "status": target.value,
- "payload": payload,
- }
- @staticmethod
- def guard_accept(
- task: TaskRecord,
- validation: Optional[ValidationReport],
- ledger: TaskLedger,
- ) -> None:
- if task.status != TaskStatus.AWAITING_DECISION:
- raise TaskConflict("Accept requires awaiting_decision")
- if not validation or validation.status != ValidationRunStatus.COMPLETED:
- raise TaskConflict("Accept requires a completed validation")
- if validation.spec_version != task.current_spec_version:
- raise TaskConflict("Validation belongs to an obsolete TaskSpec version")
- if validation.verdict != ValidationVerdict.PASSED:
- raise TaskConflict("Only a passed validation can be accepted")
- if not task.attempt_ids or validation.attempt_id != task.attempt_ids[-1]:
- raise TaskConflict("Validation does not belong to the current attempt")
- attempt = ledger.attempts[validation.attempt_id]
- if validation.snapshot_id != attempt.snapshot_id:
- raise TaskConflict("Validation does not belong to the current snapshot")
- if TaskGraph.nonterminal_descendants(ledger, task):
- raise TaskConflict(
- "Task cannot be accepted while descendants are non-terminal"
- )
- @staticmethod
- def guard_replan(
- task: TaskRecord,
- validation: Optional[ValidationReport],
- ledger: TaskLedger,
- ) -> None:
- if task.status != TaskStatus.AWAITING_DECISION:
- raise TaskConflict("Repair requires awaiting_decision")
- if not validation or validation.status != ValidationRunStatus.COMPLETED:
- raise TaskConflict("Repair requires a completed validation")
- if validation.spec_version != task.current_spec_version:
- raise TaskConflict(
- "Repair validation belongs to an obsolete TaskSpec version"
- )
- if not task.attempt_ids or validation.attempt_id != task.attempt_ids[-1]:
- raise TaskConflict(
- "Repair validation does not belong to the current attempt"
- )
- attempt = ledger.attempts[validation.attempt_id]
- if validation.snapshot_id != attempt.snapshot_id:
- raise TaskConflict(
- "Repair validation does not belong to the current snapshot"
- )
- if validation.verdict == ValidationVerdict.PASSED:
- raise TaskConflict("Passed work should be accepted, not repaired")
- @staticmethod
- def _task(ledger: TaskLedger, task_id: str) -> TaskRecord:
- try:
- return ledger.tasks[task_id]
- except KeyError as exc:
- raise ValueError(f"Task not found: {task_id}") from exc
- @staticmethod
- def _require_replan_state(task: TaskRecord, action: str) -> None:
- if task.status not in (
- TaskStatus.AWAITING_DECISION,
- TaskStatus.NEEDS_REPLAN,
- ):
- raise TaskConflict(f"{action} requires awaiting_decision or needs_replan")
- @staticmethod
- def _guard_replace_or_cancel(
- ledger: TaskLedger,
- task: TaskRecord,
- is_root: bool,
- action: str,
- ) -> None:
- if is_root:
- raise TaskConflict(f"The root task cannot be {action}")
- if task.status in TERMINAL_TASK_STATUSES:
- raise TaskConflict("Task is already terminal")
- if TaskGraph.nonterminal_descendants(ledger, task):
- raise TaskConflict(f"Task with non-terminal descendants cannot be {action}")
- __all__ = ["DecisionEngine"]
|