maybe testable

This commit is contained in:
Evan
2026-06-03 14:32:10 +01:00
parent f5862a5579
commit c0beb6fedd
17 changed files with 111 additions and 120 deletions

View File

@@ -14,6 +14,7 @@ pub mod storage;
pub mod task;
use crate::last_value::lv_submodule;
use crate::mailbox::mailbox_module;
use crate::networking::networking_submodule;
use crate::pidfile::pidfile_submodule;
use crate::session::session_submodule;
@@ -178,6 +179,7 @@ fn main_module(m: &Bound<'_, PyModule>) -> PyResult<()> {
session_submodule(m)?;
storage_submodule(m)?;
task_submodule(m)?;
mailbox_module(m)?;
Ok(())
}

View File

@@ -49,3 +49,8 @@ impl Mailbox {
})
}
}
pub fn mailbox_module(m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_class::<Mailbox>()?;
Ok(())
}

View File

@@ -129,10 +129,10 @@ from exo.api.types.openai_responses import (
)
from exo.master.image_store import ImageStore
from exo.master.placement import (
#add_instance_to_placements,
#cancel_unnecessary_downloads,
#delete_instance,
#get_transition_events,
# add_instance_to_placements,
# cancel_unnecessary_downloads,
# delete_instance,
# get_transition_events,
place_instance,
)
from exo.shared.apply import apply
@@ -162,7 +162,6 @@ from exo.shared.types.chunks import (
)
from exo.shared.types.commands import (
CancelDownload,
Command,
DeleteDownload,
DeleteInstance,
DownloadCommand,
@@ -263,7 +262,7 @@ class API:
self.node_id: NodeId = node_id
self.last_completed_election: int = 0
self.port = port
self.aggregator = session_handle.last_value_aggregator("node_metrics")
self.aggregator = session_handle.last_value_aggregator("nodes")
self.storage = session_handle.storage_interface()
self.task_requester = session_handle.task_requester()
# TODO: Mail sender?
@@ -451,7 +450,6 @@ class API:
mail = JoinInstance(instance=instance)
self._sh.send_mail(nodes, mail.model_dump_json())
# TODO: wait until we see all nodes have appeared
return CreateInstanceResponse(
message="Command received.",
@@ -1920,13 +1918,6 @@ class API:
if removed > 0:
logger.debug(f"Cleaned up {removed} expired images")
async def _send(self, command: Command):
while self.paused:
await self.paused_ev.wait()
await self.command_sender.send(
ForwarderCommand(origin=self._system_id, command=command)
)
async def _send_download(self, command: DownloadCommand):
await self.download_command_sender.send(
ForwarderDownloadCommand(origin=self._system_id, command=command)

View File

@@ -90,7 +90,7 @@ class DownloadCoordinator:
def _publisher_for_model(self, model_id: ModelId) -> LVPublisher:
if (publisher := self._download_publishers.get(model_id)) is None:
publisher = self.session_handle.last_value_publisher(
f"node_metrics/{self.node_id}/downloads/{model_id}"
f"nodes/{self.node_id}/downloads/{model_id}"
)
self._download_publishers[model_id] = publisher

View File

@@ -128,7 +128,7 @@ class Node:
local_event_receiver=router.receiver(topics.LOCAL_EVENTS),
command_receiver=router.receiver(topics.COMMANDS),
download_command_sender=router.sender(topics.DOWNLOAD_COMMANDS),
aggregator=session_handle.last_value_aggregator("node_metrics"),
aggregator=session_handle.last_value_aggregator("nodes"),
storage=session_handle.storage_interface(),
)
@@ -234,7 +234,7 @@ class Node:
download_command_sender=self.router.sender(
topics.DOWNLOAD_COMMANDS
),
aggregator=self._sh.last_value_aggregator("node_metrics"),
aggregator=self._sh.last_value_aggregator("nodes"),
storage=self._sh.storage_interface(),
)
self._tg.start_soon(self.master.run)

View File

@@ -10,7 +10,6 @@ from exo.master.placement import (
cancel_unnecessary_downloads,
delete_instance,
get_transition_events,
place_instance,
)
from exo.master.placement_utils import find_ip_prioritised
from exo.routing.event_router import (
@@ -26,10 +25,8 @@ from exo.shared.types.commands import (
ForwarderDownloadCommand,
ImageEdits,
ImageGeneration,
PlaceInstance,
RequestEventLog,
TaskCancelled,
TaskFinished,
TestCommand,
TextGeneration,
)
@@ -43,7 +40,6 @@ from exo.shared.types.events import (
NodeGatheredInfo,
NodeTimedOut,
TaskCreated,
TaskDeleted,
TaskStatusUpdated,
TraceEventData,
TracesCollected,
@@ -376,22 +372,6 @@ class Master:
)
)
generated_events.extend(transition_events)
case PlaceInstance():
state = self.state.with_aggregator(self.aggregator)
placement = place_instance(
command,
state.topology,
state.instances,
state.node_memory,
state.node_network,
state.node_backends,
download_status=state.downloads,
node_rdma_ctl=state.node_rdma_ctl,
)
transition_events = get_transition_events(
self.state.instances, placement, self.state.tasks
)
generated_events.extend(transition_events)
case CreateInstance():
state = self.state.with_aggregator(self.aggregator)
placement = add_instance_to_placements(
@@ -419,18 +399,6 @@ class Master:
logger.warning(
f"Nonexistent command {command.cancelled_command_id} cancelled"
)
case TaskFinished():
if (
task_id := self.command_task_mapping.pop(
command.finished_command_id, None
)
) is not None:
generated_events.append(TaskDeleted(task_id=task_id))
else:
logger.warning(
f"Finished command {command.finished_command_id} finished"
)
case RequestEventLog():
# We should just be able to send everything, since other buffers will ignore old messages
# rate limit to 1000 at a time
@@ -442,6 +410,8 @@ class Master:
await self._send_indexed_event(
IndexedEvent(idx=i, event=event)
)
case other:
logger.warning(f"ONE SLIPPED THROUGH {other}")
for event in generated_events:
await self.event_sender.send(event)
except Exception as e:

View File

@@ -294,7 +294,6 @@ def place_instance(
)
def delete_instance(
command: DeleteInstance,
current_instances: Mapping[InstanceId, Instance],

View File

@@ -17,7 +17,6 @@ from exo.shared.types.events import (
RunnerStatusUpdated,
TaskAcknowledged,
TaskCreated,
TaskDeleted,
TaskFailed,
TaskStatusUpdated,
TestEvent,
@@ -84,8 +83,6 @@ def event_apply(event: Event, state: State) -> State:
return apply_runner_status_updated(event, state)
case TaskCreated():
return apply_task_created(event, state)
case TaskDeleted():
return apply_task_deleted(event, state)
case TaskFailed():
return apply_task_failed(event, state)
case TaskStatusUpdated():
@@ -144,13 +141,6 @@ def apply_task_created(event: TaskCreated, state: State) -> State:
return state.model_copy(update={"tasks": new_tasks})
def apply_task_deleted(event: TaskDeleted, state: State) -> State:
new_tasks: Mapping[TaskId, Task] = {
tid: task for tid, task in state.tasks.items() if tid != event.task_id
}
return state.model_copy(update={"tasks": new_tasks})
def apply_task_status_updated(event: TaskStatusUpdated, state: State) -> State:
if event.task_id not in state.tasks:
# maybe should raise

View File

@@ -100,8 +100,10 @@ class ForwarderDownloadCommand(FrozenModel):
origin: SystemId
command: DownloadCommand
class JoinInstance(TaggedModel):
# TODO: strip this down to less data
instance: Instance
Mail = JoinInstance

View File

@@ -11,7 +11,11 @@ from exo.shared.models import model_cards
from exo.shared.topology import Topology
from exo.shared.types.backends import Backend
from exo.shared.types.common import NodeId
from exo.shared.types.events import NodeDownloadProgress, NodeGatheredInfo
from exo.shared.types.events import (
NodeDownloadProgress,
NodeGatheredInfo,
RunnerStatusUpdated,
)
from exo.shared.types.profiling import (
DiskUsage,
MemoryUsage,
@@ -39,6 +43,7 @@ from exo.utils.info_gatherer.info_gatherer import (
from exo.utils.pydantic_ext import FrozenModel
_DOWNLOAD_PROGRESS_ADAPTER = TypeAdapter[DownloadProgress](DownloadProgress)
_RUNNER_STATUS_ADAPTER = TypeAdapter[RunnerStatus](RunnerStatus)
_GATHERED_INFO_ADAPTER = TypeAdapter[GatheredInfo](GatheredInfo)
@@ -143,6 +148,7 @@ class State(FrozenModel):
def with_aggregator(self, aggregator: LVAggregator) -> "State":
from exo.shared.apply import event_apply
state = self.model_copy()
values = aggregator.dump()
node_ids = {NodeId(key.split("/")[0]) for key in values}
@@ -177,6 +183,11 @@ class State(FrozenModel):
if len(parts) >= 3 and parts[1] == "downloads":
progress = _DOWNLOAD_PROGRESS_ADAPTER.validate_json(value)
event = NodeDownloadProgress(download_progress=progress)
elif len(parts) == 5 and parts[1] == "runners":
runner_status = _RUNNER_STATUS_ADAPTER.validate_json(value)
event = RunnerStatusUpdated(
runner_status=runner_status, runner_id=RunnerId(parts[3])
)
else:
data = _GATHERED_INFO_ADAPTER.validate_json(value)
node_id = NodeId(parts[0])

View File

@@ -88,6 +88,7 @@ class ImageEdits(BaseTask): # emitted by Master
class Shutdown(BaseTask): # emitted by Worker
runner_id: RunnerId
Task = (
CreateRunner
| DownloadModel

View File

@@ -33,7 +33,9 @@ class BaseInstance(TaggedModel):
yield rid
def primary_output_node(self) -> NodeId:
return self.shard_assignments.shards[self.shard_assignments.primary_output_node].node_id
return self.shard_assignments.shards[
self.shard_assignments.primary_output_node
].node_id
class MlxRingInstance(BaseInstance):

View File

@@ -107,6 +107,4 @@ class TensorShardMetadata(BaseShardMetadata):
pass
ShardMetadata = (
PipelineShardMetadata | CfgShardMetadata | TensorShardMetadata
)
ShardMetadata = PipelineShardMetadata | CfgShardMetadata | TensorShardMetadata

View File

@@ -411,7 +411,7 @@ class InfoGatherer:
async def send(self, info: GatheredInfo):
if (tag := info.tag()) not in self.info_senders:
self.info_senders[tag] = self.session_handle.last_value_publisher(
f"node_metrics/{self.node_id}/{tag}"
f"nodes/{self.node_id}/{tag}"
)
await self.info_senders[tag].put(info.model_dump_json())

View File

@@ -3,9 +3,9 @@ from datetime import datetime, timezone
import anyio
from anyio import fail_after, to_thread
from exo_rs import LVAggregator, SessionHandle
from exo_rs import LVAggregator, Mailbox, SessionHandle
from loguru import logger
from pydantic import ValidationError
from pydantic import TypeAdapter, ValidationError
from exo.download.download_utils import is_read_only_model_dir, resolve_existing_model
from exo.routing.event_router import (
@@ -20,6 +20,8 @@ from exo.shared.types.commands import (
DeleteInstance,
ForwarderCommand,
ForwarderDownloadCommand,
JoinInstance,
Mail,
StartDownload,
)
from exo.shared.types.common import NodeId, SystemId
@@ -47,7 +49,7 @@ from exo.shared.types.tasks import (
)
from exo.shared.types.topology import Connection, SocketConnection
from exo.shared.types.worker.downloads import DownloadCompleted
from exo.shared.types.worker.instances import InstanceId
from exo.shared.types.worker.instances import BoundInstance, InstanceId
from exo.shared.types.worker.runners import RunnerId
from exo.utils.channels import Receiver, Sender
from exo.utils.info_gatherer.info_gatherer import GatheredInfo, InfoGatherer
@@ -92,8 +94,10 @@ class Worker:
self._stopped: anyio.Event = anyio.Event()
self._sh: SessionHandle = session_handle
self.aggregator: LVAggregator = session_handle.last_value_aggregator(
"node_metrics"
"nodes"
)
self.mailbox: Mailbox = session_handle.mailbox(node_id)
self.desired_runners: dict[RunnerId, BoundInstance] = {}
async def run(self):
logger.info("Starting Worker")
@@ -107,6 +111,7 @@ class Worker:
tg.start_soon(self._event_applier)
tg.start_soon(self._poll_connection_updates)
tg.start_soon(self._reconcile_custom_cards)
tg.start_soon(self._listen_to_mailbox)
except* (EventRouterBrokenResourceError, EventRouterClosedResourceError):
# Event router has been closed (try-star syntax handles error groups)
pass
@@ -120,6 +125,26 @@ class Worker:
runner.shutdown()
self._stopped.set()
async def _listen_to_mailbox(self):
ta = TypeAdapter[Mail](Mail)
while (mail := await self.mailbox.recv()) is not None:
try:
mail = ta.validate_json(mail)
except ValidationError:
logger.warning(f"discarding corrupt mail {mail}")
continue
match mail:
case JoinInstance(instance=instance):
# no backoff - moving that inside the supervisor
for runner_id in instance.runners_for(self.node_id):
bound_instance = BoundInstance(
instance=instance,
bound_node_id=self.node_id,
bound_runner_id=runner_id,
)
self.desired_runners[runner_id] = bound_instance
async def _forward_info(self, recv: Receiver[GatheredInfo]):
with recv as info_stream:
async for info in info_stream:
@@ -164,13 +189,16 @@ class Worker:
async def plan_step(self):
while True:
await anyio.sleep(0.1)
state = self.state.with_aggregator(self.aggregator)
task: Task | None = plan(
self.node_id,
self.runners,
self.state.downloads,
self.state.instances,
self.state.runners,
self.state.tasks,
state.downloads, # comes from with_agg
{
value.instance.instance_id: value.instance
for value in self.desired_runners.values()
}, # comes from mailbox
state.runners, # comes from with_agg
self._instance_backoff,
self._download_backoff,
)
@@ -183,10 +211,11 @@ class Worker:
logger.warning(
f"Instance {iid} exceeded {EXO_MAX_INSTANCE_RETRIES} retries, requesting deletion"
)
self.desired_runners.pop(task.bound_instance.bound_runner_id)
await self.command_sender.send(
ForwarderCommand(
origin=self._system_id,
command=DeleteInstance(instance_id=iid),
command=DeleteInstance(),
)
)
continue
@@ -300,6 +329,9 @@ class Worker:
task_assignment_subscriber=self._sh.last_value_subscriber(
f"task_assignments/{task.instance_id}/*"
),
runner_status_publisher=self._sh.last_value_publisher(
f"nodes/{self.node_id}/runners/{task.bound_instance.bound_runner_id}/status"
),
task_responder=task_responder,
)
self.runners[task.bound_instance.bound_runner_id] = runner

View File

@@ -4,7 +4,6 @@ from collections.abc import Mapping, Sequence
from exo.shared.types.common import ModelId, NodeId
from exo.shared.types.tasks import (
CancelTask,
ConnectToGroup,
CreateRunner,
DownloadModel,
@@ -12,8 +11,6 @@ from exo.shared.types.tasks import (
Shutdown,
StartWarmup,
Task,
TaskId,
TaskStatus,
)
from exo.shared.types.worker.downloads import (
DownloadCompleted,
@@ -43,17 +40,17 @@ def plan(
# Runners is expected to be FRESH and so should not come from state
runners: Mapping[RunnerId, RunnerSupervisor],
global_download_status: Mapping[NodeId, Sequence[DownloadProgress]],
instances: Mapping[InstanceId, Instance],
desired_instances: Mapping[InstanceId, Instance],
all_runners: Mapping[RunnerId, RunnerStatus], # all global
tasks: Mapping[TaskId, Task],
instance_backoff: KeyedBackoff[InstanceId],
download_backoff: KeyedBackoff[ModelId],
) -> Task | None:
# Python short circuiting OR logic should evaluate these sequentially.
return (
_cancel_tasks(runners, tasks)
or _kill_runner(runners, all_runners, instances)
or _create_runner(node_id, runners, all_runners, instances, instance_backoff)
_kill_runner(runners, all_runners, desired_instances)
or _create_runner(
node_id, runners, all_runners, desired_instances, instance_backoff
)
or _model_needs_download(
node_id, runners, global_download_status, download_backoff
)
@@ -294,22 +291,3 @@ def _ready_to_warmup(
return StartWarmup(instance_id=instance.instance_id)
return None
def _cancel_tasks(
runners: Mapping[RunnerId, RunnerSupervisor],
tasks: Mapping[TaskId, Task],
) -> Task | None:
for task in tasks.values():
if task.task_status != TaskStatus.Cancelled:
continue
for runner_id, runner in runners.items():
if task.instance_id != runner.bound_instance.instance.instance_id:
continue
if task.task_id in runner.cancelled:
continue
return CancelTask(
instance_id=task.instance_id,
cancelled_task_id=task.task_id,
runner_id=runner_id,
)

View File

@@ -12,7 +12,13 @@ from anyio import (
CancelScope,
ClosedResourceError,
)
from exo_rs import LVSubscriber, TaskChunkSender, TaskRequest, TaskResponder
from exo_rs import (
LVPublisher,
LVSubscriber,
TaskChunkSender,
TaskRequest,
TaskResponder,
)
from loguru import logger
from pydantic import TypeAdapter, ValidationError
@@ -20,6 +26,7 @@ from exo.shared.constants import EXO_RUNNER_STDERR_LOG, EXO_RUNNER_STDOUT_LOG
from exo.shared.types.chunks import ErrorChunk, PrefillProgressChunk
from exo.shared.types.commands import (
Command,
DeleteInstance,
TaskCancelled,
TaskFinished,
)
@@ -39,7 +46,6 @@ from exo.shared.types.events import (
RunnerStatusUpdated,
TaskAcknowledged,
TaskCreated,
TaskDeleted,
TaskStatusUpdated,
)
from exo.shared.types.tasks import (
@@ -63,7 +69,6 @@ from exo.shared.types.worker.runners import (
RunnerStatus,
RunnerWarmingUp,
)
from exo.shared.types.worker.shards import ShardMetadata
from exo.utils.async_process import AsyncProcess
from exo.utils.channels import MpReceiver, MpSender, Receiver, Sender, mp_channel
from exo.utils.fs import ensure_parent_directory_exists
@@ -216,7 +221,6 @@ class RunnerStdioHandler:
@dataclass(eq=False)
class RunnerSupervisor:
shard_metadata: ShardMetadata
bound_instance: BoundInstance
runner_process: AsyncProcess
_runner_stdio_handler: RunnerStdioHandler
@@ -225,8 +229,9 @@ class RunnerSupervisor:
_task_sender: MpSender[Task]
_event_sender: Sender[Event]
_cancel_sender: MpSender[TaskId]
_task_responder: TaskResponder | None = None
_task_assignment_subscriber: LVSubscriber | None = None
_task_responder: TaskResponder | None
_task_assignment_subscriber: LVSubscriber
runner_status_publisher: LVPublisher
_assigned_tasks: dict[TaskId, BridgeTask] = field(default_factory=dict, init=False)
_tg: TaskGroup = field(default_factory=TaskGroup, init=False)
status: RunnerStatus = field(default_factory=RunnerIdle, init=False)
@@ -251,8 +256,9 @@ class RunnerSupervisor:
*,
bound_instance: BoundInstance,
event_sender: Sender[Event],
task_assignment_subscriber: LVSubscriber | None = None,
task_responder: TaskResponder | None = None,
task_assignment_subscriber: LVSubscriber,
runner_status_publisher: LVPublisher,
task_responder: TaskResponder | None,
initialize_timeout: float = 400,
) -> Self:
ev_send, ev_recv = mp_channel[Event | RunnerTerminationError]()
@@ -274,11 +280,8 @@ class RunnerSupervisor:
stdout_rx=runner_process.stdout, stderr_rx=runner_process.stderr
)
shard_metadata = bound_instance.bound_shard
self = cls(
bound_instance=bound_instance,
shard_metadata=shard_metadata,
runner_process=runner_process,
_runner_stdio_handler=runner_stdio_handler,
initialize_timeout=initialize_timeout,
@@ -288,6 +291,7 @@ class RunnerSupervisor:
_event_sender=event_sender,
_task_responder=task_responder,
_task_assignment_subscriber=task_assignment_subscriber,
runner_status_publisher=runner_status_publisher,
)
return self
@@ -303,11 +307,10 @@ class RunnerSupervisor:
tg.start_soon(self._forward_events)
if self._task_responder is not None:
tg.start_soon(self._run_task_responder, self._task_responder)
if self._task_assignment_subscriber is not None:
tg.start_soon(
self._run_task_assignment_subscriber,
self._task_assignment_subscriber,
)
tg.start_soon(
self._run_task_assignment_subscriber,
self._task_assignment_subscriber,
)
finally:
logger.info("Runner supervisor shutting down")
if not self._cancel_watch_runner.cancel_called:
@@ -324,6 +327,7 @@ class RunnerSupervisor:
self._cancel_sender.close()
with anyio.CancelScope(shield=True):
await self.runner_status_publisher.delete()
await self.runner_process.stop()
logger.info(
f"Runner process successfully terminated: {self.runner_process.exitcode}"
@@ -336,6 +340,8 @@ class RunnerSupervisor:
instance_id = self.bound_instance.instance.instance_id
while (received := await subscriber.recv()) is not None:
key, payload = received
if payload is None:
continue
if (ids := _task_assignment_ids(key)) is None:
continue
assigned_instance_id, assigned_task_id = ids
@@ -420,6 +426,9 @@ class RunnerSupervisor:
case RunnerStatusUpdated(runner_status=runner_status):
self.status = runner_status
await self.runner_status_publisher.put(
self.status.model_dump_json()
)
await self._event_sender.send(event)
await self._reconcile_assigned_tasks()
@@ -497,6 +506,8 @@ class RunnerSupervisor:
case TaskFinished(finished_command_id=command_id):
await self._finish_bridge_task(command_id)
request.reply(command_id)
case DeleteInstance():
self.shutdown()
case _:
request.reply_err(f"Unsupported bridge command: {command}")
except Exception as exception:
@@ -571,7 +582,6 @@ class RunnerSupervisor:
return
self.bridge_command_tasks.pop(task.command_id, None)
self.bridge_chunk_senders.pop(task.command_id, None)
await self._event_sender.send(TaskDeleted(task_id=task_id))
if self._task_responder is not None:
await self._task_responder.unassign_task(task_id)
@@ -633,7 +643,7 @@ class RunnerSupervisor:
ChunkGenerated(
command_id=task.command_id,
chunk=ErrorChunk(
model=self.shard_metadata.model_card.model_id,
model=self.bound_instance.bound_shard.model_card.model_id,
diagnostics=diagnostics,
error_message=(
"Runner shutdown before completing command "