Whisper_AI / backend /app.py
TANMAY-555's picture
Upload folder using huggingface_hub
c85835e verified
Raw
History Blame Contribute Delete
8.47 kB
import json
import logging
import threading
import time
import uuid
from pathlib import Path
from typing import Any, Dict
from flask import Flask, Response, jsonify, request, send_file, send_from_directory
from flask_cors import CORS
from werkzeug.utils import secure_filename
from . import config
from .transcriber import DEVICE, COMPUTE_TYPE, cleanup_outputs, device_info, get_model, transcribe_file
logger = logging.getLogger("whisper-app")
if not logger.handlers:
logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
FRONTEND_DIR = config.BASE_DIR / "frontend"
app = Flask(__name__, static_folder=str(FRONTEND_DIR), static_url_path="")
app.config["MAX_CONTENT_LENGTH"] = config.MAX_UPLOAD_MB * 1024 * 1024
CORS(app)
# In-memory job store {job_id: {status, percent, status_text, result, error}}
_jobs: Dict[str, Dict[str, Any]] = {}
_lock = threading.Lock()
def _allowed(filename: str) -> bool:
return Path(filename).suffix.lower() in config.ALLOWED_AUDIO_EXTENSIONS
def _run_job(job_id: str, file_path: Path, params: dict) -> None:
def update(pct: int, status: str) -> None:
with _lock:
j = _jobs.get(job_id)
if j:
j["percent"] = pct
j["status_text"] = status
try:
result = transcribe_file(
audio_path=file_path,
model_size=params["model_size"],
language=params.get("language") or None,
task=params["task"],
word_timestamps=params.get("word_timestamps", False),
no_condition=params.get("no_condition", False),
fmt=params["format"],
progress_callback=update,
)
ext_map = {"txt": ".txt", "srt": ".srt", "vtt": ".vtt", "json": ".json", "tsv": ".tsv"}
ext = ext_map.get(params["format"], ".txt")
out_path = config.OUTPUT_DIR / f"{job_id}{ext}"
out_path.write_text(result["formatted"], encoding="utf-8")
lang_code = result["language"] or "unknown"
payload = {
"text": result["text"],
"segments": result["segments"],
"language": lang_code,
"language_name": config.LANGUAGE_NAMES.get(lang_code, lang_code.upper()),
"language_probability": result["language_probability"],
"duration": result["duration"],
"format": result["format"],
"model_size": result["model_size"],
"task": result["task"],
"cached": False,
"download_url": f"/api/download/{job_id}",
"filename": f"transcript{ext}",
}
with _lock:
j = _jobs.get(job_id)
if j:
j.update({"status": "completed", "percent": 100, "status_text": "Done", "result": payload})
except Exception as exc:
logger.exception("Job %s failed", job_id)
with _lock:
j = _jobs.get(job_id)
if j:
j.update({"status": "error", "percent": 100, "status_text": "Failed", "error": str(exc)})
finally:
try:
file_path.unlink(missing_ok=True)
except Exception:
pass
# ─── routes ──────────────────────────────────────────────────────────────────
@app.get("/health")
def health():
return {"status": "ok", "model_size": config.MODEL_SIZE, **device_info()}, 200
@app.get("/api/config")
def api_config():
return jsonify({
"languages": config.LANGUAGES,
"language_names": config.LANGUAGE_NAMES,
"formats": config.FORMATS,
"models": config.MODELS,
"default_model": config.DEFAULT_MODEL,
"default_format": config.DEFAULT_FORMAT,
"max_upload_mb": config.MAX_UPLOAD_MB,
"allowed_extensions": sorted(ext.lstrip(".") for ext in config.ALLOWED_AUDIO_EXTENSIONS),
"device": DEVICE,
"compute_type": COMPUTE_TYPE,
})
@app.post("/transcribe")
@app.post("/api/transcribe")
def transcribe():
if "file" not in request.files:
return jsonify({"error": "Missing 'file' field."}), 400
uploaded = request.files["file"]
if not uploaded.filename:
return jsonify({"error": "No file selected."}), 400
if not _allowed(uploaded.filename):
return jsonify({"error": "Unsupported file type."}), 400
safe = secure_filename(uploaded.filename) or "audio.bin"
job_id = uuid.uuid4().hex
tmp = config.OUTPUT_DIR / f"upload_{job_id}_{safe}"
uploaded.save(tmp)
fmt = request.form.get("format", config.DEFAULT_FORMAT)
if fmt not in config.FORMATS:
fmt = config.DEFAULT_FORMAT
model_size = request.form.get("model_size", config.DEFAULT_MODEL)
if model_size not in config.MODELS:
model_size = config.DEFAULT_MODEL
lang = request.form.get("language", "auto")
if lang == "auto":
lang = None
bool_field = lambda k: request.form.get(k) in ("on", "true", "1", "yes")
params = {
"model_size": model_size,
"language": lang,
"task": "translate" if bool_field("translate") else "transcribe",
"format": fmt,
"word_timestamps": bool_field("word_timestamps"),
"no_condition": bool_field("no_condition"),
}
with _lock:
_jobs[job_id] = {
"status": "queued",
"percent": 0,
"status_text": "Queued…",
"result": None,
"error": None,
}
threading.Thread(target=_run_job, args=(job_id, tmp, params), daemon=True).start()
return jsonify({"job_id": job_id}), 202
@app.get("/progress/<job_id>")
@app.get("/api/progress/<job_id>")
def progress_stream(job_id: str):
def generate():
deadline = time.time() + 900 # 15-min hard limit
last_pct = -1
while time.time() < deadline:
with _lock:
job = _jobs.get(job_id)
if job is None:
yield f"event: error\ndata: {json.dumps({'error': 'Job not found.'})}\n\n"
return
status = job["status"]
if status == "completed":
yield f"event: done\ndata: {json.dumps(job['result'])}\n\n"
return
if status == "error":
yield f"event: error\ndata: {json.dumps({'error': job.get('error', 'Transcription failed.')})}\n\n"
return
pct = job.get("percent", 0)
if pct != last_pct:
last_pct = pct
data = {"percent": pct, "status": job.get("status_text", "Working…"), "eta": None}
yield f"data: {json.dumps(data)}\n\n"
else:
yield ": heartbeat\n\n"
time.sleep(0.75)
yield f"event: error\ndata: {json.dumps({'error': 'Transcription timed out (15 min).'})}\n\n"
return Response(
generate(),
mimetype="text/event-stream",
headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"},
)
@app.get("/download/<job_id>")
@app.get("/api/download/<job_id>")
def download(job_id: str):
with _lock:
job = _jobs.get(job_id)
if not job or job["status"] != "completed" or not job.get("result"):
return jsonify({"error": "Not found or not ready."}), 404
result = job["result"]
ext_map = {"txt": ".txt", "srt": ".srt", "vtt": ".vtt", "json": ".json", "tsv": ".tsv"}
ext = ext_map.get(result["format"], ".txt")
out_path = config.OUTPUT_DIR / f"{job_id}{ext}"
if not out_path.exists():
return jsonify({"error": "File has been cleaned up."}), 404
return send_file(
out_path,
mimetype="text/plain; charset=utf-8",
as_attachment=True,
download_name=result["filename"],
)
@app.get("/")
def index():
return send_from_directory(FRONTEND_DIR, "index.html")