← back to Exo
fix disconnects
880a18d205694256671efe7069e821380d378a07 · 2025-12-15 15:23:13 +0000 · Evan Quiney
Co-authored-by: Ryuichi Leo Takashige <leo@exolabs.net>
Files touched
M src/exo/main.pyM src/exo/master/api.pyM src/exo/master/main.pyM src/exo/master/tests/test_master.pyM src/exo/routing/router.pyM src/exo/shared/apply.pyM src/exo/shared/election.pyM src/exo/shared/topology.pyM src/exo/shared/types/commands.pyM src/exo/shared/types/events.pyM src/exo/shared/types/state.pyM src/exo/worker/main.pyM src/exo/worker/utils/net_profile.py
Diff
commit 880a18d205694256671efe7069e821380d378a07
Author: Evan Quiney <evanev7@gmail.com>
Date: Mon Dec 15 15:23:13 2025 +0000
fix disconnects
Co-authored-by: Ryuichi Leo Takashige <leo@exolabs.net>
---
src/exo/main.py | 23 ++-------------
src/exo/master/api.py | 28 +++++++++---------
src/exo/master/main.py | 58 +++++++++++++++++++++----------------
src/exo/master/tests/test_master.py | 2 ++
src/exo/routing/router.py | 19 ++++--------
src/exo/shared/apply.py | 48 +++++++++++++++++++++++++++---
src/exo/shared/election.py | 5 ----
src/exo/shared/topology.py | 28 ++++++++++++++++++
src/exo/shared/types/commands.py | 5 ----
src/exo/shared/types/events.py | 17 +++++++----
src/exo/shared/types/state.py | 4 ++-
src/exo/worker/main.py | 55 +++++++++++++++++------------------
src/exo/worker/utils/net_profile.py | 2 +-
13 files changed, 174 insertions(+), 120 deletions(-)
diff --git a/src/exo/main.py b/src/exo/main.py
index 0f16d6c2..b859d2ce 100644
--- a/src/exo/main.py
+++ b/src/exo/main.py
@@ -1,7 +1,7 @@
import argparse
import multiprocessing as mp
import signal
-from dataclasses import dataclass
+from dataclasses import dataclass, field
from typing import Self
import anyio
@@ -16,7 +16,6 @@ from exo.routing.router import Router, get_node_id_keypair
from exo.shared.constants import EXO_LOG
from exo.shared.election import Election, ElectionResult
from exo.shared.logging import logger_cleanup, logger_setup
-from exo.shared.types.commands import KillCommand
from exo.shared.types.common import NodeId, SessionId
from exo.utils.channels import Receiver, channel
from exo.utils.pydantic_ext import CamelCaseModel
@@ -35,7 +34,7 @@ class Node:
api: API | None
node_id: NodeId
- _tg: TaskGroup | None = None
+ _tg: TaskGroup = field(init=False, default_factory=anyio.create_task_group)
@classmethod
async def create(cls, args: "Args") -> "Self":
@@ -66,7 +65,6 @@ class Node:
node_id,
session_id,
exo_shard_downloader(),
- initial_connection_messages=[],
connection_message_receiver=router.receiver(topics.CONNECTION_MESSAGES),
global_event_receiver=router.receiver(topics.GLOBAL_EVENTS),
local_event_sender=router.sender(topics.LOCAL_EVENTS),
@@ -99,9 +97,8 @@ class Node:
return cls(router, worker, election, er_recv, master, api, node_id)
async def run(self):
- async with anyio.create_task_group() as tg:
+ async with self._tg as tg:
signal.signal(signal.SIGINT, lambda _, __: self.shutdown())
- self._tg = tg
tg.start_soon(self.router.run)
tg.start_soon(self.worker.run)
tg.start_soon(self.election.run)
@@ -110,10 +107,8 @@ class Node:
if self.api:
tg.start_soon(self.api.run)
tg.start_soon(self._elect_loop)
- tg.start_soon(self._listen_for_kill_command)
def shutdown(self):
- assert self._tg
# if this is our second call to shutdown, just sys.exit
if self._tg.cancel_scope.cancel_called:
import sys
@@ -121,18 +116,7 @@ class Node:
sys.exit(1)
self._tg.cancel_scope.cancel()
- async def _listen_for_kill_command(self):
- assert self._tg
- with self.router.receiver(topics.COMMANDS) as commands:
- async for command in commands:
- match command.command:
- case KillCommand():
- self.shutdown()
- case _:
- pass
-
async def _elect_loop(self):
- assert self._tg
with self.election_result_receiver as results:
async for result in results:
# This function continues to have a lot of very specific entangled logic
@@ -187,7 +171,6 @@ class Node:
self.node_id,
result.session_id,
exo_shard_downloader(),
- initial_connection_messages=result.historic_messages,
connection_message_receiver=self.router.receiver(
topics.CONNECTION_MESSAGES
),
diff --git a/src/exo/master/api.py b/src/exo/master/api.py
index 9d65c7c1..172ae5c1 100644
--- a/src/exo/master/api.py
+++ b/src/exo/master/api.py
@@ -37,7 +37,6 @@ from exo.shared.types.commands import (
CreateInstance,
DeleteInstance,
ForwarderCommand,
- KillCommand,
TaskFinished,
)
from exo.shared.types.common import CommandId, NodeId, SessionId
@@ -92,7 +91,7 @@ class API:
# This lets us pause the API if an election is running
election_receiver: Receiver[ElectionMessage],
) -> None:
- self.state = State()
+ self._state = State()
self.command_sender = command_sender
self.global_event_receiver = global_event_receiver
self.election_receiver = election_receiver
@@ -127,13 +126,15 @@ class API:
self._tg: TaskGroup | None = None
def reset(self, new_session_id: SessionId, result_clock: int):
- self.state = State()
+ logger.info("Resetting API State")
+ self._state = State()
self.session_id = new_session_id
self.event_buffer = OrderedBuffer[Event]()
self._chat_completion_queues = {}
self.unpause(result_clock)
def unpause(self, result_clock: int):
+ logger.info("Unpausing API")
self.last_completed_election = result_clock
self.paused = False
self.paused_ev.set()
@@ -155,11 +156,10 @@ class API:
self.app.get("/models")(self.get_models)
self.app.get("/v1/models")(self.get_models)
self.app.post("/v1/chat/completions")(self.chat_completions)
- self.app.get("/state")(lambda: self.state)
- self.app.delete("/kill")(self.kill_exo)
+ self.app.get("/state")(self.state)
- async def kill_exo(self):
- await self._send(KillCommand())
+ async def state(self) -> State:
+ return self._state
async def create_instance(
self, payload: CreateInstanceTaskParams
@@ -189,12 +189,12 @@ class API:
)
def get_instance(self, instance_id: InstanceId) -> Instance:
- if instance_id not in self.state.instances:
+ if instance_id not in self._state.instances:
raise HTTPException(status_code=404, detail="Instance not found")
- return self.state.instances[instance_id]
+ return self._state.instances[instance_id]
async def delete_instance(self, instance_id: InstanceId) -> DeleteInstanceResponse:
- if instance_id not in self.state.instances:
+ if instance_id not in self._state.instances:
raise HTTPException(status_code=404, detail="Instance not found")
command = DeleteInstance(
@@ -261,7 +261,7 @@ class API:
if not any(
instance.shard_assignments.model_id == payload.model
- for instance in self.state.instances.values()
+ for instance in self._state.instances.values()
):
await self._trigger_notify_user_to_download_model(payload.model)
raise HTTPException(
@@ -281,7 +281,7 @@ class API:
"""Calculate total available memory across all nodes in bytes."""
total_available = Memory()
- for node in self.state.topology.list_nodes():
+ for node in self._state.topology.list_nodes():
if node.node_profile is not None:
total_available += node.node_profile.memory.ram_available
@@ -328,9 +328,11 @@ class API:
async def _apply_state(self):
with self.global_event_receiver as events:
async for f_event in events:
+ if f_event.origin != self.session_id.master_node_id:
+ continue
self.event_buffer.ingest(f_event.origin_idx, f_event.event)
for idx, event in self.event_buffer.drain_indexed():
- self.state = apply(self.state, IndexedEvent(event=event, idx=idx))
+ self._state = apply(self._state, IndexedEvent(event=event, idx=idx))
if (
isinstance(event, ChunkGenerated)
and event.command_id in self._chat_completion_queues
diff --git a/src/exo/master/main.py b/src/exo/master/main.py
index 5dadb5c3..149bfbd2 100644
--- a/src/exo/master/main.py
+++ b/src/exo/master/main.py
@@ -1,6 +1,6 @@
-from datetime import datetime, timezone
+from datetime import datetime, timedelta, timezone
-from anyio import create_task_group
+import anyio
from anyio.abc import TaskGroup
from loguru import logger
@@ -15,7 +15,6 @@ from exo.shared.types.commands import (
CreateInstance,
DeleteInstance,
ForwarderCommand,
- KillCommand,
RequestEventLog,
TaskFinished,
TestCommand,
@@ -26,9 +25,9 @@ from exo.shared.types.events import (
ForwarderEvent,
IndexedEvent,
InstanceDeleted,
+ NodeTimedOut,
TaskCreated,
TaskDeleted,
- TopologyEdgeDeleted,
)
from exo.shared.types.state import State
from exo.shared.types.tasks import (
@@ -59,7 +58,7 @@ class Master:
tb_only: bool = False,
):
self.state = State()
- self._tg: TaskGroup | None = None
+ self._tg: TaskGroup = anyio.create_task_group()
self.node_id = node_id
self.session_id = session_id
self.command_task_mapping: dict[CommandId, TaskId] = {}
@@ -80,11 +79,11 @@ class Master:
async def run(self):
logger.info("Starting Master")
- async with create_task_group() as tg:
- self._tg = tg
+ async with self._tg as tg:
tg.start_soon(self._event_processor)
tg.start_soon(self._command_processor)
tg.start_soon(self._loopback_processor)
+ tg.start_soon(self._plan)
self.global_event_sender.close()
self.local_event_receiver.close()
self.command_receiver.close()
@@ -92,9 +91,8 @@ class Master:
self._loopback_event_receiver.close()
async def shutdown(self):
- if self._tg:
- logger.info("Stopping Master")
- self._tg.cancel_scope.cancel()
+ logger.info("Stopping Master")
+ self._tg.cancel_scope.cancel()
async def _command_processor(self) -> None:
with self.command_receiver as commands:
@@ -104,7 +102,7 @@ class Master:
generated_events: list[Event] = []
command = forwarder_command.command
match command:
- case TestCommand() | KillCommand():
+ case TestCommand():
pass
case ChatCompletion():
instance_task_counts: dict[InstanceId, int] = {}
@@ -191,6 +189,30 @@ class Master:
except ValueError as e:
logger.opt(exception=e).warning("Error in command processor")
+ # These plan loops are the cracks showing in our event sourcing architecture - more things could be commands
+ async def _plan(self) -> None:
+ while True:
+ # kill broken instances
+ connected_node_ids = set(
+ [x.node_id for x in self.state.topology.list_nodes()]
+ )
+ for instance_id, instance in self.state.instances.items():
+ for node_id in instance.shard_assignments.node_to_runner:
+ if node_id not in connected_node_ids:
+ await self.event_sender.send(
+ InstanceDeleted(instance_id=instance_id)
+ )
+ break
+
+ # time out dead nodes
+ for node_id, time in self.state.last_seen.items():
+ now = datetime.now(tz=timezone.utc)
+ if now - time > timedelta(seconds=30):
+ logger.info(f"Manually removing node {node_id} due to inactivity")
+ await self.event_sender.send(NodeTimedOut(node_id=node_id))
+
+ await anyio.sleep(10)
+
async def _event_processor(self) -> None:
with self.local_event_receiver as local_events:
async for local_event in local_events:
@@ -209,23 +231,9 @@ class Master:
event._master_time_stamp = datetime.now(tz=timezone.utc) # pyright: ignore[reportPrivateUsage]
- # TODO: SQL <- What does this mean?
self._event_log.append(event)
await self._send_event(indexed)
- # TODO: This can be done in a better place. But for now, we use this to check if any running instances have been broken.
- if isinstance(event, TopologyEdgeDeleted):
- connected_node_ids = set(
- [x.node_id for x in self.state.topology.list_nodes()]
- )
- for instance_id, instance in self.state.instances.items():
- for node_id in instance.shard_assignments.node_to_runner:
- if node_id not in connected_node_ids:
- await self.event_sender.send(
- InstanceDeleted(instance_id=instance_id)
- )
- break
-
async def _loopback_processor(self) -> None:
# this would ideally not be necessary.
# this is WAY less hacky than how I was working around this before
diff --git a/src/exo/master/tests/test_master.py b/src/exo/master/tests/test_master.py
index a87abc34..948bcb1f 100644
--- a/src/exo/master/tests/test_master.py
+++ b/src/exo/master/tests/test_master.py
@@ -1,3 +1,4 @@
+from datetime import datetime, timezone
from typing import Sequence
import anyio
@@ -84,6 +85,7 @@ async def test_master():
session=session_id,
event=(
NodePerformanceMeasured(
+ when=str(datetime.now(tz=timezone.utc)),
node_id=node_id,
node_profile=NodePerformanceProfile(
model_id="maccy",
diff --git a/src/exo/routing/router.py b/src/exo/routing/router.py
index 21aece29..ac6073af 100644
--- a/src/exo/routing/router.py
+++ b/src/exo/routing/router.py
@@ -44,7 +44,7 @@ class TopicRouter[T: CamelCaseModel]:
self.senders: set[Sender[T]] = set()
send, recv = channel[T]()
self.receiver: Receiver[T] = recv
- self.temp_sender: Sender[T] | None = send
+ self._sender: Sender[T] = send
self.networking_sender: Sender[tuple[str, bytes]] = networking_sender
async def run(self):
@@ -68,8 +68,7 @@ class TopicRouter[T: CamelCaseModel]:
# Close all the things!
for sender in self.senders:
sender.close()
- if self.temp_sender:
- self.temp_sender.close()
+ self._sender.close()
self.receiver.close()
async def publish(self, item: T):
@@ -89,6 +88,9 @@ class TopicRouter[T: CamelCaseModel]:
async def publish_bytes(self, data: bytes):
await self.publish(self.topic.deserialize(data))
+ def new_sender(self) -> Sender[T]:
+ return self._sender.clone()
+
async def _send_out(self, item: T):
logger.trace(f"TopicRouter {self.topic.topic} sending {item}")
await self.networking_sender.send(
@@ -126,16 +128,7 @@ class Router:
# There's gotta be a way to do this without THIS many asserts
assert router is not None
assert router.topic == topic
- send: Sender[T] | None = cast(Sender[T] | None, router.temp_sender)
- if send:
- router.temp_sender = None
- return send
- try:
- sender = cast(Receiver[T], router.receiver).clone_sender()
- except ClosedResourceError:
- sender, router.receiver = cast(
- tuple[Sender[T], Receiver[CamelCaseModel]], channel[T]()
- )
+ sender = cast(TopicRouter[T], router).new_sender()
return sender
def receiver[T: CamelCaseModel](self, topic: TypedTopic[T]) -> Receiver[T]:
diff --git a/src/exo/shared/apply.py b/src/exo/shared/apply.py
index 178d2c5f..9bb597cb 100644
--- a/src/exo/shared/apply.py
+++ b/src/exo/shared/apply.py
@@ -1,5 +1,6 @@
import copy
from collections.abc import Mapping, Sequence
+from datetime import datetime
from loguru import logger
@@ -14,6 +15,7 @@ from exo.shared.types.events import (
NodeDownloadProgress,
NodeMemoryMeasured,
NodePerformanceMeasured,
+ NodeTimedOut,
RunnerDeleted,
RunnerStatusUpdated,
TaskAcknowledged,
@@ -45,6 +47,10 @@ def event_apply(event: Event, state: State) -> State:
return apply_instance_created(event, state)
case InstanceDeleted():
return apply_instance_deleted(event, state)
+ case NodeCreated():
+ return apply_topology_node_created(event, state)
+ case NodeTimedOut():
+ return apply_node_timed_out(event, state)
case NodePerformanceMeasured():
return apply_node_performance_measured(event, state)
case NodeDownloadProgress():
@@ -63,8 +69,6 @@ def event_apply(event: Event, state: State) -> State:
return apply_task_failed(event, state)
case TaskStatusUpdated():
return apply_task_status_updated(event, state)
- case NodeCreated():
- return apply_topology_node_created(event, state)
case TopologyEdgeCreated():
return apply_topology_edge_created(event, state)
case TopologyEdgeDeleted():
@@ -183,6 +187,24 @@ def apply_runner_deleted(event: RunnerDeleted, state: State) -> State:
return state.model_copy(update={"runners": new_runners})
+def apply_node_timed_out(event: NodeTimedOut, state: State) -> State:
+ topology = copy.copy(state.topology)
+ state.topology.remove_node(event.node_id)
+ node_profiles = {
+ key: value for key, value in state.node_profiles.items() if key != event.node_id
+ }
+ last_seen = {
+ key: value for key, value in state.last_seen.items() if key != event.node_id
+ }
+ return state.model_copy(
+ update={
+ "topology": topology,
+ "node_profiles": node_profiles,
+ "last_seen": last_seen,
+ }
+ )
+
+
def apply_node_performance_measured(
event: NodePerformanceMeasured, state: State
) -> State:
@@ -190,13 +212,23 @@ def apply_node_performance_measured(
**state.node_profiles,
event.node_id: event.node_profile,
}
+ last_seen: Mapping[NodeId, datetime] = {
+ **state.last_seen,
+ event.node_id: datetime.fromisoformat(event.when),
+ }
state = state.model_copy(update={"node_profiles": new_profiles})
topology = copy.copy(state.topology)
# TODO: NodeCreated
if not topology.contains_node(event.node_id):
topology.add_node(NodeInfo(node_id=event.node_id))
topology.update_node_profile(event.node_id, event.node_profile)
- return state.model_copy(update={"topology": topology})
+ return state.model_copy(
+ update={
+ "node_profiles": new_profiles,
+ "topology": topology,
+ "last_seen": last_seen,
+ }
+ )
def apply_node_memory_measured(event: NodeMemoryMeasured, state: State) -> State:
@@ -224,12 +256,20 @@ def apply_node_memory_measured(event: NodeMemoryMeasured, state: State) -> State
**state.node_profiles,
event.node_id: created,
}
+ last_seen: Mapping[NodeId, datetime] = {
+ **state.last_seen,
+ event.node_id: datetime.fromisoformat(event.when),
+ }
if not topology.contains_node(event.node_id):
topology.add_node(NodeInfo(node_id=event.node_id))
# TODO: NodeCreated
topology.update_node_profile(event.node_id, created)
return state.model_copy(
- update={"node_profiles": created_profiles, "topology": topology}
+ update={
+ "node_profiles": created_profiles,
+ "topology": topology,
+ "last_seen": last_seen,
+ }
)
updated = existing.model_copy(update={"memory": event.memory})
diff --git a/src/exo/shared/election.py b/src/exo/shared/election.py
index b4dc36b6..9d030d5e 100644
--- a/src/exo/shared/election.py
+++ b/src/exo/shared/election.py
@@ -44,7 +44,6 @@ class ElectionResult(CamelCaseModel):
session_id: SessionId
won_clock: int
is_new_master: bool
- historic_messages: list[ConnectionMessage]
class Election:
@@ -84,7 +83,6 @@ class Election:
self._campaign_cancel_scope: CancelScope | None = None
self._campaign_done: Event | None = None
self._tg: TaskGroup | None = None
- self._connection_messages: list[ConnectionMessage] = []
async def run(self):
logger.info("Starting Election")
@@ -121,7 +119,6 @@ class Election:
won_clock=em.clock,
session_id=em.proposed_session,
is_new_master=is_new_master,
- historic_messages=self._connection_messages,
)
)
@@ -188,8 +185,6 @@ class Election:
self._campaign, candidates, DEFAULT_ELECTION_TIMEOUT
)
logger.debug("Campaign started")
- self._connection_messages.append(first)
- self._connection_messages.extend(rest)
logger.debug("Connection message added")
async def _command_counter(self) -> None:
diff --git a/src/exo/shared/topology.py b/src/exo/shared/topology.py
index 7413161f..46419d72 100644
--- a/src/exo/shared/topology.py
+++ b/src/exo/shared/topology.py
@@ -55,6 +55,22 @@ class Topology:
and len(self._graph.neighbors(self._node_id_to_rx_id_map[node_id])) == 1
)
+ def neighbours(self, node_id: NodeId) -> list[NodeId]:
+ return [
+ self._rx_id_to_node_id_map[rx_id]
+ for rx_id in self._graph.neighbors(self._node_id_to_rx_id_map[node_id])
+ ]
+
+ def out_edges(self, node_id: NodeId) -> list[tuple[NodeId, Connection]]:
+ if node_id not in self._node_id_to_rx_id_map:
+ return []
+ return [
+ (self._rx_id_to_node_id_map[nid], conn)
+ for _, nid, conn in self._graph.out_edges(
+ self._node_id_to_rx_id_map[node_id]
+ )
+ ]
+
def contains_node(self, node_id: NodeId) -> bool:
return node_id in self._node_id_to_rx_id_map
@@ -112,6 +128,16 @@ class Topology:
return None
def remove_node(self, node_id: NodeId) -> None:
+ if node_id not in self._node_id_to_rx_id_map:
+ return
+
+ for connection in self.list_connections():
+ if (
+ connection.local_node_id == node_id
+ or connection.send_back_node_id == node_id
+ ):
+ self.remove_connection(connection)
+
rx_idx = self._node_id_to_rx_id_map[node_id]
self._graph.remove_node(rx_idx)
@@ -119,6 +145,8 @@ class Topology:
del self._rx_id_to_node_id_map[rx_idx]
def remove_connection(self, connection: Connection) -> None:
+ if connection not in self._edge_id_to_rx_id_map:
+ return
rx_idx = self._edge_id_to_rx_id_map[connection]
self._graph.remove_edge_from_index(rx_idx)
del self._edge_id_to_rx_id_map[connection]
diff --git a/src/exo/shared/types/commands.py b/src/exo/shared/types/commands.py
index 1ea4027a..0a584ff5 100644
--- a/src/exo/shared/types/commands.py
+++ b/src/exo/shared/types/commands.py
@@ -16,10 +16,6 @@ class TestCommand(BaseCommand):
__test__ = False
-class KillCommand(BaseCommand):
- pass
-
-
class ChatCompletion(BaseCommand):
request_params: ChatCompletionTaskParams
@@ -45,7 +41,6 @@ class RequestEventLog(BaseCommand):
Command = (
TestCommand
- | KillCommand
| RequestEventLog
| ChatCompletion
| CreateInstance
diff --git a/src/exo/shared/types/events.py b/src/exo/shared/types/events.py
index 7ad465d4..29b750ef 100644
--- a/src/exo/shared/types/events.py
+++ b/src/exo/shared/types/events.py
@@ -81,20 +81,26 @@ class NodeCreated(BaseEvent):
node_id: NodeId
-class NodePerformanceMeasured(BaseEvent):
+class NodeTimedOut(BaseEvent):
node_id: NodeId
- node_profile: NodePerformanceProfile
-class NodeDownloadProgress(BaseEvent):
- download_progress: DownloadProgress
+class NodePerformanceMeasured(BaseEvent):
+ node_id: NodeId
+ when: str # this is a manually cast datetime overrode by the master when the event is indexed, rather than the local time on the device
+ node_profile: NodePerformanceProfile
class NodeMemoryMeasured(BaseEvent):
node_id: NodeId
+ when: str # this is a manually cast datetime overrode by the master when the event is indexed, rather than the local time on the device
memory: MemoryPerformanceProfile
+class NodeDownloadProgress(BaseEvent):
+ download_progress: DownloadProgress
+
+
class ChunkGenerated(BaseEvent):
command_id: CommandId
chunk: GenerationChunk
@@ -119,11 +125,12 @@ Event = (
| InstanceDeleted
| RunnerStatusUpdated
| RunnerDeleted
+ | NodeCreated
+ | NodeTimedOut
| NodePerformanceMeasured
| NodeMemoryMeasured
| NodeDownloadProgress
| ChunkGenerated
- | NodeCreated
| TopologyEdgeCreated
| TopologyEdgeDeleted
)
diff --git a/src/exo/shared/types/state.py b/src/exo/shared/types/state.py
index efdb5bcb..58b14d2e 100644
--- a/src/exo/shared/types/state.py
+++ b/src/exo/shared/types/state.py
@@ -1,4 +1,5 @@
from collections.abc import Mapping, Sequence
+from datetime import datetime
from typing import Any, cast
from pydantic import ConfigDict, Field, field_serializer, field_validator
@@ -35,7 +36,8 @@ class State(CamelCaseModel):
downloads: Mapping[NodeId, Sequence[DownloadProgress]] = {}
tasks: Mapping[TaskId, Task] = {}
node_profiles: Mapping[NodeId, NodePerformanceProfile] = {}
- topology: Topology = Topology()
+ last_seen: Mapping[NodeId, datetime] = {}
+ topology: Topology = Field(default_factory=Topology)
last_event_applied_idx: int = Field(default=-1, ge=-1)
@field_serializer("topology", mode="plain")
diff --git a/src/exo/worker/main.py b/src/exo/worker/main.py
index 6028c2b4..a5c049dc 100644
--- a/src/exo/worker/main.py
+++ b/src/exo/worker/main.py
@@ -1,3 +1,4 @@
+from datetime import datetime, timezone
from random import random
import anyio
@@ -50,7 +51,7 @@ from exo.worker.download.shard_downloader import RepoDownloadProgress, ShardDown
from exo.worker.plan import plan
from exo.worker.runner.runner_supervisor import RunnerSupervisor
from exo.worker.utils import start_polling_memory_metrics, start_polling_node_metrics
-from exo.worker.utils.net_profile import connect_all
+from exo.worker.utils.net_profile import check_reachable
class Worker:
@@ -60,7 +61,6 @@ class Worker:
session_id: SessionId,
shard_downloader: ShardDownloader,
*,
- initial_connection_messages: list[ConnectionMessage],
connection_message_receiver: Receiver[ConnectionMessage],
global_event_receiver: Receiver[ForwarderEvent],
local_event_sender: Sender[ForwarderEvent],
@@ -80,7 +80,6 @@ class Worker:
self.command_sender = command_sender
self.connection_message_receiver = connection_message_receiver
self.event_buffer = OrderedBuffer[Event]()
- self._initial_connection_messages = initial_connection_messages
self.out_for_delivery: dict[EventId, ForwarderEvent] = {}
self.state: State = State()
@@ -104,7 +103,9 @@ class Worker:
) -> None:
await self.event_sender.send(
NodePerformanceMeasured(
- node_id=self.node_id, node_profile=node_performance_profile
+ node_id=self.node_id,
+ node_profile=node_performance_profile,
+ when=str(datetime.now(tz=timezone.utc)),
),
)
@@ -112,7 +113,11 @@ class Worker:
memory_profile: MemoryPerformanceProfile,
) -> None:
await self.event_sender.send(
- NodeMemoryMeasured(node_id=self.node_id, memory=memory_profile)
+ NodeMemoryMeasured(
+ node_id=self.node_id,
+ memory=memory_profile,
+ when=str(datetime.now(tz=timezone.utc)),
+ )
)
# END CLEANUP
@@ -128,12 +133,6 @@ class Worker:
tg.start_soon(self._event_applier)
tg.start_soon(self._forward_events)
tg.start_soon(self._poll_connection_updates)
- # TODO: This is a little gross, but not too bad
- for msg in self._initial_connection_messages:
- await self.event_sender.send(
- self._convert_connection_message_to_event(msg)
- )
- self._initial_connection_messages = []
# Actual shutdown code - waits for all tasks to complete before executing.
self.local_event_sender.close()
@@ -143,9 +142,11 @@ class Worker:
async def _event_applier(self):
with self.global_event_receiver as events:
- async for event in events:
- self.event_buffer.ingest(event.origin_idx, event.event)
- event_id = event.event.event_id
+ async for f_event in events:
+ if f_event.origin != self.session_id.master_node_id:
+ continue
+ self.event_buffer.ingest(f_event.origin_idx, f_event.event)
+ event_id = f_event.event.event_id
if event_id in self.out_for_delivery:
del self.out_for_delivery[event_id]
@@ -167,15 +168,8 @@ class Worker:
elif indexed_events and self._nack_cancel_scope:
self._nack_cancel_scope.cancel()
- flag = False
for idx, event in indexed_events:
self.state = apply(self.state, IndexedEvent(idx=idx, event=event))
- if event_relevant_to_worker(event, self):
- flag = True
-
- # 3. If we've found a "relevant" event, run a plan -> op -> execute cycle.
- if flag:
- pass
async def plan_step(self):
while True:
@@ -420,23 +414,28 @@ class Worker:
while True:
# TODO: EdgeDeleted
edges = set(self.state.topology.list_connections())
- conns = await connect_all(self.state.topology)
+ conns = await check_reachable(self.state.topology)
for nid in conns:
for ip in conns[nid]:
edge = Connection(
local_node_id=self.node_id,
send_back_node_id=nid,
+ # nonsense multiaddr
send_back_multiaddr=Multiaddr(address=f"/ip4/{ip}/tcp/8000")
if "." in ip
+ # nonsense multiaddr
else Multiaddr(address=f"/ip6/{ip}/tcp/8000"),
)
if edge not in edges:
- logger.debug(f"manually discovered {edge=}")
+ logger.debug(f"ping discovered {edge=}")
await self.event_sender.send(TopologyEdgeCreated(edge=edge))
- await anyio.sleep(10)
-
+ for nid, conn in self.state.topology.out_edges(self.node_id):
+ if (
+ nid not in conns
+ or conn.send_back_multiaddr.ip_address not in conns.get(nid, set())
+ ):
+ logger.debug(f"ping failed to discover {conn=}")
+ await self.event_sender.send(TopologyEdgeDeleted(edge=conn))
-def event_relevant_to_worker(event: Event, worker: Worker):
- # TODO
- return True
+ await anyio.sleep(10)
diff --git a/src/exo/worker/utils/net_profile.py b/src/exo/worker/utils/net_profile.py
index 923048b0..1c8c5fe4 100644
--- a/src/exo/worker/utils/net_profile.py
+++ b/src/exo/worker/utils/net_profile.py
@@ -27,7 +27,7 @@ async def check_reachability(
out[target_node_id].add(target_ip)
-async def connect_all(topology: Topology) -> dict[NodeId, set[str]]:
+async def check_reachable(topology: Topology) -> dict[NodeId, set[str]]:
reachable: dict[NodeId, set[str]] = {}
async with create_task_group() as tg:
for node in topology.list_nodes():
← 70298ce0 Negative index nack request
·
back to Exo
·
backport the dashboard to staging 09593c5e →