← back to Exo
add system ids (#1536)
dc8d42b4dc7586465b55b9ae9812ab35104cbe65 · 2026-02-20 15:41:59 +0000 · Evan Quiney
addresses some election edge cases where a new worker with an old master
would get stuck on the old workers buffer index - we now use new system
ids each time we instantiate a node, and each event-producing system has
a unique system id for its lifespan (until the master moves).
Files touched
M src/exo/download/coordinator.pyA src/exo/download/tests/test_coordinator_ack.pyM src/exo/main.pyM src/exo/master/api.pyM src/exo/master/main.pyM src/exo/master/tests/test_master.pyM src/exo/routing/topics.pyM src/exo/shared/tests/test_election.pyM src/exo/shared/types/commands.pyM src/exo/shared/types/common.pyM src/exo/shared/types/events.pyM src/exo/worker/main.py
Diff
commit dc8d42b4dc7586465b55b9ae9812ab35104cbe65
Author: Evan Quiney <evanev7@gmail.com>
Date: Fri Feb 20 15:41:59 2026 +0000
add system ids (#1536)
addresses some election edge cases where a new worker with an old master
would get stuck on the old workers buffer index - we now use new system
ids each time we instantiate a node, and each event-producing system has
a unique system id for its lifespan (until the master moves).
---
src/exo/download/coordinator.py | 43 ++++++++---
src/exo/download/tests/test_coordinator_ack.py | 98 ++++++++++++++++++++++++++
src/exo/main.py | 22 +++---
src/exo/master/api.py | 21 +++---
src/exo/master/main.py | 22 +++---
src/exo/master/tests/test_master.py | 18 ++---
src/exo/routing/topics.py | 7 +-
src/exo/shared/tests/test_election.py | 4 +-
src/exo/shared/types/commands.py | 6 +-
src/exo/shared/types/common.py | 4 ++
src/exo/shared/types/events.py | 13 +++-
src/exo/worker/main.py | 29 ++++----
12 files changed, 214 insertions(+), 73 deletions(-)
diff --git a/src/exo/download/coordinator.py b/src/exo/download/coordinator.py
index f2b44495..0aa2374a 100644
--- a/src/exo/download/coordinator.py
+++ b/src/exo/download/coordinator.py
@@ -1,7 +1,7 @@
import asyncio
import socket
from dataclasses import dataclass, field
-from typing import Iterator
+from random import random
import anyio
from anyio import current_time
@@ -22,10 +22,13 @@ from exo.shared.types.commands import (
ForwarderDownloadCommand,
StartDownload,
)
-from exo.shared.types.common import NodeId, SessionId
+from exo.shared.types.common import NodeId, SessionId, SystemId
from exo.shared.types.events import (
Event,
- ForwarderEvent,
+ EventId,
+ # TODO(evan): just for acks, should delete this ASAP
+ GlobalForwarderEvent,
+ LocalForwarderEvent,
NodeDownloadProgress,
)
from exo.shared.types.worker.downloads import (
@@ -45,10 +48,16 @@ class DownloadCoordinator:
session_id: SessionId
shard_downloader: ShardDownloader
download_command_receiver: Receiver[ForwarderDownloadCommand]
- local_event_sender: Sender[ForwarderEvent]
- event_index_counter: Iterator[int]
+ local_event_sender: Sender[LocalForwarderEvent]
+
+ # ack stuff
+ _global_event_receiver: Receiver[GlobalForwarderEvent]
+ _out_for_delivery: dict[EventId, LocalForwarderEvent] = field(default_factory=dict)
+
offline: bool = False
+ _system_id: SystemId = field(default_factory=SystemId)
+
# Local state
download_status: dict[ModelId, DownloadProgress] = field(default_factory=dict)
active_downloads: dict[ModelId, asyncio.Task[None]] = field(default_factory=dict)
@@ -119,6 +128,8 @@ class DownloadCoordinator:
tg.start_soon(self._command_processor)
tg.start_soon(self._forward_events)
tg.start_soon(self._emit_existing_download_progress)
+ tg.start_soon(self._resend_out_for_delivery)
+ tg.start_soon(self._clear_ofd)
if not self.offline:
tg.start_soon(self._check_internet_connection)
@@ -153,6 +164,20 @@ class DownloadCoordinator:
def shutdown(self) -> None:
self._tg.cancel_scope.cancel()
+ # directly copied from worker
+ async def _resend_out_for_delivery(self) -> None:
+ # This can also be massively tightened, we should check events are at least a certain age before resending.
+ # Exponential backoff would also certainly help here.
+ while True:
+ await anyio.sleep(1 + random())
+ for event in self._out_for_delivery.copy().values():
+ await self.local_event_sender.send(event)
+
+ async def _clear_ofd(self) -> None:
+ with self._global_event_receiver as events:
+ async for event in events:
+ self._out_for_delivery.pop(event.event.event_id, None)
+
async def _command_processor(self) -> None:
with self.download_command_receiver as commands:
async for cmd in commands:
@@ -298,19 +323,21 @@ class DownloadCoordinator:
del self.download_status[model_id]
async def _forward_events(self) -> None:
+ idx = 0
with self.event_receiver as events:
async for event in events:
- idx = next(self.event_index_counter)
- fe = ForwarderEvent(
+ fe = LocalForwarderEvent(
origin_idx=idx,
- origin=self.node_id,
+ origin=self._system_id,
session=self.session_id,
event=event,
)
+ idx += 1
logger.debug(
f"DownloadCoordinator published event {idx}: {str(event)[:100]}"
)
await self.local_event_sender.send(fe)
+ self._out_for_delivery[event.event_id] = fe
async def _emit_existing_download_progress(self) -> None:
try:
diff --git a/src/exo/download/tests/test_coordinator_ack.py b/src/exo/download/tests/test_coordinator_ack.py
new file mode 100644
index 00000000..e99ec912
--- /dev/null
+++ b/src/exo/download/tests/test_coordinator_ack.py
@@ -0,0 +1,98 @@
+from typing import Any
+
+import anyio
+import pytest
+
+from exo.download.coordinator import DownloadCoordinator
+from exo.download.shard_downloader import NoopShardDownloader
+from exo.shared.models.model_cards import ModelCard, ModelTask
+from exo.shared.types.common import ModelId, NodeId, SessionId
+from exo.shared.types.events import (
+ GlobalForwarderEvent,
+ LocalForwarderEvent,
+ NodeDownloadProgress,
+)
+from exo.shared.types.memory import Memory
+from exo.shared.types.worker.downloads import (
+ DownloadPending,
+)
+from exo.shared.types.worker.shards import PipelineShardMetadata
+from exo.utils.channels import channel
+
+# Use the built‑in NoopShardDownloader directly – it already implements the required abstract interface.
+# No additional subclass is needed for this test.
+
+
+@pytest.mark.anyio
+async def test_ack_behaviour():
+ # Create channels (type Any for simplicity)
+ _, command_receiver = channel[Any]()
+ local_sender, _ = channel[Any]()
+ global_sender, global_receiver = channel[Any]()
+
+ # Minimal identifiers
+ node_id = NodeId()
+ session_id = SessionId(master_node_id=node_id, election_clock=0)
+
+ # Create a dummy model card and shard metadata
+ model_id = ModelId("test/model")
+ model_card = ModelCard(
+ model_id=model_id,
+ storage_size=Memory.from_bytes(0),
+ n_layers=1,
+ hidden_size=1,
+ supports_tensor=True,
+ tasks=[ModelTask.TextGeneration],
+ )
+ shard = PipelineShardMetadata(
+ model_card=model_card,
+ device_rank=0,
+ world_size=1,
+ start_layer=0,
+ end_layer=1,
+ n_layers=1,
+ )
+
+ # Instantiate the coordinator with the dummy downloader
+ coord = DownloadCoordinator(
+ node_id=node_id,
+ session_id=session_id,
+ shard_downloader=NoopShardDownloader(),
+ download_command_receiver=command_receiver,
+ local_event_sender=local_sender,
+ _global_event_receiver=global_receiver,
+ )
+
+ async with anyio.create_task_group() as tg:
+ # Start the forwarding and ack‑clearing loops
+ tg.start_soon(coord._forward_events) # pyright: ignore[reportPrivateUsage]
+ tg.start_soon(coord._clear_ofd) # pyright: ignore[reportPrivateUsage]
+
+ # Send a pending download progress event via the internal event sender
+ pending = DownloadPending(
+ node_id=node_id,
+ shard_metadata=shard,
+ model_directory="/tmp/model",
+ )
+ await coord.event_sender.send(NodeDownloadProgress(download_progress=pending))
+ # Allow the forwarder to process the event
+ await anyio.sleep(0.1)
+
+ # There should be exactly one entry awaiting ACK
+ assert len(coord._out_for_delivery) == 1 # pyright: ignore[reportPrivateUsage]
+ # Retrieve the stored LocalForwarderEvent
+ stored_fe: LocalForwarderEvent = next(iter(coord._out_for_delivery.values())) # pyright: ignore[reportPrivateUsage]
+ # Simulate receiving a global ack for this event
+ ack = GlobalForwarderEvent(
+ origin_idx=0,
+ origin=node_id,
+ session=session_id,
+ event=stored_fe.event,
+ )
+ await global_sender.send(ack)
+ # Give the clear‑ofd task a moment to process the ack
+ await anyio.sleep(0.1)
+ # The out‑for‑delivery map should now be empty
+ assert len(coord._out_for_delivery) == 0 # pyright: ignore[reportPrivateUsage]
+ # Cancel background tasks
+ tg.cancel_scope.cancel()
diff --git a/src/exo/main.py b/src/exo/main.py
index 27c78165..60729cc6 100644
--- a/src/exo/main.py
+++ b/src/exo/main.py
@@ -1,11 +1,10 @@
import argparse
-import itertools
import multiprocessing as mp
import os
import resource
import signal
from dataclasses import dataclass, field
-from typing import Iterator, Self
+from typing import Self
import anyio
from anyio.abc import TaskGroup
@@ -38,12 +37,11 @@ class Node:
api: API | None
node_id: NodeId
- event_index_counter: Iterator[int]
offline: bool
_tg: TaskGroup = field(init=False, default_factory=anyio.create_task_group)
@classmethod
- async def create(cls, args: "Args") -> "Self":
+ async def create(cls, args: "Args") -> Self:
keypair = get_node_id_keypair()
node_id = NodeId(keypair.to_node_id())
session_id = SessionId(master_node_id=node_id, election_clock=0)
@@ -57,9 +55,6 @@ class Node:
logger.info(f"Starting node {node_id}")
- # Create shared event index counter for Worker and DownloadCoordinator
- event_index_counter = itertools.count()
-
# Create DownloadCoordinator (unless --no-downloads)
if not args.no_downloads:
download_coordinator = DownloadCoordinator(
@@ -68,8 +63,9 @@ class Node:
exo_shard_downloader(),
download_command_receiver=router.receiver(topics.DOWNLOAD_COMMANDS),
local_event_sender=router.sender(topics.LOCAL_EVENTS),
- event_index_counter=event_index_counter,
offline=args.offline,
+ # TODO(evan): remove
+ _global_event_receiver=router.receiver(topics.GLOBAL_EVENTS),
)
else:
download_coordinator = None
@@ -95,7 +91,6 @@ class Node:
local_event_sender=router.sender(topics.LOCAL_EVENTS),
command_sender=router.sender(topics.COMMANDS),
download_command_sender=router.sender(topics.DOWNLOAD_COMMANDS),
- event_index_counter=event_index_counter,
)
else:
worker = None
@@ -133,7 +128,6 @@ class Node:
master,
api,
node_id,
- event_index_counter,
args.offline,
)
@@ -212,8 +206,6 @@ class Node:
)
if result.is_new_master:
await anyio.sleep(0)
- # Fresh counter for new session (buffer expects indices from 0)
- self.event_index_counter = itertools.count()
if self.download_coordinator:
self.download_coordinator.shutdown()
self.download_coordinator = DownloadCoordinator(
@@ -224,8 +216,11 @@ class Node:
topics.DOWNLOAD_COMMANDS
),
local_event_sender=self.router.sender(topics.LOCAL_EVENTS),
- event_index_counter=self.event_index_counter,
offline=self.offline,
+ # TODO(evan): remove
+ _global_event_receiver=self.router.receiver(
+ topics.GLOBAL_EVENTS
+ ),
)
self._tg.start_soon(self.download_coordinator.run)
if self.worker:
@@ -242,7 +237,6 @@ class Node:
download_command_sender=self.router.sender(
topics.DOWNLOAD_COMMANDS
),
- event_index_counter=self.event_index_counter,
)
self._tg.start_soon(self.worker.run)
if self.api:
diff --git a/src/exo/master/api.py b/src/exo/master/api.py
index e23ed1d0..15d1e142 100644
--- a/src/exo/master/api.py
+++ b/src/exo/master/api.py
@@ -140,11 +140,11 @@ from exo.shared.types.commands import (
TaskFinished,
TextGeneration,
)
-from exo.shared.types.common import CommandId, Id, NodeId, SessionId
+from exo.shared.types.common import CommandId, Id, NodeId, SessionId, SystemId
from exo.shared.types.events import (
ChunkGenerated,
Event,
- ForwarderEvent,
+ GlobalForwarderEvent,
IndexedEvent,
TracesMerged,
)
@@ -197,8 +197,7 @@ class API:
session_id: SessionId,
*,
port: int,
- # Ideally this would be a MasterForwarderEvent but type system says no :(
- global_event_receiver: Receiver[ForwarderEvent],
+ global_event_receiver: Receiver[GlobalForwarderEvent],
command_sender: Sender[ForwarderCommand],
download_command_sender: Sender[ForwarderDownloadCommand],
# This lets us pause the API if an election is running
@@ -206,6 +205,7 @@ class API:
) -> None:
self.state = State()
self._event_log = DiskEventLog(_API_EVENT_LOG_DIR)
+ self._system_id = SystemId()
self.command_sender = command_sender
self.download_command_sender = download_command_sender
self.global_event_receiver = global_event_receiver
@@ -257,6 +257,7 @@ class API:
self._event_log.close()
self._event_log = DiskEventLog(_API_EVENT_LOG_DIR)
self.state = State()
+ self._system_id = SystemId()
self.session_id = new_session_id
self.event_buffer = OrderedBuffer[Event]()
self._text_generation_queues = {}
@@ -589,7 +590,7 @@ class API:
command = TaskCancelled(cancelled_command_id=command_id)
with anyio.CancelScope(shield=True):
await self.command_sender.send(
- ForwarderCommand(origin=self.node_id, command=command)
+ ForwarderCommand(origin=self._system_id, command=command)
)
raise
finally:
@@ -937,7 +938,7 @@ class API:
command = TaskCancelled(cancelled_command_id=command_id)
with anyio.CancelScope(shield=True):
await self.command_sender.send(
- ForwarderCommand(origin=self.node_id, command=command)
+ ForwarderCommand(origin=self._system_id, command=command)
)
raise
finally:
@@ -1023,7 +1024,7 @@ class API:
command = TaskCancelled(cancelled_command_id=command_id)
with anyio.CancelScope(shield=True):
await self.command_sender.send(
- ForwarderCommand(origin=self.node_id, command=command)
+ ForwarderCommand(origin=self._system_id, command=command)
)
raise
finally:
@@ -1621,6 +1622,8 @@ class API:
async def _apply_state(self):
with self.global_event_receiver as events:
async for f_event in events:
+ if f_event.session != self.session_id:
+ continue
if f_event.origin != self.session_id.master_node_id:
continue
self.event_buffer.ingest(f_event.origin_idx, f_event.event)
@@ -1684,12 +1687,12 @@ class API:
while self.paused:
await self.paused_ev.wait()
await self.command_sender.send(
- ForwarderCommand(origin=self.node_id, command=command)
+ ForwarderCommand(origin=self._system_id, command=command)
)
async def _send_download(self, command: DownloadCommand):
await self.download_command_sender.send(
- ForwarderDownloadCommand(origin=self.node_id, command=command)
+ ForwarderDownloadCommand(origin=self._system_id, command=command)
)
async def start_download(
diff --git a/src/exo/master/main.py b/src/exo/master/main.py
index 9c7cf578..196fae99 100644
--- a/src/exo/master/main.py
+++ b/src/exo/master/main.py
@@ -29,13 +29,14 @@ from exo.shared.types.commands import (
TestCommand,
TextGeneration,
)
-from exo.shared.types.common import CommandId, NodeId, SessionId
+from exo.shared.types.common import CommandId, NodeId, SessionId, SystemId
from exo.shared.types.events import (
Event,
- ForwarderEvent,
+ GlobalForwarderEvent,
IndexedEvent,
InputChunkReceived,
InstanceDeleted,
+ LocalForwarderEvent,
NodeGatheredInfo,
NodeTimedOut,
TaskCreated,
@@ -71,8 +72,8 @@ class Master:
session_id: SessionId,
*,
command_receiver: Receiver[ForwarderCommand],
- local_event_receiver: Receiver[ForwarderEvent],
- global_event_sender: Sender[ForwarderEvent],
+ local_event_receiver: Receiver[LocalForwarderEvent],
+ global_event_sender: Sender[GlobalForwarderEvent],
download_command_sender: Sender[ForwarderDownloadCommand],
):
self.state = State()
@@ -87,10 +88,11 @@ class Master:
send, recv = channel[Event]()
self.event_sender: Sender[Event] = send
self._loopback_event_receiver: Receiver[Event] = recv
- self._loopback_event_sender: Sender[ForwarderEvent] = (
+ self._loopback_event_sender: Sender[LocalForwarderEvent] = (
local_event_receiver.clone_sender()
)
- self._multi_buffer = MultiSourceBuffer[NodeId, Event]()
+ self._system_id = SystemId()
+ self._multi_buffer = MultiSourceBuffer[SystemId, Event]()
self._event_log = DiskEventLog(EXO_EVENT_LOG_DIR / "master")
self._pending_traces: dict[TaskId, dict[int, list[TraceEventData]]] = {}
self._expected_ranks: dict[TaskId, set[int]] = {}
@@ -288,7 +290,7 @@ class Master:
):
await self.download_command_sender.send(
ForwarderDownloadCommand(
- origin=self.node_id, command=cmd
+ origin=self._system_id, command=cmd
)
)
generated_events.extend(transition_events)
@@ -414,8 +416,8 @@ class Master:
with self._loopback_event_receiver as events:
async for event in events:
await self._loopback_event_sender.send(
- ForwarderEvent(
- origin=NodeId(f"master_{self.node_id}"),
+ LocalForwarderEvent(
+ origin=self._system_id,
origin_idx=local_index,
session=self.session_id,
event=event,
@@ -427,7 +429,7 @@ class Master:
async def _send_event(self, event: IndexedEvent):
# Convenience method since this line is ugly
await self.global_event_sender.send(
- ForwarderEvent(
+ GlobalForwarderEvent(
origin=self.node_id,
origin_idx=event.idx,
session=self.session_id,
diff --git a/src/exo/master/tests/test_master.py b/src/exo/master/tests/test_master.py
index fcf71ee4..21fabb21 100644
--- a/src/exo/master/tests/test_master.py
+++ b/src/exo/master/tests/test_master.py
@@ -15,11 +15,12 @@ from exo.shared.types.commands import (
PlaceInstance,
TextGeneration,
)
-from exo.shared.types.common import ModelId, NodeId, SessionId
+from exo.shared.types.common import ModelId, NodeId, SessionId, SystemId
from exo.shared.types.events import (
- ForwarderEvent,
+ GlobalForwarderEvent,
IndexedEvent,
InstanceCreated,
+ LocalForwarderEvent,
NodeGatheredInfo,
TaskCreated,
)
@@ -45,9 +46,9 @@ async def test_master():
node_id = NodeId(keypair.to_node_id())
session_id = SessionId(master_node_id=node_id, election_clock=0)
- ge_sender, global_event_receiver = channel[ForwarderEvent]()
+ ge_sender, global_event_receiver = channel[GlobalForwarderEvent]()
command_sender, co_receiver = channel[ForwarderCommand]()
- local_event_sender, le_receiver = channel[ForwarderEvent]()
+ local_event_sender, le_receiver = channel[LocalForwarderEvent]()
fcds, _fcdr = channel[ForwarderDownloadCommand]()
all_events: list[IndexedEvent] = []
@@ -75,13 +76,12 @@ async def test_master():
async with anyio.create_task_group() as tg:
tg.start_soon(master.run)
- sender_node_id = NodeId(f"{keypair.to_node_id()}_sender")
# inject a NodeGatheredInfo event
logger.info("inject a NodeGatheredInfo event")
await local_event_sender.send(
- ForwarderEvent(
+ LocalForwarderEvent(
origin_idx=0,
- origin=sender_node_id,
+ origin=SystemId("Worker"),
session=session_id,
event=(
NodeGatheredInfo(
@@ -108,7 +108,7 @@ async def test_master():
logger.info("inject a CreateInstance Command")
await command_sender.send(
ForwarderCommand(
- origin=node_id,
+ origin=SystemId("API"),
command=(
PlaceInstance(
command_id=CommandId(),
@@ -133,7 +133,7 @@ async def test_master():
logger.info("inject a TextGeneration Command")
await command_sender.send(
ForwarderCommand(
- origin=node_id,
+ origin=SystemId("API"),
command=(
TextGeneration(
command_id=CommandId(),
diff --git a/src/exo/routing/topics.py b/src/exo/routing/topics.py
index 5a122a95..5d95a2a1 100644
--- a/src/exo/routing/topics.py
+++ b/src/exo/routing/topics.py
@@ -5,7 +5,8 @@ from exo.routing.connection_message import ConnectionMessage
from exo.shared.election import ElectionMessage
from exo.shared.types.commands import ForwarderCommand, ForwarderDownloadCommand
from exo.shared.types.events import (
- ForwarderEvent,
+ GlobalForwarderEvent,
+ LocalForwarderEvent,
)
from exo.utils.pydantic_ext import CamelCaseModel
@@ -36,8 +37,8 @@ class TypedTopic[T: CamelCaseModel]:
return self.model_type.model_validate_json(b.decode("utf-8"))
-GLOBAL_EVENTS = TypedTopic("global_events", PublishPolicy.Always, ForwarderEvent)
-LOCAL_EVENTS = TypedTopic("local_events", PublishPolicy.Always, ForwarderEvent)
+GLOBAL_EVENTS = TypedTopic("global_events", PublishPolicy.Always, GlobalForwarderEvent)
+LOCAL_EVENTS = TypedTopic("local_events", PublishPolicy.Always, LocalForwarderEvent)
COMMANDS = TypedTopic("commands", PublishPolicy.Always, ForwarderCommand)
ELECTION_MESSAGES = TypedTopic(
"election_messages", PublishPolicy.Always, ElectionMessage
diff --git a/src/exo/shared/tests/test_election.py b/src/exo/shared/tests/test_election.py
index 49550601..cc84deec 100644
--- a/src/exo/shared/tests/test_election.py
+++ b/src/exo/shared/tests/test_election.py
@@ -4,7 +4,7 @@ from anyio import create_task_group, fail_after, move_on_after
from exo.routing.connection_message import ConnectionMessage, ConnectionMessageType
from exo.shared.election import Election, ElectionMessage, ElectionResult
from exo.shared.types.commands import ForwarderCommand, TestCommand
-from exo.shared.types.common import NodeId, SessionId
+from exo.shared.types.common import NodeId, SessionId, SystemId
from exo.utils.channels import channel
# ======= #
@@ -384,7 +384,7 @@ async def test_tie_breaker_prefers_node_with_more_commands_seen() -> None:
# Pump local commands so our commands_seen is high before the round starts
for _ in range(50):
await co_tx.send(
- ForwarderCommand(origin=NodeId("SOMEONE"), command=TestCommand())
+ ForwarderCommand(origin=SystemId("SOMEONE"), command=TestCommand())
)
# Trigger a round at clock=1 with a peer of equal seniority but fewer commands
diff --git a/src/exo/shared/types/commands.py b/src/exo/shared/types/commands.py
index 09c135aa..e1c37288 100644
--- a/src/exo/shared/types/commands.py
+++ b/src/exo/shared/types/commands.py
@@ -6,7 +6,7 @@ from exo.shared.types.api import (
ImageGenerationTaskParams,
)
from exo.shared.types.chunks import InputImageChunk
-from exo.shared.types.common import CommandId, NodeId
+from exo.shared.types.common import CommandId, NodeId, SystemId
from exo.shared.types.text_generation import TextGenerationTaskParams
from exo.shared.types.worker.instances import Instance, InstanceId, InstanceMeta
from exo.shared.types.worker.shards import Sharding, ShardMetadata
@@ -100,10 +100,10 @@ Command = (
class ForwarderCommand(CamelCaseModel):
- origin: NodeId
+ origin: SystemId
command: Command
class ForwarderDownloadCommand(CamelCaseModel):
- origin: NodeId
+ origin: SystemId
command: DownloadCommand
diff --git a/src/exo/shared/types/common.py b/src/exo/shared/types/common.py
index 5db51cef..b5674b4b 100644
--- a/src/exo/shared/types/common.py
+++ b/src/exo/shared/types/common.py
@@ -25,6 +25,10 @@ class NodeId(Id):
pass
+class SystemId(Id):
+ pass
+
+
class ModelId(Id):
def normalize(self) -> str:
return self.replace("/", "--")
diff --git a/src/exo/shared/types/events.py b/src/exo/shared/types/events.py
index 5cf93d0c..c7ad2574 100644
--- a/src/exo/shared/types/events.py
+++ b/src/exo/shared/types/events.py
@@ -5,7 +5,7 @@ from pydantic import Field
from exo.shared.topology import Connection
from exo.shared.types.chunks import GenerationChunk, InputImageChunk
-from exo.shared.types.common import CommandId, Id, NodeId, SessionId
+from exo.shared.types.common import CommandId, Id, NodeId, SessionId, SystemId
from exo.shared.types.tasks import Task, TaskId, TaskStatus
from exo.shared.types.worker.downloads import DownloadProgress
from exo.shared.types.worker.instances import Instance, InstanceId
@@ -162,10 +162,19 @@ class IndexedEvent(CamelCaseModel):
event: Event
-class ForwarderEvent(CamelCaseModel):
+class GlobalForwarderEvent(CamelCaseModel):
"""An event the forwarder will serialize and send over the network"""
origin_idx: int = Field(ge=0)
origin: NodeId
session: SessionId
event: Event
+
+
+class LocalForwarderEvent(CamelCaseModel):
+ """An event the forwarder will serialize and send over the network"""
+
+ origin_idx: int = Field(ge=0)
+ origin: SystemId
+ session: SessionId
+ event: Event
diff --git a/src/exo/worker/main.py b/src/exo/worker/main.py
index 3bafedaa..7782ce2d 100644
--- a/src/exo/worker/main.py
+++ b/src/exo/worker/main.py
@@ -1,7 +1,6 @@
from collections import defaultdict
from datetime import datetime, timezone
from random import random
-from typing import Iterator
import anyio
from anyio import CancelScope, create_task_group, fail_after
@@ -17,13 +16,14 @@ from exo.shared.types.commands import (
RequestEventLog,
StartDownload,
)
-from exo.shared.types.common import CommandId, NodeId, SessionId
+from exo.shared.types.common import CommandId, NodeId, SessionId, SystemId
from exo.shared.types.events import (
Event,
EventId,
- ForwarderEvent,
+ GlobalForwarderEvent,
IndexedEvent,
InputChunkReceived,
+ LocalForwarderEvent,
NodeGatheredInfo,
TaskCreated,
TaskStatusUpdated,
@@ -58,24 +58,22 @@ class Worker:
node_id: NodeId,
session_id: SessionId,
*,
- global_event_receiver: Receiver[ForwarderEvent],
- local_event_sender: Sender[ForwarderEvent],
+ global_event_receiver: Receiver[GlobalForwarderEvent],
+ local_event_sender: Sender[LocalForwarderEvent],
# This is for requesting updates. It doesn't need to be a general command sender right now,
# but I think it's the correct way to be thinking about commands
command_sender: Sender[ForwarderCommand],
download_command_sender: Sender[ForwarderDownloadCommand],
- event_index_counter: Iterator[int],
):
self.node_id: NodeId = node_id
self.session_id: SessionId = session_id
self.global_event_receiver = global_event_receiver
self.local_event_sender = local_event_sender
- self.event_index_counter = event_index_counter
self.command_sender = command_sender
self.download_command_sender = download_command_sender
self.event_buffer = OrderedBuffer[Event]()
- self.out_for_delivery: dict[EventId, ForwarderEvent] = {}
+ self.out_for_delivery: dict[EventId, LocalForwarderEvent] = {}
self.state: State = State()
self.runners: dict[RunnerId, RunnerSupervisor] = {}
@@ -86,6 +84,8 @@ class Worker:
self._nack_base_seconds: float = 0.5
self._nack_cap_seconds: float = 10.0
+ self._system_id = SystemId()
+
self.event_sender, self.event_receiver = channel[Event]()
# Buffer for input image chunks (for image editing)
@@ -132,6 +132,8 @@ class Worker:
async def _event_applier(self):
with self.global_event_receiver as events:
async for f_event in events:
+ if f_event.session != self.session_id:
+ continue
if f_event.origin != self.session_id.master_node_id:
continue
self.event_buffer.ingest(f_event.origin_idx, f_event.event)
@@ -212,7 +214,7 @@ class Worker:
await self.download_command_sender.send(
ForwarderDownloadCommand(
- origin=self.node_id,
+ origin=self._system_id,
command=StartDownload(
target_node_id=self.node_id,
shard_metadata=shard,
@@ -317,7 +319,7 @@ class Worker:
)
await self.command_sender.send(
ForwarderCommand(
- origin=self.node_id,
+ origin=self._system_id,
command=RequestEventLog(since_idx=since_idx),
)
)
@@ -344,15 +346,16 @@ class Worker:
return runner
async def _forward_events(self) -> None:
+ idx = 0
with self.event_receiver as events:
async for event in events:
- idx = next(self.event_index_counter)
- fe = ForwarderEvent(
+ fe = LocalForwarderEvent(
origin_idx=idx,
- origin=self.node_id,
+ origin=self._system_id,
session=self.session_id,
event=event,
)
+ idx += 1
logger.debug(f"Worker published event {idx}: {str(event)[:100]}")
await self.local_event_sender.send(fe)
self.out_for_delivery[event.event_id] = fe
← d484b062 bench: add download timing to bench output (#1566)
·
back to Exo
·
Refactor runner into separate runners (#1570) 1a2b8b04 →