fix topology disconnects and add heartbeat
Co-authored-by: Gelu Vrabie <[email protected]>
This commit is contained in:
committed by
GitHub
co-authored by
Gelu Vrabie
parent
dbd0bdc34b
commit
b88abf1cc2
+10
-4
@@ -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
@@ -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})
|
||||
@@ -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
@@ -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]
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user