mirror of
https://github.com/exo-explore/exo.git
synced 2026-09-08 11:35:40 -04:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f10e6ab1b7 | ||
|
|
7f229425d6 | ||
|
|
62fd3ae36c | ||
|
|
c0d4fd7fa7 | ||
|
|
7a25b4186d | ||
|
|
dbc6286066 | ||
|
|
e1a79a0918 | ||
|
|
ea04549692 | ||
|
|
a2c2a0fc92 |
No files matched your search
+130
-34
@@ -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
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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()
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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 == {}
|
||||
@@ -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
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
Reference in new issue
Block a user