Initial commit: audio2text 双语字幕生成服务
- 音频/视频转双语(英/中)SRT 字幕,Docker 容器化,CPU 开发/GPU 生产同一份代码 - faster-whisper ASR(词级时间戳) + 断句时间戳重算 + NLLB 翻译(模型不共驻) - 分片上传(断点续传) + SQLite 持久化 + 主页/历史/日志页面 - 历史页文件名搜索;缓存定时清理(默认保留7天,可配置) - 双 Dockerfile(cpu/gpu) + setup/start/stop 脚本
This commit is contained in:
0
app/services/__init__.py
Normal file
0
app/services/__init__.py
Normal file
65
app/services/asr_service.py
Normal file
65
app/services/asr_service.py
Normal file
@@ -0,0 +1,65 @@
|
||||
"""语音识别服务:faster-whisper,输出带词级时间戳的 segments。
|
||||
|
||||
CPU dev: tiny.en + int8;GPU prod: large-v3-turbo + float16。同一份代码。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from pathlib import Path
|
||||
|
||||
from ..config import get_settings
|
||||
from .model_manager import get_model_manager
|
||||
from .types import Segment, Word
|
||||
|
||||
logger = logging.getLogger("audio2text.asr")
|
||||
|
||||
|
||||
def transcribe(wav_path: Path) -> list[Segment]:
|
||||
"""转写 wav,返回 segments(含词级时间戳)。
|
||||
|
||||
Args:
|
||||
wav_path: 16kHz mono PCM wav
|
||||
|
||||
Returns:
|
||||
list[Segment],每个 Segment 带词级 words(若 word_timestamps 启用)。
|
||||
"""
|
||||
s = get_settings().asr
|
||||
if not wav_path.is_file():
|
||||
raise FileNotFoundError(f"音频不存在:{wav_path}")
|
||||
|
||||
model = get_model_manager().get_asr()
|
||||
logger.debug("开始转写 %s(model=%s language=%s)", wav_path.name, s.model, s.language)
|
||||
|
||||
segments_gen, info = model.transcribe(
|
||||
str(wav_path),
|
||||
language=s.language,
|
||||
word_timestamps=s.word_timestamps,
|
||||
vad_filter=s.vad_filter,
|
||||
beam_size=5,
|
||||
)
|
||||
logger.debug(
|
||||
"音频时长 %.1fs,检测语言=%s(置信度 %.2f)",
|
||||
info.duration, info.language, info.language_probability,
|
||||
)
|
||||
|
||||
segments: list[Segment] = []
|
||||
for seg in segments_gen:
|
||||
words: list[Word] = []
|
||||
if s.word_timestamps and getattr(seg, "words", None):
|
||||
for w in seg.words:
|
||||
words.append(Word(
|
||||
text=w.word.strip(),
|
||||
start=float(w.start),
|
||||
end=float(w.end),
|
||||
probability=float(getattr(w, "probability", 1.0)),
|
||||
))
|
||||
segments.append(Segment(
|
||||
text=seg.text.strip(),
|
||||
start=float(seg.start),
|
||||
end=float(seg.end),
|
||||
words=words,
|
||||
))
|
||||
logger.debug("转写完成:%d 段,%d 词。",
|
||||
len(segments), sum(len(s.words) for s in segments))
|
||||
return segments
|
||||
194
app/services/cache_cleaner.py
Normal file
194
app/services/cache_cleaner.py
Normal file
@@ -0,0 +1,194 @@
|
||||
"""缓存清理:删除超期任务产物 + 孤儿目录,并删除对应 DB 记录。
|
||||
|
||||
「缓存」指每个任务落盘的产物:
|
||||
<output_dir>/task_<id>/ 双语 / 英文 / 中文字幕
|
||||
<work_dir>/task_<id>.wav 中间音频(若 keep_audio=true 未被管线删掉)
|
||||
<upload_dir>/yyyy/mm/<uuid>.<ext> 保留的原始视频(若 delete_original_after_extract=false)
|
||||
|
||||
超期 = `created_at` 早于 `now - cache_retention_days`。超期任务的全部产物连同
|
||||
Task / UploadSession 行一起删除,避免历史页出现指向已删文件的死链接。
|
||||
|
||||
另做一次孤儿扫描:output_dir / work_dir 下存在但无对应 Task 的目录(进程崩溃残留),
|
||||
按目录 mtime 判超期后删除,让磁盘不被异常退出留下的碎片占满。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import shutil
|
||||
import time
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from pathlib import Path
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from ..config import get_settings
|
||||
from ..database import get_session_local
|
||||
from ..models.task import Task
|
||||
from ..models.upload_session import UploadSession
|
||||
|
||||
logger = logging.getLogger("audio2text.cache")
|
||||
|
||||
|
||||
def _now_utc() -> datetime:
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
|
||||
def _aware(dt: datetime | None) -> datetime | None:
|
||||
"""SQLite 存 naive datetime,统一补 UTC。"""
|
||||
if dt is None:
|
||||
return None
|
||||
return dt if dt.tzinfo is not None else dt.replace(tzinfo=timezone.utc)
|
||||
|
||||
|
||||
def _rm_tree(path: Path) -> bool:
|
||||
"""删目录或文件,失败仅警告不抛。返回是否实际删除。"""
|
||||
try:
|
||||
if path.is_dir():
|
||||
shutil.rmtree(path, ignore_errors=False)
|
||||
return True
|
||||
if path.is_file():
|
||||
path.unlink(missing_ok=True)
|
||||
return True
|
||||
except Exception as exc: # pragma: no cover
|
||||
logger.warning("清理 %s 失败:%s", path, exc)
|
||||
return False
|
||||
|
||||
|
||||
def purge_expired_cache(db: Session | None = None) -> dict:
|
||||
"""执行一次清理,返回统计 dict。
|
||||
|
||||
可传入已有 Session(如复用请求会话),不传则自建一个并关闭。
|
||||
"""
|
||||
own_db = db is None
|
||||
if own_db:
|
||||
db = get_session_local()()
|
||||
try:
|
||||
return _purge(db)
|
||||
finally:
|
||||
if own_db:
|
||||
db.close()
|
||||
|
||||
|
||||
def _purge(db: Session) -> dict:
|
||||
s = get_settings()
|
||||
retention = max(0, s.storage.cache_retention_days)
|
||||
cutoff = _now_utc() - timedelta(days=retention)
|
||||
upload_root = s.upload_dir()
|
||||
work_root = s.work_dir()
|
||||
output_root = s.output_dir()
|
||||
|
||||
stats = {"tasks": 0, "outputs": 0, "audio": 0, "videos": 0,
|
||||
"orphans": 0, "retention_days": retention}
|
||||
|
||||
if retention <= 0:
|
||||
logger.info("缓存清理已禁用(cache_retention_days=%s)。", retention)
|
||||
return stats
|
||||
|
||||
# ---------- 1. 超期任务 ----------
|
||||
tasks = db.query(Task).all()
|
||||
for task in tasks:
|
||||
created = _aware(task.created_at)
|
||||
if created is None or created >= cutoff:
|
||||
continue
|
||||
# 输出目录
|
||||
out_dir = output_root / f"task_{task.id}"
|
||||
if out_dir.is_dir() and _rm_tree(out_dir):
|
||||
stats["outputs"] += 1
|
||||
# 中间音频
|
||||
wav = work_root / f"task_{task.id}.wav"
|
||||
if wav.is_file() and _rm_tree(wav):
|
||||
stats["audio"] += 1
|
||||
# 保留的原始视频(若未在提取后删除)
|
||||
if task.source_path:
|
||||
src = upload_root / task.source_path
|
||||
if src.is_file() and _rm_tree(src):
|
||||
stats["videos"] += 1
|
||||
# 删 DB 记录(先删关联的 UploadSession,再删 Task)
|
||||
db.query(UploadSession).filter(UploadSession.task_id == task.id).delete()
|
||||
db.delete(task)
|
||||
stats["tasks"] += 1
|
||||
logger.info(
|
||||
"清理超期任务 id=%s file=%s created=%s",
|
||||
task.id, task.filename, created.isoformat(),
|
||||
)
|
||||
if stats["tasks"]:
|
||||
db.commit()
|
||||
|
||||
# ---------- 2. 孤儿目录扫描 ----------
|
||||
existing_ids = {row[0] for row in db.query(Task.id).all()}
|
||||
stats["orphans"] += _purge_orphans(output_root, "task_", cutoff, existing_ids)
|
||||
stats["orphans"] += _purge_orphans(work_root, "task_", cutoff, existing_ids, suffix=".wav")
|
||||
|
||||
logger.info(
|
||||
"缓存清理完成:任务 %d(输出 %d / 音频 %d / 视频 %d)+ 孤儿 %d,保留期 %d 天",
|
||||
stats["tasks"], stats["outputs"], stats["audio"], stats["videos"],
|
||||
stats["orphans"], retention,
|
||||
)
|
||||
return stats
|
||||
|
||||
|
||||
def _purge_orphans(root: Path, prefix: str, cutoff: datetime,
|
||||
existing_ids: set[int], suffix: str | None = None) -> int:
|
||||
"""删 root 下名为 `prefix<id>` 但无对应 Task 的孤儿条目。
|
||||
|
||||
output_dir 下是目录(task_<id>);work_dir 下可能是 task_<id>.wav 文件。
|
||||
用 mtime 判超期,避免误删刚产生的中间产物。
|
||||
"""
|
||||
if not root.is_dir():
|
||||
return 0
|
||||
n = 0
|
||||
for entry in root.iterdir():
|
||||
name = entry.name
|
||||
if not name.startswith(prefix):
|
||||
continue
|
||||
rest = name[len(prefix):]
|
||||
if suffix:
|
||||
if not rest.endswith(suffix):
|
||||
continue
|
||||
rest = rest[: -len(suffix)]
|
||||
try:
|
||||
tid = int(rest)
|
||||
except ValueError:
|
||||
continue # 不是 task_<id> 命名,跳过
|
||||
if tid in existing_ids:
|
||||
continue
|
||||
# 孤儿:按 mtime 判超期
|
||||
try:
|
||||
mtime = datetime.fromtimestamp(entry.stat().st_mtime, tz=timezone.utc)
|
||||
except OSError:
|
||||
continue
|
||||
if mtime >= cutoff:
|
||||
continue
|
||||
if _rm_tree(entry):
|
||||
n += 1
|
||||
logger.info("清理孤儿缓存 %s(mtime=%s)", entry, mtime.isoformat())
|
||||
return n
|
||||
|
||||
|
||||
def run_forever(interval_hours: int) -> None:
|
||||
"""后台线程入口:先跑一次,再按间隔循环。供 main.py lifespan 拉起。"""
|
||||
interval = max(1, interval_hours) * 3600
|
||||
logger.info("缓存清理调度启动:间隔 %d 小时,保留期 %d 天。",
|
||||
interval_hours, get_settings().storage.cache_retention_days)
|
||||
# 启动时先跑一次
|
||||
try:
|
||||
purge_expired_cache()
|
||||
except Exception as exc: # pragma: no cover
|
||||
logger.warning("启动缓存清理失败:%s", exc)
|
||||
while True:
|
||||
time.sleep(interval)
|
||||
try:
|
||||
purge_expired_cache()
|
||||
except Exception as exc: # pragma: no cover
|
||||
logger.warning("定时缓存清理失败:%s", exc)
|
||||
|
||||
|
||||
if __name__ == "__main__": # pragma: no cover
|
||||
# 容器内手动触发:python -m app.services.cache_cleaner
|
||||
import json
|
||||
|
||||
logging.basicConfig(level=logging.INFO,
|
||||
format="%(asctime)s %(levelname)s %(name)s: %(message)s")
|
||||
result = purge_expired_cache()
|
||||
print(json.dumps(result, ensure_ascii=False))
|
||||
64
app/services/ffmpeg_service.py
Normal file
64
app/services/ffmpeg_service.py
Normal file
@@ -0,0 +1,64 @@
|
||||
"""ffmpeg 服务:从视频/音频文件提取 16kHz 单声道 PCM wav。
|
||||
|
||||
16kHz mono PCM 正是 Whisper 的标准输入,省去模型内重采样。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import shutil
|
||||
import subprocess
|
||||
from pathlib import Path
|
||||
|
||||
logger = logging.getLogger("audio2text.ffmpeg")
|
||||
|
||||
|
||||
class FFmpegError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
def extract_audio(
|
||||
src: Path,
|
||||
out_wav: Path,
|
||||
sample_rate: int = 16000,
|
||||
) -> Path:
|
||||
"""提取音频为 16kHz 单声道 PCM wav。
|
||||
|
||||
Args:
|
||||
src: 输入视频/音频文件
|
||||
out_wav: 输出 wav 路径
|
||||
sample_rate: 采样率,默认 16000(Whisper 标准)
|
||||
|
||||
Returns:
|
||||
out_wav 路径
|
||||
|
||||
Raises:
|
||||
FFmpegError: ffmpeg 不可用或提取失败
|
||||
"""
|
||||
if shutil.which("ffmpeg") is None:
|
||||
raise FFmpegError("ffmpeg 未安装;容器内应通过 apt 装好。")
|
||||
if not src.is_file():
|
||||
raise FFmpegError(f"源文件不存在:{src}")
|
||||
|
||||
out_wav.parent.mkdir(parents=True, exist_ok=True)
|
||||
# -vn 去视频;-ac 1 单声道;-ar 16k 采样率;-c:a pcm_s16le 16bit PCM
|
||||
cmd = [
|
||||
"ffmpeg", "-y", "-loglevel", "error",
|
||||
"-i", str(src),
|
||||
"-vn", "-ac", "1", "-ar", str(sample_rate),
|
||||
"-c:a", "pcm_s16le",
|
||||
str(out_wav),
|
||||
]
|
||||
logger.debug("提取音频:%s -> %s", src.name, out_wav.name)
|
||||
logger.debug("ffmpeg 命令:%s", " ".join(cmd))
|
||||
try:
|
||||
result = subprocess.run(cmd, capture_output=True, text=True, timeout=3600)
|
||||
except subprocess.TimeoutExpired as exc:
|
||||
raise FFmpegError(f"ffmpeg 超时(>1h):{src}") from exc
|
||||
if result.returncode != 0:
|
||||
raise FFmpegError(
|
||||
f"ffmpeg 失败 (code={result.returncode}): {result.stderr.strip()[:500]}"
|
||||
)
|
||||
if not out_wav.is_file():
|
||||
raise FFmpegError(f"ffmpeg 未生成输出文件:{out_wav}")
|
||||
return out_wav
|
||||
109
app/services/log_buffer.py
Normal file
109
app/services/log_buffer.py
Normal file
@@ -0,0 +1,109 @@
|
||||
"""内存日志缓冲:捕获所有 audio2text.* 日志到有界 deque,供 /logs 页面查询。
|
||||
|
||||
设计:
|
||||
- MemoryLogHandler 挂到 `audio2text` logger,子 logger 的记录经传播自动汇入。
|
||||
- deque(maxlen=N) 有界,旧记录自动淘汰,避免内存无限增长。
|
||||
- emit() 在 logging 内部锁下执行,deque.append 线程安全。
|
||||
- 查询时按级别过滤(≥指定级别)、取最近 N 条,返回结构化 dict。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from collections import deque
|
||||
from datetime import datetime, timezone
|
||||
|
||||
# 级别名 → 数值,用于过滤
|
||||
_LEVELS = {
|
||||
"debug": logging.DEBUG,
|
||||
"info": logging.INFO,
|
||||
"warning": logging.WARNING,
|
||||
"error": logging.ERROR,
|
||||
}
|
||||
|
||||
|
||||
class _BoundedLogRecord:
|
||||
"""从 LogRecord 提取显示所需字段,避免持有完整对象引用。"""
|
||||
|
||||
__slots__ = ("ts", "level_no", "level_name", "logger_name", "message", "traceback")
|
||||
|
||||
def __init__(self, record: logging.LogRecord) -> None:
|
||||
self.ts: float = record.created
|
||||
self.level_no: int = record.levelno
|
||||
self.level_name: str = record.levelname
|
||||
self.logger_name: str = record.name
|
||||
# getMessage() 应用 %-style 懒格式化
|
||||
self.message: str = record.getMessage()
|
||||
# exc_info 存的是 (type, value, tb) 三元组,格式化为字符串
|
||||
self.traceback: str | None = None
|
||||
if record.exc_info:
|
||||
import traceback as _tb
|
||||
self.traceback = "".join(_tb.format_exception(*record.exc_info))
|
||||
|
||||
|
||||
class MemoryLogHandler(logging.Handler):
|
||||
"""把日志记录存入有界 deque,供 /api/logs 查询。"""
|
||||
|
||||
def __init__(self, buffer_size: int = 2000) -> None:
|
||||
super().__init__()
|
||||
self._buffer: deque[_BoundedLogRecord] = deque(maxlen=buffer_size)
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def emit(self, record: logging.LogRecord) -> None:
|
||||
try:
|
||||
entry = _BoundedLogRecord(record)
|
||||
with self._lock:
|
||||
self._buffer.append(entry)
|
||||
except Exception: # pragma: no cover —— logging 不应抛
|
||||
self.handleError(record)
|
||||
|
||||
def get_records(
|
||||
self, level: str = "info", tail: int = 200,
|
||||
) -> list[dict]:
|
||||
"""返回最近 tail 条、≥指定级别的日志(JSON 友好的 dict 列表)。
|
||||
|
||||
Args:
|
||||
level: debug | info | warning | error(最小级别)
|
||||
tail: 最多返回条数
|
||||
"""
|
||||
min_level = _LEVELS.get(level.lower(), logging.INFO)
|
||||
with self._lock:
|
||||
snapshot = list(self._buffer)
|
||||
# 过滤 + 取最近 tail 条
|
||||
filtered = [r for r in snapshot if r.level_no >= min_level]
|
||||
result = []
|
||||
for r in filtered[-tail:]:
|
||||
ts_str = datetime.fromtimestamp(r.ts, tz=timezone.utc).strftime(
|
||||
"%Y-%m-%d %H:%M:%S"
|
||||
)
|
||||
result.append({
|
||||
"ts": ts_str,
|
||||
"level": r.level_name,
|
||||
"logger": r.logger_name,
|
||||
"msg": r.message,
|
||||
"traceback": r.traceback,
|
||||
})
|
||||
return result
|
||||
|
||||
def clear(self) -> None:
|
||||
with self._lock:
|
||||
self._buffer.clear()
|
||||
|
||||
|
||||
# 进程级单例
|
||||
_buffer: MemoryLogHandler | None = None
|
||||
|
||||
|
||||
def get_log_buffer() -> MemoryLogHandler:
|
||||
global _buffer
|
||||
if _buffer is None:
|
||||
_buffer = MemoryLogHandler()
|
||||
return _buffer
|
||||
|
||||
|
||||
def init_log_buffer(buffer_size: int) -> MemoryLogHandler:
|
||||
"""(重新)创建缓冲并返回,供 main.py 启动时按配置初始化。"""
|
||||
global _buffer
|
||||
_buffer = MemoryLogHandler(buffer_size=buffer_size)
|
||||
return _buffer
|
||||
136
app/services/model_manager.py
Normal file
136
app/services/model_manager.py
Normal file
@@ -0,0 +1,136 @@
|
||||
"""模型管理器:ASR 与翻译模型不共驻,任一时刻 GPU 上只有一个模型。
|
||||
|
||||
设计:
|
||||
- 单例 ModelManager 跟踪当前加载的模型类型(none / asr / translator)
|
||||
- get_asr(): 若翻译器在内存 → 先卸载(del + gc + empty_cache)→ 加载 faster-whisper
|
||||
- get_translator(): 若 ASR 在内存 → 先卸载 → 加载 NLLB
|
||||
- 翻译阶段独占显存,可用大 batch_size;ASR 阶段同理
|
||||
|
||||
这样在 24G 3090 上无需担心显存叠加,CPU dev 时也省内存。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import gc
|
||||
import logging
|
||||
import threading
|
||||
from typing import Any
|
||||
|
||||
from ..config import get_settings
|
||||
|
||||
logger = logging.getLogger("audio2text.models")
|
||||
|
||||
# 全局单例 + 锁:模型加载/卸载必须串行
|
||||
_lock = threading.Lock()
|
||||
|
||||
|
||||
def _free_memory() -> None:
|
||||
"""释放 Python 对象与 GPU 缓存。"""
|
||||
gc.collect()
|
||||
try:
|
||||
import torch
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.synchronize()
|
||||
except Exception: # pragma: no cover
|
||||
pass
|
||||
|
||||
|
||||
class ModelManager:
|
||||
"""ASR / 翻译模型的不共驻管理器。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._asr: Any = None # faster_whisper.WhisperModel
|
||||
self._translator: Any = None # transformers pipeline
|
||||
self._current: str = "none" # none | asr | translator
|
||||
|
||||
# ---------------- ASR ----------------
|
||||
|
||||
def get_asr(self) -> Any:
|
||||
"""返回已加载的 faster-whisper 模型;必要时先卸载翻译器。"""
|
||||
with _lock:
|
||||
if self._asr is not None:
|
||||
return self._asr
|
||||
if self._translator is not None:
|
||||
self._unload_translator_locked()
|
||||
s = get_settings().asr
|
||||
logger.debug("加载 ASR 模型 model=%s device=%s compute_type=%s",
|
||||
s.model, s.device, s.compute_type)
|
||||
from faster_whisper import WhisperModel
|
||||
# device/compute_type 组合:cpu+int8 / cuda+float16
|
||||
self._asr = WhisperModel(
|
||||
s.model, device=s.device, compute_type=s.compute_type,
|
||||
)
|
||||
self._current = "asr"
|
||||
logger.debug("ASR 模型已就绪。")
|
||||
return self._asr
|
||||
|
||||
def unload_asr(self) -> None:
|
||||
with _lock:
|
||||
self._unload_asr_locked()
|
||||
|
||||
def _unload_asr_locked(self) -> None:
|
||||
if self._asr is None:
|
||||
return
|
||||
logger.debug("卸载 ASR 模型(释放显存供翻译器独占)。")
|
||||
# faster-whisper 模型无显式 close,del 即可
|
||||
del self._asr
|
||||
self._asr = None
|
||||
self._current = "none" if self._translator is None else "translator"
|
||||
_free_memory()
|
||||
|
||||
# ---------------- 翻译器 ----------------
|
||||
|
||||
def get_translator(self) -> Any:
|
||||
"""返回已加载的 NLLB 翻译 pipeline;必要时先卸载 ASR。"""
|
||||
with _lock:
|
||||
if self._translator is not None:
|
||||
return self._translator
|
||||
if self._asr is not None:
|
||||
self._unload_asr_locked()
|
||||
s = get_settings().translation
|
||||
logger.debug("加载翻译模型 model=%s device=%s", s.model, s.device)
|
||||
from transformers import pipeline
|
||||
self._translator = pipeline(
|
||||
"translation",
|
||||
model=s.model,
|
||||
device=s.device,
|
||||
src_lang=s.src_lang,
|
||||
tgt_lang=s.tgt_lang,
|
||||
)
|
||||
self._current = "translator"
|
||||
logger.debug("翻译模型已就绪(独占显存,可用大 batch)。")
|
||||
return self._translator
|
||||
|
||||
def unload_translator(self) -> None:
|
||||
with _lock:
|
||||
self._unload_translator_locked()
|
||||
|
||||
def _unload_translator_locked(self) -> None:
|
||||
if self._translator is None:
|
||||
return
|
||||
logger.debug("卸载翻译模型。")
|
||||
# 释放 pipeline 持有的 model + tokenizer
|
||||
mdl = getattr(self._translator, "model", None)
|
||||
tok = getattr(self._translator, "tokenizer", None)
|
||||
del self._translator, mdl, tok
|
||||
self._translator = None
|
||||
self._current = "none" if self._asr is None else "asr"
|
||||
_free_memory()
|
||||
|
||||
# ---------------- 状态 ----------------
|
||||
|
||||
@property
|
||||
def current(self) -> str:
|
||||
return self._current
|
||||
|
||||
|
||||
# 进程级单例
|
||||
_manager: ModelManager | None = None
|
||||
|
||||
|
||||
def get_model_manager() -> ModelManager:
|
||||
global _manager
|
||||
if _manager is None:
|
||||
_manager = ModelManager()
|
||||
return _manager
|
||||
153
app/services/pipeline.py
Normal file
153
app/services/pipeline.py
Normal file
@@ -0,0 +1,153 @@
|
||||
"""转写管线编排:提取音频 → ASR → 断句 → 翻译 → 写 SRT。
|
||||
|
||||
任务状态机:
|
||||
queued → extracting → transcribing → segmenting → translating → done
|
||||
任一步失败 → failed
|
||||
|
||||
模型不共驻:ASR 与翻译分阶段加载,翻译时先卸载 Whisper 释放显存跑大 batch。
|
||||
管线在后台线程跑(每个任务一个线程),通过 DB 更新状态与进度。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
import traceback
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
|
||||
from ..config import get_settings
|
||||
from ..database import get_session_local
|
||||
from ..models.task import Task
|
||||
from . import ffmpeg_service, asr_service, segmenter, translate_service, srt_writer
|
||||
from .types import Subtitle
|
||||
|
||||
logger = logging.getLogger("audio2text.pipeline")
|
||||
|
||||
# 进度锚点(各阶段在 0-100 中的占比)
|
||||
P_EXTRACT = 5.0
|
||||
P_TRANSCRIBE_START = 5.0
|
||||
P_TRANSCRIBE_END = 55.0
|
||||
P_SEGMENT_START = 55.0
|
||||
P_SEGMENT_END = 60.0
|
||||
P_TRANSLATE_START = 60.0
|
||||
P_TRANSLATE_END = 98.0
|
||||
P_DONE = 100.0
|
||||
|
||||
|
||||
def enqueue_task(task_id: int) -> None:
|
||||
"""把任务交给后台线程处理(非阻塞,供 upload_service.complete 调用)。"""
|
||||
t = threading.Thread(target=_run_task, args=(task_id,), daemon=True)
|
||||
t.start()
|
||||
logger.info("任务 %d 已入队(后台线程 %s)。", task_id, t.name)
|
||||
|
||||
|
||||
def _run_task(task_id: int) -> None:
|
||||
"""后台执行完整管线。所有异常都被捕获并写入 task.error。"""
|
||||
db = get_session_local()()
|
||||
try:
|
||||
task = db.get(Task, task_id)
|
||||
if task is None:
|
||||
logger.error("任务 %d 不存在。", task_id)
|
||||
return
|
||||
_pipeline(db, task)
|
||||
except Exception as exc:
|
||||
logger.exception("任务 %d 失败:%s", task_id, exc)
|
||||
_mark_failed(db, task_id, str(exc))
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def _pipeline(db, task: Task) -> None:
|
||||
s = get_settings()
|
||||
src = s.upload_dir() / task.source_path
|
||||
|
||||
# ---------- 1. 提取音频 ----------
|
||||
_set_status(db, task, "extracting", P_EXTRACT)
|
||||
wav = s.work_dir() / f"task_{task.id}.wav"
|
||||
ffmpeg_service.extract_audio(src, wav)
|
||||
|
||||
# 按配置决定是否删原始视频(提取成功后)
|
||||
if s.processing.delete_original_after_extract and src.is_file():
|
||||
try:
|
||||
src.unlink()
|
||||
logger.info("已删除原始视频 %s(delete_original_after_extract=true)。", src.name)
|
||||
except OSError as exc:
|
||||
logger.warning("删除原始视频失败 %s: %s", src, exc)
|
||||
|
||||
# ---------- 2. 语音识别 ----------
|
||||
_set_status(db, task, "transcribing", P_TRANSCRIBE_START)
|
||||
segments = asr_service.transcribe(wav)
|
||||
_set_status(db, task, "transcribing", P_TRANSCRIBE_END,
|
||||
note=f"识别出 {len(segments)} 段")
|
||||
|
||||
# ---------- 3. 断句 + 时间戳重算 ----------
|
||||
_set_status(db, task, "segmenting", P_SEGMENT_START)
|
||||
subs = segmenter.resegment(segments)
|
||||
_set_status(db, task, "segmenting", P_SEGMENT_END,
|
||||
note=f"重组为 {len(subs)} 条字幕")
|
||||
|
||||
# ---------- 4. 翻译 ----------
|
||||
_set_status(db, task, "translating", P_TRANSLATE_START)
|
||||
# 翻译阶段:model_manager 会自动卸载 ASR、加载翻译器(独占显存)
|
||||
zh_texts = translate_service.translate(subs)
|
||||
_set_status(db, task, "translating", P_TRANSLATE_END,
|
||||
note=f"翻译 {len(zh_texts)} 条")
|
||||
|
||||
# ---------- 5. 写 SRT ----------
|
||||
out_dir = s.output_dir()
|
||||
stem = Path(task.filename).stem
|
||||
en_path = out_dir / f"task_{task.id}/{stem}.en.srt"
|
||||
zh_path = out_dir / f"task_{task.id}/{stem}.zh.srt"
|
||||
bi_path = out_dir / f"task_{task.id}/{stem}.srt"
|
||||
|
||||
srt_writer.write_srt(subs, en_path)
|
||||
# 中文 SRT(用译文 + 同时间戳)
|
||||
zh_subs = [Subtitle(text=zh, start=sub.start, end=sub.end)
|
||||
for zh, sub in zip(zh_texts, subs)]
|
||||
srt_writer.write_srt(zh_subs, zh_path)
|
||||
srt_writer.write_bilingual_srt(subs, zh_texts, bi_path)
|
||||
|
||||
# 记录相对路径
|
||||
task.en_srt_path = str(en_path.relative_to(out_dir))
|
||||
task.zh_srt_path = str(zh_path.relative_to(out_dir))
|
||||
task.bilingual_srt_path = str(bi_path.relative_to(out_dir))
|
||||
task.status = "done"
|
||||
task.progress = P_DONE
|
||||
task.updated_at = datetime.now(timezone.utc)
|
||||
db.commit()
|
||||
|
||||
# 清理中间音频
|
||||
if not s.processing.keep_audio and wav.is_file():
|
||||
try:
|
||||
wav.unlink()
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
logger.info("任务 %d 完成:%s", task.id, bi_path.name)
|
||||
|
||||
|
||||
# ---------------- DB 状态更新 ----------------
|
||||
|
||||
def _set_status(db, task: Task, status: str, progress: float, note: str = "") -> None:
|
||||
task.status = status
|
||||
task.progress = progress
|
||||
task.updated_at = datetime.now(timezone.utc)
|
||||
db.commit()
|
||||
if note:
|
||||
logger.info("任务 %d [%s %.0f%%] %s", task.id, status, progress, note)
|
||||
else:
|
||||
logger.info("任务 %d [%s %.0f%%]", task.id, status, progress)
|
||||
|
||||
|
||||
def _mark_failed(db, task_id: int, error: str) -> None:
|
||||
try:
|
||||
task = db.get(Task, task_id)
|
||||
if task is None:
|
||||
return
|
||||
task.status = "failed"
|
||||
task.error = error[:2000]
|
||||
task.updated_at = datetime.now(timezone.utc)
|
||||
db.commit()
|
||||
except Exception: # pragma: no cover
|
||||
logger.error("写入失败状态时又失败:\n%s", traceback.format_exc())
|
||||
22
app/services/reaper.py
Normal file
22
app/services/reaper.py
Normal file
@@ -0,0 +1,22 @@
|
||||
"""后台 reaper:清理被放弃的分片上传会话与临时文件。供 main.py 启动时调用。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from ..database import get_session_local
|
||||
from .upload_service import UploadService
|
||||
|
||||
logger = logging.getLogger("audio2text.reaper")
|
||||
|
||||
|
||||
def reap_stale_sessions() -> int:
|
||||
"""执行一次过期会话清理。"""
|
||||
db = get_session_local()()
|
||||
try:
|
||||
return UploadService(db).reap_stale_sessions()
|
||||
except Exception as exc: # pragma: no cover
|
||||
logger.warning("reaper 执行失败:%s", exc)
|
||||
return 0
|
||||
finally:
|
||||
db.close()
|
||||
233
app/services/segmenter.py
Normal file
233
app/services/segmenter.py
Normal file
@@ -0,0 +1,233 @@
|
||||
"""断句 + 时间戳重算(纯算法,零模型开销)。
|
||||
|
||||
Whisper 原始 segment 的断句通常很混乱:每段不一定是完整句子,时间戳也不对齐句界。
|
||||
本模块基于词级时间戳重新断句,得到规范的字幕条目。
|
||||
|
||||
两路策略:
|
||||
1. 精确路(有 word_timestamps):按句末标点(. ! ? ;)切句,超长句再按逗号拆,
|
||||
每条字幕的时间戳直接取首词.start ~ 末词.end,精确无误。
|
||||
2. 匀速估算路(无 word_timestamps):段内按字符数比例分配时间——
|
||||
句start = 段start + (前缀字符数 / 段总字符数) × 段时长。
|
||||
即「短时匀速」假设,无需大模型。
|
||||
|
||||
最后做 SRT 规范化:单条 1–7 秒、≤2 行、每行 ≤42 字符。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
|
||||
from ..config import get_settings
|
||||
from .types import Segment, Subtitle, Word
|
||||
|
||||
logger = logging.getLogger("audio2text.segmenter")
|
||||
|
||||
# 句末标点:句号、感叹号、问号、分号
|
||||
_SENTENCE_END = re.compile(r"[.!?;]+")
|
||||
# 句内停顿:逗号、冒号、破折号
|
||||
_CLAUSE_BREAK = re.compile(r"[,:\-—]+")
|
||||
|
||||
|
||||
def resegment(segments: list[Segment]) -> list[Subtitle]:
|
||||
"""把 ASR segments 重组为规范字幕条目。
|
||||
|
||||
Args:
|
||||
segments: asr_service.transcribe 的输出(可能含词级时间戳)
|
||||
|
||||
Returns:
|
||||
list[Subtitle],已做时长/字数规范化。
|
||||
"""
|
||||
cfg = get_settings().segmentation
|
||||
max_words = cfg.max_words_per_line
|
||||
max_dur = cfg.max_duration_seconds
|
||||
min_dur = cfg.min_duration_seconds
|
||||
max_chars = cfg.max_chars_per_line
|
||||
|
||||
# 第一步:把所有词串成大列表(精确路)或退化为段级(估算路)
|
||||
all_words: list[Word] = []
|
||||
has_word_ts = True
|
||||
for seg in segments:
|
||||
if not seg.words:
|
||||
has_word_ts = False
|
||||
break
|
||||
all_words.extend(seg.words)
|
||||
|
||||
if has_word_ts and all_words:
|
||||
subs = _resegment_by_words(all_words, max_words, max_dur)
|
||||
else:
|
||||
logger.warning("无词级时间戳,退化为匀速估算路。")
|
||||
subs = _resegment_by_estimate(segments, max_words, max_dur)
|
||||
|
||||
# 规范化:合并过短条目、拆分行宽
|
||||
subs = _normalize(subs, min_dur, max_chars)
|
||||
logger.debug("断句完成:%d 条字幕。", len(subs))
|
||||
return subs
|
||||
|
||||
|
||||
# ---------------- 精确路:按词级时间戳 ----------------
|
||||
|
||||
def _resegment_by_words(
|
||||
words: list[Word], max_words: int, max_dur: float,
|
||||
) -> list[Subtitle]:
|
||||
"""按句末标点切句,超长句按逗号拆,时间戳取首末词。"""
|
||||
subs: list[Subtitle] = []
|
||||
# 当前句的词
|
||||
current: list[Word] = []
|
||||
|
||||
def flush(wlist: list[Word]) -> None:
|
||||
if not wlist:
|
||||
return
|
||||
text = " ".join(w.text for w in wlist).strip()
|
||||
if not text:
|
||||
return
|
||||
subs.append(Subtitle(
|
||||
text=text,
|
||||
start=wlist[0].start,
|
||||
end=wlist[-1].end,
|
||||
))
|
||||
|
||||
for w in words:
|
||||
current.append(w)
|
||||
# 句末标点 → 收尾
|
||||
if _SENTENCE_END.search(w.text):
|
||||
_maybe_split_and_flush(current, max_words, max_dur, flush)
|
||||
current = []
|
||||
continue
|
||||
# 超长(词数或时长)→ 优先在最近的逗号处断
|
||||
cur_dur = (current[-1].end - current[0].start) if len(current) > 1 else 0
|
||||
if len(current) >= max_words or cur_dur >= max_dur:
|
||||
_maybe_split_and_flush(current, max_words, max_dur, flush)
|
||||
current = []
|
||||
|
||||
flush(current)
|
||||
return subs
|
||||
|
||||
|
||||
def _maybe_split_and_flush(
|
||||
wlist: list[Word], max_words: int, max_dur: float, flush,
|
||||
) -> None:
|
||||
"""若 wlist 过长,在逗号处再拆;否则整条 flush。"""
|
||||
if len(wlist) <= max_words and (len(wlist) <= 1 or
|
||||
wlist[-1].end - wlist[0].start < max_dur):
|
||||
flush(wlist)
|
||||
return
|
||||
# 找逗号断点
|
||||
parts: list[list[Word]] = []
|
||||
cur: list[Word] = []
|
||||
for w in wlist:
|
||||
cur.append(w)
|
||||
if _CLAUSE_BREAK.search(w.text) and len(cur) >= max_words // 2:
|
||||
parts.append(cur)
|
||||
cur = []
|
||||
if cur:
|
||||
parts.append(cur)
|
||||
# 若逗号拆不开(无逗号),强制按 max_words 等分
|
||||
if len(parts) == 1 and len(parts[0]) > max_words:
|
||||
parts = [parts[0][i:i + max_words] for i in range(0, len(parts[0]), max_words)]
|
||||
for p in parts:
|
||||
flush(p)
|
||||
|
||||
|
||||
# ---------------- 匀速估算路:段内按字符比例 ----------------
|
||||
|
||||
def _resegment_by_estimate(
|
||||
segments: list[Segment], max_words: int, max_dur: float,
|
||||
) -> list[Subtitle]:
|
||||
"""无词级时间戳时:先按文本断句,再按字符数比例估算时间戳。"""
|
||||
subs: list[Subtitle] = []
|
||||
for seg in segments:
|
||||
text = seg.text.strip()
|
||||
if not text:
|
||||
continue
|
||||
dur = seg.end - seg.start
|
||||
# 按句末标点切
|
||||
sentences = _split_sentences(text)
|
||||
if not sentences:
|
||||
sentences = [text]
|
||||
# 段内按字符数比例分配时间
|
||||
total_chars = sum(len(s) for s in sentences) or 1
|
||||
cursor = seg.start
|
||||
for sent in sentences:
|
||||
frac = len(sent) / total_chars
|
||||
est_end = cursor + dur * frac
|
||||
# 超长句再按逗号拆(时间按字符比例再分)
|
||||
if len(sent.split()) > max_words or (est_end - cursor) > max_dur:
|
||||
for clause in _split_clauses(sent):
|
||||
cfrac = len(clause) / len(sent) if len(sent) else 1
|
||||
c_end = cursor + (est_end - cursor) * cfrac
|
||||
subs.append(Subtitle(text=clause.strip(), start=cursor, end=c_end))
|
||||
cursor = c_end
|
||||
else:
|
||||
subs.append(Subtitle(text=sent.strip(), start=cursor, end=est_end))
|
||||
cursor = est_end
|
||||
return subs
|
||||
|
||||
|
||||
def _split_sentences(text: str) -> list[str]:
|
||||
"""按句末标点切句,保留标点。"""
|
||||
parts = _SENTENCE_END.split(text)
|
||||
marks = _SENTENCE_END.findall(text)
|
||||
out = []
|
||||
for i, p in enumerate(parts):
|
||||
p = p.strip()
|
||||
if not p:
|
||||
continue
|
||||
out.append(p + (marks[i] if i < len(marks) else ""))
|
||||
return out
|
||||
|
||||
|
||||
def _split_clauses(sentence: str) -> list[str]:
|
||||
"""按逗号/冒号拆子句,保留标点。"""
|
||||
parts = _CLAUSE_BREAK.split(sentence)
|
||||
marks = _CLAUSE_BREAK.findall(sentence)
|
||||
out = []
|
||||
for i, p in enumerate(parts):
|
||||
p = p.strip()
|
||||
if not p:
|
||||
continue
|
||||
out.append(p + (marks[i - 1] if 0 < i <= len(marks) else ""))
|
||||
return out
|
||||
|
||||
|
||||
# ---------------- 规范化 ----------------
|
||||
|
||||
def _normalize(subs: list[Subtitle], min_dur: float, max_chars: int) -> list[Subtitle]:
|
||||
"""合并过短条目、拆分过宽行。"""
|
||||
# 1. 合并过短(< min_dur 且非末尾)
|
||||
merged: list[Subtitle] = []
|
||||
for s in subs:
|
||||
if merged and (s.end - s.start) < min_dur:
|
||||
prev = merged[-1]
|
||||
prev.text = (prev.text + " " + s.text).strip()
|
||||
prev.end = s.end
|
||||
else:
|
||||
merged.append(Subtitle(text=s.text, start=s.start, end=s.end))
|
||||
|
||||
# 2. 拆分超过 max_chars 的行(按词折行,不改时间戳)
|
||||
out: list[Subtitle] = []
|
||||
for s in merged:
|
||||
if len(s.text) <= max_chars:
|
||||
out.append(s)
|
||||
continue
|
||||
lines = _wrap_text(s.text, max_chars)
|
||||
out.append(Subtitle(text=lines, start=s.start, end=s.end))
|
||||
return out
|
||||
|
||||
|
||||
def _wrap_text(text: str, max_chars: int) -> str:
|
||||
"""按词折成 ≤2 行,每行 ≤max_chars 字符(SRT 规范)。"""
|
||||
words = text.split()
|
||||
lines: list[str] = []
|
||||
cur = ""
|
||||
for idx, w in enumerate(words):
|
||||
if cur and len(cur) + 1 + len(w) > max_chars:
|
||||
lines.append(cur)
|
||||
# 已有一行 + 当前行,合并剩余为一行
|
||||
cur = " ".join([w] + words[idx + 1:])
|
||||
break
|
||||
else:
|
||||
cur = (cur + " " + w).strip() if cur else w
|
||||
if cur:
|
||||
lines.append(cur)
|
||||
return "\n".join(lines[:2])
|
||||
68
app/services/srt_writer.py
Normal file
68
app/services/srt_writer.py
Normal file
@@ -0,0 +1,68 @@
|
||||
"""SRT 字幕文件写入与双语合并。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from .types import Subtitle
|
||||
|
||||
|
||||
def _format_ts(seconds: float) -> str:
|
||||
"""秒 → SRT 时间戳 HH:MM:SS,mmm。"""
|
||||
if seconds < 0:
|
||||
seconds = 0.0
|
||||
ms = int(round((seconds - int(seconds)) * 1000))
|
||||
s = int(seconds) % 60
|
||||
m = (int(seconds) // 60) % 60
|
||||
h = int(seconds) // 3600
|
||||
if ms == 1000: # 四舍五入进位
|
||||
ms = 0
|
||||
s += 1
|
||||
if s == 60:
|
||||
s = 0
|
||||
m += 1
|
||||
if m == 60:
|
||||
m = 0
|
||||
h += 1
|
||||
return f"{h:02d}:{m:02d}:{s:02d},{ms:03d}"
|
||||
|
||||
|
||||
def write_srt(subtitles: list[Subtitle], path: Path) -> Path:
|
||||
"""写单语 SRT。"""
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
lines: list[str] = []
|
||||
for i, sub in enumerate(subtitles, 1):
|
||||
lines.append(str(i))
|
||||
lines.append(f"{_format_ts(sub.start)} --> {_format_ts(sub.end)}")
|
||||
lines.append(sub.text)
|
||||
lines.append("")
|
||||
path.write_text("\n".join(lines), encoding="utf-8")
|
||||
return path
|
||||
|
||||
|
||||
def write_bilingual_srt(
|
||||
en_subs: list[Subtitle],
|
||||
zh_texts: list[str],
|
||||
path: Path,
|
||||
) -> Path:
|
||||
"""写双语合并 SRT:英文在上、中文在下,同一时间戳。
|
||||
|
||||
Args:
|
||||
en_subs: 英文字幕条目
|
||||
zh_texts: 与 en_subs 等长、顺序对应的中文译文
|
||||
path: 输出路径
|
||||
"""
|
||||
if len(en_subs) != len(zh_texts):
|
||||
raise ValueError(
|
||||
f"英文字幕数({len(en_subs)}) 与中文译文数({len(zh_texts)}) 不一致"
|
||||
)
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
lines: list[str] = []
|
||||
for i, (sub, zh) in enumerate(zip(en_subs, zh_texts), 1):
|
||||
lines.append(str(i))
|
||||
lines.append(f"{_format_ts(sub.start)} --> {_format_ts(sub.end)}")
|
||||
lines.append(sub.text)
|
||||
lines.append(zh)
|
||||
lines.append("")
|
||||
path.write_text("\n".join(lines), encoding="utf-8")
|
||||
return path
|
||||
60
app/services/translate_service.py
Normal file
60
app/services/translate_service.py
Normal file
@@ -0,0 +1,60 @@
|
||||
"""翻译服务:NLLB-200,英译中。
|
||||
|
||||
通过 model_manager 加载,确保 ASR 已卸载、翻译器独占显存,从而可用大 batch_size。
|
||||
按字幕条目批量翻译,保留索引对应。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from ..config import get_settings
|
||||
from .model_manager import get_model_manager
|
||||
from .types import Subtitle
|
||||
|
||||
logger = logging.getLogger("audio2text.translate")
|
||||
|
||||
|
||||
def translate(subtitles: list[Subtitle]) -> list[str]:
|
||||
"""批量翻译英文字幕为中文。
|
||||
|
||||
Args:
|
||||
subtitles: 断句后的英文字幕条目
|
||||
|
||||
Returns:
|
||||
list[str],与 subtitles 等长、顺序对应的中文译文。
|
||||
单条翻译失败时该位置回退为原英文。
|
||||
"""
|
||||
if not subtitles:
|
||||
return []
|
||||
|
||||
s = get_settings().translation
|
||||
pipe = get_model_manager().get_translator()
|
||||
batch = s.batch_size
|
||||
max_len = s.max_length
|
||||
|
||||
# 取纯文本(去掉折行),避免翻译把换行符当语义
|
||||
texts = [sub.text.replace("\n", " ").strip() for sub in subtitles]
|
||||
logger.debug("开始翻译 %d 条字幕(batch_size=%d)...", len(texts), batch)
|
||||
|
||||
results: list[str] = []
|
||||
for i in range(0, len(texts), batch):
|
||||
chunk = texts[i:i + batch]
|
||||
try:
|
||||
out = pipe(chunk, max_length=max_len)
|
||||
for item in out:
|
||||
# pipeline 返回 [{"translation_text": "..."}]
|
||||
results.append(item.get("translation_text", "").strip())
|
||||
except Exception as exc: # pragma: no cover
|
||||
logger.warning("第 %d-%d 批翻译失败,逐条重试:%s", i, i + len(chunk), exc)
|
||||
for t in chunk:
|
||||
try:
|
||||
out = pipe([t], max_length=max_len)
|
||||
results.append(out[0].get("translation_text", "").strip())
|
||||
except Exception:
|
||||
results.append(t) # 回退原文
|
||||
if (i // batch + 1) % 5 == 0:
|
||||
logger.debug("已翻译 %d/%d 条。", min(i + len(chunk), len(texts)), len(texts))
|
||||
|
||||
logger.debug("翻译完成:%d 条。", len(results))
|
||||
return results
|
||||
38
app/services/types.py
Normal file
38
app/services/types.py
Normal file
@@ -0,0 +1,38 @@
|
||||
"""管线各阶段共享的数据传输对象(DTO)。
|
||||
|
||||
独立于任何模型加载逻辑——segmenter(纯算法)和 srt_writer(纯 IO)可以只依赖
|
||||
本模块,不拉入 faster_whisper / torch。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
@dataclass
|
||||
class Word:
|
||||
"""ASR 识别出的单个词,带时间戳。"""
|
||||
|
||||
text: str
|
||||
start: float # 秒
|
||||
end: float
|
||||
probability: float = 1.0
|
||||
|
||||
|
||||
@dataclass
|
||||
class Segment:
|
||||
"""ASR 输出的一段文本,可能含词级时间戳。"""
|
||||
|
||||
text: str
|
||||
start: float
|
||||
end: float
|
||||
words: list[Word]
|
||||
|
||||
|
||||
@dataclass
|
||||
class Subtitle:
|
||||
"""断句后的一条字幕:英文文本 + 起止时间戳。"""
|
||||
|
||||
text: str
|
||||
start: float
|
||||
end: float
|
||||
249
app/services/upload_service.py
Normal file
249
app/services/upload_service.py
Normal file
@@ -0,0 +1,249 @@
|
||||
"""分片上传服务:会话管理 + 分片落盘 + 拼接 + 创建转写任务。
|
||||
|
||||
存储布局::
|
||||
|
||||
<work_dir>/<upload_id>/0.part 分片暂存
|
||||
<work_dir>/<upload_id>/1.part
|
||||
...
|
||||
<upload_dir>/<yyyy>/<mm>/<uuid>.<ext> complete 后的正式视频
|
||||
|
||||
与 server 的区别:视频无需 sha256 去重(每个视频都转写),complete 直接创建 Task。
|
||||
管线触发由 controller 调用 pipeline.enqueue_task,本服务不依赖 pipeline。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import shutil
|
||||
import uuid
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from pathlib import Path
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from ..config import get_settings
|
||||
from ..models.task import Task
|
||||
from ..models.upload_session import UploadSession
|
||||
from ..schemas.task import (
|
||||
ChunkUploadResponse,
|
||||
CompleteResponse,
|
||||
CreateSessionRequest,
|
||||
CreateSessionResponse,
|
||||
SessionStatusResponse,
|
||||
)
|
||||
|
||||
logger = logging.getLogger("audio2text.upload")
|
||||
|
||||
|
||||
class UploadService:
|
||||
def __init__(self, db: Session) -> None:
|
||||
self.db = db
|
||||
s = get_settings()
|
||||
self.upload_root = s.upload_dir()
|
||||
self.work_root = s.work_dir()
|
||||
self.chunk_bytes = s.storage.chunk_bytes
|
||||
self.session_ttl = s.storage.chunk_session_ttl_seconds
|
||||
|
||||
# ---------------- 会话生命周期 ----------------
|
||||
|
||||
def create_session(self, body: CreateSessionRequest) -> CreateSessionResponse:
|
||||
upload_id = uuid.uuid4().hex
|
||||
session = UploadSession(
|
||||
upload_id=upload_id,
|
||||
filename=body.filename,
|
||||
size_bytes=body.size_bytes,
|
||||
chunk_size=body.chunk_size,
|
||||
total_chunks=body.total_chunks,
|
||||
uploaded_chunks=[],
|
||||
status="pending",
|
||||
)
|
||||
self.db.add(session)
|
||||
self.db.commit()
|
||||
self._session_dir(upload_id).mkdir(parents=True, exist_ok=True)
|
||||
return CreateSessionResponse(
|
||||
upload_id=upload_id,
|
||||
filename=body.filename,
|
||||
size_bytes=body.size_bytes,
|
||||
chunk_size=body.chunk_size,
|
||||
total_chunks=body.total_chunks,
|
||||
)
|
||||
|
||||
def get_status(self, upload_id: str) -> SessionStatusResponse:
|
||||
session = self._require_session(upload_id)
|
||||
return SessionStatusResponse(
|
||||
upload_id=session.upload_id,
|
||||
filename=session.filename,
|
||||
size_bytes=session.size_bytes,
|
||||
chunk_size=session.chunk_size,
|
||||
total_chunks=session.total_chunks,
|
||||
uploaded_chunks=list(session.uploaded_chunks or []),
|
||||
completed=(session.status == "completed"),
|
||||
task_id=session.task_id,
|
||||
)
|
||||
|
||||
# ---------------- 分片写入 ----------------
|
||||
|
||||
def write_chunk(self, upload_id: str, index: int, data: bytes) -> list[int]:
|
||||
session = self._require_session(upload_id)
|
||||
self._validate_index(session, index)
|
||||
|
||||
session_dir = self._session_dir(upload_id)
|
||||
session_dir.mkdir(parents=True, exist_ok=True)
|
||||
chunk_path = session_dir / f"{index}.part"
|
||||
|
||||
try:
|
||||
with chunk_path.open("wb") as out:
|
||||
out.write(data)
|
||||
out.flush()
|
||||
os.fsync(out.fileno())
|
||||
except Exception:
|
||||
chunk_path.unlink(missing_ok=True)
|
||||
raise
|
||||
|
||||
uploaded = list(session.uploaded_chunks or [])
|
||||
if index not in uploaded:
|
||||
uploaded.append(index)
|
||||
session.uploaded_chunks = uploaded
|
||||
session.updated_at = datetime.now(timezone.utc)
|
||||
self.db.commit()
|
||||
return sorted(uploaded)
|
||||
|
||||
# ---------------- 拼接 + 创建任务 ----------------
|
||||
|
||||
def complete(self, upload_id: str) -> CompleteResponse:
|
||||
session = self._require_session(upload_id)
|
||||
|
||||
# 幂等:已 complete 则返回已建任务
|
||||
if session.status == "completed" and session.task_id is not None:
|
||||
task = self.db.get(Task, session.task_id)
|
||||
if task is not None:
|
||||
return CompleteResponse(
|
||||
task_id=task.id, filename=task.filename,
|
||||
size_bytes=session.size_bytes, status=task.status,
|
||||
)
|
||||
|
||||
uploaded = set(session.uploaded_chunks or [])
|
||||
missing = [i for i in range(session.total_chunks) if i not in uploaded]
|
||||
if missing:
|
||||
raise HTTPException(
|
||||
409,
|
||||
f"分片不齐全:缺失 {len(missing)} 个,例如 {sorted(missing)[:10]}",
|
||||
)
|
||||
|
||||
final_path = self._assemble(session)
|
||||
rel = str(final_path.relative_to(self.upload_root))
|
||||
|
||||
task = Task(
|
||||
filename=session.filename,
|
||||
source_path=rel,
|
||||
status="queued",
|
||||
progress=0.0,
|
||||
)
|
||||
self.db.add(task)
|
||||
session.status = "completed"
|
||||
session.final_path = rel
|
||||
session.updated_at = datetime.now(timezone.utc)
|
||||
self.db.commit()
|
||||
self.db.refresh(task)
|
||||
# 正向关联:session → task(替代旧的 source_path 反向查找)
|
||||
session.task_id = task.id
|
||||
self.db.commit()
|
||||
|
||||
# 清理分片暂存
|
||||
self._cleanup_session_dir(upload_id)
|
||||
|
||||
logger.info("上传完成 task_id=%s file=%s size=%d", task.id, session.filename, session.size_bytes)
|
||||
return CompleteResponse(
|
||||
task_id=task.id, filename=task.filename,
|
||||
size_bytes=session.size_bytes, status=task.status,
|
||||
)
|
||||
|
||||
# ---------------- 过期会话清理 ----------------
|
||||
|
||||
def reap_stale_sessions(self) -> int:
|
||||
"""清理被放弃的会话:pending 且 updated_at 超过 ttl。
|
||||
|
||||
时间比较统一用 aware UTC datetime,避免 naive datetime 的 .timestamp()
|
||||
按本地时区算导致的偏移。
|
||||
"""
|
||||
cutoff = datetime.now(timezone.utc) - timedelta(seconds=self.session_ttl)
|
||||
sessions = (
|
||||
self.db.query(UploadSession)
|
||||
.filter(UploadSession.status == "pending")
|
||||
.all()
|
||||
)
|
||||
n = 0
|
||||
for session in sessions:
|
||||
updated = session.updated_at
|
||||
if updated is None:
|
||||
continue
|
||||
# SQLite 存 naive datetime,统一补 UTC 后比较
|
||||
if updated.tzinfo is None:
|
||||
updated = updated.replace(tzinfo=timezone.utc)
|
||||
if updated < cutoff:
|
||||
self._cleanup_session_dir(session.upload_id)
|
||||
self.db.delete(session)
|
||||
n += 1
|
||||
logger.info("清理过期分片会话 upload_id=%s file=%s", session.upload_id, session.filename)
|
||||
if n:
|
||||
self.db.commit()
|
||||
return n
|
||||
|
||||
# ---------------- 内部 ----------------
|
||||
|
||||
def _require_session(self, upload_id: str) -> UploadSession:
|
||||
if not upload_id:
|
||||
raise HTTPException(400, "upload_id 不能为空")
|
||||
session = (
|
||||
self.db.query(UploadSession)
|
||||
.filter(UploadSession.upload_id == upload_id)
|
||||
.first()
|
||||
)
|
||||
if session is None:
|
||||
raise HTTPException(404, f"会话不存在或已过期:{upload_id}")
|
||||
return session
|
||||
|
||||
@staticmethod
|
||||
def _validate_index(session: UploadSession, index: int) -> None:
|
||||
if index < 0 or index >= session.total_chunks:
|
||||
raise HTTPException(400, f"分片下标越界:{index} 不在 [0, {session.total_chunks})")
|
||||
|
||||
def _session_dir(self, upload_id: str) -> Path:
|
||||
return self.work_root / upload_id
|
||||
|
||||
def _assemble(self, session: UploadSession) -> Path:
|
||||
"""按 index 顺序拼接全部分片为正式视频文件。"""
|
||||
ext = Path(session.filename).suffix or ".mp4"
|
||||
now = datetime.now(timezone.utc)
|
||||
sub = self.upload_root / f"{now:%Y}" / f"{now:%m}"
|
||||
sub.mkdir(parents=True, exist_ok=True)
|
||||
final = sub / f"{uuid.uuid4().hex}{ext}"
|
||||
part_path = final.with_suffix(final.suffix + ".part")
|
||||
|
||||
session_dir = self._session_dir(session.upload_id)
|
||||
try:
|
||||
with part_path.open("wb") as out:
|
||||
for index in range(session.total_chunks):
|
||||
chunk_path = session_dir / f"{index}.part"
|
||||
if not chunk_path.is_file():
|
||||
raise HTTPException(409, f"拼接时发现分片缺失:{index}.part")
|
||||
with chunk_path.open("rb") as src:
|
||||
while buf := src.read(self.chunk_bytes):
|
||||
out.write(buf)
|
||||
out.flush()
|
||||
os.fsync(out.fileno())
|
||||
os.replace(part_path, final)
|
||||
except Exception:
|
||||
part_path.unlink(missing_ok=True)
|
||||
raise
|
||||
return final
|
||||
|
||||
def _cleanup_session_dir(self, upload_id: str) -> None:
|
||||
session_dir = self._session_dir(upload_id)
|
||||
try:
|
||||
if session_dir.exists():
|
||||
shutil.rmtree(session_dir, ignore_errors=True)
|
||||
except Exception as exc: # pragma: no cover
|
||||
logger.warning("清理会话目录失败 upload_id=%s: %s", upload_id, exc)
|
||||
Reference in New Issue
Block a user