Add worker gateway and websocket transport flow

This commit is contained in:
2026-07-03 21:30:17 +03:00
parent b681a90eae
commit 8423d8e8ad
8 changed files with 727 additions and 12 deletions
+69
View File
@@ -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
+124
View File
@@ -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"
+4 -1
View File
@@ -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