Compare commits

...
Author SHA1 Message Date
Jamie Pine df50b8a925 add --force-reinstall --no-deps to torchaudio CUDA install 2026-03-17 00:07:45 -07:00
Jamie Pine a672ac5279 remove voicebox.icns from tracking and add to gitignore 2026-03-17 00:07:19 -07:00
Jamie Pine d35e6f0cc5 fix sample upload blocking the event loop and causing server timeouts
Move audio validation and saving to thread pool so librosa/ffmpeg decoding
doesn't block the async event loop. Combine validate + load into a single
pass to avoid decoding the file twice. Add 50 MB upload limit and chunked
reads to prevent unbounded memory allocation.

Closes #278
2026-03-16 23:29:18 -07:00
Jamie Pine b1069b4521 upgrade CUDA backend build from cu121 to cu126
cu121 only ships kernels up to SM 9.0 (Ada Lovelace). RTX 50-series
(Blackwell, SM 12.0) and RTX 6000 Pro need cu126 which includes SM 12.0
support while remaining backward compatible with older GPUs.

Closes #289
2026-03-16 23:23:55 -07:00
Jamie Pine f9e1aa153d handle client disconnects in SSE and streaming endpoints
Wrap SSE generators with BrokenPipeError/ConnectionResetError handling
so client disconnects during generation status polling, download progress,
or audio streaming don't produce unhandled Errno 32 errors.

Closes #248
2026-03-16 23:17:09 -07:00
Jamie Pine 01800f196f upgrade pip before installing deps in Docker build
Fixes hash mismatch when pip resolves Qwen3-TTS transitive deps.

Closes #286
2026-03-16 23:13:50 -07:00
Jamie Pine 606da1c894 fix generation list not updating on completion
Use refetchQueries instead of invalidateQueries for more reliable history
refresh. Add history refetch to SSE onerror handler so dropped connections
don't leave the list stale. Reset page to 0 in HistoryTable when a pending
generation completes.

Closes #231
2026-03-16 22:54:11 -07:00
Jamie Pine 664178f0cf fix error detail serialization producing [object Object] in error messages
Closes #290
2026-03-16 22:47:01 -07:00
Jamie Pine f1541701fb add model selection and expanded language support to /transcribe endpoint
Closes #233
2026-03-16 22:44:28 -07:00
Jamie PineandGitHub a2adc3b506 Merge pull request #293 from jamiepine/fix/audio-player-freeze
Fix audio player freezing and improve UX
2026-03-16 22:28:39 -07:00
Jamie PineandGitHub 15ba824472 Merge pull request #294 from jamiepine/feat/settings-overhaul
Settings overhaul: routed sub-tabs, server logs, changelog, about page
2026-03-16 13:07:39 -07:00
Jamie Pine 2ad4776a76 fix review feedback: restart race, listener cleanup, stable keys, accessibility 2026-03-16 13:06:57 -07:00
Jamie Pine a8469b39f1 fix about license 2026-03-16 12:31:33 -07:00
29 changed files with 243 additions and 117 deletions
+3 -3
View File
@@ -189,10 +189,10 @@ jobs:
pip install -r backend/requirements.txt pip install -r backend/requirements.txt
pip install --no-deps chatterbox-tts pip install --no-deps chatterbox-tts
- name: Install PyTorch with CUDA 12.1 - name: Install PyTorch with CUDA 12.6
run: | run: |
pip install torch --index-url https://download.pytorch.org/whl/cu121 --force-reinstall --no-deps pip install torch --index-url https://download.pytorch.org/whl/cu126 --force-reinstall --no-deps
pip install torchaudio --index-url https://download.pytorch.org/whl/cu121 pip install torchaudio --index-url https://download.pytorch.org/whl/cu126 --force-reinstall --no-deps
- name: Verify CUDA support in torch - name: Verify CUDA support in torch
run: | run: |
+1
View File
@@ -50,6 +50,7 @@ logs/
app/openapi.json app/openapi.json
tauri/src-tauri/binaries/* tauri/src-tauri/binaries/*
tauri/src-tauri/gen/Assets.car tauri/src-tauri/gen/Assets.car
tauri/src-tauri/gen/voicebox.icns
# Temporary # Temporary
tmp/ tmp/
+2
View File
@@ -31,6 +31,8 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
build-essential \ build-essential \
&& rm -rf /var/lib/apt/lists/* && rm -rf /var/lib/apt/lists/*
RUN pip install --no-cache-dir --upgrade pip
COPY backend/requirements.txt . COPY backend/requirements.txt .
RUN pip install --no-cache-dir --prefix=/install -r requirements.txt RUN pip install --no-cache-dir --prefix=/install -r requirements.txt
RUN pip install --no-cache-dir --prefix=/install \ RUN pip install --no-cache-dir --prefix=/install \
+16 -1
View File
@@ -153,7 +153,9 @@ export function HistoryTable() {
} }
}, [historyData, page]); }, [historyData, page]);
// Reset to page 0 when deletions or imports occur // Reset to page 0 when deletions, imports, or generation completions occur
const pendingCount = useGenerationStore((state) => state.pendingGenerationIds.size);
const prevPendingCountRef = useRef(pendingCount);
useEffect(() => { useEffect(() => {
if (deleteGeneration.isSuccess || importGeneration.isSuccess) { if (deleteGeneration.isSuccess || importGeneration.isSuccess) {
setPage(0); setPage(0);
@@ -161,6 +163,19 @@ export function HistoryTable() {
} }
}, [deleteGeneration.isSuccess, importGeneration.isSuccess]); }, [deleteGeneration.isSuccess, importGeneration.isSuccess]);
useEffect(() => {
// A generation finished (pending count decreased) — scroll back to show it
if (
prevPendingCountRef.current > 0 &&
pendingCount < prevPendingCountRef.current &&
page !== 0
) {
setPage(0);
setAllHistory([]);
}
prevPendingCountRef.current = pendingCount;
}, [pendingCount, page]);
// Intersection Observer for infinite scroll // Intersection Observer for infinite scroll
useEffect(() => { useEffect(() => {
const loadMoreEl = loadMoreRef.current; const loadMoreEl = loadMoreRef.current;
+1 -1
View File
@@ -124,7 +124,7 @@ export function AboutPage() {
rel="noopener noreferrer" rel="noopener noreferrer"
className="hover:text-muted-foreground/60 transition-colors" className="hover:text-muted-foreground/60 transition-colors"
> >
BSL 1.1 MIT
</a> </a>
</p> </p>
</FadeIn> </FadeIn>
@@ -57,8 +57,8 @@ function renderMarkdown(md: string): React.ReactNode[] {
} }
elements.push( elements.push(
<ul key={elements.length} className="space-y-1 my-2"> <ul key={elements.length} className="space-y-1 my-2">
{items.map((item) => ( {items.map((item, idx) => (
<li key={item} className="text-sm text-muted-foreground flex gap-2"> <li key={idx} className="text-sm text-muted-foreground flex gap-2">
<span className="text-muted-foreground/50 select-none shrink-0">&bull;</span> <span className="text-muted-foreground/50 select-none shrink-0">&bull;</span>
<span>{inlineMarkdown(item)}</span> <span>{inlineMarkdown(item)}</span>
</li> </li>
@@ -96,9 +96,9 @@ function renderTable(tableLines: string[], keyBase: number): React.ReactNode {
<table className="text-sm w-full"> <table className="text-sm w-full">
<thead> <thead>
<tr className="border-b"> <tr className="border-b">
{headers.map((h) => ( {headers.map((h, hIdx) => (
<th <th
key={h} key={hIdx}
className="text-left py-1.5 pr-4 text-muted-foreground font-medium text-xs" className="text-left py-1.5 pr-4 text-muted-foreground font-medium text-xs"
> >
{inlineMarkdown(h)} {inlineMarkdown(h)}
@@ -107,10 +107,10 @@ function renderTable(tableLines: string[], keyBase: number): React.ReactNode {
</tr> </tr>
</thead> </thead>
<tbody> <tbody>
{rows.map((row) => ( {rows.map((row, rowIdx) => (
<tr key={row.join('|')} className="border-b border-border/50"> <tr key={rowIdx} className="border-b border-border/50">
{row.map((cell) => ( {row.map((cell, cellIdx) => (
<td key={cell} className="py-1.5 pr-4 text-muted-foreground"> <td key={cellIdx} className="py-1.5 pr-4 text-muted-foreground">
{inlineMarkdown(cell)} {inlineMarkdown(cell)}
</td> </td>
))} ))}
@@ -133,6 +133,13 @@ export function GeneralPage() {
setKeepServerRunningOnClose(checked); setKeepServerRunningOnClose(checked);
platform.lifecycle.setKeepServerRunning(checked).catch((error) => { platform.lifecycle.setKeepServerRunning(checked).catch((error) => {
console.error('Failed to sync setting to Rust:', error); console.error('Failed to sync setting to Rust:', error);
setKeepServerRunningOnClose(!checked);
toast({
title: 'Failed to update setting',
description: 'Could not sync setting to backend.',
variant: 'destructive',
});
return;
}); });
toast({ toast({
title: 'Setting updated', title: 'Setting updated',
+28 -37
View File
@@ -171,17 +171,21 @@ export function GpuPage() {
}; };
}, [cudaDownloading, serverUrl, refetchCudaStatus]); }, [cudaDownloading, serverUrl, refetchCudaStatus]);
const clearHealthPolling = useCallback(() => {
if (healthPollRef.current) {
clearInterval(healthPollRef.current);
healthPollRef.current = null;
}
}, []);
const startHealthPolling = useCallback(() => { const startHealthPolling = useCallback(() => {
if (healthPollRef.current) return; clearHealthPolling();
healthPollRef.current = setInterval(async () => { healthPollRef.current = setInterval(async () => {
try { try {
const result = await apiClient.getHealth(); const result = await apiClient.getHealth();
if (result.status === 'healthy') { if (result.status === 'healthy') {
if (healthPollRef.current) { clearHealthPolling();
clearInterval(healthPollRef.current);
healthPollRef.current = null;
}
setRestartPhase('ready'); setRestartPhase('ready');
queryClient.invalidateQueries(); queryClient.invalidateQueries();
setTimeout(() => setRestartPhase('idle'), 2000); setTimeout(() => setRestartPhase('idle'), 2000);
@@ -190,7 +194,23 @@ export function GpuPage() {
// Server still down, keep polling // Server still down, keep polling
} }
}, 1000); }, 1000);
}, [queryClient]); }, [queryClient, clearHealthPolling]);
const restartServerWithPolling = useCallback(
async (errorMessage: string) => {
setRestartPhase('stopping');
try {
await platform.lifecycle.restartServer();
setRestartPhase('waiting');
startHealthPolling();
} catch (e: unknown) {
clearHealthPolling();
setRestartPhase('idle');
throw new Error(e instanceof Error ? e.message : errorMessage);
}
},
[platform, startHealthPolling, clearHealthPolling],
);
const handleDownload = async () => { const handleDownload = async () => {
setError(null); setError(null);
@@ -209,24 +229,9 @@ export function GpuPage() {
const handleRestart = async () => { const handleRestart = async () => {
setError(null); setError(null);
setRestartPhase('stopping');
try { try {
setRestartPhase('waiting'); await restartServerWithPolling('Restart failed');
startHealthPolling();
await platform.lifecycle.restartServer();
if (healthPollRef.current) {
clearInterval(healthPollRef.current);
healthPollRef.current = null;
}
setRestartPhase('ready');
queryClient.invalidateQueries();
setTimeout(() => setRestartPhase('idle'), 2000);
} catch (e: unknown) { } catch (e: unknown) {
setRestartPhase('idle');
if (healthPollRef.current) {
clearInterval(healthPollRef.current);
healthPollRef.current = null;
}
setError(e instanceof Error ? e.message : 'Restart failed'); setError(e instanceof Error ? e.message : 'Restart failed');
} }
}; };
@@ -236,22 +241,8 @@ export function GpuPage() {
setRestartPhase('stopping'); setRestartPhase('stopping');
try { try {
await apiClient.deleteCudaBackend(); await apiClient.deleteCudaBackend();
setRestartPhase('waiting'); await restartServerWithPolling('Failed to switch to CPU');
startHealthPolling();
await platform.lifecycle.restartServer();
if (healthPollRef.current) {
clearInterval(healthPollRef.current);
healthPollRef.current = null;
}
setRestartPhase('ready');
queryClient.invalidateQueries();
setTimeout(() => setRestartPhase('idle'), 2000);
} catch (e: unknown) { } catch (e: unknown) {
setRestartPhase('idle');
if (healthPollRef.current) {
clearInterval(healthPollRef.current);
healthPollRef.current = null;
}
setError(e instanceof Error ? e.message : 'Failed to switch to CPU'); setError(e instanceof Error ? e.message : 'Failed to switch to CPU');
refetchCudaStatus(); refetchCudaStatus();
} }
+1 -1
View File
@@ -96,7 +96,7 @@ export function LogsPage() {
)} )}
</div> </div>
) : ( ) : (
entries.map((entry, i) => <LogLine key={`${entry.timestamp}-${i}`} entry={entry} />) entries.map((entry) => <LogLine key={entry.id} entry={entry} />)
)} )}
</div> </div>
</div> </div>
+1 -1
View File
@@ -48,7 +48,7 @@ export function SettingRow({
<div className="min-w-0"> <div className="min-w-0">
<label <label
htmlFor={htmlFor} htmlFor={htmlFor}
className="text-sm font-medium leading-none cursor-pointer select-none" className={`text-sm font-medium leading-none select-none ${htmlFor ? 'cursor-pointer' : ''}`}
> >
{title} {title}
</label> </label>
+1 -1
View File
@@ -25,7 +25,7 @@ export function Toaster() {
<ToastClose /> <ToastClose />
</Toast> </Toast>
))} ))}
<ToastViewport className={isPlayerOpen ? 'sm:bottom-32' : ''} /> <ToastViewport className={isPlayerOpen ? 'sm:bottom-44' : ''} />
</ToastProvider> </ToastProvider>
); );
} }
+1
View File
@@ -26,6 +26,7 @@ const Toggle = React.forwardRef<HTMLButtonElement, ToggleProps>(
}} }}
className={cn( className={cn(
'relative inline-flex h-5 w-9 shrink-0 items-center rounded-full transition-colors', 'relative inline-flex h-5 w-9 shrink-0 items-center rounded-full transition-colors',
'focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-ring focus-visible:ring-offset-2',
checked ? 'bg-accent' : 'bg-muted-foreground/25', checked ? 'bg-accent' : 'bg-muted-foreground/25',
disabled ? 'opacity-50 cursor-not-allowed' : 'cursor-pointer', disabled ? 'opacity-50 cursor-not-allowed' : 'cursor-pointer',
className, className,
+35 -12
View File
@@ -32,8 +32,24 @@ import type {
TranscriptionResponse, TranscriptionResponse,
VoiceProfileCreate, VoiceProfileCreate,
VoiceProfileResponse, VoiceProfileResponse,
WhisperModelSize,
} from './types'; } from './types';
function formatErrorDetail(detail: unknown, fallback: string): string {
if (typeof detail === 'string') return detail;
if (Array.isArray(detail)) {
return detail
.map((e: Record<string, unknown>) => e.msg || e.message || JSON.stringify(e))
.join('; ');
}
if (detail && typeof detail === 'object') {
const obj = detail as Record<string, unknown>;
if (typeof obj.message === 'string') return obj.message;
return JSON.stringify(detail);
}
return fallback;
}
class ApiClient { class ApiClient {
private getBaseUrl(): string { private getBaseUrl(): string {
const serverUrl = useServerStore.getState().serverUrl; const serverUrl = useServerStore.getState().serverUrl;
@@ -54,7 +70,7 @@ class ApiClient {
const error = await response.json().catch(() => ({ const error = await response.json().catch(() => ({
detail: response.statusText, detail: response.statusText,
})); }));
throw new Error(error.detail || `HTTP error! status: ${response.status}`); throw new Error(formatErrorDetail(error.detail, `HTTP error! status: ${response.status}`));
} }
return response.json(); return response.json();
@@ -113,7 +129,7 @@ class ApiClient {
const error = await response.json().catch(() => ({ const error = await response.json().catch(() => ({
detail: response.statusText, detail: response.statusText,
})); }));
throw new Error(error.detail || `HTTP error! status: ${response.status}`); throw new Error(formatErrorDetail(error.detail, `HTTP error! status: ${response.status}`));
} }
return response.json(); return response.json();
@@ -147,7 +163,7 @@ class ApiClient {
const error = await response.json().catch(() => ({ const error = await response.json().catch(() => ({
detail: response.statusText, detail: response.statusText,
})); }));
throw new Error(error.detail || `HTTP error! status: ${response.status}`); throw new Error(formatErrorDetail(error.detail, `HTTP error! status: ${response.status}`));
} }
return response.blob(); return response.blob();
@@ -167,7 +183,7 @@ class ApiClient {
const error = await response.json().catch(() => ({ const error = await response.json().catch(() => ({
detail: response.statusText, detail: response.statusText,
})); }));
throw new Error(error.detail || `HTTP error! status: ${response.status}`); throw new Error(formatErrorDetail(error.detail, `HTTP error! status: ${response.status}`));
} }
return response.json(); return response.json();
@@ -187,7 +203,7 @@ class ApiClient {
const error = await response.json().catch(() => ({ const error = await response.json().catch(() => ({
detail: response.statusText, detail: response.statusText,
})); }));
throw new Error(error.detail || `HTTP error! status: ${response.status}`); throw new Error(formatErrorDetail(error.detail, `HTTP error! status: ${response.status}`));
} }
return response.json(); return response.json();
@@ -257,7 +273,7 @@ class ApiClient {
const error = await response.json().catch(() => ({ const error = await response.json().catch(() => ({
detail: response.statusText, detail: response.statusText,
})); }));
throw new Error(error.detail || `HTTP error! status: ${response.status}`); throw new Error(formatErrorDetail(error.detail, `HTTP error! status: ${response.status}`));
} }
return response.blob(); return response.blob();
@@ -271,7 +287,7 @@ class ApiClient {
const error = await response.json().catch(() => ({ const error = await response.json().catch(() => ({
detail: response.statusText, detail: response.statusText,
})); }));
throw new Error(error.detail || `HTTP error! status: ${response.status}`); throw new Error(formatErrorDetail(error.detail, `HTTP error! status: ${response.status}`));
} }
return response.blob(); return response.blob();
@@ -297,7 +313,7 @@ class ApiClient {
const error = await response.json().catch(() => ({ const error = await response.json().catch(() => ({
detail: response.statusText, detail: response.statusText,
})); }));
throw new Error(error.detail || `HTTP error! status: ${response.status}`); throw new Error(formatErrorDetail(error.detail, `HTTP error! status: ${response.status}`));
} }
return response.json(); return response.json();
@@ -318,12 +334,19 @@ class ApiClient {
} }
// Transcription // Transcription
async transcribeAudio(file: File, language?: LanguageCode): Promise<TranscriptionResponse> { async transcribeAudio(
file: File,
language?: LanguageCode,
model?: WhisperModelSize,
): Promise<TranscriptionResponse> {
const formData = new FormData(); const formData = new FormData();
formData.append('file', file); formData.append('file', file);
if (language) { if (language) {
formData.append('language', language); formData.append('language', language);
} }
if (model) {
formData.append('model', model);
}
const url = `${this.getBaseUrl()}/transcribe`; const url = `${this.getBaseUrl()}/transcribe`;
const response = await fetch(url, { const response = await fetch(url, {
@@ -335,7 +358,7 @@ class ApiClient {
const error = await response.json().catch(() => ({ const error = await response.json().catch(() => ({
detail: response.statusText, detail: response.statusText,
})); }));
throw new Error(error.detail || `HTTP error! status: ${response.status}`); throw new Error(formatErrorDetail(error.detail, `HTTP error! status: ${response.status}`));
} }
return response.json(); return response.json();
@@ -608,7 +631,7 @@ class ApiClient {
const error = await response.json().catch(() => ({ const error = await response.json().catch(() => ({
detail: response.statusText, detail: response.statusText,
})); }));
throw new Error(error.detail || `HTTP error! status: ${response.status}`); throw new Error(formatErrorDetail(error.detail, `HTTP error! status: ${response.status}`));
} }
return response.blob(); return response.blob();
@@ -705,7 +728,7 @@ class ApiClient {
const error = await response.json().catch(() => ({ const error = await response.json().catch(() => ({
detail: response.statusText, detail: response.statusText,
})); }));
throw new Error(error.detail || `HTTP error! status: ${response.status}`); throw new Error(formatErrorDetail(error.detail, `HTTP error! status: ${response.status}`));
} }
return response.blob(); return response.blob();
+3
View File
@@ -99,8 +99,11 @@ export interface HistoryListResponse {
total: number; total: number;
} }
export type WhisperModelSize = 'base' | 'small' | 'medium' | 'large' | 'turbo';
export interface TranscriptionRequest { export interface TranscriptionRequest {
language?: LanguageCode; language?: LanguageCode;
model?: WhisperModelSize;
} }
export interface TranscriptionResponse { export interface TranscriptionResponse {
+6 -5
View File
@@ -75,8 +75,8 @@ export function useGenerationProgress() {
currentSources.delete(id); currentSources.delete(id);
removePendingGeneration(id); removePendingGeneration(id);
// Refresh history to pick up the completed generation // Refetch history to pick up the completed generation
queryClient.invalidateQueries({ queryKey: ['history'] }); queryClient.refetchQueries({ queryKey: ['history'] });
// If this generation was queued for a story, add it now // If this generation was queued for a story, add it now
const storyId = removePendingStoryAdd(id); const storyId = removePendingStoryAdd(id);
@@ -120,7 +120,7 @@ export function useGenerationProgress() {
removePendingGeneration(id); removePendingGeneration(id);
removePendingStoryAdd(id); removePendingStoryAdd(id);
queryClient.invalidateQueries({ queryKey: ['history'] }); queryClient.refetchQueries({ queryKey: ['history'] });
toast({ toast({
title: data.status === 'not_found' ? 'Generation not found' : 'Generation failed', title: data.status === 'not_found' ? 'Generation not found' : 'Generation failed',
@@ -134,11 +134,12 @@ export function useGenerationProgress() {
}; };
source.onerror = () => { source.onerror = () => {
// EventSource auto-reconnects, but if we get repeated errors // SSE connection dropped — clean up and refresh history so any
// just clean up // completed/failed generation still appears in the list
source.close(); source.close();
currentSources.delete(id); currentSources.delete(id);
removePendingGeneration(id); removePendingGeneration(id);
queryClient.refetchQueries({ queryKey: ['history'] });
}; };
currentSources.set(id, source); currentSources.set(id, source);
+10 -2
View File
@@ -1,10 +1,18 @@
import { useMutation } from '@tanstack/react-query'; import { useMutation } from '@tanstack/react-query';
import { apiClient } from '@/lib/api/client'; import { apiClient } from '@/lib/api/client';
import type { WhisperModelSize } from '@/lib/api/types';
import type { LanguageCode } from '@/lib/constants/languages'; import type { LanguageCode } from '@/lib/constants/languages';
export function useTranscription() { export function useTranscription() {
return useMutation({ return useMutation({
mutationFn: ({ file, language }: { file: File; language?: LanguageCode }) => mutationFn: ({
apiClient.transcribeAudio(file, language), file,
language,
model,
}: {
file: File;
language?: LanguageCode;
model?: WhisperModelSize;
}) => apiClient.transcribeAudio(file, language, model),
}); });
} }
+4 -2
View File
@@ -3,7 +3,10 @@ import type { ServerLogEntry } from '@/platform/types';
const MAX_LOG_ENTRIES = 2000; const MAX_LOG_ENTRIES = 2000;
let nextLogEntryId = 0;
export interface LogEntry extends ServerLogEntry { export interface LogEntry extends ServerLogEntry {
id: number;
timestamp: number; timestamp: number;
} }
@@ -17,9 +20,8 @@ export const useLogStore = create<LogStore>((set) => ({
entries: [], entries: [],
addEntry: (entry) => addEntry: (entry) =>
set((state) => { set((state) => {
const newEntry: LogEntry = { ...entry, timestamp: Date.now() }; const newEntry: LogEntry = { ...entry, id: nextLogEntryId++, timestamp: Date.now() };
const entries = [...state.entries, newEntry]; const entries = [...state.entries, newEntry];
// Cap buffer size
if (entries.length > MAX_LOG_ENTRIES) { if (entries.length > MAX_LOG_ENTRIES) {
return { entries: entries.slice(entries.length - MAX_LOG_ENTRIES) }; return { entries: entries.slice(entries.length - MAX_LOG_ENTRIES) };
} }
+1
View File
@@ -134,6 +134,7 @@ class STTBackend(Protocol):
self, self,
audio_path: str, audio_path: str,
language: Optional[str] = None, language: Optional[str] = None,
model_size: Optional[str] = None,
) -> str: ) -> str:
""" """
Transcribe audio to text. Transcribe audio to text.
+4 -2
View File
@@ -345,18 +345,20 @@ class MLXSTTBackend:
self, self,
audio_path: str, audio_path: str,
language: Optional[str] = None, language: Optional[str] = None,
model_size: Optional[str] = None,
) -> str: ) -> str:
""" """
Transcribe audio to text. Transcribe audio to text.
Args: Args:
audio_path: Path to audio file audio_path: Path to audio file
language: Optional language hint (en or zh) language: Optional language hint
model_size: Optional model size override
Returns: Returns:
Transcribed text Transcribed text
""" """
await self.load_model_async(None) await self.load_model_async(model_size)
def _transcribe_sync(): def _transcribe_sync():
"""Run synchronous transcription in thread pool.""" """Run synchronous transcription in thread pool."""
+4 -2
View File
@@ -306,18 +306,20 @@ class PyTorchSTTBackend:
self, self,
audio_path: str, audio_path: str,
language: Optional[str] = None, language: Optional[str] = None,
model_size: Optional[str] = None,
) -> str: ) -> str:
""" """
Transcribe audio to text. Transcribe audio to text.
Args: Args:
audio_path: Path to audio file audio_path: Path to audio file
language: Optional language hint (en or zh) language: Optional language hint
model_size: Optional model size override
Returns: Returns:
Transcribed text Transcribed text
""" """
await self.load_model_async(None) await self.load_model_async(model_size)
def _transcribe_sync(): def _transcribe_sync():
"""Run synchronous transcription in thread pool.""" """Run synchronous transcription in thread pool."""
+2 -1
View File
@@ -149,7 +149,8 @@ class HistoryListResponse(BaseModel):
class TranscriptionRequest(BaseModel): class TranscriptionRequest(BaseModel):
"""Request model for audio transcription.""" """Request model for audio transcription."""
language: Optional[str] = Field(None, pattern="^(en|zh)$") language: Optional[str] = Field(None, pattern="^(en|zh|ja|ko|de|fr|ru|pt|es|it)$")
model: Optional[str] = Field(None, pattern="^(base|small|medium|large|turbo)$")
class TranscriptionResponse(BaseModel): class TranscriptionResponse(BaseModel):
+28 -19
View File
@@ -1,12 +1,15 @@
"""TTS generation endpoints.""" """TTS generation endpoints."""
import asyncio import asyncio
import logging
import uuid import uuid
from fastapi import APIRouter, Depends, HTTPException from fastapi import APIRouter, Depends, HTTPException
from fastapi.responses import StreamingResponse from fastapi.responses import StreamingResponse
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
logger = logging.getLogger(__name__)
from .. import models from .. import models
from ..services import history, profiles, tts from ..services import history, profiles, tts
from ..database import Generation as DBGeneration, VoiceProfile as DBVoiceProfile, get_db from ..database import Generation as DBGeneration, VoiceProfile as DBVoiceProfile, get_db
@@ -181,25 +184,28 @@ async def get_generation_status(generation_id: str, db: Session = Depends(get_db
import json import json
async def event_stream(): async def event_stream():
while True: try:
db.expire_all() while True:
gen = db.query(DBGeneration).filter_by(id=generation_id).first() db.expire_all()
if not gen: gen = db.query(DBGeneration).filter_by(id=generation_id).first()
yield f"data: {json.dumps({'status': 'not_found', 'id': generation_id})}\n\n" if not gen:
return yield f"data: {json.dumps({'status': 'not_found', 'id': generation_id})}\n\n"
return
payload = { payload = {
"id": gen.id, "id": gen.id,
"status": gen.status or "completed", "status": gen.status or "completed",
"duration": gen.duration, "duration": gen.duration,
"error": gen.error, "error": gen.error,
} }
yield f"data: {json.dumps(payload)}\n\n" yield f"data: {json.dumps(payload)}\n\n"
if (gen.status or "completed") in ("completed", "failed"): if (gen.status or "completed") in ("completed", "failed"):
return return
await asyncio.sleep(1) await asyncio.sleep(1)
except (BrokenPipeError, ConnectionResetError, asyncio.CancelledError):
logger.debug("SSE client disconnected for generation %s", generation_id)
return StreamingResponse( return StreamingResponse(
event_stream(), event_stream(),
@@ -265,9 +271,12 @@ async def stream_speech(
wav_bytes = tts.audio_to_wav_bytes(audio, sample_rate) wav_bytes = tts.audio_to_wav_bytes(audio, sample_rate)
async def _wav_stream(): async def _wav_stream():
chunk_size = 64 * 1024 try:
for i in range(0, len(wav_bytes), chunk_size): chunk_size = 64 * 1024
yield wav_bytes[i : i + chunk_size] for i in range(0, len(wav_bytes), chunk_size):
yield wav_bytes[i : i + chunk_size]
except (BrokenPipeError, ConnectionResetError, asyncio.CancelledError):
logger.debug("Client disconnected during audio stream")
return StreamingResponse( return StreamingResponse(
_wav_stream(), _wav_stream(),
+14 -2
View File
@@ -102,6 +102,10 @@ async def delete_profile(
return {"message": "Profile deleted successfully"} return {"message": "Profile deleted successfully"}
SAMPLE_MAX_FILE_SIZE = 50 * 1024 * 1024 # 50 MB
SAMPLE_UPLOAD_CHUNK_SIZE = 1024 * 1024 # 1 MB
@router.post("/profiles/{profile_id}/samples", response_model=models.ProfileSampleResponse) @router.post("/profiles/{profile_id}/samples", response_model=models.ProfileSampleResponse)
async def add_profile_sample( async def add_profile_sample(
profile_id: str, profile_id: str,
@@ -115,8 +119,16 @@ async def add_profile_sample(
file_suffix = _uploaded_ext if _uploaded_ext in _allowed_audio_exts else ".wav" file_suffix = _uploaded_ext if _uploaded_ext in _allowed_audio_exts else ".wav"
with tempfile.NamedTemporaryFile(suffix=file_suffix, delete=False) as tmp: with tempfile.NamedTemporaryFile(suffix=file_suffix, delete=False) as tmp:
content = await file.read() total_size = 0
tmp.write(content) while chunk := await file.read(SAMPLE_UPLOAD_CHUNK_SIZE):
total_size += len(chunk)
if total_size > SAMPLE_MAX_FILE_SIZE:
Path(tmp.name).unlink(missing_ok=True)
raise HTTPException(
status_code=413,
detail=f"File too large (max {SAMPLE_MAX_FILE_SIZE // (1024 * 1024)} MB)",
)
tmp.write(chunk)
tmp_path = tmp.name tmp_path = tmp.name
try: try:
+13 -3
View File
@@ -20,6 +20,7 @@ UPLOAD_CHUNK_SIZE = 1024 * 1024 # 1MB
async def transcribe_audio( async def transcribe_audio(
file: UploadFile = File(...), file: UploadFile = File(...),
language: str | None = Form(None), language: str | None = Form(None),
model: str | None = Form(None),
): ):
"""Transcribe audio file to text.""" """Transcribe audio file to text."""
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp: with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp:
@@ -29,14 +30,23 @@ async def transcribe_audio(
try: try:
from ..utils.audio import load_audio from ..utils.audio import load_audio
from ..backends import WHISPER_HF_REPOS
audio, sr = await asyncio.to_thread(load_audio, tmp_path) audio, sr = await asyncio.to_thread(load_audio, tmp_path)
duration = len(audio) / sr duration = len(audio) / sr
whisper_model = transcribe.get_whisper_model() whisper_model = transcribe.get_whisper_model()
model_size = whisper_model.model_size model_size = model if model else whisper_model.model_size
if not whisper_model.is_loaded() and not whisper_model._is_model_cached(model_size): valid_sizes = list(WHISPER_HF_REPOS.keys())
if model_size not in valid_sizes:
raise HTTPException(
status_code=400,
detail=f"Invalid model size '{model_size}'. Must be one of: {', '.join(valid_sizes)}",
)
already_loaded = whisper_model.is_loaded() and whisper_model.model_size == model_size
if not already_loaded and not whisper_model._is_model_cached(model_size):
progress_model_name = f"whisper-{model_size}" progress_model_name = f"whisper-{model_size}"
task_manager = get_task_manager() task_manager = get_task_manager()
@@ -59,7 +69,7 @@ async def transcribe_audio(
}, },
) )
text = await whisper_model.transcribe(tmp_path, language) text = await whisper_model.transcribe(tmp_path, language, model_size)
return models.TranscriptionResponse( return models.TranscriptionResponse(
text=text, text=text,
+8 -4
View File
@@ -22,7 +22,7 @@ from ..database import (
Generation as DBGeneration, Generation as DBGeneration,
) )
from ..models import EffectConfig from ..models import EffectConfig
from ..utils.audio import validate_reference_audio, load_audio, save_audio from ..utils.audio import validate_reference_audio, validate_and_load_reference_audio, load_audio, save_audio
from ..utils.images import validate_image, process_avatar from ..utils.images import validate_image, process_avatar
from ..utils.cache import _get_cache_dir, clear_profile_cache from ..utils.cache import _get_cache_dir, clear_profile_cache
from .tts import get_tts_model from .tts import get_tts_model
@@ -117,11 +117,16 @@ async def add_profile_sample(
Returns: Returns:
Created sample Created sample
""" """
import asyncio
profile = db.query(DBVoiceProfile).filter_by(id=profile_id).first() profile = db.query(DBVoiceProfile).filter_by(id=profile_id).first()
if not profile: if not profile:
raise ValueError(f"Profile {profile_id} not found") raise ValueError(f"Profile {profile_id} not found")
is_valid, error_msg = validate_reference_audio(audio_path) # Validate and load audio in a single pass, off the event loop
is_valid, error_msg, audio, sr = await asyncio.to_thread(
validate_and_load_reference_audio, audio_path
)
if not is_valid: if not is_valid:
raise ValueError(f"Invalid reference audio: {error_msg}") raise ValueError(f"Invalid reference audio: {error_msg}")
@@ -130,8 +135,7 @@ async def add_profile_sample(
profile_dir.mkdir(parents=True, exist_ok=True) profile_dir.mkdir(parents=True, exist_ok=True)
dest_path = profile_dir / f"{sample_id}.wav" dest_path = profile_dir / f"{sample_id}.wav"
audio, sr = load_audio(audio_path) await asyncio.to_thread(save_audio, audio, str(dest_path), sr)
save_audio(audio, str(dest_path), sr)
db_sample = DBProfileSample( db_sample = DBProfileSample(
id=sample_id, id=sample_id,
+24 -6
View File
@@ -217,22 +217,40 @@ def validate_reference_audio(
Returns: Returns:
Tuple of (is_valid, error_message) Tuple of (is_valid, error_message)
""" """
result = validate_and_load_reference_audio(
audio_path, min_duration, max_duration, min_rms
)
return (result[0], result[1])
def validate_and_load_reference_audio(
audio_path: str,
min_duration: float = 2.0,
max_duration: float = 30.0,
min_rms: float = 0.01,
) -> Tuple[bool, Optional[str], Optional[np.ndarray], Optional[int]]:
"""
Validate and load reference audio in a single pass.
Returns:
Tuple of (is_valid, error_message, audio_array, sample_rate)
"""
try: try:
audio, sr = load_audio(audio_path) audio, sr = load_audio(audio_path)
duration = len(audio) / sr duration = len(audio) / sr
if duration < min_duration: if duration < min_duration:
return False, f"Audio too short (minimum {min_duration} seconds)" return False, f"Audio too short (minimum {min_duration} seconds)", None, None
if duration > max_duration: if duration > max_duration:
return False, f"Audio too long (maximum {max_duration} seconds)" return False, f"Audio too long (maximum {max_duration} seconds)", None, None
rms = np.sqrt(np.mean(audio**2)) rms = np.sqrt(np.mean(audio**2))
if rms < min_rms: if rms < min_rms:
return False, "Audio is too quiet or silent" return False, "Audio is too quiet or silent", None, None
if np.abs(audio).max() > 0.99: if np.abs(audio).max() > 0.99:
return False, "Audio is clipping (reduce input gain)" return False, "Audio is clipping (reduce input gain)", None, None
return True, None return True, None, audio, sr
except Exception as e: except Exception as e:
return False, f"Error validating audio: {str(e)}" return False, f"Error validating audio: {str(e)}", None, None
+2
View File
@@ -246,6 +246,8 @@ class ProgressManager:
# Send heartbeat # Send heartbeat
yield ": heartbeat\n\n" yield ": heartbeat\n\n"
continue continue
except (BrokenPipeError, ConnectionResetError, asyncio.CancelledError):
logger.debug(f"SSE client disconnected from {model_name}")
finally: finally:
# Remove from listeners # Remove from listeners
if model_name in self._listeners: if model_name in self._listeners:
Binary file not shown.
+15 -4
View File
@@ -88,16 +88,27 @@ class TauriLifecycle implements PlatformLifecycle {
} }
subscribeToServerLogs(callback: (entry: ServerLogEntry) => void): () => void { subscribeToServerLogs(callback: (entry: ServerLogEntry) => void): () => void {
let disposed = false;
let unlisten: (() => void) | null = null; let unlisten: (() => void) | null = null;
listen<ServerLogEntry>('server-log', (event) => { void listen<ServerLogEntry>('server-log', (event) => {
callback(event.payload); callback(event.payload);
}).then((fn) => { })
unlisten = fn; .then((fn) => {
}); if (disposed) {
fn();
return;
}
unlisten = fn;
})
.catch((error) => {
console.error('Failed to subscribe to server logs:', error);
});
return () => { return () => {
disposed = true;
unlisten?.(); unlisten?.();
unlisten = null;
}; };
} }
} }