From 80715d81f6459be98a549d9b18853a5c7186d5ba Mon Sep 17 00:00:00 2001 From: Mikhail Date: Fri, 3 Jul 2026 21:12:58 +0300 Subject: [PATCH] Add execution engine and task lifecycle services --- docs/architecture/09_idempotency_and_retry.md | 47 ++++++ docs/backlog/01_execution_backlog.md | 13 +- src/ai_orchestrator/application/ports.py | 31 ++++ .../application/services/execution.py | 128 ++++++++++++++++ .../application/services/lifecycle.py | 137 ++++++++++++++++++ .../application/services/orchestrator.py | 4 +- src/ai_orchestrator/delivery/http/app.py | 66 ++++++++- src/ai_orchestrator/delivery/http/schemas.py | 4 + tests/integration/test_http_app.py | 20 ++- tests/unit/test_execution_engine.py | 47 ++++++ tests/unit/test_lifecycle_service.py | 76 ++++++++++ 11 files changed, 559 insertions(+), 14 deletions(-) create mode 100644 docs/architecture/09_idempotency_and_retry.md create mode 100644 src/ai_orchestrator/application/services/execution.py create mode 100644 src/ai_orchestrator/application/services/lifecycle.py create mode 100644 tests/unit/test_execution_engine.py create mode 100644 tests/unit/test_lifecycle_service.py diff --git a/docs/architecture/09_idempotency_and_retry.md b/docs/architecture/09_idempotency_and_retry.md new file mode 100644 index 0000000..bf3ce7f --- /dev/null +++ b/docs/architecture/09_idempotency_and_retry.md @@ -0,0 +1,47 @@ +# Idempotency And Retry + +## Idempotency Principles + +### Task Creation + +- внешний `create task` запрос должен поддерживать idempotency key; +- повтор одного и того же запроса не должен создавать дублирующиеся task records. + +### Confirmation Resolution + +- approve/reject одной confirmation должны быть идемпотентны; +- повтор того же решения должен возвращать текущее состояние без повторного side effect. + +### Worker Results + +- `command_id` должен быть уникальным ключом идемпотентности для результата worker command; +- повторно полученный result envelope не должен переисполнять финализацию шага. + +### MCP / Model Invocations + +- invocation records должны хранить correlation id; +- retry создает новую попытку, но с ссылкой на первичную cause chain. + +## Retry Matrix + +### Retryable + +- временные network errors; +- upstream timeout; +- parser failure weak model; +- transient MCP transport error; +- временный disconnect worker при разрешенном resume policy. + +### Non-Retryable + +- policy denial; +- unsupported capability; +- invalid request schema; +- explicit user rejection; +- deterministic contract mismatch in adapter configuration. + +## Execution Rule + +Retry инициируется только orchestrator. +Ни MCP adapter, ни worker, ни nested runner не решают самостоятельно “попробовать еще раз” вне зафиксированной policy/runtime semantics. + diff --git a/docs/backlog/01_execution_backlog.md b/docs/backlog/01_execution_backlog.md index 578bcab..b9c6d36 100644 --- a/docs/backlog/01_execution_backlog.md +++ b/docs/backlog/01_execution_backlog.md @@ -18,13 +18,13 @@ ## Package B: Core Hardening -- [ ] graph scheduler service -- [ ] node runner abstraction +- [x] graph scheduler service +- [x] node runner abstraction - [ ] reviewer contract -- [ ] finalizer contract -- [ ] resume/cancel use cases -- [ ] idempotency policy -- [ ] retry policy matrix +- [x] finalizer contract +- [x] resume/cancel use cases +- [x] idempotency policy +- [x] retry policy matrix ## Package C: Persistence @@ -57,4 +57,3 @@ - [ ] scenario tests - [ ] smoke suite - [ ] static checks in CI - diff --git a/src/ai_orchestrator/application/ports.py b/src/ai_orchestrator/application/ports.py index da944a1..b690dd2 100644 --- a/src/ai_orchestrator/application/ports.py +++ b/src/ai_orchestrator/application/ports.py @@ -8,6 +8,7 @@ from ai_orchestrator.domain.models import ( ActionDescriptor, ConfirmationRequest, ExecutionGraph, + ExecutionNode, PolicyDecision, Task, WorkerSession, @@ -100,3 +101,33 @@ class WorkerGateway(Protocol): task_context: dict[str, Any], ) -> ToolInvocationResult: ... + +class NodeRunner(Protocol): + supported_node_type: str + + def run( + self, + *, + task: Task, + node: ExecutionNode, + graph: ExecutionGraph, + ) -> dict[str, Any]: ... + + +class Reviewer(Protocol): + def review( + self, + *, + task: Task, + node: ExecutionNode, + output_data: dict[str, Any], + ) -> dict[str, Any]: ... + + +class Finalizer(Protocol): + def finalize( + self, + *, + task: Task, + graph: ExecutionGraph, + ) -> dict[str, Any]: ... diff --git a/src/ai_orchestrator/application/services/execution.py b/src/ai_orchestrator/application/services/execution.py new file mode 100644 index 0000000..b0ea417 --- /dev/null +++ b/src/ai_orchestrator/application/services/execution.py @@ -0,0 +1,128 @@ +from __future__ import annotations + +from dataclasses import dataclass, field + +from ai_orchestrator.application.ports import EventStore, Finalizer, NodeRunner, Reviewer +from ai_orchestrator.domain.enums import NodeStatus, NodeType +from ai_orchestrator.domain.events import DomainEvent +from ai_orchestrator.domain.models import ExecutionGraph, ExecutionNode, Task + + +@dataclass(slots=True) +class RunnerRegistry: + _runners: dict[str, NodeRunner] = field(default_factory=dict) + + def register(self, runner: NodeRunner) -> None: + self._runners[runner.supported_node_type] = runner + + def get(self, node_type: str) -> NodeRunner: + try: + return self._runners[node_type] + except KeyError as exc: + raise LookupError(f"No runner registered for node type: {node_type}") from exc + + +@dataclass(slots=True) +class NoOpReviewer(Reviewer): + def review( + self, + *, + task: Task, + node: ExecutionNode, + output_data: dict[str, object], + ) -> dict[str, object]: + return { + "status": "accepted", + "task_id": task.task_id, + "node_id": node.node_id, + **output_data, + } + + +@dataclass(slots=True) +class DefaultFinalizer(Finalizer): + def finalize(self, *, task: Task, graph: ExecutionGraph) -> dict[str, object]: + completed = sum(1 for node in graph.nodes if node.status == NodeStatus.COMPLETED) + failed = sum(1 for node in graph.nodes if node.status == NodeStatus.FAILED) + return { + "task_id": task.task_id, + "status": "completed", + "summary": ( + f"Task completed with {completed} completed nodes and " + f"{failed} failed nodes." + ), + "cards": [{"type": "result_card", "task_id": task.task_id}], + } + + +@dataclass(slots=True) +class PlannerNodeRunner(NodeRunner): + supported_node_type: str = NodeType.PLANNER.value + + def run( + self, + *, + task: Task, + node: ExecutionNode, + graph: ExecutionGraph, + ) -> dict[str, object]: + del graph + return { + "status": "planned", + "goal": task.goal, + "inputs": task.inputs, + } + + +@dataclass(slots=True) +class FinalizerNodeRunner(NodeRunner): + finalizer: Finalizer + supported_node_type: str = NodeType.FINALIZER.value + + def run( + self, + *, + task: Task, + node: ExecutionNode, + graph: ExecutionGraph, + ) -> dict[str, object]: + del node + return self.finalizer.finalize(task=task, graph=graph) + + +@dataclass(slots=True) +class GraphExecutionEngine: + event_store: EventStore + runner_registry: RunnerRegistry + reviewer: Reviewer + + def execute_ready_nodes(self, *, task: Task, graph: ExecutionGraph) -> list[ExecutionNode]: + executed: list[ExecutionNode] = [] + ready_nodes = graph.ready_nodes() + for node in ready_nodes: + node.mark_running() + task.mark_running(node_id=node.node_id) + self.event_store.append( + DomainEvent( + event_type="node_started", + task_id=task.task_id, + conversation_id=task.conversation_id, + node_id=node.node_id, + payload={"node_type": node.node_type.value, "attempt": node.attempts}, + ) + ) + runner = self.runner_registry.get(node.node_type.value) + output = runner.run(task=task, node=node, graph=graph) + reviewed_output = self.reviewer.review(task=task, node=node, output_data=output) + node.mark_completed(reviewed_output) + self.event_store.append( + DomainEvent( + event_type="node_completed", + task_id=task.task_id, + conversation_id=task.conversation_id, + node_id=node.node_id, + payload={"node_type": node.node_type.value, "output": reviewed_output}, + ) + ) + executed.append(node) + return executed diff --git a/src/ai_orchestrator/application/services/lifecycle.py b/src/ai_orchestrator/application/services/lifecycle.py new file mode 100644 index 0000000..05f90dc --- /dev/null +++ b/src/ai_orchestrator/application/services/lifecycle.py @@ -0,0 +1,137 @@ +from __future__ import annotations + +from dataclasses import dataclass + +from ai_orchestrator.application.ports import ( + ConfirmationRepository, + EventStore, + GraphRepository, + TaskRepository, +) +from ai_orchestrator.domain.enums import NodeStatus, TaskStatus +from ai_orchestrator.domain.events import DomainEvent +from ai_orchestrator.domain.models import ConfirmationRequest, ExecutionGraph, Task + + +@dataclass(slots=True) +class TaskLifecycleService: + task_repository: TaskRepository + graph_repository: GraphRepository + confirmation_repository: ConfirmationRepository + event_store: EventStore + + def approve_confirmation( + self, + confirmation_id: str, + comment: str | None = None, + ) -> ConfirmationRequest: + confirmation = self._require_confirmation(confirmation_id) + confirmation.approve(comment) + self.confirmation_repository.save(confirmation) + task = self._require_task(confirmation.task_id) + graph = self._require_graph(task.task_id) + node = self._require_node(graph, confirmation.node_id) + node.mark_ready() + task.mark_running(node_id=node.node_id) + self.task_repository.save(task) + self.graph_repository.save(graph) + self.event_store.append( + DomainEvent( + event_type="confirmation_approved", + task_id=task.task_id, + conversation_id=task.conversation_id, + node_id=node.node_id, + payload={"confirmation_id": confirmation.confirmation_id, "comment": comment}, + ) + ) + return confirmation + + def reject_confirmation(self, confirmation_id: str, reason: str) -> ConfirmationRequest: + confirmation = self._require_confirmation(confirmation_id) + confirmation.reject(reason) + self.confirmation_repository.save(confirmation) + task = self._require_task(confirmation.task_id) + graph = self._require_graph(task.task_id) + node = self._require_node(graph, confirmation.node_id) + node.mark_failed({"reason": reason, "status": "rejected"}) + task.fail({"reason": reason, "rejected_confirmation_id": confirmation.confirmation_id}) + self.task_repository.save(task) + self.graph_repository.save(graph) + self.event_store.append( + DomainEvent( + event_type="confirmation_rejected", + task_id=task.task_id, + conversation_id=task.conversation_id, + node_id=node.node_id, + payload={"confirmation_id": confirmation.confirmation_id, "reason": reason}, + ) + ) + return confirmation + + def resume_task(self, task_id: str) -> Task: + task = self._require_task(task_id) + if task.status not in {TaskStatus.WAITING_MANUAL, TaskStatus.WAITING_CONFIRMATION}: + return task + task.mark_running(node_id=task.current_node_id) + self.task_repository.save(task) + self.event_store.append( + DomainEvent( + event_type="task_resumed", + task_id=task.task_id, + conversation_id=task.conversation_id, + node_id=task.current_node_id, + payload={"previous_status": "waiting"}, + ) + ) + return task + + def cancel_task(self, task_id: str, reason: str) -> Task: + task = self._require_task(task_id) + graph = self._require_graph(task.task_id) + for node in graph.nodes: + if node.status in { + NodeStatus.PENDING, + NodeStatus.READY, + NodeStatus.RUNNING, + NodeStatus.WAITING_CONFIRMATION, + }: + node.status = NodeStatus.CANCELLED + task.status = TaskStatus.CANCELLED + task.result_summary = {"reason": reason, "status": "cancelled"} + self.task_repository.save(task) + self.graph_repository.save(graph) + self.event_store.append( + DomainEvent( + event_type="task_cancelled", + task_id=task.task_id, + conversation_id=task.conversation_id, + node_id=task.current_node_id, + payload={"reason": reason}, + ) + ) + return task + + def _require_task(self, task_id: str) -> Task: + task = self.task_repository.get(task_id) + if task is None: + raise LookupError(f"Task not found: {task_id}") + return task + + def _require_graph(self, task_id: str) -> ExecutionGraph: + graph = self.graph_repository.get(task_id) + if graph is None: + raise LookupError(f"Graph not found for task: {task_id}") + return graph + + def _require_confirmation(self, confirmation_id: str) -> ConfirmationRequest: + confirmation = self.confirmation_repository.get(confirmation_id) + if confirmation is None: + raise LookupError(f"Confirmation not found: {confirmation_id}") + return confirmation + + @staticmethod + def _require_node(graph: ExecutionGraph, node_id: str): + for node in graph.nodes: + if node.node_id == node_id: + return node + raise LookupError(f"Node not found in graph: {node_id}") diff --git a/src/ai_orchestrator/application/services/orchestrator.py b/src/ai_orchestrator/application/services/orchestrator.py index bb8e597..91e2420 100644 --- a/src/ai_orchestrator/application/services/orchestrator.py +++ b/src/ai_orchestrator/application/services/orchestrator.py @@ -86,6 +86,9 @@ class OrchestratorService: ) return graph + def get_task(self, task_id: str) -> Task: + return self._require_task(task_id) + def evaluate_action( self, *, @@ -131,4 +134,3 @@ class OrchestratorService: if task is None: raise LookupError(f"Task not found: {task_id}") return task - diff --git a/src/ai_orchestrator/delivery/http/app.py b/src/ai_orchestrator/delivery/http/app.py index 9981190..2ed6f94 100644 --- a/src/ai_orchestrator/delivery/http/app.py +++ b/src/ai_orchestrator/delivery/http/app.py @@ -3,6 +3,15 @@ from __future__ import annotations from fastapi import FastAPI, HTTPException from sse_starlette.sse import EventSourceResponse +from ai_orchestrator.application.services.execution import ( + DefaultFinalizer, + FinalizerNodeRunner, + GraphExecutionEngine, + NoOpReviewer, + PlannerNodeRunner, + RunnerRegistry, +) +from ai_orchestrator.application.services.lifecycle import TaskLifecycleService from ai_orchestrator.application.services.orchestrator import CreateTaskRequest, OrchestratorService from ai_orchestrator.application.services.workers import RegisterWorkerRequest, WorkerService from ai_orchestrator.config import AppSettings @@ -12,6 +21,7 @@ from ai_orchestrator.delivery.http.schemas import ( ConfirmationApproveRequest, ConfirmationRejectRequest, CreateTaskRequestSchema, + TaskActionRequest, TaskStatusResponse, WorkerListItem, WorkerListResponse, @@ -39,6 +49,14 @@ def create_app(settings: AppSettings | None = None) -> FastAPI: policy_evaluator = StaticProjectPolicyEvaluator( projects={project_config.project.id: project_config} ) + runner_registry = RunnerRegistry() + runner_registry.register(PlannerNodeRunner()) + runner_registry.register(FinalizerNodeRunner(finalizer=DefaultFinalizer())) + execution_engine = GraphExecutionEngine( + event_store=event_store, + runner_registry=runner_registry, + reviewer=NoOpReviewer(), + ) orchestrator = OrchestratorService( task_repository=task_repository, graph_repository=graph_repository, @@ -46,6 +64,12 @@ def create_app(settings: AppSettings | None = None) -> FastAPI: event_store=event_store, policy_evaluator=policy_evaluator, ) + lifecycle_service = TaskLifecycleService( + task_repository=task_repository, + graph_repository=graph_repository, + confirmation_repository=confirmation_repository, + event_store=event_store, + ) app = FastAPI(title="AI Orchestrator", version="0.1.0") app.state.task_repository = task_repository @@ -54,6 +78,8 @@ def create_app(settings: AppSettings | None = None) -> FastAPI: app.state.worker_repository = worker_repository app.state.event_store = event_store app.state.orchestrator = orchestrator + app.state.execution_engine = execution_engine + app.state.lifecycle_service = lifecycle_service app.state.worker_service = WorkerService( worker_repository=worker_repository, event_store=event_store, @@ -74,7 +100,8 @@ def create_app(settings: AppSettings | None = None) -> FastAPI: requested_mode=payload.mode, ) ) - orchestrator.plan_task(task.task_id) + graph = orchestrator.plan_task(task.task_id) + execution_engine.execute_ready_nodes(task=task, graph=graph) return ChatResponse( conversation_id=task.conversation_id or "conv_default", task_id=task.task_id, @@ -94,6 +121,7 @@ def create_app(settings: AppSettings | None = None) -> FastAPI: ) ) graph = orchestrator.plan_task(task.task_id) + execution_engine.execute_ready_nodes(task=task, graph=graph) return TaskStatusResponse( task_id=task.task_id, status=task.status.value, @@ -142,8 +170,7 @@ def create_app(settings: AppSettings | None = None) -> FastAPI: confirmation = confirmation_repository.get(confirmation_id) if confirmation is None: raise HTTPException(status_code=404, detail="Confirmation not found") - confirmation.approve(payload.comment) - confirmation_repository.save(confirmation) + lifecycle_service.approve_confirmation(confirmation_id, payload.comment) return {"status": confirmation.status.value} @app.post("/confirmations/{confirmation_id}/reject") @@ -154,10 +181,39 @@ def create_app(settings: AppSettings | None = None) -> FastAPI: confirmation = confirmation_repository.get(confirmation_id) if confirmation is None: raise HTTPException(status_code=404, detail="Confirmation not found") - confirmation.reject(payload.reason) - confirmation_repository.save(confirmation) + lifecycle_service.reject_confirmation(confirmation_id, payload.reason) return {"status": confirmation.status.value} + @app.post("/tasks/{task_id}/resume", response_model=TaskStatusResponse) + def resume_task(task_id: str, payload: TaskActionRequest) -> TaskStatusResponse: + del payload + task = lifecycle_service.resume_task(task_id) + graph = graph_repository.get(task_id) + if graph is None: + raise HTTPException(status_code=404, detail="Task graph not found") + execution_engine.execute_ready_nodes(task=task, graph=graph) + completed = sum(1 for node in graph.nodes if node.status.value == "completed") + return TaskStatusResponse( + task_id=task.task_id, + status=task.status.value, + current_node=task.current_node_id, + progress={"completed": completed, "total": len(graph.nodes)}, + ) + + @app.post("/tasks/{task_id}/cancel", response_model=TaskStatusResponse) + def cancel_task(task_id: str, payload: TaskActionRequest) -> TaskStatusResponse: + task = lifecycle_service.cancel_task(task_id, payload.reason or "cancelled") + graph = graph_repository.get(task_id) + if graph is None: + raise HTTPException(status_code=404, detail="Task graph not found") + completed = sum(1 for node in graph.nodes if node.status.value == "completed") + return TaskStatusResponse( + task_id=task.task_id, + status=task.status.value, + current_node=task.current_node_id, + progress={"completed": completed, "total": len(graph.nodes)}, + ) + @app.get("/workers", response_model=WorkerListResponse) def list_workers() -> WorkerListResponse: return WorkerListResponse( diff --git a/src/ai_orchestrator/delivery/http/schemas.py b/src/ai_orchestrator/delivery/http/schemas.py index 70af523..ff86712 100644 --- a/src/ai_orchestrator/delivery/http/schemas.py +++ b/src/ai_orchestrator/delivery/http/schemas.py @@ -64,3 +64,7 @@ class WorkerListItem(BaseModel): class WorkerListResponse(BaseModel): workers: list[WorkerListItem] = Field(default_factory=list) + + +class TaskActionRequest(BaseModel): + reason: str | None = None diff --git a/tests/integration/test_http_app.py b/tests/integration/test_http_app.py index 42486f2..ce83a65 100644 --- a/tests/integration/test_http_app.py +++ b/tests/integration/test_http_app.py @@ -28,10 +28,28 @@ def test_post_tasks_creates_planned_task() -> None: payload = response.json() assert response.status_code == 200 - assert payload["status"] == "planned" + assert payload["status"] == "running" assert payload["progress"]["total"] == 2 +def test_cancel_endpoint_marks_task_cancelled() -> None: + client = TestClient(app) + + created = client.post( + "/tasks", + json={ + "project_id": "default", + "goal": "Cancel me", + "inputs": {}, + "execution_mode": "agent_graph", + }, + ).json() + response = client.post(f"/tasks/{created['task_id']}/cancel", json={"reason": "stop"}) + + assert response.status_code == 200 + assert response.json()["status"] == "cancelled" + + def test_worker_registration_is_exposed_via_list_endpoint() -> None: client = TestClient(app) diff --git a/tests/unit/test_execution_engine.py b/tests/unit/test_execution_engine.py new file mode 100644 index 0000000..8d458b1 --- /dev/null +++ b/tests/unit/test_execution_engine.py @@ -0,0 +1,47 @@ +from ai_orchestrator.application.services.execution import ( + DefaultFinalizer, + FinalizerNodeRunner, + GraphExecutionEngine, + NoOpReviewer, + PlannerNodeRunner, + RunnerRegistry, +) +from ai_orchestrator.domain.enums import NodeStatus, NodeType, TaskStatus +from ai_orchestrator.domain.models import ExecutionGraph, ExecutionNode, Task +from ai_orchestrator.infrastructure.storage.memory import InMemoryEventStore + + +def test_execution_engine_runs_ready_nodes_and_records_events() -> None: + task = Task(project_id="default", goal="Run planner", inputs={}) + planner = ExecutionNode(task_id=task.task_id, node_type=NodeType.PLANNER, input_data={}) + finalizer = ExecutionNode( + task_id=task.task_id, + node_type=NodeType.FINALIZER, + input_data={}, + dependencies=[planner.node_id], + ) + graph = ExecutionGraph(task_id=task.task_id, nodes=[planner, finalizer]) + event_store = InMemoryEventStore() + registry = RunnerRegistry() + registry.register(PlannerNodeRunner()) + registry.register(FinalizerNodeRunner(finalizer=DefaultFinalizer())) + engine = GraphExecutionEngine( + event_store=event_store, + runner_registry=registry, + reviewer=NoOpReviewer(), + ) + + first_pass = engine.execute_ready_nodes(task=task, graph=graph) + second_pass = engine.execute_ready_nodes(task=task, graph=graph) + + assert len(first_pass) == 1 + assert len(second_pass) == 1 + assert planner.status == NodeStatus.COMPLETED + assert finalizer.status == NodeStatus.COMPLETED + assert task.status == TaskStatus.RUNNING + assert [event.event_type for event in event_store.items] == [ + "node_started", + "node_completed", + "node_started", + "node_completed", + ] diff --git a/tests/unit/test_lifecycle_service.py b/tests/unit/test_lifecycle_service.py new file mode 100644 index 0000000..68a91cb --- /dev/null +++ b/tests/unit/test_lifecycle_service.py @@ -0,0 +1,76 @@ +from ai_orchestrator.application.services.lifecycle import TaskLifecycleService +from ai_orchestrator.domain.enums import NodeStatus, NodeType, PolicyDecisionType, TaskStatus +from ai_orchestrator.domain.models import ( + ConfirmationRequest, + ExecutionGraph, + ExecutionNode, + PolicyDecision, + Task, +) +from ai_orchestrator.infrastructure.storage.memory import ( + InMemoryConfirmationRepository, + InMemoryEventStore, + InMemoryGraphRepository, + InMemoryTaskRepository, +) + + +def _build_service() -> tuple[ + TaskLifecycleService, + InMemoryTaskRepository, + InMemoryGraphRepository, + InMemoryConfirmationRepository, +]: + task_repository = InMemoryTaskRepository() + graph_repository = InMemoryGraphRepository() + confirmation_repository = InMemoryConfirmationRepository() + event_store = InMemoryEventStore() + service = TaskLifecycleService( + task_repository=task_repository, + graph_repository=graph_repository, + confirmation_repository=confirmation_repository, + event_store=event_store, + ) + return service, task_repository, graph_repository, confirmation_repository + + +def test_approve_confirmation_returns_task_to_running() -> None: + service, task_repository, graph_repository, confirmation_repository = _build_service() + task = Task(project_id="default", goal="Approve", inputs={}) + task.wait_for_confirmation("node_1") + task_repository.create(task) + node = ExecutionNode(task_id=task.task_id, node_type=NodeType.TOOL_CALL, input_data={}) + node.node_id = "node_1" + node.mark_waiting_confirmation() + graph_repository.save(ExecutionGraph(task_id=task.task_id, nodes=[node])) + confirmation = ConfirmationRequest( + task_id=task.task_id, + node_id=node.node_id, + decision=PolicyDecision( + decision=PolicyDecisionType.CONFIRM, + reason="Need approval", + requires_confirmation=True, + ), + ) + confirmation_repository.create(confirmation) + + updated = service.approve_confirmation(confirmation.confirmation_id, "ok") + + assert updated.status.value == "approved" + assert task_repository.get(task.task_id).status == TaskStatus.RUNNING + assert graph_repository.get(task.task_id).nodes[0].status == NodeStatus.READY + + +def test_cancel_task_marks_active_nodes_cancelled() -> None: + service, task_repository, graph_repository, _ = _build_service() + task = Task(project_id="default", goal="Cancel", inputs={}) + task.mark_running("node_1") + task_repository.create(task) + node = ExecutionNode(task_id=task.task_id, node_type=NodeType.PLANNER, input_data={}) + node.mark_running() + graph_repository.save(ExecutionGraph(task_id=task.task_id, nodes=[node])) + + cancelled = service.cancel_task(task.task_id, "stop") + + assert cancelled.status == TaskStatus.CANCELLED + assert graph_repository.get(task.task_id).nodes[0].status == NodeStatus.CANCELLED