[object Object]

← back to Exo

fix topology disconnects and add heartbeat

b88abf1cc259e03e296eed397c7ee378003dcbbc · 2025-07-28 22:00:05 +0100 · Gelu Vrabie

Co-authored-by: Gelu Vrabie <gelu@exolabs.net>

Files touched

Diff

commit b88abf1cc259e03e296eed397c7ee378003dcbbc
Author: Gelu Vrabie <gelu.vrabie.univ@gmail.com>
Date:   Mon Jul 28 22:00:05 2025 +0100

    fix topology disconnects and add heartbeat
    
    Co-authored-by: Gelu Vrabie <gelu@exolabs.net>
---
 master/main.py                 | 14 ++++++---
 shared/apply/apply.py          | 18 +++++++++++-
 shared/db/sqlite/connector.py  |  5 ++--
 shared/topology.py             | 66 ++++++++++++++++++++++++++++++------------
 shared/types/events/_events.py |  8 +++++
 shared/types/topology.py       |  8 ++---
 6 files changed, 90 insertions(+), 29 deletions(-)

diff --git a/master/main.py b/master/main.py
index 45224d66..2ce5ed8b 100644
--- a/master/main.py
+++ b/master/main.py
@@ -20,6 +20,7 @@ from shared.db.sqlite.event_log_manager import EventLogManager
 from shared.types.common import NodeId
 from shared.types.events import (
     Event,
+    Heartbeat,
     TaskCreated,
     TopologyNodeCreated,
 )
@@ -114,7 +115,6 @@ class Master:
                     next_events.extend(transition_events)
 
             await self.event_log_for_writes.append_events(next_events, origin=self.node_id)
-
         # 2. get latest events
         events = await self.event_log_for_reads.get_events_since(self.state.last_event_applied_idx)
         if len(events) == 0:
@@ -126,11 +126,16 @@ class Master:
         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()
+        
+        async def heartbeat_task():
+            while True:
+                await self.event_log_for_writes.append_events([Heartbeat(node_id=self.node_id)], origin=self.node_id)
+                await asyncio.sleep(5)
+        asyncio.create_task(heartbeat_task())
 
         # TODO: we should clean these up on shutdown
         await self.forwarder_supervisor.start_as_replica()
@@ -139,7 +144,8 @@ class Master:
         else:
             await self.election_callbacks.on_became_master()
 
-        await self.event_log_for_writes.append_events([TopologyNodeCreated(node_id=self.node_id)], origin=self.node_id)
+        role = "MASTER" if self.forwarder_supervisor.current_role == ForwarderRole.MASTER else "REPLICA"
+        await self.event_log_for_writes.append_events([TopologyNodeCreated(node_id=self.node_id, role=role)], origin=self.node_id)
         while True:
             try:
                 await self._run_event_loop_body()
diff --git a/shared/apply/apply.py b/shared/apply/apply.py
index 1386a475..25eb2f27 100644
--- a/shared/apply/apply.py
+++ b/shared/apply/apply.py
@@ -8,6 +8,7 @@ from shared.types.events import (
     ChunkGenerated,
     Event,
     EventFromEventLog,
+    Heartbeat,
     InstanceActivated,
     InstanceCreated,
     InstanceDeactivated,
@@ -28,7 +29,7 @@ from shared.types.events import (
 from shared.types.profiling import NodePerformanceProfile
 from shared.types.state import State
 from shared.types.tasks import Task, TaskId
-from shared.types.topology import Node
+from shared.types.topology import Connection, Node
 from shared.types.worker.common import NodeStatus, RunnerId
 from shared.types.worker.instances import Instance, InstanceId, InstanceStatus
 from shared.types.worker.runners import RunnerStatus
@@ -43,6 +44,10 @@ def apply(state: State, event: EventFromEventLog[Event]) -> State:
     new_state: State = event_apply(event.event, state)
     return new_state.model_copy(update={"last_event_applied_idx": event.idx_in_log})
 
+@event_apply.register(Heartbeat)
+def apply_heartbeat(event: Heartbeat, state: State) -> State:
+    return state
+
 @event_apply.register(TaskCreated)
 def apply_task_created(event: TaskCreated, state: State) -> State:
     new_tasks: Mapping[TaskId, Task] = {**state.tasks, event.task_id: event.task}
@@ -134,6 +139,8 @@ def apply_chunk_generated(event: ChunkGenerated, state: State) -> State:
 def apply_topology_node_created(event: TopologyNodeCreated, state: State) -> State:
     topology = copy.copy(state.topology)
     topology.add_node(Node(node_id=event.node_id))
+    if event.role == "MASTER":
+        topology.set_master_node_id(event.node_id)
     return state.model_copy(update={"topology": topology})
 
 @event_apply.register(TopologyEdgeCreated)
@@ -154,4 +161,13 @@ def apply_topology_edge_deleted(event: TopologyEdgeDeleted, state: State) -> Sta
     if not topology.contains_connection(event.edge):
         return state
     topology.remove_connection(event.edge)
+    opposite_edge = Connection(
+        local_node_id=event.edge.send_back_node_id,
+        send_back_node_id=event.edge.local_node_id,
+        local_multiaddr=event.edge.send_back_multiaddr,
+        send_back_multiaddr=event.edge.local_multiaddr
+    )
+    if not topology.contains_connection(opposite_edge):
+        return state.model_copy(update={"topology": topology})
+    topology.remove_connection(opposite_edge)
     return state.model_copy(update={"topology": topology})
\ No newline at end of file
diff --git a/shared/db/sqlite/connector.py b/shared/db/sqlite/connector.py
index 873a89d8..d03dbd61 100644
--- a/shared/db/sqlite/connector.py
+++ b/shared/db/sqlite/connector.py
@@ -12,6 +12,7 @@ from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine
 from sqlmodel import SQLModel
 
 from shared.types.events import Event, EventParser, NodeId
+from shared.types.events._events import Heartbeat
 from shared.types.events.components import EventFromEventLog
 
 from .types import StoredEvent
@@ -246,8 +247,8 @@ class AsyncSQLiteEventStorage:
                     session.add(stored_event)
                 
                 await session.commit()
-            
-            self._logger.debug(f"Committed batch of {len(batch)} events")
+            if len([ev for ev in batch if not isinstance(ev[0], Heartbeat)]) > 0:
+                self._logger.debug(f"Committed batch of {len(batch)} events")
         
         except Exception as e:
             self._logger.error(f"Failed to commit batch: {e}")
diff --git a/shared/topology.py b/shared/topology.py
index 0f75a214..9658d483 100644
--- a/shared/topology.py
+++ b/shared/topology.py
@@ -53,6 +53,9 @@ class Topology(TopologyProto):
         rx_id = self._graph.add_node(node)
         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 set_master_node_id(self, node_id: NodeId) -> None:
+        self.master_node_id = node_id
 
     def contains_node(self, node_id: NodeId) -> bool:
         return node_id in self._node_id_to_rx_id_map
@@ -115,18 +118,27 @@ class Topology(TopologyProto):
 
     def remove_connection(self, connection: Connection) -> None:
         rx_idx = self._edge_id_to_rx_id_map[connection]
+        print(f"removing connection: {connection}, is bridge: {self._is_bridge(connection)}")
         if self._is_bridge(connection):
-            orphan_node_ids = self._get_orphan_node_ids(connection.local_node_id, connection)
+            # Determine the reference node from which reachability is calculated.
+            # Prefer a master node if the topology knows one; otherwise fall back to
+            # the local end of the connection being removed.
+            reference_node_id: NodeId = self.master_node_id if self.master_node_id is not None else connection.local_node_id
+            orphan_node_ids = self._get_orphan_node_ids(reference_node_id, connection)
+            print(f"orphan node ids: {orphan_node_ids}")
             for orphan_node_id in orphan_node_ids:
                 orphan_node_rx_id = self._node_id_to_rx_id_map[orphan_node_id]
+                print(f"removing orphan node: {orphan_node_id}, rx_id: {orphan_node_rx_id}")
                 self._graph.remove_node(orphan_node_rx_id)
                 del self._node_id_to_rx_id_map[orphan_node_id]
-                del self._rx_id_to_node_id_map[orphan_node_rx_id]
-        else:
-            self._graph.remove_edge_from_index(rx_idx)
-            del self._edge_id_to_rx_id_map[connection]
-            if rx_idx in self._rx_id_to_node_id_map:
-                del self._rx_id_to_node_id_map[rx_idx]
+            
+        self._graph.remove_edge_from_index(rx_idx)
+        del self._edge_id_to_rx_id_map[connection]
+        if rx_idx in self._rx_id_to_node_id_map:
+            del self._rx_id_to_node_id_map[rx_idx]
+        
+        
+        print(f"topology after edge removal: {self.to_snapshot()}")
 
     def get_cycles(self) -> list[list[Node]]:
         cycle_idxs = rx.simple_cycles(self._graph)
@@ -150,24 +162,42 @@ class Topology(TopologyProto):
 
     def _is_bridge(self, connection: Connection) -> bool:
         edge_idx = self._edge_id_to_rx_id_map[connection]
-        graph_copy = self._graph.copy().to_undirected()
-        components_before = rx.number_connected_components(graph_copy)
+        graph_copy: rx.PyDiGraph[Node, Connection] = self._graph.copy()
+        components_before = rx.strongly_connected_components(graph_copy)
 
         graph_copy.remove_edge_from_index(edge_idx)
-        components_after = rx.number_connected_components(graph_copy)
+        components_after = rx.strongly_connected_components(graph_copy)
 
         return components_after > components_before
 
     def _get_orphan_node_ids(self, master_node_id: NodeId, connection: Connection) -> list[NodeId]:
+        """Return node_ids that become unreachable from `master_node_id` once `connection` is removed.
+
+        A node is considered *orphaned* if there exists **no directed path** from
+        the master node to that node after deleting the edge identified by
+        ``connection``. This definition is strictly weaker than being in a
+        different *strongly* connected component and more appropriate for
+        directed networks where information only needs to flow *outwards* from
+        the master.
+        """
         edge_idx = self._edge_id_to_rx_id_map[connection]
-        graph_copy = self._graph.copy().to_undirected()
+        # Operate on a copy so the original topology remains intact while we
+        # compute reachability.
+        graph_copy: rx.PyDiGraph[Node, Connection] = self._graph.copy()
         graph_copy.remove_edge_from_index(edge_idx)
-        components = rx.connected_components(graph_copy)
 
-        orphan_node_rx_ids: set[int] = set()
-        master_node_rx_id = self._node_id_to_rx_id_map[master_node_id]
-        for component in components:
-            if master_node_rx_id not in component:
-                orphan_node_rx_ids.update(component)
+        if master_node_id not in self._node_id_to_rx_id_map:
+            # If the provided master node isn't present we conservatively treat
+            # every other node as orphaned.
+            return list(self._node_id_to_rx_id_map.keys())
 
-        return [self._rx_id_to_node_id_map[rx_id] for rx_id in orphan_node_rx_ids]
+        master_rx_id = self._node_id_to_rx_id_map[master_node_id]
+
+        # Nodes reachable by following outgoing edges from the master.
+        reachable_rx_ids: set[int] = set(rx.descendants(graph_copy, master_rx_id))
+        reachable_rx_ids.add(master_rx_id)
+
+        # Every existing node index not reachable is orphaned.
+        orphan_rx_ids = set(graph_copy.node_indices()) - reachable_rx_ids
+
+        return [self._rx_id_to_node_id_map[rx_id] for rx_id in orphan_rx_ids if rx_id in self._rx_id_to_node_id_map]
diff --git a/shared/types/events/_events.py b/shared/types/events/_events.py
index 668b556d..6ae7d005 100644
--- a/shared/types/events/_events.py
+++ b/shared/types/events/_events.py
@@ -43,6 +43,9 @@ class _EventType(str, Enum):
     Here are all the unique kinds of events that can be sent over the network.
     """
 
+    # Heartbeat Events
+    Heartbeat = "Heartbeat"
+
     # Task Events
     TaskCreated = "TaskCreated"
     TaskStateUpdated = "TaskStateUpdated"
@@ -95,6 +98,9 @@ class _BaseEvent[T: _EventType](BaseModel):
         """
         return True
 
+class Heartbeat(_BaseEvent[_EventType.Heartbeat]):
+    event_type: Literal[_EventType.Heartbeat] = _EventType.Heartbeat
+    node_id: NodeId
 
 class TaskCreated(_BaseEvent[_EventType.TaskCreated]):
     event_type: Literal[_EventType.TaskCreated] = _EventType.TaskCreated
@@ -170,6 +176,7 @@ class ChunkGenerated(_BaseEvent[_EventType.ChunkGenerated]):
 class TopologyNodeCreated(_BaseEvent[_EventType.TopologyNodeCreated]):
     event_type: Literal[_EventType.TopologyNodeCreated] = _EventType.TopologyNodeCreated
     node_id: NodeId
+    role: Literal["MASTER", "REPLICA"]
 
 class TopologyEdgeCreated(_BaseEvent[_EventType.TopologyEdgeCreated]):
     event_type: Literal[_EventType.TopologyEdgeCreated] = _EventType.TopologyEdgeCreated
@@ -192,6 +199,7 @@ class TopologyEdgeDeleted(_BaseEvent[_EventType.TopologyEdgeDeleted]):
 
 
 _Event = Union[
+    Heartbeat,
     TaskCreated,
     TaskStateUpdated,
     TaskDeleted,
diff --git a/shared/types/topology.py b/shared/types/topology.py
index 029db17f..1b9a20bc 100644
--- a/shared/types/topology.py
+++ b/shared/types/topology.py
@@ -22,8 +22,8 @@ class Connection(BaseModel):
             (
                 self.local_node_id,
                 self.send_back_node_id,
-                self.local_multiaddr.address,
-                self.send_back_multiaddr.address,
+                self.local_multiaddr.ipv4_address,
+                self.send_back_multiaddr.ipv4_address,
             )
         )
 
@@ -33,8 +33,8 @@ class Connection(BaseModel):
         return (
                 self.local_node_id == other.local_node_id
                 and self.send_back_node_id == other.send_back_node_id
-                and self.local_multiaddr.address == other.local_multiaddr.address
-                and self.send_back_multiaddr.address == other.send_back_multiaddr.address
+                and self.local_multiaddr.ipv4_address == other.local_multiaddr.ipv4_address
+                and self.send_back_multiaddr.ipv4_address == other.send_back_multiaddr.ipv4_address
         )
 
 

← dbd0bdc3 fix ci linter  ·  back to Exo  ·  better profiling 12566865 →