diff --git a/dashboard/src/lib/stores/app.svelte.ts b/dashboard/src/lib/stores/app.svelte.ts index 352a4c07..c5a54104 100644 --- a/dashboard/src/lib/stores/app.svelte.ts +++ b/dashboard/src/lib/stores/app.svelte.ts @@ -3256,6 +3256,31 @@ class AppStore { } } + /** + * Cancel/pause an active download on a specific node + */ + async cancelDownload(nodeId: string, modelId: string): Promise { + try { + const response = await fetch("/download/cancel", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ + targetNodeId: nodeId, + modelId: modelId, + }), + }); + if (!response.ok) { + const errorText = await response.text(); + throw new Error( + `Failed to cancel download: ${response.status} - ${errorText}`, + ); + } + } catch (error) { + console.error("Error cancelling download:", error); + throw error; + } + } + /** * Delete a downloaded model from a specific node */ @@ -3477,6 +3502,8 @@ export const resetImageGenerationParams = () => // Download actions export const startDownload = (nodeId: string, shardMetadata: object) => appStore.startDownload(nodeId, shardMetadata); +export const cancelDownload = (nodeId: string, modelId: string) => + appStore.cancelDownload(nodeId, modelId); export const deleteDownload = (nodeId: string, modelId: string) => appStore.deleteDownload(nodeId, modelId); diff --git a/dashboard/src/routes/downloads/+page.svelte b/dashboard/src/routes/downloads/+page.svelte index e3420005..4bf69cb6 100644 --- a/dashboard/src/routes/downloads/+page.svelte +++ b/dashboard/src/routes/downloads/+page.svelte @@ -9,6 +9,7 @@ refreshState, lastUpdate as lastUpdateStore, startDownload, + cancelDownload, deleteDownload, } from "$lib/stores/app.svelte"; import { @@ -349,6 +350,59 @@ }); +{#snippet trashIcon()} + + + +{/snippet} + +{#snippet downloadIcon(size?: string)} + + + +{/snippet} + +{#snippet pauseIcon()} + + + +{/snippet} + +{#snippet deleteButton(nodeId: string, modelId: string)} + +{/snippet} +
@@ -486,27 +540,7 @@ {formatBytes(cell.totalBytes)} - + {@render deleteButton(col.nodeId, row.modelId)}
{:else if cell.kind === "downloading"}
{formatSpeed(cell.speed)} +
+ + {@render deleteButton(col.nodeId, row.modelId)} +
{:else if cell.kind === "pending"}
- {#if row.shardMetadata} - - {:else} - paused - {/if} + {@render downloadIcon()} + + {:else} + paused + {/if} + {@render deleteButton(col.nodeId, row.modelId)} + {:else if row.shardMetadata} {:else} ... @@ -626,29 +652,20 @@ clip-rule="evenodd" > - {#if row.shardMetadata} - - {/if} + {@render downloadIcon()} + + {/if} + {@render deleteButton(col.nodeId, row.modelId)} + {:else}
- - - + {@render downloadIcon()} {/if}
diff --git a/src/exo/api/main.py b/src/exo/api/main.py index 5635a409..685281a3 100644 --- a/src/exo/api/main.py +++ b/src/exo/api/main.py @@ -55,6 +55,8 @@ from exo.api.types import ( BenchImageGenerationResponse, BenchImageGenerationTaskParams, CancelCommandResponse, + CancelDownloadParams, + CancelDownloadResponse, ChatCompletionChoice, ChatCompletionMessage, ChatCompletionRequest, @@ -147,6 +149,7 @@ from exo.shared.types.chunks import ( ) from exo.shared.types.commands import ( AddCustomModelCard, + CancelDownload, Command, CreateInstance, DeleteCustomModelCard, @@ -351,6 +354,7 @@ class API: self.app.get("/events")(self.stream_events) self.app.post("/download/start")(self.start_download) self.app.delete("/download/{node_id}/{model_id:path}")(self.delete_download) + self.app.post("/download/cancel")(self.cancel_download) self.app.get("/v1/traces")(self.list_traces) self.app.post("/v1/traces/delete")(self.delete_traces) self.app.get("/v1/traces/{task_id}")(self.get_trace) @@ -1879,6 +1883,17 @@ class API: await self._send_download(command) return DeleteDownloadResponse(command_id=command.command_id) + async def cancel_download( + self, + payload: CancelDownloadParams, + ) -> CancelDownloadResponse: + command = CancelDownload( + target_node_id=payload.target_node_id, + model_id=payload.model_id, + ) + await self._send_download(command) + return CancelDownloadResponse(command_id=command.command_id) + @staticmethod def _get_trace_path(task_id: str) -> Path: trace_path = EXO_TRACING_CACHE_DIR / f"trace_{task_id}.json" diff --git a/src/exo/api/types/__init__.py b/src/exo/api/types/__init__.py index 5821bdf4..3a61c1cd 100644 --- a/src/exo/api/types/__init__.py +++ b/src/exo/api/types/__init__.py @@ -5,6 +5,8 @@ from .api import BenchChatCompletionResponse as BenchChatCompletionResponse from .api import BenchImageGenerationResponse as BenchImageGenerationResponse from .api import BenchImageGenerationTaskParams as BenchImageGenerationTaskParams from .api import CancelCommandResponse as CancelCommandResponse +from .api import CancelDownloadParams as CancelDownloadParams +from .api import CancelDownloadResponse as CancelDownloadResponse from .api import ChatCompletionChoice as ChatCompletionChoice from .api import ChatCompletionContentPart as ChatCompletionContentPart from .api import ChatCompletionMessage as ChatCompletionMessage diff --git a/src/exo/api/types/api.py b/src/exo/api/types/api.py index b4fa3147..709ad144 100644 --- a/src/exo/api/types/api.py +++ b/src/exo/api/types/api.py @@ -430,6 +430,15 @@ class DeleteDownloadResponse(CamelCaseModel): command_id: CommandId +class CancelDownloadParams(CamelCaseModel): + target_node_id: NodeId + model_id: ModelId + + +class CancelDownloadResponse(CamelCaseModel): + command_id: CommandId + + class TraceEventResponse(CamelCaseModel): name: str start_us: int diff --git a/src/exo/download/tests/test_cancel_download.py b/src/exo/download/tests/test_cancel_download.py new file mode 100644 index 00000000..740e28bc --- /dev/null +++ b/src/exo/download/tests/test_cancel_download.py @@ -0,0 +1,300 @@ +"""Tests for cancelling (pausing) an active download via CancelDownload command.""" + +import asyncio +import contextlib +from collections.abc import AsyncIterator, Awaitable +from datetime import timedelta +from pathlib import Path +from typing import Callable + +from exo.download.coordinator import DownloadCoordinator +from exo.download.download_utils import RepoDownloadProgress +from exo.download.impl_shard_downloader import SingletonShardDownloader +from exo.download.shard_downloader import ShardDownloader +from exo.shared.models.model_cards import ModelCard, ModelId, ModelTask +from exo.shared.types.commands import ( + CancelDownload, + ForwarderDownloadCommand, + StartDownload, +) +from exo.shared.types.common import NodeId, SystemId +from exo.shared.types.events import Event, NodeDownloadProgress +from exo.shared.types.memory import Memory +from exo.shared.types.worker.downloads import DownloadPending +from exo.shared.types.worker.shards import PipelineShardMetadata, ShardMetadata +from exo.utils.channels import Receiver, Sender, channel + +NODE_ID = NodeId("aaaaaaaa-aaaa-4aaa-8aaa-aaaaaaaaaaaa") +MODEL_ID = ModelId("test-org/test-model") + + +def _make_shard(model_id: ModelId = MODEL_ID) -> ShardMetadata: + return PipelineShardMetadata( + model_card=ModelCard( + model_id=model_id, + storage_size=Memory.from_mb(100), + n_layers=28, + hidden_size=1024, + supports_tensor=False, + tasks=[ModelTask.TextGeneration], + ), + device_rank=0, + world_size=1, + start_layer=0, + end_layer=28, + n_layers=28, + ) + + +class SlowShardDownloader(ShardDownloader): + """Fake downloader that blocks during ensure_shard until cancelled, + simulating a long-running download.""" + + def __init__(self) -> None: + self._progress_callbacks: list[ + Callable[[ShardMetadata, RepoDownloadProgress], Awaitable[None]] + ] = [] + self.download_started = asyncio.Event() + + def on_progress( + self, + callback: Callable[[ShardMetadata, RepoDownloadProgress], Awaitable[None]], + ) -> None: + self._progress_callbacks.append(callback) + + async def ensure_shard( + self, + shard: ShardMetadata, + config_only: bool = False, # noqa: ARG002 + ) -> Path: + # Fire an in-progress callback, then block forever (until cancelled) + progress = RepoDownloadProgress( + repo_id=str(shard.model_card.model_id), + repo_revision="main", + shard=shard, + completed_files=0, + total_files=1, + downloaded=Memory.from_mb(50), + downloaded_this_session=Memory.from_mb(50), + total=Memory.from_mb(100), + overall_speed=1024 * 1024, + overall_eta=timedelta(seconds=50), + status="in_progress", + ) + for cb in self._progress_callbacks: + await cb(shard, progress) + self.download_started.set() + # Block until cancelled + await asyncio.Event().wait() + return ( + Path("/fake/models") / shard.model_card.model_id.normalize() + ) # pragma: no cover + + async def get_shard_download_status( + self, + ) -> AsyncIterator[tuple[Path, RepoDownloadProgress]]: + if False: # noqa: SIM108 # empty async generator + yield ( + Path(), + RepoDownloadProgress( # pyright: ignore[reportUnreachable] + repo_id="", + repo_revision="", + shard=_make_shard(), + completed_files=0, + total_files=0, + downloaded=Memory.from_bytes(0), + downloaded_this_session=Memory.from_bytes(0), + total=Memory.from_bytes(0), + overall_speed=0, + overall_eta=timedelta(seconds=0), + status="not_started", + ), + ) + + async def get_shard_download_status_for_shard( + self, + shard: ShardMetadata, + ) -> RepoDownloadProgress: + return RepoDownloadProgress( + repo_id=str(shard.model_card.model_id), + repo_revision="main", + shard=shard, + completed_files=0, + total_files=1, + downloaded=Memory.from_bytes(0), + downloaded_this_session=Memory.from_bytes(0), + total=Memory.from_mb(100), + overall_speed=0, + overall_eta=timedelta(seconds=0), + status="not_started", + ) + + +def _setup_coordinator( + downloader: ShardDownloader, +) -> tuple[ + DownloadCoordinator, + Sender[ForwarderDownloadCommand], + Receiver[Event], +]: + cmd_send, cmd_recv = channel[ForwarderDownloadCommand]() + event_send, event_recv = channel[Event]() + wrapped = SingletonShardDownloader(downloader) + coordinator = DownloadCoordinator( + node_id=NODE_ID, + shard_downloader=wrapped, + download_command_receiver=cmd_recv, + event_sender=event_send, + ) + return coordinator, cmd_send, event_recv + + +async def _wait_for_pending( + event_recv: Receiver[Event], model_id: ModelId, timeout: float = 2.0 +) -> DownloadPending | None: + """Drain events until we see a DownloadPending for the given model, or timeout.""" + try: + async with asyncio.timeout(timeout): + while True: + event = await event_recv.receive() + if ( + isinstance(event, NodeDownloadProgress) + and isinstance(event.download_progress, DownloadPending) + and event.download_progress.shard_metadata.model_card.model_id + == model_id + ): + return event.download_progress + except TimeoutError: + return None + + +async def test_cancel_active_download_transitions_to_pending() -> None: + """Cancelling an in-progress download should emit a DownloadPending event + and remove the model from active_downloads.""" + slow_downloader = SlowShardDownloader() + coordinator, cmd_send, event_recv = _setup_coordinator(slow_downloader) + shard = _make_shard() + origin = SystemId("test") + + coordinator_task = asyncio.create_task(coordinator.run()) + try: + # Start a download + await cmd_send.send( + ForwarderDownloadCommand( + origin=origin, + command=StartDownload(target_node_id=NODE_ID, shard_metadata=shard), + ) + ) + + # Wait for the download to actually start (blocking in ensure_shard) + await asyncio.wait_for(slow_downloader.download_started.wait(), timeout=2.0) + + # Cancel the download + await cmd_send.send( + ForwarderDownloadCommand( + origin=origin, + command=CancelDownload(target_node_id=NODE_ID, model_id=MODEL_ID), + ) + ) + + # Should receive a DownloadPending event + pending = await _wait_for_pending(event_recv, MODEL_ID) + assert pending is not None, "Cancel should emit DownloadPending" + assert pending.shard_metadata.model_card.model_id == MODEL_ID + + # Give coordinator time to clean up + await asyncio.sleep(0.05) + + # Model should no longer be in active_downloads + assert MODEL_ID not in coordinator.active_downloads + # But should still be in download_status as pending + assert MODEL_ID in coordinator.download_status + assert isinstance(coordinator.download_status[MODEL_ID], DownloadPending) + finally: + await coordinator.shutdown() + coordinator_task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await coordinator_task + + +async def test_cancel_nonexistent_download_is_noop() -> None: + """Cancelling a model that isn't being downloaded should be a no-op.""" + slow_downloader = SlowShardDownloader() + coordinator, cmd_send, event_recv = _setup_coordinator(slow_downloader) + origin = SystemId("test") + + coordinator_task = asyncio.create_task(coordinator.run()) + try: + # Cancel a model that was never started + await cmd_send.send( + ForwarderDownloadCommand( + origin=origin, + command=CancelDownload(target_node_id=NODE_ID, model_id=MODEL_ID), + ) + ) + + # Should NOT receive any DownloadPending event + pending = await _wait_for_pending(event_recv, MODEL_ID, timeout=0.5) + assert pending is None, "Cancel of non-existent download should not emit events" + + # Coordinator state should be empty + assert MODEL_ID not in coordinator.active_downloads + assert MODEL_ID not in coordinator.download_status + finally: + await coordinator.shutdown() + coordinator_task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await coordinator_task + + +async def test_cancel_then_resume_download() -> None: + """After cancelling, re-issuing StartDownload should restart the download.""" + slow_downloader = SlowShardDownloader() + coordinator, cmd_send, event_recv = _setup_coordinator(slow_downloader) + shard = _make_shard() + origin = SystemId("test") + + coordinator_task = asyncio.create_task(coordinator.run()) + try: + # Start download + await cmd_send.send( + ForwarderDownloadCommand( + origin=origin, + command=StartDownload(target_node_id=NODE_ID, shard_metadata=shard), + ) + ) + await asyncio.wait_for(slow_downloader.download_started.wait(), timeout=2.0) + + # Cancel + await cmd_send.send( + ForwarderDownloadCommand( + origin=origin, + command=CancelDownload(target_node_id=NODE_ID, model_id=MODEL_ID), + ) + ) + pending = await _wait_for_pending(event_recv, MODEL_ID) + assert pending is not None, "Cancel should emit DownloadPending" + + await asyncio.sleep(0.05) + + # Reset the event so we can detect the next download start + slow_downloader.download_started.clear() + + # Resume by sending StartDownload again + await cmd_send.send( + ForwarderDownloadCommand( + origin=origin, + command=StartDownload(target_node_id=NODE_ID, shard_metadata=shard), + ) + ) + + # The download should restart + await asyncio.wait_for(slow_downloader.download_started.wait(), timeout=2.0) + assert MODEL_ID in coordinator.active_downloads, ( + "Model should be actively downloading again after resume" + ) + finally: + await coordinator.shutdown() + coordinator_task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await coordinator_task