committed by
GitHub
co-authored by
Gelu Vrabie
parent
261e575262
commit
2e4635a8f5
+14
-3
@@ -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__":
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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]
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user