From 8423d8e8ad18fd77317971afd6312cc8365cb4ae Mon Sep 17 00:00:00 2001 From: Mikhail Date: Fri, 3 Jul 2026 21:30:17 +0300 Subject: [PATCH] Add worker gateway and websocket transport flow --- docs/backlog/01_execution_backlog.md | 8 +- .../application/services/workers.py | 320 +++++++++++++++++- src/ai_orchestrator/delivery/http/app.py | 153 ++++++++- src/ai_orchestrator/delivery/http/schemas.py | 45 +++ src/ai_orchestrator/domain/models.py | 15 + tests/integration/test_http_app.py | 69 ++++ tests/unit/test_worker_gateway.py | 124 +++++++ tests/unit/test_worker_service.py | 5 +- 8 files changed, 727 insertions(+), 12 deletions(-) create mode 100644 tests/unit/test_worker_gateway.py diff --git a/docs/backlog/01_execution_backlog.md b/docs/backlog/01_execution_backlog.md index ddd126c..3da1ee0 100644 --- a/docs/backlog/01_execution_backlog.md +++ b/docs/backlog/01_execution_backlog.md @@ -38,13 +38,13 @@ - [x] model router providers - [x] MCP transport adapter -- [ ] worker WebSocket gateway -- [ ] heartbeat monitor -- [ ] capability-aware dispatch +- [x] worker WebSocket gateway +- [x] heartbeat monitor +- [x] capability-aware dispatch ## Package E: Delivery -- [ ] typed worker registration schema +- [x] typed worker registration schema - [ ] task event DTO normalization - [ ] confirmation cards - [ ] progress cards diff --git a/src/ai_orchestrator/application/services/workers.py b/src/ai_orchestrator/application/services/workers.py index 92de9ea..69c776a 100644 --- a/src/ai_orchestrator/application/services/workers.py +++ b/src/ai_orchestrator/application/services/workers.py @@ -1,8 +1,17 @@ from __future__ import annotations 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.models import WorkerSession @@ -17,13 +26,259 @@ class RegisterWorkerRequest: 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) class WorkerService: worker_repository: WorkerRepository event_store: EventStore - def register(self, request: RegisterWorkerRequest) -> WorkerSession: - worker = WorkerSession( + def register(self, request: RegisterWorkerRequest) -> RegisterWorkerResponse: + 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, name=request.name, machine=request.machine, @@ -31,6 +286,12 @@ class WorkerService: version=request.version, 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() self.worker_repository.save(worker) self.event_store.append( @@ -41,7 +302,60 @@ class WorkerService: "session_id": worker.session_id, "worker_id": worker.worker_id, "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 diff --git a/src/ai_orchestrator/delivery/http/app.py b/src/ai_orchestrator/delivery/http/app.py index 0c90fbf..1eab816 100644 --- a/src/ai_orchestrator/delivery/http/app.py +++ b/src/ai_orchestrator/delivery/http/app.py @@ -1,6 +1,8 @@ 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 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.orchestrator import CreateTaskRequest, OrchestratorService 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.delivery.http.schemas import ( ChatRequest, @@ -24,9 +33,14 @@ from ai_orchestrator.delivery.http.schemas import ( CreateTaskRequestSchema, TaskActionRequest, TaskStatusResponse, + WorkerDispatchRequest, + WorkerDispatchResponse, + WorkerHeartbeatRequest, WorkerListItem, WorkerListResponse, + WorkerProgressMessage, WorkerRegisterRequest, + WorkerResultMessage, ) from ai_orchestrator.infrastructure.config_loader import load_project_config from ai_orchestrator.infrastructure.model_router import StaticMockModelProvider @@ -47,6 +61,13 @@ def create_app(settings: AppSettings | None = None) -> FastAPI: policy_evaluator = StaticProjectPolicyEvaluator( 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( project_configs={project_config.project.id: project_config}, providers={ @@ -86,6 +107,8 @@ def create_app(settings: AppSettings | None = None) -> FastAPI: app.state.storage = storage app.state.invocation_store = invocation_store 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.execution_engine = execution_engine app.state.lifecycle_service = lifecycle_service @@ -240,7 +263,7 @@ def create_app(settings: AppSettings | None = None) -> FastAPI: @app.post("/workers/register") def register_worker(payload: WorkerRegisterRequest) -> dict[str, object]: - worker = app.state.worker_service.register( + response = app.state.worker_service.register( RegisterWorkerRequest( worker_id=payload.worker_id, name=payload.name, @@ -250,6 +273,128 @@ def create_app(settings: AppSettings | None = None) -> FastAPI: 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 diff --git a/src/ai_orchestrator/delivery/http/schemas.py b/src/ai_orchestrator/delivery/http/schemas.py index ff86712..5807d81 100644 --- a/src/ai_orchestrator/delivery/http/schemas.py +++ b/src/ai_orchestrator/delivery/http/schemas.py @@ -68,3 +68,48 @@ class WorkerListResponse(BaseModel): class TaskActionRequest(BaseModel): 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 diff --git a/src/ai_orchestrator/domain/models.py b/src/ai_orchestrator/domain/models.py index 8dd979a..a7f73de 100644 --- a/src/ai_orchestrator/domain/models.py +++ b/src/ai_orchestrator/domain/models.py @@ -186,7 +186,22 @@ class WorkerSession: def mark_online(self) -> None: self.status = WorkerSessionStatus.ONLINE self.last_heartbeat_at = _utcnow() + self.disconnected_at = None def heartbeat(self) -> None: 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 diff --git a/tests/integration/test_http_app.py b/tests/integration/test_http_app.py index ce83a65..a423737 100644 --- a/tests/integration/test_http_app.py +++ b/tests/integration/test_http_app.py @@ -69,3 +69,72 @@ def test_worker_registration_is_exposed_via_list_endpoint() -> None: assert registration.status_code == 200 assert listing.status_code == 200 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 diff --git a/tests/unit/test_worker_gateway.py b/tests/unit/test_worker_gateway.py new file mode 100644 index 0000000..3b319de --- /dev/null +++ b/tests/unit/test_worker_gateway.py @@ -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" diff --git a/tests/unit/test_worker_service.py b/tests/unit/test_worker_service.py index 033fcb8..455459f 100644 --- a/tests/unit/test_worker_service.py +++ b/tests/unit/test_worker_service.py @@ -10,7 +10,7 @@ def test_worker_registration_creates_online_session_and_event() -> None: event_store = InMemoryEventStore() service = WorkerService(worker_repository=worker_repository, event_store=event_store) - worker = service.register( + response = service.register( RegisterWorkerRequest( worker_id="worker_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_repository.get(worker.session_id) is not None assert event_store.items[-1].event_type == "worker_registered" + assert response.reconnect is False