import os import shutil import logging import numpy as np from sqlalchemy import create_engine from sqlalchemy.orm import Session from .celery_app import celery_app from ..config import settings from ..models.tables import Video, Segment, Job, VideoStatus, JobStatus from ..services.video_utils import build_output_path logger = logging.getLogger(__name__) engine = create_engine(settings.database_url.replace("+aiosqlite", "")) def _get_sync_session(): return Session(engine) @celery_app.task(bind=True, name="app.tasks.generate.generate_audio") def generate_audio(self, video_id: int, job_id: int = None): from ..services.tts import synthesize_segment, unload_model as unload_tts from ..services.mixer import mix_audio db = _get_sync_session() temp_files = [] audio_dir = None try: video = db.query(Video).filter(Video.id == video_id).first() if not video: raise ValueError(f"Video {video_id} not found") video.status = VideoStatus.generating.value db.commit() if job_id: job = db.query(Job).filter(Job.id == job_id).first() else: job = Job( video_id=video_id, job_type="generate", status=JobStatus.running.value ) db.add(job) db.commit() job.celery_task_id = self.request.id job.status = JobStatus.running.value db.commit() segments = ( db.query(Segment) .filter(Segment.video_id == video_id) .order_by(Segment.start_time) .all() ) total = len(segments) audio_dir = os.path.join(settings.temp_dir, f"tts_{video_id}") os.makedirs(audio_dir, exist_ok=True) speaker_refs = {} for seg in segments: if seg.speaker_label and seg.speaker_label not in speaker_refs: seg_audio_path = os.path.join( settings.temp_dir, f"ref_{video_id}_{seg.speaker_label.replace(' ', '_')}.wav", ) temp_files.append(seg_audio_path) _extract_segment_audio( video.filepath, seg.start_time, seg.end_time, seg_audio_path ) speaker_refs[seg.speaker_label] = { "audio_path": seg_audio_path, "text": seg.english_text, } logger.info("Starting TTS for %d segments", total) segment_data = [] for i, seg in enumerate(segments): ref = speaker_refs.get(seg.speaker_label, {}) audio_bytes = synthesize_segment( text=seg.danish_text or seg.english_text or "", reference_audio_path=ref.get("audio_path"), reference_text=ref.get("text"), speaker_label=seg.speaker_label, ) seg_audio_path = os.path.join(audio_dir, f"seg_{seg.id}.wav") with open(seg_audio_path, "wb") as f: f.write(audio_bytes) temp_files.append(seg_audio_path) import soundfile as sf danish_audio, sr = sf.read(seg_audio_path) if len(danish_audio.shape) > 1: danish_audio = danish_audio.mean(axis=1) if sr != 24000: import librosa danish_audio = librosa.resample( danish_audio, orig_sr=sr, target_sr=24000 ) segment_data.append( { "start_time": seg.start_time, "end_time": seg.end_time, "danish_audio": danish_audio.astype(np.float32), } ) pct = (i + 1) / total * 0.9 self.update_state(state="PROGRESS", meta={"progress": pct}) job.progress = pct db.commit() logger.info("Mixing audio for video %d", video_id) output_path = build_output_path(video.filepath) mix_audio(video.filepath, segment_data, output_path) video.output_path = output_path video.status = VideoStatus.done.value job.status = JobStatus.done.value job.progress = 1.0 db.commit() unload_tts() logger.info("Generation complete for video %d", video_id) return {"video_id": video_id, "output_path": output_path} except Exception as e: logger.error("Generation failed for video %d: %s", video_id, e, exc_info=True) video = db.query(Video).filter(Video.id == video_id).first() if video: video.status = VideoStatus.error.value video.error_message = str(e) if job_id: job = db.query(Job).filter(Job.id == job_id).first() else: job = ( db.query(Job) .filter(Job.video_id == video_id, Job.job_type == "generate") .first() ) if job: job.status = JobStatus.failed.value job.error_message = str(e) db.commit() raise finally: for f in temp_files: if os.path.exists(f): try: os.remove(f) except OSError: pass if audio_dir and os.path.exists(audio_dir): try: shutil.rmtree(audio_dir) except OSError: pass db.close() def _extract_segment_audio(video_path: str, start: float, end: float, output_path: str): import subprocess duration = end - start cmd = [ "ffmpeg", "-y", "-i", video_path, "-ss", str(start), "-t", str(duration), "-vn", "-acodec", "pcm_s16le", "-ar", "16000", "-ac", "1", output_path, ] subprocess.run(cmd, capture_output=True, check=True)