← back to Exo
Worker plan
a6b3ab63322d50c4bf325a6ec0e2975fef096706 · 2025-07-24 12:45:27 +0100 · Alex Cheema
Co-authored-by: Matt Beton <matthew.beton@gmail.com>
Co-authored-by: Seth Howes <71157822+sethhowes@users.noreply.github.com>
Co-authored-by: Gelu Vrabie <gelu.vrabie.univ@gmail.com>
Co-authored-by: Gelu Vrabie <gelu@exolabs.net>
Co-authored-by: Andrei Cravtov <the.andrei.cravtov@gmail.com>
Co-authored-by: Seth Howes <sethshowes@gmail.com>
Files touched
M .gitignoreM .idea/.gitignoreA shared/__init__.pyM shared/db/sqlite/event_log_manager.pyM shared/types/events/_common.pyM shared/types/events/_events.pyM shared/types/events/commands.pyA worker/NOTES.mdA worker/__init__.pyM worker/download/shard_downloader.pyM worker/main.pyM worker/runner/runner.pyA worker/tests/__init__.pyM worker/tests/conftest.pyM worker/tests/test_worker_handlers.pyA worker/tests/test_worker_integration.pyM worker/tests/test_worker_plan.pyA worker/tests/test_worker_plan_utils.pyM worker/tests/test_worker_state.py
Diff
commit a6b3ab63322d50c4bf325a6ec0e2975fef096706
Author: Alex Cheema <41707476+AlexCheema@users.noreply.github.com>
Date: Thu Jul 24 12:45:27 2025 +0100
Worker plan
Co-authored-by: Matt Beton <matthew.beton@gmail.com>
Co-authored-by: Seth Howes <71157822+sethhowes@users.noreply.github.com>
Co-authored-by: Gelu Vrabie <gelu.vrabie.univ@gmail.com>
Co-authored-by: Gelu Vrabie <gelu@exolabs.net>
Co-authored-by: Andrei Cravtov <the.andrei.cravtov@gmail.com>
Co-authored-by: Seth Howes <sethshowes@gmail.com>
---
.gitignore | 2 +-
.idea/.gitignore | 1 +
shared/__init__.py | 1 +
shared/db/sqlite/event_log_manager.py | 1 +
shared/types/events/_common.py | 5 +-
shared/types/events/_events.py | 7 +-
shared/types/events/commands.py | 4 +-
worker/NOTES.md | 2 +
worker/__init__.py | 1 +
worker/download/shard_downloader.py | 4 +-
worker/main.py | 243 +++++--
worker/runner/runner.py | 1 +
worker/tests/__init__.py | 1 +
worker/tests/conftest.py | 22 +-
worker/tests/test_worker_handlers.py | 38 +-
worker/tests/test_worker_integration.py | 57 ++
worker/tests/test_worker_plan.py | 1046 +++++++++++++++++++++++++------
worker/tests/test_worker_plan_utils.py | 212 +++++++
worker/tests/test_worker_state.py | 3 +-
19 files changed, 1392 insertions(+), 259 deletions(-)
diff --git a/.gitignore b/.gitignore
index 16f168d6..d0ef8f27 100644
--- a/.gitignore
+++ b/.gitignore
@@ -5,7 +5,7 @@ __pycache__
hosts_*.json
# hide direnv stuff
-/.direnv
+.direnv/
# TODO figure out how to properly solve the issue with these target directories showing up
networking/target/
networking/topology/target/
diff --git a/.idea/.gitignore b/.idea/.gitignore
index 13566b81..5ddb3d79 100644
--- a/.idea/.gitignore
+++ b/.idea/.gitignore
@@ -6,3 +6,4 @@
# Datasource local storage ignored files
/dataSources/
/dataSources.local.xml
+workspace.xml
\ No newline at end of file
diff --git a/shared/__init__.py b/shared/__init__.py
new file mode 100644
index 00000000..0519ecba
--- /dev/null
+++ b/shared/__init__.py
@@ -0,0 +1 @@
+
\ No newline at end of file
diff --git a/shared/db/sqlite/event_log_manager.py b/shared/db/sqlite/event_log_manager.py
index a20f3eca..266b24ff 100644
--- a/shared/db/sqlite/event_log_manager.py
+++ b/shared/db/sqlite/event_log_manager.py
@@ -24,6 +24,7 @@ class EventLogManager:
# Ensure base directory exists
EXO_HOME.mkdir(parents=True, exist_ok=True)
+ # TODO: This seems like it's a pattern to avoid an async __init__ function. But as we know, there's a better pattern for this - using a create() function, like in runner_supervisor.
async def initialize(self) -> None:
"""Initialize both connectors - call this during startup"""
# Both master and worker need both connectors
diff --git a/shared/types/events/_common.py b/shared/types/events/_common.py
index 53d2d4aa..a99af369 100644
--- a/shared/types/events/_common.py
+++ b/shared/types/events/_common.py
@@ -21,15 +21,14 @@ class EventId(NewUUID):
"""
+# Event base-class boilerplate (you should basically never touch these)
+# Only very specialised registry or serialisation/deserialization logic might need know about these
class CommandId(NewUUID):
"""
Newtype around `NewUUID` for command IDs
"""
-# Event base-class boilerplate (you should basically never touch these)
-# Only very specialised registry or serialisation/deserialization logic might need know about these
-
class _EventType(str, Enum):
"""
Here are all the unique kinds of events that can be sent over the network.
diff --git a/shared/types/events/_events.py b/shared/types/events/_events.py
index 06494877..9023567c 100644
--- a/shared/types/events/_events.py
+++ b/shared/types/events/_events.py
@@ -4,17 +4,14 @@ from pydantic import Field
from shared.topology import Connection, ConnectionProfile, Node, NodePerformanceProfile
from shared.types.common import NodeId
+from shared.types.events import CommandId
from shared.types.events.chunks import GenerationChunk
from shared.types.tasks import Task, TaskId, TaskStatus
from shared.types.worker.common import InstanceId, NodeStatus
from shared.types.worker.instances import InstanceParams, TypeOfInstance
from shared.types.worker.runners import RunnerId, RunnerStatus
-from ._common import (
- CommandId,
- _BaseEvent, # pyright: ignore[reportPrivateUsage]
- _EventType, # pyright: ignore[reportPrivateUsage]
-)
+from ._common import _BaseEvent, _EventType # pyright: ignore[reportPrivateUsage]
class TaskCreated(_BaseEvent[_EventType.TaskCreated]):
diff --git a/shared/types/events/commands.py b/shared/types/events/commands.py
index cce1b043..a4ec0e58 100644
--- a/shared/types/events/commands.py
+++ b/shared/types/events/commands.py
@@ -4,8 +4,7 @@ from typing import Annotated, Callable, Literal, Sequence
from pydantic import BaseModel, Field, TypeAdapter
from shared.types.api import ChatCompletionTaskParams
-from shared.types.events import Event
-from shared.types.events._common import CommandId
+from shared.types.events import CommandId, Event
from shared.types.state import InstanceId, State
@@ -35,6 +34,7 @@ class DeleteInstanceCommand(_BaseCommand[CommandTypes.DELETE_INSTANCE]):
command_type: Literal[CommandTypes.DELETE_INSTANCE] = CommandTypes.DELETE_INSTANCE
instance_id: InstanceId
+
Command = Annotated[
ChatCompletionCommand | CreateInstanceCommand | DeleteInstanceCommand,
Field(discriminator="command_type")
diff --git a/worker/NOTES.md b/worker/NOTES.md
new file mode 100644
index 00000000..1170d0b9
--- /dev/null
+++ b/worker/NOTES.md
@@ -0,0 +1,2 @@
+- Where should we check where the model is downloaded?
+- Error handling. How do we handle the scenario where an operation keeps failing to execute
diff --git a/worker/__init__.py b/worker/__init__.py
new file mode 100644
index 00000000..0519ecba
--- /dev/null
+++ b/worker/__init__.py
@@ -0,0 +1 @@
+
\ No newline at end of file
diff --git a/worker/download/shard_downloader.py b/worker/download/shard_downloader.py
index 0fbab318..68a095c7 100644
--- a/worker/download/shard_downloader.py
+++ b/worker/download/shard_downloader.py
@@ -54,7 +54,7 @@ class ShardDownloader(ABC):
device_rank=0,
world_size=1,
start_layer=0,
- end_layer=0,
+ end_layer=1,
n_layers=1,
),
completed_files=0,
@@ -93,7 +93,7 @@ class NoopShardDownloader(ShardDownloader):
device_rank=0,
world_size=1,
start_layer=0,
- end_layer=0,
+ end_layer=1,
n_layers=1,
),
completed_files=0,
diff --git a/worker/main.py b/worker/main.py
index b1274a70..2f0589e0 100644
--- a/worker/main.py
+++ b/worker/main.py
@@ -3,13 +3,23 @@ import os
from asyncio import Queue
from functools import partial
from logging import Logger
-from typing import AsyncGenerator, Optional
+from typing import AsyncGenerator, Callable, Optional
from pydantic import BaseModel, ConfigDict
+from shared.db.sqlite import AsyncSQLiteEventStorage
from shared.types.common import NodeId
-from shared.types.events import ChunkGenerated, Event, InstanceId, RunnerStatusUpdated
+from shared.types.events import (
+ ChunkGenerated,
+ Event,
+ InstanceCreated,
+ InstanceId,
+ RunnerStatusUpdated,
+ TaskStateUpdated,
+)
+from shared.types.events.components import EventFromEventLog
from shared.types.state import State
+from shared.types.tasks import TaskStatus
from shared.types.worker.common import RunnerId
from shared.types.worker.downloads import (
DownloadCompleted,
@@ -17,6 +27,7 @@ from shared.types.worker.downloads import (
DownloadOngoing,
DownloadProgressData,
)
+from shared.types.worker.instances import TypeOfInstance
from shared.types.worker.mlx import Host
from shared.types.worker.ops import (
AssignRunnerOp,
@@ -64,16 +75,35 @@ class AssignedRunner(BaseModel):
runner_status=self.status,
)
+# TODO: This should all be shared with the master.
+type ApplyFromEventLog = Callable[[State, EventFromEventLog[Event]], State]
+def get_apply_fn() -> ApplyFromEventLog:
+ # TODO: this needs to be done in a nice type-safe way
+ def _apply_instance_created(state: State, event_from_log: InstanceCreated) -> State:
+ return state
+
+ def apply_fn(state: State, event_from_log: EventFromEventLog[Event]) -> State:
+ if isinstance(event_from_log.event, InstanceCreated):
+ next_state = _apply_instance_created(state, event_from_log.event)
+ else:
+ raise ValueError(f"Unknown event type: {event_from_log.event}")
+ next_state.last_event_applied_idx = event_from_log.idx_in_log
+ return next_state
+
+ return apply_fn
+
class Worker:
def __init__(
self,
node_id: NodeId,
initial_state: State,
logger: Logger,
+ worker_events: AsyncSQLiteEventStorage | None,
):
- self.node_id = node_id
- self.state = initial_state
- self.logger = logger
+ self.node_id: NodeId = node_id
+ self.state: State = initial_state
+ self.worker_events: AsyncSQLiteEventStorage | None = worker_events
+ self.logger: Logger = logger
self.assigned_runners: dict[RunnerId, AssignedRunner] = {}
self._task: asyncio.Task[None] | None = None
@@ -82,15 +112,21 @@ class Worker:
@property
def _is_running(self) -> bool:
return self._task is not None and not self._task.done()
+
+ @property
+ def exception(self) -> Exception | None:
+ if self._task is not None:
+ self._task.exception()
+ # We don't start immediately on init - for testing purposes it is useful to have an 'inactive' worker.
async def start(self):
self._task = asyncio.create_task(self._loop())
async def stop(self):
if not self._is_running:
raise RuntimeError("Worker is not running")
-
- assert self._task is not None
+
+ assert self._task is not None
self._task.cancel()
@@ -118,13 +154,13 @@ class Worker:
self, op: UnassignRunnerOp
) -> AsyncGenerator[Event, None]:
if op.runner_id not in self.assigned_runners:
- return
+ return
# We can try to do a graceful shutdown of the runner.
- runner: RunnerSupervisor | None = self.assigned_runners[op.runner_id].runner
+ runner: RunnerSupervisor | None = self.assigned_runners[op.runner_id].runner
if runner is not None:
await runner.astop()
-
+
# This is all we really need:
del self.assigned_runners[op.runner_id]
@@ -174,7 +210,7 @@ class Worker:
downloaded_bytes=0
)
)
- )
+ )
self.assigned_runners[op.runner_id] = AssignedRunner(
runner_id=op.runner_id,
@@ -188,7 +224,7 @@ class Worker:
yield assigned_runner.status_update_event()
# Download it!
- # TODO: we probably want download progress as part of a callback that gets passed to the downloader.
+ # TODO: we probably want download progress as part of a callback that gets passed to the downloader.
try:
assert assigned_runner.is_downloaded
@@ -209,22 +245,19 @@ class Worker:
assigned_runner.status = ReadyRunnerStatus()
yield assigned_runner.status_update_event()
-# Plan:
-# First get a single inference running
-# Then build boilerplate for passing callback when mlx is in the 'ready' state
-# Then figure out if we can do what's needed with events. But this is a little challenging because it depends on Alex's code.
- async def _execute_chat_completion_op(
+
+ async def _execute_task_op(
self, op: ExecuteTaskOp
) -> AsyncGenerator[Event, None]:
'''
- This is the entry point for a chat completion starting.
+ This is the entry point for a chat completion starting.
While there is only one execute function, it will get called in different ways for runner 0 and runner [1, 2, 3, ...].
Runners [1, 2, 3, ...] will run this method when a task is in 'pending' state.
Runner 0 will run this method when a task is in 'running' state.
TODO: How do we handle the logic of ensuring that n-1 nodes have started their execution before allowing the 0'th runner to start?
This is still a little unclear to me.
'''
- assigned_runner = self.assigned_runners[op.runner_id]
+ assigned_runner = self.assigned_runners[op.runner_id]
async def inner_execute(queue: asyncio.Queue[Event]) -> None:
assert assigned_runner.runner is not None
@@ -234,27 +267,46 @@ class Worker:
# Called when the MLX process has been kicked off
assigned_runner.status = RunningRunnerStatus()
await queue.put(assigned_runner.status_update_event())
-
-
+
+ if assigned_runner.shard_metadata.device_rank == 0:
+ await queue.put(TaskStateUpdated(
+ task_id=op.task.task_id,
+ task_status=TaskStatus.RUNNING,
+ ))
+
try:
async for chunk in assigned_runner.runner.stream_response(
- task=op.task,
+ task=op.task,
request_started_callback=partial(running_callback, queue)):
- await queue.put(ChunkGenerated(
- # todo: at some point we will no longer have a bijection between task_id and row_id.
- # So we probably want to store a mapping between these two in our Worker object.
- command_id=chunk.command_id,
- chunk=chunk
+ if assigned_runner.shard_metadata.device_rank == 0:
+ await queue.put(ChunkGenerated(
+ # todo: at some point we will no longer have a bijection between task_id and row_id.
+ # So we probably want to store a mapping between these two in our Worker object.
+ command_id=chunk.command_id,
+ chunk=chunk
+ ))
+
+ if assigned_runner.shard_metadata.device_rank == 0:
+ await queue.put(TaskStateUpdated(
+ task_id=op.task.task_id,
+ task_status=TaskStatus.COMPLETE,
))
-
+
# After a successful inference:
assigned_runner.status = LoadedRunnerStatus()
await queue.put(assigned_runner.status_update_event())
+
except Exception as e:
# TODO: What log level?
self.logger.log(2, f'Runner failed whilst running inference task. Task: {op.task}. Error: {e}')
+ if assigned_runner.shard_metadata.device_rank == 0:
+ await queue.put(TaskStateUpdated(
+ task_id=op.task.task_id,
+ task_status=TaskStatus.FAILED,
+ ))
+
assigned_runner.runner = None
assigned_runner.status = FailedRunnerStatus(error_message=str(e))
await queue.put(assigned_runner.status_update_event())
@@ -292,7 +344,7 @@ class Worker:
case RunnerOpType.DOWNLOAD:
event_generator = self._execute_download_op(op)
case RunnerOpType.CHAT_COMPLETION:
- event_generator = self._execute_chat_completion_op(op)
+ event_generator = self._execute_task_op(op)
async for event in event_generator:
yield event
@@ -300,10 +352,67 @@ class Worker:
## Planning logic
def plan(self, state: State) -> RunnerOp | None:
# Compare state to worker 'mood'
-
- # First spin things down
-
- # Then spin things up
+
+ # First, unassign assigned runners that are no longer in the state.
+ for runner_id, _ in self.assigned_runners.items():
+ if runner_id not in state.runners:
+ return UnassignRunnerOp(runner_id=runner_id)
+
+ # Then spin down active runners
+ for _instance_id, instance in state.instances.items():
+ for node_id, runner_id in instance.instance_params.shard_assignments.node_to_runner.items():
+ if node_id != self.node_id:
+ continue
+
+ # We spin down a runner if it's meant to be inactive and it's Loaded.
+ if runner_id in self.assigned_runners and \
+ isinstance(self.assigned_runners[runner_id].status, LoadedRunnerStatus) and \
+ instance.instance_type == TypeOfInstance.INACTIVE:
+ return RunnerDownOp(runner_id=runner_id)
+
+ # If we are part of an instance that has a dead node - and we aren't the dead node - we should spin down
+ # TODO: We need to limit number of retries if we keep failing.
+ for _instance_id, instance in state.instances.items():
+ if self.node_id in instance.instance_params.shard_assignments.node_to_runner:
+ other_node_in_instance_has_failed = False
+ for runner_id in instance.instance_params.shard_assignments.runner_to_shard:
+ if isinstance(state.runners[runner_id], FailedRunnerStatus) and \
+ runner_id not in self.assigned_runners:
+ other_node_in_instance_has_failed= True
+
+ if other_node_in_instance_has_failed:
+ # Spin down *our* runner
+ return RunnerDownOp(runner_id=instance.instance_params.shard_assignments.node_to_runner[self.node_id])
+
+ # If we are failed - and *all of the other nodes have spun down* - then we can spin down too.
+ for _instance_id, instance in state.instances.items():
+ if self.node_id in instance.instance_params.shard_assignments.node_to_runner and \
+ isinstance(state.runners[instance.instance_params.shard_assignments.node_to_runner[self.node_id]], FailedRunnerStatus):
+
+ num_spundown_nodes = 0
+ for runner_id in instance.instance_params.shard_assignments.runner_to_shard:
+ if isinstance(state.runners[runner_id], ReadyRunnerStatus) and \
+ runner_id not in self.assigned_runners:
+ num_spundown_nodes += 1
+
+ if num_spundown_nodes == next(iter(instance.instance_params.shard_assignments.runner_to_shard.values())).world_size - 1:
+ # All the other nodes are spun down - so now we can spin down too.
+ # This also catches the case of 1-node. If there's one node in the instance then we should spin down straight away
+ return RunnerDownOp(runner_id=instance.instance_params.shard_assignments.node_to_runner[self.node_id])
+
+ # Then assign runners we do want
+ for instance_id, instance in state.instances.items():
+ for node_id, runner_id in instance.instance_params.shard_assignments.node_to_runner.items():
+ if node_id != self.node_id:
+ continue
+
+ if runner_id not in self.assigned_runners:
+ return AssignRunnerOp(
+ runner_id=runner_id,
+ instance_id=instance_id,
+ shard_metadata=instance.instance_params.shard_assignments.runner_to_shard[runner_id],
+ hosts=instance.instance_params.hosts
+ )
# Then make sure things are downloading.
for instance_id, instance in state.instances.items():
@@ -327,24 +436,80 @@ class Worker:
hosts=instance.instance_params.hosts
)
+ # Then spin up 'ready' runners that should be active
+ for _instance_id, instance in state.instances.items():
+ if self.node_id in instance.instance_params.shard_assignments.node_to_runner and \
+ self.assigned_runners[instance.instance_params.shard_assignments.node_to_runner[self.node_id]].runner is None and \
+ instance.instance_type == TypeOfInstance.ACTIVE:
+
+ # We are part of this instance, we want it up but it hasn't been spun up yet.
+ # Need to assert all other runners are ready before we can spin up.
+ ready_to_spin = True
+ for runner_id in instance.instance_params.shard_assignments.node_to_runner.values():
+ if state.runners[runner_id].runner_status != RunnerStatusType.Ready:
+ ready_to_spin = False
+ if ready_to_spin:
+ return RunnerUpOp(runner_id=instance.instance_params.shard_assignments.node_to_runner[self.node_id])
+ # Then make sure things are running based on tasks.
+ for instance_id, instance in state.instances.items():
+ for node_id, runner_id in instance.instance_params.shard_assignments.node_to_runner.items():
+ if node_id != self.node_id:
+ continue
+ assert runner_id in self.assigned_runners
+ runner = self.assigned_runners[runner_id]
+ if runner.status.runner_status != RunnerStatusType.Loaded:
+ continue # The only previous state to get to Running is from Loaded
+
+ for _, task in state.tasks.items():
+ if task.instance_id == instance_id:
+ if (runner.shard_metadata.device_rank >= 1 or runner.shard_metadata.world_size == 1):
+ return ExecuteTaskOp(runner_id=runner_id, task=task)
+ else:
+ # We already know our own status is Loaded. We are rank 0,
+ # so let's check that all the other runners are running - ready for us to fire the prompt.
+ running_runner_count = 0
+ for other_runner_id, other_runner_status in state.runners.items():
+ if other_runner_id in instance.instance_params.shard_assignments.node_to_runner.values() and \
+ isinstance(other_runner_status, RunningRunnerStatus):
+ running_runner_count += 1
+
+ if running_runner_count == runner.shard_metadata.world_size - 1:
+ return ExecuteTaskOp(runner_id=runner_id, task=task)
- # Finally, chat completion.
return None
+ async def event_publisher(self, event: Event) -> None:
+ assert self.worker_events is not None
+ await self.worker_events.append_events([event], self.node_id)
+
# Handle state updates
async def _loop(self):
+ assert self.worker_events is not None
+ self.apply_fn = get_apply_fn()
+
while True:
- state_copy = self.state.model_copy(deep=False)
- op: RunnerOp | None = self.plan(state_copy)
+ # ToDo: Where do we update state? Do we initialize it from scratch & read all events in, or do we preload the state?
+
+ # 1. get latest events
+ events = await self.worker_events.get_events_since(self.state.last_event_applied_idx)
+ if len(events) == 0:
+ await asyncio.sleep(0.01)
+ continue
+
+ # 2. for each event, apply it to the state and run sagas
+ for event_from_log in events:
+ self.state = self.apply_fn(self.state, event_from_log)
+
+ # 3. based on the updated state, we plan & execute an operation.
+ op: RunnerOp | None = self.plan(self.state)
# run the op, synchronously blocking for now
if op is not None:
async for event in self._execute_op(op):
- print(event)
- # self.event_publisher(event)
+ await self.event_publisher(event)
await asyncio.sleep(0.01)
@@ -352,7 +517,7 @@ class Worker:
# TODO: Handle resource monitoring (write-only)
async def main():
-
+
print("Hello from worker!")
diff --git a/worker/runner/runner.py b/worker/runner/runner.py
index eebb9a5b..99d6a2e5 100644
--- a/worker/runner/runner.py
+++ b/worker/runner/runner.py
@@ -121,6 +121,7 @@ async def main():
case ChatTaskMessage(task_data=task):
runner_print(f"received chat request: {task}")
# Ensure we have a chat-completion task subtype
+ # TODO: this is a hack, why are we only looking at the first message? should have a tokenizer
prompt = task.messages[0]
if prompt.content is not None and 'EXO RUNNER MUST FAIL' in prompt.content:
raise Exception('Artificial runner exception - for testing purposes only.')
diff --git a/worker/tests/__init__.py b/worker/tests/__init__.py
new file mode 100644
index 00000000..0519ecba
--- /dev/null
+++ b/worker/tests/__init__.py
@@ -0,0 +1 @@
+
\ No newline at end of file
diff --git a/worker/tests/conftest.py b/worker/tests/conftest.py
index 7e4f003c..0182e9c2 100644
--- a/worker/tests/conftest.py
+++ b/worker/tests/conftest.py
@@ -1,4 +1,3 @@
-import asyncio
import uuid
from logging import Logger, getLogger
from pathlib import Path
@@ -6,6 +5,7 @@ from typing import Callable
import pytest
+from shared.db.sqlite.event_log_manager import EventLogConfig, EventLogManager
from shared.types.api import ChatCompletionMessage, ChatCompletionTaskParams
from shared.types.common import NodeId
from shared.types.models import ModelId, ModelMetadata
@@ -115,9 +115,14 @@ def chat_completion_task(completion_create_params: ChatCompletionTaskParams) ->
)
@pytest.fixture
-def state():
+def node_id() -> NodeId:
+ """Shared node ID for tests"""
+ return NodeId(uuid.uuid4())
+
+@pytest.fixture
+def state(node_id: NodeId):
node_status={
- NodeId(uuid.uuid4()): NodeStatus.Idle
+ node_id: NodeStatus.Idle
}
return State(
@@ -155,14 +160,15 @@ def instance(pipeline_shard_meta: Callable[[int, int], PipelineShardMetadata], h
return _instance
@pytest.fixture
-def worker(state: State, logger: Logger):
- return Worker(NodeId(uuid.uuid4()), state, logger)
+async def worker(node_id: NodeId, state: State, logger: Logger):
+ event_log_manager = EventLogManager(EventLogConfig(), logger)
+ await event_log_manager.initialize()
+
+ return Worker(node_id, state, logger, worker_events=event_log_manager.global_events)
@pytest.fixture
async def worker_with_assigned_runner(worker: Worker, instance: Callable[[NodeId], Instance]):
"""Fixture that provides a worker with an already assigned runner."""
- await worker.start()
- await asyncio.sleep(0.01)
instance_obj: Instance = instance(worker.node_id)
@@ -196,4 +202,4 @@ async def worker_with_running_runner(worker_with_assigned_runner: tuple[Worker,
assert supervisor is not None
assert supervisor.healthy
- return worker, runner_id, instance_obj
\ No newline at end of file
+ return worker, runner_id, instance_obj
diff --git a/worker/tests/test_worker_handlers.py b/worker/tests/test_worker_handlers.py
index 20823c5e..02f77234 100644
--- a/worker/tests/test_worker_handlers.py
+++ b/worker/tests/test_worker_handlers.py
@@ -1,15 +1,19 @@
## Tests for worker state handlers
-import asyncio
from pathlib import Path
from typing import Callable
import pytest
from shared.types.common import NodeId
-from shared.types.events import ChunkGenerated, Event, RunnerStatusUpdated
+from shared.types.events import (
+ ChunkGenerated,
+ Event,
+ RunnerStatusUpdated,
+ TaskStateUpdated,
+)
from shared.types.events.chunks import TokenChunk
-from shared.types.tasks import Task
+from shared.types.tasks import Task, TaskStatus
from shared.types.worker.common import RunnerId
from shared.types.worker.instances import Instance
from shared.types.worker.ops import (
@@ -36,9 +40,6 @@ def user_message():
@pytest.mark.asyncio
async def test_assign_op(worker: Worker, instance: Callable[[NodeId], Instance], tmp_path: Path):
- await worker.start()
- await asyncio.sleep(0.01)
-
instance_obj: Instance = instance(worker.node_id)
runner_id: RunnerId | None = None
for x in instance_obj.instance_params.shard_assignments.runner_to_shard:
@@ -167,15 +168,24 @@ async def test_execute_task_op(
assert len(events) > 20
+ print(f'{events=}')
+
+
assert isinstance(events[0], RunnerStatusUpdated)
assert isinstance(events[0].runner_status, RunningRunnerStatus)
+ assert isinstance(events[1], TaskStateUpdated)
+ assert events[1].task_status == TaskStatus.RUNNING # It tried to start.
+
+ assert isinstance(events[-2], TaskStateUpdated)
+ assert events[-2].task_status == TaskStatus.COMPLETE # It tried to start.
+
assert isinstance(events[-1], RunnerStatusUpdated)
assert isinstance(events[-1].runner_status, LoadedRunnerStatus) # It should not have failed.
gen_events: list[ChunkGenerated] = [x for x in events if isinstance(x, ChunkGenerated)]
text_chunks: list[TokenChunk] = [x.chunk for x in gen_events if isinstance(x.chunk, TokenChunk)]
- assert len(text_chunks) == len(events) - 2
+ assert len(text_chunks) == len(events) - 4
output_text = ''.join([x.text for x in text_chunks])
assert '42' in output_text
@@ -202,10 +212,18 @@ async def test_execute_task_fails(
async for event in worker._execute_op(execute_task_op): # type: ignore[misc]
events.append(event)
- assert len(events) == 2
+ assert len(events) == 4
+
+ print(events)
assert isinstance(events[0], RunnerStatusUpdated)
assert isinstance(events[0].runner_status, RunningRunnerStatus) # It tried to start.
- assert isinstance(events[-1], RunnerStatusUpdated)
- assert isinstance(events[-1].runner_status, FailedRunnerStatus) # It should have failed.
\ No newline at end of file
+ assert isinstance(events[1], TaskStateUpdated)
+ assert events[1].task_status == TaskStatus.RUNNING # It tried to start.
+
+ assert isinstance(events[2], TaskStateUpdated)
+ assert events[2].task_status == TaskStatus.FAILED # Task marked as failed.
+
+ assert isinstance(events[3], RunnerStatusUpdated)
+ assert isinstance(events[3].runner_status, FailedRunnerStatus) # It should have failed.
\ No newline at end of file
diff --git a/worker/tests/test_worker_integration.py b/worker/tests/test_worker_integration.py
new file mode 100644
index 00000000..7e8e5a99
--- /dev/null
+++ b/worker/tests/test_worker_integration.py
@@ -0,0 +1,57 @@
+import asyncio
+from logging import Logger
+from typing import Callable, Final
+from uuid import UUID
+
+from shared.db.sqlite.event_log_manager import EventLogConfig, EventLogManager
+from shared.types.common import NodeId
+from shared.types.events import InstanceCreated
+from shared.types.models import ModelId
+from shared.types.state import State
+from shared.types.tasks import TaskId
+from shared.types.worker.common import InstanceId, RunnerId
+from shared.types.worker.instances import Instance
+from worker.main import Worker
+
+MASTER_NODE_ID = NodeId(uuid=UUID("ffffffff-aaaa-4aaa-8aaa-aaaaaaaaaaaa"))
+NODE_A: Final[NodeId] = NodeId(uuid=UUID("aaaaaaaa-aaaa-4aaa-8aaa-aaaaaaaaaaaa"))
+NODE_B: Final[NodeId] = NodeId(uuid=UUID("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()
+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()
+
+async def test_runner_spin_up(instance: Callable[[NodeId], Instance]):
+ # TODO.
+ return
+ node_id = NodeId()
+ logger = Logger('worker_test_logger')
+ event_log_manager = EventLogManager(EventLogConfig(), logger)
+ await event_log_manager.initialize()
+
+ global_events = event_log_manager.global_events
+
+ worker = Worker(node_id, State(), logger=logger, worker_events=global_events)
+ await worker.start()
+
+ instance_value = instance(node_id)
+
+ await global_events.append_events(
+ [
+ InstanceCreated(
+ instance_id=instance_value.instance_id,
+ instance_params=instance_value.instance_params,
+ instance_type=instance_value.instance_type
+ )
+ ],
+ origin=MASTER_NODE_ID
+ )
+
+ await asyncio.sleep(0.1)
+
+ assert worker.assigned_runners
\ No newline at end of file
diff --git a/worker/tests/test_worker_plan.py b/worker/tests/test_worker_plan.py
index 953b0fab..f27c5652 100644
--- a/worker/tests/test_worker_plan.py
+++ b/worker/tests/test_worker_plan.py
@@ -1,186 +1,836 @@
from __future__ import annotations
import logging
-from dataclasses import dataclass
+import tempfile
from pathlib import Path
-from typing import Callable, Final, List, Optional, Type
import pytest
-from shared.types.common import NodeId
-from shared.types.models import ModelId
+from shared.types.api import ChatCompletionMessage
from shared.types.state import State
-
-# WorkerState import below after RunnerCase definition to avoid forward reference issues
-from shared.types.worker.common import InstanceId, NodeStatus, RunnerId
-from shared.types.worker.downloads import DownloadOngoing, DownloadProgressData
+from shared.types.tasks import (
+ ChatCompletionTask,
+ ChatCompletionTaskParams,
+ TaskStatus,
+ TaskType,
+)
+from shared.types.worker.common import NodeStatus
+from shared.types.worker.downloads import DownloadPending
from shared.types.worker.instances import Instance, InstanceParams, TypeOfInstance
-from shared.types.worker.ops import DownloadOp
+from shared.types.worker.ops import (
+ AssignRunnerOp,
+ DownloadOp,
+ ExecuteTaskOp,
+ RunnerDownOp,
+ RunnerUpOp,
+ UnassignRunnerOp,
+)
from shared.types.worker.runners import (
+ AssignedRunnerStatus,
DownloadingRunnerStatus,
+ FailedRunnerStatus,
+ LoadedRunnerStatus,
ReadyRunnerStatus,
- RunnerStatus,
+ RunningRunnerStatus,
ShardAssignments,
)
from shared.types.worker.shards import PipelineShardMetadata
from worker.download.download_utils import build_model_path
-from worker.main import AssignedRunner, Worker
+from worker.main import Worker
+
+from .test_worker_plan_utils import (
+ INSTANCE_1_ID,
+ MODEL_A_ID,
+ NODE_A,
+ NODE_B,
+ RUNNER_1_ID,
+ RUNNER_2_ID,
+ TASK_1_ID,
+ InProcessRunner,
+ OverrideAssignedRunner,
+ PlanTestCase,
+ make_downloading_status,
+ make_model_meta,
+ make_shard_metadata,
+)
+"""
+The idea with these tests is to define declaratively the input and expected output of the worker.plan function.
+
+We initialize a Worker with InProcessRunners. We then construct a State which gets passed to Worker.plan.
+We then check what operation is returned by Worker.plan.
+"""
+
+def _get_test_cases(tmp_path: Path) -> list[PlanTestCase]:
+ # The `model_path` for `RUNNER_1_ID` must exist for the `DownloadOp` test case to pass validation.
+ (tmp_path / f"model_for_runner_{RUNNER_1_ID}").mkdir(exist_ok=True, parents=True)
+ model_a_meta = make_model_meta(MODEL_A_ID)
+ return [
+ PlanTestCase(
+ description="no runners -> no-op",
+ in_process_runners=[],
+ state=State(node_status={NODE_A: NodeStatus.Idle}, instances={}, runners={}),
+ expected_op=None,
+ ),
-@dataclass(slots=True, frozen=True)
-class RunnerCase:
- """Important, minimal state for a *single* runner relevant to planning."""
+ # I don't think this should ever happen, as if it's currently downloading then the worker loop will be blocked
+ # Potentially useful for future compatibility when worker becomes non-blocking
+ PlanTestCase(
+ description="runner state assigned, runner is assigned and downloading -> no-op",
+ in_process_runners=[
+ InProcessRunner(
+ runner_id=RUNNER_1_ID,
+ instance_id=INSTANCE_1_ID,
+ model_id=MODEL_A_ID,
+ status=make_downloading_status(NODE_A),
+ downloaded=False,
+ )
+ ],
+ state=State(
+ node_status={NODE_A: NodeStatus.Idle},
+ instances={},
+ runners={RUNNER_1_ID: make_downloading_status(NODE_A)},
+ ),
+ expected_op=None,
+ ),
- status: RunnerStatus
- downloaded: bool # Does the model shard already exist on disk?
+ PlanTestCase(
+ description="runner state downloading, runner is downloading -> no-op",
+ in_process_runners=[
+ InProcessRunner(
+ runner_id=RUNNER_1_ID,
+ instance_id=INSTANCE_1_ID,
+ model_id=MODEL_A_ID,
+ status=make_downloading_status(NODE_A),
+ downloaded=False,
+ )
+ ],
+ state=State(
+ node_status={NODE_A: NodeStatus.Idle},
+ instances={
+ INSTANCE_1_ID: Instance(
+ instance_type=TypeOfInstance.INACTIVE,
+ instance_id=INSTANCE_1_ID,
+ instance_params=InstanceParams(
+ shard_assignments=ShardAssignments(
+ model_id=MODEL_A_ID,
+ runner_to_shard={
+ RUNNER_1_ID: make_shard_metadata(device_rank=0, world_size=1)
+ },
+ node_to_runner={NODE_A: RUNNER_1_ID}
+ ),
+ hosts=[]
+ ),
+ )
+ },
+ runners={RUNNER_1_ID: make_downloading_status(NODE_A)},
+ ),
+ expected_op=None,
+ ),
+ PlanTestCase(
+ description="ready runner, model present -> no-op",
+ in_process_runners=[
+ InProcessRunner(
+ runner_id=RUNNER_1_ID,
+ instance_id=INSTANCE_1_ID,
+ model_id=MODEL_A_ID,
+ status=ReadyRunnerStatus(),
+ downloaded=True,
+ )
+ ],
+ state=State(
+ node_status={NODE_A: NodeStatus.Idle},
+ instances={
+ INSTANCE_1_ID: Instance(
+ instance_type=TypeOfInstance.INACTIVE,
+ instance_id=INSTANCE_1_ID,
+ instance_params=InstanceParams(
+ shard_assignments=ShardAssignments(
+ model_id=MODEL_A_ID,
+ runner_to_shard={
+ RUNNER_1_ID: make_shard_metadata(device_rank=0, world_size=1)
+ },
+ node_to_runner={NODE_A: RUNNER_1_ID}
+ ),
+ hosts=[]
+ ),
+ )
+ },
+ runners={RUNNER_1_ID: ReadyRunnerStatus()},
+ ),
+ expected_op=None,
+ ),
-@dataclass(slots=True, frozen=True)
-class PlanTestCase:
- """Table-driven description of an entire planning scenario."""
+ PlanTestCase(
+ description="runner assigned and not in state -> AssignRunnerOp",
+ in_process_runners=[],
+ state=State(
+ node_status={NODE_A: NodeStatus.Idle},
+ instances={
+ INSTANCE_1_ID: Instance(
+ instance_type=TypeOfInstance.ACTIVE, # Either active or inactive should yield the same.
+ instance_id=INSTANCE_1_ID,
+ instance_params=InstanceParams(
+ shard_assignments=ShardAssignments(
+ model_id=MODEL_A_ID,
+ runner_to_shard={
+ RUNNER_1_ID: make_shard_metadata(device_rank=0, world_size=1)
+ },
+ node_to_runner={NODE_A: RUNNER_1_ID}
+ ),
+ hosts=[]
+ ),
+ )
+ },
+ runners={RUNNER_1_ID: AssignedRunnerStatus()},
+ ),
+ expected_op=AssignRunnerOp(
+ instance_id=INSTANCE_1_ID,
+ runner_id=RUNNER_1_ID,
+ shard_metadata=PipelineShardMetadata(
+ device_rank=0,
+ world_size=1,
+ model_meta=model_a_meta,
+ start_layer=0,
+ end_layer=1,
+ n_layers=1,
+ ),
+ hosts=[]
+ ),
+ ),
- description: str
- runners: List[RunnerCase]
- # If we expect an op, specify the precise type and the index of the runner it targets.
- expected_op_type: Optional[Type[DownloadOp]] # Currently only DownloadOp handled.
- expected_op_runner_idx: Optional[int] = None
- # Allow overriding the WorkerState passed to Worker.plan. When None, a default state
- # is constructed from `runners` via helper `_build_worker_state`.
- worker_state_override: Optional[State] = None
+ PlanTestCase(
+ description="runner assigned but no longer in state -> UnassignRunnerOp",
+ in_process_runners=[
+ InProcessRunner(
+ runner_id=RUNNER_1_ID,
+ instance_id=INSTANCE_1_ID,
+ model_id=MODEL_A_ID,
+ status=AssignedRunnerStatus(),
+ downloaded=False,
+ )
+ ],
+ state=State(node_status={NODE_A: NodeStatus.Idle}, instances={}, runners={}),
+ expected_op=UnassignRunnerOp(runner_id=RUNNER_1_ID),
+ ),
- def id(self) -> str: # noqa: D401
- return self.description.replace(" ", "_")
+ PlanTestCase(
+ description="runner state assigned, runner is assigned, not downloaded -> expect DownloadOp",
+ in_process_runners=[
+ InProcessRunner(
+ runner_id=RUNNER_1_ID,
+ instance_id=INSTANCE_1_ID,
+ model_id=MODEL_A_ID,
+ status=AssignedRunnerStatus(),
+ downloaded=False,
+ )
+ ],
+ state=State(
+ node_status={NODE_A: NodeStatus.Idle},
+ instances={
+ INSTANCE_1_ID: Instance(
+ instance_type=TypeOfInstance.ACTIVE,
+ instance_id=INSTANCE_1_ID,
+ instance_params=InstanceParams(
+ shard_assignments=ShardAssignments(
+ model_id=MODEL_A_ID,
+ runner_to_shard={
+ RUNNER_1_ID: make_shard_metadata(device_rank=0, world_size=1)
+ },
+ node_to_runner={NODE_A: RUNNER_1_ID}
+ ),
+ hosts=[]
+ ),
+ )
+ },
+ runners={RUNNER_1_ID: AssignedRunnerStatus()},
+ ),
+ expected_op=DownloadOp(
+ runner_id=RUNNER_1_ID,
+ instance_id=INSTANCE_1_ID,
+ shard_metadata=PipelineShardMetadata(
+ device_rank=0,
+ world_size=1,
+ model_meta=model_a_meta,
+ start_layer=0,
+ end_layer=1,
+ n_layers=1,
+ ),
+ hosts=[],
+ ),
+ ),
+ PlanTestCase(
+ description="ready runner (and state up) -> expect RunnerUpOp",
+ in_process_runners=[
+ InProcessRunner(
+ runner_id=RUNNER_1_ID,
+ instance_id=INSTANCE_1_ID,
+ model_id=MODEL_A_ID,
+ status=ReadyRunnerStatus(),
+ downloaded=True,
+ )
+ ],
+ state=State(
+ node_status={NODE_A: NodeStatus.Idle},
+ instances={
+ INSTANCE_1_ID: Instance(
+ instance_type=TypeOfInstance.ACTIVE,
+ instance_id=INSTANCE_1_ID,
+ instance_params=InstanceParams(
+ shard_assignments=ShardAssignments(
+ model_id=MODEL_A_ID,
+ runner_to_shard={
+ RUNNER_1_ID: make_shard_metadata(device_rank=0, world_size=1)
+ },
+ node_to_runner={NODE_A: RUNNER_1_ID}
+ ),
+ hosts=[]
+ ),
+ )
+ },
+ runners={RUNNER_1_ID: ReadyRunnerStatus()},
+ tasks={},
+ ),
+ expected_op=RunnerUpOp(runner_id=RUNNER_1_ID),
+ ),
-def _make_downloading_status(node_id: NodeId) -> DownloadingRunnerStatus:
- """Factory for a *Downloading* status with placeholder progress."""
- return DownloadingRunnerStatus(
- download_progress=DownloadOngoing(
- node_id=node_id,
- download_progress=DownloadProgressData(total_bytes=1, downloaded_bytes=0),
- )
- )
+ PlanTestCase(
+ description="1 ready, 1 downloading (and state up) -> no-op",
+ in_process_runners=[
+ InProcessRunner(
+ runner_id=RUNNER_1_ID,
+ instance_id=INSTANCE_1_ID,
+ model_id=MODEL_A_ID,
+ status=ReadyRunnerStatus(),
+ downloaded=True,
+ device_rank=0,
+ ),
+ InProcessRunner(
+ runner_id=RUNNER_2_ID,
+ instance_id=INSTANCE_1_ID,
+ model_id=MODEL_A_ID,
+ status=DownloadingRunnerStatus(
+ download_progress=DownloadPending(node_id=NODE_A)
+ ),
+ downloaded=False,
+ device_rank=1,
+ ),
+ ],
+ state=State(
+ node_status={NODE_A: NodeStatus.Idle, NODE_B: NodeStatus.Idle},
+ instances={
+ INSTANCE_1_ID: Instance(
+ instance_type=TypeOfInstance.ACTIVE,
+ instance_id=INSTANCE_1_ID,
+ instance_params=InstanceParams(
+ shard_assignments=ShardAssignments(
+ model_id=MODEL_A_ID,
+ runner_to_shard={
+ RUNNER_1_ID: make_shard_metadata(device_rank=0, world_size=2),
+ RUNNER_2_ID: make_shard_metadata(device_rank=1, world_size=2)
+ },
+ node_to_runner={NODE_A: RUNNER_1_ID, NODE_B: RUNNER_2_ID}
+ ),
+ hosts=[]
+ ),
+ )
+ },
+ runners={RUNNER_1_ID: ReadyRunnerStatus(), RUNNER_2_ID: DownloadingRunnerStatus(download_progress=DownloadPending(node_id=NODE_A))},
+ tasks={TASK_1_ID: ChatCompletionTask(task_id=TASK_1_ID, task_type=TaskType.CHAT_COMPLETION, task_status=TaskStatus.PENDING, task_params=ChatCompletionTaskParams(model=str(MODEL_A_ID), messages=[ChatCompletionMessage(role="user", content="Hello, world!")]), instance_id=INSTANCE_1_ID)},
+ ),
+ expected_op=None
+ ),
+ PlanTestCase(
+ description="2 ready runners (and state up) -> expect RunnerUpOp",
+ in_process_runners=[
+ InProcessRunner(
+ runner_id=RUNNER_1_ID,
+ instance_id=INSTANCE_1_ID,
+ model_id=MODEL_A_ID,
+ status=ReadyRunnerStatus(),
+ downloaded=True,
+ device_rank=0,
+ ),
+ InProcessRunner(
+ runner_id=RUNNER_2_ID,
+ instance_id=INSTANCE_1_ID,
+ model_id=MODEL_A_ID,
+ status=ReadyRunnerStatus(),
+ downloaded=True,
+ device_rank=1,
+ ),
+ ],
+ state=State(
+ node_status={NODE_A: NodeStatus.Idle, NODE_B: NodeStatus.Idle},
+ instances={
+ INSTANCE_1_ID: Instance(
+ instance_type=TypeOfInstance.ACTIVE,
+ instance_id=INSTANCE_1_ID,
+ instance_params=InstanceParams(
+ shard_assignments=ShardAssignments(
+ model_id=MODEL_A_ID,
+ runner_to_shard={
+ RUNNER_1_ID: make_shard_metadata(device_rank=0, world_size=2),
+ RUNNER_2_ID: make_shard_metadata(device_rank=1, world_size=2)
+ },
+ node_to_runner={NODE_A: RUNNER_1_ID, NODE_B: RUNNER_2_ID}
+ ),
+ hosts=[]
+ ),
+ )
+ },
+ runners={RUNNER_1_ID: ReadyRunnerStatus(), RUNNER_2_ID: ReadyRunnerStatus()},
+ tasks={TASK_1_ID: ChatCompletionTask(task_id=TASK_1_ID, task_type=TaskType.CHAT_COMPLETION, task_status=TaskStatus.PENDING, task_params=ChatCompletionTaskParams(model=str(MODEL_A_ID), messages=[ChatCompletionMessage(role="user", content="Hello, world!")]), instance_id=INSTANCE_1_ID)},
+ ),
+ expected_op=RunnerUpOp(runner_id=RUNNER_1_ID)
+ ),
-# ---------------------------------------------------------------------------
-# Scenarios
-# ---------------------------------------------------------------------------
+ PlanTestCase(
+ description="loaded runner (and state down) -> expect RunnerDownOp",
+ in_process_runners=[
+ InProcessRunner(
+ runner_id=RUNNER_1_ID,
+ instance_id=INSTANCE_1_ID,
+ model_id=MODEL_A_ID,
+ status=LoadedRunnerStatus(),
+ downloaded=True,
+ )
+ ],
+ state=State(
+ node_status={NODE_A: NodeStatus.Idle},
+ instances={
+ INSTANCE_1_ID: Instance(
+ instance_type=TypeOfInstance.INACTIVE,
+ instance_id=INSTANCE_1_ID,
+ instance_params=InstanceParams(
+ shard_assignments=ShardAssignments(
+ model_id=MODEL_A_ID,
+ runner_to_shard={
+ RUNNER_1_ID: make_shard_metadata(device_rank=0, world_size=1)
+ },
+ node_to_runner={NODE_A: RUNNER_1_ID}
+ ),
+ hosts=[]
+ ),
+ )
+ },
+ runners={RUNNER_1_ID: LoadedRunnerStatus()},
+ tasks={},
+ ),
+ expected_op=RunnerDownOp(runner_id=RUNNER_1_ID),
+ ),
-TEST_CASES: Final[List[PlanTestCase]] = [
- PlanTestCase(
- description="no runners ⇢ no-op",
- runners=[],
- expected_op_type=None,
- expected_op_runner_idx=None,
- ),
- PlanTestCase(
- description="single ready runner, model missing ⇢ expect DownloadOp",
- runners=[
- RunnerCase(status=ReadyRunnerStatus(), downloaded=False),
- ],
- expected_op_type=DownloadOp,
- expected_op_runner_idx=0,
- ),
- PlanTestCase(
- description="runner already downloading ⇢ no-op",
- runners=[
- RunnerCase(status=_make_downloading_status(NodeId()), downloaded=False),
- ],
- expected_op_type=None,
- expected_op_runner_idx=None,
- ),
- PlanTestCase(
- description="ready runner, model present ⇢ no-op",
- runners=[
- RunnerCase(status=ReadyRunnerStatus(), downloaded=True),
- ],
- expected_op_type=None,
- expected_op_runner_idx=None,
- ),
- PlanTestCase(
- description="instance for other node ⇢ no-op",
- runners=[
- RunnerCase(status=ReadyRunnerStatus(), downloaded=False),
- ],
- expected_op_type=None,
- expected_op_runner_idx=None,
- worker_state_override=State(
- node_status={NodeId(): NodeStatus.Idle},
- instances={},
+ PlanTestCase(
+ description="failed runner (and state down) -> expect RunnerDownOp",
+ in_process_runners=[
+ InProcessRunner(
+ runner_id=RUNNER_1_ID,
+ instance_id=INSTANCE_1_ID,
+ model_id=MODEL_A_ID,
+ status=FailedRunnerStatus(),
+ downloaded=True,
+ )
+ ],
+ state=State(
+ node_status={NODE_A: NodeStatus.Idle},
+ instances={
+ INSTANCE_1_ID: Instance(
+ instance_type=TypeOfInstance.INACTIVE,
+ instance_id=INSTANCE_1_ID,
+ instance_params=InstanceParams(
+ shard_assignments=ShardAssignments(
+ model_id=MODEL_A_ID,
+ runner_to_shard={
+ RUNNER_1_ID: make_shard_metadata(device_rank=0, world_size=1)
+ },
+ node_to_runner={NODE_A: RUNNER_1_ID}
+ ),
+ hosts=[]
+ ),
+ )
+ },
+ runners={RUNNER_1_ID: FailedRunnerStatus()},
+ tasks={},
+ ),
+ expected_op=RunnerDownOp(runner_id=RUNNER_1_ID),
),
- ),
-]
+ PlanTestCase(
+ description="loaded runner, model present, task pending -> expect ExecuteTaskOp",
+ in_process_runners=[
+ InProcessRunner(
+ runner_id=RUNNER_1_ID,
+ instance_id=INSTANCE_1_ID,
+ model_id=MODEL_A_ID,
+ status=LoadedRunnerStatus(),
+ downloaded=True,
+ )
+ ],
+ state=State(
+ node_status={NODE_A: NodeStatus.Idle},
+ instances={
+ INSTANCE_1_ID: Instance(
+ instance_type=TypeOfInstance.ACTIVE,
+ instance_id=INSTANCE_1_ID,
+ instance_params=InstanceParams(
+ shard_assignments=ShardAssignments(
+ model_id=MODEL_A_ID,
+ runner_to_shard={
+ RUNNER_1_ID: make_shard_metadata(device_rank=0, world_size=1)
+ },
+ node_to_runner={NODE_A: RUNNER_1_ID}
+ ),
+ hosts=[]
+ ),
+ )
+ },
+ runners={RUNNER_1_ID: LoadedRunnerStatus()},
+ tasks={TASK_1_ID: ChatCompletionTask(task_id=TASK_1_ID, task_type=TaskType.CHAT_COMPLETION, task_status=TaskStatus.PENDING, task_params=ChatCompletionTaskParams(model=str(MODEL_A_ID), messages=[ChatCompletionMessage(role="user", content="Hello, world!")]), instance_id=INSTANCE_1_ID)},
+ ),
+ expected_op=ExecuteTaskOp(runner_id=RUNNER_1_ID, task=ChatCompletionTask(
+ task_id=TASK_1_ID,
+ instance_id=INSTANCE_1_ID,
+ task_type=TaskType.CHAT_COMPLETION,
+ task_status=TaskStatus.PENDING,
+ task_params=ChatCompletionTaskParams(
+ model=str(MODEL_A_ID),
+ messages=[ChatCompletionMessage(role="user", content="Hello, world!")]
+ ),
+ )),
+ ),
-# ---------------------------------------------------------------------------
-# Shared factory helpers
-# ---------------------------------------------------------------------------
+ PlanTestCase(
+ # We should only run rank 0 once all other ranks are running.
+ description="two loaded runners & task, i'm rank 0 -> no-op",
+ in_process_runners=[
+ InProcessRunner(
+ runner_id=RUNNER_1_ID,
+ instance_id=INSTANCE_1_ID,
+ model_id=MODEL_A_ID,
+ status=LoadedRunnerStatus(),
+ downloaded=True,
+ device_rank=0,
+ ),
+ InProcessRunner(
+ runner_id=RUNNER_2_ID,
+ instance_id=INSTANCE_1_ID,
+ model_id=MODEL_A_ID,
+ status=LoadedRunnerStatus(),
+ downloaded=True,
+ device_rank=1,
+ ),
+ ],
+ state=State(
+ node_status={NODE_A: NodeStatus.Idle, NODE_B: NodeStatus.Idle},
+ instances={
+ INSTANCE_1_ID: Instance(
+ instance_type=TypeOfInstance.ACTIVE,
+ instance_id=INSTANCE_1_ID,
+ instance_params=InstanceParams(
+ shard_assignments=ShardAssignments(
+ model_id=MODEL_A_ID,
+ runner_to_shard={
+ RUNNER_1_ID: make_shard_metadata(device_rank=0, world_size=2),
+ RUNNER_2_ID: make_shard_metadata(device_rank=1, world_size=2)
+ },
+ node_to_runner={NODE_A: RUNNER_1_ID, NODE_B: RUNNER_2_ID}
+ ),
+ hosts=[]
+ ),
+ )
+ },
+ runners={RUNNER_1_ID: LoadedRunnerStatus(), RUNNER_2_ID: LoadedRunnerStatus()},
+ tasks={TASK_1_ID: ChatCompletionTask(task_id=TASK_1_ID, task_type=TaskType.CHAT_COMPLETION, task_status=TaskStatus.PENDING, task_params=ChatCompletionTaskParams(model=str(MODEL_A_ID), messages=[ChatCompletionMessage(role="user", content="Hello, world!")]), instance_id=INSTANCE_1_ID)},
+ ),
+ expected_op=None
+ ),
+ PlanTestCase(
+ description="two loaded runners & task, i'm rank 1 -> expect ExecuteTaskOp on rank 1",
+ in_process_runners=[
+ InProcessRunner(
+ runner_id=RUNNER_1_ID,
+ instance_id=INSTANCE_1_ID,
+ model_id=MODEL_A_ID,
+ status=LoadedRunnerStatus(),
+ downloaded=True,
+ device_rank=1,
+ ),
+ InProcessRunner(
+ runner_id=RUNNER_2_ID,
+ instance_id=INSTANCE_1_ID,
+ model_id=MODEL_A_ID,
+ status=LoadedRunnerStatus(),
+ downloaded=True,
+ device_rank=0,
+ ),
+ ],
+ state=State(
+ node_status={NODE_A: NodeStatus.Idle, NODE_B: NodeStatus.Idle},
+ instances={
+ INSTANCE_1_ID: Instance(
+ instance_type=TypeOfInstance.ACTIVE,
+ instance_id=INSTANCE_1_ID,
+ instance_params=InstanceParams(
+ shard_assignments=ShardAssignments(
+ model_id=MODEL_A_ID,
+ runner_to_shard={
+ RUNNER_1_ID: make_shard_metadata(device_rank=1, world_size=2),
+ RUNNER_2_ID: make_shard_metadata(device_rank=0, world_size=2)
+ },
+ node_to_runner={NODE_A: RUNNER_1_ID, NODE_B: RUNNER_2_ID}
+ ),
+ hosts=[]
+ ),
+ )
+ },
+ runners={RUNNER_1_ID: LoadedRunnerStatus(), RUNNER_2_ID: LoadedRunnerStatus()},
+ tasks={TASK_1_ID: ChatCompletionTask(task_id=TASK_1_ID, task_type=TaskType.CHAT_COMPLETION, task_status=TaskStatus.PENDING, task_params=ChatCompletionTaskParams(model=str(MODEL_A_ID), messages=[ChatCompletionMessage(role="user", content="Hello, world!")]), instance_id=INSTANCE_1_ID)},
+ ),
+ expected_op=ExecuteTaskOp(
+ runner_id=RUNNER_1_ID,
+ task=ChatCompletionTask(
+ task_id=TASK_1_ID,
+ instance_id=INSTANCE_1_ID,
+ task_type=TaskType.CHAT_COMPLETION,
+ task_params=ChatCompletionTaskParams(
+ model=str(MODEL_A_ID),
+ messages=[ChatCompletionMessage(role="user", content="Hello, world!")],
+ ),
+ task_status=TaskStatus.PENDING,
+ ),
+ ),
+ ),
-@dataclass(frozen=True, slots=True)
-class RunnerContext:
- runner_id: RunnerId
- instance_id: InstanceId
- shard_metadata: PipelineShardMetadata
- instance_params: InstanceParams
-
-
-# TODO: generalize this it's in conftest.
-def _build_worker_state(
- *,
- tmp_path: Path,
- node_id: NodeId,
- pipeline_shard_metadata: PipelineShardMetadata,
- runner_cases: List[RunnerCase],
-) -> tuple[State, List[RunnerContext]]:
- """Construct a WorkerState plus per-runner context objects."""
-
- instances: dict[InstanceId, Instance] = {}
- runner_contexts: list[RunnerContext] = []
-
- for idx, _ in enumerate(runner_cases):
- runner_id = RunnerId()
- instance_id = InstanceId()
- model_id = ModelId()
-
- # Unique sub-directory per runner to allow selective `downloaded` mocking.
- model_subdir = tmp_path / f"runner_{idx}"
- model_subdir.mkdir(exist_ok=True)
-
- shard_assignments = ShardAssignments(
- model_id=model_id,
- runner_to_shard={runner_id: pipeline_shard_metadata},
- node_to_runner={node_id: runner_id},
- )
+ PlanTestCase(
+ description="rank 1 loaded, rank 0 ready, i'm rank 0 -> expect ExecuteTaskOp on rank 0",
+ in_process_runners=[
+ InProcessRunner(
+ runner_id=RUNNER_1_ID,
+ instance_id=INSTANCE_1_ID,
+ model_id=MODEL_A_ID,
+ status=LoadedRunnerStatus(),
+ downloaded=True,
+ device_rank=0,
+ ),
+ InProcessRunner(
+ runner_id=RUNNER_2_ID,
+ instance_id=INSTANCE_1_ID,
+ model_id=MODEL_A_ID,
+ status=RunningRunnerStatus(),
+ downloaded=True,
+ device_rank=1,
+ ),
+ ],
+ state=State(
+ node_status={NODE_A: NodeStatus.Idle, NODE_B: NodeStatus.Running},
+ instances={
+ INSTANCE_1_ID: Instance(
+ instance_type=TypeOfInstance.ACTIVE,
+ instance_id=INSTANCE_1_ID,
+ instance_params=InstanceParams(
+ shard_assignments=ShardAssignments(
+ model_id=MODEL_A_ID,
+ runner_to_shard={
+ RUNNER_1_ID: make_shard_metadata(device_rank=0, world_size=2),
+ RUNNER_2_ID: make_shard_metadata(device_rank=1, world_size=2)
+ },
+ node_to_runner={NODE_A: RUNNER_1_ID, NODE_B: RUNNER_2_ID}
+ ),
+ hosts=[]
+ ),
+ )
+ },
+ runners={RUNNER_1_ID: LoadedRunnerStatus(), RUNNER_2_ID: RunningRunnerStatus()},
+ tasks={TASK_1_ID: ChatCompletionTask(task_id=TASK_1_ID, task_type=TaskType.CHAT_COMPLETION, task_status=TaskStatus.PENDING, task_params=ChatCompletionTaskParams(model=str(MODEL_A_ID), messages=[ChatCompletionMessage(role="user", content="Hello, world!")]), instance_id=INSTANCE_1_ID)},
+ ),
+ expected_op=ExecuteTaskOp(
+ runner_id=RUNNER_1_ID,
+ task=ChatCompletionTask(
+ task_id=TASK_1_ID,
+ instance_id=INSTANCE_1_ID,
+ task_type=TaskType.CHAT_COMPLETION,
+ task_params=ChatCompletionTaskParams(
+ model=str(MODEL_A_ID),
+ messages=[ChatCompletionMessage(role="user", content="Hello, world!")],
+ ),
+ task_status=TaskStatus.PENDING,
+ ),
+ ),
+ ),
- instance_params = InstanceParams(
- shard_assignments=shard_assignments,
- hosts=[],
- )
+ PlanTestCase(
+ description="other runner failed -> RunnerDownOp",
+ in_process_runners=[
+ InProcessRunner(
+ runner_id=RUNNER_1_ID,
+ instance_id=INSTANCE_1_ID,
+ model_id=MODEL_A_ID,
+ status=LoadedRunnerStatus(),
+ downloaded=True,
+ device_rank=0,
+ ),
+ InProcessRunner(
+ runner_id=RUNNER_2_ID,
+ instance_id=INSTANCE_1_ID,
+ model_id=MODEL_A_ID,
+ status=FailedRunnerStatus(),
+ downloaded=True,
+ device_rank=1,
+ ),
+ ],
+ state=State(
+ node_status={NODE_A: NodeStatus.Idle, NODE_B: NodeStatus.Idle},
+ instances={
+ INSTANCE_1_ID: Instance(
+ instance_type=TypeOfInstance.ACTIVE,
+ instance_id=INSTANCE_1_ID,
+ instance_params=InstanceParams(
+ shard_assignments=ShardAssignments(
+ model_id=MODEL_A_ID,
+ runner_to_shard={
+ RUNNER_1_ID: make_shard_metadata(device_rank=0, world_size=2),
+ RUNNER_2_ID: make_shard_metadata(device_rank=1, world_size=2)
+ },
+ node_to_runner={NODE_A: RUNNER_1_ID, NODE_B: RUNNER_2_ID}
+ ),
+ hosts=[]
+ ),
+ )
+ },
+ runners={RUNNER_1_ID: LoadedRunnerStatus(), RUNNER_2_ID: FailedRunnerStatus()},
+ ),
+ expected_op=RunnerDownOp(runner_id=RUNNER_1_ID)
+ ),
- instance = Instance(
- instance_id=instance_id,
- instance_params=instance_params,
- instance_type=TypeOfInstance.ACTIVE,
- )
+ PlanTestCase(
+ description="this runner failed (1 node) -> RunnerDownOp",
+ in_process_runners=[
+ InProcessRunner(
+ runner_id=RUNNER_1_ID,
+ instance_id=INSTANCE_1_ID,
+ model_id=MODEL_A_ID,
+ status=FailedRunnerStatus(),
+ downloaded=True,
+ device_rank=0,
+ ),
+ ],
+ state=State(
+ node_status={NODE_A: NodeStatus.Idle, NODE_B: NodeStatus.Idle},
+ instances={
+ INSTANCE_1_ID: Instance(
+ instance_type=TypeOfInstance.ACTIVE,
+ instance_id=INSTANCE_1_ID,
+ instance_params=InstanceParams(
+ shard_assignments=ShardAssignments(
+ model_id=MODEL_A_ID,
+ runner_to_shard={
+ RUNNER_1_ID: make_shard_metadata(device_rank=0, world_size=1),
+ },
+ node_to_runner={NODE_A: RUNNER_1_ID}
+ ),
+ hosts=[]
+ ),
+ )
+ },
+ runners={RUNNER_1_ID: FailedRunnerStatus()},
+ ),
+ expected_op=RunnerDownOp(runner_id=RUNNER_1_ID)
+ ),
- instances[instance_id] = instance
- runner_contexts.append(
- RunnerContext(
- runner_id=runner_id,
- instance_id=instance_id,
- shard_metadata=pipeline_shard_metadata,
- instance_params=instance_params,
- )
- )
+ PlanTestCase(
+ description="this runner failed (2 nodes) -> no-op",
+ in_process_runners=[
+ InProcessRunner(
+ runner_id=RUNNER_1_ID,
+ instance_id=INSTANCE_1_ID,
+ model_id=MODEL_A_ID,
+ status=FailedRunnerStatus(),
+ downloaded=True,
+ device_rank=0,
+ ),
+ InProcessRunner(
+ runner_id=RUNNER_2_ID,
+ instance_id=INSTANCE_1_ID,
+ model_id=MODEL_A_ID,
+ status=LoadedRunnerStatus(),
+ downloaded=True,
+ device_rank=1,
+ ),
+ ],
+ state=State(
+ node_status={NODE_A: NodeStatus.Idle, NODE_B: NodeStatus.Idle},
+ instances={
+ INSTANCE_1_ID: Instance(
+ instance_type=TypeOfInstance.ACTIVE,
+ instance_id=INSTANCE_1_ID,
+ instance_params=InstanceParams(
+ shard_assignments=ShardAssignments(
+ model_id=MODEL_A_ID,
+ runner_to_shard={
+ RUNNER_1_ID: make_shard_metadata(device_rank=0, world_size=2),
+ RUNNER_2_ID: make_shard_metadata(device_rank=1, world_size=2)
+ },
+ node_to_runner={NODE_A: RUNNER_1_ID, NODE_B: RUNNER_2_ID}
+ ),
+ hosts=[]
+ ),
+ )
+ },
+ runners={RUNNER_1_ID: FailedRunnerStatus(), RUNNER_2_ID: LoadedRunnerStatus()},
+ ),
+ expected_op=None
+ ),
- worker_state = State(
- node_status={node_id: NodeStatus.Idle},
- instances=instances,
- )
+ PlanTestCase(
+ description="this node failed, other node spun down -> RunnerDownOp",
+ in_process_runners=[
+ InProcessRunner(
+ runner_id=RUNNER_1_ID,
+ instance_id=INSTANCE_1_ID,
+ model_id=MODEL_A_ID,
+ status=FailedRunnerStatus(),
+ downloaded=True,
+ device_rank=0,
+ ),
+ InProcessRunner(
+ runner_id=RUNNER_2_ID,
+ instance_id=INSTANCE_1_ID,
+ model_id=MODEL_A_ID,
+ status=ReadyRunnerStatus(),
+ downloaded=True,
+ device_rank=1,
+ ),
+ ],
+ state=State(
+ node_status={NODE_A: NodeStatus.Idle, NODE_B: NodeStatus.Idle},
+ instances={
+ INSTANCE_1_ID: Instance(
+ instance_type=TypeOfInstance.ACTIVE,
+ instance_id=INSTANCE_1_ID,
+ instance_params=InstanceParams(
+ shard_assignments=ShardAssignments(
+ model_id=MODEL_A_ID,
+ runner_to_shard={
+ RUNNER_1_ID: make_shard_metadata(device_rank=0, world_size=2),
+ RUNNER_2_ID: make_shard_metadata(device_rank=1, world_size=2)
+ },
+ node_to_runner={NODE_A: RUNNER_1_ID, NODE_B: RUNNER_2_ID}
+ ),
+ hosts=[]
+ ),
+ )
+ },
+ runners={RUNNER_1_ID: FailedRunnerStatus(), RUNNER_2_ID: ReadyRunnerStatus()},
+ ),
+ expected_op=RunnerDownOp(runner_id=RUNNER_1_ID)
+ ),
- return worker_state, runner_contexts
+ ]
# ---------------------------------------------------------------------------
@@ -189,46 +839,80 @@ def _build_worker_state(
# Pre-compute readable identifiers for each case to avoid lambda typing issues.
-@pytest.mark.parametrize("case", TEST_CASES, ids=[case.id() for case in TEST_CASES])
-def test_worker_plan(case: PlanTestCase, tmp_path: Path, monkeypatch: pytest.MonkeyPatch, pipeline_shard_meta: Callable[..., PipelineShardMetadata]) -> None:
+@pytest.mark.parametrize(
+ "case",
+ # We use a factory to delay test case generation until tmp_path is available.
+ [pytest.param(c, id=c.id()) for c in _get_test_cases(Path(tempfile.TemporaryDirectory().name))],
+)
+def test_worker_plan(case: PlanTestCase, tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
"""Exercise Worker.plan across declarative scenarios."""
- # Fresh identifier for isolation of node
- node_id = NodeId()
+ print(f"----- case: {case.description}")
- # Assemble WorkerState and surrounding objects ---------------------------------------
- worker_state, runner_contexts = _build_worker_state(
- tmp_path=tmp_path,
- node_id=node_id,
- pipeline_shard_metadata=pipeline_shard_meta(1, 0),
- runner_cases=case.runners,
- )
+ # Regenerate test cases with the actual tmp_path fixture
+ test_cases = {c.description: c for c in _get_test_cases(tmp_path)}
+ case = test_cases[case.description]
- # Replace with explicit override if provided by the scenario.
- if case.worker_state_override is not None:
- worker_state = case.worker_state_override
+ node_id = NODE_A
+ initial_state = State(
+ node_status={node_id: NodeStatus.Idle},
+ instances={},
+ runners={},
+ tasks={},
+ )
logger = logging.getLogger("test_worker_plan")
- worker = Worker(node_id=node_id, initial_state=worker_state, logger=logger)
+ worker = Worker(node_id=node_id, initial_state=initial_state, worker_events=None, logger=logger)
- # Build assigned_runners and a path→downloaded lookup --------------------------------
path_downloaded_map: dict[str, bool] = {}
- for idx, runner_case in enumerate(case.runners):
- runner_status = runner_case.status
- ctx = runner_contexts[idx]
+ runner_config: InProcessRunner
+ for runner_config in case.in_process_runners:
+
+ model_path = tmp_path / f"model_for_runner_{runner_config.runner_id}"
+ model_path.mkdir(exist_ok=True, parents=True)
+
+ if len(case.state.instances) == 1:
+ instance_id = next(iter(case.state.instances))
+
+ shard_assignments = case.state.instances[instance_id].instance_params.shard_assignments
+ shard_metadata = shard_assignments.runner_to_shard[runner_config.runner_id]
+
+ # Only add this runner if it belongs to our node
+ runner_node = None
+ for node, runner in shard_assignments.node_to_runner.items():
+ if runner == runner_config.runner_id:
+ runner_node = node
+ break
+
+ if runner_node != node_id:
+ # This runner belongs to a different node, skip it
+ continue
+
+ elif len(case.state.instances) == 0:
+ shard_metadata = PipelineShardMetadata(
+ device_rank=runner_config.device_rank,
+ world_size=1,
+ model_meta=make_model_meta(runner_config.model_id),
+ start_layer=0,
+ end_layer=1,
+ n_layers=1,
+ )
+ else:
+ raise Exception('test_worker_plan not currently designed to have more than 1 instance.')
+
- assigned_runner = AssignedRunner(
- runner_id=ctx.runner_id,
- instance_id=ctx.instance_id,
- shard_metadata=ctx.shard_metadata,
- hosts=ctx.instance_params.hosts,
- status=runner_status,
+ assigned_runner = OverrideAssignedRunner(
+ runner_id=runner_config.runner_id,
+ instance_id=runner_config.instance_id,
+ shard_metadata=shard_metadata,
+ hosts=[],
+ status=runner_config.status,
runner=None,
+ downloaded=runner_config.downloaded
)
- worker.assigned_runners[ctx.runner_id] = assigned_runner
-
- path_downloaded_map[str(build_model_path(ctx.shard_metadata.model_meta.model_id))] = runner_case.downloaded
+ worker.assigned_runners[runner_config.runner_id] = assigned_runner
+ path_downloaded_map[str(build_model_path(shard_metadata.model_meta.model_id))] = runner_config.downloaded
# Stub filesystem existence check ------------------------------------------------------
from worker import main as worker_main # local import for module-scoped os
@@ -238,19 +922,5 @@ def test_worker_plan(case: PlanTestCase, tmp_path: Path, monkeypatch: pytest.Mon
monkeypatch.setattr(worker_main.os.path, "exists", _fake_exists)
- # Plan and assert ----------------------------------------------------------------------
- op = worker.plan(worker_state)
-
- if case.expected_op_type is None:
- assert op is None, f"Unexpected op {op} for scenario: {case.description}"
- else:
- assert isinstance(op, case.expected_op_type), (
- f"Expected {case.expected_op_type.__name__}, got {type(op).__name__ if op else 'None'}"
- )
-
- assert case.expected_op_runner_idx is not None, "Runner index must be set when expecting an op"
- target_ctx = runner_contexts[case.expected_op_runner_idx]
-
- assert op.runner_id == target_ctx.runner_id
- assert op.instance_id == target_ctx.instance_id
- assert op.shard_metadata == target_ctx.shard_metadata
+ op = worker.plan(case.state)
+ assert op == case.expected_op
diff --git a/worker/tests/test_worker_plan_utils.py b/worker/tests/test_worker_plan_utils.py
new file mode 100644
index 00000000..05298efd
--- /dev/null
+++ b/worker/tests/test_worker_plan_utils.py
@@ -0,0 +1,212 @@
+from __future__ import annotations
+
+from dataclasses import dataclass
+from pathlib import Path
+from typing import Final, List, Optional, override
+from uuid import UUID
+
+from shared.types.common import NodeId
+from shared.types.models import ModelId, ModelMetadata
+from shared.types.state import State
+from shared.types.tasks import TaskId
+from shared.types.worker.common import InstanceId, NodeStatus, RunnerId
+from shared.types.worker.downloads import DownloadOngoing, DownloadProgressData
+from shared.types.worker.instances import Instance, InstanceParams, TypeOfInstance
+from shared.types.worker.ops import RunnerOp
+from shared.types.worker.runners import (
+ AssignedRunnerStatus,
+ DownloadingRunnerStatus,
+ RunnerStatus,
+ ShardAssignments,
+)
+from shared.types.worker.shards import PipelineShardMetadata
+from worker.download.model_cards import MODEL_CARDS, ModelCard
+from worker.main import AssignedRunner
+
+NODE_A: Final[NodeId] = NodeId(uuid=UUID("aaaaaaaa-aaaa-4aaa-8aaa-aaaaaaaaaaaa"))
+NODE_B: Final[NodeId] = NodeId(uuid=UUID("bbbbbbbb-bbbb-4bbb-8bbb-bbbbbbbbbbbb"))
+
+# Define constant IDs for deterministic test cases
+RUNNER_1_ID: Final[RunnerId] = RunnerId(uuid=UUID("cccccccc-aaaa-4aaa-8aaa-aaaaaaaaaaaa"))
+INSTANCE_1_ID: Final[InstanceId] = InstanceId()
+RUNNER_2_ID: Final[RunnerId] = RunnerId(uuid=UUID("dddddddd-aaaa-4aaa-8aaa-aaaaaaaaaaaa"))
+INSTANCE_2_ID: Final[InstanceId] = InstanceId()
+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()
+
+@dataclass(slots=True, frozen=True)
+class InProcessRunner:
+ """Minimal description of a runner's in-process state."""
+ # TODO: Rename to InProcessRunnerConfig and create a constructor for OverrideAssignedRunner.
+
+ runner_id: RunnerId
+ instance_id: InstanceId
+ model_id: ModelId
+ status: RunnerStatus
+ downloaded: bool
+ device_rank: int = 0
+
+# Helper class to override the is_downloaded property to whatever is specified by InProcessRunner
+class OverrideAssignedRunner(AssignedRunner):
+ downloaded: bool
+
+ @property
+ @override
+ def is_downloaded(self) -> bool:
+ return self.downloaded
+
+
+@dataclass(slots=True, frozen=True)
+class PlanTestCase:
+ """Table-driven description of an entire planning scenario."""
+
+ description: str
+ state: State
+ in_process_runners: List[InProcessRunner]
+ expected_op: Optional[RunnerOp]
+
+ def id(self) -> str: # noqa: D401
+ return self.description.replace(" ", "_")
+
+
+def make_shard_metadata(device_rank: int, world_size: int, model_id: ModelId = MODEL_A_ID) -> PipelineShardMetadata:
+ """Create PipelineShardMetadata with proper layer assignments based on device_rank and world_size."""
+ total_layers = world_size # For simplicity in tests, total_layers = world_size
+
+ if world_size == 1:
+ start_layer = 0
+ end_layer = 1
+ n_layers = 1
+ else:
+ # For multi-device setup, each device gets one layer
+ start_layer = device_rank
+ end_layer = device_rank + 1
+ n_layers = total_layers
+
+ return PipelineShardMetadata(
+ device_rank=device_rank,
+ world_size=world_size,
+ model_meta=make_model_meta(model_id),
+ start_layer=start_layer,
+ end_layer=end_layer,
+ n_layers=n_layers,
+ )
+
+
+def make_downloading_status(node_id: NodeId) -> DownloadingRunnerStatus:
+ """Factory for a *Downloading* status with placeholder progress."""
+ return DownloadingRunnerStatus(
+ download_progress=DownloadOngoing(
+ node_id=node_id,
+ download_progress=DownloadProgressData(total_bytes=1, downloaded_bytes=0),
+ )
+ )
+
+def make_model_meta(
+ model_id: str
+) -> ModelMetadata:
+ model_card: ModelCard
+ for card in MODEL_CARDS.values():
+ if card.repo_id == model_id:
+ model_card = card
+
+ return ModelMetadata(
+ model_id=model_id,
+ pretty_name=model_card.id,
+ storage_size_kilobytes=10**6,
+ n_layers=16,
+ )
+
+ raise Exception(f'Unknown model_id passed: {model_id}')
+
+ ## Alternatively, if we are ok for this method to be async:
+ # await _get_model_meta(model_id)
+
+
+def create_worker_state(
+ *,
+ node_id: NodeId,
+ runner_configs: list[tuple[RunnerId, InstanceId, ModelId]],
+ tmp_path: Path,
+) -> State:
+ """Create a test `State` based on a list of runner configurations."""
+ instances: dict[InstanceId, Instance] = {}
+ for runner_id, instance_id, model_id in runner_configs:
+ model_path = tmp_path / f"model_for_runner_{runner_id}"
+ model_path.mkdir(exist_ok=True, parents=True)
+
+ shard_metadata = PipelineShardMetadata(
+ device_rank=0,
+ world_size=1,
+ model_meta=make_model_meta(model_id),
+ start_layer=0,
+ end_layer=1,
+ n_layers=1,
+ )
+ shard_assignments = ShardAssignments(
+ model_id=model_id,
+ runner_to_shard={runner_id: shard_metadata},
+ node_to_runner={node_id: runner_id},
+ )
+ instance_params = InstanceParams(
+ shard_assignments=shard_assignments,
+ hosts=[],
+ )
+ instance = Instance(
+ instance_id=instance_id,
+ instance_params=instance_params,
+ instance_type=TypeOfInstance.ACTIVE,
+ )
+ instances[instance_id] = instance
+
+ return State(
+ node_status={node_id: NodeStatus.Idle},
+ instances=instances,
+ runners={runner_id: AssignedRunnerStatus() for runner_id, _, _ in runner_configs},
+ tasks={},
+ )
+
+
+def make_instance(
+ instance_id: InstanceId,
+ model_id: ModelId,
+ tmp_path: Path,
+ runner_specs: list[tuple[RunnerId, NodeId, int]],
+) -> Instance:
+ """Creates an instance with one or more runners."""
+ runner_to_shard: dict[RunnerId, PipelineShardMetadata] = {}
+ node_to_runner: dict[NodeId, RunnerId] = {}
+ world_size = len(runner_specs)
+
+ for runner_id, node_id, device_rank in runner_specs:
+ model_path = tmp_path / f"model_for_runner_{runner_id}"
+ model_path.mkdir(exist_ok=True, parents=True)
+
+ shard_metadata = PipelineShardMetadata(
+ device_rank=device_rank,
+ world_size=world_size,
+ model_meta=make_model_meta(model_id),
+ start_layer=0,
+ end_layer=1,
+ n_layers=1,
+ )
+ runner_to_shard[runner_id] = shard_metadata
+ node_to_runner[node_id] = runner_id
+
+ shard_assignments = ShardAssignments(
+ model_id=model_id,
+ runner_to_shard=runner_to_shard,
+ node_to_runner=node_to_runner,
+ )
+ instance_params = InstanceParams(
+ shard_assignments=shard_assignments,
+ hosts=[],
+ )
+ return Instance(
+ instance_id=instance_id,
+ instance_params=instance_params,
+ instance_type=TypeOfInstance.ACTIVE,
+ )
+
+### For worker plan tests
\ No newline at end of file
diff --git a/worker/tests/test_worker_state.py b/worker/tests/test_worker_state.py
index 99f154d7..1d010101 100644
--- a/worker/tests/test_worker_state.py
+++ b/worker/tests/test_worker_state.py
@@ -1,6 +1,7 @@
## Tests for worker state differentials
## When the worker state changes, this should be reflected by a worker intention.
+
import asyncio
from typing import Callable
from uuid import uuid4
@@ -19,7 +20,7 @@ async def test_worker_runs_and_stops(worker: Worker):
await worker.start()
await asyncio.sleep(0.01)
- assert worker._is_running # type: ignore
+ assert worker._is_running, worker._task.exception() # type: ignore
await worker.stop()
await asyncio.sleep(0.01)
← 56d35657 Add apply functions
·
back to Exo
·
Fix tests 5097493a →