fix topology disconnects and add heartbeat

Co-authored-by: Gelu Vrabie <[email protected]>
This commit is contained in:
Gelu Vrabie
2025-07-28 22:00:05 +01:00
committed by GitHub
co-authored by Gelu Vrabie
parent dbd0bdc34b
commit b88abf1cc2
6 changed files with 90 additions and 29 deletions
+10 -4
View File
@@ -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()
+17 -1
View File
@@ -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})
+3 -2
View File
@@ -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}")
+48 -18
View File
@@ -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]
+8
View File
@@ -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,
+4 -4
View File
@@ -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
)