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
|
||||
|
||||
- [ ] 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
|
||||
|
||||
|
||||
@@ -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]: ...
|
||||
|
||||
@@ -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
|
||||
|
||||
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
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -64,3 +64,7 @@ class WorkerListItem(BaseModel):
|
||||
|
||||
class WorkerListResponse(BaseModel):
|
||||
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()
|
||||
|
||||
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)
|
||||
|
||||
|
||||
@@ -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