[object Object]

← back to Exo

fix disconnects

880a18d205694256671efe7069e821380d378a07 · 2025-12-15 15:23:13 +0000 · Evan Quiney

Co-authored-by: Ryuichi Leo Takashige <leo@exolabs.net>

Files touched

Diff

commit 880a18d205694256671efe7069e821380d378a07
Author: Evan Quiney <evanev7@gmail.com>
Date:   Mon Dec 15 15:23:13 2025 +0000

    fix disconnects
    
    Co-authored-by: Ryuichi Leo Takashige <leo@exolabs.net>
---
 src/exo/main.py                     | 23 ++-------------
 src/exo/master/api.py               | 28 +++++++++---------
 src/exo/master/main.py              | 58 +++++++++++++++++++++----------------
 src/exo/master/tests/test_master.py |  2 ++
 src/exo/routing/router.py           | 19 ++++--------
 src/exo/shared/apply.py             | 48 +++++++++++++++++++++++++++---
 src/exo/shared/election.py          |  5 ----
 src/exo/shared/topology.py          | 28 ++++++++++++++++++
 src/exo/shared/types/commands.py    |  5 ----
 src/exo/shared/types/events.py      | 17 +++++++----
 src/exo/shared/types/state.py       |  4 ++-
 src/exo/worker/main.py              | 55 +++++++++++++++++------------------
 src/exo/worker/utils/net_profile.py |  2 +-
 13 files changed, 174 insertions(+), 120 deletions(-)

diff --git a/src/exo/main.py b/src/exo/main.py
index 0f16d6c2..b859d2ce 100644
--- a/src/exo/main.py
+++ b/src/exo/main.py
@@ -1,7 +1,7 @@
 import argparse
 import multiprocessing as mp
 import signal
-from dataclasses import dataclass
+from dataclasses import dataclass, field
 from typing import Self
 
 import anyio
@@ -16,7 +16,6 @@ from exo.routing.router import Router, get_node_id_keypair
 from exo.shared.constants import EXO_LOG
 from exo.shared.election import Election, ElectionResult
 from exo.shared.logging import logger_cleanup, logger_setup
-from exo.shared.types.commands import KillCommand
 from exo.shared.types.common import NodeId, SessionId
 from exo.utils.channels import Receiver, channel
 from exo.utils.pydantic_ext import CamelCaseModel
@@ -35,7 +34,7 @@ class Node:
     api: API | None
 
     node_id: NodeId
-    _tg: TaskGroup | None = None
+    _tg: TaskGroup = field(init=False, default_factory=anyio.create_task_group)
 
     @classmethod
     async def create(cls, args: "Args") -> "Self":
@@ -66,7 +65,6 @@ class Node:
             node_id,
             session_id,
             exo_shard_downloader(),
-            initial_connection_messages=[],
             connection_message_receiver=router.receiver(topics.CONNECTION_MESSAGES),
             global_event_receiver=router.receiver(topics.GLOBAL_EVENTS),
             local_event_sender=router.sender(topics.LOCAL_EVENTS),
@@ -99,9 +97,8 @@ class Node:
         return cls(router, worker, election, er_recv, master, api, node_id)
 
     async def run(self):
-        async with anyio.create_task_group() as tg:
+        async with self._tg as tg:
             signal.signal(signal.SIGINT, lambda _, __: self.shutdown())
-            self._tg = tg
             tg.start_soon(self.router.run)
             tg.start_soon(self.worker.run)
             tg.start_soon(self.election.run)
@@ -110,10 +107,8 @@ class Node:
             if self.api:
                 tg.start_soon(self.api.run)
             tg.start_soon(self._elect_loop)
-            tg.start_soon(self._listen_for_kill_command)
 
     def shutdown(self):
-        assert self._tg
         # if this is our second call to shutdown, just sys.exit
         if self._tg.cancel_scope.cancel_called:
             import sys
@@ -121,18 +116,7 @@ class Node:
             sys.exit(1)
         self._tg.cancel_scope.cancel()
 
-    async def _listen_for_kill_command(self):
-        assert self._tg
-        with self.router.receiver(topics.COMMANDS) as commands:
-            async for command in commands:
-                match command.command:
-                    case KillCommand():
-                        self.shutdown()
-                    case _:
-                        pass
-
     async def _elect_loop(self):
-        assert self._tg
         with self.election_result_receiver as results:
             async for result in results:
                 # This function continues to have a lot of very specific entangled logic
@@ -187,7 +171,6 @@ class Node:
                             self.node_id,
                             result.session_id,
                             exo_shard_downloader(),
-                            initial_connection_messages=result.historic_messages,
                             connection_message_receiver=self.router.receiver(
                                 topics.CONNECTION_MESSAGES
                             ),
diff --git a/src/exo/master/api.py b/src/exo/master/api.py
index 9d65c7c1..172ae5c1 100644
--- a/src/exo/master/api.py
+++ b/src/exo/master/api.py
@@ -37,7 +37,6 @@ from exo.shared.types.commands import (
     CreateInstance,
     DeleteInstance,
     ForwarderCommand,
-    KillCommand,
     TaskFinished,
 )
 from exo.shared.types.common import CommandId, NodeId, SessionId
@@ -92,7 +91,7 @@ class API:
         # This lets us pause the API if an election is running
         election_receiver: Receiver[ElectionMessage],
     ) -> None:
-        self.state = State()
+        self._state = State()
         self.command_sender = command_sender
         self.global_event_receiver = global_event_receiver
         self.election_receiver = election_receiver
@@ -127,13 +126,15 @@ class API:
         self._tg: TaskGroup | None = None
 
     def reset(self, new_session_id: SessionId, result_clock: int):
-        self.state = State()
+        logger.info("Resetting API State")
+        self._state = State()
         self.session_id = new_session_id
         self.event_buffer = OrderedBuffer[Event]()
         self._chat_completion_queues = {}
         self.unpause(result_clock)
 
     def unpause(self, result_clock: int):
+        logger.info("Unpausing API")
         self.last_completed_election = result_clock
         self.paused = False
         self.paused_ev.set()
@@ -155,11 +156,10 @@ class API:
         self.app.get("/models")(self.get_models)
         self.app.get("/v1/models")(self.get_models)
         self.app.post("/v1/chat/completions")(self.chat_completions)
-        self.app.get("/state")(lambda: self.state)
-        self.app.delete("/kill")(self.kill_exo)
+        self.app.get("/state")(self.state)
 
-    async def kill_exo(self):
-        await self._send(KillCommand())
+    async def state(self) -> State:
+        return self._state
 
     async def create_instance(
         self, payload: CreateInstanceTaskParams
@@ -189,12 +189,12 @@ class API:
         )
 
     def get_instance(self, instance_id: InstanceId) -> Instance:
-        if instance_id not in self.state.instances:
+        if instance_id not in self._state.instances:
             raise HTTPException(status_code=404, detail="Instance not found")
-        return self.state.instances[instance_id]
+        return self._state.instances[instance_id]
 
     async def delete_instance(self, instance_id: InstanceId) -> DeleteInstanceResponse:
-        if instance_id not in self.state.instances:
+        if instance_id not in self._state.instances:
             raise HTTPException(status_code=404, detail="Instance not found")
 
         command = DeleteInstance(
@@ -261,7 +261,7 @@ class API:
 
         if not any(
             instance.shard_assignments.model_id == payload.model
-            for instance in self.state.instances.values()
+            for instance in self._state.instances.values()
         ):
             await self._trigger_notify_user_to_download_model(payload.model)
             raise HTTPException(
@@ -281,7 +281,7 @@ class API:
         """Calculate total available memory across all nodes in bytes."""
         total_available = Memory()
 
-        for node in self.state.topology.list_nodes():
+        for node in self._state.topology.list_nodes():
             if node.node_profile is not None:
                 total_available += node.node_profile.memory.ram_available
 
@@ -328,9 +328,11 @@ class API:
     async def _apply_state(self):
         with self.global_event_receiver as events:
             async for f_event in events:
+                if f_event.origin != self.session_id.master_node_id:
+                    continue
                 self.event_buffer.ingest(f_event.origin_idx, f_event.event)
                 for idx, event in self.event_buffer.drain_indexed():
-                    self.state = apply(self.state, IndexedEvent(event=event, idx=idx))
+                    self._state = apply(self._state, IndexedEvent(event=event, idx=idx))
                     if (
                         isinstance(event, ChunkGenerated)
                         and event.command_id in self._chat_completion_queues
diff --git a/src/exo/master/main.py b/src/exo/master/main.py
index 5dadb5c3..149bfbd2 100644
--- a/src/exo/master/main.py
+++ b/src/exo/master/main.py
@@ -1,6 +1,6 @@
-from datetime import datetime, timezone
+from datetime import datetime, timedelta, timezone
 
-from anyio import create_task_group
+import anyio
 from anyio.abc import TaskGroup
 from loguru import logger
 
@@ -15,7 +15,6 @@ from exo.shared.types.commands import (
     CreateInstance,
     DeleteInstance,
     ForwarderCommand,
-    KillCommand,
     RequestEventLog,
     TaskFinished,
     TestCommand,
@@ -26,9 +25,9 @@ from exo.shared.types.events import (
     ForwarderEvent,
     IndexedEvent,
     InstanceDeleted,
+    NodeTimedOut,
     TaskCreated,
     TaskDeleted,
-    TopologyEdgeDeleted,
 )
 from exo.shared.types.state import State
 from exo.shared.types.tasks import (
@@ -59,7 +58,7 @@ class Master:
         tb_only: bool = False,
     ):
         self.state = State()
-        self._tg: TaskGroup | None = None
+        self._tg: TaskGroup = anyio.create_task_group()
         self.node_id = node_id
         self.session_id = session_id
         self.command_task_mapping: dict[CommandId, TaskId] = {}
@@ -80,11 +79,11 @@ class Master:
     async def run(self):
         logger.info("Starting Master")
 
-        async with create_task_group() as tg:
-            self._tg = tg
+        async with self._tg as tg:
             tg.start_soon(self._event_processor)
             tg.start_soon(self._command_processor)
             tg.start_soon(self._loopback_processor)
+            tg.start_soon(self._plan)
         self.global_event_sender.close()
         self.local_event_receiver.close()
         self.command_receiver.close()
@@ -92,9 +91,8 @@ class Master:
         self._loopback_event_receiver.close()
 
     async def shutdown(self):
-        if self._tg:
-            logger.info("Stopping Master")
-            self._tg.cancel_scope.cancel()
+        logger.info("Stopping Master")
+        self._tg.cancel_scope.cancel()
 
     async def _command_processor(self) -> None:
         with self.command_receiver as commands:
@@ -104,7 +102,7 @@ class Master:
                     generated_events: list[Event] = []
                     command = forwarder_command.command
                     match command:
-                        case TestCommand() | KillCommand():
+                        case TestCommand():
                             pass
                         case ChatCompletion():
                             instance_task_counts: dict[InstanceId, int] = {}
@@ -191,6 +189,30 @@ class Master:
                 except ValueError as e:
                     logger.opt(exception=e).warning("Error in command processor")
 
+    # These plan loops are the cracks showing in our event sourcing architecture - more things could be commands
+    async def _plan(self) -> None:
+        while True:
+            # kill broken instances
+            connected_node_ids = set(
+                [x.node_id for x in self.state.topology.list_nodes()]
+            )
+            for instance_id, instance in self.state.instances.items():
+                for node_id in instance.shard_assignments.node_to_runner:
+                    if node_id not in connected_node_ids:
+                        await self.event_sender.send(
+                            InstanceDeleted(instance_id=instance_id)
+                        )
+                        break
+
+            # time out dead nodes
+            for node_id, time in self.state.last_seen.items():
+                now = datetime.now(tz=timezone.utc)
+                if now - time > timedelta(seconds=30):
+                    logger.info(f"Manually removing node {node_id} due to inactivity")
+                    await self.event_sender.send(NodeTimedOut(node_id=node_id))
+
+            await anyio.sleep(10)
+
     async def _event_processor(self) -> None:
         with self.local_event_receiver as local_events:
             async for local_event in local_events:
@@ -209,23 +231,9 @@ class Master:
 
                     event._master_time_stamp = datetime.now(tz=timezone.utc)  # pyright: ignore[reportPrivateUsage]
 
-                    # TODO: SQL <- What does this mean?
                     self._event_log.append(event)
                     await self._send_event(indexed)
 
-                    # TODO: This can be done in a better place. But for now, we use this to check if any running instances have been broken.
-                    if isinstance(event, TopologyEdgeDeleted):
-                        connected_node_ids = set(
-                            [x.node_id for x in self.state.topology.list_nodes()]
-                        )
-                        for instance_id, instance in self.state.instances.items():
-                            for node_id in instance.shard_assignments.node_to_runner:
-                                if node_id not in connected_node_ids:
-                                    await self.event_sender.send(
-                                        InstanceDeleted(instance_id=instance_id)
-                                    )
-                                    break
-
     async def _loopback_processor(self) -> None:
         # this would ideally not be necessary.
         # this is WAY less hacky than how I was working around this before
diff --git a/src/exo/master/tests/test_master.py b/src/exo/master/tests/test_master.py
index a87abc34..948bcb1f 100644
--- a/src/exo/master/tests/test_master.py
+++ b/src/exo/master/tests/test_master.py
@@ -1,3 +1,4 @@
+from datetime import datetime, timezone
 from typing import Sequence
 
 import anyio
@@ -84,6 +85,7 @@ async def test_master():
                 session=session_id,
                 event=(
                     NodePerformanceMeasured(
+                        when=str(datetime.now(tz=timezone.utc)),
                         node_id=node_id,
                         node_profile=NodePerformanceProfile(
                             model_id="maccy",
diff --git a/src/exo/routing/router.py b/src/exo/routing/router.py
index 21aece29..ac6073af 100644
--- a/src/exo/routing/router.py
+++ b/src/exo/routing/router.py
@@ -44,7 +44,7 @@ class TopicRouter[T: CamelCaseModel]:
         self.senders: set[Sender[T]] = set()
         send, recv = channel[T]()
         self.receiver: Receiver[T] = recv
-        self.temp_sender: Sender[T] | None = send
+        self._sender: Sender[T] = send
         self.networking_sender: Sender[tuple[str, bytes]] = networking_sender
 
     async def run(self):
@@ -68,8 +68,7 @@ class TopicRouter[T: CamelCaseModel]:
         # Close all the things!
         for sender in self.senders:
             sender.close()
-        if self.temp_sender:
-            self.temp_sender.close()
+        self._sender.close()
         self.receiver.close()
 
     async def publish(self, item: T):
@@ -89,6 +88,9 @@ class TopicRouter[T: CamelCaseModel]:
     async def publish_bytes(self, data: bytes):
         await self.publish(self.topic.deserialize(data))
 
+    def new_sender(self) -> Sender[T]:
+        return self._sender.clone()
+
     async def _send_out(self, item: T):
         logger.trace(f"TopicRouter {self.topic.topic} sending {item}")
         await self.networking_sender.send(
@@ -126,16 +128,7 @@ class Router:
         # There's gotta be a way to do this without THIS many asserts
         assert router is not None
         assert router.topic == topic
-        send: Sender[T] | None = cast(Sender[T] | None, router.temp_sender)
-        if send:
-            router.temp_sender = None
-            return send
-        try:
-            sender = cast(Receiver[T], router.receiver).clone_sender()
-        except ClosedResourceError:
-            sender, router.receiver = cast(
-                tuple[Sender[T], Receiver[CamelCaseModel]], channel[T]()
-            )
+        sender = cast(TopicRouter[T], router).new_sender()
         return sender
 
     def receiver[T: CamelCaseModel](self, topic: TypedTopic[T]) -> Receiver[T]:
diff --git a/src/exo/shared/apply.py b/src/exo/shared/apply.py
index 178d2c5f..9bb597cb 100644
--- a/src/exo/shared/apply.py
+++ b/src/exo/shared/apply.py
@@ -1,5 +1,6 @@
 import copy
 from collections.abc import Mapping, Sequence
+from datetime import datetime
 
 from loguru import logger
 
@@ -14,6 +15,7 @@ from exo.shared.types.events import (
     NodeDownloadProgress,
     NodeMemoryMeasured,
     NodePerformanceMeasured,
+    NodeTimedOut,
     RunnerDeleted,
     RunnerStatusUpdated,
     TaskAcknowledged,
@@ -45,6 +47,10 @@ def event_apply(event: Event, state: State) -> State:
             return apply_instance_created(event, state)
         case InstanceDeleted():
             return apply_instance_deleted(event, state)
+        case NodeCreated():
+            return apply_topology_node_created(event, state)
+        case NodeTimedOut():
+            return apply_node_timed_out(event, state)
         case NodePerformanceMeasured():
             return apply_node_performance_measured(event, state)
         case NodeDownloadProgress():
@@ -63,8 +69,6 @@ def event_apply(event: Event, state: State) -> State:
             return apply_task_failed(event, state)
         case TaskStatusUpdated():
             return apply_task_status_updated(event, state)
-        case NodeCreated():
-            return apply_topology_node_created(event, state)
         case TopologyEdgeCreated():
             return apply_topology_edge_created(event, state)
         case TopologyEdgeDeleted():
@@ -183,6 +187,24 @@ def apply_runner_deleted(event: RunnerDeleted, state: State) -> State:
     return state.model_copy(update={"runners": new_runners})
 
 
+def apply_node_timed_out(event: NodeTimedOut, state: State) -> State:
+    topology = copy.copy(state.topology)
+    state.topology.remove_node(event.node_id)
+    node_profiles = {
+        key: value for key, value in state.node_profiles.items() if key != event.node_id
+    }
+    last_seen = {
+        key: value for key, value in state.last_seen.items() if key != event.node_id
+    }
+    return state.model_copy(
+        update={
+            "topology": topology,
+            "node_profiles": node_profiles,
+            "last_seen": last_seen,
+        }
+    )
+
+
 def apply_node_performance_measured(
     event: NodePerformanceMeasured, state: State
 ) -> State:
@@ -190,13 +212,23 @@ def apply_node_performance_measured(
         **state.node_profiles,
         event.node_id: event.node_profile,
     }
+    last_seen: Mapping[NodeId, datetime] = {
+        **state.last_seen,
+        event.node_id: datetime.fromisoformat(event.when),
+    }
     state = state.model_copy(update={"node_profiles": new_profiles})
     topology = copy.copy(state.topology)
     # TODO: NodeCreated
     if not topology.contains_node(event.node_id):
         topology.add_node(NodeInfo(node_id=event.node_id))
     topology.update_node_profile(event.node_id, event.node_profile)
-    return state.model_copy(update={"topology": topology})
+    return state.model_copy(
+        update={
+            "node_profiles": new_profiles,
+            "topology": topology,
+            "last_seen": last_seen,
+        }
+    )
 
 
 def apply_node_memory_measured(event: NodeMemoryMeasured, state: State) -> State:
@@ -224,12 +256,20 @@ def apply_node_memory_measured(event: NodeMemoryMeasured, state: State) -> State
             **state.node_profiles,
             event.node_id: created,
         }
+        last_seen: Mapping[NodeId, datetime] = {
+            **state.last_seen,
+            event.node_id: datetime.fromisoformat(event.when),
+        }
         if not topology.contains_node(event.node_id):
             topology.add_node(NodeInfo(node_id=event.node_id))
             # TODO: NodeCreated
         topology.update_node_profile(event.node_id, created)
         return state.model_copy(
-            update={"node_profiles": created_profiles, "topology": topology}
+            update={
+                "node_profiles": created_profiles,
+                "topology": topology,
+                "last_seen": last_seen,
+            }
         )
 
     updated = existing.model_copy(update={"memory": event.memory})
diff --git a/src/exo/shared/election.py b/src/exo/shared/election.py
index b4dc36b6..9d030d5e 100644
--- a/src/exo/shared/election.py
+++ b/src/exo/shared/election.py
@@ -44,7 +44,6 @@ class ElectionResult(CamelCaseModel):
     session_id: SessionId
     won_clock: int
     is_new_master: bool
-    historic_messages: list[ConnectionMessage]
 
 
 class Election:
@@ -84,7 +83,6 @@ class Election:
         self._campaign_cancel_scope: CancelScope | None = None
         self._campaign_done: Event | None = None
         self._tg: TaskGroup | None = None
-        self._connection_messages: list[ConnectionMessage] = []
 
     async def run(self):
         logger.info("Starting Election")
@@ -121,7 +119,6 @@ class Election:
                 won_clock=em.clock,
                 session_id=em.proposed_session,
                 is_new_master=is_new_master,
-                historic_messages=self._connection_messages,
             )
         )
 
@@ -188,8 +185,6 @@ class Election:
                     self._campaign, candidates, DEFAULT_ELECTION_TIMEOUT
                 )
                 logger.debug("Campaign started")
-                self._connection_messages.append(first)
-                self._connection_messages.extend(rest)
                 logger.debug("Connection message added")
 
     async def _command_counter(self) -> None:
diff --git a/src/exo/shared/topology.py b/src/exo/shared/topology.py
index 7413161f..46419d72 100644
--- a/src/exo/shared/topology.py
+++ b/src/exo/shared/topology.py
@@ -55,6 +55,22 @@ class Topology:
             and len(self._graph.neighbors(self._node_id_to_rx_id_map[node_id])) == 1
         )
 
+    def neighbours(self, node_id: NodeId) -> list[NodeId]:
+        return [
+            self._rx_id_to_node_id_map[rx_id]
+            for rx_id in self._graph.neighbors(self._node_id_to_rx_id_map[node_id])
+        ]
+
+    def out_edges(self, node_id: NodeId) -> list[tuple[NodeId, Connection]]:
+        if node_id not in self._node_id_to_rx_id_map:
+            return []
+        return [
+            (self._rx_id_to_node_id_map[nid], conn)
+            for _, nid, conn in self._graph.out_edges(
+                self._node_id_to_rx_id_map[node_id]
+            )
+        ]
+
     def contains_node(self, node_id: NodeId) -> bool:
         return node_id in self._node_id_to_rx_id_map
 
@@ -112,6 +128,16 @@ class Topology:
             return None
 
     def remove_node(self, node_id: NodeId) -> None:
+        if node_id not in self._node_id_to_rx_id_map:
+            return
+
+        for connection in self.list_connections():
+            if (
+                connection.local_node_id == node_id
+                or connection.send_back_node_id == node_id
+            ):
+                self.remove_connection(connection)
+
         rx_idx = self._node_id_to_rx_id_map[node_id]
         self._graph.remove_node(rx_idx)
 
@@ -119,6 +145,8 @@ class Topology:
         del self._rx_id_to_node_id_map[rx_idx]
 
     def remove_connection(self, connection: Connection) -> None:
+        if connection not in self._edge_id_to_rx_id_map:
+            return
         rx_idx = self._edge_id_to_rx_id_map[connection]
         self._graph.remove_edge_from_index(rx_idx)
         del self._edge_id_to_rx_id_map[connection]
diff --git a/src/exo/shared/types/commands.py b/src/exo/shared/types/commands.py
index 1ea4027a..0a584ff5 100644
--- a/src/exo/shared/types/commands.py
+++ b/src/exo/shared/types/commands.py
@@ -16,10 +16,6 @@ class TestCommand(BaseCommand):
     __test__ = False
 
 
-class KillCommand(BaseCommand):
-    pass
-
-
 class ChatCompletion(BaseCommand):
     request_params: ChatCompletionTaskParams
 
@@ -45,7 +41,6 @@ class RequestEventLog(BaseCommand):
 
 Command = (
     TestCommand
-    | KillCommand
     | RequestEventLog
     | ChatCompletion
     | CreateInstance
diff --git a/src/exo/shared/types/events.py b/src/exo/shared/types/events.py
index 7ad465d4..29b750ef 100644
--- a/src/exo/shared/types/events.py
+++ b/src/exo/shared/types/events.py
@@ -81,20 +81,26 @@ class NodeCreated(BaseEvent):
     node_id: NodeId
 
 
-class NodePerformanceMeasured(BaseEvent):
+class NodeTimedOut(BaseEvent):
     node_id: NodeId
-    node_profile: NodePerformanceProfile
 
 
-class NodeDownloadProgress(BaseEvent):
-    download_progress: DownloadProgress
+class NodePerformanceMeasured(BaseEvent):
+    node_id: NodeId
+    when: str  # this is a manually cast datetime overrode by the master when the event is indexed, rather than the local time on the device
+    node_profile: NodePerformanceProfile
 
 
 class NodeMemoryMeasured(BaseEvent):
     node_id: NodeId
+    when: str  # this is a manually cast datetime overrode by the master when the event is indexed, rather than the local time on the device
     memory: MemoryPerformanceProfile
 
 
+class NodeDownloadProgress(BaseEvent):
+    download_progress: DownloadProgress
+
+
 class ChunkGenerated(BaseEvent):
     command_id: CommandId
     chunk: GenerationChunk
@@ -119,11 +125,12 @@ Event = (
     | InstanceDeleted
     | RunnerStatusUpdated
     | RunnerDeleted
+    | NodeCreated
+    | NodeTimedOut
     | NodePerformanceMeasured
     | NodeMemoryMeasured
     | NodeDownloadProgress
     | ChunkGenerated
-    | NodeCreated
     | TopologyEdgeCreated
     | TopologyEdgeDeleted
 )
diff --git a/src/exo/shared/types/state.py b/src/exo/shared/types/state.py
index efdb5bcb..58b14d2e 100644
--- a/src/exo/shared/types/state.py
+++ b/src/exo/shared/types/state.py
@@ -1,4 +1,5 @@
 from collections.abc import Mapping, Sequence
+from datetime import datetime
 from typing import Any, cast
 
 from pydantic import ConfigDict, Field, field_serializer, field_validator
@@ -35,7 +36,8 @@ class State(CamelCaseModel):
     downloads: Mapping[NodeId, Sequence[DownloadProgress]] = {}
     tasks: Mapping[TaskId, Task] = {}
     node_profiles: Mapping[NodeId, NodePerformanceProfile] = {}
-    topology: Topology = Topology()
+    last_seen: Mapping[NodeId, datetime] = {}
+    topology: Topology = Field(default_factory=Topology)
     last_event_applied_idx: int = Field(default=-1, ge=-1)
 
     @field_serializer("topology", mode="plain")
diff --git a/src/exo/worker/main.py b/src/exo/worker/main.py
index 6028c2b4..a5c049dc 100644
--- a/src/exo/worker/main.py
+++ b/src/exo/worker/main.py
@@ -1,3 +1,4 @@
+from datetime import datetime, timezone
 from random import random
 
 import anyio
@@ -50,7 +51,7 @@ from exo.worker.download.shard_downloader import RepoDownloadProgress, ShardDown
 from exo.worker.plan import plan
 from exo.worker.runner.runner_supervisor import RunnerSupervisor
 from exo.worker.utils import start_polling_memory_metrics, start_polling_node_metrics
-from exo.worker.utils.net_profile import connect_all
+from exo.worker.utils.net_profile import check_reachable
 
 
 class Worker:
@@ -60,7 +61,6 @@ class Worker:
         session_id: SessionId,
         shard_downloader: ShardDownloader,
         *,
-        initial_connection_messages: list[ConnectionMessage],
         connection_message_receiver: Receiver[ConnectionMessage],
         global_event_receiver: Receiver[ForwarderEvent],
         local_event_sender: Sender[ForwarderEvent],
@@ -80,7 +80,6 @@ class Worker:
         self.command_sender = command_sender
         self.connection_message_receiver = connection_message_receiver
         self.event_buffer = OrderedBuffer[Event]()
-        self._initial_connection_messages = initial_connection_messages
         self.out_for_delivery: dict[EventId, ForwarderEvent] = {}
 
         self.state: State = State()
@@ -104,7 +103,9 @@ class Worker:
         ) -> None:
             await self.event_sender.send(
                 NodePerformanceMeasured(
-                    node_id=self.node_id, node_profile=node_performance_profile
+                    node_id=self.node_id,
+                    node_profile=node_performance_profile,
+                    when=str(datetime.now(tz=timezone.utc)),
                 ),
             )
 
@@ -112,7 +113,11 @@ class Worker:
             memory_profile: MemoryPerformanceProfile,
         ) -> None:
             await self.event_sender.send(
-                NodeMemoryMeasured(node_id=self.node_id, memory=memory_profile)
+                NodeMemoryMeasured(
+                    node_id=self.node_id,
+                    memory=memory_profile,
+                    when=str(datetime.now(tz=timezone.utc)),
+                )
             )
 
         # END CLEANUP
@@ -128,12 +133,6 @@ class Worker:
             tg.start_soon(self._event_applier)
             tg.start_soon(self._forward_events)
             tg.start_soon(self._poll_connection_updates)
-            # TODO: This is a little gross, but not too bad
-            for msg in self._initial_connection_messages:
-                await self.event_sender.send(
-                    self._convert_connection_message_to_event(msg)
-                )
-            self._initial_connection_messages = []
 
         # Actual shutdown code - waits for all tasks to complete before executing.
         self.local_event_sender.close()
@@ -143,9 +142,11 @@ class Worker:
 
     async def _event_applier(self):
         with self.global_event_receiver as events:
-            async for event in events:
-                self.event_buffer.ingest(event.origin_idx, event.event)
-                event_id = event.event.event_id
+            async for f_event in events:
+                if f_event.origin != self.session_id.master_node_id:
+                    continue
+                self.event_buffer.ingest(f_event.origin_idx, f_event.event)
+                event_id = f_event.event.event_id
                 if event_id in self.out_for_delivery:
                     del self.out_for_delivery[event_id]
 
@@ -167,15 +168,8 @@ class Worker:
                 elif indexed_events and self._nack_cancel_scope:
                     self._nack_cancel_scope.cancel()
 
-                flag = False
                 for idx, event in indexed_events:
                     self.state = apply(self.state, IndexedEvent(idx=idx, event=event))
-                    if event_relevant_to_worker(event, self):
-                        flag = True
-
-                # 3. If we've found a "relevant" event, run a plan -> op -> execute cycle.
-                if flag:
-                    pass
 
     async def plan_step(self):
         while True:
@@ -420,23 +414,28 @@ class Worker:
         while True:
             # TODO: EdgeDeleted
             edges = set(self.state.topology.list_connections())
-            conns = await connect_all(self.state.topology)
+            conns = await check_reachable(self.state.topology)
             for nid in conns:
                 for ip in conns[nid]:
                     edge = Connection(
                         local_node_id=self.node_id,
                         send_back_node_id=nid,
+                        # nonsense multiaddr
                         send_back_multiaddr=Multiaddr(address=f"/ip4/{ip}/tcp/8000")
                         if "." in ip
+                        # nonsense multiaddr
                         else Multiaddr(address=f"/ip6/{ip}/tcp/8000"),
                     )
                     if edge not in edges:
-                        logger.debug(f"manually discovered {edge=}")
+                        logger.debug(f"ping discovered {edge=}")
                         await self.event_sender.send(TopologyEdgeCreated(edge=edge))
 
-            await anyio.sleep(10)
-
+            for nid, conn in self.state.topology.out_edges(self.node_id):
+                if (
+                    nid not in conns
+                    or conn.send_back_multiaddr.ip_address not in conns.get(nid, set())
+                ):
+                    logger.debug(f"ping failed to discover {conn=}")
+                    await self.event_sender.send(TopologyEdgeDeleted(edge=conn))
 
-def event_relevant_to_worker(event: Event, worker: Worker):
-    # TODO
-    return True
+            await anyio.sleep(10)
diff --git a/src/exo/worker/utils/net_profile.py b/src/exo/worker/utils/net_profile.py
index 923048b0..1c8c5fe4 100644
--- a/src/exo/worker/utils/net_profile.py
+++ b/src/exo/worker/utils/net_profile.py
@@ -27,7 +27,7 @@ async def check_reachability(
         out[target_node_id].add(target_ip)
 
 
-async def connect_all(topology: Topology) -> dict[NodeId, set[str]]:
+async def check_reachable(topology: Topology) -> dict[NodeId, set[str]]:
     reachable: dict[NodeId, set[str]] = {}
     async with create_task_group() as tg:
         for node in topology.list_nodes():

← 70298ce0 Negative index nack request  ·  back to Exo  ·  backport the dashboard to staging 09593c5e →