feat(ui): show live scan progress and ETA for cache refresh

Broadcast typed scan_progress messages over /ws/fetch-progress from the
manual refresh/rebuild paths of ModelScanner and RecipeScanner, and
render percent, processed/total, current file name and an EMA-smoothed
ETA in the loading overlay. Hardcoded refresh strings move to i18n
(common.scanProgress); WS connection failure falls back to the previous
static loading behavior.
This commit is contained in:
Will Miao
2026-09-03 11:38:27 +08:00
parent da71985c3e
commit 14da8a6f17
19 changed files with 1668 additions and 32 deletions
+21
View File
@@ -50,6 +50,27 @@
"mb": "MB", "mb": "MB",
"gb": "GB", "gb": "GB",
"tb": "TB" "tb": "TB"
},
"scanProgress": {
"refreshing": "[TODO: Translate] Refreshing {type}s...",
"fullRebuilding": "[TODO: Translate] Full rebuild {type}s...",
"actionRefresh": "[TODO: Translate] Refresh",
"actionFullRebuild": "[TODO: Translate] Full rebuild",
"actionRefreshLower": "[TODO: Translate] refresh",
"actionRebuildLower": "[TODO: Translate] rebuild",
"stages": {
"scan_folders": "[TODO: Translate] Scanning folders...",
"count_models": "[TODO: Translate] Found {total} files",
"process_models": "[TODO: Translate] Processing models",
"reconcile_scan": "[TODO: Translate] Checking for changes...",
"process_new": "[TODO: Translate] Processing new models",
"finalizing": "[TODO: Translate] Finalizing..."
},
"eta": {
"lessThanMinute": "[TODO: Translate] Less than a minute remaining",
"minutes": "[TODO: Translate] ~{minutes} min remaining",
"hours": "[TODO: Translate] ~{hours} hr {minutes} min remaining"
}
} }
}, },
"onboarding": { "onboarding": {
+21
View File
@@ -50,6 +50,27 @@
"mb": "MB", "mb": "MB",
"gb": "GB", "gb": "GB",
"tb": "TB" "tb": "TB"
},
"scanProgress": {
"refreshing": "Refreshing {type}s...",
"fullRebuilding": "Full rebuild {type}s...",
"actionRefresh": "Refresh",
"actionFullRebuild": "Full rebuild",
"actionRefreshLower": "refresh",
"actionRebuildLower": "rebuild",
"stages": {
"scan_folders": "Scanning folders...",
"count_models": "Found {total} files",
"process_models": "Processing models",
"reconcile_scan": "Checking for changes...",
"process_new": "Processing new models",
"finalizing": "Finalizing..."
},
"eta": {
"lessThanMinute": "Less than a minute remaining",
"minutes": "~{minutes} min remaining",
"hours": "~{hours} hr {minutes} min remaining"
}
} }
}, },
"onboarding": { "onboarding": {
+21
View File
@@ -50,6 +50,27 @@
"mb": "MB", "mb": "MB",
"gb": "GB", "gb": "GB",
"tb": "TB" "tb": "TB"
},
"scanProgress": {
"refreshing": "[TODO: Translate] Refreshing {type}s...",
"fullRebuilding": "[TODO: Translate] Full rebuild {type}s...",
"actionRefresh": "[TODO: Translate] Refresh",
"actionFullRebuild": "[TODO: Translate] Full rebuild",
"actionRefreshLower": "[TODO: Translate] refresh",
"actionRebuildLower": "[TODO: Translate] rebuild",
"stages": {
"scan_folders": "[TODO: Translate] Scanning folders...",
"count_models": "[TODO: Translate] Found {total} files",
"process_models": "[TODO: Translate] Processing models",
"reconcile_scan": "[TODO: Translate] Checking for changes...",
"process_new": "[TODO: Translate] Processing new models",
"finalizing": "[TODO: Translate] Finalizing..."
},
"eta": {
"lessThanMinute": "[TODO: Translate] Less than a minute remaining",
"minutes": "[TODO: Translate] ~{minutes} min remaining",
"hours": "[TODO: Translate] ~{hours} hr {minutes} min remaining"
}
} }
}, },
"onboarding": { "onboarding": {
+21
View File
@@ -50,6 +50,27 @@
"mb": "Mo", "mb": "Mo",
"gb": "Go", "gb": "Go",
"tb": "To" "tb": "To"
},
"scanProgress": {
"refreshing": "[TODO: Translate] Refreshing {type}s...",
"fullRebuilding": "[TODO: Translate] Full rebuild {type}s...",
"actionRefresh": "[TODO: Translate] Refresh",
"actionFullRebuild": "[TODO: Translate] Full rebuild",
"actionRefreshLower": "[TODO: Translate] refresh",
"actionRebuildLower": "[TODO: Translate] rebuild",
"stages": {
"scan_folders": "[TODO: Translate] Scanning folders...",
"count_models": "[TODO: Translate] Found {total} files",
"process_models": "[TODO: Translate] Processing models",
"reconcile_scan": "[TODO: Translate] Checking for changes...",
"process_new": "[TODO: Translate] Processing new models",
"finalizing": "[TODO: Translate] Finalizing..."
},
"eta": {
"lessThanMinute": "[TODO: Translate] Less than a minute remaining",
"minutes": "[TODO: Translate] ~{minutes} min remaining",
"hours": "[TODO: Translate] ~{hours} hr {minutes} min remaining"
}
} }
}, },
"onboarding": { "onboarding": {
+21
View File
@@ -50,6 +50,27 @@
"mb": "MB", "mb": "MB",
"gb": "GB", "gb": "GB",
"tb": "TB" "tb": "TB"
},
"scanProgress": {
"refreshing": "[TODO: Translate] Refreshing {type}s...",
"fullRebuilding": "[TODO: Translate] Full rebuild {type}s...",
"actionRefresh": "[TODO: Translate] Refresh",
"actionFullRebuild": "[TODO: Translate] Full rebuild",
"actionRefreshLower": "[TODO: Translate] refresh",
"actionRebuildLower": "[TODO: Translate] rebuild",
"stages": {
"scan_folders": "[TODO: Translate] Scanning folders...",
"count_models": "[TODO: Translate] Found {total} files",
"process_models": "[TODO: Translate] Processing models",
"reconcile_scan": "[TODO: Translate] Checking for changes...",
"process_new": "[TODO: Translate] Processing new models",
"finalizing": "[TODO: Translate] Finalizing..."
},
"eta": {
"lessThanMinute": "[TODO: Translate] Less than a minute remaining",
"minutes": "[TODO: Translate] ~{minutes} min remaining",
"hours": "[TODO: Translate] ~{hours} hr {minutes} min remaining"
}
} }
}, },
"onboarding": { "onboarding": {
+21
View File
@@ -50,6 +50,27 @@
"mb": "MB", "mb": "MB",
"gb": "GB", "gb": "GB",
"tb": "TB" "tb": "TB"
},
"scanProgress": {
"refreshing": "[TODO: Translate] Refreshing {type}s...",
"fullRebuilding": "[TODO: Translate] Full rebuild {type}s...",
"actionRefresh": "[TODO: Translate] Refresh",
"actionFullRebuild": "[TODO: Translate] Full rebuild",
"actionRefreshLower": "[TODO: Translate] refresh",
"actionRebuildLower": "[TODO: Translate] rebuild",
"stages": {
"scan_folders": "[TODO: Translate] Scanning folders...",
"count_models": "[TODO: Translate] Found {total} files",
"process_models": "[TODO: Translate] Processing models",
"reconcile_scan": "[TODO: Translate] Checking for changes...",
"process_new": "[TODO: Translate] Processing new models",
"finalizing": "[TODO: Translate] Finalizing..."
},
"eta": {
"lessThanMinute": "[TODO: Translate] Less than a minute remaining",
"minutes": "[TODO: Translate] ~{minutes} min remaining",
"hours": "[TODO: Translate] ~{hours} hr {minutes} min remaining"
}
} }
}, },
"onboarding": { "onboarding": {
+21
View File
@@ -50,6 +50,27 @@
"mb": "MB", "mb": "MB",
"gb": "GB", "gb": "GB",
"tb": "TB" "tb": "TB"
},
"scanProgress": {
"refreshing": "[TODO: Translate] Refreshing {type}s...",
"fullRebuilding": "[TODO: Translate] Full rebuild {type}s...",
"actionRefresh": "[TODO: Translate] Refresh",
"actionFullRebuild": "[TODO: Translate] Full rebuild",
"actionRefreshLower": "[TODO: Translate] refresh",
"actionRebuildLower": "[TODO: Translate] rebuild",
"stages": {
"scan_folders": "[TODO: Translate] Scanning folders...",
"count_models": "[TODO: Translate] Found {total} files",
"process_models": "[TODO: Translate] Processing models",
"reconcile_scan": "[TODO: Translate] Checking for changes...",
"process_new": "[TODO: Translate] Processing new models",
"finalizing": "[TODO: Translate] Finalizing..."
},
"eta": {
"lessThanMinute": "[TODO: Translate] Less than a minute remaining",
"minutes": "[TODO: Translate] ~{minutes} min remaining",
"hours": "[TODO: Translate] ~{hours} hr {minutes} min remaining"
}
} }
}, },
"onboarding": { "onboarding": {
+21
View File
@@ -50,6 +50,27 @@
"mb": "МБ", "mb": "МБ",
"gb": "ГБ", "gb": "ГБ",
"tb": "ТБ" "tb": "ТБ"
},
"scanProgress": {
"refreshing": "[TODO: Translate] Refreshing {type}s...",
"fullRebuilding": "[TODO: Translate] Full rebuild {type}s...",
"actionRefresh": "[TODO: Translate] Refresh",
"actionFullRebuild": "[TODO: Translate] Full rebuild",
"actionRefreshLower": "[TODO: Translate] refresh",
"actionRebuildLower": "[TODO: Translate] rebuild",
"stages": {
"scan_folders": "[TODO: Translate] Scanning folders...",
"count_models": "[TODO: Translate] Found {total} files",
"process_models": "[TODO: Translate] Processing models",
"reconcile_scan": "[TODO: Translate] Checking for changes...",
"process_new": "[TODO: Translate] Processing new models",
"finalizing": "[TODO: Translate] Finalizing..."
},
"eta": {
"lessThanMinute": "[TODO: Translate] Less than a minute remaining",
"minutes": "[TODO: Translate] ~{minutes} min remaining",
"hours": "[TODO: Translate] ~{hours} hr {minutes} min remaining"
}
} }
}, },
"onboarding": { "onboarding": {
+21
View File
@@ -50,6 +50,27 @@
"mb": "MB", "mb": "MB",
"gb": "GB", "gb": "GB",
"tb": "TB" "tb": "TB"
},
"scanProgress": {
"refreshing": "[TODO: Translate] Refreshing {type}s...",
"fullRebuilding": "[TODO: Translate] Full rebuild {type}s...",
"actionRefresh": "[TODO: Translate] Refresh",
"actionFullRebuild": "[TODO: Translate] Full rebuild",
"actionRefreshLower": "[TODO: Translate] refresh",
"actionRebuildLower": "[TODO: Translate] rebuild",
"stages": {
"scan_folders": "[TODO: Translate] Scanning folders...",
"count_models": "[TODO: Translate] Found {total} files",
"process_models": "[TODO: Translate] Processing models",
"reconcile_scan": "[TODO: Translate] Checking for changes...",
"process_new": "[TODO: Translate] Processing new models",
"finalizing": "[TODO: Translate] Finalizing..."
},
"eta": {
"lessThanMinute": "[TODO: Translate] Less than a minute remaining",
"minutes": "[TODO: Translate] ~{minutes} min remaining",
"hours": "[TODO: Translate] ~{hours} hr {minutes} min remaining"
}
} }
}, },
"onboarding": { "onboarding": {
+21
View File
@@ -50,6 +50,27 @@
"mb": "MB", "mb": "MB",
"gb": "GB", "gb": "GB",
"tb": "TB" "tb": "TB"
},
"scanProgress": {
"refreshing": "[TODO: Translate] Refreshing {type}s...",
"fullRebuilding": "[TODO: Translate] Full rebuild {type}s...",
"actionRefresh": "[TODO: Translate] Refresh",
"actionFullRebuild": "[TODO: Translate] Full rebuild",
"actionRefreshLower": "[TODO: Translate] refresh",
"actionRebuildLower": "[TODO: Translate] rebuild",
"stages": {
"scan_folders": "[TODO: Translate] Scanning folders...",
"count_models": "[TODO: Translate] Found {total} files",
"process_models": "[TODO: Translate] Processing models",
"reconcile_scan": "[TODO: Translate] Checking for changes...",
"process_new": "[TODO: Translate] Processing new models",
"finalizing": "[TODO: Translate] Finalizing..."
},
"eta": {
"lessThanMinute": "[TODO: Translate] Less than a minute remaining",
"minutes": "[TODO: Translate] ~{minutes} min remaining",
"hours": "[TODO: Translate] ~{hours} hr {minutes} min remaining"
}
} }
}, },
"onboarding": { "onboarding": {
+135 -14
View File
@@ -66,6 +66,14 @@ def _is_hidden_relative_path(rel_path: str) -> bool:
# requests (modal open + autocomplete) do not re-walk the model roots. # requests (modal open + autocomplete) do not re-walk the model roots.
ALL_FOLDERS_CACHE_TTL_SECONDS = 5.0 ALL_FOLDERS_CACHE_TTL_SECONDS = 5.0
# Maps a scanner model type to the manager page type used in progress
# broadcasts (e.g. 'lora' -> 'loras').
PAGE_TYPE_MAP = {
'lora': 'loras',
'checkpoint': 'checkpoints',
'embedding': 'embeddings',
}
def _is_pending_delete_path(path: str) -> bool: def _is_pending_delete_path(path: str) -> bool:
"""Return True when any path component is the pending-delete staging dir.""" """Return True when any path component is the pending-delete staging dir."""
@@ -149,6 +157,38 @@ class ModelScanner:
# Register this service # Register this service
asyncio.create_task(self._register_service()) asyncio.create_task(self._register_service())
@property
def page_type(self) -> str:
"""Manager page type used in progress broadcasts (e.g. 'loras')."""
return PAGE_TYPE_MAP.get(self.model_type, self.model_type)
async def _broadcast_scan_progress(
self,
status: str,
stage: str,
progress: int,
full_rebuild: bool,
**extra: Any,
) -> None:
"""Broadcast manual-refresh scan progress on the generic WS channel.
Best-effort only: broadcast failures must never affect the scan itself.
"""
payload: Dict[str, Any] = {
'type': 'scan_progress',
'status': status,
'model_type': self.model_type,
'pageType': self.page_type,
'stage': stage,
'full_rebuild': full_rebuild,
'progress': progress,
}
payload.update(extra)
try:
await ws_manager.broadcast(payload)
except Exception as exc: # pragma: no cover - defensive logging
logger.error(f"Error broadcasting scan progress for {self.model_type}: {exc}")
@property @property
def cache_version(self) -> int: def cache_version(self) -> int:
"""Monotonic version counter for the in-memory cache. """Monotonic version counter for the in-memory cache.
@@ -434,12 +474,7 @@ class ModelScanner:
self._is_initializing = True self._is_initializing = True
# Determine the page type based on model type # Determine the page type based on model type
page_type_map = { page_type = self.page_type
'lora': 'loras',
'checkpoint': 'checkpoints',
'embedding': 'embeddings'
}
page_type = page_type_map.get(self.model_type, self.model_type)
# First, try to load from cache # First, try to load from cache
await ws_manager.broadcast_init_progress({ await ws_manager.broadcast_init_progress({
@@ -804,7 +839,7 @@ class ModelScanner:
last_progress_time = time.time() last_progress_time = time.time()
last_progress_percent = 0 last_progress_percent = 0
async def progress_callback(processed_files: int, expected_total: int) -> None: async def progress_callback(processed_files: int, expected_total: int, current_name: str = '') -> None:
nonlocal last_progress_time, last_progress_percent nonlocal last_progress_time, last_progress_percent
if expected_total <= 0: if expected_total <= 0:
@@ -871,32 +906,84 @@ class ModelScanner:
async def _initialize_cache(self) -> None: async def _initialize_cache(self) -> None:
"""Initialize or refresh the cache""" """Initialize or refresh the cache"""
self._is_initializing = True # Set flag self._is_initializing = True # Set flag
last_progress_percent = 0
try: try:
start_time = time.time() start_time = time.time()
await self._broadcast_scan_progress('started', 'scan_folders', 0, True)
# Manually trigger a symlink rescan during a full rebuild. # Manually trigger a symlink rescan during a full rebuild.
# This ensures that any new symlink mappings are correctly picked up. # This ensures that any new symlink mappings are correctly picked up.
config.rebuild_symlink_cache() config.rebuild_symlink_cache()
# Determine the page type based on model type # Count files in a thread so the event loop stays responsive
loop = asyncio.get_running_loop()
total_files = await loop.run_in_executor(None, self._count_model_files)
await self._broadcast_scan_progress(
'processing', 'count_models', 1, True,
processed=0, total=total_files,
)
last_progress_time = time.time()
async def progress_callback(processed_files: int, expected_total: int, current_name: str = '') -> None:
nonlocal last_progress_time, last_progress_percent
if expected_total <= 0:
return
current_time = time.time()
progress_percent = min(99, int(1 + (processed_files / expected_total) * 98))
if progress_percent <= last_progress_percent:
return
if current_time - last_progress_time <= 0.5 and processed_files != expected_total:
return
last_progress_percent = progress_percent
last_progress_time = current_time
await self._broadcast_scan_progress(
'processing', 'process_models', progress_percent, True,
processed=processed_files, total=expected_total,
current_name=current_name,
)
# Scan for new data # Scan for new data
scan_result = await self._gather_model_data() scan_result = await self._gather_model_data(
total_files=total_files,
progress_callback=progress_callback,
)
if not self.is_cancelled(): if not self.is_cancelled():
await self._broadcast_scan_progress('finalizing', 'finalizing', 99, True)
await self._apply_scan_result(scan_result) await self._apply_scan_result(scan_result)
await self._save_persistent_cache(scan_result) await self._save_persistent_cache(scan_result)
await self._sync_download_history(scan_result.raw_data, source='scan') await self._sync_download_history(scan_result.raw_data, source='scan')
await self._broadcast_scan_progress(
'completed', 'finalizing', 100, True,
elapsed_seconds=time.time() - start_time,
)
logger.info( logger.info(
f"{self.model_type.capitalize()} Scanner: Cache initialization completed in {time.time() - start_time:.2f} seconds, " f"{self.model_type.capitalize()} Scanner: Cache initialization completed in {time.time() - start_time:.2f} seconds, "
f"found {len(scan_result.raw_data)} models" f"found {len(scan_result.raw_data)} models"
) )
else: else:
await self._broadcast_scan_progress(
'cancelled', 'process_models', last_progress_percent, True,
elapsed_seconds=time.time() - start_time,
)
logger.info( logger.info(
f"{self.model_type.capitalize()} Scanner: Cache initialization cancelled " f"{self.model_type.capitalize()} Scanner: Cache initialization cancelled "
f"after {time.time() - start_time:.2f} seconds" f"after {time.time() - start_time:.2f} seconds"
) )
except Exception as e: except Exception as e:
logger.error(f"{self.model_type.capitalize()} Scanner: Error initializing cache: {e}") logger.error(f"{self.model_type.capitalize()} Scanner: Error initializing cache: {e}")
await self._broadcast_scan_progress(
'error', 'process_models', last_progress_percent, True,
error=str(e),
)
# Ensure cache is at least an empty structure on error # Ensure cache is at least an empty structure on error
if self._cache is None: if self._cache is None:
self._cache = ModelCache( self._cache = ModelCache(
@@ -915,6 +1002,8 @@ class ModelScanner:
start_time = time.time() start_time = time.time()
logger.info(f"{self.model_type.capitalize()} Scanner: Starting fast cache reconciliation...") logger.info(f"{self.model_type.capitalize()} Scanner: Starting fast cache reconciliation...")
await self._broadcast_scan_progress('started', 'reconcile_scan', 0, False)
# Get current cached file paths # Get current cached file paths
cached_paths = {item['file_path'] for item in self._cache.raw_data} cached_paths = {item['file_path'] for item in self._cache.raw_data}
path_to_item = {item['file_path']: item for item in self._cache.raw_data} path_to_item = {item['file_path']: item for item in self._cache.raw_data}
@@ -987,6 +1076,10 @@ class ModelScanner:
await asyncio.sleep(0) await asyncio.sleep(0)
if self.is_cancelled(): if self.is_cancelled():
logger.info(f"{self.model_type.capitalize()} Scanner: Reconcile scan cancelled") logger.info(f"{self.model_type.capitalize()} Scanner: Reconcile scan cancelled")
await self._broadcast_scan_progress(
'cancelled', 'reconcile_scan', 0, False,
elapsed_seconds=time.time() - start_time,
)
return return
# Process new files in batches # Process new files in batches
@@ -994,10 +1087,14 @@ class ModelScanner:
if new_files: if new_files:
logger.info(f"{self.model_type.capitalize()} Scanner: Found {len(new_files)} new files to process") logger.info(f"{self.model_type.capitalize()} Scanner: Found {len(new_files)} new files to process")
batch_size = 50 batch_size = 50
for i in range(0, len(new_files), batch_size): total_new = len(new_files)
processed_new = 0
last_progress_time = time.time()
for i in range(0, total_new, batch_size):
batch = new_files[i:i+batch_size] batch = new_files[i:i+batch_size]
for path in batch: for path in batch:
logger.info(f"{self.model_type.capitalize()} Scanner: Processing {path}") logger.info(f"{self.model_type.capitalize()} Scanner: Processing {path}")
processed_new += 1
try: try:
# Find the appropriate root path for this file # Find the appropriate root path for this file
root_path = None root_path = None
@@ -1054,8 +1151,23 @@ class ModelScanner:
except Exception as e: except Exception as e:
logger.error(f"Error adding {path} to cache: {e}") logger.error(f"Error adding {path} to cache: {e}")
current_time = time.time()
if current_time - last_progress_time > 0.5 or processed_new == total_new:
last_progress_time = current_time
await self._broadcast_scan_progress(
'processing', 'process_new',
min(99, int(1 + (processed_new / total_new) * 98)), False,
processed=processed_new, total=total_new,
current_name=os.path.basename(path),
)
if self.is_cancelled(): if self.is_cancelled():
logger.info(f"{self.model_type.capitalize()} Scanner: Reconcile processing cancelled") logger.info(f"{self.model_type.capitalize()} Scanner: Reconcile processing cancelled")
await self._broadcast_scan_progress(
'cancelled', 'process_new',
min(99, int(1 + (processed_new / total_new) * 98)), False,
elapsed_seconds=time.time() - start_time,
)
return return
# Find missing files (in cache but not in filesystem) # Find missing files (in cache but not in filesystem)
@@ -1121,8 +1233,17 @@ class ModelScanner:
await self._persist_current_cache() await self._persist_current_cache()
logger.info(f"{self.model_type.capitalize()} Scanner: Cache reconciliation completed in {time.time() - start_time:.2f} seconds. Added {total_added}, removed {total_removed} models.") logger.info(f"{self.model_type.capitalize()} Scanner: Cache reconciliation completed in {time.time() - start_time:.2f} seconds. Added {total_added}, removed {total_removed} models.")
await self._broadcast_scan_progress(
'completed', 'process_new', 100, False,
added=total_added, removed=total_removed,
elapsed_seconds=time.time() - start_time,
)
except Exception as e: except Exception as e:
logger.error(f"{self.model_type.capitalize()} Scanner: Error reconciling cache: {e}", exc_info=True) logger.error(f"{self.model_type.capitalize()} Scanner: Error reconciling cache: {e}", exc_info=True)
await self._broadcast_scan_progress(
'error', 'reconcile_scan', 0, False,
error=str(e),
)
finally: finally:
self._is_initializing = False # Unset flag self._is_initializing = False # Unset flag
self.bump_cache_version() self.bump_cache_version()
@@ -1498,7 +1619,7 @@ class ModelScanner:
self, self,
*, *,
total_files: int = 0, total_files: int = 0,
progress_callback: Optional[Callable[[int, int], Awaitable[None]]] = None progress_callback: Optional[Callable[[int, int, str], Awaitable[None]]] = None
) -> CacheBuildResult: ) -> CacheBuildResult:
"""Collect metadata for all model files.""" """Collect metadata for all model files."""
@@ -1510,11 +1631,11 @@ class ModelScanner:
processed_real_files: Set[str] = set() processed_real_files: Set[str] = set()
visited_real_dirs: Set[str] = set() visited_real_dirs: Set[str] = set()
async def handle_progress() -> None: async def handle_progress(current_name: str = '') -> None:
if progress_callback is None: if progress_callback is None:
return return
try: try:
await progress_callback(processed_files, total_files) await progress_callback(processed_files, total_files, current_name)
except Exception as exc: # pragma: no cover - defensive logging except Exception as exc: # pragma: no cover - defensive logging
logger.error(f"Error reporting progress for {self.model_type}: {exc}") logger.error(f"Error reporting progress for {self.model_type}: {exc}")
@@ -1580,7 +1701,7 @@ class ModelScanner:
for tag in result.get('tags') or []: for tag in result.get('tags') or []:
tags_count[tag] = tags_count.get(tag, 0) + 1 tags_count[tag] = tags_count.get(tag, 0) + 1
await handle_progress() await handle_progress(entry.name)
await asyncio.sleep(0) await asyncio.sleep(0)
if self.is_cancelled(): if self.is_cancelled():
return return
+102 -4
View File
@@ -1753,7 +1753,36 @@ class RecipeScanner:
# Mark initialization as complete regardless of outcome # Mark initialization as complete regardless of outcome
self._is_initializing = False self._is_initializing = False
def _initialize_recipe_cache_sync(self): async def _broadcast_scan_progress(
self,
status: str,
stage: str,
progress: int,
full_rebuild: bool,
**extra: Any,
) -> None:
"""Broadcast manual-refresh scan progress on the generic WS channel.
Mirrors ``ModelScanner._broadcast_scan_progress`` so the recipes page
can reuse the same frontend contract. Best-effort only: broadcast
failures must never affect the scan itself.
"""
payload: Dict[str, Any] = {
'type': 'scan_progress',
'status': status,
'model_type': 'recipe',
'pageType': 'recipes',
'stage': stage,
'full_rebuild': full_rebuild,
'progress': progress,
}
payload.update(extra)
try:
await ws_manager.broadcast(payload)
except Exception as exc: # pragma: no cover - defensive logging
logger.error(f"Error broadcasting scan progress for recipe: {exc}")
def _initialize_recipe_cache_sync(self, report_progress: bool = False):
"""Synchronous version of recipe cache initialization for thread pool execution. """Synchronous version of recipe cache initialization for thread pool execution.
Uses persistent cache for fast startup when available: Uses persistent cache for fast startup when available:
@@ -1761,8 +1790,14 @@ class RecipeScanner:
2. Reconcile with filesystem (check mtime/size for changes) 2. Reconcile with filesystem (check mtime/size for changes)
3. Fall back to full directory scan if cache miss or reconciliation fails 3. Fall back to full directory scan if cache miss or reconciliation fails
4. Persist results for next startup 4. Persist results for next startup
Args:
report_progress: When True (manual force-refresh only), broadcast
scan_progress messages during the full directory scan. Startup
initialization leaves this False and behaves as before.
""" """
loop = None loop = None
scan_start_time: Optional[float] = None
try: try:
# Ensure cache exists to avoid None reference errors # Ensure cache exists to avoid None reference errors
if self._cache is None: if self._cache is None:
@@ -1844,7 +1879,17 @@ class RecipeScanner:
# Fall back to full directory scan # Fall back to full directory scan
logger.info("Recipe cache miss: performing full directory scan") logger.info("Recipe cache miss: performing full directory scan")
recipes, json_paths = self._full_directory_scan_sync(recipes_dir) if report_progress:
scan_start_time = time.time()
# Broadcast from the worker thread via its own event loop,
# mirroring ModelScanner._initialize_cache_sync.
loop.run_until_complete(
self._broadcast_scan_progress('started', 'scan_folders', 0, True)
)
recipes, json_paths = self._full_directory_scan_sync(
recipes_dir,
progress_loop=loop if report_progress else None,
)
self._json_path_map = json_paths self._json_path_map = json_paths
# Update cache with the collected data # Update cache with the collected data
@@ -1858,12 +1903,30 @@ class RecipeScanner:
recipes, json_paths, self._cache.image_id_map recipes, json_paths, self._cache.image_id_map
) )
if report_progress:
loop.run_until_complete(
self._broadcast_scan_progress(
'completed', 'finalizing', 100, True,
elapsed_seconds=time.time() - (scan_start_time or time.time()),
total=len(recipes),
)
)
return self._cache return self._cache
except Exception as e: except Exception as e:
logger.error(f"Error in thread-based recipe cache initialization: {e}") logger.error(f"Error in thread-based recipe cache initialization: {e}")
import traceback import traceback
traceback.print_exc(file=sys.stderr) traceback.print_exc(file=sys.stderr)
if report_progress and loop is not None:
try:
loop.run_until_complete(
self._broadcast_scan_progress(
'error', 'process_models', 0, True, error=str(e)
)
)
except Exception: # pragma: no cover - defensive logging
logger.error("Error broadcasting recipe scan failure", exc_info=True)
return self._cache if hasattr(self, "_cache") else None return self._cache if hasattr(self, "_cache") else None
finally: finally:
# Clean up the event loop # Clean up the event loop
@@ -2017,12 +2080,16 @@ class RecipeScanner:
return updated return updated
def _full_directory_scan_sync( def _full_directory_scan_sync(
self, recipes_dir: str self,
recipes_dir: str,
progress_loop: Optional[asyncio.AbstractEventLoop] = None,
) -> Tuple[List[Dict[str, Any]], Dict[str, str]]: ) -> Tuple[List[Dict[str, Any]], Dict[str, str]]:
"""Perform a full synchronous directory scan for recipes. """Perform a full synchronous directory scan for recipes.
Args: Args:
recipes_dir: Path to the recipes directory. recipes_dir: Path to the recipes directory.
progress_loop: When set (manual force-refresh only), broadcast
scan_progress messages through this thread-local event loop.
Returns: Returns:
Tuple of (recipes list, json_paths dict). Tuple of (recipes list, json_paths dict).
@@ -2037,6 +2104,17 @@ class RecipeScanner:
if file.lower().endswith(".recipe.json"): if file.lower().endswith(".recipe.json"):
recipe_files.append(os.path.join(root, file)) recipe_files.append(os.path.join(root, file))
total_files = len(recipe_files)
if progress_loop is not None:
progress_loop.run_until_complete(
self._broadcast_scan_progress(
'processing', 'count_models', 1, True,
processed=0, total=total_files,
)
)
last_progress_time = time.time()
# Process each recipe file # Process each recipe file
for i, recipe_path in enumerate(recipe_files): for i, recipe_path in enumerate(recipe_files):
recipe_data = self._load_recipe_file_sync(recipe_path) recipe_data = self._load_recipe_file_sync(recipe_path)
@@ -2044,6 +2122,23 @@ class RecipeScanner:
recipe_id = str(recipe_data.get("id", "")) recipe_id = str(recipe_data.get("id", ""))
recipes.append(recipe_data) recipes.append(recipe_data)
json_paths[recipe_id] = recipe_path json_paths[recipe_id] = recipe_path
if progress_loop is not None and total_files > 0:
processed = i + 1
current_time = time.time()
# Throttle to one update per 0.5s; always send the final one.
if (
processed == total_files
or current_time - last_progress_time > 0.5
):
last_progress_time = current_time
progress_percent = min(99, int(1 + (processed / total_files) * 98))
progress_loop.run_until_complete(
self._broadcast_scan_progress(
'processing', 'process_models', progress_percent, True,
processed=processed, total=total_files,
current_name=os.path.basename(recipe_path),
)
)
# Periodically release GIL so the event loop thread can run # Periodically release GIL so the event loop thread can run
if i % 100 == 0: if i % 100 == 0:
time.sleep(0) time.sleep(0)
@@ -2613,11 +2708,14 @@ class RecipeScanner:
start_time = time.time() start_time = time.time()
# Run the heavy lifting in a thread pool same path # Run the heavy lifting in a thread pool same path
# used by initialize_in_background(). # used by initialize_in_background(). Pass
# report_progress=True so manual refreshes broadcast
# scan_progress updates; startup init keeps it off.
loop = asyncio.get_event_loop() loop = asyncio.get_event_loop()
cache = await loop.run_in_executor( cache = await loop.run_in_executor(
None, None,
self._initialize_recipe_cache_sync, self._initialize_recipe_cache_sync,
True,
) )
if cache is not None: if cache is not None:
self._cache = cache self._cache = cache
+108 -7
View File
@@ -12,6 +12,11 @@ import {
} from './apiConfig.js'; } from './apiConfig.js';
import { resetAndReload } from './modelApiFactory.js'; import { resetAndReload } from './modelApiFactory.js';
import { sidebarManager } from '../components/SidebarManager.js'; import { sidebarManager } from '../components/SidebarManager.js';
// Shared scan ETA helpers live in a dependency-light module so pages that do
// not use BaseModelApiClient (e.g. recipes) can reuse them without pulling
// this module's import cycle (modelApiFactory -> loraApi -> baseModelApi).
import { createScanEtaTracker, formatScanRemainingTime } from '../utils/scanEtaUtils.js';
export { createScanEtaTracker, formatScanRemainingTime };
/** /**
* Abstract base class for all model API clients * Abstract base class for all model API clients
@@ -507,23 +512,67 @@ export class BaseModelApiClient {
async refreshModels(fullRebuild = false) { async refreshModels(fullRebuild = false) {
const abortController = new AbortController(); const abortController = new AbortController();
try { const displayName = this.apiConfig.config.displayName;
state.loadingManager.show( const singularName = this.apiConfig.config.singularName;
`${fullRebuild ? 'Full rebuild' : 'Refreshing'} ${this.apiConfig.config.displayName}s...`, const actionText = translate(
0 fullRebuild ? 'common.scanProgress.actionFullRebuild' : 'common.scanProgress.actionRefresh',
{},
fullRebuild ? 'Full rebuild' : 'Refresh'
); );
const actionLowerText = translate(
fullRebuild ? 'common.scanProgress.actionRebuildLower' : 'common.scanProgress.actionRefreshLower',
{},
fullRebuild ? 'rebuild' : 'refresh'
);
const initialMessage = translate(
fullRebuild ? 'common.scanProgress.fullRebuilding' : 'common.scanProgress.refreshing',
{ type: displayName },
`${fullRebuild ? 'Full rebuild' : 'Refreshing'} ${displayName}s...`
);
const etaTracker = createScanEtaTracker();
let ws = null;
const handleScanProgress = (data) => {
if (typeof data.progress === 'number') {
state.loadingManager.setProgress(data.progress);
}
let statusText = translate(
`common.scanProgress.stages.${data.stage}`,
{ total: data.total },
data.stage || ''
);
if (data.status === 'processing' && data.total > 0) {
statusText += ` (${data.processed}/${data.total})`;
if (data.current_name) {
statusText += ` ${data.current_name}`;
}
const etaText = etaTracker.update(data.processed, data.total);
if (etaText) {
statusText += ` | ${etaText}`;
}
}
state.loadingManager.setStatus(statusText);
};
try {
state.loadingManager.show(initialMessage, 0);
state.loadingManager.showCancelButton(() => { state.loadingManager.showCancelButton(() => {
this.cancelTask(); this.cancelTask();
abortController.abort(); abortController.abort();
}); });
// Connect to the shared progress channel for live scan updates.
// Failure to connect must not block the refresh itself — fall back
// to the plain loading indicator.
ws = await this._connectScanProgressSocket(handleScanProgress, singularName);
const url = new URL(this.apiConfig.endpoints.scan, window.location.origin); const url = new URL(this.apiConfig.endpoints.scan, window.location.origin);
url.searchParams.append('full_rebuild', fullRebuild); url.searchParams.append('full_rebuild', fullRebuild);
const response = await fetch(url, { signal: abortController.signal }); const response = await fetch(url, { signal: abortController.signal });
if (!response.ok) { if (!response.ok) {
throw new Error(`Failed to refresh ${this.apiConfig.config.displayName}s: ${response.status} ${response.statusText}`); throw new Error(`Failed to refresh ${displayName}s: ${response.status} ${response.statusText}`);
} }
const data = await response.json(); const data = await response.json();
@@ -534,20 +583,69 @@ export class BaseModelApiClient {
resetAndReload(true); resetAndReload(true);
showToast('toast.api.refreshComplete', { action: fullRebuild ? 'Full rebuild' : 'Refresh' }, 'success'); showToast('toast.api.refreshComplete', { action: actionText }, 'success');
} catch (error) { } catch (error) {
if (error.name === 'AbortError') { if (error.name === 'AbortError') {
showToast('toast.api.operationCancelled', {}, 'info'); showToast('toast.api.operationCancelled', {}, 'info');
return; return;
} }
console.error('Refresh failed:', error); console.error('Refresh failed:', error);
showToast('toast.api.refreshFailed', { action: fullRebuild ? 'rebuild' : 'refresh', type: this.apiConfig.config.displayName }, 'error'); showToast('toast.api.refreshFailed', { action: actionLowerText, type: displayName }, 'error');
} finally { } finally {
if (ws) {
ws.close();
}
state.loadingManager.hide(); state.loadingManager.hide();
state.loadingManager.restoreProgressBar(); state.loadingManager.restoreProgressBar();
} }
} }
/**
* Connect to the shared fetch-progress WebSocket for scan progress updates.
* Returns null when the connection cannot be established (silent fallback).
* @param {Function} onScanProgress - Handler for scan_progress messages
* @param {string} singularName - Model type filter (e.g. 'lora')
* @returns {Promise<WebSocket|null>}
*/
async _connectScanProgressSocket(onScanProgress, singularName) {
let socket = null;
try {
const wsProtocol = window.location.protocol === 'https:' ? 'wss://' : 'ws://';
socket = new WebSocket(`${wsProtocol}${window.location.host}${WS_ENDPOINTS.fetchProgress}`);
await new Promise((resolve, reject) => {
socket.onopen = resolve;
socket.onerror = reject;
});
socket.onmessage = (event) => {
let data;
try {
data = JSON.parse(event.data);
} catch (parseError) {
return;
}
// Only handle scan progress for this client's model type;
// other operations share this channel and must be ignored.
if (data.type !== 'scan_progress' || data.model_type !== singularName) {
return;
}
onScanProgress(data);
};
return socket;
} catch (error) {
if (socket) {
try {
socket.close();
} catch (closeError) {
// Ignore close errors during fallback
}
}
return null;
}
}
async refreshSingleModelMetadata(filePath) { async refreshSingleModelMetadata(filePath) {
try { try {
state.loadingManager.showSimpleLoading('Refreshing metadata...'); state.loadingManager.showSimpleLoading('Refreshing metadata...');
@@ -605,6 +703,9 @@ export class BaseModelApiClient {
ws.onmessage = (event) => { ws.onmessage = (event) => {
const data = JSON.parse(event.data); const data = JSON.parse(event.data);
// Scan progress shares this channel; it is handled by refreshModels
if (data.type === 'scan_progress') return;
switch (data.status) { switch (data.status) {
case 'started': case 'started':
loading.setStatus('Starting metadata fetch...'); loading.setStatus('Starting metadata fetch...');
+100 -5
View File
@@ -1,7 +1,12 @@
import { RecipeCard } from '../components/RecipeCard.js'; import { RecipeCard } from '../components/RecipeCard.js';
import { state, getCurrentPageState } from '../state/index.js'; import { state, getCurrentPageState } from '../state/index.js';
import { showToast } from '../utils/uiHelpers.js'; import { showToast } from '../utils/uiHelpers.js';
import { translate } from '../utils/i18nHelpers.js';
import { captureScrollPosition, restoreScrollPosition } from '../utils/infiniteScroll.js'; import { captureScrollPosition, restoreScrollPosition } from '../utils/infiniteScroll.js';
import { WS_ENDPOINTS } from './apiConfig.js';
// Import from the dependency-light utils module, not baseModelApi.js, to
// avoid the baseModelApi <-> modelApiFactory import cycle on this page.
import { createScanEtaTracker } from '../utils/scanEtaUtils.js';
const RECIPE_ENDPOINTS = { const RECIPE_ENDPOINTS = {
list: '/api/lm/recipes', list: '/api/lm/recipes',
@@ -333,11 +338,53 @@ export async function syncChanges() {
} }
export async function refreshRecipes(fullRebuild = true) { export async function refreshRecipes(fullRebuild = true) {
const actionLabel = fullRebuild ? 'Rebuilding recipe cache' : 'Refreshing recipes'; const actionText = translate(
const actionToast = fullRebuild ? 'Full rebuild' : 'Refresh'; fullRebuild ? 'common.scanProgress.actionFullRebuild' : 'common.scanProgress.actionRefresh',
{},
fullRebuild ? 'Full rebuild' : 'Refresh'
);
const actionLowerText = translate(
fullRebuild ? 'common.scanProgress.actionRebuildLower' : 'common.scanProgress.actionRefreshLower',
{},
fullRebuild ? 'rebuild' : 'refresh'
);
const initialMessage = translate(
fullRebuild ? 'common.scanProgress.fullRebuilding' : 'common.scanProgress.refreshing',
{ type: RECIPE_SIDEBAR_CONFIG.config.displayName },
`${fullRebuild ? 'Full rebuild' : 'Refreshing'} Recipes...`
);
const etaTracker = createScanEtaTracker();
let ws = null;
const handleScanProgress = (data) => {
if (typeof data.progress === 'number') {
state.loadingManager.setProgress(data.progress);
}
let statusText = translate(
`common.scanProgress.stages.${data.stage}`,
{ total: data.total },
data.stage || ''
);
if (data.status === 'processing' && data.total > 0) {
statusText += ` (${data.processed}/${data.total})`;
if (data.current_name) {
statusText += ` ${data.current_name}`;
}
const etaText = etaTracker.update(data.processed, data.total);
if (etaText) {
statusText += ` | ${etaText}`;
}
}
state.loadingManager.setStatus(statusText);
};
try { try {
state.loadingManager.show(`${actionLabel}...`, 0); state.loadingManager.show(initialMessage, 0);
// Connect to the shared progress channel for live scan updates.
// Failure to connect must not block the refresh itself — fall back
// to the plain loading indicator.
ws = await connectScanProgressSocket(handleScanProgress);
const url = new URL(RECIPE_ENDPOINTS.scan, window.location.origin); const url = new URL(RECIPE_ENDPOINTS.scan, window.location.origin);
url.searchParams.append('full_rebuild', fullRebuild); url.searchParams.append('full_rebuild', fullRebuild);
@@ -356,16 +403,64 @@ export async function refreshRecipes(fullRebuild = true) {
await resetAndReload(false); await resetAndReload(false);
showToast('toast.api.refreshComplete', { action: actionToast }, 'success'); showToast('toast.api.refreshComplete', { action: actionText }, 'success');
} catch (error) { } catch (error) {
console.error('Error refreshing recipes:', error); console.error('Error refreshing recipes:', error);
showToast('toast.api.refreshFailed', { action: fullRebuild ? 'rebuild' : 'refresh', type: 'recipe' }, 'error'); showToast('toast.api.refreshFailed', { action: actionLowerText, type: 'recipe' }, 'error');
} finally { } finally {
if (ws) {
ws.close();
}
state.loadingManager.hide(); state.loadingManager.hide();
state.loadingManager.restoreProgressBar(); state.loadingManager.restoreProgressBar();
} }
} }
/**
* Connect to the shared fetch-progress WebSocket for recipe scan progress.
* Returns null when the connection cannot be established (silent fallback).
* @param {Function} onScanProgress - Handler for scan_progress messages
* @returns {Promise<WebSocket|null>}
*/
async function connectScanProgressSocket(onScanProgress) {
let socket = null;
try {
const wsProtocol = window.location.protocol === 'https:' ? 'wss://' : 'ws://';
socket = new WebSocket(`${wsProtocol}${window.location.host}${WS_ENDPOINTS.fetchProgress}`);
await new Promise((resolve, reject) => {
socket.onopen = resolve;
socket.onerror = reject;
});
socket.onmessage = (event) => {
let data;
try {
data = JSON.parse(event.data);
} catch (parseError) {
return;
}
// Only handle recipe scan progress; other operations share this
// channel and must be ignored.
if (data.type !== 'scan_progress' || data.model_type !== 'recipe') {
return;
}
onScanProgress(data);
};
return socket;
} catch (error) {
if (socket) {
try {
socket.close();
} catch (closeError) {
// Ignore close errors during fallback
}
}
return null;
}
}
/** /**
* Load more recipes with pagination - updated to work with VirtualScroller * Load more recipes with pagination - updated to work with VirtualScroller
* @param {boolean} resetPage - Whether to reset to the first page * @param {boolean} resetPage - Whether to reset to the first page
+61
View File
@@ -0,0 +1,61 @@
import { translate } from './i18nHelpers.js';
/**
* Format a remaining-time estimate for scan progress display.
* @param {number} remainingMs - Estimated remaining time in milliseconds
* @returns {string} Localized ETA text
*/
export function formatScanRemainingTime(remainingMs) {
if (remainingMs < 60000) {
return translate('common.scanProgress.eta.lessThanMinute', {}, 'Less than a minute remaining');
}
if (remainingMs < 3600000) {
const minutes = Math.round(remainingMs / 60000);
return translate('common.scanProgress.eta.minutes', { minutes }, `~${minutes} min remaining`);
}
const hours = Math.floor(remainingMs / 3600000);
const minutes = Math.round((remainingMs % 3600000) / 60000);
return translate('common.scanProgress.eta.hours', { hours, minutes }, `~${hours} hr ${minutes} min remaining`);
}
/**
* Create an ETA tracker for scan progress. Uses an exponential moving
* average (0.7/0.3) over the observed per-file processing time, mirroring
* the estimator in components/initialization.js.
* @returns {{ update: (processed: number, total: number) => (string|null) }}
*/
export function createScanEtaTracker() {
let startTime = null;
let lastProcessed = 0;
let averageMsPerFile = null;
return {
/**
* Update with the latest counters.
* @returns {string|null} Localized ETA text, or null when not applicable
*/
update(processed, total) {
if (!total || total <= 0 || processed >= total) {
return null;
}
const now = Date.now();
if (startTime === null) {
// First sample only anchors the timer; not enough data yet
startTime = now;
lastProcessed = processed;
return translate('initialization.estimatingTime', {}, 'Estimating time...');
}
if (processed > lastProcessed) {
const msPerFile = (now - startTime) / processed;
averageMsPerFile = averageMsPerFile === null
? msPerFile
: averageMsPerFile * 0.7 + msPerFile * 0.3;
lastProcessed = processed;
}
if (averageMsPerFile === null) {
return translate('initialization.estimatingTime', {}, 'Estimating time...');
}
return formatScanRemainingTime((total - lastProcessed) * averageMsPerFile);
}
};
}
@@ -0,0 +1,348 @@
import { describe, it, beforeEach, afterEach, expect, vi } from 'vitest';
const {
BASE_MODEL_API_MODULE,
STATE_MODULE,
UI_HELPERS_MODULE,
I18N_MODULE,
STORAGE_MODULE,
API_CONFIG_MODULE,
API_FACTORY_MODULE,
SIDEBAR_MANAGER_MODULE,
} = vi.hoisted(() => ({
BASE_MODEL_API_MODULE: new URL('../../../static/js/api/baseModelApi.js', import.meta.url).pathname,
STATE_MODULE: new URL('../../../static/js/state/index.js', import.meta.url).pathname,
UI_HELPERS_MODULE: new URL('../../../static/js/utils/uiHelpers.js', import.meta.url).pathname,
I18N_MODULE: new URL('../../../static/js/utils/i18nHelpers.js', import.meta.url).pathname,
STORAGE_MODULE: new URL('../../../static/js/utils/storageHelpers.js', import.meta.url).pathname,
API_CONFIG_MODULE: new URL('../../../static/js/api/apiConfig.js', import.meta.url).pathname,
API_FACTORY_MODULE: new URL('../../../static/js/api/modelApiFactory.js', import.meta.url).pathname,
SIDEBAR_MANAGER_MODULE: new URL('../../../static/js/components/SidebarManager.js', import.meta.url).pathname,
}));
const showToastMock = vi.fn();
const showMock = vi.fn();
const showCancelButtonMock = vi.fn();
const hideMock = vi.fn();
const restoreProgressBarMock = vi.fn();
const setProgressMock = vi.fn();
const setStatusMock = vi.fn();
const resetAndReloadMock = vi.fn();
vi.mock(STATE_MODULE, () => ({
state: {
loadingManager: {
show: showMock,
showCancelButton: showCancelButtonMock,
hide: hideMock,
restoreProgressBar: restoreProgressBarMock,
setProgress: setProgressMock,
setStatus: setStatusMock,
},
},
getCurrentPageState: vi.fn(() => ({})),
}));
vi.mock(UI_HELPERS_MODULE, () => ({
showToast: showToastMock,
}));
vi.mock(I18N_MODULE, () => ({
translate: vi.fn((key, params, fallback) => {
if (fallback) {
return Object.entries(params || {}).reduce(
(text, [name, value]) => text.replaceAll(`{${name}}`, value),
fallback
);
}
return key;
}),
}));
vi.mock(STORAGE_MODULE, () => ({
getStorageItem: vi.fn(),
getSessionItem: vi.fn(),
removeSessionItem: vi.fn(),
saveMapToStorage: vi.fn(),
}));
vi.mock(API_CONFIG_MODULE, () => ({
getCompleteApiConfig: vi.fn(() => ({
endpoints: { scan: '/api/lm/loras/scan' },
config: { displayName: 'LoRA', singularName: 'lora' },
})),
getCurrentModelType: vi.fn(() => 'loras'),
isValidModelType: vi.fn(() => true),
DOWNLOAD_ENDPOINTS: {},
HF_ENDPOINTS: {},
WS_ENDPOINTS: { fetchProgress: '/ws/fetch-progress' },
}));
vi.mock(API_FACTORY_MODULE, () => ({
resetAndReload: resetAndReloadMock,
}));
vi.mock(SIDEBAR_MANAGER_MODULE, () => ({
sidebarManager: { refresh: vi.fn() },
}));
class FakeWebSocket {
static instances = [];
static failNextConnection = false;
constructor(url) {
this.url = url;
this.onopen = null;
this.onerror = null;
this.onmessage = null;
this.close = vi.fn();
FakeWebSocket.instances.push(this);
const shouldFail = FakeWebSocket.failNextConnection;
FakeWebSocket.failNextConnection = false;
queueMicrotask(() => {
if (shouldFail) {
this.onerror?.(new Error('connection refused'));
} else {
this.onopen?.();
}
});
}
emit(data) {
this.onmessage?.({ data: JSON.stringify(data) });
}
}
async function createClient() {
const { BaseModelApiClient } = await import(BASE_MODEL_API_MODULE);
class TestClient extends BaseModelApiClient {}
return new TestClient('loras');
}
async function flushMicrotasks() {
await new Promise((resolve) => setTimeout(resolve, 0));
}
describe('BaseModelApiClient.refreshModels scan progress', () => {
beforeEach(() => {
showToastMock.mockReset();
showMock.mockReset();
showCancelButtonMock.mockReset();
hideMock.mockReset();
restoreProgressBarMock.mockReset();
setProgressMock.mockReset();
setStatusMock.mockReset();
resetAndReloadMock.mockReset();
FakeWebSocket.instances = [];
FakeWebSocket.failNextConnection = false;
vi.stubGlobal('WebSocket', FakeWebSocket);
});
afterEach(() => {
delete global.fetch;
vi.unstubAllGlobals();
vi.restoreAllMocks();
});
function mockFetchPending() {
let resolveFetch;
global.fetch = vi.fn(() => new Promise((resolve) => { resolveFetch = resolve; }));
return {
resolveOk: (payload = { status: 'success' }) =>
resolveFetch({ ok: true, json: async () => payload }),
};
}
async function startRefresh(client, fullRebuild = false) {
const promise = client.refreshModels(fullRebuild);
await vi.waitFor(() => {
expect(FakeWebSocket.instances.length).toBe(1);
});
await flushMicrotasks();
const socket = FakeWebSocket.instances[0];
await vi.waitFor(() => {
expect(socket.onmessage).toBeTruthy();
});
return { promise, socket };
}
it('shows scan progress updates from the WebSocket channel', async () => {
const fetchControl = mockFetchPending();
const client = await createClient();
const { promise, socket } = await startRefresh(client);
expect(socket.url).toBe(`ws://${window.location.host}/ws/fetch-progress`);
socket.emit({
type: 'scan_progress',
status: 'started',
stage: 'scan_folders',
model_type: 'lora',
pageType: 'loras',
full_rebuild: false,
progress: 0,
});
socket.emit({
type: 'scan_progress',
status: 'processing',
stage: 'process_models',
model_type: 'lora',
pageType: 'loras',
full_rebuild: false,
progress: 50,
processed: 5,
total: 10,
current_name: 'style.safetensors',
});
expect(setProgressMock).toHaveBeenCalledWith(0);
expect(setProgressMock).toHaveBeenCalledWith(50);
const lastStatus = setStatusMock.mock.calls.at(-1)[0];
expect(lastStatus).toContain('(5/10)');
expect(lastStatus).toContain('style.safetensors');
// First ETA sample only anchors the timer
expect(lastStatus).toContain('Estimating time...');
fetchControl.resolveOk();
await promise;
expect(resetAndReloadMock).toHaveBeenCalledWith(true);
expect(showToastMock).toHaveBeenCalledWith(
'toast.api.refreshComplete',
{ action: 'Refresh' },
'success'
);
expect(socket.close).toHaveBeenCalled();
expect(hideMock).toHaveBeenCalled();
});
it('ignores messages for other types or other model types', async () => {
const fetchControl = mockFetchPending();
const client = await createClient();
const { promise, socket } = await startRefresh(client);
socket.emit({
type: 'scan_progress',
status: 'processing',
stage: 'process_models',
model_type: 'checkpoint',
progress: 33,
processed: 1,
total: 3,
});
socket.emit({
type: 'example_images_progress',
status: 'running',
model_type: 'lora',
progress: 66,
processed: 2,
total: 3,
});
expect(setProgressMock).not.toHaveBeenCalled();
expect(setStatusMock).not.toHaveBeenCalled();
fetchControl.resolveOk();
await promise;
});
it('falls back to plain loading when the WebSocket connection fails', async () => {
FakeWebSocket.failNextConnection = true;
global.fetch = vi.fn().mockResolvedValue({
ok: true,
json: async () => ({ status: 'success' }),
});
const client = await createClient();
await client.refreshModels(true);
expect(global.fetch).toHaveBeenCalled();
const [url] = global.fetch.mock.calls[0];
expect(url.searchParams.get('full_rebuild')).toBe('true');
expect(showMock).toHaveBeenCalledWith('Full rebuild LoRAs...', 0);
expect(showToastMock).toHaveBeenCalledWith(
'toast.api.refreshComplete',
{ action: 'Full rebuild' },
'success'
);
});
it('computes an ETA with EMA smoothing once enough samples arrive', async () => {
const fetchControl = mockFetchPending();
let now = 1000;
vi.spyOn(Date, 'now').mockImplementation(() => now);
const client = await createClient();
const { promise, socket } = await startRefresh(client);
const emitProcessing = (processed, total) => socket.emit({
type: 'scan_progress',
status: 'processing',
stage: 'process_models',
model_type: 'lora',
progress: Math.floor((processed / total) * 100),
processed,
total,
});
// First sample anchors the timer
emitProcessing(1, 10);
expect(setStatusMock.mock.calls.at(-1)[0]).toContain('Estimating time...');
// 100s elapsed for 2 files -> 50s per file -> 400s remaining -> ~7 min
now = 101000;
emitProcessing(2, 10);
expect(setStatusMock.mock.calls.at(-1)[0]).toContain('~7 min remaining');
// 110s elapsed for 4 files -> EMA = 50000*0.7 + 27500*0.3 = 43250ms/file
// remaining 6 files -> 259.5s -> ~4 min
now = 111000;
emitProcessing(4, 10);
expect(setStatusMock.mock.calls.at(-1)[0]).toContain('~4 min remaining');
fetchControl.resolveOk();
await promise;
});
it('shows the cancelled toast when the server reports cancellation', async () => {
const fetchControl = mockFetchPending();
const client = await createClient();
const { promise } = await startRefresh(client);
fetchControl.resolveOk({ status: 'cancelled' });
await promise;
expect(showToastMock).toHaveBeenCalledWith('toast.api.operationCancelled', {}, 'info');
expect(resetAndReloadMock).not.toHaveBeenCalled();
});
});
describe('createScanEtaTracker / formatScanRemainingTime', () => {
it('estimates remaining time from EMA of per-file cost', async () => {
const { createScanEtaTracker } = await import(BASE_MODEL_API_MODULE);
let now = 0;
vi.spyOn(Date, 'now').mockImplementation(() => now);
const tracker = createScanEtaTracker();
expect(tracker.update(1, 10)).toBe('Estimating time...');
now = 60000; // 60s for 3 files -> 20s/file -> 7 * 20s = 140s -> ~2 min
expect(tracker.update(3, 10)).toBe('~2 min remaining');
now = 61000; // tiny delta keeps EMA near 20s/file
expect(tracker.update(4, 10)).toBe('~2 min remaining');
// Done: no ETA
expect(tracker.update(10, 10)).toBeNull();
expect(tracker.update(0, 0)).toBeNull();
vi.restoreAllMocks();
});
it('formats hours and sub-minute remainders', async () => {
const { formatScanRemainingTime } = await import(BASE_MODEL_API_MODULE);
expect(formatScanRemainingTime(30000)).toBe('Less than a minute remaining');
expect(formatScanRemainingTime(5 * 60000)).toBe('~5 min remaining');
expect(formatScanRemainingTime(3600000 + 30 * 60000)).toBe('~1 hr 30 min remaining');
});
});
@@ -0,0 +1,285 @@
import { describe, it, beforeEach, afterEach, expect, vi } from 'vitest';
const showToastMock = vi.hoisted(() => vi.fn());
const loadingManagerMock = vi.hoisted(() => ({
show: vi.fn(),
hide: vi.fn(),
restoreProgressBar: vi.fn(),
setProgress: vi.fn(),
setStatus: vi.fn(),
}));
const virtualScrollerMock = vi.hoisted(() => ({
refreshWithData: vi.fn(),
}));
const getCurrentPageStateMock = vi.hoisted(() => vi.fn());
const etaUpdateMock = vi.hoisted(() => vi.fn(() => 'ETA soon'));
vi.mock('../../../static/js/components/RecipeCard.js', () => ({
RecipeCard: vi.fn(() => ({ element: document.createElement('div') })),
}));
vi.mock('../../../static/js/state/index.js', () => ({
state: {
loadingManager: loadingManagerMock,
virtualScroller: virtualScrollerMock,
},
getCurrentPageState: getCurrentPageStateMock,
}));
vi.mock('../../../static/js/utils/uiHelpers.js', () => ({
showToast: showToastMock,
}));
vi.mock('../../../static/js/utils/i18nHelpers.js', () => ({
translate: vi.fn((key, params, fallback) => {
if (fallback) {
return Object.entries(params || {}).reduce(
(text, [name, value]) => text.replaceAll(`{${name}}`, value),
fallback
);
}
return key;
}),
}));
vi.mock('../../../static/js/utils/infiniteScroll.js', () => ({
captureScrollPosition: vi.fn(),
restoreScrollPosition: vi.fn(),
}));
vi.mock('../../../static/js/api/apiConfig.js', () => ({
WS_ENDPOINTS: { fetchProgress: '/ws/fetch-progress' },
}));
vi.mock('../../../static/js/utils/scanEtaUtils.js', () => ({
createScanEtaTracker: () => ({ update: etaUpdateMock }),
}));
import { refreshRecipes } from '../../../static/js/api/recipeApi.js';
class FakeWebSocket {
static instances = [];
static failNextConnection = false;
constructor(url) {
this.url = url;
this.onopen = null;
this.onerror = null;
this.onmessage = null;
this.close = vi.fn();
FakeWebSocket.instances.push(this);
const shouldFail = FakeWebSocket.failNextConnection;
FakeWebSocket.failNextConnection = false;
queueMicrotask(() => {
if (shouldFail) {
this.onerror?.(new Error('connection refused'));
} else {
this.onopen?.();
}
});
}
emit(data) {
this.onmessage?.({ data: JSON.stringify(data) });
}
}
async function flushMicrotasks() {
await new Promise((resolve) => setTimeout(resolve, 0));
}
describe('refreshRecipes scan progress', () => {
beforeEach(() => {
vi.clearAllMocks();
getCurrentPageStateMock.mockReturnValue({
pageSize: 50,
currentPage: 1,
hasMore: true,
isLoading: false,
sortBy: 'date:desc',
showFavoritesOnly: false,
activeFolder: null,
searchOptions: { recursive: true },
customFilter: { active: false },
filters: {},
});
FakeWebSocket.instances = [];
FakeWebSocket.failNextConnection = false;
vi.stubGlobal('WebSocket', FakeWebSocket);
});
afterEach(() => {
delete global.fetch;
vi.unstubAllGlobals();
});
function mockFetchPendingScan() {
let resolveScan;
global.fetch = vi.fn((input) => {
const url = String(input);
if (url.includes('/scan')) {
return new Promise((resolve) => { resolveScan = resolve; });
}
// Recipe list reload after the scan completes
return Promise.resolve({
ok: true,
json: async () => ({ items: [], total: 0, total_pages: 0 }),
});
});
return {
resolveOk: (payload = { status: 'success' }) =>
resolveScan({ ok: true, json: async () => payload }),
resolveNotOk: () =>
resolveScan({ ok: false, status: 500, statusText: 'Server Error' }),
};
}
async function startRefresh(fullRebuild = true) {
const promise = refreshRecipes(fullRebuild);
await vi.waitFor(() => {
expect(FakeWebSocket.instances.length).toBe(1);
});
await flushMicrotasks();
const socket = FakeWebSocket.instances[0];
await vi.waitFor(() => {
expect(socket.onmessage).toBeTruthy();
});
return { promise, socket };
}
it('shows scan progress updates from the WebSocket channel', async () => {
const fetchControl = mockFetchPendingScan();
const { promise, socket } = await startRefresh();
expect(socket.url).toBe(`ws://${window.location.host}/ws/fetch-progress`);
socket.emit({
type: 'scan_progress',
status: 'started',
stage: 'scan_folders',
model_type: 'recipe',
pageType: 'recipes',
full_rebuild: true,
progress: 0,
});
socket.emit({
type: 'scan_progress',
status: 'processing',
stage: 'process_models',
model_type: 'recipe',
pageType: 'recipes',
full_rebuild: true,
progress: 50,
processed: 5,
total: 10,
current_name: 'style.recipe.json',
});
expect(loadingManagerMock.setProgress).toHaveBeenCalledWith(0);
expect(loadingManagerMock.setProgress).toHaveBeenCalledWith(50);
const lastStatus = loadingManagerMock.setStatus.mock.calls.at(-1)[0];
expect(lastStatus).toContain('(5/10)');
expect(lastStatus).toContain('style.recipe.json');
expect(lastStatus).toContain('ETA soon');
expect(etaUpdateMock).toHaveBeenCalledWith(5, 10);
fetchControl.resolveOk();
await promise;
expect(showToastMock).toHaveBeenCalledWith(
'toast.api.refreshComplete',
{ action: 'Full rebuild' },
'success'
);
expect(socket.close).toHaveBeenCalled();
expect(loadingManagerMock.hide).toHaveBeenCalled();
});
it('ignores messages for other types or other model types', async () => {
const fetchControl = mockFetchPendingScan();
const { promise, socket } = await startRefresh();
socket.emit({
type: 'scan_progress',
status: 'processing',
stage: 'process_models',
model_type: 'lora',
progress: 33,
processed: 1,
total: 3,
});
socket.emit({
type: 'example_images_progress',
status: 'running',
model_type: 'recipe',
progress: 66,
processed: 2,
total: 3,
});
expect(loadingManagerMock.setProgress).not.toHaveBeenCalled();
expect(loadingManagerMock.setStatus).not.toHaveBeenCalled();
fetchControl.resolveOk();
await promise;
});
it('falls back to plain loading when the WebSocket connection fails', async () => {
FakeWebSocket.failNextConnection = true;
global.fetch = vi.fn((input) => {
const url = String(input);
if (url.includes('/scan')) {
return Promise.resolve({
ok: true,
json: async () => ({ status: 'success' }),
});
}
return Promise.resolve({
ok: true,
json: async () => ({ items: [], total: 0, total_pages: 0 }),
});
});
await refreshRecipes(false);
expect(global.fetch).toHaveBeenCalled();
const [url] = global.fetch.mock.calls[0];
expect(url.searchParams.get('full_rebuild')).toBe('false');
expect(loadingManagerMock.show).toHaveBeenCalledWith('Refreshing Recipes...', 0);
expect(showToastMock).toHaveBeenCalledWith(
'toast.api.refreshComplete',
{ action: 'Refresh' },
'success'
);
});
it('shows the cancelled toast when the server reports cancellation', async () => {
const fetchControl = mockFetchPendingScan();
const { promise } = await startRefresh();
fetchControl.resolveOk({ status: 'cancelled' });
await promise;
expect(showToastMock).toHaveBeenCalledWith('toast.api.operationCancelled', {}, 'info');
expect(showToastMock).not.toHaveBeenCalledWith(
'toast.api.refreshComplete',
expect.anything(),
expect.anything()
);
});
it('reports refresh failures through the error toast', async () => {
const fetchControl = mockFetchPendingScan();
const { promise } = await startRefresh();
fetchControl.resolveNotOk();
await promise;
expect(showToastMock).toHaveBeenCalledWith(
'toast.api.refreshFailed',
{ action: 'rebuild', type: 'recipe' },
'error'
);
expect(loadingManagerMock.hide).toHaveBeenCalled();
});
});
+186
View File
@@ -30,10 +30,14 @@ from py.utils.models import BaseModelMetadata
class RecordingWebSocketManager: class RecordingWebSocketManager:
def __init__(self) -> None: def __init__(self) -> None:
self.payloads: List[Dict[str, Any]] = [] self.payloads: List[Dict[str, Any]] = []
self.broadcasts: List[Dict[str, Any]] = []
async def broadcast_init_progress(self, payload: Dict[str, Any]) -> None: async def broadcast_init_progress(self, payload: Dict[str, Any]) -> None:
self.payloads.append(payload) self.payloads.append(payload)
async def broadcast(self, payload: Dict[str, Any]) -> None:
self.broadcasts.append(payload)
def _normalize_path(path: Path) -> str: def _normalize_path(path: Path) -> str:
return str(path).replace(os.sep, "/") return str(path).replace(os.sep, "/")
@@ -1395,3 +1399,185 @@ async def test_get_all_folders_invalidated_after_move(tmp_path: Path):
assert "new" in all_folders assert "new" in all_folders
assert "new/deep" in all_folders assert "new/deep" in all_folders
assert set(cache.folders) <= set(all_folders) assert set(cache.folders) <= set(all_folders)
@pytest.mark.asyncio
async def test_initialize_cache_broadcasts_scan_progress(tmp_path: Path, monkeypatch):
_create_files(tmp_path)
scanner = DummyScanner(tmp_path)
ws_stub = RecordingWebSocketManager()
monkeypatch.setattr(model_scanner, "ws_manager", ws_stub)
await scanner._initialize_cache()
messages = ws_stub.broadcasts
assert messages, "expected scan_progress broadcasts"
started = messages[0]
assert started["type"] == "scan_progress"
assert started["status"] == "started"
assert started["stage"] == "scan_folders"
assert started["progress"] == 0
assert started["model_type"] == "dummy"
assert started["pageType"] == "dummy"
assert started["full_rebuild"] is True
count_messages = [m for m in messages if m["stage"] == "count_models"]
assert count_messages and count_messages[0]["total"] == 3
process_messages = [
m for m in messages
if m["stage"] == "process_models" and m["status"] == "processing"
]
assert process_messages, "expected at least one process_models update"
final_process = process_messages[-1]
assert final_process["processed"] == 3
assert final_process["total"] == 3
assert final_process["current_name"].endswith(".txt")
for message in process_messages:
assert 0 < message["progress"] <= 99
stages = [m["stage"] for m in messages]
assert "finalizing" in stages
completed = messages[-1]
assert completed["status"] == "completed"
assert completed["progress"] == 100
assert completed["elapsed_seconds"] >= 0
@pytest.mark.asyncio
async def test_initialize_cache_broadcasts_cancelled(tmp_path: Path, monkeypatch):
_create_files(tmp_path)
scanner = DummyScanner(tmp_path)
ws_stub = RecordingWebSocketManager()
monkeypatch.setattr(model_scanner, "ws_manager", ws_stub)
original_process = DummyScanner._process_model_file
async def cancelling_process(self, file_path, root_path, **kwargs):
scanner.cancel_task()
return await original_process(self, file_path, root_path, **kwargs)
monkeypatch.setattr(DummyScanner, "_process_model_file", cancelling_process)
await scanner._initialize_cache()
messages = ws_stub.broadcasts
assert messages[0]["status"] == "started"
assert messages[-1]["status"] == "cancelled"
assert messages[-1]["elapsed_seconds"] >= 0
assert not any(m["status"] == "completed" for m in messages)
@pytest.mark.asyncio
async def test_initialize_cache_broadcasts_error(tmp_path: Path, monkeypatch):
scanner = DummyScanner(tmp_path)
ws_stub = RecordingWebSocketManager()
monkeypatch.setattr(model_scanner, "ws_manager", ws_stub)
async def raising_gather(**_kwargs):
raise RuntimeError("boom")
monkeypatch.setattr(scanner, "_gather_model_data", raising_gather)
await scanner._initialize_cache()
messages = ws_stub.broadcasts
assert messages[0]["status"] == "started"
assert messages[-1]["status"] == "error"
assert messages[-1]["error"] == "boom"
@pytest.mark.asyncio
async def test_reconcile_cache_broadcasts_scan_progress(tmp_path: Path, monkeypatch):
_create_files(tmp_path)
scanner = DummyScanner(tmp_path)
await scanner._initialize_cache()
ws_stub = RecordingWebSocketManager()
monkeypatch.setattr(model_scanner, "ws_manager", ws_stub)
new_file = tmp_path / "three.txt"
new_file.write_text("three", encoding="utf-8")
await scanner._reconcile_cache()
messages = ws_stub.broadcasts
assert messages, "expected scan_progress broadcasts"
started = messages[0]
assert started["type"] == "scan_progress"
assert started["status"] == "started"
assert started["stage"] == "reconcile_scan"
assert started["progress"] == 0
assert started["full_rebuild"] is False
process_messages = [
m for m in messages
if m["stage"] == "process_new" and m["status"] == "processing"
]
assert process_messages, "expected process_new progress updates"
assert process_messages[-1]["processed"] == 1
assert process_messages[-1]["total"] == 1
assert process_messages[-1]["current_name"] == "three.txt"
completed = messages[-1]
assert completed["status"] == "completed"
assert completed["progress"] == 100
assert completed["added"] == 1
assert completed["removed"] == 0
assert completed["elapsed_seconds"] >= 0
@pytest.mark.asyncio
async def test_reconcile_cache_broadcasts_cancelled(tmp_path: Path, monkeypatch):
_create_files(tmp_path)
scanner = DummyScanner(tmp_path)
await scanner._initialize_cache()
ws_stub = RecordingWebSocketManager()
monkeypatch.setattr(model_scanner, "ws_manager", ws_stub)
new_file = tmp_path / "four.txt"
new_file.write_text("four", encoding="utf-8")
original_process = DummyScanner._process_model_file
async def cancelling_process(self, file_path, root_path, **kwargs):
scanner.cancel_task()
return await original_process(self, file_path, root_path, **kwargs)
monkeypatch.setattr(DummyScanner, "_process_model_file", cancelling_process)
await scanner._reconcile_cache()
messages = ws_stub.broadcasts
assert messages[0]["status"] == "started"
assert messages[-1]["status"] == "cancelled"
assert messages[-1]["elapsed_seconds"] >= 0
assert not any(m["status"] == "completed" for m in messages)
@pytest.mark.asyncio
async def test_reconcile_cache_broadcasts_error(tmp_path: Path, monkeypatch):
_create_files(tmp_path)
scanner = DummyScanner(tmp_path)
await scanner._initialize_cache()
ws_stub = RecordingWebSocketManager()
monkeypatch.setattr(model_scanner, "ws_manager", ws_stub)
def raising_walk(*_args, **_kwargs):
raise RuntimeError("walk failed")
monkeypatch.setattr(model_scanner.os, "walk", raising_walk)
await scanner._reconcile_cache()
messages = ws_stub.broadcasts
assert messages[0]["status"] == "started"
assert messages[-1]["status"] == "error"
assert messages[-1]["error"] == "walk failed"
+131
View File
@@ -9,6 +9,7 @@ import pytest
from py.config import config from py.config import config
from py.services import model_scanner as model_scanner_module from py.services import model_scanner as model_scanner_module
from py.services import recipe_scanner as recipe_scanner_module
from py.services.model_cache import ModelCache from py.services.model_cache import ModelCache
from py.services.model_hash_index import ModelHashIndex from py.services.model_hash_index import ModelHashIndex
from py.services.model_scanner import CacheBuildResult, ModelScanner from py.services.model_scanner import CacheBuildResult, ModelScanner
@@ -4965,3 +4966,133 @@ async def test_find_all_duplicate_recipes_include_prompt_missing_gen_params(reci
groups = await scanner.find_all_duplicate_recipes(include_prompt=True) groups = await scanner.find_all_duplicate_recipes(include_prompt=True)
# Recipes without gen_params/prompt normalize to empty prompt and match # Recipes without gen_params/prompt normalize to empty prompt and match
assert groups == {"abc:0.8\x1f": ["r1", "r2"]} assert groups == {"abc:0.8\x1f": ["r1", "r2"]}
class RecordingRecipeWebSocketManager:
"""Minimal ws_manager stand-in that records broadcasts."""
def __init__(self) -> None:
self.payloads: list[Dict[str, Any]] = []
self.broadcasts: list[Dict[str, Any]] = []
async def broadcast_init_progress(self, payload: Dict[str, Any]) -> None:
self.payloads.append(payload)
async def broadcast(self, payload: Dict[str, Any]) -> None:
self.broadcasts.append(payload)
def _write_progress_recipe_files(recipes_dir: Path, count: int) -> None:
recipes_dir.mkdir(parents=True, exist_ok=True)
for idx in range(count):
recipe_path = recipes_dir / f"progress-recipe-{idx}.recipe.json"
recipe_path.write_text(
json.dumps(
{
"id": f"progress-recipe-{idx}",
"file_path": str(recipes_dir / f"img-{idx}.png"),
"title": f"Recipe {idx}",
"modified": 0.0,
"created_date": 0.0,
"loras": [],
}
),
encoding="utf-8",
)
@pytest.mark.asyncio
async def test_force_refresh_broadcasts_scan_progress(
tmp_path: Path, monkeypatch, recipe_scanner
):
scanner, _stub = recipe_scanner
recipes_dir = Path(config.loras_roots[0]) / "recipes"
_write_progress_recipe_files(recipes_dir, 3)
ws_stub = RecordingRecipeWebSocketManager()
monkeypatch.setattr(recipe_scanner_module, "ws_manager", ws_stub)
await scanner.get_cached_data(force_refresh=True)
# Wait for the FTS index build so no background task outlives the loop.
if scanner._fts_index_task:
await scanner._fts_index_task
messages = ws_stub.broadcasts
assert messages, "expected scan_progress broadcasts"
started = messages[0]
assert started["type"] == "scan_progress"
assert started["status"] == "started"
assert started["stage"] == "scan_folders"
assert started["progress"] == 0
assert started["model_type"] == "recipe"
assert started["pageType"] == "recipes"
assert started["full_rebuild"] is True
count_messages = [m for m in messages if m["stage"] == "count_models"]
assert count_messages and count_messages[0]["total"] == 3
process_messages = [
m
for m in messages
if m["stage"] == "process_models" and m["status"] == "processing"
]
assert process_messages, "expected at least one process_models update"
final_process = process_messages[-1]
assert final_process["processed"] == 3
assert final_process["total"] == 3
assert final_process["current_name"].endswith(".recipe.json")
for message in process_messages:
assert 0 < message["progress"] <= 99
completed = messages[-1]
assert completed["status"] == "completed"
assert completed["progress"] == 100
assert completed["elapsed_seconds"] >= 0
assert completed["total"] == 3
def test_sync_init_without_report_progress_does_not_broadcast(
tmp_path: Path, monkeypatch, recipe_scanner
):
"""Startup path (initialize_in_background) must not emit scan_progress."""
scanner, _stub = recipe_scanner
recipes_dir = Path(config.loras_roots[0]) / "recipes"
_write_progress_recipe_files(recipes_dir, 2)
ws_stub = RecordingRecipeWebSocketManager()
monkeypatch.setattr(recipe_scanner_module, "ws_manager", ws_stub)
# Invalidate the persistent cache so the sync path performs a full
# directory scan, exactly like a force refresh but without progress
# reporting (this is how initialize_in_background invokes it).
scanner._persistent_cache.save_cache([], {})
scanner._initialize_recipe_cache_sync()
assert ws_stub.broadcasts == []
def test_sync_init_reports_error_broadcast(
tmp_path: Path, monkeypatch, recipe_scanner
):
scanner, _stub = recipe_scanner
recipes_dir = Path(config.loras_roots[0]) / "recipes"
_write_progress_recipe_files(recipes_dir, 1)
ws_stub = RecordingRecipeWebSocketManager()
monkeypatch.setattr(recipe_scanner_module, "ws_manager", ws_stub)
scanner._persistent_cache.save_cache([], {})
def raising_scan(self, recipes_dir, progress_loop=None):
raise RuntimeError("boom")
monkeypatch.setattr(RecipeScanner, "_full_directory_scan_sync", raising_scan)
scanner._initialize_recipe_cache_sync(report_progress=True)
messages = ws_stub.broadcasts
assert messages[0]["status"] == "started"
assert messages[-1]["status"] == "error"
assert messages[-1]["error"] == "boom"