diff --git a/app/src/App.tsx b/app/src/App.tsx
index f8d8f02c..85d51597 100644
--- a/app/src/App.tsx
+++ b/app/src/App.tsx
@@ -13,7 +13,9 @@ import { TitleBarDragRegion } from '@/components/TitleBarDragRegion';
import { UpdateNotification } from '@/components/UpdateNotification';
import { Toaster } from '@/components/ui/toaster';
import { ProfileList } from '@/components/VoiceProfiles/ProfileList';
-import { isTauri, isMacOS, setupWindowCloseHandler, startServer } from '@/lib/tauri';
+import { useModelDownloadToast } from '@/lib/hooks/useModelDownloadToast';
+import { useRestoreActiveTasks, MODEL_DISPLAY_NAMES } from '@/lib/hooks/useRestoreActiveTasks';
+import { isMacOS, isTauri, setupWindowCloseHandler, startServer } from '@/lib/tauri';
// Track if server is starting to prevent duplicate starts
let serverStarting = false;
@@ -46,6 +48,9 @@ function App() {
const [serverReady, setServerReady] = useState(false);
const [loadingMessageIndex, setLoadingMessageIndex] = useState(0);
+ // Monitor active downloads/generations and show toasts for them
+ const activeDownloads = useRestoreActiveTasks();
+
// Setup window close handler and auto-start server when running in Tauri (production only)
useEffect(() => {
if (!isTauri()) {
@@ -118,27 +123,27 @@ function App() {
return (
-
-
-
-

-
-
-
+
+
+
+
+
+
+
+
);
}
@@ -202,9 +207,41 @@ function App() {
{/* Audio Player - always visible except on settings */}
{activeTab !== 'settings' &&
}
+ {/* Show download toasts for any active downloads (from anywhere) */}
+ {activeDownloads.map((download) => {
+ const displayName = MODEL_DISPLAY_NAMES[download.model_name] || download.model_name;
+ return (
+
+ );
+ })}
+
);
}
+/**
+ * Component that restores a download toast for a specific model.
+ */
+function DownloadToastRestorer({
+ modelName,
+ displayName,
+}: {
+ modelName: string;
+ displayName: string;
+}) {
+ // Use the download toast hook to restore the toast
+ useModelDownloadToast({
+ modelName,
+ displayName,
+ enabled: true,
+ });
+
+ return null;
+}
+
export default App;
diff --git a/app/src/components/AudioPlayer/AudioPlayer.tsx b/app/src/components/AudioPlayer/AudioPlayer.tsx
index 81fb51f0..1e125e13 100644
--- a/app/src/components/AudioPlayer/AudioPlayer.tsx
+++ b/app/src/components/AudioPlayer/AudioPlayer.tsx
@@ -16,18 +16,19 @@ export function AudioPlayer() {
duration,
volume,
isLooping,
+ shouldRestart,
setIsPlaying,
setCurrentTime,
setDuration,
setVolume,
toggleLoop,
+ clearRestartFlag,
} = usePlayerStore();
const waveformRef = useRef
(null);
const wavesurferRef = useRef(null);
const loadingRef = useRef(false);
const previousAudioIdRef = useRef(null);
- const previousCurrentTimeRef = useRef(0);
const hasInitializedRef = useRef(false);
const [isLoading, setIsLoading] = useState(false);
const [error, setError] = useState(null);
@@ -367,51 +368,30 @@ export function AudioPlayer() {
if (audioId !== previousAudioIdRef.current && previousAudioIdRef.current !== null) {
hasInitializedRef.current = false;
}
+ if (audioId !== null) {
+ previousAudioIdRef.current = audioId;
+ }
}, [duration, audioId]);
- // Handle clicking the same audio again - always restart from beginning
- // When setAudio is called with the same audioId, it sets currentTime to 0 in the store
- // but WaveSurfer's actual position is still wherever it was. We detect this mismatch and reset.
+ // Handle restart flag - when history item is clicked again, restart from beginning
useEffect(() => {
const wavesurfer = wavesurferRef.current;
- if (!wavesurfer || !audioId || duration === 0 || !hasInitializedRef.current) {
- // Update the refs even if we don't process
- if (audioId !== null) {
- previousAudioIdRef.current = audioId;
- }
- previousCurrentTimeRef.current = currentTime;
+ if (!wavesurfer || !shouldRestart || duration === 0) {
return;
}
- const previousAudioId = previousAudioIdRef.current;
- const previousCurrentTime = previousCurrentTimeRef.current;
+ // Reset to beginning and play
+ console.log('Restarting current audio from beginning');
+ wavesurfer.seekTo(0);
+ wavesurfer.play().catch((error) => {
+ console.error('Failed to play after restart:', error);
+ setIsPlaying(false);
+ setError(`Playback error: ${error instanceof Error ? error.message : String(error)}`);
+ });
- // Check if the same audio was clicked again
- // This happens when:
- // 1. audioId matches the previous one (same audio)
- // 2. currentTime was reset from a non-zero value to 0 (setAudio was called)
- // 3. WaveSurfer is not at the beginning (needs reset)
- const wasResetToZero = previousCurrentTime > 0.1 && currentTime < 0.1;
- const isSameAudio = audioId === previousAudioId;
- const wavesurferPosition = wavesurfer.getCurrentTime();
- const wavesurferNotAtStart = wavesurferPosition > 0.1;
-
- // Update refs for next time
- previousAudioIdRef.current = audioId;
- previousCurrentTimeRef.current = currentTime;
-
- // If same audio was clicked (reset to 0) and WaveSurfer is not at start, reset it
- if (isSameAudio && wasResetToZero && wavesurferNotAtStart) {
- // Reset to beginning and play
- console.log('Same audio clicked again, resetting to beginning');
- wavesurfer.seekTo(0);
- wavesurfer.play().catch((error) => {
- console.error('Failed to play after reset:', error);
- setIsPlaying(false);
- setError(`Playback error: ${error instanceof Error ? error.message : String(error)}`);
- });
- }
- }, [audioId, duration, currentTime, setIsPlaying]);
+ // Clear the restart flag
+ clearRestartFlag();
+ }, [shouldRestart, duration, setIsPlaying, clearRestartFlag]);
// Handle loop - WaveSurfer handles this via the 'finish' event
diff --git a/app/src/components/Generation/GenerationForm.tsx b/app/src/components/Generation/GenerationForm.tsx
index 37c3c246..d11d36d0 100644
--- a/app/src/components/Generation/GenerationForm.tsx
+++ b/app/src/components/Generation/GenerationForm.tsx
@@ -25,6 +25,7 @@ import {
import { Textarea } from '@/components/ui/textarea';
import { useToast } from '@/components/ui/use-toast';
import { apiClient } from '@/lib/api/client';
+import { LANGUAGE_CODES, LANGUAGE_OPTIONS, type LanguageCode } from '@/lib/constants/languages';
import { useGeneration } from '@/lib/hooks/useGeneration';
import { useModelDownloadToast } from '@/lib/hooks/useModelDownloadToast';
import { useProfile } from '@/lib/hooks/useProfiles';
@@ -34,7 +35,7 @@ import { useUIStore } from '@/stores/uiStore';
const generationSchema = z.object({
text: z.string().min(1, 'Text is required').max(5000),
- language: z.enum(['en', 'zh']),
+ language: z.enum(LANGUAGE_CODES as [LanguageCode, ...LanguageCode[]]),
seed: z.number().int().optional(),
modelSize: z.enum(['1.7B', '0.6B']).optional(),
instruct: z.string().max(500).optional(),
@@ -214,8 +215,11 @@ export function GenerationForm() {
- English
- Chinese
+ {LANGUAGE_OPTIONS.map((lang) => (
+
+ {lang.label}
+
+ ))}
diff --git a/app/src/components/History/HistoryTable.tsx b/app/src/components/History/HistoryTable.tsx
index 4c0bec95..c0e3f7f0 100644
--- a/app/src/components/History/HistoryTable.tsx
+++ b/app/src/components/History/HistoryTable.tsx
@@ -31,6 +31,7 @@ export function HistoryTable() {
const deleteGeneration = useDeleteGeneration();
const setAudio = usePlayerStore((state) => state.setAudio);
+ const restartCurrentAudio = usePlayerStore((state) => state.restartCurrentAudio);
const currentAudioId = usePlayerStore((state) => state.audioId);
const isPlaying = usePlayerStore((state) => state.isPlaying);
const audioUrl = usePlayerStore((state) => state.audioUrl);
@@ -49,9 +50,14 @@ export function HistoryTable() {
}, []);
const handlePlay = (audioId: string, text: string) => {
- const audioUrl = apiClient.getAudioUrl(audioId);
- // If clicking the same audio that's playing, it will be handled by the player
- setAudio(audioUrl, audioId, text.substring(0, 50));
+ // If clicking the same audio, restart it from the beginning
+ if (currentAudioId === audioId) {
+ restartCurrentAudio();
+ } else {
+ // Otherwise, load the new audio
+ const audioUrl = apiClient.getAudioUrl(audioId);
+ setAudio(audioUrl, audioId, text.substring(0, 50));
+ }
};
const handleDownload = (audioId: string, text: string) => {
diff --git a/app/src/components/ServerSettings/ModelManagement.tsx b/app/src/components/ServerSettings/ModelManagement.tsx
index 2ab4224f..3b160e5b 100644
--- a/app/src/components/ServerSettings/ModelManagement.tsx
+++ b/app/src/components/ServerSettings/ModelManagement.tsx
@@ -54,10 +54,10 @@ export function ModelManagement() {
return apiClient.triggerModelDownload(modelName);
},
onSuccess: () => {
- // Refetch status after a delay to see progress
- setTimeout(() => {
- queryClient.invalidateQueries({ queryKey: ['modelStatus'] });
- }, 1000);
+ // Download completed - clear state and refetch status
+ setDownloadingModel(null);
+ setDownloadingDisplayName(null);
+ queryClient.invalidateQueries({ queryKey: ['modelStatus'] });
},
onError: (error: Error) => {
setDownloadingModel(null);
@@ -68,13 +68,6 @@ export function ModelManagement() {
variant: 'destructive',
});
},
- onSettled: () => {
- // Clear downloading state after a delay to allow progress to show
- setTimeout(() => {
- setDownloadingModel(null);
- setDownloadingDisplayName(null);
- }, 2000);
- },
});
const deleteMutation = useMutation({
diff --git a/app/src/components/Sidebar.tsx b/app/src/components/Sidebar.tsx
index e9db6a29..81cb218d 100644
--- a/app/src/components/Sidebar.tsx
+++ b/app/src/components/Sidebar.tsx
@@ -1,4 +1,4 @@
-import { Home, Loader2, Settings } from 'lucide-react';
+import { Volume2, Loader2, Settings } from 'lucide-react';
import voiceboxLogo from '@/assets/voicebox-logo.png';
import { cn } from '@/lib/utils/cn';
import { useGenerationStore } from '@/stores/generationStore';
@@ -11,7 +11,7 @@ interface SidebarProps {
}
const tabs = [
- { id: 'main', icon: Home, label: 'Main' },
+ { id: 'main', icon: Volume2, label: 'Main' },
{ id: 'settings', icon: Settings, label: 'Settings' },
];
diff --git a/app/src/components/VoiceProfiles/ProfileForm.tsx b/app/src/components/VoiceProfiles/ProfileForm.tsx
index 9ba3ef69..d01681d6 100644
--- a/app/src/components/VoiceProfiles/ProfileForm.tsx
+++ b/app/src/components/VoiceProfiles/ProfileForm.tsx
@@ -30,6 +30,7 @@ import {
import { Textarea } from '@/components/ui/textarea';
import { Tabs, TabsContent, TabsList, TabsTrigger } from '@/components/ui/tabs';
import { useToast } from '@/components/ui/use-toast';
+import { LANGUAGE_CODES, LANGUAGE_OPTIONS, type LanguageCode } from '@/lib/constants/languages';
import {
useCreateProfile,
useProfile,
@@ -68,7 +69,7 @@ const profileSchema = z
.object({
name: z.string().min(1, 'Name is required').max(100),
description: z.string().max(500).optional(),
- language: z.enum(['en', 'zh']),
+ language: z.enum(LANGUAGE_CODES as [LanguageCode, ...LanguageCode[]]),
// Sample fields - only required when creating (not editing)
sampleFile: z.instanceof(File).optional(),
referenceText: z.string().max(1000).optional(),
@@ -186,7 +187,7 @@ export function ProfileForm() {
form.reset({
name: editingProfile.name,
description: editingProfile.description || '',
- language: editingProfile.language as 'en' | 'zh',
+ language: editingProfile.language as LanguageCode,
sampleFile: undefined,
referenceText: undefined,
});
@@ -214,7 +215,7 @@ export function ProfileForm() {
}
try {
- const language = form.getValues('language') as 'en' | 'zh' | undefined;
+ const language = form.getValues('language');
const result = await transcribe.mutateAsync({ file, language });
form.setValue('referenceText', result.text, { shouldValidate: true });
@@ -405,8 +406,11 @@ export function ProfileForm() {
- English
- Chinese
+ {LANGUAGE_OPTIONS.map((lang) => (
+
+ {lang.label}
+
+ ))}
diff --git a/app/src/components/ui/progress.tsx b/app/src/components/ui/progress.tsx
index cebbc856..6d79c1d9 100644
--- a/app/src/components/ui/progress.tsx
+++ b/app/src/components/ui/progress.tsx
@@ -12,7 +12,7 @@ const Progress = React.forwardRef<
{...props}
>
diff --git a/app/src/components/ui/toaster.tsx b/app/src/components/ui/toaster.tsx
index 694ea016..ae772ab6 100644
--- a/app/src/components/ui/toaster.tsx
+++ b/app/src/components/ui/toaster.tsx
@@ -15,7 +15,7 @@ export function Toaster() {
{toasts.map(({ id, title, description, action, ...props }) => (
-
+
{title && {title}}
{description && {description}}
diff --git a/app/src/lib/api/client.ts b/app/src/lib/api/client.ts
index 94115a2a..4694b100 100644
--- a/app/src/lib/api/client.ts
+++ b/app/src/lib/api/client.ts
@@ -12,6 +12,7 @@ import type {
HealthResponse,
ModelStatusListResponse,
ModelDownloadRequest,
+ ActiveTasksResponse,
} from './types';
class ApiClient {
@@ -225,6 +226,11 @@ class ApiClient {
method: 'DELETE',
});
}
+
+ // Task Management
+ async getActiveTasks(): Promise
{
+ return this.request('/tasks/active');
+ }
}
export const apiClient = new ApiClient();
diff --git a/app/src/lib/api/types.ts b/app/src/lib/api/types.ts
index bde0c1d0..a473f29e 100644
--- a/app/src/lib/api/types.ts
+++ b/app/src/lib/api/types.ts
@@ -105,3 +105,21 @@ export interface ModelStatusListResponse {
export interface ModelDownloadRequest {
model_name: string;
}
+
+export interface ActiveDownloadTask {
+ model_name: string;
+ status: string;
+ started_at: string;
+}
+
+export interface ActiveGenerationTask {
+ task_id: string;
+ profile_id: string;
+ text_preview: string;
+ started_at: string;
+}
+
+export interface ActiveTasksResponse {
+ downloads: ActiveDownloadTask[];
+ generations: ActiveGenerationTask[];
+}
diff --git a/app/src/lib/constants/languages.ts b/app/src/lib/constants/languages.ts
new file mode 100644
index 00000000..9ffc396f
--- /dev/null
+++ b/app/src/lib/constants/languages.ts
@@ -0,0 +1,26 @@
+/**
+ * Supported languages for Qwen3-TTS
+ * Based on: https://github.com/QwenLM/Qwen3-TTS
+ */
+
+export const SUPPORTED_LANGUAGES = {
+ zh: 'Chinese',
+ en: 'English',
+ ja: 'Japanese',
+ ko: 'Korean',
+ de: 'German',
+ fr: 'French',
+ ru: 'Russian',
+ pt: 'Portuguese',
+ es: 'Spanish',
+ it: 'Italian',
+} as const;
+
+export type LanguageCode = keyof typeof SUPPORTED_LANGUAGES;
+
+export const LANGUAGE_CODES = Object.keys(SUPPORTED_LANGUAGES) as LanguageCode[];
+
+export const LANGUAGE_OPTIONS = LANGUAGE_CODES.map((code) => ({
+ value: code,
+ label: SUPPORTED_LANGUAGES[code],
+}));
diff --git a/app/src/lib/hooks/useRestoreActiveTasks.tsx b/app/src/lib/hooks/useRestoreActiveTasks.tsx
new file mode 100644
index 00000000..063e6bcb
--- /dev/null
+++ b/app/src/lib/hooks/useRestoreActiveTasks.tsx
@@ -0,0 +1,87 @@
+import { useCallback, useEffect, useRef, useState } from 'react';
+import { apiClient } from '@/lib/api/client';
+import { useGenerationStore } from '@/stores/generationStore';
+import type { ActiveDownloadTask } from '@/lib/api/types';
+
+// Polling interval in milliseconds
+const POLL_INTERVAL = 2000;
+
+/**
+ * Hook to monitor active tasks (downloads and generations).
+ * Polls the server periodically to catch downloads triggered from anywhere
+ * (transcription, generation, explicit download, etc.).
+ *
+ * Returns the active downloads so components can render download toasts.
+ */
+export function useRestoreActiveTasks() {
+ const [activeDownloads, setActiveDownloads] = useState([]);
+ const setIsGenerating = useGenerationStore((state) => state.setIsGenerating);
+ const setActiveGenerationId = useGenerationStore((state) => state.setActiveGenerationId);
+
+ // Track which downloads we've seen to detect new ones
+ const seenDownloadsRef = useRef>(new Set());
+
+ const fetchActiveTasks = useCallback(async () => {
+ try {
+ const tasks = await apiClient.getActiveTasks();
+
+ // Update generation state
+ if (tasks.generations.length > 0) {
+ setIsGenerating(true);
+ setActiveGenerationId(tasks.generations[0].task_id);
+ } else {
+ // Only clear if we were tracking a generation
+ const currentId = useGenerationStore.getState().activeGenerationId;
+ if (currentId) {
+ setIsGenerating(false);
+ setActiveGenerationId(null);
+ }
+ }
+
+ // Update active downloads
+ // Keep track of all active downloads (including new ones)
+ const currentDownloadNames = new Set(tasks.downloads.map((d) => d.model_name));
+
+ // Remove completed downloads from our seen set
+ for (const name of seenDownloadsRef.current) {
+ if (!currentDownloadNames.has(name)) {
+ seenDownloadsRef.current.delete(name);
+ }
+ }
+
+ // Add new downloads to seen set
+ for (const download of tasks.downloads) {
+ seenDownloadsRef.current.add(download.model_name);
+ }
+
+ setActiveDownloads(tasks.downloads);
+ } catch (error) {
+ // Silently fail - server might be temporarily unavailable
+ console.debug('Failed to fetch active tasks:', error);
+ }
+ }, [setIsGenerating, setActiveGenerationId]);
+
+ useEffect(() => {
+ // Fetch immediately on mount
+ fetchActiveTasks();
+
+ // Poll for active tasks
+ const interval = setInterval(fetchActiveTasks, POLL_INTERVAL);
+
+ return () => clearInterval(interval);
+ }, [fetchActiveTasks]);
+
+ return activeDownloads;
+}
+
+/**
+ * Map model names to display names for download toasts.
+ */
+export const MODEL_DISPLAY_NAMES: Record = {
+ 'qwen-tts-1.7B': 'Qwen TTS 1.7B',
+ 'qwen-tts-0.6B': 'Qwen TTS 0.6B',
+ 'whisper-base': 'Whisper Base',
+ 'whisper-small': 'Whisper Small',
+ 'whisper-medium': 'Whisper Medium',
+ 'whisper-large': 'Whisper Large',
+};
diff --git a/app/src/stores/generationStore.ts b/app/src/stores/generationStore.ts
index c507f5f3..c0d63383 100644
--- a/app/src/stores/generationStore.ts
+++ b/app/src/stores/generationStore.ts
@@ -2,10 +2,14 @@ import { create } from 'zustand';
interface GenerationState {
isGenerating: boolean;
+ activeGenerationId: string | null;
setIsGenerating: (generating: boolean) => void;
+ setActiveGenerationId: (id: string | null) => void;
}
export const useGenerationStore = create((set) => ({
isGenerating: false,
+ activeGenerationId: null,
setIsGenerating: (generating) => set({ isGenerating: generating }),
+ setActiveGenerationId: (id) => set({ activeGenerationId: id }),
}));
diff --git a/app/src/stores/playerStore.ts b/app/src/stores/playerStore.ts
index ded4617b..006e0d85 100644
--- a/app/src/stores/playerStore.ts
+++ b/app/src/stores/playerStore.ts
@@ -9,6 +9,7 @@ interface PlayerState {
duration: number;
volume: number;
isLooping: boolean;
+ shouldRestart: boolean;
setAudio: (url: string, id: string, title?: string) => void;
setIsPlaying: (playing: boolean) => void;
@@ -16,6 +17,8 @@ interface PlayerState {
setDuration: (duration: number) => void;
setVolume: (volume: number) => void;
toggleLoop: () => void;
+ restartCurrentAudio: () => void;
+ clearRestartFlag: () => void;
reset: () => void;
}
@@ -28,6 +31,7 @@ export const usePlayerStore = create((set) => ({
duration: 0,
volume: 1,
isLooping: false,
+ shouldRestart: false,
setAudio: (url, id, title) =>
set({
@@ -36,12 +40,15 @@ export const usePlayerStore = create((set) => ({
title: title || null,
currentTime: 0,
isPlaying: false,
+ shouldRestart: false,
}),
setIsPlaying: (playing) => set({ isPlaying: playing }),
setCurrentTime: (time) => set({ currentTime: time }),
setDuration: (duration) => set({ duration }),
setVolume: (volume) => set({ volume }),
toggleLoop: () => set((state) => ({ isLooping: !state.isLooping })),
+ restartCurrentAudio: () => set({ shouldRestart: true }),
+ clearRestartFlag: () => set({ shouldRestart: false }),
reset: () =>
set({
audioUrl: null,
@@ -51,5 +58,6 @@ export const usePlayerStore = create((set) => ({
currentTime: 0,
duration: 0,
isLooping: false,
+ shouldRestart: false,
}),
}));
diff --git a/backend/main.py b/backend/main.py
index 19a90c6f..152ddec3 100644
--- a/backend/main.py
+++ b/backend/main.py
@@ -10,6 +10,7 @@ from fastapi.responses import FileResponse, StreamingResponse
from fastapi.staticfiles import StaticFiles
from sqlalchemy.orm import Session
from typing import List, Optional
+from datetime import datetime
import uvicorn
import argparse
import torch
@@ -21,6 +22,7 @@ import uuid
from . import database, models, profiles, history, tts, transcribe, config, export_import
from .database import get_db, Generation as DBGeneration, VoiceProfile as DBVoiceProfile
from .utils.progress import get_progress_manager
+from .utils.tasks import get_task_manager
app = FastAPI(
title="voicebox API",
@@ -300,7 +302,17 @@ async def generate_speech(
db: Session = Depends(get_db),
):
"""Generate speech from text using a voice profile."""
+ task_manager = get_task_manager()
+ generation_id = str(uuid.uuid4())
+
try:
+ # Start tracking generation
+ task_manager.start_generation(
+ task_id=generation_id,
+ profile_id=data.profile_id,
+ text=data.text,
+ )
+
# Get profile
profile = await profiles.get_profile(data.profile_id, db)
if not profile:
@@ -329,7 +341,6 @@ async def generate_speech(
duration = len(audio) / sample_rate
# Save audio
- generation_id = str(uuid.uuid4())
audio_path = config.get_generations_dir() / f"{generation_id}.wav"
from .utils.audio import save_audio
@@ -347,11 +358,16 @@ async def generate_speech(
instruct=data.instruct,
)
+ # Mark generation as complete
+ task_manager.complete_generation(generation_id)
+
return generation
except ValueError as e:
+ task_manager.complete_generation(generation_id)
raise HTTPException(status_code=400, detail=str(e))
except Exception as e:
+ task_manager.complete_generation(generation_id)
raise HTTPException(status_code=500, detail=str(e))
@@ -742,6 +758,8 @@ async def trigger_model_download(request: models.ModelDownloadRequest):
"""Trigger download of a specific model."""
import asyncio
+ task_manager = get_task_manager()
+
model_configs = {
"qwen-tts-1.7B": {
"model_size": "1.7B",
@@ -775,12 +793,20 @@ async def trigger_model_download(request: models.ModelDownloadRequest):
config = model_configs[request.model_name]
try:
+ # Start tracking download
+ task_manager.start_download(request.model_name)
+
# Trigger download by loading the model (which will download if not cached)
# Run in background to avoid blocking
await asyncio.to_thread(config["load_func"])
+ # Mark download as complete
+ task_manager.complete_download(request.model_name)
+
return {"message": f"Model {request.model_name} download started"}
except Exception as e:
+ # Mark download as failed
+ task_manager.error_download(request.model_name, str(e))
raise HTTPException(status_code=500, detail=str(e))
@@ -866,6 +892,72 @@ async def delete_model(model_name: str):
raise HTTPException(status_code=500, detail=f"Failed to delete model: {str(e)}")
+# ============================================
+# TASK MANAGEMENT
+# ============================================
+
+@app.get("/tasks/active", response_model=models.ActiveTasksResponse)
+async def get_active_tasks():
+ """Return all currently active downloads and generations."""
+ task_manager = get_task_manager()
+ progress_manager = get_progress_manager()
+
+ # Get active downloads from both task manager and progress manager
+ # Task manager tracks which downloads are active
+ # Progress manager has the actual progress data
+ active_downloads = []
+ task_manager_downloads = task_manager.get_active_downloads()
+ progress_active = progress_manager.get_all_active()
+
+ # Combine data from both sources
+ download_map = {task.model_name: task for task in task_manager_downloads}
+ progress_map = {p["model_name"]: p for p in progress_active}
+
+ # Create unified list
+ all_model_names = set(download_map.keys()) | set(progress_map.keys())
+ for model_name in all_model_names:
+ task = download_map.get(model_name)
+ progress = progress_map.get(model_name)
+
+ if task:
+ active_downloads.append(models.ActiveDownloadTask(
+ model_name=model_name,
+ status=task.status,
+ started_at=task.started_at,
+ ))
+ elif progress:
+ # Progress exists but no task - create from progress data
+ timestamp_str = progress.get("timestamp")
+ if timestamp_str:
+ try:
+ started_at = datetime.fromisoformat(timestamp_str.replace('Z', '+00:00'))
+ except (ValueError, AttributeError):
+ started_at = datetime.utcnow()
+ else:
+ started_at = datetime.utcnow()
+
+ active_downloads.append(models.ActiveDownloadTask(
+ model_name=model_name,
+ status=progress.get("status", "downloading"),
+ started_at=started_at,
+ ))
+
+ # Get active generations
+ active_generations = []
+ for gen_task in task_manager.get_active_generations():
+ active_generations.append(models.ActiveGenerationTask(
+ task_id=gen_task.task_id,
+ profile_id=gen_task.profile_id,
+ text_preview=gen_task.text_preview,
+ started_at=gen_task.started_at,
+ ))
+
+ return models.ActiveTasksResponse(
+ downloads=active_downloads,
+ generations=active_generations,
+ )
+
+
# ============================================
# STARTUP & SHUTDOWN
# ============================================
diff --git a/backend/models.py b/backend/models.py
index 594b6778..caeeebba 100644
--- a/backend/models.py
+++ b/backend/models.py
@@ -11,7 +11,7 @@ class VoiceProfileCreate(BaseModel):
"""Request model for creating a voice profile."""
name: str = Field(..., min_length=1, max_length=100)
description: Optional[str] = Field(None, max_length=500)
- language: str = Field(default="en", pattern="^(en|zh)$")
+ language: str = Field(default="en", pattern="^(zh|en|ja|ko|de|fr|ru|pt|es|it)$")
class VoiceProfileResponse(BaseModel):
@@ -47,7 +47,7 @@ class GenerationRequest(BaseModel):
"""Request model for voice generation."""
profile_id: str
text: str = Field(..., min_length=1, max_length=5000)
- language: str = Field(default="en", pattern="^(en|zh)$")
+ language: str = Field(default="en", pattern="^(zh|en|ja|ko|de|fr|ru|pt|es|it)$")
seed: Optional[int] = Field(None, ge=0)
model_size: Optional[str] = Field(default="1.7B", pattern="^(1\\.7B|0\\.6B)$")
instruct: Optional[str] = Field(None, max_length=500)
@@ -138,3 +138,24 @@ class ModelStatusListResponse(BaseModel):
class ModelDownloadRequest(BaseModel):
"""Request model for triggering model download."""
model_name: str
+
+
+class ActiveDownloadTask(BaseModel):
+ """Response model for active download task."""
+ model_name: str
+ status: str
+ started_at: datetime
+
+
+class ActiveGenerationTask(BaseModel):
+ """Response model for active generation task."""
+ task_id: str
+ profile_id: str
+ text_preview: str
+ started_at: datetime
+
+
+class ActiveTasksResponse(BaseModel):
+ """Response model for active tasks."""
+ downloads: List[ActiveDownloadTask]
+ generations: List[ActiveGenerationTask]
diff --git a/backend/transcribe.py b/backend/transcribe.py
index dfdd2360..6f21966b 100644
--- a/backend/transcribe.py
+++ b/backend/transcribe.py
@@ -9,6 +9,7 @@ import numpy as np
from pathlib import Path
from .utils.progress import get_progress_manager
from .utils.hf_progress import HFProgressTracker, create_hf_progress_callback
+from .utils.tasks import get_task_manager
class WhisperModel:
@@ -55,8 +56,21 @@ class WhisperModel:
progress_manager = get_progress_manager()
progress_model_name = f"whisper-{model_size}"
+ # Start tracking download task
+ task_manager = get_task_manager()
+ task_manager.start_download(progress_model_name)
+
print(f"Loading Whisper model {model_size} on {self.device}...")
+ # Initialize progress state to show download has started
+ progress_manager.update_progress(
+ model_name=progress_model_name,
+ current=0,
+ total=1, # Set to 1 initially, will be updated by callback
+ filename="",
+ status="downloading",
+ )
+
# Set up progress callback
progress_callback = create_hf_progress_callback(progress_model_name, progress_manager)
tracker = HFProgressTracker(progress_callback)
@@ -71,13 +85,17 @@ class WhisperModel:
# Mark as complete
progress_manager.mark_complete(progress_model_name)
+ task_manager.complete_download(progress_model_name)
print(f"Whisper model {model_size} loaded successfully")
except Exception as e:
print(f"Error loading Whisper model: {e}")
progress_manager = get_progress_manager()
- progress_manager.mark_error(f"whisper-{model_size}", str(e))
+ task_manager = get_task_manager()
+ progress_model_name = f"whisper-{model_size}"
+ progress_manager.mark_error(progress_model_name, str(e))
+ task_manager.error_download(progress_model_name, str(e))
raise
async def load_model_async(self, model_size: Optional[str] = None):
diff --git a/backend/tts.py b/backend/tts.py
index 70961135..f2e9783e 100644
--- a/backend/tts.py
+++ b/backend/tts.py
@@ -14,6 +14,7 @@ from .utils.cache import get_cache_key, get_cached_voice_prompt, cache_voice_pro
from .utils.audio import normalize_audio
from .utils.progress import get_progress_manager
from .utils.hf_progress import HFProgressTracker, create_hf_progress_callback
+from .utils.tasks import get_task_manager
from . import config
@@ -111,6 +112,10 @@ class TTSModel:
if model_path.startswith("Qwen/"):
print(f"Loading TTS model {model_size} on {self.device}...")
+ # Start tracking download task
+ task_manager = get_task_manager()
+ task_manager.start_download(model_name)
+
# Initialize progress state to show download has started
progress_manager.update_progress(
model_name=model_name,
@@ -135,6 +140,7 @@ class TTSModel:
# Mark as complete
progress_manager.mark_complete(model_name)
+ task_manager.complete_download(model_name)
else:
# Local model, no download needed
print(f"Loading TTS model {model_size} on {self.device}...")
@@ -152,13 +158,19 @@ class TTSModel:
except ImportError as e:
print(f"Error: qwen_tts package not found. Install with: pip install git+https://github.com/QwenLM/Qwen3-TTS.git")
progress_manager = get_progress_manager()
- progress_manager.mark_error(f"qwen-tts-{model_size}", str(e))
+ task_manager = get_task_manager()
+ model_name = f"qwen-tts-{model_size}"
+ progress_manager.mark_error(model_name, str(e))
+ task_manager.error_download(model_name, str(e))
raise
except Exception as e:
print(f"Error loading TTS model: {e}")
print(f"Tip: The model will be automatically downloaded from HuggingFace Hub on first use.")
progress_manager = get_progress_manager()
- progress_manager.mark_error(f"qwen-tts-{model_size}", str(e))
+ task_manager = get_task_manager()
+ model_name = f"qwen-tts-{model_size}"
+ progress_manager.mark_error(model_name, str(e))
+ task_manager.error_download(model_name, str(e))
raise
async def load_model_async(self, model_size: Optional[str] = None):
diff --git a/backend/utils/hf_progress.py b/backend/utils/hf_progress.py
index b935208a..ed9aeb19 100644
--- a/backend/utils/hf_progress.py
+++ b/backend/utils/hf_progress.py
@@ -5,116 +5,124 @@ HuggingFace Hub download progress tracking.
from typing import Optional, Callable
from contextlib import contextmanager
import threading
+import sys
class HFProgressTracker:
- """Tracks HuggingFace Hub download progress by intercepting hf_hub_download and snapshot_download."""
+ """Tracks HuggingFace Hub download progress by intercepting tqdm."""
def __init__(self, progress_callback: Optional[Callable] = None):
self.progress_callback = progress_callback
- self._original_hf_hub_download = None
- self._original_snapshot_download = None
+ self._original_tqdm_class = None
self._lock = threading.Lock()
self._total_downloaded = 0
self._total_size = 0
self._file_sizes = {} # Track sizes of individual files
self._file_downloaded = {} # Track downloaded bytes per file
self._current_filename = ""
+ self._active_tqdms = {} # Track active tqdm instances
- def _tracked_hf_hub_download(self, *args, **kwargs):
- """Wrapper for hf_hub_download with progress tracking."""
- import huggingface_hub
+ def _create_tracked_tqdm_class(self):
+ """Create a tqdm subclass that tracks progress."""
+ tracker = self
+ original_tqdm = self._original_tqdm_class
- # Get original callback if present
- original_resume_callback = kwargs.get("resume_download", None)
-
- # Extract filename if available
- filename = kwargs.get("filename", "")
- if not filename and len(args) > 1:
- filename = args[1] if isinstance(args[1], str) else ""
-
- with self._lock:
- self._current_filename = filename
-
- def combined_callback(downloaded: int, total: int):
- """Combined callback that tracks progress."""
- # Update per-file tracking
- with self._lock:
- if filename:
- self._file_sizes[filename] = total
- self._file_downloaded[filename] = downloaded
+ class TrackedTqdm(original_tqdm):
+ """A tqdm subclass that reports progress to our tracker."""
+
+ def __init__(self, *args, **kwargs):
+ # Extract filename from desc before passing to parent
+ desc = kwargs.get("desc", "")
+ if not desc and args:
+ first_arg = args[0]
+ if isinstance(first_arg, str):
+ desc = first_arg
- # Calculate totals across all files
- self._total_size = sum(self._file_sizes.values())
- self._total_downloaded = sum(self._file_downloaded.values())
+ filename = ""
+ if desc:
+ # Try to extract filename from description
+ # HuggingFace Hub uses format like "model.safetensors: 0%|..."
+ if ":" in desc:
+ filename = desc.split(":")[0].strip()
+ else:
+ filename = desc.strip()
+
+ # Filter out non-standard kwargs that huggingface_hub might pass
+ # These are custom kwargs that tqdm doesn't understand
+ filtered_kwargs = {}
+ # Known tqdm kwargs - pass these through
+ tqdm_kwargs = {
+ 'iterable', 'desc', 'total', 'leave', 'file', 'ncols', 'mininterval',
+ 'maxinterval', 'miniters', 'ascii', 'disable', 'unit', 'unit_scale',
+ 'dynamic_ncols', 'smoothing', 'bar_format', 'initial', 'position',
+ 'postfix', 'unit_divisor', 'write_bytes', 'lock_args', 'nrows',
+ 'colour', 'color', 'delay', 'gui', 'disable_default', 'pos'
+ }
+ for key, value in kwargs.items():
+ if key in tqdm_kwargs:
+ filtered_kwargs[key] = value
+
+ # Try to initialize with filtered kwargs, fall back to all kwargs if that fails
+ try:
+ super().__init__(*args, **filtered_kwargs)
+ except TypeError:
+ # If filtering failed, try with all kwargs (maybe tqdm version accepts them)
+ super().__init__(*args, **kwargs)
+
+ self._tracker_filename = filename or "unknown"
+
+ with tracker._lock:
+ if filename:
+ tracker._current_filename = filename
+ tracker._active_tqdms[id(self)] = {
+ "filename": self._tracker_filename,
+ }
- # Call original callback if present
- if original_resume_callback:
- original_resume_callback(downloaded, total)
+ def update(self, n=1):
+ result = super().update(n)
+
+ # Report progress
+ with tracker._lock:
+ if id(self) in tracker._active_tqdms:
+ filename = tracker._active_tqdms[id(self)]["filename"]
+ current = getattr(self, "n", 0)
+ total = getattr(self, "total", 0)
+
+ if total and total > 0:
+ # Update per-file tracking
+ tracker._file_sizes[filename] = total
+ tracker._file_downloaded[filename] = current
+
+ # Calculate totals across all files
+ tracker._total_size = sum(tracker._file_sizes.values())
+ tracker._total_downloaded = sum(tracker._file_downloaded.values())
+
+ # Call progress callback
+ if tracker.progress_callback:
+ tracker.progress_callback(
+ tracker._total_downloaded,
+ tracker._total_size,
+ filename
+ )
+
+ return result
- # Call our progress callback
- if self.progress_callback:
- with self._lock:
- # Pass filename for better progress display
- self.progress_callback(self._total_downloaded, self._total_size, filename)
+ def close(self):
+ with tracker._lock:
+ if id(self) in tracker._active_tqdms:
+ del tracker._active_tqdms[id(self)]
+ return super().close()
- # Replace callback
- kwargs["resume_download"] = combined_callback
-
- # Call original download
- return self._original_hf_hub_download(*args, **kwargs)
-
- def _tracked_snapshot_download(self, *args, **kwargs):
- """Wrapper for snapshot_download with progress tracking."""
- import huggingface_hub
-
- # snapshot_download also uses resume_download callback
- original_resume_callback = kwargs.get("resume_download", None)
-
- def combined_callback(downloaded: int, total: int):
- """Combined callback that tracks progress."""
- with self._lock:
- # For snapshot_download, we track overall progress
- if total > 0:
- self._total_size = max(self._total_size, total)
- self._total_downloaded = downloaded
-
- # Call original callback if present
- if original_resume_callback:
- original_resume_callback(downloaded, total)
-
- # Call our progress callback
- if self.progress_callback:
- with self._lock:
- self.progress_callback(self._total_downloaded, self._total_size, "")
-
- # Replace callback
- kwargs["resume_download"] = combined_callback
-
- # Call original download
- return self._original_snapshot_download(*args, **kwargs)
-
- def _tracked_tqdm_update(self, n=1):
- """Track tqdm updates for progress."""
- if self._original_tqdm:
- # Get current tqdm instance
- import tqdm
- # Try to get progress info from tqdm
- # This is a fallback if hf_hub_download callback doesn't work
- pass
+ return TrackedTqdm
@contextmanager
def patch_download(self):
- """Context manager to patch hf_hub_download and snapshot_download for progress tracking."""
+ """Context manager to patch tqdm for progress tracking."""
try:
- import huggingface_hub
- self._original_hf_hub_download = huggingface_hub.hf_hub_download
+ import tqdm as tqdm_module
- # Also patch snapshot_download if available (used by from_pretrained)
- try:
- self._original_snapshot_download = huggingface_hub.snapshot_download
- except AttributeError:
- self._original_snapshot_download = None
+ # Store original tqdm class
+ self._original_tqdm_class = tqdm_module.tqdm
# Reset totals
with self._lock:
@@ -123,29 +131,62 @@ class HFProgressTracker:
self._file_sizes = {}
self._file_downloaded = {}
self._current_filename = ""
+ self._active_tqdms = {}
- # Patch the functions
- huggingface_hub.hf_hub_download = self._tracked_hf_hub_download
- if self._original_snapshot_download:
- huggingface_hub.snapshot_download = self._tracked_snapshot_download
+ # Create our tracked tqdm class
+ tracked_tqdm = self._create_tracked_tqdm_class()
+
+ # Patch tqdm.tqdm
+ tqdm_module.tqdm = tracked_tqdm
+
+ # Also patch tqdm.auto.tqdm if it exists (used by huggingface_hub)
+ self._original_tqdm_auto = None
+ if hasattr(tqdm_module, "auto") and hasattr(tqdm_module.auto, "tqdm"):
+ self._original_tqdm_auto = tqdm_module.auto.tqdm
+ tqdm_module.auto.tqdm = tracked_tqdm
+
+ # Patch in sys.modules to catch already-imported references
+ self._patched_modules = {}
+ for module_name in list(sys.modules.keys()):
+ if "huggingface" in module_name or module_name.startswith("tqdm"):
+ try:
+ module = sys.modules[module_name]
+ if hasattr(module, "tqdm"):
+ attr = getattr(module, "tqdm")
+ # Only patch if it's the original tqdm class (not already patched)
+ if attr is self._original_tqdm_class or (
+ hasattr(attr, "__name__") and attr.__name__ == "tqdm"
+ ):
+ self._patched_modules[module_name] = attr
+ setattr(module, "tqdm", tracked_tqdm)
+ except (AttributeError, TypeError):
+ pass
yield
+
except ImportError:
- # If huggingface_hub not available, just yield without patching
+ # If tqdm not available, just yield without patching
yield
finally:
- # Restore original functions
- if self._original_hf_hub_download:
+ # Restore original tqdm
+ if self._original_tqdm_class:
try:
- import huggingface_hub
- huggingface_hub.hf_hub_download = self._original_hf_hub_download
- except ImportError:
- pass
-
- if self._original_snapshot_download:
- try:
- import huggingface_hub
- huggingface_hub.snapshot_download = self._original_snapshot_download
+ import tqdm as tqdm_module
+ tqdm_module.tqdm = self._original_tqdm_class
+
+ if self._original_tqdm_auto:
+ tqdm_module.auto.tqdm = self._original_tqdm_auto
+
+ # Restore patched modules
+ for module_name, original in self._patched_modules.items():
+ try:
+ module = sys.modules.get(module_name)
+ if module and original:
+ setattr(module, "tqdm", original)
+ except (AttributeError, TypeError):
+ pass
+ self._patched_modules = {}
+
except (ImportError, AttributeError):
pass
diff --git a/backend/utils/progress.py b/backend/utils/progress.py
index 68360c16..879ba6a5 100644
--- a/backend/utils/progress.py
+++ b/backend/utils/progress.py
@@ -2,7 +2,7 @@
Progress tracking for model downloads using Server-Sent Events.
"""
-from typing import Optional, Callable, Dict
+from typing import Optional, Callable, Dict, List
from fastapi.responses import StreamingResponse
import asyncio
import json
@@ -58,6 +58,15 @@ class ProgressManager:
"""Get current progress for a model."""
return self._progress.get(model_name)
+ def get_all_active(self) -> List[Dict]:
+ """Get all active downloads (status is 'downloading' or 'extracting')."""
+ active = []
+ for model_name, progress in self._progress.items():
+ status = progress.get("status", "")
+ if status in ("downloading", "extracting"):
+ active.append(progress.copy())
+ return active
+
def create_progress_callback(self, model_name: str, filename: Optional[str] = None):
"""
Create a progress callback function for HuggingFace downloads.
diff --git a/backend/utils/tasks.py b/backend/utils/tasks.py
new file mode 100644
index 00000000..05b8e019
--- /dev/null
+++ b/backend/utils/tasks.py
@@ -0,0 +1,93 @@
+"""
+Task tracking for active downloads and generations.
+"""
+
+from typing import Optional, Dict, List
+from datetime import datetime
+from dataclasses import dataclass, field
+
+
+@dataclass
+class DownloadTask:
+ """Represents an active download task."""
+ model_name: str
+ status: str = "downloading" # downloading, extracting, complete, error
+ started_at: datetime = field(default_factory=datetime.utcnow)
+ error: Optional[str] = None
+
+
+@dataclass
+class GenerationTask:
+ """Represents an active generation task."""
+ task_id: str
+ profile_id: str
+ text_preview: str # First 50 chars of text
+ started_at: datetime = field(default_factory=datetime.utcnow)
+
+
+class TaskManager:
+ """Manages active downloads and generations."""
+
+ def __init__(self):
+ self._active_downloads: Dict[str, DownloadTask] = {}
+ self._active_generations: Dict[str, GenerationTask] = {}
+
+ def start_download(self, model_name: str) -> None:
+ """Mark a download as started."""
+ self._active_downloads[model_name] = DownloadTask(
+ model_name=model_name,
+ status="downloading",
+ )
+
+ def complete_download(self, model_name: str) -> None:
+ """Mark a download as complete."""
+ if model_name in self._active_downloads:
+ del self._active_downloads[model_name]
+
+ def error_download(self, model_name: str, error: str) -> None:
+ """Mark a download as failed."""
+ if model_name in self._active_downloads:
+ self._active_downloads[model_name].status = "error"
+ self._active_downloads[model_name].error = error
+
+ def start_generation(self, task_id: str, profile_id: str, text: str) -> None:
+ """Mark a generation as started."""
+ text_preview = text[:50] + "..." if len(text) > 50 else text
+ self._active_generations[task_id] = GenerationTask(
+ task_id=task_id,
+ profile_id=profile_id,
+ text_preview=text_preview,
+ )
+
+ def complete_generation(self, task_id: str) -> None:
+ """Mark a generation as complete."""
+ if task_id in self._active_generations:
+ del self._active_generations[task_id]
+
+ def get_active_downloads(self) -> List[DownloadTask]:
+ """Get all active downloads."""
+ return list(self._active_downloads.values())
+
+ def get_active_generations(self) -> List[GenerationTask]:
+ """Get all active generations."""
+ return list(self._active_generations.values())
+
+ def is_download_active(self, model_name: str) -> bool:
+ """Check if a download is active."""
+ return model_name in self._active_downloads
+
+ def is_generation_active(self, task_id: str) -> bool:
+ """Check if a generation is active."""
+ return task_id in self._active_generations
+
+
+# Global task manager instance
+_task_manager: Optional[TaskManager] = None
+
+
+def get_task_manager() -> TaskManager:
+ """Get or create the global task manager."""
+ global _task_manager
+ if _task_manager is None:
+ _task_manager = TaskManager()
+ return _task_manager
diff --git a/backend/utils/validation.py b/backend/utils/validation.py
index e5b4e0cf..3637da86 100644
--- a/backend/utils/validation.py
+++ b/backend/utils/validation.py
@@ -29,17 +29,20 @@ def validate_text(text: str, max_length: int = 5000) -> Tuple[bool, Optional[str
def validate_language(language: str) -> Tuple[bool, Optional[str]]:
"""
Validate language code.
-
+
+ Supported languages for Qwen3-TTS:
+ Chinese, English, Japanese, Korean, German, French, Russian, Portuguese, Spanish, Italian
+
Args:
language: Language code
-
+
Returns:
Tuple of (is_valid, error_message)
"""
- valid_languages = ["en", "zh"]
+ valid_languages = ["zh", "en", "ja", "ko", "de", "fr", "ru", "pt", "es", "it"]
if language not in valid_languages:
return False, f"Invalid language (must be one of: {', '.join(valid_languages)})"
-
+
return True, None
diff --git a/tauri/src-tauri/gen/Assets.car b/tauri/src-tauri/gen/Assets.car
index 0d31b441..e5667d52 100644
Binary files a/tauri/src-tauri/gen/Assets.car and b/tauri/src-tauri/gen/Assets.car differ