[object Object]

← back to Exo

Fixes for running this end to end

57ca487fdefb7b7b7715d637f847a3b30354c396 · 2025-07-28 10:51:03 +0100 · Alex Cheema

Co-authored-by: Gelu Vrabie <gelu.vrabie.univ@gmail.com>
Co-authored-by: Gelu Vrabie <gelu@exolabs.net>

Files touched

Diff

commit 57ca487fdefb7b7b7715d637f847a3b30354c396
Author: Alex Cheema <41707476+AlexCheema@users.noreply.github.com>
Date:   Mon Jul 28 10:51:03 2025 +0100

    Fixes for running this end to end
    
    Co-authored-by: Gelu Vrabie <gelu.vrabie.univ@gmail.com>
    Co-authored-by: Gelu Vrabie <gelu@exolabs.net>
---
 master/api.py                            |  56 ++++++++++-------
 master/forwarder_supervisor.py           |   6 ++
 master/main.py                           |  18 ++++--
 master/placement.py                      |   6 +-
 master/tests/test_forwarder_manager.py   |  10 +--
 master/tests/test_master.py              | 101 ++++++++++++++++++++++++++++---
 master/tests/test_placement.py           |  14 +++--
 master/tests/test_placement_utils.py     |  26 ++++----
 rust/discovery/Cargo.toml                |   1 +
 rust/discovery/src/lib.rs                |   6 +-
 rust/exo_pyo3_bindings/src/discovery.rs  |  87 ++++++++++++++++++--------
 shared/apply/apply.py                    |   5 ++
 shared/topology.py                       |   9 ++-
 shared/types/tasks.py                    |   3 +-
 worker/download/download_utils.py        |  26 +++++---
 worker/download/impl_shard_downloader.py |  10 +++
 worker/download/shard_downloader.py      |  19 ++++++
 worker/main.py                           |  98 +++++++++++++++++++++++-------
 worker/runner/runner_supervisor.py       |   2 +-
 worker/tests/conftest.py                 |  10 ++-
 worker/tests/test_worker_integration.py  |  11 ++--
 worker/tests/test_worker_plan.py         |  34 +++++------
 worker/tests/test_worker_plan_utils.py   |  16 +----
 23 files changed, 412 insertions(+), 162 deletions(-)

diff --git a/master/api.py b/master/api.py
index 387f2e5d..5b4ea986 100644
--- a/master/api.py
+++ b/master/api.py
@@ -29,6 +29,7 @@ from shared.types.events.commands import (
     DeleteInstanceCommand,
 )
 from shared.types.events.components import EventFromEventLog
+from shared.types.models import ModelMetadata
 from shared.types.state import State
 from shared.types.tasks import ChatCompletionTaskParams
 from shared.types.worker.common import InstanceId
@@ -37,9 +38,9 @@ from shared.types.worker.instances import Instance
 
 def chunk_to_response(chunk: TokenChunk) -> ChatCompletionResponse:
     return ChatCompletionResponse(
-        id='abc',
+        id=chunk.command_id,
         created=int(time.time()),
-        model='idk',
+        model=chunk.model,
         choices=[
             StreamingChoiceResponse(
                 index=0,
@@ -52,6 +53,12 @@ def chunk_to_response(chunk: TokenChunk) -> ChatCompletionResponse:
         ]
     )
 
+async def resolve_model_meta(model_id: str) -> ModelMetadata:
+    if model_id in MODEL_CARDS:
+        model_card = MODEL_CARDS[model_id]
+        return model_card.metadata
+    else:
+        return await get_model_meta(model_id)
 
 @final
 class API:
@@ -67,7 +74,7 @@ class API:
         # self._app.get("/topology/control_plane")(self.get_control_plane_topology)
         # self._app.get("/topology/data_plane")(self.get_data_plane_topology)
         # self._app.get("/instances/list")(self.list_instances)
-        self._app.post("/instances/create")(self.create_instance)
+        self._app.post("/instance")(self.create_instance)
         self._app.get("/instance/{instance_id}")(self.get_instance)
         self._app.delete("/instance/{instance_id}")(self.delete_instance)
         # self._app.get("/model/{model_id}/metadata")(self.get_model_data)
@@ -92,11 +99,7 @@ class API:
     #     return {"message": "Hello, World!"}
 
     async def create_instance(self, payload: CreateInstanceTaskParams) -> CreateInstanceResponse:
-        if payload.model_id in MODEL_CARDS:
-            model_card = MODEL_CARDS[payload.model_id]
-            model_meta = model_card.metadata
-        else:
-            model_meta = await get_model_meta(payload.model_id)
+        model_meta = await resolve_model_meta(payload.model_id)
 
         command = CreateInstanceCommand(
             command_id=CommandId(),
@@ -139,22 +142,12 @@ class API:
 
     # def get_instances_by_model(self, model_id: ModelId) -> list[Instance]: ...
 
-    async def _generate_chat_stream(self, payload: ChatCompletionTaskParams) -> AsyncGenerator[str, None]:
+    async def _generate_chat_stream(self, command_id: CommandId) -> AsyncGenerator[str, None]:
         """Generate chat completion stream as JSON strings."""
+
         events = await self.global_events.get_events_since(0)
         prev_idx = await self.global_events.get_last_idx()
 
-        # At the moment, we just create the task in the API.
-        # In the future, a `Request` will be created here and they will be bundled into `Task` objects by the master.
-        command_id=CommandId()
-
-        request = ChatCompletionCommand(
-            command_id=command_id,
-            command_type=CommandType.CHAT_COMPLETION,
-            request_params=payload,
-        )
-        self.command_buffer.append(request)
-
         finished = False
         while not finished:
             await asyncio.sleep(0.01)
@@ -177,10 +170,29 @@ class API:
 
         return
 
+    async def _trigger_notify_user_to_download_model(self, model_id: str) -> None:
+        print("TODO: we should send a notification to the user to download the model")
+
     async def chat_completions(self, payload: ChatCompletionTaskParams) -> StreamingResponse:
         """Handle chat completions with proper streaming response."""
+        model_meta = await resolve_model_meta(payload.model)
+        payload.model = model_meta.model_id
+
+        for instance in self.get_state().instances.values():
+            if instance.shard_assignments.model_id == payload.model:
+                break
+        else:
+            await self._trigger_notify_user_to_download_model(payload.model)
+            raise HTTPException(status_code=404, detail=f"No instance found for model {payload.model}")
+
+        command = ChatCompletionCommand(
+            command_id=CommandId(),
+            command_type=CommandType.CHAT_COMPLETION,
+            request_params=payload,
+        )
+        self.command_buffer.append(command)
         return StreamingResponse(
-            self._generate_chat_stream(payload),
+            self._generate_chat_stream(command.command_id),
             media_type="text/plain"
         )
 
@@ -195,4 +207,4 @@ def start_fastapi_server(
 ):
     api = API(command_buffer, global_events, get_state)
 
-    uvicorn.run(api.app, host=host, port=port)
\ No newline at end of file
+    uvicorn.run(api.app, host=host, port=port)
diff --git a/master/forwarder_supervisor.py b/master/forwarder_supervisor.py
index 93a0bab0..d00f4418 100644
--- a/master/forwarder_supervisor.py
+++ b/master/forwarder_supervisor.py
@@ -10,6 +10,7 @@ from shared.constants import (
     LIBP2P_GLOBAL_EVENTS_TOPIC,
     LIBP2P_WORKER_EVENTS_TOPIC,
 )
+from shared.types.common import NodeId
 
 
 class ForwarderRole(str, Enum):
@@ -35,10 +36,12 @@ class ForwarderSupervisor:
     
     def __init__(
         self, 
+        node_id: NodeId,
         forwarder_binary_path: Path,
         logger: Logger,
         health_check_interval: float = 5.0
     ):
+        self.node_id = node_id
         self._binary_path = forwarder_binary_path
         self._logger = logger
         self._health_check_interval = health_check_interval
@@ -108,6 +111,9 @@ class ForwarderSupervisor:
             f'{pairs}',
             stdout=None,
             stderr=None,
+            env={
+                "FORWARDER_NODE_ID": str(self.node_id),
+            }
         )
         
         self._logger.info(f"Starting forwarder with forwarding pairs: {pairs}")
diff --git a/master/main.py b/master/main.py
index 24868af7..6417b9c4 100644
--- a/master/main.py
+++ b/master/main.py
@@ -39,6 +39,7 @@ class Master:
     def __init__(self, node_id_keypair: Keypair, node_id: NodeId, command_buffer: list[Command],
                  global_events: AsyncSQLiteEventStorage, worker_events: AsyncSQLiteEventStorage,
                  forwarder_binary_path: Path, logger: logging.Logger):
+        self.state = State()
         self.node_id = node_id
         self.command_buffer = command_buffer
         self.global_events = global_events
@@ -46,16 +47,22 @@ class Master:
         self.discovery_supervisor = DiscoverySupervisor(
             node_id_keypair,
             node_id,
-            global_events,
+            # TODO: needs to be more general for when we have master election
+            worker_events if os.getenv('EXO_RUN_AS_REPLICA') in set(['TRUE', 'true', '1']) else global_events,
             logger
         )
         self.forwarder_supervisor = ForwarderSupervisor(
+            self.node_id,
             forwarder_binary_path=forwarder_binary_path,
             logger=logger
         )
         self.election_callbacks = ElectionCallbacks(self.forwarder_supervisor, logger)
         self.logger = logger
 
+    @property
+    def event_log_for_reads(self) -> AsyncSQLiteEventStorage:
+        return self.global_events
+
     @property
     def event_log_for_writes(self) -> AsyncSQLiteEventStorage:
         if self.forwarder_supervisor.current_role == ForwarderRole.MASTER:
@@ -89,8 +96,9 @@ class Master:
                     next_events.append(TaskCreated(
                         task_id=task_id,
                         task=ChatCompletionTask(
-                            task_id=task_id,
                             task_type=TaskType.CHAT_COMPLETION,
+                            task_id=task_id,
+                            command_id=next_command.command_id,
                             instance_id=matching_instance.instance_id,
                             task_status=TaskStatus.PENDING,
                             task_params=next_command.request_params
@@ -108,16 +116,18 @@ class Master:
             await self.event_log_for_writes.append_events(next_events, origin=self.node_id)
 
         # 2. get latest events
-        events = await self.global_events.get_events_since(self.state.last_event_applied_idx)
+        events = await self.event_log_for_reads.get_events_since(self.state.last_event_applied_idx)
         if len(events) == 0:
             await asyncio.sleep(0.01)
             return
+        self.logger.info(f"got events: {events}")
 
         # 3. for each event, apply it to the state
         for event_from_log in events:
+            print(f"applying event: {event_from_log}")
             self.state = apply(self.state, event_from_log)
 
-        self.logger.info(f"state: {self.state}")
+        self.logger.info(f"state: {self.state.model_dump_json()}")
 
     async def run(self):
         self.state = await self._get_state_snapshot()
diff --git a/master/placement.py b/master/placement.py
index 7137938f..cd3320cc 100644
--- a/master/placement.py
+++ b/master/placement.py
@@ -28,10 +28,10 @@ def get_instance_placements(
     
     all_nodes = list(topology.list_nodes())
     cycles = topology.get_cycles()
-    nodes_in_cycles = {node.node_id for cycle in cycles for node in cycle}
-    singleton_cycles = [[node] for node in all_nodes if node.node_id not in nodes_in_cycles]
+    # we can also always just have a node on its own
+    singleton_cycles = [[node] for node in all_nodes]
     candidate_cycles = cycles + singleton_cycles
-    cycles_with_sufficient_memory = filter_cycles_by_memory(candidate_cycles, command.model_meta.storage_size_kilobytes)
+    cycles_with_sufficient_memory = filter_cycles_by_memory(candidate_cycles, command.model_meta.storage_size_kilobytes * 1024)
     if not cycles_with_sufficient_memory:
         raise ValueError("No cycles found with sufficient memory")
 
diff --git a/master/tests/test_forwarder_manager.py b/master/tests/test_forwarder_manager.py
index 0160362b..c9413c52 100644
--- a/master/tests/test_forwarder_manager.py
+++ b/master/tests/test_forwarder_manager.py
@@ -24,6 +24,7 @@ from shared.constants import (
     LIBP2P_GLOBAL_EVENTS_TOPIC,
     LIBP2P_WORKER_EVENTS_TOPIC,
 )
+from shared.types.common import NodeId
 
 # Mock forwarder script content
 MOCK_FORWARDER_SCRIPT = '''#!/usr/bin/env python3
@@ -182,7 +183,7 @@ class TestForwardersupervisorBasic:
         # Set environment
         os.environ.update(mock_env_vars)
         
-        supervisor = ForwarderSupervisor(mock_forwarder_script, test_logger)
+        supervisor = ForwarderSupervisor(NodeId(), mock_forwarder_script, test_logger)
         await supervisor.start_as_replica()
         
         # Track the process for cleanup
@@ -224,7 +225,7 @@ class TestForwardersupervisorBasic:
         """Test changing role from replica to master."""
         os.environ.update(mock_env_vars)
         
-        supervisor = ForwarderSupervisor(mock_forwarder_script, test_logger)
+        supervisor = ForwarderSupervisor(NodeId(), mock_forwarder_script, test_logger)
         await supervisor.start_as_replica()
         
         if supervisor.process:
@@ -268,7 +269,7 @@ class TestForwardersupervisorBasic:
         """Test that setting the same role twice doesn't restart the process."""
         os.environ.update(mock_env_vars)
         
-        supervisor = ForwarderSupervisor(mock_forwarder_script, test_logger)
+        supervisor = ForwarderSupervisor(NodeId(), mock_forwarder_script, test_logger)
         await supervisor.start_as_replica()
         
         original_pid = supervisor.process_pid
@@ -300,6 +301,7 @@ class TestForwardersupervisorBasic:
         os.environ.update(mock_env_vars)
         
         supervisor = ForwarderSupervisor(
+            NodeId(),
             mock_forwarder_script,
             test_logger,
             health_check_interval=0.5  # Faster health checks for testing
@@ -346,7 +348,7 @@ class TestForwardersupervisorBasic:
         """Test behavior when forwarder binary doesn't exist."""
         nonexistent_path = temp_dir / "nonexistent_forwarder"
         
-        supervisor = ForwarderSupervisor(nonexistent_path, test_logger)
+        supervisor = ForwarderSupervisor(NodeId(), nonexistent_path, test_logger)
         
         # Should raise FileNotFoundError
         with pytest.raises(FileNotFoundError):
diff --git a/master/tests/test_master.py b/master/tests/test_master.py
index 4c4d23e4..767481e9 100644
--- a/master/tests/test_master.py
+++ b/master/tests/test_master.py
@@ -14,9 +14,27 @@ from shared.db.sqlite.event_log_manager import EventLogManager
 from shared.types.api import ChatCompletionMessage, ChatCompletionTaskParams
 from shared.types.common import NodeId
 from shared.types.events import TaskCreated
-from shared.types.events._events import TopologyNodeCreated
-from shared.types.events.commands import ChatCompletionCommand, Command, CommandId
+from shared.types.events._events import (
+    InstanceCreated,
+    NodePerformanceMeasured,
+    TopologyNodeCreated,
+)
+from shared.types.events.commands import (
+    ChatCompletionCommand,
+    Command,
+    CommandId,
+    CreateInstanceCommand,
+)
+from shared.types.models import ModelMetadata
+from shared.types.profiling import (
+    MemoryPerformanceProfile,
+    NodePerformanceProfile,
+    SystemPerformanceProfile,
+)
 from shared.types.tasks import ChatCompletionTask, TaskStatus, TaskType
+from shared.types.worker.common import InstanceId
+from shared.types.worker.instances import Instance, InstanceStatus, ShardAssignments
+from shared.types.worker.shards import PartitionStrategy, PipelineShardMetadata
 
 
 def _create_forwarder_dummy_binary() -> Path:
@@ -42,9 +60,40 @@ async def test_master():
     node_id_keypair = Keypair.generate_ed25519()
     node_id = NodeId(node_id_keypair.to_peer_id().to_base58())
     master = Master(node_id_keypair, node_id, command_buffer=command_buffer, global_events=global_events,
-                    forwarder_binary_path=forwarder_binary_path, logger=logger, worker_events=event_log_manager.worker_events)
+                    forwarder_binary_path=forwarder_binary_path, logger=logger, worker_events=global_events)
     asyncio.create_task(master.run())
+    # wait for initial topology event
+    while len(list(master.state.topology.list_nodes())) == 0:
+        print("waiting")
+        await asyncio.sleep(0.001)
+    # inject a NodePerformanceProfile event
+    await event_log_manager.global_events.append_events([
+        NodePerformanceMeasured(
+            node_id=node_id,
+            node_profile=NodePerformanceProfile(
+                model_id="maccy",
+                chip_id="arm",
+                memory=MemoryPerformanceProfile(ram_total=678948*1024, ram_available=678948*1024, swap_total=0, swap_available=0),
+                network_interfaces=[],
+                system=SystemPerformanceProfile(flops_fp16=0)
+            )
+        )
+    ], origin=node_id)
+    while len(master.state.node_profiles) == 0:
+        await asyncio.sleep(0.001)
 
+    command_buffer.append(CreateInstanceCommand(
+        command_id=CommandId(),
+        instance_id=InstanceId(),
+        model_meta=ModelMetadata(
+            model_id="llama-3.2-1b",
+            pretty_name="Llama 3.2 1B",
+            n_layers=16,
+            storage_size_kilobytes=678948
+        )
+    ))
+    while len(master.state.instances.keys()) == 0:
+        await asyncio.sleep(0.001)
     command_buffer.append(
         ChatCompletionCommand(
             command_id=CommandId(),
@@ -54,20 +103,52 @@ async def test_master():
             )
         )
     )
-    while len(await global_events.get_events_since(0)) == 0:
+    while len(await global_events.get_events_since(0)) < 4:
         await asyncio.sleep(0.001)
 
     events = await global_events.get_events_since(0)
-    assert len(events) == 2
+    print(events)
+    assert len(events) == 4
     assert events[0].idx_in_log == 1
     assert isinstance(events[0].event, TopologyNodeCreated)
-    assert isinstance(events[1].event, TaskCreated)
-    assert events[1].event == TaskCreated(
-        task_id=events[1].event.task_id,
+    assert isinstance(events[1].event, NodePerformanceMeasured)
+    assert isinstance(events[2].event, InstanceCreated)
+    runner_id = list(events[2].event.instance.shard_assignments.runner_to_shard.keys())[0]
+    assert events[2].event == InstanceCreated(
+        instance=Instance(
+            instance_id=events[2].event.instance.instance_id,
+            instance_type=InstanceStatus.ACTIVE,
+            shard_assignments=ShardAssignments(
+                model_id="llama-3.2-1b",
+                runner_to_shard={
+                    (runner_id): PipelineShardMetadata(
+                        partition_strategy=PartitionStrategy.pipeline,
+                        start_layer=0,
+                        end_layer=16,
+                        n_layers=16,
+                        model_meta=ModelMetadata(
+                            model_id="llama-3.2-1b",
+                            pretty_name="Llama 3.2 1B",
+                            n_layers=16,
+                            storage_size_kilobytes=678948
+                        ),
+                        device_rank=0,
+                        world_size=1
+                    )
+                },
+                node_to_runner={node_id: runner_id}
+            ),
+            hosts=[]
+        )
+    )
+    assert isinstance(events[3].event, TaskCreated)
+    assert events[3].event == TaskCreated(
+        task_id=events[3].event.task_id,
         task=ChatCompletionTask(
-            task_id=events[1].event.task_id,
+            task_id=events[3].event.task_id,
+            command_id=events[3].event.task.command_id,
             task_type=TaskType.CHAT_COMPLETION,
-            instance_id=events[1].event.task.instance_id,
+            instance_id=events[3].event.task.instance_id,
             task_status=TaskStatus.PENDING,
             task_params=ChatCompletionTaskParams(
                 model="llama-3.2-1b",
diff --git a/master/tests/test_placement.py b/master/tests/test_placement.py
index d51d16b1..1ab8a9ef 100644
--- a/master/tests/test_placement.py
+++ b/master/tests/test_placement.py
@@ -64,6 +64,8 @@ def test_get_instance_placements_create_instance(
     create_node: Callable[[int, NodeId | None], Node],
     create_connection: Callable[[NodeId, NodeId], Connection]
 ):
+    # TODO: this test is not exactly what we want. if a model can fit on one node, it should be placed there.
+    # TODO: right now we assume it will be placed across all nodes.
     # arrange
     model_meta.n_layers = total_layers
     
@@ -75,9 +77,9 @@ def test_get_instance_placements_create_instance(
     node_id_a = NodeId()
     node_id_b = NodeId()
     node_id_c = NodeId()
-    topology.add_node(create_node(available_memory[0], node_id_a))
-    topology.add_node(create_node(available_memory[1], node_id_b))
-    topology.add_node(create_node(available_memory[2], node_id_c))
+    topology.add_node(create_node(available_memory[0]*1024, node_id_a))
+    topology.add_node(create_node(available_memory[1]*1024, node_id_b))
+    topology.add_node(create_node(available_memory[2]*1024, node_id_c))
     topology.add_connection(create_connection(node_id_a, node_id_b))
     topology.add_connection(create_connection(node_id_b, node_id_c))
     topology.add_connection(create_connection(node_id_c, node_id_a))
@@ -113,7 +115,7 @@ def test_get_instance_placements_one_node_exact_fit(
 ) -> None:
     topology = Topology()
     node_id = NodeId()
-    topology.add_node(create_node(1000, node_id))
+    topology.add_node(create_node(1000*1024, node_id))
     create_instance_command = CreateInstanceCommand(
         command_id=CommandId(),
         model_meta=ModelMetadata(
@@ -139,7 +141,7 @@ def test_get_instance_placements_one_node_fits_with_extra_memory(
 ) -> None:
     topology = Topology()
     node_id = NodeId()
-    topology.add_node(create_node(1001, node_id))
+    topology.add_node(create_node(1001*1024, node_id))
     create_instance_command = CreateInstanceCommand(
         command_id=CommandId(),
         model_meta=ModelMetadata(
@@ -165,7 +167,7 @@ def test_get_instance_placements_one_node_not_fit(
 ) -> None:
     topology = Topology()
     node_id = NodeId()
-    topology.add_node(create_node(1000, node_id))
+    topology.add_node(create_node(1000*1024, node_id))
     create_instance_command = CreateInstanceCommand(
         command_id=CommandId(),
         model_meta=ModelMetadata(
diff --git a/master/tests/test_placement_utils.py b/master/tests/test_placement_utils.py
index 2ef84cd1..f9c286d9 100644
--- a/master/tests/test_placement_utils.py
+++ b/master/tests/test_placement_utils.py
@@ -24,8 +24,8 @@ def test_filter_cycles_by_memory(topology: Topology, create_node: Callable[[int,
     node1_id = NodeId()
     node2_id = NodeId()
 
-    node1 = create_node(1000, node1_id)
-    node2 = create_node(1000, node2_id)
+    node1 = create_node(1000*1024, node1_id)
+    node2 = create_node(1000*1024, node2_id)
     
     topology.add_node(node1)
     topology.add_node(node2)
@@ -52,8 +52,8 @@ def test_filter_cycles_by_insufficient_memory(topology: Topology, create_node: C
     node1_id = NodeId()
     node2_id = NodeId()
 
-    node1 = create_node(1000, node1_id)
-    node2 = create_node(1000, node2_id)
+    node1 = create_node(1000*1024, node1_id)
+    node2 = create_node(1000*1024, node2_id)
 
     topology.add_node(node1)
     topology.add_node(node2)
@@ -77,9 +77,9 @@ def test_filter_multiple_cycles_by_memory(topology: Topology, create_node: Calla
     node_b_id = NodeId()
     node_c_id = NodeId()
     
-    node_a = create_node(500, node_a_id)
-    node_b = create_node(500, node_b_id)
-    node_c = create_node(1000, node_c_id)
+    node_a = create_node(500*1024, node_a_id)
+    node_b = create_node(500*1024, node_b_id)
+    node_c = create_node(1000*1024, node_c_id)
     
     topology.add_node(node_a)
     topology.add_node(node_b)
@@ -107,9 +107,9 @@ def test_get_smallest_cycles(topology: Topology, create_node: Callable[[int, Nod
     node_b_id = NodeId()
     node_c_id = NodeId()
     
-    node_a = create_node(500, node_a_id)
-    node_b = create_node(500, node_b_id)
-    node_c = create_node(1000, node_c_id)
+    node_a = create_node(500*1024, node_a_id)
+    node_b = create_node(500*1024, node_b_id)
+    node_c = create_node(1000*1024, node_c_id)
 
     topology.add_node(node_a)
     topology.add_node(node_b)
@@ -139,9 +139,9 @@ def test_get_shard_assignments(topology: Topology, create_node: Callable[[int, N
     node_b_id = NodeId()
     node_c_id = NodeId()
     
-    node_a = create_node(available_memory[0], node_a_id)
-    node_b = create_node(available_memory[1], node_b_id)
-    node_c = create_node(available_memory[2], node_c_id)
+    node_a = create_node(available_memory[0]*1024, node_a_id)
+    node_b = create_node(available_memory[1]*1024, node_b_id)
+    node_c = create_node(available_memory[2]*1024, node_c_id)
 
     topology.add_node(node_a)
     topology.add_node(node_b)
diff --git a/rust/discovery/Cargo.toml b/rust/discovery/Cargo.toml
index 6ca9ef17..ff94a8be 100644
--- a/rust/discovery/Cargo.toml
+++ b/rust/discovery/Cargo.toml
@@ -33,6 +33,7 @@ thiserror = { workspace = true }
 #itertools = { workspace = true }
 tracing-subscriber = { version = "0.3.19", features = ["default", "env-filter"] }
 keccak-const = { workspace = true }
+log = "0.4"
 
 # Networking
 libp2p = { workspace = true, features = ["full"] }
\ No newline at end of file
diff --git a/rust/discovery/src/lib.rs b/rust/discovery/src/lib.rs
index 17cb78ca..bcc1075a 100644
--- a/rust/discovery/src/lib.rs
+++ b/rust/discovery/src/lib.rs
@@ -41,6 +41,8 @@ pub(crate) mod private {
 /// Create and configure a swarm, and start listening to all ports/OS.
 #[inline]
 pub fn discovery_swarm(keypair: identity::Keypair) -> alias::AnyResult<Swarm<DiscoveryBehaviour>> {
+    let peer_id = keypair.public().to_peer_id();
+    log::info!("RUST: Creating discovery swarm with peer_id: {}", peer_id);
     let mut swarm = SwarmBuilder::with_existing_identity(keypair)
         .with_tokio()
         .with_other_transport(discovery_transport)?
@@ -49,7 +51,9 @@ pub fn discovery_swarm(keypair: identity::Keypair) -> alias::AnyResult<Swarm<Dis
 
     // Listen on all interfaces and whatever port the OS assigns
     // swarm.listen_on("/ip4/0.0.0.0/udp/0/quic-v1".parse()?)?; // TODO: make this
-    swarm.listen_on("/ip4/0.0.0.0/tcp/0".parse()?)?;
+    let listen_addr = "/ip4/0.0.0.0/tcp/0".parse()?;
+    log::info!("RUST: Attempting to listen on: {}", listen_addr);
+    swarm.listen_on(listen_addr)?;
 
     Ok(swarm)
 }
diff --git a/rust/exo_pyo3_bindings/src/discovery.rs b/rust/exo_pyo3_bindings/src/discovery.rs
index 411c41b6..3ba8bbc6 100644
--- a/rust/exo_pyo3_bindings/src/discovery.rs
+++ b/rust/exo_pyo3_bindings/src/discovery.rs
@@ -120,11 +120,29 @@ async fn discovery_task(
                     Behaviour(Mdns(Discovered(list))) => {
                         for (peer_id, multiaddr) in list {
                             log::info!("RUST: mDNS discovered a new peer: {peer_id} on {multiaddr}");
-                            // TODO: this does the job of (actually) creating & maintaining connection
-                            //       but its coupled to gossipsub & also the connection isn't configured
-                            //       for setting "connection keep alive" in NetworkBehavior's ConnectionHandler
-                            //       >in future, make own small NetworkBehavior impl just to track this state
+                            let local_peer_id = *swarm.local_peer_id();
+                            // To avoid simultaneous dial races, only the lexicographically larger peer_id dials.
+                            if peer_id > local_peer_id {
+                                let dial_opts = DialOpts::peer_id(peer_id)
+                                    .addresses(vec![multiaddr.clone()].into())
+                                    .condition(libp2p::swarm::dial_opts::PeerCondition::Always)
+                                    .build();
+                                match swarm.dial(dial_opts) {
+                                    Ok(()) => log::info!("RUST: Dial initiated to {multiaddr}"),
+                                    Err(libp2p::swarm::DialError::DialPeerConditionFalse(_)) => {
+                                        // Another dial is already in progress; not an error for us.
+                                        log::debug!(
+                                            "RUST: Dial skipped because another dial is active for {peer_id}"
+                                        );
+                                    }
+                                    Err(e) => {
+                                        log::warn!("RUST: Failed to dial {multiaddr}: {e:?}");
+                                    }
+                                }
+                            }
+                            // Maintain peer in gossipsub mesh so the connection stays alive once established.
                             swarm.behaviour_mut().gossipsub.add_explicit_peer(&peer_id);
+                            log::info!("RUST: Added peer {peer_id} to gossipsub explicit peers");
                         }
                     }
                     Behaviour(Mdns(Expired(list))) => {
@@ -149,6 +167,7 @@ async fn discovery_task(
                         concurrent_dial_errors,
                         established_in: _established_in,
                     } => {
+                        log::info!("RUST: ConnectionEstablished event - peer_id: {peer_id}, connection_id: {connection_id:?}, endpoint: {endpoint:?}");
                         // log any connection errors
                         if let Some(concurrent_dial_errors) = concurrent_dial_errors {
                             for (multiaddr, error) in concurrent_dial_errors {
@@ -156,17 +175,21 @@ async fn discovery_task(
                             }
                         }
 
-                        // TODO: right now we assume we are using TCP/IP which treats all nodes
-                        //       as both dialers AND listeners. This means for each connection you will actually
-                        //       see TWO duplicate Connected events => Dialer & Listener
-                        //       SO ignore the Dialer & extract the info we need from Listener
-                        //       HOWEVER this makes the swarm implicitly rely on TCP/IP, so is brittle to changes
-                        //       e.g. adding QUIC protocol or something
-                        //       >As soon as we add anything other than TCP/IP, this must be updated or there will be broken code
-                        let ConnectedPoint::Listener { local_addr, send_back_addr } = endpoint else {
-                            log::warn!("Ignoring `ConnectedPoint::Dialer` event because for TCP/IP it has a dual `ConnectedPoint::Listener` event: {endpoint:?}");
-                            continue;
+                        // Extract addresses based on endpoint type
+                        let (local_addr, send_back_addr) = match &endpoint {
+                            ConnectedPoint::Listener { local_addr, send_back_addr } => {
+                                log::info!("RUST: Connection established (Listener) - local_addr: {local_addr}, send_back_addr: {send_back_addr}");
+                                (local_addr.clone(), send_back_addr.clone())
+                            },
+                            ConnectedPoint::Dialer { address, .. } => {
+                                log::info!("RUST: Connection established (Dialer) - remote_addr: {address}");
+                                // For dialer, we use the dialed address as both local and send_back
+                                // This isn't perfect but allows both sides to be notified
+                                (address.clone(), address.clone())
+                            }
                         };
+                        
+                        log::info!("RUST: Number of connected callbacks: {}", connected_callbacks.len());
 
 
                         // trigger callback on connected peer
@@ -180,22 +203,27 @@ async fn discovery_task(
                         }
                     },
                     ConnectionClosed { peer_id, connection_id, endpoint, num_established, cause }  => {
+                        log::info!("RUST: ConnectionClosed event - peer_id: {peer_id}, connection_id: {connection_id:?}, endpoint: {endpoint:?}, num_established: {num_established}");
                         // log any connection errors
                         if let Some(cause) = cause {
                             log::error!("Connection error: cause={cause:?}");
                         }
 
-                        // TODO: right now we assume we are using TCP/IP which treats all nodes
-                        //       as both dialers AND listeners. This means for each connection you will actually
-                        //       see TWO duplicate Connected events => Dialer & Listener
-                        //       SO ignore the Dialer & extract the info we need from Listener
-                        //       HOWEVER this makes the swarm implicitly rely on TCP/IP, so is brittle to changes
-                        //       e.g. adding QUIC protocol or something
-                        //       >As soon as we add anything other than TCP/IP, this must be updated or there will be broken code
-                        let ConnectedPoint::Listener { local_addr, send_back_addr } = endpoint else {
-                            log::warn!("Ignoring `ConnectedPoint::Dialer` event because for TCP/IP it has a dual `ConnectedPoint::Listener` event: {endpoint:?}");
-                            continue;
+                        // Extract addresses based on endpoint type
+                        let (local_addr, send_back_addr) = match &endpoint {
+                            ConnectedPoint::Listener { local_addr, send_back_addr } => {
+                                log::info!("RUST: Connection closed (Listener) - local_addr: {local_addr}, send_back_addr: {send_back_addr}");
+                                (local_addr.clone(), send_back_addr.clone())
+                            },
+                            ConnectedPoint::Dialer { address, .. } => {
+                                log::info!("RUST: Connection closed (Dialer) - remote_addr: {address}");
+                                // For dialer, we use the dialed address as both local and send_back
+                                // This isn't perfect but allows both sides to be notified
+                                (address.clone(), address.clone())
+                            }
                         };
+                        
+                        log::info!("RUST: Number of disconnected callbacks: {}", disconnected_callbacks.len());
 
                         // trigger callback on connected peer
                         for disconnected_callback in &disconnected_callbacks {
@@ -207,8 +235,13 @@ async fn discovery_task(
                             });
                         }
                     }
+                    NewListenAddr { address, .. } => {
+                        log::info!("RUST: Local node is listening on {address}");
+                        let local_peer = swarm.local_peer_id();
+                        log::info!("RUST: Local peer_id: {local_peer}");
+                    }
                     e => {
-                        log::info!("RUST: Other event {e:?}");
+                        log::debug!("RUST: Other event {e:?}");
                     }
                 }
             }
@@ -258,15 +291,19 @@ impl PyDiscoveryService {
 
         // get identity
         let identity = identity.borrow().0.clone();
+        log::info!("RUST: Creating DiscoveryService with keypair");
 
         // create discovery swarm (within tokio context!! or it crashes)
         let swarm = get_runtime()
             .block_on(async { discovery_swarm(identity) })
             .pyerr()?;
+        log::info!("RUST: Discovery swarm created successfully");
 
         // spawn tokio task
         get_runtime().spawn(async move {
+            log::info!("RUST: Starting discovery task");
             discovery_task(receiver, swarm).await;
+            log::info!("RUST: Discovery task ended");
         });
         Ok(Self::new(sender))
     }
diff --git a/shared/apply/apply.py b/shared/apply/apply.py
index b5f49538..1386a475 100644
--- a/shared/apply/apply.py
+++ b/shared/apply/apply.py
@@ -114,6 +114,9 @@ def apply_runner_deleted(event: RunnerDeleted, state: State) -> State:
 def apply_node_performance_measured(event: NodePerformanceMeasured, state: State) -> State:
     new_profiles: Mapping[NodeId, NodePerformanceProfile] = {**state.node_profiles, event.node_id: event.node_profile}
     state = state.model_copy(update={"node_profiles": new_profiles})
+    if not state.topology.contains_node(event.node_id):
+        # TODO: figure out why this is happening in the first place
+        return state
     topology = copy.copy(state.topology)
     topology.update_node_profile(event.node_id, event.node_profile)
     return state.model_copy(update={"topology": topology})
@@ -148,5 +151,7 @@ def apply_topology_edge_replaced_atomically(event: TopologyEdgeReplacedAtomicall
 @event_apply.register(TopologyEdgeDeleted)
 def apply_topology_edge_deleted(event: TopologyEdgeDeleted, state: State) -> State:
     topology = copy.copy(state.topology)
+    if not topology.contains_connection(event.edge):
+        return state
     topology.remove_connection(event.edge)
     return state.model_copy(update={"topology": topology})
\ No newline at end of file
diff --git a/shared/topology.py b/shared/topology.py
index d007e532..2263c447 100644
--- a/shared/topology.py
+++ b/shared/topology.py
@@ -63,6 +63,12 @@ class Topology(TopologyProto):
         self._node_id_to_rx_id_map[node.node_id] = rx_id
         self._rx_id_to_node_id_map[rx_id] = node.node_id
 
+    def contains_node(self, node_id: NodeId) -> bool:
+        return node_id in self._node_id_to_rx_id_map
+
+    def contains_connection(self, connection: Connection) -> bool:
+        return connection in self._edge_id_to_rx_id_map
+
     def add_connection(
             self,
             connection: Connection,
@@ -120,7 +126,8 @@ class Topology(TopologyProto):
         else:
             self._graph.remove_edge_from_index(rx_idx)
             del self._edge_id_to_rx_id_map[connection]
-            del self._rx_id_to_node_id_map[rx_idx]
+            if rx_idx in self._rx_id_to_node_id_map:
+                del self._rx_id_to_node_id_map[rx_idx]
 
     def get_cycles(self) -> list[list[Node]]:
         cycle_idxs = rx.simple_cycles(self._graph)
diff --git a/shared/types/tasks.py b/shared/types/tasks.py
index 12b0b514..00426ba9 100644
--- a/shared/types/tasks.py
+++ b/shared/types/tasks.py
@@ -4,7 +4,7 @@ from typing import Annotated, Literal
 from pydantic import BaseModel, Field
 
 from shared.types.api import ChatCompletionTaskParams
-from shared.types.common import ID
+from shared.types.common import ID, CommandId
 from shared.types.worker.common import InstanceId
 
 
@@ -26,6 +26,7 @@ class TaskStatus(str, Enum):
 class ChatCompletionTask(BaseModel):
     task_type: Literal[TaskType.CHAT_COMPLETION] = TaskType.CHAT_COMPLETION
     task_id: TaskId
+    command_id: CommandId
     instance_id: InstanceId
     task_status: TaskStatus
     task_params: ChatCompletionTaskParams
diff --git a/worker/download/download_utils.py b/worker/download/download_utils.py
index cde8f056..a5615163 100644
--- a/worker/download/download_utils.py
+++ b/worker/download/download_utils.py
@@ -198,17 +198,22 @@ async def calc_hash(path: Path, hash_type: Literal["sha1", "sha256"] = "sha1") -
       hasher.update(chunk)
   return hasher.hexdigest()
 
-async def file_meta(repo_id: str, revision: str, path: str) -> Tuple[int, str]:
-  url = urljoin(f"{get_hf_endpoint()}/{repo_id}/resolve/{revision}/", path)
+async def file_meta(repo_id: str, revision: str, path: str, redirected_location: str | None = None) -> Tuple[int, str]:
+  # NOTE: huggingface broke the E-Tag so we can no longer assume E-Tag == sha256(file)
+  url = urljoin(f"{get_hf_endpoint()}/{repo_id}/resolve/{revision}/", path) if redirected_location is None else f"{get_hf_endpoint()}{redirected_location}"
   headers = await get_auth_headers()
   async with aiohttp.ClientSession(timeout=aiohttp.ClientTimeout(total=1800, connect=60, sock_read=1800, sock_connect=60)) as session, session.head(url, headers=headers) as r:
-      content_length = int(r.headers.get('x-linked-size') or r.headers.get('content-length') or 0)
-      etag = r.headers.get('X-Linked-ETag') or r.headers.get('ETag') or r.headers.get('Etag')
-      assert content_length > 0, f"No content length for {url}"
-      assert etag is not None, f"No remote hash for {url}"
-      if  (etag[0] == '"' and etag[-1] == '"') or (etag[0] == "'" and etag[-1] == "'"):
-        etag = etag[1:-1]
-      return content_length, etag
+    if r.status == 307:
+      redirected_location = r.headers.get('Location')
+      return await file_meta(repo_id, revision, path, redirected_location)
+
+    content_length = int(r.headers.get('x-linked-size') or r.headers.get('content-length') or 0)
+    etag = r.headers.get('X-Linked-ETag') or r.headers.get('ETag') or r.headers.get('Etag')
+    assert content_length > 0, f"No content length for {url}"
+    assert etag is not None, f"No remote hash for {url}"
+    if  (etag[0] == '"' and etag[-1] == '"') or (etag[0] == "'" and etag[-1] == "'"):
+      etag = etag[1:-1]
+    return content_length, etag
 
 async def download_file_with_retry(repo_id: str, revision: str, path: str, target_dir: Path, on_progress: Callable[[int, int], None] = lambda _, __: None) -> Path:
   n_attempts = 30
@@ -243,7 +248,8 @@ async def _download_file(repo_id: str, revision: str, path: str, target_dir: Pat
         assert r.status in [200, 206], f"Failed to download {path} from {url}: {r.status}"
         async with aiofiles.open(partial_path, 'ab' if resume_byte_pos else 'wb') as f:
           while chunk := await r.content.read(8 * 1024 * 1024):
-            on_progress(n_read := n_read + await f.write(chunk), length)
+            n_read = n_read + (await f.write(chunk))
+            on_progress(n_read, length)
 
   final_hash = await calc_hash(partial_path, hash_type="sha256" if len(remote_hash) == 64 else "sha1")
   integrity = final_hash == remote_hash
diff --git a/worker/download/impl_shard_downloader.py b/worker/download/impl_shard_downloader.py
index 3843107e..dff56912 100644
--- a/worker/download/impl_shard_downloader.py
+++ b/worker/download/impl_shard_downloader.py
@@ -64,6 +64,9 @@ class SingletonShardDownloader(ShardDownloader):
     async for path, status in self.shard_downloader.get_shard_download_status():
       yield path, status
 
+  async def get_shard_download_status_for_shard(self, shard: ShardMetadata) -> RepoDownloadProgress:
+    return await self.shard_downloader.get_shard_download_status_for_shard(shard)
+
 class CachedShardDownloader(ShardDownloader):
   def __init__(self, shard_downloader: ShardDownloader):
     self.shard_downloader = shard_downloader
@@ -86,6 +89,9 @@ class CachedShardDownloader(ShardDownloader):
     async for path, status in self.shard_downloader.get_shard_download_status():
       yield path, status
 
+  async def get_shard_download_status_for_shard(self, shard: ShardMetadata) -> RepoDownloadProgress:
+    return await self.shard_downloader.get_shard_download_status_for_shard(shard)
+
 class ResumableShardDownloader(ShardDownloader):
   def __init__(self, max_parallel_downloads: int = 8):
     self.max_parallel_downloads = max_parallel_downloads
@@ -126,3 +132,7 @@ class ResumableShardDownloader(ShardDownloader):
         yield (path, progress)
       except Exception as e:
         print("Error downloading shard:", e)
+
+  async def get_shard_download_status_for_shard(self, shard: ShardMetadata) -> RepoDownloadProgress:
+    _, progress = await download_shard(shard, self.on_progress_wrapper, skip_download=True)
+    return progress
diff --git a/worker/download/shard_downloader.py b/worker/download/shard_downloader.py
index 68a095c7..27b88411 100644
--- a/worker/download/shard_downloader.py
+++ b/worker/download/shard_downloader.py
@@ -68,6 +68,10 @@ class ShardDownloader(ABC):
         )
     )
 
+  @abstractmethod
+  async def get_shard_download_status_for_shard(self, shard: ShardMetadata) -> RepoDownloadProgress:
+    ...
+
 
 class NoopShardDownloader(ShardDownloader):
   async def ensure_shard(self, shard: ShardMetadata, config_only: bool = False) -> Path:
@@ -106,3 +110,18 @@ class NoopShardDownloader(ShardDownloader):
             status="complete",
         )
     )
+
+  async def get_shard_download_status_for_shard(self, shard: ShardMetadata) -> RepoDownloadProgress:
+    return RepoDownloadProgress(
+        repo_id="noop",
+        repo_revision="noop",
+        shard=shard,
+        completed_files=0,
+        total_files=0,
+        downloaded_bytes=0,
+        downloaded_bytes_this_session=0,
+        total_bytes=0,
+        overall_speed=0,
+        overall_eta=timedelta(seconds=0),
+        status="complete",
+    )
\ No newline at end of file
diff --git a/worker/main.py b/worker/main.py
index 987da047..1275a3e6 100644
--- a/worker/main.py
+++ b/worker/main.py
@@ -1,8 +1,9 @@
 import asyncio
 import logging
-import os
 from asyncio import Queue
+from copy import deepcopy
 from functools import partial
+from time import process_time
 from typing import AsyncGenerator, Optional
 
 from pydantic import BaseModel, ConfigDict
@@ -54,7 +55,8 @@ from shared.types.worker.runners import (
 )
 from shared.types.worker.shards import ShardMetadata
 from shared.utils import get_node_id_keypair
-from worker.download.download_utils import build_model_path
+from worker.download.impl_shard_downloader import exo_shard_downloader
+from worker.download.shard_downloader import RepoDownloadProgress, ShardDownloader
 from worker.runner.runner_supervisor import RunnerSupervisor
 from worker.utils.profile import start_polling_node_metrics
 
@@ -70,15 +72,15 @@ class AssignedRunner(BaseModel):
 
     model_config = ConfigDict(arbitrary_types_allowed=True)
 
-    @property
-    def is_downloaded(self) -> bool:
-        # TODO: Do this properly with huggingface validating each of the files.
-        return os.path.exists(build_model_path(self.shard_metadata.model_meta.model_id))
-
+    is_downloaded: bool = False
+    
+    def set_is_downloaded(self, is_downloaded: bool) -> None:
+        self.is_downloaded = is_downloaded
+    
     def status_update_event(self) -> RunnerStatusUpdated:
         return RunnerStatusUpdated(
             runner_id=self.runner_id,
-            runner_status=self.status,
+            runner_status=deepcopy(self.status),
         )
 
 class Worker:
@@ -86,11 +88,13 @@ class Worker:
         self,
         node_id: NodeId,
         logger: logging.Logger,
+        shard_downloader: ShardDownloader,
         worker_events: AsyncSQLiteEventStorage | None,
         global_events: AsyncSQLiteEventStorage | None,
     ):
         self.node_id: NodeId = node_id
         self.state: State = State()
+        self.shard_downloader: ShardDownloader = shard_downloader
         self.worker_events: AsyncSQLiteEventStorage | None = worker_events # worker_events is None in some tests.
         self.global_events: AsyncSQLiteEventStorage | None = global_events
         self.logger: logging.Logger = logger
@@ -183,12 +187,26 @@ class Worker:
         The model needs assigning and then downloading.
         This op moves the runner from Assigned -> Downloading -> Ready state.
         '''
+        
+        initial_progress = await self.shard_downloader.get_shard_download_status_for_shard(op.shard_metadata)
+        if initial_progress.status == "complete":
+            self.assigned_runners[op.runner_id].set_is_downloaded(True)
+            self.assigned_runners[op.runner_id].status = DownloadingRunnerStatus(
+                download_progress=DownloadCompleted(
+                    node_id=self.node_id,
+                )
+            )
+            yield self.assigned_runners[op.runner_id].status_update_event()
+            self.assigned_runners[op.runner_id].status = ReadyRunnerStatus()
+            yield self.assigned_runners[op.runner_id].status_update_event()
+            return
+        
         initial_status = DownloadingRunnerStatus(
             download_progress=DownloadOngoing(
                 node_id=self.node_id,
                 download_progress=DownloadProgressData(
-                    total_bytes=1, # tmp
-                    downloaded_bytes=0
+                    total_bytes=initial_progress.total_bytes,
+                    downloaded_bytes=initial_progress.downloaded_bytes
                 )
             )
         )
@@ -206,25 +224,53 @@ class Worker:
 
         # Download it!
         # TODO: we probably want download progress as part of a callback that gets passed to the downloader.
+        download_progress_queue: asyncio.Queue[RepoDownloadProgress] = asyncio.Queue()
+        def download_progress_callback(shard: ShardMetadata, progress: RepoDownloadProgress) -> None:
+            download_progress_queue.put_nowait(progress)
 
-        try:
-            assert assigned_runner.is_downloaded
-            assigned_runner.status = DownloadingRunnerStatus(
-                download_progress=DownloadCompleted(
-                    node_id=self.node_id,
+
+        self.shard_downloader.on_progress(download_progress_callback)
+
+        asyncio.create_task(self.shard_downloader.ensure_shard(op.shard_metadata))
+
+        timeout_secs = 10 * 60
+        start_time = process_time()
+        last_yield_progress = start_time
+        while process_time() - start_time < timeout_secs:
+            progress: RepoDownloadProgress = await download_progress_queue.get()
+            if progress.status == "complete":
+                assigned_runner.status = DownloadingRunnerStatus(
+                    download_progress=DownloadCompleted(
+                        node_id=self.node_id,
+                    )
                 )
-            )
-        except Exception as e:
+                yield assigned_runner.status_update_event()
+                assigned_runner.set_is_downloaded(True)
+                assigned_runner.status = ReadyRunnerStatus()
+                yield assigned_runner.status_update_event()
+                break
+            elif progress.status == "in_progress":
+                if process_time() - last_yield_progress > 1:
+                    assigned_runner.status = DownloadingRunnerStatus(
+                        download_progress=DownloadOngoing(
+                            node_id=self.node_id,
+                            download_progress=DownloadProgressData(
+                                total_bytes=progress.total_bytes,
+                                downloaded_bytes=progress.downloaded_bytes,
+                            )
+                        )
+                    )
+                    yield assigned_runner.status_update_event()
+                    last_yield_progress = process_time()
+        else:
             assigned_runner.status = DownloadingRunnerStatus(
                 download_progress=DownloadFailed(
                     node_id=self.node_id,
-                    error_message=str(e)
+                    error_message=f"Timeout downloading model: {op.shard_metadata.model_meta.model_id}"
                 )
             )
-        yield assigned_runner.status_update_event()
+            yield assigned_runner.status_update_event()
 
-        assigned_runner.status = ReadyRunnerStatus()
-        yield assigned_runner.status_update_event()
 
     async def _execute_task_op(
         self, op: ExecuteTaskOp
@@ -383,6 +429,7 @@ class Worker:
         for _instance_id, instance in state.instances.items():
             if self.node_id in instance.shard_assignments.node_to_runner and \
                 instance.shard_assignments.node_to_runner[self.node_id] in state.runners and \
+                instance.shard_assignments.node_to_runner[self.node_id] in self.assigned_runners and \
                 isinstance(self.assigned_runners[instance.shard_assignments.node_to_runner[self.node_id]].status, FailedRunnerStatus):
 
                 num_spundown_nodes = 0
@@ -484,6 +531,7 @@ class Worker:
     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)
+        print(f"published event: {event}")
 
     # Handle state updates
     async def run(self):
@@ -500,6 +548,8 @@ class Worker:
 
             # 3. based on the updated state, we plan & execute an operation.
             op: RunnerOp | None = self.plan(self.state)
+            if op is not None:
+                self.logger.info(f"!!! plan result: {op}")
 
             # run the op, synchronously blocking for now
             if op is not None:
@@ -507,7 +557,8 @@ class Worker:
                     await self.event_publisher(event)
 
             await asyncio.sleep(0.01)
-            self.logger.info(f"state: {self.state}")
+            if len(events) > 0:
+                self.logger.info(f"state: {self.state}")
 
 
 async def main():
@@ -522,6 +573,7 @@ async def main():
 
     event_log_manager = EventLogManager(EventLogConfig(), logger)
     await event_log_manager.initialize()
+    shard_downloader = exo_shard_downloader()
     
     # TODO: add profiling etc to resource monitor
     async def resource_monitor_callback(node_performance_profile: NodePerformanceProfile) -> None:
@@ -530,7 +582,7 @@ async def main():
         )
     asyncio.create_task(start_polling_node_metrics(callback=resource_monitor_callback))
 
-    worker = Worker(node_id, logger, event_log_manager.worker_events, event_log_manager.global_events)
+    worker = Worker(node_id, logger, shard_downloader, event_log_manager.worker_events, event_log_manager.global_events)
 
     await worker.run()
 
diff --git a/worker/runner/runner_supervisor.py b/worker/runner/runner_supervisor.py
index 54d380d2..43b515dc 100644
--- a/worker/runner/runner_supervisor.py
+++ b/worker/runner/runner_supervisor.py
@@ -197,7 +197,7 @@ class RunnerSupervisor:
                         text=text, token=token, finish_reason=finish_reason
                     ):
                         yield TokenChunk(
-                            command_id=CommandId(task.task_id),
+                            command_id=CommandId(task.command_id),
                             idx=token,
                             model=self.model_shard_meta.model_meta.model_id,
                             text=text,
diff --git a/worker/tests/conftest.py b/worker/tests/conftest.py
index ad76fdab..9ef65c3d 100644
--- a/worker/tests/conftest.py
+++ b/worker/tests/conftest.py
@@ -9,7 +9,7 @@ from shared.db.sqlite.connector import AsyncSQLiteEventStorage
 from shared.db.sqlite.event_log_manager import EventLogConfig, EventLogManager
 from shared.models.model_meta import get_model_meta
 from shared.types.api import ChatCompletionMessage, ChatCompletionTaskParams
-from shared.types.common import NodeId
+from shared.types.common import CommandId, NodeId
 from shared.types.models import ModelId, ModelMetadata
 from shared.types.state import State
 from shared.types.tasks import (
@@ -27,6 +27,7 @@ from shared.types.worker.ops import (
 )
 from shared.types.worker.runners import RunnerId, ShardAssignments
 from shared.types.worker.shards import PipelineShardMetadata
+from worker.download.shard_downloader import NoopShardDownloader
 from worker.main import Worker
 
 
@@ -103,6 +104,7 @@ def chat_completion_task(completion_create_params: ChatCompletionTaskParams):
     def _chat_completion_task(instance_id: InstanceId) -> ChatCompletionTask:
         return ChatCompletionTask(
             task_id=TaskId(),
+            command_id=CommandId(),
             instance_id=instance_id,
             task_type=TaskType.CHAT_COMPLETION,
             task_status=TaskStatus.PENDING,
@@ -153,9 +155,10 @@ def instance(pipeline_shard_meta: Callable[[int, int], PipelineShardMetadata], h
 @pytest.fixture
 async def worker(node_id: NodeId, logger: Logger):
     event_log_manager = EventLogManager(EventLogConfig(), logger)
+    shard_downloader = NoopShardDownloader()
     await event_log_manager.initialize()
 
-    return Worker(node_id, logger, worker_events=event_log_manager.global_events, global_events=event_log_manager.global_events)
+    return Worker(node_id, logger, shard_downloader, worker_events=event_log_manager.global_events, global_events=event_log_manager.global_events)
 
 @pytest.fixture
 async def worker_with_assigned_runner(worker: Worker, instance: Callable[[InstanceId, NodeId, RunnerId], Instance]):
@@ -204,7 +207,8 @@ def worker_running(logger: Logger) -> Callable[[NodeId], Awaitable[tuple[Worker,
         global_events = event_log_manager.global_events
         await global_events.delete_all_events()
 
-        worker = Worker(node_id, logger=logger, worker_events=global_events, global_events=global_events)
+        shard_downloader = NoopShardDownloader()
+        worker = Worker(node_id, logger=logger, shard_downloader=shard_downloader, worker_events=global_events, global_events=global_events)
         asyncio.create_task(worker.run())
 
         return worker, global_events
diff --git a/worker/tests/test_worker_integration.py b/worker/tests/test_worker_integration.py
index 3041080c..acd28735 100644
--- a/worker/tests/test_worker_integration.py
+++ b/worker/tests/test_worker_integration.py
@@ -32,6 +32,7 @@ from shared.types.worker.runners import (
     # RunningRunnerStatus,
 )
 from shared.types.worker.shards import PipelineShardMetadata
+from worker.download.shard_downloader import NoopShardDownloader
 from worker.main import AssignedRunner, Worker
 from worker.tests.test_worker_integration_utils import read_streaming_response
 
@@ -269,14 +270,15 @@ async def test_2_runner_inference(
     ):
     event_log_manager = EventLogManager(EventLogConfig(), logger)
     await event_log_manager.initialize()
+    shard_downloader = NoopShardDownloader()
 
     global_events = event_log_manager.global_events
     await global_events.delete_all_events()
 
-    worker1 = Worker(NODE_A, logger=logger, worker_events=global_events, global_events=global_events)
+    worker1 = Worker(NODE_A, logger=logger, shard_downloader=shard_downloader, worker_events=global_events, global_events=global_events)
     asyncio.create_task(worker1.run())
 
-    worker2 = Worker(NODE_B, logger=logger, worker_events=global_events, global_events=global_events)
+    worker2 = Worker(NODE_B, logger=logger, shard_downloader=shard_downloader, worker_events=global_events, global_events=global_events)
     asyncio.create_task(worker2.run())
 
     ## Instance
@@ -348,14 +350,15 @@ async def test_runner_respawn(
     ):
     event_log_manager = EventLogManager(EventLogConfig(), logger)
     await event_log_manager.initialize()
+    shard_downloader = NoopShardDownloader()
 
     global_events = event_log_manager.global_events
     await global_events.delete_all_events()
 
-    worker1 = Worker(NODE_A, logger=logger, worker_events=global_events, global_events=global_events)
+    worker1 = Worker(NODE_A, logger=logger, shard_downloader=shard_downloader, worker_events=global_events, global_events=global_events)
     asyncio.create_task(worker1.run())
 
-    worker2 = Worker(NODE_B, logger=logger, worker_events=global_events, global_events=global_events)
+    worker2 = Worker(NODE_B, logger=logger, shard_downloader=shard_downloader, worker_events=global_events, global_events=global_events)
     asyncio.create_task(worker2.run())
 
     ## Instance
diff --git a/worker/tests/test_worker_plan.py b/worker/tests/test_worker_plan.py
index 120e3895..040d47ee 100644
--- a/worker/tests/test_worker_plan.py
+++ b/worker/tests/test_worker_plan.py
@@ -36,9 +36,11 @@ from shared.types.worker.runners import (
 )
 from shared.types.worker.shards import PipelineShardMetadata
 from worker.download.download_utils import build_model_path
-from worker.main import Worker
+from worker.download.shard_downloader import NoopShardDownloader
+from worker.main import AssignedRunner, Worker
 
 from .test_worker_plan_utils import (
+    COMMAND_1_ID,
     INSTANCE_1_ID,
     MODEL_A_ID,
     NODE_A,
@@ -47,7 +49,6 @@ from .test_worker_plan_utils import (
     RUNNER_2_ID,
     TASK_1_ID,
     InProcessRunner,
-    OverrideAssignedRunner,
     PlanTestCase,
     make_downloading_status,
     make_model_meta,
@@ -339,7 +340,7 @@ def _get_test_cases(tmp_path: Path) -> list[PlanTestCase]:
                     )
                 },
                 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)},
+                tasks={TASK_1_ID: ChatCompletionTask(task_id=TASK_1_ID, command_id=COMMAND_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
         ),
@@ -382,7 +383,7 @@ def _get_test_cases(tmp_path: Path) -> list[PlanTestCase]:
                     )
                 },
                 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)},
+                tasks={TASK_1_ID: ChatCompletionTask(task_id=TASK_1_ID, command_id=COMMAND_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)
         ),
@@ -484,6 +485,7 @@ def _get_test_cases(tmp_path: Path) -> list[PlanTestCase]:
                 tasks={
                     TASK_1_ID: ChatCompletionTask(
                         task_id=TASK_1_ID,
+                        command_id=COMMAND_1_ID,
                         task_type=TaskType.CHAT_COMPLETION,
                         task_status=TaskStatus.PENDING,
                         task_params=ChatCompletionTaskParams(
@@ -501,6 +503,7 @@ def _get_test_cases(tmp_path: Path) -> list[PlanTestCase]:
             ),
             expected_op=ExecuteTaskOp(runner_id=RUNNER_1_ID, task=ChatCompletionTask(
                 task_id=TASK_1_ID,
+                command_id=COMMAND_1_ID,
                 instance_id=INSTANCE_1_ID,
                 task_type=TaskType.CHAT_COMPLETION,
                 task_status=TaskStatus.PENDING,
@@ -550,7 +553,7 @@ def _get_test_cases(tmp_path: Path) -> list[PlanTestCase]:
                     )
                 },
                 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)},
+                tasks={TASK_1_ID: ChatCompletionTask(task_id=TASK_1_ID, command_id=COMMAND_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
         ),
@@ -593,12 +596,13 @@ def _get_test_cases(tmp_path: Path) -> list[PlanTestCase]:
                     )
                 },
                 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)},
+                tasks={TASK_1_ID: ChatCompletionTask(task_id=TASK_1_ID, command_id=COMMAND_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,
+                    command_id=COMMAND_1_ID,
                     instance_id=INSTANCE_1_ID,
                     task_type=TaskType.CHAT_COMPLETION,
                     task_params=ChatCompletionTaskParams(
@@ -648,12 +652,13 @@ def _get_test_cases(tmp_path: Path) -> list[PlanTestCase]:
                     )
                 },
                 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)},
+                tasks={TASK_1_ID: ChatCompletionTask(task_id=TASK_1_ID, command_id=COMMAND_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,
+                    command_id=COMMAND_1_ID,
                     instance_id=INSTANCE_1_ID,
                     task_type=TaskType.CHAT_COMPLETION,
                     task_params=ChatCompletionTaskParams(
@@ -851,7 +856,8 @@ def test_worker_plan(case: PlanTestCase, tmp_path: Path, monkeypatch: pytest.Mon
     node_id = NODE_A
 
     logger = logging.getLogger("test_worker_plan")
-    worker = Worker(node_id=node_id, worker_events=None, global_events=None, logger=logger)
+    shard_downloader = NoopShardDownloader()
+    worker = Worker(node_id=node_id, shard_downloader=shard_downloader, worker_events=None, global_events=None, logger=logger)
 
     path_downloaded_map: dict[str, bool] = {}
 
@@ -891,25 +897,17 @@ def test_worker_plan(case: PlanTestCase, tmp_path: Path, monkeypatch: pytest.Mon
             raise Exception('test_worker_plan not currently designed to have more than 1 instance.')
 
 
-        assigned_runner = OverrideAssignedRunner(
+        assigned_runner = AssignedRunner(
             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
+            is_downloaded=runner_config.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
-
-    def _fake_exists(path: str | Path) -> bool:  # noqa: ANN001  – match os.path.exists signature
-        return path_downloaded_map.get(str(path), False)
-
-    monkeypatch.setattr(worker_main.os.path, "exists", _fake_exists)
-
     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
index b0c81fad..84d92ab0 100644
--- a/worker/tests/test_worker_plan_utils.py
+++ b/worker/tests/test_worker_plan_utils.py
@@ -2,10 +2,10 @@ from __future__ import annotations
 
 from dataclasses import dataclass
 from pathlib import Path
-from typing import Final, List, Optional, override
+from typing import Final, List, Optional
 
 from shared.models.model_cards import MODEL_CARDS, ModelCard
-from shared.types.common import NodeId
+from shared.types.common import CommandId, NodeId
 from shared.types.models import ModelId, ModelMetadata
 from shared.types.state import State
 from shared.types.tasks import TaskId
@@ -20,7 +20,6 @@ from shared.types.worker.runners import (
     ShardAssignments,
 )
 from shared.types.worker.shards import PipelineShardMetadata
-from worker.main import AssignedRunner
 
 NODE_A: Final[NodeId] = NodeId("aaaaaaaa-aaaa-4aaa-8aaa-aaaaaaaaaaaa")
 NODE_B: Final[NodeId] = NodeId("bbbbbbbb-bbbb-4bbb-8bbb-bbbbbbbbbbbb")
@@ -33,11 +32,11 @@ 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()
+COMMAND_1_ID: Final[CommandId] = CommandId()
 
 @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
@@ -46,15 +45,6 @@ class InProcessRunner:
     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:

← b687dec6 Discovery integration master  ·  back to Exo  ·  fix placement tests b285a9f0 →