mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-09-18 14:20:42 -07:00
/transcribe wrote every upload to a temp file named .wav regardless of its real format. librosa picks its decoder from the extension, so any non-wav upload failed with "could not open/decode file" even though the format is one the app handles elsewhere. profiles.py already solves this for voice samples by keeping the uploaded extension when it is one of the audio types it accepts, and falling back to .wav otherwise. Same approach here, same set. The fallback means an unknown or missing extension behaves exactly as it does today.
92 lines
3.2 KiB
Python
92 lines
3.2 KiB
Python
"""Transcription endpoints."""
|
|
|
|
import asyncio
|
|
import tempfile
|
|
from pathlib import Path
|
|
|
|
from fastapi import APIRouter, File, Form, HTTPException, UploadFile
|
|
|
|
from .. import models
|
|
from ..services import transcribe
|
|
from ..services.task_queue import create_background_task
|
|
from ..utils.tasks import get_task_manager
|
|
|
|
router = APIRouter()
|
|
|
|
UPLOAD_CHUNK_SIZE = 1024 * 1024 # 1MB
|
|
|
|
# Same set profiles.py accepts for voice samples. librosa picks its decoder from the
|
|
# file extension, so the temp file has to keep the uploaded one.
|
|
ALLOWED_AUDIO_EXTS = {".wav", ".mp3", ".m4a", ".ogg", ".flac", ".aac", ".webm", ".opus"}
|
|
|
|
|
|
@router.post("/transcribe", response_model=models.TranscriptionResponse)
|
|
async def transcribe_audio(
|
|
file: UploadFile = File(...),
|
|
language: str | None = Form(None),
|
|
model: str | None = Form(None),
|
|
):
|
|
"""Transcribe audio file to text."""
|
|
uploaded_ext = Path(file.filename or "").suffix.lower()
|
|
file_suffix = uploaded_ext if uploaded_ext in ALLOWED_AUDIO_EXTS else ".wav"
|
|
|
|
with tempfile.NamedTemporaryFile(suffix=file_suffix, delete=False) as tmp:
|
|
while chunk := await file.read(UPLOAD_CHUNK_SIZE):
|
|
tmp.write(chunk)
|
|
tmp_path = tmp.name
|
|
|
|
try:
|
|
from ..utils.audio import load_audio
|
|
from ..backends import WHISPER_HF_REPOS
|
|
|
|
audio, sr = await asyncio.to_thread(load_audio, tmp_path)
|
|
duration = len(audio) / sr
|
|
|
|
whisper_model = transcribe.get_whisper_model()
|
|
model_size = model if model else whisper_model.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}"
|
|
task_manager = get_task_manager()
|
|
|
|
async def download_whisper_background():
|
|
try:
|
|
await whisper_model.load_model_async(model_size)
|
|
task_manager.complete_download(progress_model_name)
|
|
except Exception as e:
|
|
task_manager.error_download(progress_model_name, str(e))
|
|
|
|
task_manager.start_download(progress_model_name)
|
|
create_background_task(download_whisper_background())
|
|
|
|
raise HTTPException(
|
|
status_code=202,
|
|
detail={
|
|
"message": f"Whisper model {model_size} is being downloaded. Please wait and try again.",
|
|
"model_name": progress_model_name,
|
|
"downloading": True,
|
|
},
|
|
)
|
|
|
|
text = await whisper_model.transcribe(tmp_path, language, model_size)
|
|
|
|
return models.TranscriptionResponse(
|
|
text=text,
|
|
duration=duration,
|
|
)
|
|
|
|
except HTTPException:
|
|
raise
|
|
except Exception as e:
|
|
raise HTTPException(status_code=500, detail=str(e))
|
|
finally:
|
|
Path(tmp_path).unlink(missing_ok=True)
|