diff --git a/src/exo/main.py b/src/exo/main.py index 035c4c9a..7d340c6c 100644 --- a/src/exo/main.py +++ b/src/exo/main.py @@ -90,7 +90,6 @@ class Node: worker = Worker( node_id, session_id, - connection_message_receiver=router.receiver(topics.CONNECTION_MESSAGES), global_event_receiver=router.receiver(topics.GLOBAL_EVENTS), local_event_sender=router.sender(topics.LOCAL_EVENTS), command_sender=router.sender(topics.COMMANDS), @@ -227,9 +226,6 @@ class Node: self.worker = Worker( self.node_id, result.session_id, - connection_message_receiver=self.router.receiver( - topics.CONNECTION_MESSAGES - ), global_event_receiver=self.router.receiver( topics.GLOBAL_EVENTS ), diff --git a/src/exo/worker/main.py b/src/exo/worker/main.py index 2aea72be..78ad7050 100644 --- a/src/exo/worker/main.py +++ b/src/exo/worker/main.py @@ -7,7 +7,6 @@ from anyio import CancelScope, create_task_group, fail_after from anyio.abc import TaskGroup from loguru import logger -from exo.routing.connection_message import ConnectionMessage, ConnectionMessageType from exo.shared.apply import apply from exo.shared.models.model_cards import ModelId from exo.shared.types.api import ImageEditsTaskParams @@ -57,7 +56,6 @@ class Worker: node_id: NodeId, session_id: SessionId, *, - connection_message_receiver: Receiver[ConnectionMessage], global_event_receiver: Receiver[ForwarderEvent], local_event_sender: Sender[ForwarderEvent], # This is for requesting updates. It doesn't need to be a general command sender right now, @@ -74,7 +72,6 @@ class Worker: self.event_index_counter = event_index_counter self.command_sender = command_sender self.download_command_sender = download_command_sender - self.connection_message_receiver = connection_message_receiver self.event_buffer = OrderedBuffer[Event]() self.out_for_delivery: dict[EventId, ForwarderEvent] = {} @@ -105,7 +102,6 @@ class Worker: tg.start_soon(info_gatherer.run) tg.start_soon(self._forward_info, info_recv) tg.start_soon(self.plan_step) - tg.start_soon(self._connection_message_event_writer) tg.start_soon(self._resend_out_for_delivery) tg.start_soon(self._event_applier) tg.start_soon(self._forward_events) @@ -279,41 +275,6 @@ class Worker: instance = self.state.instances[task.instance_id] return instance.shard_assignments.node_to_runner[self.node_id] - async def _connection_message_event_writer(self): - with self.connection_message_receiver as connection_messages: - async for msg in connection_messages: - await self.event_sender.send( - self._convert_connection_message_to_event(msg) - ) - - def _convert_connection_message_to_event(self, msg: ConnectionMessage): - match msg.connection_type: - case ConnectionMessageType.Connected: - return TopologyEdgeCreated( - conn=Connection( - source=self.node_id, - sink=msg.node_id, - edge=SocketConnection( - sink_multiaddr=Multiaddr( - address=f"/ip4/{msg.remote_ipv4}/tcp/{msg.remote_tcp_port}" - ), - ), - ), - ) - - case ConnectionMessageType.Disconnected: - return TopologyEdgeDeleted( - conn=Connection( - source=self.node_id, - sink=msg.node_id, - edge=SocketConnection( - sink_multiaddr=Multiaddr( - address=f"/ip4/{msg.remote_ipv4}/tcp/{msg.remote_tcp_port}" - ), - ), - ), - ) - async def _nack_request(self, since_idx: int) -> None: # We request all events after (and including) the missing index. # This function is started whenever we receive an event that is out of sequence.