77 lines
2.9 KiB
Python
77 lines
2.9 KiB
Python
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
|