Compare commits

...
21 changed files with 1134 additions and 93 deletions

No files matched your search

+130 -34
View File
@@ -121,6 +121,8 @@ from exo.api.types.openai_responses import (
)
from exo.master.image_store import ImageStore
from exo.master.placement import place_instance as get_instance_placements
from exo.routing.event_router import EventRouter
from exo.routing.snapshot_receiver import SnapshotReceiver
from exo.shared.apply import apply
from exo.shared.constants import (
DASHBOARD_DIR,
@@ -164,6 +166,7 @@ from exo.shared.types.commands import (
ImageEdits,
ImageGeneration,
PlaceInstance,
RequestSnapshot,
SendInputChunk,
SetInstanceLink,
StartDownload,
@@ -171,16 +174,17 @@ from exo.shared.types.commands import (
TaskFinished,
TextGeneration,
)
from exo.shared.types.common import CommandId, Id, NodeId, SystemId
from exo.shared.types.common import CommandId, Id, NodeId, SessionId, SystemId
from exo.shared.types.events import (
ChunkGenerated,
Event,
IndexedEvent,
InstanceDeleted,
TracesMerged,
TransientEvent,
)
from exo.shared.types.instance_link import InstanceLink, InstanceLinkId
from exo.shared.types.memory import Memory
from exo.shared.types.snapshots import SnapshotChunk
from exo.shared.types.state import State
from exo.shared.types.tasks import (
ImageEdits as ImageEditsTask,
@@ -207,6 +211,8 @@ from exo.utils.task_group import TaskGroup
_API_EVENT_LOG_DIR = EXO_EVENT_LOG_DIR / "api"
ONBOARDING_COMPLETE_FILE = EXO_CACHE_HOME / "onboarding_complete"
_SNAPSHOT_FETCH_TIMEOUT_SECONDS = 30
def _format_to_content_type(image_format: Literal["png", "jpeg", "webp"] | None) -> str:
return f"image/{image_format or 'png'}"
@@ -236,9 +242,13 @@ class API:
def __init__(
self,
node_id: NodeId,
session_id: SessionId,
*,
port: int,
event_router: EventRouter,
event_receiver: Receiver[IndexedEvent],
transient_event_receiver: Receiver[TransientEvent],
snapshot_chunk_receiver: Receiver[SnapshotChunk],
command_sender: Sender[ForwarderCommand],
download_command_sender: Sender[ForwarderDownloadCommand],
# This lets us pause the API if an election is running
@@ -247,9 +257,13 @@ class API:
self.state = State()
self._event_log = DiskEventLog(_API_EVENT_LOG_DIR)
self._system_id = SystemId()
self.session_id = session_id
self.event_router = event_router
self.command_sender = command_sender
self.download_command_sender = download_command_sender
self.event_receiver = event_receiver
self.transient_event_receiver = transient_event_receiver
self.snapshot_chunk_receiver = snapshot_chunk_receiver
self.election_receiver = election_receiver
self.node_id: NodeId = node_id
self.last_completed_election: int = 0
@@ -288,21 +302,38 @@ class API:
self._image_generation_queues: dict[
CommandId, Sender[ImageChunk | ErrorChunk]
] = {}
self._observed_generation_commands: set[CommandId] = set()
self._image_store = ImageStore(EXO_IMAGE_CACHE_DIR)
self._tg: TaskGroup = TaskGroup()
def reset(self, result_clock: int, event_receiver: Receiver[IndexedEvent]):
def reset(
self,
result_clock: int,
session_id: SessionId,
event_router: EventRouter,
event_receiver: Receiver[IndexedEvent],
transient_event_receiver: Receiver[TransientEvent],
snapshot_chunk_receiver: Receiver[SnapshotChunk],
):
logger.info("Resetting API State")
self._event_log.close()
self._event_log = DiskEventLog(_API_EVENT_LOG_DIR)
self.state = State()
self._system_id = SystemId()
self.session_id = session_id
self.event_router = event_router
self._text_generation_queues = {}
self._image_generation_queues = {}
self._observed_generation_commands = set()
self.unpause(result_clock)
self.event_receiver.close()
self.event_receiver = event_receiver
self._tg.start_soon(self._apply_state)
self.transient_event_receiver.close()
self.transient_event_receiver = transient_event_receiver
self.snapshot_chunk_receiver.close()
self.snapshot_chunk_receiver = snapshot_chunk_receiver
self._tg.start_soon(self._bootstrap_then_apply_state)
self._tg.start_soon(self._apply_transient)
def unpause(self, result_clock: int):
logger.info("Unpausing API")
@@ -1836,7 +1867,9 @@ class API:
try:
async with self._tg as tg:
logger.info("Starting API")
tg.start_soon(self._apply_state)
tg.start_soon(self._bootstrap_then_apply_state)
tg.start_soon(self._apply_transient)
tg.start_soon(self._reconcile_streams)
tg.start_soon(self._pause_on_new_election)
tg.start_soon(self._cleanup_expired_images)
print_startup_banner(self.port)
@@ -1850,6 +1883,8 @@ class API:
self._event_log.close()
self.command_sender.close()
self.event_receiver.close()
self.transient_event_receiver.close()
self.snapshot_chunk_receiver.close()
async def run_api(self, ev: anyio.Event):
cfg = Config()
@@ -1865,47 +1900,108 @@ class API:
shutdown_trigger=ev.wait,
)
async def _bootstrap_then_apply_state(self):
await self._fetch_snapshot()
await self._apply_state()
async def _fetch_snapshot(self) -> None:
receiver = SnapshotReceiver(self.node_id, self.session_id)
await self.command_sender.send(
ForwarderCommand(
origin=self._system_id,
command=RequestSnapshot(requester_node_id=self.node_id),
)
)
with anyio.move_on_after(_SNAPSHOT_FETCH_TIMEOUT_SECONDS):
with self.snapshot_chunk_receiver as chunks:
async for chunk in chunks:
received = receiver.ingest(chunk)
if received is None:
continue
self.state = received.state
self.event_router.set_buffer_start(
received.last_event_applied_idx + 1
)
logger.info(
f"API bootstrapped from snapshot at idx "
f"{received.last_event_applied_idx}"
)
return
logger.info(
"API: no snapshot received before timeout; falling back to full event-log replay"
)
async def _apply_state(self):
with self.event_receiver as events:
async for i_event in events:
if i_event.idx <= self.state.last_event_applied_idx:
continue
self._event_log.append(i_event.event)
self.state = apply(self.state, i_event)
event = i_event.event
if isinstance(event, ChunkGenerated):
if queue := self._image_generation_queues.get(
event.command_id, None
):
assert isinstance(event.chunk, ImageChunk)
try:
await queue.send(event.chunk)
except (BrokenResourceError, ClosedResourceError):
self._image_generation_queues.pop(event.command_id, None)
if queue := self._text_generation_queues.get(
event.command_id, None
):
assert not isinstance(event.chunk, ImageChunk)
try:
await queue.send(event.chunk)
except (BrokenResourceError, ClosedResourceError):
self._text_generation_queues.pop(event.command_id, None)
if isinstance(event, InstanceDeleted):
self._close_streams_for_instance(event.instance_id)
if isinstance(event, TracesMerged):
self._save_merged_trace(event)
def _close_streams_for_instance(self, instance_id: InstanceId) -> None:
"""Close any active generation streams for commands running on the given instance."""
for task in self.state.tasks.values():
if task.instance_id != instance_id:
continue
if not isinstance(
async def _apply_transient(self) -> None:
with self.transient_event_receiver as events:
async for event in events:
if isinstance(event, ChunkGenerated):
await self._dispatch_chunk(event)
async def _dispatch_chunk(self, event: ChunkGenerated) -> None:
if queue := self._image_generation_queues.get(event.command_id, None):
assert isinstance(event.chunk, ImageChunk)
try:
await queue.send(event.chunk)
except (BrokenResourceError, ClosedResourceError):
self._image_generation_queues.pop(event.command_id, None)
if queue := self._text_generation_queues.get(event.command_id, None):
assert not isinstance(event.chunk, ImageChunk)
try:
await queue.send(event.chunk)
except (BrokenResourceError, ClosedResourceError):
self._text_generation_queues.pop(event.command_id, None)
async def _reconcile_streams(self) -> None:
while True:
await anyio.sleep(1)
self._reconcile_streams_once()
def _reconcile_streams_once(self) -> None:
generation_tasks = [
task
for task in self.state.tasks.values()
if isinstance(
task, (TextGenerationTask, ImageGenerationTask, ImageEditsTask)
):
continue
if sender := self._text_generation_queues.pop(task.command_id, None):
)
]
state_command_ids = {task.command_id for task in generation_tasks}
self._observed_generation_commands.update(state_command_ids)
live_command_ids = {
task.command_id
for task in generation_tasks
if task.instance_id in self.state.instances
}
queued_command_ids = set(self._text_generation_queues) | set(
self._image_generation_queues
)
stale_command_ids = (
self._observed_generation_commands - live_command_ids
) & queued_command_ids
self._close_streams_for_commands(stale_command_ids)
self._observed_generation_commands = (
self._observed_generation_commands & queued_command_ids
) | state_command_ids
def _close_streams_for_commands(self, command_ids: set[CommandId]) -> None:
for command_id in command_ids:
if sender := self._text_generation_queues.pop(command_id, None):
sender.close()
if sender := self._image_generation_queues.pop(task.command_id, None):
if sender := self._image_generation_queues.pop(command_id, None):
sender.close()
def _save_merged_trace(self, event: TracesMerged) -> None:
@@ -0,0 +1,191 @@
# pyright: reportPrivateUsage=false
import hashlib
import anyio
import pytest
import zstandard
from exo.api.main import API
from exo.routing.event_router import EventRouter
from exo.shared.models.model_cards import ModelId
from exo.shared.types.chunks import (
ErrorChunk,
PrefillProgressChunk,
TokenChunk,
ToolCallChunk,
)
from exo.shared.types.commands import ForwarderCommand, RequestSnapshot
from exo.shared.types.common import CommandId, NodeId, SessionId, SystemId
from exo.shared.types.events import (
ChunkGenerated,
Event,
GlobalForwarderEvent,
IndexedEvent,
LocalForwarderEvent,
TestEvent,
TransientEvent,
)
from exo.shared.types.snapshots import SnapshotChunk, SnapshotTransferId
from exo.shared.types.state import State
from exo.utils.channels import Receiver, Sender, channel
class _FakeEventLog:
def __init__(self) -> None:
self.appended: list[Event] = []
def append(self, event: Event) -> None:
self.appended.append(event)
def _snapshot_chunk(
state: State, *, requester_node_id: NodeId, session_id: SessionId
) -> SnapshotChunk:
body = zstandard.ZstdCompressor().compress(state.model_dump_json().encode("utf-8"))
return SnapshotChunk.from_data(
data=body,
transfer_id=SnapshotTransferId("transfer-1"),
requester_node_id=requester_node_id,
session_id=session_id,
schema_version=state.schema_version,
last_event_applied_idx=state.last_event_applied_idx,
chunk_index=0,
total_chunks=1,
sha256_hex=hashlib.sha256(body).hexdigest(),
)
def _api(
node_id: NodeId, session_id: SessionId
) -> tuple[
API,
EventRouter,
Receiver[ForwarderCommand],
Sender[SnapshotChunk],
Sender[TransientEvent],
Sender[IndexedEvent],
_FakeEventLog,
]:
router_command_sender, _router_command_receiver = channel[ForwarderCommand]()
_global_event_sender, global_event_receiver = channel[GlobalForwarderEvent]()
local_event_sender, _local_event_receiver = channel[LocalForwarderEvent]()
event_router = EventRouter(
session_id=session_id,
command_sender=router_command_sender,
external_inbound=global_event_receiver,
external_outbound=local_event_sender,
)
event_sender, event_receiver = channel[IndexedEvent]()
transient_sender, transient_receiver = channel[TransientEvent]()
command_sender, command_receiver = channel[ForwarderCommand]()
snapshot_sender, snapshot_receiver = channel[SnapshotChunk]()
api = object.__new__(API)
api.node_id = node_id
api.session_id = session_id
api.event_router = event_router
api.event_receiver = event_receiver
api.transient_event_receiver = transient_receiver
api.snapshot_chunk_receiver = snapshot_receiver
api.command_sender = command_sender
api._system_id = SystemId("api-system")
api.state = State()
event_log = _FakeEventLog()
api._event_log = event_log # pyright: ignore[reportAttributeAccessIssue]
api._image_generation_queues = {}
api._text_generation_queues = {}
return (
api,
event_router,
command_receiver,
snapshot_sender,
transient_sender,
event_sender,
event_log,
)
@pytest.mark.asyncio
async def test_api_fetch_snapshot_applies_state_and_fast_forwards_router() -> None:
node_id = NodeId("api")
session_id = SessionId(master_node_id=NodeId("master"), election_clock=1)
(
api,
event_router,
command_receiver,
snapshot_sender,
_transient_sender,
_event_sender,
_event_log,
) = _api(node_id, session_id)
state = State(last_event_applied_idx=7)
async with anyio.create_task_group() as tg:
tg.start_soon(api._fetch_snapshot)
command = await command_receiver.receive()
assert isinstance(command.command, RequestSnapshot)
assert command.command.requester_node_id == node_id
await snapshot_sender.send(
_snapshot_chunk(state, requester_node_id=node_id, session_id=session_id)
)
assert api.state.last_event_applied_idx == 7
assert event_router.event_buffer.next_idx_to_release == 8
@pytest.mark.asyncio
async def test_api_apply_state_ignores_events_covered_by_snapshot() -> None:
node_id = NodeId("api")
session_id = SessionId(master_node_id=NodeId("master"), election_clock=1)
(
api,
_event_router,
_command_receiver,
_snapshot_sender,
_transient_sender,
event_sender,
event_log,
) = _api(node_id, session_id)
api.state = State(last_event_applied_idx=7)
async with anyio.create_task_group() as tg:
tg.start_soon(api._apply_state)
await event_sender.send(IndexedEvent(idx=7, event=TestEvent()))
await event_sender.send(IndexedEvent(idx=8, event=TestEvent()))
while api.state.last_event_applied_idx != 8:
await anyio.sleep(0.001)
tg.cancel_scope.cancel()
assert len(event_log.appended) == 1
@pytest.mark.asyncio
async def test_api_apply_transient_dispatches_generated_chunks() -> None:
node_id = NodeId("api")
session_id = SessionId(master_node_id=NodeId("master"), election_clock=1)
(
api,
_event_router,
_command_receiver,
_snapshot_sender,
transient_sender,
_event_sender,
_event_log,
) = _api(node_id, session_id)
command_id = CommandId("cmd-a")
chunk_sender, chunk_receiver = channel[
TokenChunk | ErrorChunk | ToolCallChunk | PrefillProgressChunk
]()
api._text_generation_queues[command_id] = chunk_sender
chunk = ErrorChunk(model=ModelId("test-model"), error_message="test chunk")
async with anyio.create_task_group() as tg:
tg.start_soon(api._apply_transient)
await transient_sender.send(ChunkGenerated(command_id=command_id, chunk=chunk))
assert await chunk_receiver.receive() == chunk
tg.cancel_scope.cancel()
@@ -1,11 +1,11 @@
# pyright: reportUnusedFunction=false, reportAny=false
"""Tests that InstanceDeleted events close active generation streams."""
"""Tests that streaming queues reconcile against durable State."""
from unittest.mock import MagicMock
from exo.api.main import API
from exo.api.types import ImageGenerationTaskParams
from exo.shared.types.common import CommandId, ModelId
from exo.shared.types.common import CommandId, ModelId, NodeId
from exo.shared.types.state import State
from exo.shared.types.tasks import ImageGeneration, TextGeneration
from exo.shared.types.text_generation import (
@@ -13,15 +13,16 @@ from exo.shared.types.text_generation import (
InputMessageContent,
TextGenerationTaskParams,
)
from exo.shared.types.worker.instances import InstanceId
from exo.shared.types.worker.instances import InstanceId, MlxRingInstance
from exo.shared.types.worker.runners import ShardAssignments
def _make_api_with_state(state: State) -> API:
"""Create a minimal API instance with pre-set state."""
api = object.__new__(API)
api.state = state
api._text_generation_queues = {} # pyright: ignore[reportPrivateUsage]
api._image_generation_queues = {} # pyright: ignore[reportPrivateUsage]
api._observed_generation_commands = set() # pyright: ignore[reportPrivateUsage]
return api
@@ -38,45 +39,90 @@ def _make_text_gen_task(
)
def test_close_streams_for_deleted_instance() -> None:
"""Deleting an instance closes the text generation sender for commands on that instance."""
def _make_instance(instance_id: InstanceId) -> MlxRingInstance:
return MlxRingInstance(
instance_id=instance_id,
shard_assignments=ShardAssignments(
model_id=ModelId("test-model"),
node_to_runner={},
runner_to_shard={},
),
hosts_by_node={NodeId("node-1"): []},
ephemeral_port=1,
)
def test_reconcile_closes_stream_when_task_instance_is_missing() -> None:
instance_id = InstanceId("inst-1")
command_id = CommandId("cmd-1")
task = _make_text_gen_task(instance_id, command_id)
state = State(tasks={task.task_id: task})
api = _make_api_with_state(state)
api = _make_api_with_state(State(tasks={task.task_id: task}, instances={}))
sender = MagicMock()
api._text_generation_queues[command_id] = sender # pyright: ignore[reportPrivateUsage]
api._close_streams_for_instance(instance_id) # pyright: ignore[reportPrivateUsage]
api._reconcile_streams_once() # pyright: ignore[reportPrivateUsage]
sender.close.assert_called_once()
assert command_id not in api._text_generation_queues # pyright: ignore[reportPrivateUsage]
def test_close_streams_ignores_unrelated_instances() -> None:
"""Deleting an instance does NOT close streams for commands on other instances."""
target_id = InstanceId("inst-delete")
other_id = InstanceId("inst-keep")
other_cmd = CommandId("cmd-keep")
other_task = _make_text_gen_task(other_id, other_cmd)
state = State(tasks={other_task.task_id: other_task})
api = _make_api_with_state(state)
def test_reconcile_keeps_stream_for_live_task_instance() -> None:
instance_id = InstanceId("inst-live")
command_id = CommandId("cmd-live")
task = _make_text_gen_task(instance_id, command_id)
api = _make_api_with_state(
State(
tasks={task.task_id: task},
instances={instance_id: _make_instance(instance_id)},
)
)
sender = MagicMock()
api._text_generation_queues[other_cmd] = sender # pyright: ignore[reportPrivateUsage]
api._text_generation_queues[command_id] = sender # pyright: ignore[reportPrivateUsage]
api._close_streams_for_instance(target_id) # pyright: ignore[reportPrivateUsage]
api._reconcile_streams_once() # pyright: ignore[reportPrivateUsage]
sender.close.assert_not_called()
assert other_cmd in api._text_generation_queues # pyright: ignore[reportPrivateUsage]
assert command_id in api._text_generation_queues # pyright: ignore[reportPrivateUsage]
def test_close_streams_for_deleted_instance_image_generation() -> None:
"""Deleting an instance closes the image generation sender for commands on that instance."""
def test_reconcile_does_not_close_command_before_state_observes_it() -> None:
command_id = CommandId("cmd-not-created-yet")
api = _make_api_with_state(State())
sender = MagicMock()
api._text_generation_queues[command_id] = sender # pyright: ignore[reportPrivateUsage]
api._reconcile_streams_once() # pyright: ignore[reportPrivateUsage]
sender.close.assert_not_called()
assert command_id in api._text_generation_queues # pyright: ignore[reportPrivateUsage]
def test_reconcile_closes_stream_after_observed_task_leaves_state() -> None:
instance_id = InstanceId("inst-live")
command_id = CommandId("cmd-deleted")
task = _make_text_gen_task(instance_id, command_id)
api = _make_api_with_state(
State(
tasks={task.task_id: task},
instances={instance_id: _make_instance(instance_id)},
)
)
sender = MagicMock()
api._text_generation_queues[command_id] = sender # pyright: ignore[reportPrivateUsage]
api._reconcile_streams_once() # pyright: ignore[reportPrivateUsage]
api.state = State(instances={instance_id: _make_instance(instance_id)})
api._reconcile_streams_once() # pyright: ignore[reportPrivateUsage]
sender.close.assert_called_once()
assert command_id not in api._text_generation_queues # pyright: ignore[reportPrivateUsage]
def test_reconcile_closes_image_stream_when_task_instance_is_missing() -> None:
instance_id = InstanceId("inst-img")
command_id = CommandId("cmd-img")
task = ImageGeneration(
@@ -84,14 +130,12 @@ def test_close_streams_for_deleted_instance_image_generation() -> None:
command_id=command_id,
task_params=ImageGenerationTaskParams(prompt="a cat", model="test-model"),
)
state = State(tasks={task.task_id: task})
api = _make_api_with_state(state)
api = _make_api_with_state(State(tasks={task.task_id: task}, instances={}))
sender = MagicMock()
api._image_generation_queues[command_id] = sender # pyright: ignore[reportPrivateUsage]
api._close_streams_for_instance(instance_id) # pyright: ignore[reportPrivateUsage]
api._reconcile_streams_once() # pyright: ignore[reportPrivateUsage]
sender.close.assert_called_once()
assert command_id not in api._image_generation_queues # pyright: ignore[reportPrivateUsage]
+47 -1
View File
@@ -17,6 +17,7 @@ from exo.download.impl_shard_downloader import exo_shard_downloader
from exo.master.main import Master
from exo.routing.event_router import EventRouter
from exo.routing.router import Router, get_node_id_keypair
from exo.routing.transient_router import TransientRouter
from exo.shared.constants import EXO_LOG
from exo.shared.election import Election, ElectionResult
from exo.shared.logging import logger_cleanup, logger_setup
@@ -31,6 +32,7 @@ from exo.worker.main import Worker
class Node:
router: Router
event_router: EventRouter
transient_router: TransientRouter
download_coordinator: DownloadCoordinator | None
worker: Worker | None
election: Election # Every node participates in election, as we do want a node to become master even if it isn't a master candidate if no master candidates are present.
@@ -55,6 +57,7 @@ class Node:
)
await router.register_topic(topics.GLOBAL_EVENTS)
await router.register_topic(topics.LOCAL_EVENTS)
await router.register_topic(topics.TRANSIENT_EVENTS)
await router.register_topic(topics.COMMANDS)
await router.register_topic(topics.ELECTION_MESSAGES)
await router.register_topic(topics.CONNECTION_MESSAGES)
@@ -66,6 +69,12 @@ class Node:
external_outbound=router.sender(topics.LOCAL_EVENTS),
external_inbound=router.receiver(topics.GLOBAL_EVENTS),
)
transient_router = TransientRouter(
node_id=node_id,
session_id=session_id,
external_outbound=router.sender(topics.TRANSIENT_EVENTS),
external_inbound=router.receiver(topics.TRANSIENT_EVENTS),
)
logger.info(f"Starting node {node_id}")
@@ -84,8 +93,12 @@ class Node:
if args.spawn_api:
api = API(
node_id,
session_id,
port=args.api_port,
event_router=event_router,
event_receiver=event_router.receiver(),
transient_event_receiver=transient_router.receiver(),
snapshot_chunk_receiver=router.receiver(topics.SNAPSHOT_RESPONSES),
command_sender=router.sender(topics.COMMANDS),
download_command_sender=router.sender(topics.DOWNLOAD_COMMANDS),
election_receiver=router.receiver(topics.ELECTION_MESSAGES),
@@ -96,8 +109,13 @@ class Node:
if not args.no_worker:
worker = Worker(
node_id,
session_id,
event_router=event_router,
event_receiver=event_router.receiver(),
event_sender=event_router.sender(),
transient_event_receiver=transient_router.receiver(),
transient_event_sender=transient_router.sender(),
snapshot_chunk_receiver=router.receiver(topics.SNAPSHOT_RESPONSES),
command_sender=router.sender(topics.COMMANDS),
download_command_sender=router.sender(topics.DOWNLOAD_COMMANDS),
api_port=args.api_port,
@@ -110,6 +128,8 @@ class Node:
node_id,
session_id,
event_sender=event_router.sender(),
transient_event_receiver=transient_router.receiver(),
transient_event_sender=transient_router.sender(),
global_event_sender=router.sender(topics.GLOBAL_EVENTS),
local_event_receiver=router.receiver(topics.LOCAL_EVENTS),
command_receiver=router.receiver(topics.COMMANDS),
@@ -134,6 +154,7 @@ class Node:
return cls(
router,
event_router,
transient_router,
download_coordinator,
worker,
election,
@@ -151,6 +172,7 @@ class Node:
signal.signal(signal.SIGTERM, lambda _, __: self.shutdown())
tg.start_soon(self.router.run)
tg.start_soon(self.event_router.run)
tg.start_soon(self.transient_router.run)
tg.start_soon(self.election.run)
if self.download_coordinator:
tg.start_soon(self.download_coordinator.run)
@@ -194,6 +216,13 @@ class Node:
self.router.receiver(topics.GLOBAL_EVENTS),
self.router.sender(topics.LOCAL_EVENTS),
)
self.transient_router.shutdown()
self.transient_router = TransientRouter(
node_id=self.node_id,
session_id=result.session_id,
external_outbound=self.router.sender(topics.TRANSIENT_EVENTS),
external_inbound=self.router.receiver(topics.TRANSIENT_EVENTS),
)
if (
result.session_id.master_node_id == self.node_id
@@ -209,6 +238,8 @@ class Node:
self.node_id,
result.session_id,
event_sender=self.event_router.sender(),
transient_event_receiver=self.transient_router.receiver(),
transient_event_sender=self.transient_router.sender(),
global_event_sender=self.router.sender(topics.GLOBAL_EVENTS),
local_event_receiver=self.router.receiver(topics.LOCAL_EVENTS),
command_receiver=self.router.receiver(topics.COMMANDS),
@@ -251,8 +282,15 @@ class Node:
# TODO: add profiling etc to resource monitor
self.worker = Worker(
self.node_id,
result.session_id,
event_router=self.event_router,
event_receiver=self.event_router.receiver(),
event_sender=self.event_router.sender(),
transient_event_receiver=self.transient_router.receiver(),
transient_event_sender=self.transient_router.sender(),
snapshot_chunk_receiver=self.router.receiver(
topics.SNAPSHOT_RESPONSES
),
command_sender=self.router.sender(topics.COMMANDS),
download_command_sender=self.router.sender(
topics.DOWNLOAD_COMMANDS
@@ -261,8 +299,16 @@ class Node:
)
self._tg.start_soon(self.worker.run)
if self.api:
self.api.reset(result.won_clock, self.event_router.receiver())
self.api.reset(
result.won_clock,
result.session_id,
self.event_router,
self.event_router.receiver(),
self.transient_router.receiver(),
self.router.receiver(topics.SNAPSHOT_RESPONSES),
)
self._tg.start_soon(self.event_router.run)
self._tg.start_soon(self.transient_router.run)
else:
if self.api:
self.api.unpause(result.won_clock)
+7
View File
@@ -55,6 +55,7 @@ from exo.shared.types.events import (
TraceEventData,
TracesCollected,
TracesMerged,
TransientEvent,
)
from exo.shared.types.instance_link import InstanceLink
from exo.shared.types.snapshots import SnapshotChunk, SnapshotTransferId
@@ -136,6 +137,8 @@ class Master:
*,
command_receiver: Receiver[ForwarderCommand],
event_sender: Sender[Event],
transient_event_receiver: Receiver[TransientEvent],
transient_event_sender: Sender[TransientEvent],
local_event_receiver: Receiver[LocalForwarderEvent],
global_event_sender: Sender[GlobalForwarderEvent],
snapshot_chunk_sender: Sender[SnapshotChunk],
@@ -148,6 +151,8 @@ class Master:
self.command_task_mapping: dict[CommandId, TaskId] = {}
self.command_receiver = command_receiver
self.local_event_receiver = local_event_receiver
self.transient_event_receiver = transient_event_receiver
self.transient_event_sender = transient_event_sender
self.global_event_sender = global_event_sender
self.snapshot_chunk_sender = snapshot_chunk_sender
self.download_command_sender = download_command_sender
@@ -170,6 +175,8 @@ class Master:
self._event_log.close()
self.global_event_sender.close()
self.local_event_receiver.close()
self.transient_event_receiver.close()
self.transient_event_sender.close()
self.snapshot_chunk_sender.close()
self.command_receiver.close()
+7
View File
@@ -26,6 +26,7 @@ from exo.shared.types.events import (
LocalForwarderEvent,
NodeGatheredInfo,
TaskCreated,
TransientEvent,
)
from exo.shared.types.memory import Memory
from exo.shared.types.profiling import (
@@ -59,6 +60,7 @@ async def test_master():
local_event_sender, le_receiver = channel[LocalForwarderEvent]()
fcds, _fcdr = channel[ForwarderDownloadCommand]()
ev_send, ev_recv = channel[Event]()
transient_send, transient_recv = channel[TransientEvent]()
snapshot_chunk_send, _snapshot_chunk_recv = channel[SnapshotChunk]()
async def mock_event_router():
@@ -93,6 +95,8 @@ async def test_master():
node_id,
session_id,
event_sender=ev_send,
transient_event_receiver=transient_recv,
transient_event_sender=transient_send,
global_event_sender=ge_sender,
local_event_receiver=le_receiver,
command_receiver=co_receiver,
@@ -249,12 +253,15 @@ async def test_master_serves_snapshot_for_current_state():
ForwarderDownloadCommand
]()
event_sender, _event_receiver = channel[Event]()
transient_event_sender, transient_event_receiver = channel[TransientEvent]()
snapshot_chunk_sender, snapshot_chunk_receiver = channel[SnapshotChunk]()
master = Master(
node_id,
session_id,
event_sender=event_sender,
transient_event_receiver=transient_event_receiver,
transient_event_sender=transient_event_sender,
global_event_sender=ge_sender,
local_event_receiver=local_event_receiver,
command_receiver=command_receiver,
@@ -0,0 +1,69 @@
import anyio
import pytest
from exo.routing import topics
from exo.routing.transient_router import TransientRouter
from exo.shared.types.common import NodeId, SessionId
from exo.shared.types.events import (
GlobalForwarderTransientEvent,
TaskAcknowledged,
)
from exo.shared.types.tasks import TaskId
from exo.utils.channels import channel
def test_transient_topic_round_trips_wrapped_event() -> None:
wrapped = GlobalForwarderTransientEvent(
origin=NodeId("worker"),
session=SessionId(master_node_id=NodeId("master"), election_clock=1),
event=TaskAcknowledged(task_id=TaskId("task-1")),
)
restored = topics.TRANSIENT_EVENTS.deserialize(
topics.TRANSIENT_EVENTS.serialize(wrapped)
)
assert restored == wrapped
@pytest.mark.asyncio
async def test_transient_router_publishes_and_dispatches_session_events() -> None:
node_id = NodeId("worker")
session_id = SessionId(master_node_id=NodeId("master"), election_clock=1)
external_outbound_sender, external_outbound_receiver = channel[
GlobalForwarderTransientEvent
]()
external_inbound_sender, external_inbound_receiver = channel[
GlobalForwarderTransientEvent
]()
router = TransientRouter(
node_id=node_id,
session_id=session_id,
external_outbound=external_outbound_sender,
external_inbound=external_inbound_receiver,
)
local_sender = router.sender()
local_receiver = router.receiver()
event = TaskAcknowledged(task_id=TaskId("task-1"))
async with anyio.create_task_group() as tg:
tg.start_soon(router.run)
await local_sender.send(event)
wrapped = await external_outbound_receiver.receive()
assert wrapped.origin == node_id
assert wrapped.session == session_id
assert wrapped.event == event
await external_inbound_sender.send(wrapped)
assert await local_receiver.receive() == event
stale_session = SessionId(master_node_id=NodeId("master"), election_clock=2)
await external_inbound_sender.send(
wrapped.model_copy(update={"session": stale_session})
)
await anyio.sleep(0)
assert local_receiver.collect() == []
router.shutdown()
tg.cancel_scope.cancel()
+4
View File
@@ -6,6 +6,7 @@ from exo.shared.election import ElectionMessage
from exo.shared.types.commands import ForwarderCommand, ForwarderDownloadCommand
from exo.shared.types.events import (
GlobalForwarderEvent,
GlobalForwarderTransientEvent,
LocalForwarderEvent,
)
from exo.shared.types.snapshots import SnapshotChunk
@@ -40,6 +41,9 @@ class TypedTopic[T: FrozenModel]:
GLOBAL_EVENTS = TypedTopic("global_events", PublishPolicy.Always, GlobalForwarderEvent)
LOCAL_EVENTS = TypedTopic("local_events", PublishPolicy.Always, LocalForwarderEvent)
TRANSIENT_EVENTS = TypedTopic(
"transient_events", PublishPolicy.Always, GlobalForwarderTransientEvent
)
COMMANDS = TypedTopic("commands", PublishPolicy.Always, ForwarderCommand)
ELECTION_MESSAGES = TypedTopic(
"election_messages", PublishPolicy.Always, ElectionMessage
+84
View File
@@ -0,0 +1,84 @@
from dataclasses import dataclass, field
from anyio import BrokenResourceError, ClosedResourceError
from loguru import logger
from exo.shared.types.common import NodeId, SessionId
from exo.shared.types.events import (
GlobalForwarderTransientEvent,
TransientEvent,
)
from exo.utils.channels import Receiver, Sender, channel
from exo.utils.task_group import TaskGroup
@dataclass
class TransientRouter:
"""Routes unordered, non-durable events over the transient-events topic."""
node_id: NodeId
session_id: SessionId
external_outbound: Sender[GlobalForwarderTransientEvent]
external_inbound: Receiver[GlobalForwarderTransientEvent]
_outbound: list[Sender[TransientEvent]] = field(init=False, default_factory=list)
_inbound: list[Receiver[TransientEvent]] = field(init=False, default_factory=list)
_tg: TaskGroup = field(init=False, default_factory=TaskGroup)
def sender(self) -> Sender[TransientEvent]:
send, recv = channel[TransientEvent]()
if self._tg.is_running():
self._tg.start_soon(self._publish, recv)
else:
self._inbound.append(recv)
return send
def receiver(self) -> Receiver[TransientEvent]:
assert not self._tg.is_running(), (
"TransientRouter receivers must be registered before run()"
)
send, recv = channel[TransientEvent]()
self._outbound.append(send)
return recv
def shutdown(self) -> None:
self._tg.cancel_tasks()
async def run(self) -> None:
try:
async with self._tg as tg:
for recv in self._inbound:
tg.start_soon(self._publish, recv)
tg.start_soon(self._dispatch_inbound)
finally:
self.external_outbound.close()
for send in self._outbound:
send.close()
async def _publish(self, recv: Receiver[TransientEvent]) -> None:
with recv as events:
async for event in events:
await self.external_outbound.send(
GlobalForwarderTransientEvent(
origin=self.node_id,
session=self.session_id,
event=event,
)
)
async def _dispatch_inbound(self) -> None:
with self.external_inbound as wrapped_events:
async for wrapped in wrapped_events:
if wrapped.session != self.session_id:
continue
stale: set[int] = set()
for index, send in enumerate(self._outbound):
try:
await send.send(wrapped.event)
except (ClosedResourceError, BrokenResourceError):
stale.add(index)
if stale:
for index in sorted(stale, reverse=True):
self._outbound.pop(index)
logger.debug(
f"TransientRouter dropped {len(stale)} closed receivers"
)
+25 -3
View File
@@ -4,7 +4,8 @@ from datetime import datetime
from loguru import logger
from exo.shared.types.common import NodeId
from exo.shared.models.model_cards import ModelCard
from exo.shared.types.common import ModelId, NodeId
from exo.shared.types.events import (
ChunkGenerated,
CustomModelCardAdded,
@@ -81,10 +82,12 @@ def event_apply(event: Event, state: State) -> State:
| TaskAcknowledged()
| TracesCollected()
| TracesMerged()
| CustomModelCardAdded()
| CustomModelCardDeleted()
): # Pass-through events that don't modify state
return state
case CustomModelCardAdded():
return apply_custom_model_card_added(event, state)
case CustomModelCardDeleted():
return apply_custom_model_card_deleted(event, state)
case InstanceCreated():
return apply_instance_created(event, state)
case InstanceDeleted():
@@ -477,3 +480,22 @@ def apply_topology_edge_deleted(event: TopologyEdgeDeleted, state: State) -> Sta
topology.remove_connection(event.conn)
# TODO: Clean up removing the reverse connection
return state.model_copy(update={"topology": topology})
def apply_custom_model_card_added(event: CustomModelCardAdded, state: State) -> State:
new_cards: Mapping[ModelId, ModelCard] = {
**state.custom_model_cards,
event.model_card.model_id: event.model_card,
}
return state.model_copy(update={"custom_model_cards": new_cards})
def apply_custom_model_card_deleted(
event: CustomModelCardDeleted, state: State
) -> State:
new_cards: Mapping[ModelId, ModelCard] = {
model_id: card
for model_id, card in state.custom_model_cards.items()
if model_id != event.model_id
}
return state.model_copy(update={"custom_model_cards": new_cards})
@@ -0,0 +1,44 @@
from exo.shared.apply import apply
from exo.shared.models.model_cards import ModelCard, ModelTask
from exo.shared.types.common import ModelId
from exo.shared.types.events import (
CustomModelCardAdded,
CustomModelCardDeleted,
IndexedEvent,
)
from exo.shared.types.memory import Memory
from exo.shared.types.state import State
def _model_card(model_id: ModelId) -> ModelCard:
return ModelCard(
model_id=model_id,
n_layers=1,
storage_size=Memory.from_bytes(1),
hidden_size=1,
supports_tensor=True,
tasks=[ModelTask.TextGeneration],
)
def test_custom_model_card_added_is_reduced_into_state() -> None:
card = _model_card(ModelId("custom/model"))
state = apply(
State(),
IndexedEvent(idx=0, event=CustomModelCardAdded(model_card=card)),
)
assert state.custom_model_cards == {card.model_id: card}
def test_custom_model_card_deleted_removes_card_from_state() -> None:
card = _model_card(ModelId("custom/model"))
state = State(custom_model_cards={card.model_id: card}, last_event_applied_idx=0)
state = apply(
state,
IndexedEvent(idx=1, event=CustomModelCardDeleted(model_id=card.model_id)),
)
assert state.custom_model_cards == {}
+11
View File
@@ -172,6 +172,9 @@ Event = (
)
TransientEvent = TaskAcknowledged | ChunkGenerated | TracesCollected | TracesMerged
class IndexedEvent(FrozenModel):
"""An event indexed by the master, with a globally unique index"""
@@ -195,3 +198,11 @@ class LocalForwarderEvent(FrozenModel):
origin: SystemId
session: SessionId
event: Event
class GlobalForwarderTransientEvent(FrozenModel):
"""An unordered, non-durable event published to the cluster."""
origin: NodeId
session: SessionId
event: TransientEvent
+6 -1
View File
@@ -5,9 +5,10 @@ from typing import Any, cast
from pydantic import ConfigDict, Field, field_serializer, field_validator
from pydantic.alias_generators import to_camel
from exo.shared.models.model_cards import ModelCard
from exo.shared.topology import Topology, TopologySnapshot
from exo.shared.types.chunks import InputImageChunk
from exo.shared.types.common import CommandId, NodeId
from exo.shared.types.common import CommandId, ModelId, NodeId
from exo.shared.types.instance_link import InstanceLink, InstanceLinkId
from exo.shared.types.profiling import (
DiskUsage,
@@ -72,6 +73,10 @@ class State(FrozenModel):
instance_links: Mapping[InstanceLinkId, InstanceLink] = {}
prefill_server_ports: Mapping[RunnerId, int] = {}
# User-added model cards. Workers can reconcile their on-disk custom card
# cache from this state after snapshot bootstrap.
custom_model_cards: Mapping[ModelId, ModelCard] = {}
@field_serializer("topology", mode="plain")
def _encode_topology(self, value: Topology) -> TopologySnapshot:
return value.to_snapshot()
+4
View File
@@ -29,6 +29,10 @@ class KeyedBackoff[K]:
"""Return the number of recorded attempts for a key."""
return self._attempts.get(key, 0)
def tracked_keys(self) -> set[K]:
"""Return keys that currently have recorded backoff state."""
return set(self._attempts) | set(self._last_time)
def reset(self, key: K) -> None:
"""Reset backoff state for a key (e.g., on success)."""
self._attempts.pop(key, None)
+13
View File
@@ -0,0 +1,13 @@
from exo.utils.keyed_backoff import KeyedBackoff
def test_tracked_keys_reports_and_resets_backoff_state() -> None:
backoff = KeyedBackoff[str]()
backoff.record_attempt("instance-a")
assert backoff.tracked_keys() == {"instance-a"}
backoff.reset("instance-a")
assert backoff.tracked_keys() == set()
+104 -22
View File
@@ -8,31 +8,38 @@ from loguru import logger
from exo.api.types import ImageEditsTaskParams
from exo.download.download_utils import is_read_only_model_dir, resolve_existing_model
from exo.routing.event_router import EventRouter
from exo.routing.snapshot_receiver import SnapshotReceiver
from exo.shared.apply import apply
from exo.shared.constants import EXO_MAX_INSTANCE_RETRIES
from exo.shared.models.model_cards import ModelId, add_to_card_cache, delete_custom_card
from exo.shared.models.model_cards import (
ModelCard,
ModelId,
add_to_card_cache,
delete_custom_card,
)
from exo.shared.types.chunks import InputImageChunk
from exo.shared.types.commands import (
DeleteInstance,
ForwarderCommand,
ForwarderDownloadCommand,
RequestSnapshot,
StartDownload,
)
from exo.shared.types.common import CommandId, NodeId, SystemId
from exo.shared.types.common import CommandId, NodeId, SessionId, SystemId
from exo.shared.types.events import (
CustomModelCardAdded,
CustomModelCardDeleted,
Event,
IndexedEvent,
InstanceDeleted,
NodeDownloadProgress,
NodeGatheredInfo,
TaskCreated,
TaskStatusUpdated,
TopologyEdgeCreated,
TopologyEdgeDeleted,
TransientEvent,
)
from exo.shared.types.multiaddr import Multiaddr
from exo.shared.types.snapshots import SnapshotChunk
from exo.shared.types.state import State
from exo.shared.types.tasks import (
CancelTask,
@@ -58,14 +65,21 @@ from exo.utils.task_group import TaskGroup
from exo.worker.plan import plan
from exo.worker.runner.supervisor import RunnerSupervisor
_SNAPSHOT_FETCH_TIMEOUT_SECONDS = 30
class Worker:
def __init__(
self,
node_id: NodeId,
session_id: SessionId,
*,
event_router: EventRouter,
event_receiver: Receiver[IndexedEvent],
event_sender: Sender[Event],
transient_event_receiver: Receiver[TransientEvent],
transient_event_sender: Sender[TransientEvent],
snapshot_chunk_receiver: Receiver[SnapshotChunk],
# 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],
@@ -73,8 +87,13 @@ class Worker:
api_port: int,
):
self.node_id: NodeId = node_id
self.session_id: SessionId = session_id
self.event_router = event_router
self.event_receiver = event_receiver
self.event_sender = event_sender
self.transient_event_receiver = transient_event_receiver
self.transient_event_sender = transient_event_sender
self.snapshot_chunk_receiver = snapshot_chunk_receiver
self.command_sender = command_sender
self.download_command_sender = download_command_sender
self.api_port = api_port
@@ -94,6 +113,7 @@ class Worker:
self._instance_backoff: KeyedBackoff[InstanceId] = KeyedBackoff(
base=0.5, cap=10.0
)
self._synced_custom_cards: dict[ModelId, ModelCard] = {}
self._stopped: anyio.Event = anyio.Event()
async def run(self):
@@ -104,21 +124,61 @@ class Worker:
try:
async with self._tg as tg:
tg.start_soon(info_gatherer.run)
tg.start_soon(self._forward_info, info_recv)
tg.start_soon(self.plan_step)
tg.start_soon(self._event_applier)
tg.start_soon(self._poll_connection_updates)
tg.start_soon(self._bootstrap_then_run, info_gatherer, info_recv)
finally:
# Actual shutdown code - waits for all tasks to complete before executing.
logger.info("Stopping Worker")
self.event_sender.close()
self.transient_event_receiver.close()
self.transient_event_sender.close()
self.snapshot_chunk_receiver.close()
self.command_sender.close()
self.download_command_sender.close()
for runner in self.runners.values():
runner.shutdown()
self._stopped.set()
async def _bootstrap_then_run(
self, info_gatherer: InfoGatherer, info_recv: Receiver[GatheredInfo]
) -> None:
await self._fetch_snapshot()
self._sync_input_views_from_state()
self._tg.start_soon(info_gatherer.run)
self._tg.start_soon(self._forward_info, info_recv)
self._tg.start_soon(self.plan_step)
self._tg.start_soon(self._event_applier)
self._tg.start_soon(self._reconcile_instance_backoff)
self._tg.start_soon(self._reconcile_custom_cards)
self._tg.start_soon(self._poll_connection_updates)
async def _fetch_snapshot(self) -> None:
receiver = SnapshotReceiver(self.node_id, self.session_id)
await self.command_sender.send(
ForwarderCommand(
origin=self._system_id,
command=RequestSnapshot(requester_node_id=self.node_id),
)
)
with anyio.move_on_after(_SNAPSHOT_FETCH_TIMEOUT_SECONDS):
with self.snapshot_chunk_receiver as chunks:
async for chunk in chunks:
received = receiver.ingest(chunk)
if received is None:
continue
self.state = received.state
self.event_router.set_buffer_start(
received.last_event_applied_idx + 1
)
logger.info(
f"Worker bootstrapped from snapshot at idx "
f"{received.last_event_applied_idx}"
)
return
logger.info(
"No snapshot received before timeout; falling back to full event-log replay"
)
async def _forward_info(self, recv: Receiver[GatheredInfo]):
with recv as info_stream:
async for info in info_stream:
@@ -133,22 +193,43 @@ class Worker:
async def _event_applier(self):
with self.event_receiver as events:
async for event in events:
if event.idx <= self.state.last_event_applied_idx:
continue
# 2. for each event, apply it to the state
self.state = apply(self.state, event=event)
event = event.event
if isinstance(event, InstanceDeleted):
self._instance_backoff.reset(event.instance_id)
if isinstance(event, CustomModelCardAdded):
await event.model_card.save_to_custom_dir()
add_to_card_cache(event.model_card)
if isinstance(event, CustomModelCardDeleted):
await delete_custom_card(event.model_id)
self._sync_input_views_from_state()
async def _reconcile_instance_backoff(self) -> None:
while True:
await anyio.sleep(1)
self._reconcile_instance_backoff_once()
def _reconcile_instance_backoff_once(self) -> None:
live_instances = set(self.state.instances)
for instance_id in self._instance_backoff.tracked_keys():
if instance_id not in live_instances:
self._instance_backoff.reset(instance_id)
async def _reconcile_custom_cards(self) -> None:
while True:
await anyio.sleep(1)
await self._sync_custom_cards_from_state()
async def _sync_custom_cards_from_state(self) -> None:
target = dict(self.state.custom_model_cards)
for model_id, card in target.items():
if self._synced_custom_cards.get(model_id) == card:
continue
await card.save_to_custom_dir()
add_to_card_cache(card)
self._synced_custom_cards[model_id] = card
for model_id in list(self._synced_custom_cards):
if model_id in target:
continue
await delete_custom_card(model_id)
self._synced_custom_cards.pop(model_id, None)
def _sync_input_views_from_state(self) -> None:
self.input_chunk_buffer = {
command_id: dict(chunks)
@@ -357,6 +438,7 @@ class Worker:
runner = RunnerSupervisor.create(
bound_instance=task.bound_instance,
event_sender=self.event_sender.clone(),
transient_event_sender=self.transient_event_sender.clone(),
)
self.runners[task.bound_instance.bound_runner_id] = runner
self._tg.start_soon(runner.run)
+11 -2
View File
@@ -19,6 +19,7 @@ from exo.shared.types.events import (
RunnerStatusUpdated,
TaskAcknowledged,
TaskStatusUpdated,
TransientEvent,
)
from exo.shared.types.tasks import (
CANCEL_ALL_TASKS,
@@ -58,6 +59,7 @@ class RunnerSupervisor:
_ev_recv: MpReceiver[Event]
_task_sender: MpSender[Task]
_event_sender: Sender[Event]
_transient_event_sender: Sender[TransientEvent]
_cancel_sender: MpSender[TaskId]
_tg: TaskGroup = field(default_factory=TaskGroup, init=False)
status: RunnerStatus = field(default_factory=RunnerIdle, init=False)
@@ -75,6 +77,7 @@ class RunnerSupervisor:
*,
bound_instance: BoundInstance,
event_sender: Sender[Event],
transient_event_sender: Sender[TransientEvent],
initialize_timeout: float = 400,
) -> Self:
ev_send, ev_recv = mp_channel[Event]()
@@ -104,6 +107,7 @@ class RunnerSupervisor:
_task_sender=task_sender,
_cancel_sender=cancel_sender,
_event_sender=event_sender,
_transient_event_sender=transient_event_sender,
)
return self
@@ -124,6 +128,8 @@ class RunnerSupervisor:
self._task_sender.close()
with contextlib.suppress(ClosedResourceError):
self._event_sender.close()
with contextlib.suppress(ClosedResourceError):
self._transient_event_sender.close()
with contextlib.suppress(ClosedResourceError):
self._cancel_sender.send(CANCEL_ALL_TASKS)
with contextlib.suppress(ClosedResourceError):
@@ -235,7 +241,10 @@ class RunnerSupervisor:
)
self.in_progress.pop(event.task_id, None)
self.completed.add(event.task_id)
await self._event_sender.send(event)
if isinstance(event, ChunkGenerated):
await self._transient_event_sender.send(event)
else:
await self._event_sender.send(event)
except (ClosedResourceError, BrokenResourceError) as e:
await self._check_runner(e)
finally:
@@ -275,7 +284,7 @@ class RunnerSupervisor:
for task in self.in_progress.values():
if isinstance(task, (TextGeneration, ImageGeneration, ImageEdits)):
with anyio.CancelScope(shield=True):
await self._event_sender.send(
await self._transient_event_sender.send(
ChunkGenerated(
command_id=task.command_id,
chunk=ErrorChunk(
@@ -7,7 +7,12 @@ import pytest
from exo.shared.models.model_cards import ModelId
from exo.shared.types.chunks import ErrorChunk
from exo.shared.types.common import CommandId, NodeId
from exo.shared.types.events import ChunkGenerated, Event, RunnerStatusUpdated
from exo.shared.types.events import (
ChunkGenerated,
Event,
RunnerStatusUpdated,
TransientEvent,
)
from exo.shared.types.tasks import Task, TaskId, TextGeneration
from exo.shared.types.text_generation import (
InputMessage,
@@ -43,6 +48,7 @@ class _DeadProcess:
@pytest.mark.asyncio
async def test_check_runner_emits_error_chunk_for_inflight_text_generation() -> None:
event_sender, event_receiver = channel[Event]()
transient_sender, transient_receiver = channel[TransientEvent]()
task_sender, _ = mp_channel[Task]()
cancel_sender, _ = mp_channel[TaskId]()
_, ev_recv = mp_channel[Event]()
@@ -62,6 +68,7 @@ async def test_check_runner_emits_error_chunk_for_inflight_text_generation() ->
_ev_recv=ev_recv,
_task_sender=task_sender,
_event_sender=event_sender,
_transient_event_sender=transient_sender,
_cancel_sender=cancel_sender,
)
@@ -81,7 +88,7 @@ async def test_check_runner_emits_error_chunk_for_inflight_text_generation() ->
await supervisor._check_runner(RuntimeError("boom")) # pyright: ignore[reportPrivateUsage]
got_chunk = await event_receiver.receive()
got_chunk = await transient_receiver.receive()
got_status = await event_receiver.receive()
assert isinstance(got_chunk, ChunkGenerated)
@@ -93,5 +100,57 @@ async def test_check_runner_emits_error_chunk_for_inflight_text_generation() ->
assert isinstance(got_status.runner_status, RunnerFailed)
event_sender.close()
transient_sender.close()
with anyio.move_on_after(0.1):
await event_receiver.aclose()
with anyio.move_on_after(0.1):
await transient_receiver.aclose()
@pytest.mark.asyncio
async def test_forward_events_routes_generated_chunks_to_transient() -> None:
event_sender, event_receiver = channel[Event]()
transient_sender, transient_receiver = channel[TransientEvent]()
task_sender, _ = mp_channel[Task]()
cancel_sender, _ = mp_channel[TaskId]()
ev_send, ev_recv = mp_channel[Event]()
bound_instance: BoundInstance = get_bound_mlx_ring_instance(
instance_id=InstanceId("instance-a"),
model_id=ModelId("mlx-community/Llama-3.2-1B-Instruct-4bit"),
runner_id=RunnerId("runner-a"),
node_id=NodeId("node-a"),
)
supervisor = RunnerSupervisor(
shard_metadata=bound_instance.bound_shard,
bound_instance=bound_instance,
runner_process=cast("mp.Process", cast(object, _DeadProcess())),
initialize_timeout=400,
_ev_recv=ev_recv,
_task_sender=task_sender,
_event_sender=event_sender,
_transient_event_sender=transient_sender,
_cancel_sender=cancel_sender,
)
command_id = CommandId("cmd-a")
chunk = ChunkGenerated(
command_id=command_id,
chunk=ErrorChunk(
model=bound_instance.bound_shard.model_card.model_id,
error_message="test chunk",
),
)
async with anyio.create_task_group() as tg:
tg.start_soon(supervisor._forward_events) # pyright: ignore[reportPrivateUsage]
ev_send.send(chunk)
assert await transient_receiver.receive() == chunk
tg.cancel_scope.cancel()
assert event_receiver.collect() == []
event_sender.close()
transient_sender.close()
ev_send.close()
@@ -0,0 +1,78 @@
# pyright: reportPrivateUsage=false
import pytest
import exo.worker.main as worker_main
from exo.shared.models.model_cards import ModelCard, ModelTask
from exo.shared.types.common import ModelId
from exo.shared.types.memory import Memory
from exo.shared.types.state import State
from exo.worker.main import Worker
def _model_card(model_id: ModelId) -> ModelCard:
return ModelCard(
model_id=model_id,
n_layers=1,
storage_size=Memory.from_bytes(1),
hidden_size=1,
supports_tensor=True,
tasks=[ModelTask.TextGeneration],
)
def _worker() -> Worker:
worker = object.__new__(Worker)
worker.state = State()
worker._synced_custom_cards = {}
return worker
@pytest.mark.asyncio
async def test_worker_syncs_custom_cards_from_state(
monkeypatch: pytest.MonkeyPatch,
) -> None:
saved: list[ModelId] = []
cached: list[ModelId] = []
async def save_to_custom_dir(card: ModelCard) -> None:
saved.append(card.model_id)
def add_to_card_cache(card: ModelCard) -> None:
cached.append(card.model_id)
monkeypatch.setattr(ModelCard, "save_to_custom_dir", save_to_custom_dir)
monkeypatch.setattr(worker_main, "add_to_card_cache", add_to_card_cache)
card = _model_card(ModelId("custom/model"))
worker = _worker()
worker.state = State(custom_model_cards={card.model_id: card})
await worker._sync_custom_cards_from_state()
await worker._sync_custom_cards_from_state()
assert saved == [card.model_id]
assert cached == [card.model_id]
assert worker._synced_custom_cards == {card.model_id: card}
@pytest.mark.asyncio
async def test_worker_deletes_custom_cards_missing_from_state(
monkeypatch: pytest.MonkeyPatch,
) -> None:
deleted: list[ModelId] = []
async def delete_custom_card(model_id: ModelId) -> bool:
deleted.append(model_id)
return True
monkeypatch.setattr(worker_main, "delete_custom_card", delete_custom_card)
card = _model_card(ModelId("custom/model"))
worker = _worker()
worker._synced_custom_cards = {card.model_id: card}
await worker._sync_custom_cards_from_state()
assert deleted == [card.model_id]
assert worker._synced_custom_cards == {}
@@ -0,0 +1,36 @@
# pyright: reportPrivateUsage=false
from exo.shared.types.common import ModelId, NodeId
from exo.shared.types.state import State
from exo.shared.types.worker.instances import InstanceId, MlxRingInstance
from exo.shared.types.worker.runners import ShardAssignments
from exo.utils.keyed_backoff import KeyedBackoff
from exo.worker.main import Worker
def _make_instance(instance_id: InstanceId) -> MlxRingInstance:
return MlxRingInstance(
instance_id=instance_id,
shard_assignments=ShardAssignments(
model_id=ModelId("test-model"),
node_to_runner={},
runner_to_shard={},
),
hosts_by_node={NodeId("node-1"): []},
ephemeral_port=1,
)
def test_worker_reconciles_instance_backoff_from_state() -> None:
live_instance_id = InstanceId("inst-live")
deleted_instance_id = InstanceId("inst-deleted")
worker = object.__new__(Worker)
worker.state = State(instances={live_instance_id: _make_instance(live_instance_id)})
worker._instance_backoff = KeyedBackoff[InstanceId]()
worker._instance_backoff.record_attempt(live_instance_id)
worker._instance_backoff.record_attempt(deleted_instance_id)
worker._reconcile_instance_backoff_once()
assert worker._instance_backoff.attempts(live_instance_id) == 1
assert worker._instance_backoff.attempts(deleted_instance_id) == 0
@@ -0,0 +1,130 @@
# pyright: reportPrivateUsage=false
import hashlib
import anyio
import pytest
import zstandard
from exo.routing.event_router import EventRouter
from exo.shared.types.commands import (
ForwarderCommand,
ForwarderDownloadCommand,
RequestSnapshot,
)
from exo.shared.types.common import NodeId, SessionId
from exo.shared.types.events import (
Event,
GlobalForwarderEvent,
IndexedEvent,
LocalForwarderEvent,
TestEvent,
TransientEvent,
)
from exo.shared.types.snapshots import SnapshotChunk, SnapshotTransferId
from exo.shared.types.state import State
from exo.utils.channels import Receiver, Sender, channel
from exo.worker.main import Worker
def _snapshot_chunk(
state: State, *, requester_node_id: NodeId, session_id: SessionId
) -> SnapshotChunk:
body = zstandard.ZstdCompressor().compress(state.model_dump_json().encode("utf-8"))
return SnapshotChunk.from_data(
data=body,
transfer_id=SnapshotTransferId("transfer-1"),
requester_node_id=requester_node_id,
session_id=session_id,
schema_version=state.schema_version,
last_event_applied_idx=state.last_event_applied_idx,
chunk_index=0,
total_chunks=1,
sha256_hex=hashlib.sha256(body).hexdigest(),
)
def _worker(
node_id: NodeId, session_id: SessionId
) -> tuple[
Worker,
EventRouter,
Receiver[ForwarderCommand],
Sender[SnapshotChunk],
Sender[IndexedEvent],
]:
router_command_sender, _router_command_receiver = channel[ForwarderCommand]()
_global_event_sender, global_event_receiver = channel[GlobalForwarderEvent]()
local_event_sender, _local_event_receiver = channel[LocalForwarderEvent]()
event_router = EventRouter(
session_id=session_id,
command_sender=router_command_sender,
external_inbound=global_event_receiver,
external_outbound=local_event_sender,
)
event_sender, event_receiver = channel[IndexedEvent]()
local_event_output_sender, _local_event_output_receiver = channel[Event]()
transient_sender, transient_receiver = channel[TransientEvent]()
command_sender, command_receiver = channel[ForwarderCommand]()
download_command_sender, _download_command_receiver = channel[
ForwarderDownloadCommand
]()
snapshot_sender, snapshot_receiver = channel[SnapshotChunk]()
worker = Worker(
node_id,
session_id,
event_router=event_router,
event_receiver=event_receiver,
event_sender=local_event_output_sender,
transient_event_receiver=transient_receiver,
transient_event_sender=transient_sender,
snapshot_chunk_receiver=snapshot_receiver,
command_sender=command_sender,
download_command_sender=download_command_sender,
api_port=52415,
)
return worker, event_router, command_receiver, snapshot_sender, event_sender
@pytest.mark.asyncio
async def test_worker_fetch_snapshot_applies_state_and_fast_forwards_router() -> None:
node_id = NodeId("worker")
session_id = SessionId(master_node_id=NodeId("master"), election_clock=1)
worker, event_router, command_receiver, snapshot_sender, _event_sender = _worker(
node_id, session_id
)
state = State(last_event_applied_idx=7)
async with anyio.create_task_group() as tg:
tg.start_soon(worker._fetch_snapshot)
command = await command_receiver.receive()
assert isinstance(command.command, RequestSnapshot)
assert command.command.requester_node_id == node_id
await snapshot_sender.send(
_snapshot_chunk(state, requester_node_id=node_id, session_id=session_id)
)
assert worker.state.last_event_applied_idx == 7
assert event_router.event_buffer.next_idx_to_release == 8
@pytest.mark.asyncio
async def test_worker_event_applier_ignores_events_covered_by_snapshot() -> None:
node_id = NodeId("worker")
session_id = SessionId(master_node_id=NodeId("master"), election_clock=1)
worker, _event_router, _command_receiver, _snapshot_sender, event_sender = _worker(
node_id, session_id
)
worker.state = State(last_event_applied_idx=7)
async with anyio.create_task_group() as tg:
tg.start_soon(worker._event_applier)
await event_sender.send(IndexedEvent(idx=7, event=TestEvent()))
await event_sender.send(IndexedEvent(idx=8, event=TestEvent()))
while worker.state.last_event_applied_idx != 8:
await anyio.sleep(0.001)
tg.cancel_scope.cancel()