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 ( InMemoryArtifactStore, InMemoryEventStore, InMemoryInvocationStore, InMemoryWorkerRepository, ) def _build_gateway() -> tuple[ CapabilityAwareWorkerGateway, WorkerService, InMemoryWorkerRepository, InMemoryEventStore, InMemoryInvocationStore, InMemoryArtifactStore, ]: worker_repository = InMemoryWorkerRepository() event_store = InMemoryEventStore() invocation_store = InMemoryInvocationStore() artifact_store = InMemoryArtifactStore() service = WorkerService(worker_repository=worker_repository, event_store=event_store) gateway = CapabilityAwareWorkerGateway( worker_repository=worker_repository, event_store=event_store, invocation_store=invocation_store, artifact_store=artifact_store, connection_manager=InMemoryWorkerConnectionManager(), ) return gateway, service, worker_repository, event_store, invocation_store, artifact_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"