Add worker gateway and websocket transport flow
This commit is contained in:
@@ -38,13 +38,13 @@
|
|||||||
|
|
||||||
- [x] model router providers
|
- [x] model router providers
|
||||||
- [x] MCP transport adapter
|
- [x] MCP transport adapter
|
||||||
- [ ] worker WebSocket gateway
|
- [x] worker WebSocket gateway
|
||||||
- [ ] heartbeat monitor
|
- [x] heartbeat monitor
|
||||||
- [ ] capability-aware dispatch
|
- [x] capability-aware dispatch
|
||||||
|
|
||||||
## Package E: Delivery
|
## Package E: Delivery
|
||||||
|
|
||||||
- [ ] typed worker registration schema
|
- [x] typed worker registration schema
|
||||||
- [ ] task event DTO normalization
|
- [ ] task event DTO normalization
|
||||||
- [ ] confirmation cards
|
- [ ] confirmation cards
|
||||||
- [ ] progress cards
|
- [ ] progress cards
|
||||||
|
|||||||
@@ -1,8 +1,17 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
|
from datetime import datetime
|
||||||
|
from uuid import uuid4
|
||||||
|
|
||||||
from ai_orchestrator.application.ports import EventStore, WorkerRepository
|
from ai_orchestrator.application.ports import (
|
||||||
|
EventStore,
|
||||||
|
InvocationStore,
|
||||||
|
ToolInvocationRecord,
|
||||||
|
ToolInvocationResult,
|
||||||
|
WorkerGateway,
|
||||||
|
WorkerRepository,
|
||||||
|
)
|
||||||
from ai_orchestrator.domain.events import DomainEvent
|
from ai_orchestrator.domain.events import DomainEvent
|
||||||
from ai_orchestrator.domain.models import WorkerSession
|
from ai_orchestrator.domain.models import WorkerSession
|
||||||
|
|
||||||
@@ -17,13 +26,259 @@ class RegisterWorkerRequest:
|
|||||||
capabilities: list[str]
|
capabilities: list[str]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class WorkerCommandEnvelope:
|
||||||
|
command_id: str
|
||||||
|
task_id: str
|
||||||
|
tool: str
|
||||||
|
args: dict[str, object]
|
||||||
|
timeout_ms: int
|
||||||
|
policy_context: dict[str, object]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class WorkerProgressUpdate:
|
||||||
|
command_id: str
|
||||||
|
task_id: str
|
||||||
|
tool: str
|
||||||
|
percent: int
|
||||||
|
message: str
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class WorkerCommandResult:
|
||||||
|
command_id: str
|
||||||
|
task_id: str
|
||||||
|
tool: str
|
||||||
|
status: str
|
||||||
|
started_at: datetime
|
||||||
|
finished_at: datetime
|
||||||
|
duration_ms: int
|
||||||
|
stdout: str = ""
|
||||||
|
stderr: str = ""
|
||||||
|
result: dict[str, object] | None = None
|
||||||
|
artifacts: list[dict[str, object]] | None = None
|
||||||
|
error: dict[str, object] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class RegisterWorkerResponse:
|
||||||
|
worker: WorkerSession
|
||||||
|
reconnect: bool = False
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class DispatchWorkerCommandRequest:
|
||||||
|
project_id: str
|
||||||
|
worker_session_id: str
|
||||||
|
command_name: str
|
||||||
|
args: dict[str, object]
|
||||||
|
task_context: dict[str, object]
|
||||||
|
timeout_ms: int = 30000
|
||||||
|
policy_context: dict[str, object] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class InMemoryWorkerConnectionManager:
|
||||||
|
queues: dict[str, list[WorkerCommandEnvelope]]
|
||||||
|
active_commands: dict[str, tuple[str, WorkerCommandEnvelope]]
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.queues = {}
|
||||||
|
self.active_commands = {}
|
||||||
|
|
||||||
|
def attach(self, session_id: str) -> None:
|
||||||
|
self.queues.setdefault(session_id, [])
|
||||||
|
|
||||||
|
def detach(self, session_id: str) -> None:
|
||||||
|
self.queues.setdefault(session_id, [])
|
||||||
|
|
||||||
|
def enqueue(self, session_id: str, command: WorkerCommandEnvelope) -> None:
|
||||||
|
self.queues.setdefault(session_id, []).append(command)
|
||||||
|
self.active_commands[command.command_id] = (session_id, command)
|
||||||
|
|
||||||
|
def poll(self, session_id: str) -> list[WorkerCommandEnvelope]:
|
||||||
|
queue = self.queues.setdefault(session_id, [])
|
||||||
|
commands = list(queue)
|
||||||
|
queue.clear()
|
||||||
|
return commands
|
||||||
|
|
||||||
|
def get(self, command_id: str) -> tuple[str, WorkerCommandEnvelope] | None:
|
||||||
|
return self.active_commands.get(command_id)
|
||||||
|
|
||||||
|
def complete(self, command_id: str) -> tuple[str, WorkerCommandEnvelope] | None:
|
||||||
|
return self.active_commands.pop(command_id, None)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class CapabilityAwareWorkerGateway(WorkerGateway):
|
||||||
|
worker_repository: WorkerRepository
|
||||||
|
event_store: EventStore
|
||||||
|
invocation_store: InvocationStore
|
||||||
|
connection_manager: InMemoryWorkerConnectionManager
|
||||||
|
|
||||||
|
def dispatch(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
project_id: str,
|
||||||
|
worker_session_id: str,
|
||||||
|
command_name: str,
|
||||||
|
args: dict[str, object],
|
||||||
|
task_context: dict[str, object],
|
||||||
|
) -> ToolInvocationResult:
|
||||||
|
del project_id
|
||||||
|
worker = self._require_worker(worker_session_id)
|
||||||
|
if command_name not in worker.capabilities:
|
||||||
|
raise ValueError(f"Worker does not support capability: {command_name}")
|
||||||
|
command = WorkerCommandEnvelope(
|
||||||
|
command_id=f"cmd_{uuid4().hex}",
|
||||||
|
task_id=str(task_context["task_id"]),
|
||||||
|
tool=command_name,
|
||||||
|
args=args,
|
||||||
|
timeout_ms=int(task_context.get("timeout_ms", 30000)),
|
||||||
|
policy_context=dict(task_context.get("policy_context", {})),
|
||||||
|
)
|
||||||
|
worker.mark_busy(task_id=command.task_id)
|
||||||
|
self.worker_repository.save(worker)
|
||||||
|
self.connection_manager.enqueue(worker_session_id, command)
|
||||||
|
self.invocation_store.save_tool_invocation(
|
||||||
|
ToolInvocationRecord(
|
||||||
|
invocation_id=f"tinv_{uuid4().hex}",
|
||||||
|
task_id=command.task_id,
|
||||||
|
node_id=str(task_context["node_id"]) if task_context.get("node_id") else None,
|
||||||
|
source_type="worker",
|
||||||
|
source_id=worker_session_id,
|
||||||
|
tool_name=command.tool,
|
||||||
|
status="dispatched",
|
||||||
|
request={
|
||||||
|
"command_id": command.command_id,
|
||||||
|
"args": command.args,
|
||||||
|
"timeout_ms": command.timeout_ms,
|
||||||
|
"policy_context": command.policy_context,
|
||||||
|
},
|
||||||
|
response={"state": "queued"},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
self.event_store.append(
|
||||||
|
DomainEvent(
|
||||||
|
event_type="worker_command_dispatched",
|
||||||
|
task_id=command.task_id,
|
||||||
|
node_id=str(task_context["node_id"]) if task_context.get("node_id") else None,
|
||||||
|
payload={
|
||||||
|
"worker_session_id": worker_session_id,
|
||||||
|
"command_id": command.command_id,
|
||||||
|
"tool": command.tool,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return ToolInvocationResult(
|
||||||
|
status="queued",
|
||||||
|
content={
|
||||||
|
"command_id": command.command_id,
|
||||||
|
"worker_session_id": worker_session_id,
|
||||||
|
"tool": command.tool,
|
||||||
|
},
|
||||||
|
logs=[],
|
||||||
|
)
|
||||||
|
|
||||||
|
def record_progress(self, worker_session_id: str, update: WorkerProgressUpdate) -> None:
|
||||||
|
worker = self._require_worker(worker_session_id)
|
||||||
|
worker.heartbeat()
|
||||||
|
self.worker_repository.save(worker)
|
||||||
|
self.event_store.append(
|
||||||
|
DomainEvent(
|
||||||
|
event_type="worker_command_progress",
|
||||||
|
task_id=update.task_id,
|
||||||
|
payload={
|
||||||
|
"worker_session_id": worker_session_id,
|
||||||
|
"command_id": update.command_id,
|
||||||
|
"tool": update.tool,
|
||||||
|
"percent": update.percent,
|
||||||
|
"message": update.message,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
def complete_command(
|
||||||
|
self,
|
||||||
|
worker_session_id: str,
|
||||||
|
result: WorkerCommandResult,
|
||||||
|
) -> ToolInvocationResult:
|
||||||
|
worker = self._require_worker(worker_session_id)
|
||||||
|
self.connection_manager.complete(result.command_id)
|
||||||
|
worker.mark_online()
|
||||||
|
self.worker_repository.save(worker)
|
||||||
|
self.invocation_store.save_tool_invocation(
|
||||||
|
ToolInvocationRecord(
|
||||||
|
invocation_id=f"tinv_{uuid4().hex}",
|
||||||
|
task_id=result.task_id,
|
||||||
|
node_id=None,
|
||||||
|
source_type="worker_result",
|
||||||
|
source_id=worker_session_id,
|
||||||
|
tool_name=result.tool,
|
||||||
|
status=result.status,
|
||||||
|
request={
|
||||||
|
"command_id": result.command_id,
|
||||||
|
"started_at": result.started_at.isoformat(),
|
||||||
|
},
|
||||||
|
response={
|
||||||
|
"finished_at": result.finished_at.isoformat(),
|
||||||
|
"duration_ms": result.duration_ms,
|
||||||
|
"stdout": result.stdout,
|
||||||
|
"stderr": result.stderr,
|
||||||
|
"result": result.result or {},
|
||||||
|
"artifacts": result.artifacts or [],
|
||||||
|
"error": result.error,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
self.event_store.append(
|
||||||
|
DomainEvent(
|
||||||
|
event_type="worker_command_completed",
|
||||||
|
task_id=result.task_id,
|
||||||
|
payload={
|
||||||
|
"worker_session_id": worker_session_id,
|
||||||
|
"command_id": result.command_id,
|
||||||
|
"tool": result.tool,
|
||||||
|
"status": result.status,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return ToolInvocationResult(
|
||||||
|
status=result.status,
|
||||||
|
content=result.result or {},
|
||||||
|
artifacts=result.artifacts or [],
|
||||||
|
logs=[item for item in [result.stdout, result.stderr] if item],
|
||||||
|
error=result.error,
|
||||||
|
)
|
||||||
|
|
||||||
|
def poll_commands(self, worker_session_id: str) -> list[WorkerCommandEnvelope]:
|
||||||
|
self._require_worker(worker_session_id)
|
||||||
|
return self.connection_manager.poll(worker_session_id)
|
||||||
|
|
||||||
|
def _require_worker(self, worker_session_id: str) -> WorkerSession:
|
||||||
|
worker = self.worker_repository.get(worker_session_id)
|
||||||
|
if worker is None:
|
||||||
|
raise LookupError(f"Worker session not found: {worker_session_id}")
|
||||||
|
return worker
|
||||||
|
|
||||||
|
|
||||||
@dataclass(slots=True)
|
@dataclass(slots=True)
|
||||||
class WorkerService:
|
class WorkerService:
|
||||||
worker_repository: WorkerRepository
|
worker_repository: WorkerRepository
|
||||||
event_store: EventStore
|
event_store: EventStore
|
||||||
|
|
||||||
def register(self, request: RegisterWorkerRequest) -> WorkerSession:
|
def register(self, request: RegisterWorkerRequest) -> RegisterWorkerResponse:
|
||||||
worker = WorkerSession(
|
reconnect = False
|
||||||
|
existing = next(
|
||||||
|
(
|
||||||
|
item
|
||||||
|
for item in self.worker_repository.list_active()
|
||||||
|
if item.worker_id == request.worker_id and item.machine == request.machine
|
||||||
|
),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
worker = existing or WorkerSession(
|
||||||
worker_id=request.worker_id,
|
worker_id=request.worker_id,
|
||||||
name=request.name,
|
name=request.name,
|
||||||
machine=request.machine,
|
machine=request.machine,
|
||||||
@@ -31,6 +286,12 @@ class WorkerService:
|
|||||||
version=request.version,
|
version=request.version,
|
||||||
capabilities=request.capabilities,
|
capabilities=request.capabilities,
|
||||||
)
|
)
|
||||||
|
if existing is not None:
|
||||||
|
reconnect = True
|
||||||
|
worker.name = request.name
|
||||||
|
worker.os = request.os
|
||||||
|
worker.version = request.version
|
||||||
|
worker.capabilities = request.capabilities
|
||||||
worker.mark_online()
|
worker.mark_online()
|
||||||
self.worker_repository.save(worker)
|
self.worker_repository.save(worker)
|
||||||
self.event_store.append(
|
self.event_store.append(
|
||||||
@@ -41,7 +302,60 @@ class WorkerService:
|
|||||||
"session_id": worker.session_id,
|
"session_id": worker.session_id,
|
||||||
"worker_id": worker.worker_id,
|
"worker_id": worker.worker_id,
|
||||||
"capabilities": worker.capabilities,
|
"capabilities": worker.capabilities,
|
||||||
|
"reconnect": reconnect,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
return RegisterWorkerResponse(worker=worker, reconnect=reconnect)
|
||||||
|
|
||||||
|
def heartbeat(self, session_id: str) -> WorkerSession:
|
||||||
|
worker = self._require_worker(session_id)
|
||||||
|
worker.heartbeat()
|
||||||
|
self.worker_repository.save(worker)
|
||||||
|
self.event_store.append(
|
||||||
|
DomainEvent(
|
||||||
|
event_type="worker_heartbeat",
|
||||||
|
task_id="system",
|
||||||
|
payload={"session_id": worker.session_id, "worker_id": worker.worker_id},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return worker
|
||||||
|
|
||||||
|
def disconnect(self, session_id: str) -> WorkerSession:
|
||||||
|
worker = self._require_worker(session_id)
|
||||||
|
worker.mark_disconnected()
|
||||||
|
self.worker_repository.save(worker)
|
||||||
|
self.event_store.append(
|
||||||
|
DomainEvent(
|
||||||
|
event_type="worker_disconnected",
|
||||||
|
task_id="system",
|
||||||
|
payload={"session_id": worker.session_id, "worker_id": worker.worker_id},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return worker
|
||||||
|
|
||||||
|
def mark_stale_workers(self, stale_before: datetime) -> list[WorkerSession]:
|
||||||
|
updated: list[WorkerSession] = []
|
||||||
|
for worker in self.worker_repository.list_active():
|
||||||
|
if (
|
||||||
|
worker.status != worker.status.__class__.DISCONNECTED
|
||||||
|
and worker.last_heartbeat_at < stale_before
|
||||||
|
):
|
||||||
|
worker.mark_stale()
|
||||||
|
self.worker_repository.save(worker)
|
||||||
|
updated.append(worker)
|
||||||
|
if updated:
|
||||||
|
self.event_store.append(
|
||||||
|
DomainEvent(
|
||||||
|
event_type="worker_sessions_stale",
|
||||||
|
task_id="system",
|
||||||
|
payload={"session_ids": [worker.session_id for worker in updated]},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return updated
|
||||||
|
|
||||||
|
def _require_worker(self, session_id: str) -> WorkerSession:
|
||||||
|
worker = self.worker_repository.get(session_id)
|
||||||
|
if worker is None:
|
||||||
|
raise LookupError(f"Worker session not found: {session_id}")
|
||||||
return worker
|
return worker
|
||||||
|
|||||||
@@ -1,6 +1,8 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from fastapi import FastAPI, HTTPException
|
from datetime import datetime
|
||||||
|
|
||||||
|
from fastapi import FastAPI, HTTPException, WebSocket, WebSocketDisconnect
|
||||||
from sse_starlette.sse import EventSourceResponse
|
from sse_starlette.sse import EventSourceResponse
|
||||||
|
|
||||||
from ai_orchestrator.application.services.execution import (
|
from ai_orchestrator.application.services.execution import (
|
||||||
@@ -14,7 +16,14 @@ from ai_orchestrator.application.services.execution import (
|
|||||||
from ai_orchestrator.application.services.lifecycle import TaskLifecycleService
|
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.router import ConfigurableModelRouter
|
from ai_orchestrator.application.services.router import ConfigurableModelRouter
|
||||||
from ai_orchestrator.application.services.workers import RegisterWorkerRequest, WorkerService
|
from ai_orchestrator.application.services.workers import (
|
||||||
|
CapabilityAwareWorkerGateway,
|
||||||
|
InMemoryWorkerConnectionManager,
|
||||||
|
RegisterWorkerRequest,
|
||||||
|
WorkerCommandResult,
|
||||||
|
WorkerProgressUpdate,
|
||||||
|
WorkerService,
|
||||||
|
)
|
||||||
from ai_orchestrator.config import AppSettings
|
from ai_orchestrator.config import AppSettings
|
||||||
from ai_orchestrator.delivery.http.schemas import (
|
from ai_orchestrator.delivery.http.schemas import (
|
||||||
ChatRequest,
|
ChatRequest,
|
||||||
@@ -24,9 +33,14 @@ from ai_orchestrator.delivery.http.schemas import (
|
|||||||
CreateTaskRequestSchema,
|
CreateTaskRequestSchema,
|
||||||
TaskActionRequest,
|
TaskActionRequest,
|
||||||
TaskStatusResponse,
|
TaskStatusResponse,
|
||||||
|
WorkerDispatchRequest,
|
||||||
|
WorkerDispatchResponse,
|
||||||
|
WorkerHeartbeatRequest,
|
||||||
WorkerListItem,
|
WorkerListItem,
|
||||||
WorkerListResponse,
|
WorkerListResponse,
|
||||||
|
WorkerProgressMessage,
|
||||||
WorkerRegisterRequest,
|
WorkerRegisterRequest,
|
||||||
|
WorkerResultMessage,
|
||||||
)
|
)
|
||||||
from ai_orchestrator.infrastructure.config_loader import load_project_config
|
from ai_orchestrator.infrastructure.config_loader import load_project_config
|
||||||
from ai_orchestrator.infrastructure.model_router import StaticMockModelProvider
|
from ai_orchestrator.infrastructure.model_router import StaticMockModelProvider
|
||||||
@@ -47,6 +61,13 @@ 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}
|
||||||
)
|
)
|
||||||
|
worker_connection_manager = InMemoryWorkerConnectionManager()
|
||||||
|
worker_gateway = CapabilityAwareWorkerGateway(
|
||||||
|
worker_repository=worker_repository,
|
||||||
|
event_store=event_store,
|
||||||
|
invocation_store=invocation_store,
|
||||||
|
connection_manager=worker_connection_manager,
|
||||||
|
)
|
||||||
model_router = ConfigurableModelRouter(
|
model_router = ConfigurableModelRouter(
|
||||||
project_configs={project_config.project.id: project_config},
|
project_configs={project_config.project.id: project_config},
|
||||||
providers={
|
providers={
|
||||||
@@ -86,6 +107,8 @@ def create_app(settings: AppSettings | None = None) -> FastAPI:
|
|||||||
app.state.storage = storage
|
app.state.storage = storage
|
||||||
app.state.invocation_store = invocation_store
|
app.state.invocation_store = invocation_store
|
||||||
app.state.model_router = model_router
|
app.state.model_router = model_router
|
||||||
|
app.state.worker_gateway = worker_gateway
|
||||||
|
app.state.worker_connection_manager = worker_connection_manager
|
||||||
app.state.orchestrator = orchestrator
|
app.state.orchestrator = orchestrator
|
||||||
app.state.execution_engine = execution_engine
|
app.state.execution_engine = execution_engine
|
||||||
app.state.lifecycle_service = lifecycle_service
|
app.state.lifecycle_service = lifecycle_service
|
||||||
@@ -240,7 +263,7 @@ def create_app(settings: AppSettings | None = None) -> FastAPI:
|
|||||||
|
|
||||||
@app.post("/workers/register")
|
@app.post("/workers/register")
|
||||||
def register_worker(payload: WorkerRegisterRequest) -> dict[str, object]:
|
def register_worker(payload: WorkerRegisterRequest) -> dict[str, object]:
|
||||||
worker = app.state.worker_service.register(
|
response = app.state.worker_service.register(
|
||||||
RegisterWorkerRequest(
|
RegisterWorkerRequest(
|
||||||
worker_id=payload.worker_id,
|
worker_id=payload.worker_id,
|
||||||
name=payload.name,
|
name=payload.name,
|
||||||
@@ -250,6 +273,128 @@ def create_app(settings: AppSettings | None = None) -> FastAPI:
|
|||||||
capabilities=payload.capabilities,
|
capabilities=payload.capabilities,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
return {"session_id": worker.session_id, "status": worker.status.value}
|
worker_connection_manager.attach(response.worker.session_id)
|
||||||
|
return {
|
||||||
|
"session_id": response.worker.session_id,
|
||||||
|
"status": response.worker.status.value,
|
||||||
|
"reconnect": response.reconnect,
|
||||||
|
}
|
||||||
|
|
||||||
|
@app.post("/workers/{session_id}/heartbeat")
|
||||||
|
def heartbeat_worker(session_id: str, payload: WorkerHeartbeatRequest) -> dict[str, str]:
|
||||||
|
del payload
|
||||||
|
worker = app.state.worker_service.heartbeat(session_id)
|
||||||
|
return {"status": worker.status.value}
|
||||||
|
|
||||||
|
@app.post("/workers/{session_id}/commands", response_model=WorkerDispatchResponse)
|
||||||
|
def dispatch_worker_command(
|
||||||
|
session_id: str,
|
||||||
|
payload: WorkerDispatchRequest,
|
||||||
|
) -> WorkerDispatchResponse:
|
||||||
|
result = worker_gateway.dispatch(
|
||||||
|
project_id="default",
|
||||||
|
worker_session_id=session_id,
|
||||||
|
command_name=payload.command_name,
|
||||||
|
args=payload.args,
|
||||||
|
task_context={
|
||||||
|
"task_id": payload.task_id,
|
||||||
|
"node_id": payload.node_id,
|
||||||
|
"timeout_ms": payload.timeout_ms,
|
||||||
|
"policy_context": payload.policy_context,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
return WorkerDispatchResponse(
|
||||||
|
status=result.status,
|
||||||
|
command_id=str(result.content["command_id"]),
|
||||||
|
worker_session_id=str(result.content["worker_session_id"]),
|
||||||
|
tool=str(result.content["tool"]),
|
||||||
|
)
|
||||||
|
|
||||||
|
@app.get("/workers/{session_id}/commands")
|
||||||
|
def poll_worker_commands(session_id: str) -> dict[str, list[dict[str, object]]]:
|
||||||
|
commands = worker_gateway.poll_commands(session_id)
|
||||||
|
return {
|
||||||
|
"commands": [
|
||||||
|
{
|
||||||
|
"command_id": command.command_id,
|
||||||
|
"task_id": command.task_id,
|
||||||
|
"tool": command.tool,
|
||||||
|
"args": command.args,
|
||||||
|
"timeout_ms": command.timeout_ms,
|
||||||
|
"policy_context": command.policy_context,
|
||||||
|
}
|
||||||
|
for command in commands
|
||||||
|
]
|
||||||
|
}
|
||||||
|
|
||||||
|
@app.websocket("/workers/ws/{session_id}")
|
||||||
|
async def worker_websocket(session_id: str, websocket: WebSocket) -> None:
|
||||||
|
await websocket.accept()
|
||||||
|
worker_connection_manager.attach(session_id)
|
||||||
|
try:
|
||||||
|
while True:
|
||||||
|
payload = await websocket.receive_json()
|
||||||
|
event_type = payload.get("event_type")
|
||||||
|
if event_type == "heartbeat":
|
||||||
|
worker = app.state.worker_service.heartbeat(session_id)
|
||||||
|
await websocket.send_json({"type": "ack", "status": worker.status.value})
|
||||||
|
elif event_type == "poll":
|
||||||
|
commands = worker_gateway.poll_commands(session_id)
|
||||||
|
await websocket.send_json(
|
||||||
|
{
|
||||||
|
"type": "commands",
|
||||||
|
"commands": [
|
||||||
|
{
|
||||||
|
"command_id": command.command_id,
|
||||||
|
"task_id": command.task_id,
|
||||||
|
"tool": command.tool,
|
||||||
|
"args": command.args,
|
||||||
|
"timeout_ms": command.timeout_ms,
|
||||||
|
"policy_context": command.policy_context,
|
||||||
|
}
|
||||||
|
for command in commands
|
||||||
|
],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
elif event_type == "progress":
|
||||||
|
message = WorkerProgressMessage.model_validate(payload)
|
||||||
|
worker_gateway.record_progress(
|
||||||
|
session_id,
|
||||||
|
WorkerProgressUpdate(
|
||||||
|
command_id=message.command_id,
|
||||||
|
task_id=message.task_id,
|
||||||
|
tool=message.tool,
|
||||||
|
percent=message.percent,
|
||||||
|
message=message.message,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
await websocket.send_json({"type": "ack", "status": "progress_recorded"})
|
||||||
|
elif event_type == "result":
|
||||||
|
message = WorkerResultMessage.model_validate(payload)
|
||||||
|
result = worker_gateway.complete_command(
|
||||||
|
session_id,
|
||||||
|
WorkerCommandResult(
|
||||||
|
command_id=message.command_id,
|
||||||
|
task_id=message.task_id,
|
||||||
|
tool=message.tool,
|
||||||
|
status=message.status,
|
||||||
|
started_at=datetime.fromisoformat(message.started_at),
|
||||||
|
finished_at=datetime.fromisoformat(message.finished_at),
|
||||||
|
duration_ms=message.duration_ms,
|
||||||
|
stdout=message.stdout,
|
||||||
|
stderr=message.stderr,
|
||||||
|
result=message.result,
|
||||||
|
artifacts=message.artifacts,
|
||||||
|
error=message.error,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
await websocket.send_json({"type": "ack", "status": result.status})
|
||||||
|
else:
|
||||||
|
await websocket.send_json(
|
||||||
|
{"type": "error", "message": "Unsupported event_type"}
|
||||||
|
)
|
||||||
|
except WebSocketDisconnect:
|
||||||
|
app.state.worker_service.disconnect(session_id)
|
||||||
|
worker_connection_manager.detach(session_id)
|
||||||
|
|
||||||
return app
|
return app
|
||||||
|
|||||||
@@ -68,3 +68,48 @@ class WorkerListResponse(BaseModel):
|
|||||||
|
|
||||||
class TaskActionRequest(BaseModel):
|
class TaskActionRequest(BaseModel):
|
||||||
reason: str | None = None
|
reason: str | None = None
|
||||||
|
|
||||||
|
|
||||||
|
class WorkerDispatchRequest(BaseModel):
|
||||||
|
task_id: str
|
||||||
|
node_id: str | None = None
|
||||||
|
command_name: str
|
||||||
|
args: dict[str, Any] = Field(default_factory=dict)
|
||||||
|
timeout_ms: int = 30000
|
||||||
|
policy_context: dict[str, Any] = Field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
|
class WorkerDispatchResponse(BaseModel):
|
||||||
|
status: str
|
||||||
|
command_id: str
|
||||||
|
worker_session_id: str
|
||||||
|
tool: str
|
||||||
|
|
||||||
|
|
||||||
|
class WorkerHeartbeatRequest(BaseModel):
|
||||||
|
event_type: str = "heartbeat"
|
||||||
|
|
||||||
|
|
||||||
|
class WorkerProgressMessage(BaseModel):
|
||||||
|
event_type: str = "progress"
|
||||||
|
command_id: str
|
||||||
|
task_id: str
|
||||||
|
tool: str
|
||||||
|
percent: int
|
||||||
|
message: str
|
||||||
|
|
||||||
|
|
||||||
|
class WorkerResultMessage(BaseModel):
|
||||||
|
event_type: str = "result"
|
||||||
|
command_id: str
|
||||||
|
task_id: str
|
||||||
|
tool: str
|
||||||
|
status: str
|
||||||
|
started_at: str
|
||||||
|
finished_at: str
|
||||||
|
duration_ms: int
|
||||||
|
stdout: str = ""
|
||||||
|
stderr: str = ""
|
||||||
|
result: dict[str, Any] = Field(default_factory=dict)
|
||||||
|
artifacts: list[dict[str, Any]] = Field(default_factory=list)
|
||||||
|
error: dict[str, Any] | None = None
|
||||||
|
|||||||
@@ -186,7 +186,22 @@ class WorkerSession:
|
|||||||
def mark_online(self) -> None:
|
def mark_online(self) -> None:
|
||||||
self.status = WorkerSessionStatus.ONLINE
|
self.status = WorkerSessionStatus.ONLINE
|
||||||
self.last_heartbeat_at = _utcnow()
|
self.last_heartbeat_at = _utcnow()
|
||||||
|
self.disconnected_at = None
|
||||||
|
|
||||||
def heartbeat(self) -> None:
|
def heartbeat(self) -> None:
|
||||||
self.last_heartbeat_at = _utcnow()
|
self.last_heartbeat_at = _utcnow()
|
||||||
|
|
||||||
|
def mark_busy(self, task_id: str | None = None) -> None:
|
||||||
|
self.status = WorkerSessionStatus.BUSY
|
||||||
|
self.current_task_id = task_id
|
||||||
|
self.last_heartbeat_at = _utcnow()
|
||||||
|
|
||||||
|
def mark_stale(self) -> None:
|
||||||
|
self.status = WorkerSessionStatus.STALE
|
||||||
|
self.last_heartbeat_at = _utcnow()
|
||||||
|
|
||||||
|
def mark_disconnected(self) -> None:
|
||||||
|
self.status = WorkerSessionStatus.DISCONNECTED
|
||||||
|
self.disconnected_at = _utcnow()
|
||||||
|
self.current_task_id = None
|
||||||
|
self.last_heartbeat_at = self.disconnected_at
|
||||||
|
|||||||
@@ -69,3 +69,72 @@ def test_worker_registration_is_exposed_via_list_endpoint() -> None:
|
|||||||
assert registration.status_code == 200
|
assert registration.status_code == 200
|
||||||
assert listing.status_code == 200
|
assert listing.status_code == 200
|
||||||
assert any(worker["worker_id"] == "worker_home_pc" for worker in listing.json()["workers"])
|
assert any(worker["worker_id"] == "worker_home_pc" for worker in listing.json()["workers"])
|
||||||
|
|
||||||
|
|
||||||
|
def test_worker_command_dispatch_and_poll_cycle() -> None:
|
||||||
|
client = TestClient(app)
|
||||||
|
|
||||||
|
registration = client.post(
|
||||||
|
"/workers/register",
|
||||||
|
json={
|
||||||
|
"worker_id": "worker_home_pc",
|
||||||
|
"name": "Home PC",
|
||||||
|
"capabilities": ["file.read", "command.run"],
|
||||||
|
"version": "0.1.0",
|
||||||
|
"machine": "DESKTOP-1",
|
||||||
|
"os": "windows",
|
||||||
|
},
|
||||||
|
).json()
|
||||||
|
dispatch = client.post(
|
||||||
|
f"/workers/{registration['session_id']}/commands",
|
||||||
|
json={
|
||||||
|
"task_id": "task_1",
|
||||||
|
"node_id": "node_1",
|
||||||
|
"command_name": "file.read",
|
||||||
|
"args": {"path": "D:/test.txt"},
|
||||||
|
"timeout_ms": 30000,
|
||||||
|
"policy_context": {"approved": True},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
polled = client.get(f"/workers/{registration['session_id']}/commands")
|
||||||
|
|
||||||
|
assert dispatch.status_code == 200
|
||||||
|
assert dispatch.json()["status"] == "queued"
|
||||||
|
assert len(polled.json()["commands"]) == 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_worker_websocket_heartbeat_and_poll() -> None:
|
||||||
|
client = TestClient(app)
|
||||||
|
registration = client.post(
|
||||||
|
"/workers/register",
|
||||||
|
json={
|
||||||
|
"worker_id": "worker_socket_pc",
|
||||||
|
"name": "Socket PC",
|
||||||
|
"capabilities": ["file.read"],
|
||||||
|
"version": "0.1.0",
|
||||||
|
"machine": "DESKTOP-2",
|
||||||
|
"os": "windows",
|
||||||
|
},
|
||||||
|
).json()
|
||||||
|
client.post(
|
||||||
|
f"/workers/{registration['session_id']}/commands",
|
||||||
|
json={
|
||||||
|
"task_id": "task_2",
|
||||||
|
"node_id": "node_2",
|
||||||
|
"command_name": "file.read",
|
||||||
|
"args": {"path": "D:/socket.txt"},
|
||||||
|
"timeout_ms": 30000,
|
||||||
|
"policy_context": {"approved": True},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
with client.websocket_connect(f"/workers/ws/{registration['session_id']}") as websocket:
|
||||||
|
websocket.send_json({"event_type": "heartbeat"})
|
||||||
|
heartbeat_ack = websocket.receive_json()
|
||||||
|
websocket.send_json({"event_type": "poll"})
|
||||||
|
commands = websocket.receive_json()
|
||||||
|
|
||||||
|
assert heartbeat_ack["type"] == "ack"
|
||||||
|
assert heartbeat_ack["status"] in {"online", "busy"}
|
||||||
|
assert commands["type"] == "commands"
|
||||||
|
assert len(commands["commands"]) == 1
|
||||||
|
|||||||
@@ -0,0 +1,124 @@
|
|||||||
|
from datetime import UTC, datetime
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from ai_orchestrator.application.services.workers import (
|
||||||
|
CapabilityAwareWorkerGateway,
|
||||||
|
InMemoryWorkerConnectionManager,
|
||||||
|
RegisterWorkerRequest,
|
||||||
|
WorkerCommandResult,
|
||||||
|
WorkerService,
|
||||||
|
)
|
||||||
|
from ai_orchestrator.infrastructure.storage.memory import (
|
||||||
|
InMemoryEventStore,
|
||||||
|
InMemoryInvocationStore,
|
||||||
|
InMemoryWorkerRepository,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _build_gateway() -> tuple[
|
||||||
|
CapabilityAwareWorkerGateway,
|
||||||
|
WorkerService,
|
||||||
|
InMemoryWorkerRepository,
|
||||||
|
InMemoryEventStore,
|
||||||
|
InMemoryInvocationStore,
|
||||||
|
]:
|
||||||
|
worker_repository = InMemoryWorkerRepository()
|
||||||
|
event_store = InMemoryEventStore()
|
||||||
|
invocation_store = InMemoryInvocationStore()
|
||||||
|
service = WorkerService(worker_repository=worker_repository, event_store=event_store)
|
||||||
|
gateway = CapabilityAwareWorkerGateway(
|
||||||
|
worker_repository=worker_repository,
|
||||||
|
event_store=event_store,
|
||||||
|
invocation_store=invocation_store,
|
||||||
|
connection_manager=InMemoryWorkerConnectionManager(),
|
||||||
|
)
|
||||||
|
return gateway, service, worker_repository, event_store, invocation_store
|
||||||
|
|
||||||
|
|
||||||
|
def test_worker_gateway_dispatches_and_completes_command() -> None:
|
||||||
|
gateway, service, worker_repository, event_store, invocation_store = _build_gateway()
|
||||||
|
response = service.register(
|
||||||
|
RegisterWorkerRequest(
|
||||||
|
worker_id="worker_home_pc",
|
||||||
|
name="Home PC",
|
||||||
|
machine="DESKTOP-1",
|
||||||
|
os="windows",
|
||||||
|
version="0.1.0",
|
||||||
|
capabilities=["file.read"],
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
queued = gateway.dispatch(
|
||||||
|
project_id="default",
|
||||||
|
worker_session_id=response.worker.session_id,
|
||||||
|
command_name="file.read",
|
||||||
|
args={"path": "D:/test.txt"},
|
||||||
|
task_context={"task_id": "task_1", "node_id": "node_1", "timeout_ms": 30000},
|
||||||
|
)
|
||||||
|
commands = gateway.poll_commands(response.worker.session_id)
|
||||||
|
completed = gateway.complete_command(
|
||||||
|
response.worker.session_id,
|
||||||
|
WorkerCommandResult(
|
||||||
|
command_id=commands[0].command_id,
|
||||||
|
task_id="task_1",
|
||||||
|
tool="file.read",
|
||||||
|
status="success",
|
||||||
|
started_at=datetime.now(UTC),
|
||||||
|
finished_at=datetime.now(UTC),
|
||||||
|
duration_ms=100,
|
||||||
|
result={"content": "ok"},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert queued.status == "queued"
|
||||||
|
assert len(commands) == 1
|
||||||
|
assert completed.status == "success"
|
||||||
|
assert worker_repository.get(response.worker.session_id).status.value == "online"
|
||||||
|
assert len(invocation_store.list_tool_invocations("task_1")) == 2
|
||||||
|
assert event_store.items[-1].event_type == "worker_command_completed"
|
||||||
|
|
||||||
|
|
||||||
|
def test_worker_gateway_rejects_unsupported_capability() -> None:
|
||||||
|
gateway, service, _, _, _ = _build_gateway()
|
||||||
|
response = service.register(
|
||||||
|
RegisterWorkerRequest(
|
||||||
|
worker_id="worker_home_pc",
|
||||||
|
name="Home PC",
|
||||||
|
machine="DESKTOP-1",
|
||||||
|
os="windows",
|
||||||
|
version="0.1.0",
|
||||||
|
capabilities=["file.read"],
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(ValueError):
|
||||||
|
gateway.dispatch(
|
||||||
|
project_id="default",
|
||||||
|
worker_session_id=response.worker.session_id,
|
||||||
|
command_name="command.run",
|
||||||
|
args={},
|
||||||
|
task_context={"task_id": "task_1"},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_worker_service_heartbeat_and_stale_tracking() -> None:
|
||||||
|
_, service, worker_repository, event_store, _ = _build_gateway()
|
||||||
|
response = service.register(
|
||||||
|
RegisterWorkerRequest(
|
||||||
|
worker_id="worker_home_pc",
|
||||||
|
name="Home PC",
|
||||||
|
machine="DESKTOP-1",
|
||||||
|
os="windows",
|
||||||
|
version="0.1.0",
|
||||||
|
capabilities=["file.read"],
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
worker = service.heartbeat(response.worker.session_id)
|
||||||
|
worker.last_heartbeat_at = datetime(2000, 1, 1, tzinfo=UTC)
|
||||||
|
worker_repository.save(worker)
|
||||||
|
stale = service.mark_stale_workers(datetime(2001, 1, 1, tzinfo=UTC))
|
||||||
|
|
||||||
|
assert stale[0].status.value == "stale"
|
||||||
|
assert event_store.items[-1].event_type == "worker_sessions_stale"
|
||||||
@@ -10,7 +10,7 @@ def test_worker_registration_creates_online_session_and_event() -> None:
|
|||||||
event_store = InMemoryEventStore()
|
event_store = InMemoryEventStore()
|
||||||
service = WorkerService(worker_repository=worker_repository, event_store=event_store)
|
service = WorkerService(worker_repository=worker_repository, event_store=event_store)
|
||||||
|
|
||||||
worker = service.register(
|
response = service.register(
|
||||||
RegisterWorkerRequest(
|
RegisterWorkerRequest(
|
||||||
worker_id="worker_home_pc",
|
worker_id="worker_home_pc",
|
||||||
name="Home PC",
|
name="Home PC",
|
||||||
@@ -21,6 +21,9 @@ def test_worker_registration_creates_online_session_and_event() -> None:
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
worker = response.worker
|
||||||
|
|
||||||
assert worker.status.value == "online"
|
assert worker.status.value == "online"
|
||||||
assert worker_repository.get(worker.session_id) is not None
|
assert worker_repository.get(worker.session_id) is not None
|
||||||
assert event_store.items[-1].event_type == "worker_registered"
|
assert event_store.items[-1].event_type == "worker_registered"
|
||||||
|
assert response.reconnect is False
|
||||||
|
|||||||
Reference in New Issue
Block a user