[object Object]

← back to Exo

add node started event

2e4635a8f543ebd60f37a755e0c28f36081ade68 · 2025-07-26 19:12:26 +0100 · Gelu Vrabie

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

Files touched

Diff

commit 2e4635a8f543ebd60f37a755e0c28f36081ade68
Author: Gelu Vrabie <gelu.vrabie.univ@gmail.com>
Date:   Sat Jul 26 19:12:26 2025 +0100

    add node started event
    
    Co-authored-by: Gelu Vrabie <gelu@exolabs.net>
---
 master/main.py                       | 17 ++++++++++++++---
 master/tests/test_master.py          | 16 +++++++++-------
 master/tests/test_placement.py       |  6 +++---
 master/tests/test_placement_utils.py | 26 +++++++++++++-------------
 master/tests/test_topology.py        | 32 ++++++++++++++++----------------
 shared/apply/apply.py                |  8 ++++++++
 shared/topology.py                   | 14 +++++++-------
 shared/types/events/_events.py       |  5 +++++
 shared/types/topology.py             |  2 +-
 9 files changed, 76 insertions(+), 50 deletions(-)

diff --git a/master/main.py b/master/main.py
index 3c1e8a57..c755cf75 100644
--- a/master/main.py
+++ b/master/main.py
@@ -19,6 +19,7 @@ from shared.types.common import NodeId
 from shared.types.events import (
     Event,
     TaskCreated,
+    TopologyNodeCreated,
 )
 from shared.types.events.commands import (
     ChatCompletionCommand,
@@ -32,10 +33,11 @@ from shared.types.worker.instances import Instance
 
 
 class Master:
-    def __init__(self, node_id: NodeId, command_buffer: list[Command], global_events: AsyncSQLiteEventStorage, forwarder_binary_path: Path, logger: logging.Logger):
+    def __init__(self, node_id: NodeId, command_buffer: list[Command], global_events: AsyncSQLiteEventStorage, worker_events: AsyncSQLiteEventStorage, forwarder_binary_path: Path, logger: logging.Logger):
         self.node_id = node_id
         self.command_buffer = command_buffer
         self.global_events = global_events
+        self.worker_events = worker_events
         self.forwarder_supervisor = ForwarderSupervisor(
             forwarder_binary_path=forwarder_binary_path,
             logger=logger
@@ -43,6 +45,13 @@ class Master:
         self.election_callbacks = ElectionCallbacks(self.forwarder_supervisor, logger)
         self.logger = logger
 
+    @property
+    def event_log_for_writes(self) -> AsyncSQLiteEventStorage:
+        if self.forwarder_supervisor.current_role == ForwarderRole.MASTER:
+            return self.global_events
+        else:
+            return self.worker_events
+
     async def _get_state_snapshot(self) -> State:
         # TODO: for now start from scratch every time, but we can optimize this by keeping a snapshot on disk so we don't have to re-apply all events
         return State()
@@ -85,7 +94,7 @@ class Master:
                     transition_events = get_transition_events(self.state.instances, placement)
                     next_events.extend(transition_events)
 
-            await self.global_events.append_events(next_events, origin=self.node_id)
+            await self.event_log_for_writes.append_events(next_events, origin=self.node_id)
 
         # 2. get latest events
         events = await self.global_events.get_events_since(self.state.last_event_applied_idx)
@@ -109,6 +118,7 @@ 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)
         while True:
             try:
                 await self._run_event_loop_body()
@@ -133,6 +143,7 @@ async def main():
     event_log_manager = EventLogManager(EventLogConfig(), logger=logger)
     await event_log_manager.initialize()
     global_events: AsyncSQLiteEventStorage = event_log_manager.global_events
+    worker_events: AsyncSQLiteEventStorage = event_log_manager.worker_events
 
     command_buffer: List[Command] = []
 
@@ -152,7 +163,7 @@ async def main():
     api_thread.start()
     logger.info('Running FastAPI server in a separate thread. Listening on port 8000.')
 
-    master = Master(node_id, command_buffer, global_events, forwarder_binary_path=Path("./build/forwarder"), logger=logger)
+    master = Master(node_id, command_buffer, global_events, worker_events, forwarder_binary_path=Path("./build/forwarder"), logger=logger)
     await master.run()
 
 if __name__ == "__main__":
diff --git a/master/tests/test_master.py b/master/tests/test_master.py
index f8fc6558..5445c967 100644
--- a/master/tests/test_master.py
+++ b/master/tests/test_master.py
@@ -13,6 +13,7 @@ from shared.db.sqlite.event_log_manager import EventLogManager
 from shared.types.api import ChatCompletionMessage, ChatCompletionTaskParams
 from shared.types.common import NodeId
 from shared.types.events import TaskCreated
+from shared.types.events._events import TopologyNodeCreated
 from shared.types.events.commands import ChatCompletionCommand, Command, CommandId
 from shared.types.tasks import ChatCompletionTask, TaskStatus, TaskType
 
@@ -38,7 +39,7 @@ async def test_master():
     forwarder_binary_path = _create_forwarder_dummy_binary()
 
     node_id = NodeId("aaaaaaaa-aaaa-4aaa-8aaa-aaaaaaaaaaaa")
-    master = Master(node_id, command_buffer=command_buffer, global_events=global_events, forwarder_binary_path=forwarder_binary_path, logger=logger)
+    master = Master(node_id, command_buffer=command_buffer, global_events=global_events, worker_events=event_log_manager.worker_events, forwarder_binary_path=forwarder_binary_path, logger=logger)
     asyncio.create_task(master.run())
 
     command_buffer.append(
@@ -54,15 +55,16 @@ async def test_master():
         await asyncio.sleep(0.001)
 
     events = await global_events.get_events_since(0)
-    assert len(events) == 1
+    assert len(events) == 2
     assert events[0].idx_in_log == 1
-    assert isinstance(events[0].event, TaskCreated)
-    assert events[0].event == TaskCreated(
-        task_id=events[0].event.task_id,
+    assert isinstance(events[0].event, TopologyNodeCreated)
+    assert isinstance(events[1].event, TaskCreated)
+    assert events[1].event == TaskCreated(
+        task_id=events[1].event.task_id,
         task=ChatCompletionTask(
-            task_id=events[0].event.task_id,
+            task_id=events[1].event.task_id,
             task_type=TaskType.CHAT_COMPLETION,
-            instance_id=events[0].event.task.instance_id,
+            instance_id=events[1].event.task.instance_id,
             task_status=TaskStatus.PENDING,
             task_params=ChatCompletionTaskParams(
                 model="llama-3.2-1b",
diff --git a/master/tests/test_placement.py b/master/tests/test_placement.py
index 3218297e..9bef8116 100644
--- a/master/tests/test_placement.py
+++ b/master/tests/test_placement.py
@@ -75,9 +75,9 @@ def test_get_instance_placements_create_instance(
     node_id_a = NodeId()
     node_id_b = NodeId()
     node_id_c = NodeId()
-    topology.add_node(create_node(available_memory[0], node_id_a), node_id_a)
-    topology.add_node(create_node(available_memory[1], node_id_b), node_id_b)
-    topology.add_node(create_node(available_memory[2], node_id_c), node_id_c)
+    topology.add_node(create_node(available_memory[0], node_id_a))
+    topology.add_node(create_node(available_memory[1], node_id_b))
+    topology.add_node(create_node(available_memory[2], node_id_c))
     topology.add_connection(create_connection(node_id_a, node_id_b))
     topology.add_connection(create_connection(node_id_b, node_id_c))
     topology.add_connection(create_connection(node_id_c, node_id_a))
diff --git a/master/tests/test_placement_utils.py b/master/tests/test_placement_utils.py
index 7dce222f..2ef84cd1 100644
--- a/master/tests/test_placement_utils.py
+++ b/master/tests/test_placement_utils.py
@@ -27,8 +27,8 @@ def test_filter_cycles_by_memory(topology: Topology, create_node: Callable[[int,
     node1 = create_node(1000, node1_id)
     node2 = create_node(1000, node2_id)
     
-    topology.add_node(node1, node1_id)
-    topology.add_node(node2, node2_id)
+    topology.add_node(node1)
+    topology.add_node(node2)
     
     connection1 = create_connection(node1_id, node2_id)
     connection2 = create_connection(node2_id, node1_id)
@@ -55,8 +55,8 @@ def test_filter_cycles_by_insufficient_memory(topology: Topology, create_node: C
     node1 = create_node(1000, node1_id)
     node2 = create_node(1000, node2_id)
 
-    topology.add_node(node1, node1_id)
-    topology.add_node(node2, node2_id)
+    topology.add_node(node1)
+    topology.add_node(node2)
 
     connection1 = create_connection(node1_id, node2_id)
     connection2 = create_connection(node2_id, node1_id)
@@ -81,9 +81,9 @@ def test_filter_multiple_cycles_by_memory(topology: Topology, create_node: Calla
     node_b = create_node(500, node_b_id)
     node_c = create_node(1000, node_c_id)
     
-    topology.add_node(node_a, node_a_id)
-    topology.add_node(node_b, node_b_id)
-    topology.add_node(node_c, node_c_id)
+    topology.add_node(node_a)
+    topology.add_node(node_b)
+    topology.add_node(node_c)
     
     topology.add_connection(create_connection(node_a_id, node_b_id))
     topology.add_connection(create_connection(node_b_id, node_a_id))
@@ -111,9 +111,9 @@ def test_get_smallest_cycles(topology: Topology, create_node: Callable[[int, Nod
     node_b = create_node(500, node_b_id)
     node_c = create_node(1000, node_c_id)
 
-    topology.add_node(node_a, node_a_id)
-    topology.add_node(node_b, node_b_id)
-    topology.add_node(node_c, node_c_id)
+    topology.add_node(node_a)
+    topology.add_node(node_b)
+    topology.add_node(node_c)
 
     topology.add_connection(create_connection(node_a_id, node_b_id))
     topology.add_connection(create_connection(node_b_id, node_c_id))
@@ -143,9 +143,9 @@ def test_get_shard_assignments(topology: Topology, create_node: Callable[[int, N
     node_b = create_node(available_memory[1], node_b_id)
     node_c = create_node(available_memory[2], node_c_id)
 
-    topology.add_node(node_a, node_a_id)
-    topology.add_node(node_b, node_b_id)
-    topology.add_node(node_c, node_c_id)
+    topology.add_node(node_a)
+    topology.add_node(node_b)
+    topology.add_node(node_c)
 
     topology.add_connection(create_connection(node_a_id, node_b_id))
     topology.add_connection(create_connection(node_b_id, node_c_id))
diff --git a/master/tests/test_topology.py b/master/tests/test_topology.py
index 1e395d2e..e5790c0a 100644
--- a/master/tests/test_topology.py
+++ b/master/tests/test_topology.py
@@ -32,7 +32,7 @@ def test_add_node(topology: Topology, node_profile: NodePerformanceProfile):
     node_id = NodeId()
 
     # act
-    topology.add_node(Node(node_id=node_id, node_profile=node_profile), node_id=node_id)
+    topology.add_node(Node(node_id=node_id, node_profile=node_profile))
 
     # assert
     data = topology.get_node_profile(node_id)
@@ -41,8 +41,8 @@ def test_add_node(topology: Topology, node_profile: NodePerformanceProfile):
 
 def test_add_connection(topology: Topology, node_profile: NodePerformanceProfile, connection: Connection):
     # arrange
-    topology.add_node(Node(node_id=connection.source_node_id, node_profile=node_profile), node_id=connection.source_node_id)
-    topology.add_node(Node(node_id=connection.sink_node_id, node_profile=node_profile), node_id=connection.sink_node_id)
+    topology.add_node(Node(node_id=connection.source_node_id, node_profile=node_profile))
+    topology.add_node(Node(node_id=connection.sink_node_id, node_profile=node_profile))
     topology.add_connection(connection)
 
     # act
@@ -53,8 +53,8 @@ def test_add_connection(topology: Topology, node_profile: NodePerformanceProfile
 
 def test_update_node_profile(topology: Topology, node_profile: NodePerformanceProfile, connection: Connection):
     # arrange
-    topology.add_node(Node(node_id=connection.source_node_id, node_profile=node_profile), node_id=connection.source_node_id)
-    topology.add_node(Node(node_id=connection.sink_node_id, node_profile=node_profile), node_id=connection.sink_node_id)
+    topology.add_node(Node(node_id=connection.source_node_id, node_profile=node_profile))
+    topology.add_node(Node(node_id=connection.sink_node_id, node_profile=node_profile))
     topology.add_connection(connection)
 
     new_node_profile = NodePerformanceProfile(model_id="test", chip_id="test", memory=MemoryPerformanceProfile(ram_total=1000, ram_available=1000, swap_total=1000, swap_available=1000), network_interfaces=[], system=SystemPerformanceProfile(flops_fp16=1000))
@@ -68,8 +68,8 @@ def test_update_node_profile(topology: Topology, node_profile: NodePerformancePr
 
 def test_update_connection_profile(topology: Topology, node_profile: NodePerformanceProfile, connection: Connection):
     # arrange
-    topology.add_node(Node(node_id=connection.source_node_id, node_profile=node_profile), node_id=connection.source_node_id)
-    topology.add_node(Node(node_id=connection.sink_node_id, node_profile=node_profile), node_id=connection.sink_node_id)
+    topology.add_node(Node(node_id=connection.source_node_id, node_profile=node_profile))
+    topology.add_node(Node(node_id=connection.sink_node_id, node_profile=node_profile))
     topology.add_connection(connection)
 
     new_connection_profile = ConnectionProfile(throughput=2000, latency=2000, jitter=2000)
@@ -84,8 +84,8 @@ def test_update_connection_profile(topology: Topology, node_profile: NodePerform
 
 def test_remove_connection_still_connected(topology: Topology, node_profile: NodePerformanceProfile, connection: Connection):
     # arrange
-    topology.add_node(Node(node_id=connection.source_node_id, node_profile=node_profile), node_id=connection.source_node_id)
-    topology.add_node(Node(node_id=connection.sink_node_id, node_profile=node_profile), node_id=connection.sink_node_id)
+    topology.add_node(Node(node_id=connection.source_node_id, node_profile=node_profile))
+    topology.add_node(Node(node_id=connection.sink_node_id, node_profile=node_profile))
     topology.add_connection(connection)
 
     # act
@@ -103,9 +103,9 @@ def test_remove_connection_bridge(topology: Topology, node_profile: NodePerforma
     node_a_id = NodeId()
     node_b_id = NodeId()
     
-    topology.add_node(Node(node_id=master_id, node_profile=node_profile), node_id=master_id)
-    topology.add_node(Node(node_id=node_a_id, node_profile=node_profile), node_id=node_a_id)
-    topology.add_node(Node(node_id=node_b_id, node_profile=node_profile), node_id=node_b_id)
+    topology.add_node(Node(node_id=master_id, node_profile=node_profile))
+    topology.add_node(Node(node_id=node_a_id, node_profile=node_profile))
+    topology.add_node(Node(node_id=node_b_id, node_profile=node_profile))
     
     connection_master_to_a = Connection(
         source_node_id=master_id,
@@ -143,8 +143,8 @@ def test_remove_connection_bridge(topology: Topology, node_profile: NodePerforma
 
 def test_remove_node_still_connected(topology: Topology, node_profile: NodePerformanceProfile, connection: Connection):
     # arrange
-    topology.add_node(Node(node_id=connection.source_node_id, node_profile=node_profile), node_id=connection.source_node_id)
-    topology.add_node(Node(node_id=connection.sink_node_id, node_profile=node_profile), node_id=connection.sink_node_id)
+    topology.add_node(Node(node_id=connection.source_node_id, node_profile=node_profile))
+    topology.add_node(Node(node_id=connection.sink_node_id, node_profile=node_profile))
     topology.add_connection(connection)
 
     # act
@@ -157,8 +157,8 @@ def test_remove_node_still_connected(topology: Topology, node_profile: NodePerfo
 
 def test_list_nodes(topology: Topology, node_profile: NodePerformanceProfile, connection: Connection):
     # arrange
-    topology.add_node(Node(node_id=connection.source_node_id, node_profile=node_profile), node_id=connection.source_node_id)
-    topology.add_node(Node(node_id=connection.sink_node_id, node_profile=node_profile), node_id=connection.sink_node_id)
+    topology.add_node(Node(node_id=connection.source_node_id, node_profile=node_profile))
+    topology.add_node(Node(node_id=connection.sink_node_id, node_profile=node_profile))
     topology.add_connection(connection)
 
     # act
diff --git a/shared/apply/apply.py b/shared/apply/apply.py
index 8a333aba..85289c00 100644
--- a/shared/apply/apply.py
+++ b/shared/apply/apply.py
@@ -22,11 +22,13 @@ from shared.types.events import (
     TopologyEdgeCreated,
     TopologyEdgeDeleted,
     TopologyEdgeReplacedAtomically,
+    TopologyNodeCreated,
     WorkerStatusUpdated,
 )
 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.worker.common import NodeStatus, RunnerId
 from shared.types.worker.instances import Instance, InstanceId, InstanceStatus
 from shared.types.worker.runners import RunnerStatus
@@ -122,6 +124,12 @@ def apply_worker_status_updated(event: WorkerStatusUpdated, state: State) -> Sta
 def apply_chunk_generated(event: ChunkGenerated, state: State) -> State:
     return state
 
+@event_apply.register(TopologyNodeCreated)
+def apply_topology_node_created(event: TopologyNodeCreated, state: State) -> State:
+    topology = copy.copy(state.topology)
+    topology.add_node(Node(node_id=event.node_id))
+    return state.model_copy(update={"topology": topology})
+
 @event_apply.register(TopologyEdgeCreated)
 def apply_topology_edge_created(event: TopologyEdgeCreated, state: State) -> State:
     topology = copy.copy(state.topology)
diff --git a/shared/topology.py b/shared/topology.py
index 0e40905d..52e2f9cd 100644
--- a/shared/topology.py
+++ b/shared/topology.py
@@ -49,19 +49,19 @@ class Topology(TopologyProto):
 
         for node in snapshot.nodes:
             with contextlib.suppress(ValueError):
-                topology.add_node(node, node.node_id)
+                topology.add_node(node)
 
         for connection in snapshot.connections:
             topology.add_connection(connection)
 
         return topology
 
-    def add_node(self, node: Node, node_id: NodeId) -> None:
-        if node_id in self._node_id_to_rx_id_map:
+    def add_node(self, node: Node) -> None:
+        if node.node_id in self._node_id_to_rx_id_map:
             raise ValueError("Node already exists")
         rx_id = self._graph.add_node(node)
-        self._node_id_to_rx_id_map[node_id] = rx_id
-        self._rx_id_to_node_id_map[rx_id] = node_id
+        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 add_connection(
@@ -69,9 +69,9 @@ class Topology(TopologyProto):
         connection: Connection,
     ) -> None:
         if connection.source_node_id not in self._node_id_to_rx_id_map:
-            self.add_node(Node(node_id=connection.source_node_id), node_id=connection.source_node_id)
+            self.add_node(Node(node_id=connection.source_node_id))
         if connection.sink_node_id not in self._node_id_to_rx_id_map:
-            self.add_node(Node(node_id=connection.sink_node_id), node_id=connection.sink_node_id)
+            self.add_node(Node(node_id=connection.sink_node_id))
 
         src_id = self._node_id_to_rx_id_map[connection.source_node_id]
         sink_id = self._node_id_to_rx_id_map[connection.sink_node_id]
diff --git a/shared/types/events/_events.py b/shared/types/events/_events.py
index e28f55c3..20d4c6c5 100644
--- a/shared/types/events/_events.py
+++ b/shared/types/events/_events.py
@@ -66,6 +66,7 @@ class _EventType(str, Enum):
     NodePerformanceMeasured = "NodePerformanceMeasured"
 
     # Topology Events
+    TopologyNodeCreated = "TopologyNodeCreated"
     TopologyEdgeCreated = "TopologyEdgeCreated"
     TopologyEdgeReplacedAtomically = "TopologyEdgeReplacedAtomically"
     TopologyEdgeDeleted = "TopologyEdgeDeleted"
@@ -166,6 +167,9 @@ class ChunkGenerated(_BaseEvent[_EventType.ChunkGenerated]):
     command_id: CommandId
     chunk: GenerationChunk
 
+class TopologyNodeCreated(_BaseEvent[_EventType.TopologyNodeCreated]):
+    event_type: Literal[_EventType.TopologyNodeCreated] = _EventType.TopologyNodeCreated
+    node_id: NodeId
 
 class TopologyEdgeCreated(_BaseEvent[_EventType.TopologyEdgeCreated]):
     event_type: Literal[_EventType.TopologyEdgeCreated] = _EventType.TopologyEdgeCreated
@@ -196,6 +200,7 @@ _Event = Union[
     NodePerformanceMeasured,
     WorkerStatusUpdated,
     ChunkGenerated,
+    TopologyNodeCreated,
     TopologyEdgeCreated,
     TopologyEdgeReplacedAtomically,
     TopologyEdgeDeleted,
diff --git a/shared/types/topology.py b/shared/types/topology.py
index 0dac5c08..f6e170af 100644
--- a/shared/types/topology.py
+++ b/shared/types/topology.py
@@ -41,7 +41,7 @@ class Node(BaseModel):
 
 
 class TopologyProto(Protocol):
-    def add_node(self, node: Node, node_id: NodeId) -> None: ...
+    def add_node(self, node: Node) -> None: ...
 
     def add_connection(
         self,

← 261e5752 Serialize topology  ·  back to Exo  ·  Inference Integration Test 93330f02 →