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