diff --git a/rust/exo_rs/src/lib.rs b/rust/exo_rs/src/lib.rs index 5efd9e793..322adcb60 100644 --- a/rust/exo_rs/src/lib.rs +++ b/rust/exo_rs/src/lib.rs @@ -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(()) } diff --git a/rust/exo_rs/src/mailbox.rs b/rust/exo_rs/src/mailbox.rs index 3fffdfcf5..42ce1b1de 100644 --- a/rust/exo_rs/src/mailbox.rs +++ b/rust/exo_rs/src/mailbox.rs @@ -49,3 +49,8 @@ impl Mailbox { }) } } + +pub fn mailbox_module(m: &Bound<'_, PyModule>) -> PyResult<()> { + m.add_class::()?; + Ok(()) +} diff --git a/src/exo/api/main.py b/src/exo/api/main.py index 648fd00f7..34b17e6f0 100644 --- a/src/exo/api/main.py +++ b/src/exo/api/main.py @@ -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) diff --git a/src/exo/download/coordinator.py b/src/exo/download/coordinator.py index 3dcd305f0..ed5e89585 100644 --- a/src/exo/download/coordinator.py +++ b/src/exo/download/coordinator.py @@ -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 diff --git a/src/exo/main.py b/src/exo/main.py index 4554bf1ed..d606ed6d4 100644 --- a/src/exo/main.py +++ b/src/exo/main.py @@ -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) diff --git a/src/exo/master/main.py b/src/exo/master/main.py index f4fc78aa1..d64e0a279 100644 --- a/src/exo/master/main.py +++ b/src/exo/master/main.py @@ -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: diff --git a/src/exo/master/placement.py b/src/exo/master/placement.py index ed5b76bf7..2495a5700 100644 --- a/src/exo/master/placement.py +++ b/src/exo/master/placement.py @@ -294,7 +294,6 @@ def place_instance( ) - def delete_instance( command: DeleteInstance, current_instances: Mapping[InstanceId, Instance], diff --git a/src/exo/shared/apply.py b/src/exo/shared/apply.py index d02a830d9..9ed4b7811 100644 --- a/src/exo/shared/apply.py +++ b/src/exo/shared/apply.py @@ -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 diff --git a/src/exo/shared/types/commands.py b/src/exo/shared/types/commands.py index 6e6856d84..c50f6939d 100644 --- a/src/exo/shared/types/commands.py +++ b/src/exo/shared/types/commands.py @@ -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 diff --git a/src/exo/shared/types/state.py b/src/exo/shared/types/state.py index 106d51d89..7461bc0ac 100644 --- a/src/exo/shared/types/state.py +++ b/src/exo/shared/types/state.py @@ -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]) diff --git a/src/exo/shared/types/tasks.py b/src/exo/shared/types/tasks.py index f188a13bc..a5fec8cc3 100644 --- a/src/exo/shared/types/tasks.py +++ b/src/exo/shared/types/tasks.py @@ -88,6 +88,7 @@ class ImageEdits(BaseTask): # emitted by Master class Shutdown(BaseTask): # emitted by Worker runner_id: RunnerId + Task = ( CreateRunner | DownloadModel diff --git a/src/exo/shared/types/worker/instances.py b/src/exo/shared/types/worker/instances.py index ccc141364..eb0e54f7e 100644 --- a/src/exo/shared/types/worker/instances.py +++ b/src/exo/shared/types/worker/instances.py @@ -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): diff --git a/src/exo/shared/types/worker/shards.py b/src/exo/shared/types/worker/shards.py index 0a873bec1..20d16d4b6 100644 --- a/src/exo/shared/types/worker/shards.py +++ b/src/exo/shared/types/worker/shards.py @@ -107,6 +107,4 @@ class TensorShardMetadata(BaseShardMetadata): pass -ShardMetadata = ( - PipelineShardMetadata | CfgShardMetadata | TensorShardMetadata -) +ShardMetadata = PipelineShardMetadata | CfgShardMetadata | TensorShardMetadata diff --git a/src/exo/utils/info_gatherer/info_gatherer.py b/src/exo/utils/info_gatherer/info_gatherer.py index 3d644c574..252254adb 100644 --- a/src/exo/utils/info_gatherer/info_gatherer.py +++ b/src/exo/utils/info_gatherer/info_gatherer.py @@ -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()) diff --git a/src/exo/worker/main.py b/src/exo/worker/main.py index acdaf0af2..eddde8c65 100644 --- a/src/exo/worker/main.py +++ b/src/exo/worker/main.py @@ -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 diff --git a/src/exo/worker/plan.py b/src/exo/worker/plan.py index 02113426f..63dff75dc 100644 --- a/src/exo/worker/plan.py +++ b/src/exo/worker/plan.py @@ -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, - ) diff --git a/src/exo/worker/runner/supervisor.py b/src/exo/worker/runner/supervisor.py index e16201299..30918aadf 100644 --- a/src/exo/worker/runner/supervisor.py +++ b/src/exo/worker/runner/supervisor.py @@ -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 "