← back to Exo
Fix download tests
36a5d75efd5b41bc3a21b8ea92073268a74f5547 · 2025-07-28 12:51:10 +0100 · Matt Beton
Files touched
M worker/main.pyM worker/tests/conftest.pyM worker/tests/test_serdes.pyM worker/tests/test_supervisor.pyM worker/tests/test_worker_handlers.pyM worker/tests/test_worker_integration.py
Diff
commit 36a5d75efd5b41bc3a21b8ea92073268a74f5547
Author: Matt Beton <matthew.beton@gmail.com>
Date: Mon Jul 28 12:51:10 2025 +0100
Fix download tests
---
worker/main.py | 7 ++--
worker/tests/conftest.py | 6 ++--
worker/tests/test_serdes.py | 6 ++--
worker/tests/test_supervisor.py | 15 +++++----
worker/tests/test_worker_handlers.py | 23 ++++++++-----
worker/tests/test_worker_integration.py | 59 +++++++++++++++++++++------------
6 files changed, 70 insertions(+), 46 deletions(-)
diff --git a/worker/main.py b/worker/main.py
index 4c40d826..42cf9850 100644
--- a/worker/main.py
+++ b/worker/main.py
@@ -44,6 +44,7 @@ from shared.types.worker.ops import (
UnassignRunnerOp,
)
from shared.types.worker.runners import (
+ AssignedRunnerStatus,
DownloadingRunnerStatus,
FailedRunnerStatus,
LoadedRunnerStatus,
@@ -115,7 +116,7 @@ class Worker:
instance_id=op.instance_id,
shard_metadata=op.shard_metadata,
hosts=op.hosts,
- status=ReadyRunnerStatus(),
+ status=AssignedRunnerStatus(),
runner=None,
)
@@ -232,6 +233,7 @@ class Worker:
asyncio.create_task(self.shard_downloader.ensure_shard(op.shard_metadata))
+ # TODO: Dynamic timeout, timeout on no packet update received.
timeout_secs = 10 * 60
start_time = process_time()
last_yield_progress = start_time
@@ -472,7 +474,8 @@ class Worker:
runner = self.assigned_runners[runner_id]
if not runner.is_downloaded:
- if runner.status.runner_status == RunnerStatusType.Downloading:
+ if runner.status.runner_status == RunnerStatusType.Downloading: # Forward compatibility
+ # TODO: If failed status then we retry
return None
else:
return DownloadOp(
diff --git a/worker/tests/conftest.py b/worker/tests/conftest.py
index 1808323b..2548fd05 100644
--- a/worker/tests/conftest.py
+++ b/worker/tests/conftest.py
@@ -101,9 +101,9 @@ def completion_create_params(user_message: str) -> ChatCompletionTaskParams:
@pytest.fixture
def chat_completion_task(completion_create_params: ChatCompletionTaskParams):
- def _chat_completion_task(instance_id: InstanceId) -> ChatCompletionTask:
+ def _chat_completion_task(instance_id: InstanceId, task_id: TaskId) -> ChatCompletionTask:
return ChatCompletionTask(
- task_id=TaskId(),
+ task_id=task_id,
command_id=CommandId(),
instance_id=instance_id,
task_type=TaskType.CHAT_COMPLETION,
@@ -145,7 +145,7 @@ def instance(pipeline_shard_meta: Callable[[int, int], PipelineShardMetadata], h
)
return Instance(
- instance_id=InstanceId(),
+ instance_id=instance_id,
instance_type=InstanceStatus.ACTIVE,
shard_assignments=shard_assignments,
hosts=hosts_one
diff --git a/worker/tests/test_serdes.py b/worker/tests/test_serdes.py
index 37fe515a..fd5fdeb7 100644
--- a/worker/tests/test_serdes.py
+++ b/worker/tests/test_serdes.py
@@ -3,8 +3,8 @@ from typing import Callable, TypeVar
from pydantic import BaseModel, TypeAdapter
+from shared.types.tasks import Task, TaskId
from shared.types.common import Host
-from shared.types.tasks import Task
from shared.types.worker.commands_runner import (
ChatTaskMessage,
RunnerMessageTypeAdapter,
@@ -38,9 +38,9 @@ def test_supervisor_setup_message_serdes(
def test_supervisor_task_message_serdes(
- chat_completion_task: Callable[[InstanceId], Task],
+ chat_completion_task: Callable[[InstanceId, TaskId], Task],
):
- task = chat_completion_task(InstanceId())
+ task = chat_completion_task(InstanceId(), TaskId())
task_message = ChatTaskMessage(
task_data=task.task_params,
)
diff --git a/worker/tests/test_supervisor.py b/worker/tests/test_supervisor.py
index 77cebdf1..1db5a7a2 100644
--- a/worker/tests/test_supervisor.py
+++ b/worker/tests/test_supervisor.py
@@ -10,6 +10,7 @@ from shared.types.events.chunks import TokenChunk
from shared.types.tasks import (
ChatCompletionTaskParams,
Task,
+ TaskId,
TaskType,
)
from shared.types.worker.common import InstanceId
@@ -27,7 +28,7 @@ def user_message():
async def test_supervisor_single_node_response(
pipeline_shard_meta: Callable[..., PipelineShardMetadata],
hosts: Callable[..., list[Host]],
- chat_completion_task: Callable[[InstanceId], Task],
+ chat_completion_task: Callable[[InstanceId, TaskId], Task],
tmp_path: Path,
):
"""Test that asking for the capital of France returns 'Paris' in the response"""
@@ -45,7 +46,7 @@ async def test_supervisor_single_node_response(
full_response = ""
stop_reason: FinishReason | None = None
- async for chunk in supervisor.stream_response(task=chat_completion_task(instance_id)):
+ async for chunk in supervisor.stream_response(task=chat_completion_task(instance_id, TaskId())):
if isinstance(chunk, TokenChunk):
full_response += chunk.text
if chunk.finish_reason:
@@ -65,7 +66,7 @@ async def test_supervisor_single_node_response(
async def test_supervisor_two_node_response(
pipeline_shard_meta: Callable[..., PipelineShardMetadata],
hosts: Callable[..., list[Host]],
- chat_completion_task: Callable[[InstanceId], Task],
+ chat_completion_task: Callable[[InstanceId, TaskId], Task],
tmp_path: Path,
):
"""Test that asking for the capital of France returns 'Paris' in the response"""
@@ -88,13 +89,13 @@ async def test_supervisor_two_node_response(
async def collect_response_0():
nonlocal full_response_0
- async for chunk in supervisor_0.stream_response(task=chat_completion_task(instance_id)):
+ async for chunk in supervisor_0.stream_response(task=chat_completion_task(instance_id, TaskId())):
if isinstance(chunk, TokenChunk):
full_response_0 += chunk.text
async def collect_response_1():
nonlocal full_response_1
- async for chunk in supervisor_1.stream_response(task=chat_completion_task(instance_id)):
+ async for chunk in supervisor_1.stream_response(task=chat_completion_task(instance_id, TaskId())):
if isinstance(chunk, TokenChunk):
full_response_1 += chunk.text
@@ -121,7 +122,7 @@ async def test_supervisor_two_node_response(
async def test_supervisor_early_stopping(
pipeline_shard_meta: Callable[..., PipelineShardMetadata],
hosts: Callable[..., list[Host]],
- chat_completion_task: Callable[[InstanceId], Task],
+ chat_completion_task: Callable[[InstanceId, TaskId], Task],
tmp_path: Path,
):
"""Test that asking for the capital of France returns 'Paris' in the response"""
@@ -133,7 +134,7 @@ async def test_supervisor_early_stopping(
hosts=hosts(1, offset=10),
)
- task = chat_completion_task(instance_id)
+ task = chat_completion_task(instance_id, TaskId())
max_tokens = 50
assert task.task_type == TaskType.CHAT_COMPLETION
diff --git a/worker/tests/test_worker_handlers.py b/worker/tests/test_worker_handlers.py
index ef5c634e..ed2fed95 100644
--- a/worker/tests/test_worker_handlers.py
+++ b/worker/tests/test_worker_handlers.py
@@ -14,7 +14,7 @@ from shared.types.events import (
TaskStateUpdated,
)
from shared.types.events.chunks import TokenChunk
-from shared.types.tasks import Task, TaskStatus
+from shared.types.tasks import Task, TaskId, TaskStatus
from shared.types.worker.common import RunnerId
from shared.types.worker.instances import Instance, InstanceId
from shared.types.worker.ops import (
@@ -26,6 +26,7 @@ from shared.types.worker.ops import (
UnassignRunnerOp,
)
from shared.types.worker.runners import (
+ AssignedRunnerStatus,
FailedRunnerStatus,
LoadedRunnerStatus,
ReadyRunnerStatus,
@@ -59,11 +60,11 @@ async def test_assign_op(worker: Worker, instance: Callable[[InstanceId, NodeId,
# We should have a status update saying 'starting'.
assert len(events) == 1
assert isinstance(events[0], RunnerStatusUpdated)
- assert isinstance(events[0].runner_status, ReadyRunnerStatus)
+ assert isinstance(events[0].runner_status, AssignedRunnerStatus)
# And the runner should be assigned
assert runner_id in worker.assigned_runners
- assert isinstance(worker.assigned_runners[runner_id].status, ReadyRunnerStatus)
+ assert isinstance(worker.assigned_runners[runner_id].status, AssignedRunnerStatus)
@pytest.mark.asyncio
async def test_unassign_op(worker_with_assigned_runner: tuple[Worker, RunnerId, Instance], tmp_path: Path):
@@ -84,7 +85,11 @@ async def test_unassign_op(worker_with_assigned_runner: tuple[Worker, RunnerId,
assert isinstance(events[0], RunnerDeleted)
@pytest.mark.asyncio
-async def test_runner_up_op(worker_with_assigned_runner: tuple[Worker, RunnerId, Instance], chat_completion_task: Callable[[InstanceId], Task], tmp_path: Path):
+async def test_runner_up_op(
+ worker_with_assigned_runner: tuple[Worker, RunnerId, Instance],
+ chat_completion_task: Callable[[InstanceId, TaskId], Task],
+ tmp_path: Path
+ ):
worker, runner_id, _ = worker_with_assigned_runner
runner_up_op = RunnerUpOp(runner_id=runner_id)
@@ -104,7 +109,7 @@ async def test_runner_up_op(worker_with_assigned_runner: tuple[Worker, RunnerId,
full_response = ''
- async for chunk in supervisor.stream_response(task=chat_completion_task(InstanceId())):
+ async for chunk in supervisor.stream_response(task=chat_completion_task(InstanceId(), TaskId())):
if isinstance(chunk, TokenChunk):
full_response += chunk.text
@@ -153,12 +158,12 @@ async def test_download_op(worker_with_assigned_runner: tuple[Worker, RunnerId,
@pytest.mark.asyncio
async def test_execute_task_op(
worker_with_running_runner: tuple[Worker, RunnerId, Instance],
- chat_completion_task: Callable[[InstanceId], Task], tmp_path: Path):
+ chat_completion_task: Callable[[InstanceId, TaskId], Task], tmp_path: Path):
worker, runner_id, _ = worker_with_running_runner
execute_task_op = ExecuteTaskOp(
runner_id=runner_id,
- task=chat_completion_task(InstanceId())
+ task=chat_completion_task(InstanceId(), TaskId())
)
events: list[Event] = []
@@ -196,10 +201,10 @@ async def test_execute_task_op(
@pytest.mark.asyncio
async def test_execute_task_fails(
worker_with_running_runner: tuple[Worker, RunnerId, Instance],
- chat_completion_task: Callable[[InstanceId], Task], tmp_path: Path):
+ chat_completion_task: Callable[[InstanceId, TaskId], Task], tmp_path: Path):
worker, runner_id, _ = worker_with_running_runner
- task = chat_completion_task(InstanceId())
+ task = chat_completion_task(InstanceId(), TaskId())
messages = task.task_params.messages
messages[0].content = 'Artificial prompt: EXO RUNNER MUST FAIL'
diff --git a/worker/tests/test_worker_integration.py b/worker/tests/test_worker_integration.py
index cbd6a681..63e3abbd 100644
--- a/worker/tests/test_worker_integration.py
+++ b/worker/tests/test_worker_integration.py
@@ -25,10 +25,12 @@ from shared.types.worker.instances import (
ShardAssignments,
)
from shared.types.worker.runners import (
+ AssignedRunnerStatus,
+ DownloadingRunnerStatus,
+ # RunningRunnerStatus,
FailedRunnerStatus,
LoadedRunnerStatus,
ReadyRunnerStatus,
- # RunningRunnerStatus,
)
from shared.types.worker.shards import PipelineShardMetadata
from worker.download.shard_downloader import NoopShardDownloader
@@ -40,13 +42,14 @@ NODE_A: Final[NodeId] = NodeId("aaaaaaaa-aaaa-4aaa-8aaa-aaaaaaaaaaaa")
NODE_B: Final[NodeId] = NodeId("bbbbbbbb-bbbb-4bbb-8bbb-bbbbbbbbbbbb")
# Define constant IDs for deterministic test cases
-RUNNER_1_ID: Final[RunnerId] = RunnerId()
-INSTANCE_1_ID: Final[InstanceId] = InstanceId()
-RUNNER_2_ID: Final[RunnerId] = RunnerId()
-INSTANCE_2_ID: Final[InstanceId] = InstanceId()
+RUNNER_1_ID: Final[RunnerId] = RunnerId("11111111-1111-4111-8111-111111111111")
+INSTANCE_1_ID: Final[InstanceId] = InstanceId("22222222-2222-4222-8222-222222222222")
+RUNNER_2_ID: Final[RunnerId] = RunnerId("33333333-3333-4333-8333-333333333333")
+INSTANCE_2_ID: Final[InstanceId] = InstanceId("44444444-4444-4444-8444-444444444444")
MODEL_A_ID: Final[ModelId] = 'mlx-community/Llama-3.2-1B-Instruct-4bit'
MODEL_B_ID: Final[ModelId] = 'mlx-community/Llama-3.2-1B-Instruct-4bit'
-TASK_1_ID: Final[TaskId] = TaskId()
+TASK_1_ID: Final[TaskId] = TaskId("55555555-5555-4555-8555-555555555555")
+TASK_2_ID: Final[TaskId] = TaskId("66666666-6666-4666-8666-666666666666")
@pytest.fixture
def user_message():
@@ -82,9 +85,15 @@ async def test_runner_assigned(
# Ensure the correct events have been emitted
events = await global_events.get_events_since(0)
- assert len(events) == 2
+ print(events)
+ assert len(events) >= 4 # len(events) is 4 if it's already downloaded. It is > 4 if there have to be download events.
+
assert isinstance(events[1].event, RunnerStatusUpdated)
- assert isinstance(events[1].event.runner_status, ReadyRunnerStatus)
+ assert isinstance(events[1].event.runner_status, AssignedRunnerStatus)
+ assert isinstance(events[2].event, RunnerStatusUpdated)
+ assert isinstance(events[2].event.runner_status, DownloadingRunnerStatus)
+ assert isinstance(events[-1].event, RunnerStatusUpdated)
+ assert isinstance(events[-1].event.runner_status, ReadyRunnerStatus)
# Ensure state is correct
assert isinstance(worker.state.runners[RUNNER_1_ID], ReadyRunnerStatus)
@@ -92,7 +101,7 @@ async def test_runner_assigned(
async def test_runner_assigned_active(
worker_running: Callable[[NodeId], Awaitable[tuple[Worker, AsyncSQLiteEventStorage]]],
instance: Callable[[InstanceId, NodeId, RunnerId], Instance],
- chat_completion_task: Callable[[InstanceId], Task]
+ chat_completion_task: Callable[[InstanceId, TaskId], Task]
):
worker, global_events = await worker_running(NODE_A)
@@ -116,9 +125,15 @@ async def test_runner_assigned_active(
# Ensure the correct events have been emitted
events = await global_events.get_events_since(0)
- assert len(events) == 3
+ assert len(events) >= 5 # len(events) is 5 if it's already downloaded. It is > 5 if there have to be download events.
+ assert isinstance(events[1].event, RunnerStatusUpdated)
+ assert isinstance(events[1].event.runner_status, AssignedRunnerStatus)
assert isinstance(events[2].event, RunnerStatusUpdated)
- assert isinstance(events[2].event.runner_status, LoadedRunnerStatus)
+ assert isinstance(events[2].event.runner_status, DownloadingRunnerStatus)
+ assert isinstance(events[-2].event, RunnerStatusUpdated)
+ assert isinstance(events[-2].event.runner_status, ReadyRunnerStatus)
+ assert isinstance(events[-1].event, RunnerStatusUpdated)
+ assert isinstance(events[-1].event.runner_status, LoadedRunnerStatus)
# Ensure state is correct
assert isinstance(worker.state.runners[RUNNER_1_ID], LoadedRunnerStatus)
@@ -130,7 +145,7 @@ async def test_runner_assigned_active(
full_response = ''
- async for chunk in supervisor.stream_response(task=chat_completion_task(INSTANCE_1_ID)):
+ async for chunk in supervisor.stream_response(task=chat_completion_task(INSTANCE_1_ID, TASK_1_ID)):
if isinstance(chunk, TokenChunk):
full_response += chunk.text
@@ -194,9 +209,9 @@ async def test_runner_unassigns(
# Ensure the correct events have been emitted (creation)
events = await global_events.get_events_since(0)
- assert len(events) == 3
- assert isinstance(events[2].event, RunnerStatusUpdated)
- assert isinstance(events[2].event.runner_status, LoadedRunnerStatus)
+ assert len(events) >= 5
+ assert isinstance(events[-1].event, RunnerStatusUpdated)
+ assert isinstance(events[-1].event.runner_status, LoadedRunnerStatus)
# Ensure state is correct
print(worker.state)
@@ -223,14 +238,14 @@ async def test_runner_unassigns(
async def test_runner_inference(
worker_running: Callable[[NodeId], Awaitable[tuple[Worker, AsyncSQLiteEventStorage]]],
instance: Callable[[InstanceId, NodeId, RunnerId], Instance],
- chat_completion_task: Callable[[InstanceId], Task]
+ chat_completion_task: Callable[[InstanceId, TaskId], Task]
):
_worker, global_events = await worker_running(NODE_A)
instance_value: Instance = instance(INSTANCE_1_ID, NODE_A, RUNNER_1_ID)
instance_value.instance_type = InstanceStatus.ACTIVE
- task: Task = chat_completion_task(INSTANCE_1_ID)
+ task: Task = chat_completion_task(INSTANCE_1_ID, TASK_1_ID)
await global_events.append_events(
[
InstanceCreated(
@@ -265,7 +280,7 @@ async def test_2_runner_inference(
logger: Logger,
pipeline_shard_meta: Callable[[int, int], PipelineShardMetadata],
hosts: Callable[[int], list[Host]],
- chat_completion_task: Callable[[InstanceId], Task]
+ chat_completion_task: Callable[[InstanceId, TaskId], Task]
):
event_log_manager = EventLogManager(EventLogConfig(), logger)
await event_log_manager.initialize()
@@ -302,7 +317,7 @@ async def test_2_runner_inference(
hosts=hosts(2)
)
- task = chat_completion_task(INSTANCE_1_ID)
+ task = chat_completion_task(INSTANCE_1_ID, TASK_1_ID)
await global_events.append_events(
[
InstanceCreated(
@@ -345,7 +360,7 @@ async def test_runner_respawn(
logger: Logger,
pipeline_shard_meta: Callable[[int, int], PipelineShardMetadata],
hosts: Callable[[int], list[Host]],
- chat_completion_task: Callable[[InstanceId], Task]
+ chat_completion_task: Callable[[InstanceId, TaskId], Task]
):
event_log_manager = EventLogManager(EventLogConfig(), logger)
await event_log_manager.initialize()
@@ -382,7 +397,7 @@ async def test_runner_respawn(
hosts=hosts(2)
)
- task = chat_completion_task(INSTANCE_1_ID)
+ task = chat_completion_task(INSTANCE_1_ID, TASK_1_ID)
await global_events.append_events(
[
InstanceCreated(
@@ -442,7 +457,7 @@ async def test_runner_respawn(
assert isinstance(event, RunnerStatusUpdated)
assert isinstance(event.runner_status, LoadedRunnerStatus)
- task = chat_completion_task(INSTANCE_1_ID)
+ task = chat_completion_task(INSTANCE_1_ID, TASK_2_ID)
await global_events.append_events(
[
TaskCreated(
← e9b80360 Add Multiaddr type and refactor Hosts type for creating shar
·
back to Exo
·
fix forwarder supervisor tests c3c8ddbc →