mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-09-18 14:20:42 -07:00
Enhance contribution guidelines and improve FloatingGenerateBox component
- Updated CONTRIBUTING.md to include instructions for building with a local Qwen3-TTS development version, facilitating easier testing and development. - Refactored FloatingGenerateBox component to streamline the rendering of text and instruct fields, improving code readability and maintainability. - Added functionality to handle auto-resizing of text areas based on content changes, enhancing user experience. - Improved event handling for keyboard interactions in StoryTrackEditor, allowing for play/pause functionality with the spacebar. - Introduced a MiniSamplePlayer component in SampleList for better audio playback control, including play, pause, and seek features. - Implemented sample update functionality in the backend, allowing users to edit reference text for audio samples, with appropriate error handling and user feedback.
This commit is contained in:
+38
-12
@@ -26,7 +26,7 @@ class ProgressManager:
|
||||
):
|
||||
"""
|
||||
Update progress for a model download.
|
||||
|
||||
|
||||
Args:
|
||||
model_name: Name of the model (e.g., "qwen-tts-1.7B", "whisper-base")
|
||||
current: Current bytes downloaded
|
||||
@@ -34,8 +34,11 @@ class ProgressManager:
|
||||
filename: Current file being downloaded
|
||||
status: Status string (downloading, extracting, complete, error)
|
||||
"""
|
||||
import logging
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
progress_pct = (current / total * 100) if total > 0 else 0
|
||||
|
||||
|
||||
self._progress[model_name] = {
|
||||
"model_name": model_name,
|
||||
"current": current,
|
||||
@@ -45,14 +48,18 @@ class ProgressManager:
|
||||
"status": status,
|
||||
"timestamp": datetime.now().isoformat(),
|
||||
}
|
||||
|
||||
|
||||
# Notify all listeners
|
||||
if model_name in self._listeners:
|
||||
listener_count = len(self._listeners.get(model_name, []))
|
||||
if listener_count > 0:
|
||||
logger.debug(f"Notifying {listener_count} listeners for {model_name}: {progress_pct:.1f}% ({filename})")
|
||||
for queue in self._listeners[model_name]:
|
||||
try:
|
||||
queue.put_nowait(self._progress[model_name].copy())
|
||||
except asyncio.QueueFull:
|
||||
pass
|
||||
logger.warning(f"Queue full for {model_name}, dropping update")
|
||||
else:
|
||||
logger.debug(f"No listeners for {model_name}, progress update stored: {progress_pct:.1f}%")
|
||||
|
||||
def get_progress(self, model_name: str) -> Optional[Dict]:
|
||||
"""Get current progress for a model."""
|
||||
@@ -98,30 +105,40 @@ class ProgressManager:
|
||||
async def subscribe(self, model_name: str):
|
||||
"""
|
||||
Subscribe to progress updates for a model.
|
||||
|
||||
|
||||
Yields progress updates as Server-Sent Events.
|
||||
"""
|
||||
import logging
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
queue = asyncio.Queue(maxsize=10)
|
||||
|
||||
|
||||
# Add to listeners
|
||||
if model_name not in self._listeners:
|
||||
self._listeners[model_name] = []
|
||||
self._listeners[model_name].append(queue)
|
||||
|
||||
|
||||
logger.info(f"SSE client subscribed to {model_name}, total listeners: {len(self._listeners[model_name])}")
|
||||
|
||||
try:
|
||||
# Send initial progress if available
|
||||
if model_name in self._progress:
|
||||
logger.info(f"Sending initial progress for {model_name}: {self._progress[model_name].get('status')}")
|
||||
yield f"data: {json.dumps(self._progress[model_name])}\n\n"
|
||||
|
||||
else:
|
||||
logger.info(f"No initial progress available for {model_name}")
|
||||
|
||||
# Stream updates
|
||||
while True:
|
||||
try:
|
||||
# Wait for update with timeout
|
||||
progress = await asyncio.wait_for(queue.get(), timeout=1.0)
|
||||
logger.debug(f"Sending progress update for {model_name}: {progress.get('status')} - {progress.get('progress', 0):.1f}%")
|
||||
yield f"data: {json.dumps(progress)}\n\n"
|
||||
|
||||
|
||||
# Stop if complete or error
|
||||
if progress.get("status") in ("complete", "error"):
|
||||
logger.info(f"Download {progress.get('status')} for {model_name}, closing SSE connection")
|
||||
break
|
||||
except asyncio.TimeoutError:
|
||||
# Send heartbeat
|
||||
@@ -133,32 +150,41 @@ class ProgressManager:
|
||||
self._listeners[model_name].remove(queue)
|
||||
if not self._listeners[model_name]:
|
||||
del self._listeners[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):
|
||||
"""Mark a model download as complete."""
|
||||
import logging
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
if model_name in self._progress:
|
||||
self._progress[model_name]["status"] = "complete"
|
||||
self._progress[model_name]["progress"] = 100.0
|
||||
logger.info(f"Marked {model_name} as complete")
|
||||
# Notify listeners
|
||||
if model_name in self._listeners:
|
||||
for queue in self._listeners[model_name]:
|
||||
try:
|
||||
queue.put_nowait(self._progress[model_name].copy())
|
||||
except asyncio.QueueFull:
|
||||
pass
|
||||
logger.warning(f"Queue full when marking {model_name} complete")
|
||||
|
||||
def mark_error(self, model_name: str, error: str):
|
||||
"""Mark a model download as failed."""
|
||||
import logging
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
if model_name in self._progress:
|
||||
self._progress[model_name]["status"] = "error"
|
||||
self._progress[model_name]["error"] = error
|
||||
logger.error(f"Marked {model_name} as error: {error}")
|
||||
# Notify listeners
|
||||
if model_name in self._listeners:
|
||||
for queue in self._listeners[model_name]:
|
||||
try:
|
||||
queue.put_nowait(self._progress[model_name].copy())
|
||||
except asyncio.QueueFull:
|
||||
pass
|
||||
logger.warning(f"Queue full when marking {model_name} error")
|
||||
|
||||
|
||||
# Global progress manager instance
|
||||
|
||||
Reference in New Issue
Block a user