Enhance ProgressManager for thread safety and event loop integration

- Added thread-safe mechanisms to the ProgressManager for handling model download progress updates.
- Introduced a main event loop setter to ensure safe operations from background threads.
- Improved listener notification to handle updates in a thread-safe manner.
- Updated methods to ensure thread safety when accessing progress data.
This commit is contained in:
Jamie Pine
2026-01-29 16:23:46 -08:00
parent fadb57164e
commit 462f104494
2 changed files with 122 additions and 47 deletions
+12 -4
View File
@@ -122,9 +122,9 @@ async def health():
model_downloaded = True model_downloaded = True
break break
except (ImportError, Exception): except (ImportError, Exception):
# Method 2: Check cache directory (using HuggingFace's OS-specific cache location) # Method 2: Check cache directory (using HuggingFace's OS-specific cache location)
cache_dir = hf_constants.HF_HUB_CACHE cache_dir = hf_constants.HF_HUB_CACHE
repo_cache = Path(cache_dir) / "models--" + default_model_id.replace("/", "--") repo_cache = Path(cache_dir) / ("models--" + default_model_id.replace("/", "--"))
if repo_cache.exists(): if repo_cache.exists():
has_model_files = ( has_model_files = (
any(repo_cache.rglob("*.bin")) or any(repo_cache.rglob("*.bin")) or
@@ -1197,7 +1197,7 @@ async def get_model_status():
if not downloaded: if not downloaded:
try: try:
cache_dir = hf_constants.HF_HUB_CACHE cache_dir = hf_constants.HF_HUB_CACHE
repo_cache = Path(cache_dir) / "models--" + config["hf_repo_id"].replace("/", "--") repo_cache = Path(cache_dir) / ("models--" + config["hf_repo_id"].replace("/", "--"))
if repo_cache.exists(): if repo_cache.exists():
# Check for model files (bin, safetensors, or other common model files) # Check for model files (bin, safetensors, or other common model files)
@@ -1493,6 +1493,14 @@ async def startup_event():
print(f"Database initialized at {database._db_path}") print(f"Database initialized at {database._db_path}")
print(f"GPU available: {_get_gpu_status()}") print(f"GPU available: {_get_gpu_status()}")
# Initialize progress manager with main event loop for thread-safe operations
try:
progress_manager = get_progress_manager()
progress_manager._set_main_loop(asyncio.get_running_loop())
print("Progress manager initialized with event loop")
except Exception as e:
print(f"Warning: Could not initialize progress manager event loop: {e}")
# Ensure HuggingFace cache directory exists # Ensure HuggingFace cache directory exists
try: try:
from huggingface_hub import constants as hf_constants from huggingface_hub import constants as hf_constants
+110 -43
View File
@@ -6,16 +6,55 @@ from typing import Optional, Callable, Dict, List
from fastapi.responses import StreamingResponse from fastapi.responses import StreamingResponse
import asyncio import asyncio
import json import json
import threading
from datetime import datetime from datetime import datetime
class ProgressManager: class ProgressManager:
"""Manages download progress for multiple models.""" """Manages download progress for multiple models.
Thread-safe: can be called from background threads (e.g., via asyncio.to_thread).
"""
def __init__(self): def __init__(self):
self._progress: Dict[str, Dict] = {} self._progress: Dict[str, Dict] = {}
self._listeners: Dict[str, list] = {} self._listeners: Dict[str, list] = {}
self._lock = threading.Lock() # Thread-safe lock for progress dict
self._main_loop: Optional[asyncio.AbstractEventLoop] = None
def _set_main_loop(self, loop: asyncio.AbstractEventLoop):
"""Set the main event loop for thread-safe operations."""
self._main_loop = loop
def _notify_listeners_threadsafe(self, model_name: str, progress_data: Dict):
"""Notify listeners in a thread-safe manner."""
import logging
logger = logging.getLogger(__name__)
if model_name not in self._listeners:
return
for queue in self._listeners[model_name]:
try:
# Check if we're in the main event loop thread
try:
running_loop = asyncio.get_running_loop()
# We're in an async context, can use put_nowait directly
queue.put_nowait(progress_data.copy())
except RuntimeError:
# Not in async context (running in background thread)
# Use call_soon_threadsafe to safely put on queue
if self._main_loop and self._main_loop.is_running():
self._main_loop.call_soon_threadsafe(
lambda q=queue, d=progress_data.copy(): q.put_nowait(d) if not q.full() else None
)
else:
logger.debug(f"No main loop available for {model_name}, skipping notification")
except asyncio.QueueFull:
logger.warning(f"Queue full for {model_name}, dropping update")
except Exception as e:
logger.warning(f"Error notifying listener for {model_name}: {e}")
def update_progress( def update_progress(
self, self,
model_name: str, model_name: str,
@@ -26,6 +65,8 @@ class ProgressManager:
): ):
""" """
Update progress for a model download. Update progress for a model download.
Thread-safe: can be called from background threads.
Args: Args:
model_name: Name of the model (e.g., "qwen-tts-1.7B", "whisper-base") model_name: Name of the model (e.g., "qwen-tts-1.7B", "whisper-base")
@@ -39,7 +80,7 @@ class ProgressManager:
progress_pct = (current / total * 100) if total > 0 else 0 progress_pct = (current / total * 100) if total > 0 else 0
self._progress[model_name] = { progress_data = {
"model_name": model_name, "model_name": model_name,
"current": current, "current": current,
"total": total, "total": total,
@@ -48,30 +89,33 @@ class ProgressManager:
"status": status, "status": status,
"timestamp": datetime.now().isoformat(), "timestamp": datetime.now().isoformat(),
} }
# Thread-safe update of progress dict
with self._lock:
self._progress[model_name] = progress_data
# Notify all listeners # Notify all listeners (thread-safe)
listener_count = len(self._listeners.get(model_name, [])) listener_count = len(self._listeners.get(model_name, []))
if listener_count > 0: if listener_count > 0:
logger.debug(f"Notifying {listener_count} listeners for {model_name}: {progress_pct:.1f}% ({filename})") logger.debug(f"Notifying {listener_count} listeners for {model_name}: {progress_pct:.1f}% ({filename})")
for queue in self._listeners[model_name]: self._notify_listeners_threadsafe(model_name, progress_data)
try:
queue.put_nowait(self._progress[model_name].copy())
except asyncio.QueueFull:
logger.warning(f"Queue full for {model_name}, dropping update")
else: else:
logger.debug(f"No listeners for {model_name}, progress update stored: {progress_pct:.1f}%") logger.debug(f"No listeners for {model_name}, progress update stored: {progress_pct:.1f}%")
def get_progress(self, model_name: str) -> Optional[Dict]: def get_progress(self, model_name: str) -> Optional[Dict]:
"""Get current progress for a model.""" """Get current progress for a model. Thread-safe."""
return self._progress.get(model_name) with self._lock:
progress = self._progress.get(model_name)
return progress.copy() if progress else None
def get_all_active(self) -> List[Dict]: def get_all_active(self) -> List[Dict]:
"""Get all active downloads (status is 'downloading' or 'extracting').""" """Get all active downloads (status is 'downloading' or 'extracting'). Thread-safe."""
active = [] active = []
for model_name, progress in self._progress.items(): with self._lock:
status = progress.get("status", "") for model_name, progress in self._progress.items():
if status in ("downloading", "extracting"): status = progress.get("status", "")
active.append(progress.copy()) if status in ("downloading", "extracting"):
active.append(progress.copy())
return active return active
def create_progress_callback(self, model_name: str, filename: Optional[str] = None): def create_progress_callback(self, model_name: str, filename: Optional[str] = None):
@@ -110,6 +154,12 @@ class ProgressManager:
""" """
import logging import logging
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
# Store the main event loop for thread-safe operations
try:
self._main_loop = asyncio.get_running_loop()
except RuntimeError:
pass
queue = asyncio.Queue(maxsize=10) queue = asyncio.Queue(maxsize=10)
@@ -121,14 +171,19 @@ class ProgressManager:
logger.info(f"SSE client subscribed to {model_name}, total listeners: {len(self._listeners[model_name])}") logger.info(f"SSE client subscribed to {model_name}, total listeners: {len(self._listeners[model_name])}")
try: try:
# Send initial progress if available and still in progress # Send initial progress if available and still in progress (thread-safe read)
if model_name in self._progress: with self._lock:
status = self._progress[model_name].get('status') initial_progress = self._progress.get(model_name)
if initial_progress:
initial_progress = initial_progress.copy()
if initial_progress:
status = initial_progress.get('status')
# Only send initial progress if download is actually in progress # Only send initial progress if download is actually in progress
# Don't send old 'complete' or 'error' status from previous downloads # Don't send old 'complete' or 'error' status from previous downloads
if status in ('downloading', 'extracting'): if status in ('downloading', 'extracting'):
logger.info(f"Sending initial progress for {model_name}: {status}") logger.info(f"Sending initial progress for {model_name}: {status}")
yield f"data: {json.dumps(self._progress[model_name])}\n\n" yield f"data: {json.dumps(initial_progress)}\n\n"
else: else:
logger.info(f"Skipping initial progress for {model_name} (status: {status})") logger.info(f"Skipping initial progress for {model_name} (status: {status})")
else: else:
@@ -159,38 +214,50 @@ class ProgressManager:
logger.info(f"SSE client unsubscribed from {model_name}, remaining listeners: {len(self._listeners.get(model_name, []))}") logger.info(f"SSE client unsubscribed from {model_name}, remaining listeners: {len(self._listeners.get(model_name, []))}")
def mark_complete(self, model_name: str): def mark_complete(self, model_name: str):
"""Mark a model download as complete.""" """Mark a model download as complete. Thread-safe."""
import logging import logging
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
if model_name in self._progress: with self._lock:
self._progress[model_name]["status"] = "complete" if model_name in self._progress:
self._progress[model_name]["progress"] = 100.0 self._progress[model_name]["status"] = "complete"
logger.info(f"Marked {model_name} as complete") self._progress[model_name]["progress"] = 100.0
# Notify listeners progress_data = self._progress[model_name].copy()
if model_name in self._listeners: else:
for queue in self._listeners[model_name]: logger.warning(f"Cannot mark {model_name} as complete: not found in progress")
try: return
queue.put_nowait(self._progress[model_name].copy())
except asyncio.QueueFull: logger.info(f"Marked {model_name} as complete")
logger.warning(f"Queue full when marking {model_name} complete") # Notify listeners (thread-safe)
self._notify_listeners_threadsafe(model_name, progress_data)
def mark_error(self, model_name: str, error: str): def mark_error(self, model_name: str, error: str):
"""Mark a model download as failed.""" """Mark a model download as failed. Thread-safe."""
import logging import logging
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
if model_name in self._progress: with self._lock:
self._progress[model_name]["status"] = "error" if model_name in self._progress:
self._progress[model_name]["error"] = error self._progress[model_name]["status"] = "error"
logger.error(f"Marked {model_name} as error: {error}") self._progress[model_name]["error"] = error
# Notify listeners progress_data = self._progress[model_name].copy()
if model_name in self._listeners: else:
for queue in self._listeners[model_name]: # Create new progress entry for error
try: progress_data = {
queue.put_nowait(self._progress[model_name].copy()) "model_name": model_name,
except asyncio.QueueFull: "current": 0,
logger.warning(f"Queue full when marking {model_name} error") "total": 0,
"progress": 0,
"filename": None,
"status": "error",
"error": error,
"timestamp": datetime.now().isoformat(),
}
self._progress[model_name] = progress_data
logger.error(f"Marked {model_name} as error: {error}")
# Notify listeners (thread-safe)
self._notify_listeners_threadsafe(model_name, progress_data)
# Global progress manager instance # Global progress manager instance