mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-09-29 15:15:27 -07:00
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:
+12
-4
@@ -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
@@ -6,15 +6,54 @@ 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,
|
||||||
@@ -27,6 +66,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")
|
||||||
current: Current bytes downloaded
|
current: Current bytes downloaded
|
||||||
@@ -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,
|
||||||
@@ -49,29 +90,32 @@ class ProgressManager:
|
|||||||
"timestamp": datetime.now().isoformat(),
|
"timestamp": datetime.now().isoformat(),
|
||||||
}
|
}
|
||||||
|
|
||||||
# Notify all listeners
|
# Thread-safe update of progress dict
|
||||||
|
with self._lock:
|
||||||
|
self._progress[model_name] = progress_data
|
||||||
|
|
||||||
|
# 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):
|
||||||
@@ -111,6 +155,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)
|
||||||
|
|
||||||
# Add to listeners
|
# Add to listeners
|
||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user