feat: Celery setup and transcription task — Whisper + FFmpeg pipeline
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,23 @@
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
class TranscriptionSegment(BaseModel):
|
||||
start: float
|
||||
end: float
|
||||
text: str
|
||||
|
||||
|
||||
class TranscriptionResponse(BaseModel):
|
||||
id: uuid.UUID
|
||||
recording_id: uuid.UUID
|
||||
text: str
|
||||
segments: list[TranscriptionSegment]
|
||||
language: str | None
|
||||
whisper_model: str
|
||||
processing_time: float | None
|
||||
created_at: datetime
|
||||
|
||||
model_config = {"from_attributes": True}
|
||||
@@ -0,0 +1,68 @@
|
||||
import os
|
||||
import subprocess
|
||||
import time
|
||||
import uuid
|
||||
|
||||
from app.config import settings
|
||||
|
||||
_model = None
|
||||
|
||||
|
||||
def _get_model():
|
||||
global _model
|
||||
if _model is None:
|
||||
from faster_whisper import WhisperModel
|
||||
_model = WhisperModel(
|
||||
settings.whisper_model,
|
||||
device=settings.whisper_device,
|
||||
compute_type=settings.whisper_compute_type,
|
||||
)
|
||||
return _model
|
||||
|
||||
|
||||
def extract_audio(video_path: str) -> str:
|
||||
audio_path = os.path.join(
|
||||
settings.upload_dir, f"{uuid.uuid4().hex}.wav"
|
||||
)
|
||||
subprocess.run(
|
||||
[
|
||||
"ffmpeg", "-i", video_path,
|
||||
"-vn", "-acodec", "pcm_s16le",
|
||||
"-ar", "16000", "-ac", "1",
|
||||
audio_path, "-y",
|
||||
],
|
||||
check=True,
|
||||
capture_output=True,
|
||||
)
|
||||
return audio_path
|
||||
|
||||
|
||||
def transcribe_audio(audio_path: str) -> dict:
|
||||
model = _get_model()
|
||||
start_time = time.time()
|
||||
|
||||
segments_iter, info = model.transcribe(
|
||||
audio_path,
|
||||
beam_size=5,
|
||||
language=None,
|
||||
vad_filter=True,
|
||||
)
|
||||
|
||||
segments = []
|
||||
full_text_parts = []
|
||||
for segment in segments_iter:
|
||||
segments.append({
|
||||
"start": round(segment.start, 2),
|
||||
"end": round(segment.end, 2),
|
||||
"text": segment.text.strip(),
|
||||
})
|
||||
full_text_parts.append(segment.text.strip())
|
||||
|
||||
processing_time = time.time() - start_time
|
||||
|
||||
return {
|
||||
"text": " ".join(full_text_parts),
|
||||
"segments": segments,
|
||||
"language": info.language,
|
||||
"processing_time": round(processing_time, 2),
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
from celery import Celery
|
||||
from celery.schedules import crontab
|
||||
|
||||
from app.config import settings
|
||||
|
||||
celery_app = Celery(
|
||||
"meeting_protocol",
|
||||
broker=settings.redis_url,
|
||||
backend=settings.redis_url,
|
||||
)
|
||||
|
||||
celery_app.conf.update(
|
||||
task_serializer="json",
|
||||
accept_content=["json"],
|
||||
result_serializer="json",
|
||||
timezone="UTC",
|
||||
enable_utc=True,
|
||||
task_track_started=True,
|
||||
beat_schedule={
|
||||
"cleanup-old-data": {
|
||||
"task": "app.tasks.cleanup.cleanup_old_data",
|
||||
"schedule": crontab(hour=3, minute=0),
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
celery_app.autodiscover_tasks(["app.tasks"])
|
||||
@@ -0,0 +1,69 @@
|
||||
import os
|
||||
import uuid
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.tasks.celery_app import celery_app
|
||||
from app.config import settings
|
||||
from app.models.recording import Recording, RecordingStatus
|
||||
from app.models.transcription import Transcription
|
||||
from app.services.transcription import extract_audio, transcribe_audio
|
||||
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
sync_db_url = settings.database_url.replace("+asyncpg", "+psycopg2").replace(
|
||||
"postgresql+psycopg2", "postgresql"
|
||||
)
|
||||
sync_engine = create_engine(sync_db_url)
|
||||
SyncSession = sessionmaker(sync_engine)
|
||||
|
||||
|
||||
@celery_app.task(bind=True, name="app.tasks.transcription.transcribe_recording")
|
||||
def transcribe_recording(self, recording_id: str):
|
||||
rec_uuid = uuid.UUID(recording_id)
|
||||
|
||||
with SyncSession() as db:
|
||||
recording = db.execute(
|
||||
select(Recording).where(Recording.id == rec_uuid)
|
||||
).scalar_one_or_none()
|
||||
|
||||
if not recording:
|
||||
return {"error": "Recording not found"}
|
||||
|
||||
recording.status = RecordingStatus.processing
|
||||
db.commit()
|
||||
|
||||
try:
|
||||
file_path = recording.file_path
|
||||
audio_path = file_path
|
||||
|
||||
video_formats = {"mp4", "webm", "mkv", "avi", "mov"}
|
||||
if recording.format in video_formats:
|
||||
audio_path = extract_audio(file_path)
|
||||
|
||||
result = transcribe_audio(audio_path)
|
||||
|
||||
if audio_path != file_path and os.path.exists(audio_path):
|
||||
os.remove(audio_path)
|
||||
|
||||
transcription = Transcription(
|
||||
recording_id=rec_uuid,
|
||||
text=result["text"],
|
||||
segments=result["segments"],
|
||||
language=result["language"],
|
||||
whisper_model=settings.whisper_model,
|
||||
processing_time=result["processing_time"],
|
||||
)
|
||||
db.add(transcription)
|
||||
recording.status = RecordingStatus.done
|
||||
db.commit()
|
||||
db.refresh(transcription)
|
||||
|
||||
return {"transcription_id": str(transcription.id), "status": "done"}
|
||||
|
||||
except Exception as e:
|
||||
recording.status = RecordingStatus.error
|
||||
db.commit()
|
||||
raise self.retry(exc=e, max_retries=1, countdown=30)
|
||||
@@ -6,6 +6,7 @@ python-multipart==0.0.20
|
||||
# Database
|
||||
sqlalchemy[asyncio]==2.0.36
|
||||
asyncpg==0.30.0
|
||||
psycopg2-binary==2.9.10
|
||||
alembic==1.14.1
|
||||
|
||||
# Validation
|
||||
|
||||
@@ -0,0 +1,41 @@
|
||||
import os
|
||||
import pytest
|
||||
from unittest.mock import patch, MagicMock
|
||||
|
||||
from app.services.transcription import extract_audio, transcribe_audio
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_extract_audio_calls_ffmpeg():
|
||||
with patch("app.services.transcription.subprocess.run") as mock_run:
|
||||
mock_run.return_value = MagicMock(returncode=0)
|
||||
result = extract_audio("/tmp/test.mp4")
|
||||
assert result.endswith(".wav")
|
||||
mock_run.assert_called_once()
|
||||
args = mock_run.call_args
|
||||
assert "ffmpeg" in args[0][0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_transcribe_audio_returns_expected_structure():
|
||||
mock_segment = MagicMock()
|
||||
mock_segment.start = 0.0
|
||||
mock_segment.end = 5.0
|
||||
mock_segment.text = " Hello world "
|
||||
|
||||
mock_info = MagicMock()
|
||||
mock_info.language = "en"
|
||||
|
||||
with patch("app.services.transcription._get_model") as mock_model:
|
||||
mock_model.return_value.transcribe.return_value = (
|
||||
iter([mock_segment]),
|
||||
mock_info,
|
||||
)
|
||||
result = transcribe_audio("/tmp/test.wav")
|
||||
|
||||
assert "text" in result
|
||||
assert result["text"] == "Hello world"
|
||||
assert len(result["segments"]) == 1
|
||||
assert result["segments"][0]["start"] == 0.0
|
||||
assert result["language"] == "en"
|
||||
assert "processing_time" in result
|
||||
Reference in New Issue
Block a user