mirror of
https://github.com/exo-explore/exo.git
synced 2026-07-30 15:17:39 -04:00
maybe testable
This commit is contained in:
@@ -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(())
|
||||
}
|
||||
|
||||
@@ -49,3 +49,8 @@ impl Mailbox {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub fn mailbox_module(m: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
m.add_class::<Mailbox>()?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -294,7 +294,6 @@ def place_instance(
|
||||
)
|
||||
|
||||
|
||||
|
||||
def delete_instance(
|
||||
command: DeleteInstance,
|
||||
current_instances: Mapping[InstanceId, Instance],
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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])
|
||||
|
||||
@@ -88,6 +88,7 @@ class ImageEdits(BaseTask): # emitted by Master
|
||||
class Shutdown(BaseTask): # emitted by Worker
|
||||
runner_id: RunnerId
|
||||
|
||||
|
||||
Task = (
|
||||
CreateRunner
|
||||
| DownloadModel
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -107,6 +107,4 @@ class TensorShardMetadata(BaseShardMetadata):
|
||||
pass
|
||||
|
||||
|
||||
ShardMetadata = (
|
||||
PipelineShardMetadata | CfgShardMetadata | TensorShardMetadata
|
||||
)
|
||||
ShardMetadata = PipelineShardMetadata | CfgShardMetadata | TensorShardMetadata
|
||||
|
||||
@@ -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())
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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 "
|
||||
|
||||
Reference in New Issue
Block a user