Add worker gateway and websocket transport flow
This commit is contained in:
@@ -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()
|
||||
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
|
||||
|
||||
Reference in New Issue
Block a user