Add execution engine and task lifecycle services
This commit is contained in:
@@ -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.
|
||||||
|
|
||||||
@@ -18,13 +18,13 @@
|
|||||||
|
|
||||||
## Package B: Core Hardening
|
## Package B: Core Hardening
|
||||||
|
|
||||||
- [ ] graph scheduler service
|
- [x] graph scheduler service
|
||||||
- [ ] node runner abstraction
|
- [x] node runner abstraction
|
||||||
- [ ] reviewer contract
|
- [ ] reviewer contract
|
||||||
- [ ] finalizer contract
|
- [x] finalizer contract
|
||||||
- [ ] resume/cancel use cases
|
- [x] resume/cancel use cases
|
||||||
- [ ] idempotency policy
|
- [x] idempotency policy
|
||||||
- [ ] retry policy matrix
|
- [x] retry policy matrix
|
||||||
|
|
||||||
## Package C: Persistence
|
## Package C: Persistence
|
||||||
|
|
||||||
@@ -57,4 +57,3 @@
|
|||||||
- [ ] scenario tests
|
- [ ] scenario tests
|
||||||
- [ ] smoke suite
|
- [ ] smoke suite
|
||||||
- [ ] static checks in CI
|
- [ ] static checks in CI
|
||||||
|
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ from ai_orchestrator.domain.models import (
|
|||||||
ActionDescriptor,
|
ActionDescriptor,
|
||||||
ConfirmationRequest,
|
ConfirmationRequest,
|
||||||
ExecutionGraph,
|
ExecutionGraph,
|
||||||
|
ExecutionNode,
|
||||||
PolicyDecision,
|
PolicyDecision,
|
||||||
Task,
|
Task,
|
||||||
WorkerSession,
|
WorkerSession,
|
||||||
@@ -100,3 +101,33 @@ class WorkerGateway(Protocol):
|
|||||||
task_context: dict[str, Any],
|
task_context: dict[str, Any],
|
||||||
) -> ToolInvocationResult: ...
|
) -> 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]: ...
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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}")
|
||||||
@@ -86,6 +86,9 @@ class OrchestratorService:
|
|||||||
)
|
)
|
||||||
return graph
|
return graph
|
||||||
|
|
||||||
|
def get_task(self, task_id: str) -> Task:
|
||||||
|
return self._require_task(task_id)
|
||||||
|
|
||||||
def evaluate_action(
|
def evaluate_action(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -131,4 +134,3 @@ class OrchestratorService:
|
|||||||
if task is None:
|
if task is None:
|
||||||
raise LookupError(f"Task not found: {task_id}")
|
raise LookupError(f"Task not found: {task_id}")
|
||||||
return task
|
return task
|
||||||
|
|
||||||
|
|||||||
@@ -3,6 +3,15 @@ from __future__ import annotations
|
|||||||
from fastapi import FastAPI, HTTPException
|
from fastapi import FastAPI, HTTPException
|
||||||
from sse_starlette.sse import EventSourceResponse
|
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.orchestrator import CreateTaskRequest, OrchestratorService
|
||||||
from ai_orchestrator.application.services.workers import RegisterWorkerRequest, WorkerService
|
from ai_orchestrator.application.services.workers import RegisterWorkerRequest, WorkerService
|
||||||
from ai_orchestrator.config import AppSettings
|
from ai_orchestrator.config import AppSettings
|
||||||
@@ -12,6 +21,7 @@ from ai_orchestrator.delivery.http.schemas import (
|
|||||||
ConfirmationApproveRequest,
|
ConfirmationApproveRequest,
|
||||||
ConfirmationRejectRequest,
|
ConfirmationRejectRequest,
|
||||||
CreateTaskRequestSchema,
|
CreateTaskRequestSchema,
|
||||||
|
TaskActionRequest,
|
||||||
TaskStatusResponse,
|
TaskStatusResponse,
|
||||||
WorkerListItem,
|
WorkerListItem,
|
||||||
WorkerListResponse,
|
WorkerListResponse,
|
||||||
@@ -39,6 +49,14 @@ def create_app(settings: AppSettings | None = None) -> FastAPI:
|
|||||||
policy_evaluator = StaticProjectPolicyEvaluator(
|
policy_evaluator = StaticProjectPolicyEvaluator(
|
||||||
projects={project_config.project.id: project_config}
|
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(
|
orchestrator = OrchestratorService(
|
||||||
task_repository=task_repository,
|
task_repository=task_repository,
|
||||||
graph_repository=graph_repository,
|
graph_repository=graph_repository,
|
||||||
@@ -46,6 +64,12 @@ def create_app(settings: AppSettings | None = None) -> FastAPI:
|
|||||||
event_store=event_store,
|
event_store=event_store,
|
||||||
policy_evaluator=policy_evaluator,
|
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 = FastAPI(title="AI Orchestrator", version="0.1.0")
|
||||||
app.state.task_repository = task_repository
|
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.worker_repository = worker_repository
|
||||||
app.state.event_store = event_store
|
app.state.event_store = event_store
|
||||||
app.state.orchestrator = orchestrator
|
app.state.orchestrator = orchestrator
|
||||||
|
app.state.execution_engine = execution_engine
|
||||||
|
app.state.lifecycle_service = lifecycle_service
|
||||||
app.state.worker_service = WorkerService(
|
app.state.worker_service = WorkerService(
|
||||||
worker_repository=worker_repository,
|
worker_repository=worker_repository,
|
||||||
event_store=event_store,
|
event_store=event_store,
|
||||||
@@ -74,7 +100,8 @@ def create_app(settings: AppSettings | None = None) -> FastAPI:
|
|||||||
requested_mode=payload.mode,
|
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(
|
return ChatResponse(
|
||||||
conversation_id=task.conversation_id or "conv_default",
|
conversation_id=task.conversation_id or "conv_default",
|
||||||
task_id=task.task_id,
|
task_id=task.task_id,
|
||||||
@@ -94,6 +121,7 @@ def create_app(settings: AppSettings | None = None) -> FastAPI:
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
graph = orchestrator.plan_task(task.task_id)
|
graph = orchestrator.plan_task(task.task_id)
|
||||||
|
execution_engine.execute_ready_nodes(task=task, graph=graph)
|
||||||
return TaskStatusResponse(
|
return TaskStatusResponse(
|
||||||
task_id=task.task_id,
|
task_id=task.task_id,
|
||||||
status=task.status.value,
|
status=task.status.value,
|
||||||
@@ -142,8 +170,7 @@ def create_app(settings: AppSettings | None = None) -> FastAPI:
|
|||||||
confirmation = confirmation_repository.get(confirmation_id)
|
confirmation = confirmation_repository.get(confirmation_id)
|
||||||
if confirmation is None:
|
if confirmation is None:
|
||||||
raise HTTPException(status_code=404, detail="Confirmation not found")
|
raise HTTPException(status_code=404, detail="Confirmation not found")
|
||||||
confirmation.approve(payload.comment)
|
lifecycle_service.approve_confirmation(confirmation_id, payload.comment)
|
||||||
confirmation_repository.save(confirmation)
|
|
||||||
return {"status": confirmation.status.value}
|
return {"status": confirmation.status.value}
|
||||||
|
|
||||||
@app.post("/confirmations/{confirmation_id}/reject")
|
@app.post("/confirmations/{confirmation_id}/reject")
|
||||||
@@ -154,10 +181,39 @@ def create_app(settings: AppSettings | None = None) -> FastAPI:
|
|||||||
confirmation = confirmation_repository.get(confirmation_id)
|
confirmation = confirmation_repository.get(confirmation_id)
|
||||||
if confirmation is None:
|
if confirmation is None:
|
||||||
raise HTTPException(status_code=404, detail="Confirmation not found")
|
raise HTTPException(status_code=404, detail="Confirmation not found")
|
||||||
confirmation.reject(payload.reason)
|
lifecycle_service.reject_confirmation(confirmation_id, payload.reason)
|
||||||
confirmation_repository.save(confirmation)
|
|
||||||
return {"status": confirmation.status.value}
|
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)
|
@app.get("/workers", response_model=WorkerListResponse)
|
||||||
def list_workers() -> WorkerListResponse:
|
def list_workers() -> WorkerListResponse:
|
||||||
return WorkerListResponse(
|
return WorkerListResponse(
|
||||||
|
|||||||
@@ -64,3 +64,7 @@ class WorkerListItem(BaseModel):
|
|||||||
|
|
||||||
class WorkerListResponse(BaseModel):
|
class WorkerListResponse(BaseModel):
|
||||||
workers: list[WorkerListItem] = Field(default_factory=list)
|
workers: list[WorkerListItem] = Field(default_factory=list)
|
||||||
|
|
||||||
|
|
||||||
|
class TaskActionRequest(BaseModel):
|
||||||
|
reason: str | None = None
|
||||||
|
|||||||
@@ -28,10 +28,28 @@ def test_post_tasks_creates_planned_task() -> None:
|
|||||||
payload = response.json()
|
payload = response.json()
|
||||||
|
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
assert payload["status"] == "planned"
|
assert payload["status"] == "running"
|
||||||
assert payload["progress"]["total"] == 2
|
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:
|
def test_worker_registration_is_exposed_via_list_endpoint() -> None:
|
||||||
client = TestClient(app)
|
client = TestClient(app)
|
||||||
|
|
||||||
|
|||||||
@@ -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",
|
||||||
|
]
|
||||||
@@ -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
|
||||||
Reference in New Issue
Block a user