← 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
M master/main.pyM shared/apply/apply.pyM shared/db/sqlite/connector.pyM shared/topology.pyM shared/types/events/_events.pyM shared/types/topology.py
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 →