Fix download tests
This commit is contained in:
+5
-2
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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'
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user