[object Object]

← back to Exo

Fix download tests

36a5d75efd5b41bc3a21b8ea92073268a74f5547 · 2025-07-28 12:51:10 +0100 · Matt Beton

Files touched

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 →