[object Object]

← back to Exo

Reduce reliance on internet (#1363)

a0f4f363555744f2a9660679437be064bb2bb712 · 2026-02-03 20:03:29 +0000 · rltakashige

## Motivation

Offline users currently have to wait for every retry to fail before
being able to launch a model.
For users that restart clusters often or share API keys between devices,
we also spam HuggingFace with downloads every 5 minutes.
These issues are caused by _emit_existing_download_progress being
inefficient.

## Changes

- Only query HuggingFace once while EXO is running (assumption being
that a change should only be reflected on a new EXO session)
- Only query HuggingFace when there is an internet connection (polling
connectivity every 10 seconds)
- Request download progress if we switch from no connectivity ->
connected to reduce the wait.
- Reduce download progress sleep as it's no longer expensive (queries
cache most of the time).
- Reduce retries as 30 is way too many.

## Test Plan

### Manual Testing
Manually tested the behaviour.

### Automated Testing
None, should I add any? We do have some tests for this folder, but they
are probably not too helpful.

Files touched

Diff

commit a0f4f363555744f2a9660679437be064bb2bb712
Author: rltakashige <rl.takashige@gmail.com>
Date:   Tue Feb 3 20:03:29 2026 +0000

    Reduce reliance on internet (#1363)
    
    ## Motivation
    
    Offline users currently have to wait for every retry to fail before
    being able to launch a model.
    For users that restart clusters often or share API keys between devices,
    we also spam HuggingFace with downloads every 5 minutes.
    These issues are caused by _emit_existing_download_progress being
    inefficient.
    
    ## Changes
    
    - Only query HuggingFace once while EXO is running (assumption being
    that a change should only be reflected on a new EXO session)
    - Only query HuggingFace when there is an internet connection (polling
    connectivity every 10 seconds)
    - Request download progress if we switch from no connectivity ->
    connected to reduce the wait.
    - Reduce download progress sleep as it's no longer expensive (queries
    cache most of the time).
    - Reduce retries as 30 is way too many.
    
    ## Test Plan
    
    ### Manual Testing
    Manually tested the behaviour.
    
    ### Automated Testing
    None, should I add any? We do have some tests for this folder, but they
    are probably not too helpful.
---
 src/exo/download/coordinator.py           | 34 +++++++++++--
 src/exo/download/download_utils.py        | 82 ++++++++++++++++++++++++-------
 src/exo/download/impl_shard_downloader.py | 28 +++++++++--
 src/exo/download/shard_downloader.py      |  5 ++
 src/exo/shared/models/model_cards.py      |  8 +--
 5 files changed, 130 insertions(+), 27 deletions(-)

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,

← acb97127 Normalize TextGenerationTaskParams.input to list[InputMessag  ·  back to Exo  ·  feat: add custom HuggingFace model support (#1368) 20632789 →