[object Object]

← 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

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 →