diff --git a/src/exo/download/coordinator.py b/src/exo/download/coordinator.py index c2f7b9e9..f5798ad3 100644 --- a/src/exo/download/coordinator.py +++ b/src/exo/download/coordinator.py @@ -1,4 +1,5 @@ import asyncio +import socket from dataclasses import dataclass, field from typing import Iterator @@ -60,10 +61,37 @@ class DownloadCoordinator: async def run(self) -> None: logger.info("Starting DownloadCoordinator") + self._test_internet_connection() async with self._tg as tg: tg.start_soon(self._command_processor) tg.start_soon(self._forward_events) tg.start_soon(self._emit_existing_download_progress) + tg.start_soon(self._check_internet_connection) + + def _test_internet_connection(self) -> None: + try: + socket.create_connection(("1.1.1.1", 443), timeout=3).close() + self.shard_downloader.set_internet_connection(True) + except OSError: + self.shard_downloader.set_internet_connection(False) + logger.debug( + f"Internet connectivity: {self.shard_downloader.internet_connection}" + ) + + async def _check_internet_connection(self) -> None: + first_connection = True + while True: + await asyncio.sleep(10) + + # Assume that internet connection is set to False on 443 errors. + if self.shard_downloader.internet_connection: + continue + + self._test_internet_connection() + + if first_connection and self.shard_downloader.internet_connection: + first_connection = False + self._tg.start_soon(self._emit_existing_download_progress) def shutdown(self) -> None: self._tg.cancel_scope.cancel() @@ -241,7 +269,7 @@ class DownloadCoordinator: async def _emit_existing_download_progress(self) -> None: try: while True: - logger.info( + logger.debug( "DownloadCoordinator: Fetching and emitting existing download progress..." ) async for ( @@ -274,10 +302,10 @@ class DownloadCoordinator: await self.event_sender.send( NodeDownloadProgress(download_progress=status) ) - logger.info( + logger.debug( "DownloadCoordinator: Done emitting existing download progress." ) - await anyio.sleep(5 * 60) # 5 minutes + await anyio.sleep(60) except Exception as e: logger.error( f"DownloadCoordinator: Error emitting existing download progress: {e}" diff --git a/src/exo/download/download_utils.py b/src/exo/download/download_utils.py index 6dec4718..618e4f38 100644 --- a/src/exo/download/download_utils.py +++ b/src/exo/download/download_utils.py @@ -49,6 +49,10 @@ class HuggingFaceAuthenticationError(Exception): """Raised when HuggingFace returns 401/403 for a model download.""" +class HuggingFaceRateLimitError(Exception): + """429 Huggingface code""" + + async def _build_auth_error_message(status_code: int, model_id: ModelId) -> str: token = await get_hf_token() if status_code == 401 and token is None: @@ -154,49 +158,76 @@ async def seed_models(seed_dir: str | Path): logger.error(traceback.format_exc()) +_fetched_file_lists_this_session: set[str] = set() + + async def fetch_file_list_with_cache( - model_id: ModelId, revision: str = "main", recursive: bool = False + model_id: ModelId, + revision: str = "main", + recursive: bool = False, + skip_internet: bool = False, + on_connection_lost: Callable[[], None] = lambda: None, ) -> list[FileListEntry]: target_dir = (await ensure_models_dir()) / "caches" / model_id.normalize() await aios.makedirs(target_dir, exist_ok=True) cache_file = target_dir / f"{model_id.normalize()}--{revision}--file_list.json" + cache_key = f"{model_id.normalize()}--{revision}" + + if cache_key in _fetched_file_lists_this_session and await aios.path.exists( + cache_file + ): + async with aiofiles.open(cache_file, "r") as f: + return TypeAdapter(list[FileListEntry]).validate_json(await f.read()) + + if skip_internet: + if await aios.path.exists(cache_file): + async with aiofiles.open(cache_file, "r") as f: + return TypeAdapter(list[FileListEntry]).validate_json(await f.read()) + raise FileNotFoundError( + f"No internet connection and no cached file list for {model_id}" + ) - # Always try fresh first try: file_list = await fetch_file_list_with_retry( - model_id, revision, recursive=recursive + model_id, + revision, + recursive=recursive, + on_connection_lost=on_connection_lost, ) - # Update cache with fresh data async with aiofiles.open(cache_file, "w") as f: await f.write( TypeAdapter(list[FileListEntry]).dump_json(file_list).decode() ) + _fetched_file_lists_this_session.add(cache_key) return file_list except Exception as e: - # Fetch failed - try cache fallback if await aios.path.exists(cache_file): logger.warning( f"Failed to fetch file list for {model_id}, using cached data: {e}" ) async with aiofiles.open(cache_file, "r") as f: return TypeAdapter(list[FileListEntry]).validate_json(await f.read()) - # No cache available, propagate the error - raise + raise FileNotFoundError(f"Failed to fetch file list for {model_id}: {e}") from e async def fetch_file_list_with_retry( - model_id: ModelId, revision: str = "main", path: str = "", recursive: bool = False + model_id: ModelId, + revision: str = "main", + path: str = "", + recursive: bool = False, + on_connection_lost: Callable[[], None] = lambda: None, ) -> list[FileListEntry]: - n_attempts = 30 + n_attempts = 3 for attempt in range(n_attempts): try: return await _fetch_file_list(model_id, revision, path, recursive) except HuggingFaceAuthenticationError: raise except Exception as e: + on_connection_lost() if attempt == n_attempts - 1: raise e - await asyncio.sleep(min(8, 0.1 * float(2.0 ** int(attempt)))) + await asyncio.sleep(2.0**attempt) raise Exception( f"Failed to fetch file list for {model_id=} {revision=} {path=} {recursive=}" ) @@ -216,7 +247,11 @@ async def _fetch_file_list( if response.status in [401, 403]: msg = await _build_auth_error_message(response.status, model_id) raise HuggingFaceAuthenticationError(msg) - if response.status == 200: + elif response.status == 429: + raise HuggingFaceRateLimitError( + f"Couldn't download {model_id} because of HuggingFace rate limit." + ) + elif response.status == 200: data_json = await response.text() data = TypeAdapter(list[FileListEntry]).validate_json(data_json) files: list[FileListEntry] = [] @@ -249,7 +284,7 @@ def create_http_session( else: total_timeout = 1800 connect_timeout = 60 - sock_read_timeout = 1800 + sock_read_timeout = 60 sock_connect_timeout = 60 ssl_context = ssl.create_default_context( @@ -324,8 +359,9 @@ async def download_file_with_retry( path: str, target_dir: Path, on_progress: Callable[[int, int, bool], None] = lambda _, __, ___: None, + on_connection_lost: Callable[[], None] = lambda: None, ) -> Path: - n_attempts = 30 + n_attempts = 3 for attempt in range(n_attempts): try: return await _download_file( @@ -333,14 +369,19 @@ async def download_file_with_retry( ) except HuggingFaceAuthenticationError: raise - except Exception as e: - if isinstance(e, FileNotFoundError) or attempt == n_attempts - 1: + except HuggingFaceRateLimitError as e: + if attempt == n_attempts - 1: raise e logger.error( f"Download error on attempt {attempt}/{n_attempts} for {model_id=} {revision=} {path=} {target_dir=}" ) logger.error(traceback.format_exc()) - await asyncio.sleep(min(8, 0.1 * (2.0**attempt))) + await asyncio.sleep(2.0**attempt) + except Exception as e: + on_connection_lost() + if attempt == n_attempts - 1: + raise e + break raise Exception( f"Failed to download file {model_id=} {revision=} {path=} {target_dir=}" ) @@ -542,7 +583,9 @@ async def download_shard( on_progress: Callable[[ShardMetadata, RepoDownloadProgress], Awaitable[None]], max_parallel_downloads: int = 8, skip_download: bool = False, + skip_internet: bool = False, allow_patterns: list[str] | None = None, + on_connection_lost: Callable[[], None] = lambda: None, ) -> tuple[Path, RepoDownloadProgress]: if not skip_download: logger.debug(f"Downloading {shard.model_card.model_id=}") @@ -562,7 +605,11 @@ async def download_shard( all_start_time = time.time() file_list = await fetch_file_list_with_cache( - shard.model_card.model_id, revision, recursive=True + shard.model_card.model_id, + revision, + recursive=True, + skip_internet=skip_internet, + on_connection_lost=on_connection_lost, ) filtered_file_list = list( filter_repo_objects( @@ -672,6 +719,7 @@ async def download_shard( lambda curr_bytes, total_bytes, is_renamed: schedule_progress( file, curr_bytes, total_bytes, is_renamed ), + on_connection_lost=on_connection_lost, ) if not skip_download: diff --git a/src/exo/download/impl_shard_downloader.py b/src/exo/download/impl_shard_downloader.py index 1b7f5eab..0e7aea1e 100644 --- a/src/exo/download/impl_shard_downloader.py +++ b/src/exo/download/impl_shard_downloader.py @@ -1,4 +1,5 @@ import asyncio +from asyncio import create_task from collections.abc import Awaitable from pathlib import Path from typing import AsyncIterator, Callable @@ -49,6 +50,10 @@ class SingletonShardDownloader(ShardDownloader): self.shard_downloader = shard_downloader self.active_downloads: dict[ShardMetadata, asyncio.Task[Path]] = {} + def set_internet_connection(self, value: bool) -> None: + self.internet_connection = value + self.shard_downloader.set_internet_connection(value) + def on_progress( self, callback: Callable[[ShardMetadata, RepoDownloadProgress], Awaitable[None]], @@ -85,6 +90,10 @@ class CachedShardDownloader(ShardDownloader): self.shard_downloader = shard_downloader self.cache: dict[tuple[str, ShardMetadata], Path] = {} + def set_internet_connection(self, value: bool) -> None: + self.internet_connection = value + self.shard_downloader.set_internet_connection(value) + def on_progress( self, callback: Callable[[ShardMetadata, RepoDownloadProgress], Awaitable[None]], @@ -142,6 +151,8 @@ class ResumableShardDownloader(ShardDownloader): self.on_progress_wrapper, max_parallel_downloads=self.max_parallel_downloads, allow_patterns=allow_patterns, + skip_internet=not self.internet_connection, + on_connection_lost=lambda: self.set_internet_connection(False), ) return target_dir @@ -154,12 +165,23 @@ class ResumableShardDownloader(ShardDownloader): """Helper coroutine that builds the shard for a model and gets its download status.""" shard = await build_full_shard(model_id) return await download_shard( - shard, self.on_progress_wrapper, skip_download=True + shard, + self.on_progress_wrapper, + skip_download=True, + skip_internet=not self.internet_connection, + on_connection_lost=lambda: self.set_internet_connection(False), ) - # Kick off download status coroutines concurrently + semaphore = asyncio.Semaphore(self.max_parallel_downloads) + + async def download_with_semaphore( + model_card: ModelCard, + ) -> tuple[Path, RepoDownloadProgress]: + async with semaphore: + return await _status_for_model(model_card.model_id) + tasks = [ - asyncio.create_task(_status_for_model(model_card.model_id)) + create_task(download_with_semaphore(model_card)) for model_card in await get_model_cards() ] diff --git a/src/exo/download/shard_downloader.py b/src/exo/download/shard_downloader.py index 30c11d25..9dd8c324 100644 --- a/src/exo/download/shard_downloader.py +++ b/src/exo/download/shard_downloader.py @@ -16,6 +16,11 @@ from exo.shared.types.worker.shards import ( # TODO: the PipelineShardMetadata getting reinstantiated is a bit messy. Should this be a classmethod? class ShardDownloader(ABC): + internet_connection: bool = False + + def set_internet_connection(self, value: bool) -> None: + self.internet_connection = value + @abstractmethod async def ensure_shard( self, shard: ShardMetadata, config_only: bool = False diff --git a/src/exo/shared/models/model_cards.py b/src/exo/shared/models/model_cards.py index bf2a892f..58d84849 100644 --- a/src/exo/shared/models/model_cards.py +++ b/src/exo/shared/models/model_cards.py @@ -108,9 +108,9 @@ class ModelCard(CamelCaseModel): async def fetch_from_hf(model_id: ModelId) -> "ModelCard": """Fetches storage size and number of layers for a Hugging Face model, returns Pydantic ModelMeta.""" # TODO: failure if files do not exist - config_data = await get_config_data(model_id) + config_data = await fetch_config_data(model_id) num_layers = config_data.layer_count - mem_size_bytes = await get_safetensors_size(model_id) + mem_size_bytes = await fetch_safetensors_size(model_id) mc = ModelCard( model_id=ModelId(model_id), @@ -258,7 +258,7 @@ class ConfigData(BaseModel): return data -async def get_config_data(model_id: ModelId) -> ConfigData: +async def fetch_config_data(model_id: ModelId) -> ConfigData: """Downloads and parses config.json for a model.""" from exo.download.download_utils import ( download_file_with_retry, @@ -280,7 +280,7 @@ async def get_config_data(model_id: ModelId) -> ConfigData: return ConfigData.model_validate_json(await f.read()) -async def get_safetensors_size(model_id: ModelId) -> Memory: +async def fetch_safetensors_size(model_id: ModelId) -> Memory: """Gets model size from safetensors index or falls back to HF API.""" from exo.download.download_utils import ( download_file_with_retry,