add node started event

Co-authored-by: Gelu Vrabie <[email protected]>
This commit is contained in:
Gelu Vrabie
2025-07-26 19:12:26 +01:00
committed by GitHub
co-authored by Gelu Vrabie
parent 261e575262
commit 2e4635a8f5
9 changed files with 76 additions and 50 deletions
+14 -3
View File
@@ -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__":
+9 -7
View File
@@ -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",
+3 -3
View File
@@ -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))
+13 -13
View File
@@ -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))
+16 -16
View File
@@ -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
+8
View File
@@ -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)
+7 -7
View File
@@ -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]
+5
View File
@@ -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,
+1 -1
View File
@@ -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,