diff --git a/Dockerfile b/Dockerfile index e7e0bd3..a3c0dd3 100644 --- a/Dockerfile +++ b/Dockerfile @@ -100,7 +100,9 @@ VOLUME ["/data", "/models"] # "Unable to load libcudnn_ops.so.9"。放在 final/dev 而非 deps,避免 ENV 变化 # 导致 deps 的 apt/pip 层缓存失效。CPU 镜像无此目录,路径被忽略不影响。 ENV CONFIG_PATH=/app/config.yaml \ - LD_LIBRARY_PATH=/usr/local/lib/python3.12/site-packages/nvidia/cudnn/lib:/usr/local/lib/python3.12/dist-packages/nvidia/cudnn/lib:/usr/local/lib/python3.12/site-packages/nvidia/cublas/lib:/usr/local/lib/python3.12/dist-packages/nvidia/cublas/lib + LD_LIBRARY_PATH=/usr/local/lib/python3.12/site-packages/nvidia/cudnn/lib:/usr/local/lib/python3.12/dist-packages/nvidia/cudnn/lib:/usr/local/lib/python3.12/site-packages/nvidia/cublas/lib:/usr/local/lib/python3.12/dist-packages/nvidia/cublas/lib \ + HF_HUB_OFFLINE=1 \ + TRANSFORMERS_OFFLINE=1 EXPOSE 8000 CMD ["python", "-m", "uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "8000", "--reload"] @@ -118,7 +120,9 @@ COPY config.example.yaml /app/config.example.yaml # 全部走 volume,镜像本身无状态、无敏感数据 VOLUME ["/data", "/models"] ENV CONFIG_PATH=/app/config.yaml \ - LD_LIBRARY_PATH=/usr/local/lib/python3.12/site-packages/nvidia/cudnn/lib:/usr/local/lib/python3.12/dist-packages/nvidia/cudnn/lib:/usr/local/lib/python3.12/site-packages/nvidia/cublas/lib:/usr/local/lib/python3.12/dist-packages/nvidia/cublas/lib + LD_LIBRARY_PATH=/usr/local/lib/python3.12/site-packages/nvidia/cudnn/lib:/usr/local/lib/python3.12/dist-packages/nvidia/cudnn/lib:/usr/local/lib/python3.12/site-packages/nvidia/cublas/lib:/usr/local/lib/python3.12/dist-packages/nvidia/cublas/lib \ + HF_HUB_OFFLINE=1 \ + TRANSFORMERS_OFFLINE=1 EXPOSE 8000 CMD ["python", "-m", "uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "8000"] diff --git a/app/config.py b/app/config.py index 3369ed3..3d505ac 100644 --- a/app/config.py +++ b/app/config.py @@ -128,13 +128,111 @@ def _load_yaml(path: Path) -> dict: return yaml.safe_load(path.read_text(encoding="utf-8")) or {} +# ---------------- 运行时覆盖 ---------------- +# 允许通过设置页修改的配置项(点分路径 -> 类型)。config.yaml 是只读挂载, +# 改它需重启容器;运行时覆盖存 DB,进程重启后自动加载,无需重建镜像。 +# 设置页保存时调 save_setting() 写 DB + 清 lru_cache,下次 get_settings() 生效。 +_applying_overrides = False # 防递归标志:_apply_overrides 内部 DB 初始化会回调 get_settings() +_OVERIDEABLE_FIELDS: dict[str, type] = { + "asr.batch_size": int, + "asr.beam_size": int, + "translation.batch_size": int, + "translation.sort_by_length": bool, +} + + +def _apply_overrides(settings: Settings) -> Settings: + """从 DB 读取覆盖值并应用到 Settings 对象。 + + 在 lru_cache 的 get_settings() 内部调用,保证缓存的对象已含覆盖。 + DB 还没初始化时(首次 import)静默跳过,用 YAML 原值。 + + 注意:get_session_local() -> get_engine() -> _db_path() -> get_settings() + 会形成递归。用 _applying_overrides 标志阻断:递归调用直接返回当前 settings + (此时 DB 路径只需 work_dir,无覆盖也无妨)。 + """ + global _applying_overrides + if _applying_overrides: + return settings # 递归调用(_db_path 触发),直接返回 YAML 原值 + _applying_overrides = True + try: + from .database import get_session_local + from .models.setting import Setting + import json + db = get_session_local()() + try: + rows = db.query(Setting).all() + overrides = {r.key: r.value for r in rows} + finally: + db.close() + for key, type_ in _OVERIDEABLE_FIELDS.items(): + if key not in overrides: + continue + try: + val = json.loads(overrides[key]) + val = type_(val) + except (json.JSONDecodeError, ValueError, TypeError): + continue + _set_nested(settings, key, val) + except Exception: + # DB 未就绪(首次 import 时 database.py 可能还在初始化)-> 跳过,用 YAML 原值 + pass + finally: + _applying_overrides = False + return settings + + +def _set_nested(settings: Settings, key: str, val) -> None: + """按点分路径设置嵌套属性,如 'asr.batch_size' -> settings.asr.batch_size""" + parts = key.split(".") + obj = settings + for p in parts[:-1]: + obj = getattr(obj, p) + setattr(obj, parts[-1], val) + + @lru_cache(maxsize=1) def get_settings() -> Settings: + """读取 config.yaml + 应用 DB 覆盖,返回完整 Settings。 + + 结果被 lru_cache 缓存。修改设置后调 reload_settings() 清缓存, + 下次调用返回含新值的 Settings。 + """ path = Path(os.getenv("CONFIG_PATH", str(DEFAULT_CONFIG_PATH))) - return Settings.model_validate(_load_yaml(path)) + settings = Settings.model_validate(_load_yaml(path)) + return _apply_overrides(settings) def reload_settings() -> Settings: - """清缓存并重新读取,供脚本与测试使用。""" + """清缓存并重新读取(含 DB 覆盖),供设置页保存后调用。""" get_settings.cache_clear() return get_settings() + + +def save_setting(key: str, value) -> None: + """保存单个配置项覆盖到 DB + 清 lru_cache。 + + Args: + key: 点分路径,必须在 _OVERIDEABLE_FIELDS 中 + value: 要保存的值(自动 JSON 编码) + """ + import json + if key not in _OVERIDEABLE_FIELDS: + raise ValueError(f"不允许修改的配置项:{key}") + from .database import get_session_local + from .models.setting import Setting + type_ = _OVERIDEABLE_FIELDS[key] + encoded = json.dumps(type_(value)) + db = get_session_local()() + try: + row = db.get(Setting, key) + if row is None: + row = Setting(key=key, value=encoded) + db.add(row) + else: + row.value = encoded + db.commit() + finally: + db.close() + # 清缓存,让后续 get_settings() 读到新值 + get_settings.cache_clear() diff --git a/app/controllers/__init__.py b/app/controllers/__init__.py index 884ebcc..bd66627 100644 --- a/app/controllers/__init__.py +++ b/app/controllers/__init__.py @@ -1,7 +1,8 @@ """路由聚合:导出各 controller 的 router,供 main.py include。""" from .log_router import router as log_router +from .settings_router import router as settings_router from .task_router import router as task_router from .upload_router import router as upload_router -__all__ = ["log_router", "task_router", "upload_router"] +__all__ = ["log_router", "settings_router", "task_router", "upload_router"] diff --git a/app/controllers/settings_router.py b/app/controllers/settings_router.py new file mode 100644 index 0000000..b9d6eef --- /dev/null +++ b/app/controllers/settings_router.py @@ -0,0 +1,92 @@ +"""设置路由:查询/修改运行时可调参数。 + +config.yaml 是只读挂载,改它需重启容器。本路由把部分参数(batch_size 等) +存到 DB 的 setting 表,通过 config.save_setting() + reload_settings() 实现 +运行时热更新:保存后清 lru_cache,后续任务读到新值。 + +当前可调项与 config._OVERIDEABLE_FIELDS 对齐。 +""" + +from __future__ import annotations + +import logging + +from fastapi import APIRouter +from pydantic import BaseModel, Field + +from ..config import _OVERIDEABLE_FIELDS, get_settings, save_setting + +logger = logging.getLogger("audio2text.settings") +router = APIRouter(prefix="/api/settings", tags=["settings"]) + + +# ---------------- 响应 / 请求 DTO ---------------- + +class SettingsResponse(BaseModel): + """当前生效的设置值(YAML 基础 + DB 覆盖后的合并值)。""" + asr_batch_size: int + asr_beam_size: int + translation_batch_size: int + translation_sort_by_length: bool + # 不可改但展示的只读信息 + asr_model: str + asr_device: str + asr_compute_type: str + translation_model: str + translation_device: str + + +class SettingsUpdate(BaseModel): + """设置更新请求:只传要改的字段,未传的保持不变。""" + asr_batch_size: int | None = Field(default=None, ge=1, le=128) + asr_beam_size: int | None = Field(default=None, ge=1, le=10) + translation_batch_size: int | None = Field(default=None, ge=1, le=256) + translation_sort_by_length: bool | None = None + + +# ---------------- 字段映射:DTO 字段名 -> config 点分路径 ---------------- + +_FIELD_MAP: dict[str, str] = { + "asr_batch_size": "asr.batch_size", + "asr_beam_size": "asr.beam_size", + "translation_batch_size": "translation.batch_size", + "translation_sort_by_length": "translation.sort_by_length", +} + + +# ---------------- 接口 ---------------- + +@router.get("", summary="查询当前生效的设置") +def get_current_settings() -> SettingsResponse: + """返回当前生效的设置(YAML 基础 + DB 覆盖合并后的值)。""" + s = get_settings() + return SettingsResponse( + asr_batch_size=s.asr.batch_size, + asr_beam_size=s.asr.beam_size, + translation_batch_size=s.translation.batch_size, + translation_sort_by_length=s.translation.sort_by_length, + asr_model=s.asr.model, + asr_device=s.asr.device, + asr_compute_type=s.asr.compute_type, + translation_model=s.translation.model, + translation_device=s.translation.device, + ) + + +@router.put("", summary="更新设置(保存后对后续任务生效)") +def update_settings(req: SettingsUpdate) -> dict: + """保存修改的设置项到 DB,清配置缓存。 + + 只处理请求中非 None 的字段。保存后立即生效(后续任务读到新值), + 已在跑的任务不受影响(任务在各阶段开始时读配置)。 + """ + changed: dict = {} + for field, path in _FIELD_MAP.items(): + val = getattr(req, field) + if val is not None: + save_setting(path, val) + changed[field] = val + logger.info("设置已更新:%s = %s(对后续任务生效)", path, val) + if not changed: + return {"status": "no_change", "changed": {}} + return {"status": "saved", "changed": changed} diff --git a/app/controllers/task_router.py b/app/controllers/task_router.py index 369f649..cd34727 100644 --- a/app/controllers/task_router.py +++ b/app/controllers/task_router.py @@ -1,7 +1,9 @@ -"""任务路由:列表 / 状态 / 下载字幕。""" +"""任务路由:列表 / 状态 / 下载字幕 / 删除。""" from __future__ import annotations +import logging +import shutil from datetime import datetime, timezone from pathlib import Path @@ -11,10 +13,11 @@ from sqlalchemy.orm import Session from ..config import get_settings from ..database import get_db -from ..models.task import Task, STATUS_UPLOADING +from ..models.task import Task, STATUS_UPLOADING, STATUS_DONE, STATUS_FAILED from ..models.upload_session import UploadSession from ..schemas.task import TaskListResponse, TaskResponse +logger = logging.getLogger("audio2text.tasks") router = APIRouter(prefix="/api/tasks", tags=["task"]) @@ -129,3 +132,47 @@ def download_subtitle( media_type="application/x-subrip", filename=download_name, ) + + +@router.delete("/{task_id}", summary="删除任务(仅允许已完成/失败)") +def delete_task(task_id: int, db: Session = Depends(get_db)) -> dict: + """删除任务及其产物(字幕 / 中间音频 / 保留的原始视频)+ DB 记录。 + + 仅允许删除已完成(done)或失败(failed)的任务,进行中的任务不可删。 + """ + task = db.get(Task, task_id) + if task is None: + raise HTTPException(404, f"任务不存在:{task_id}") + if task.status not in (STATUS_DONE, STATUS_FAILED): + raise HTTPException(409, f"任务进行中,无法删除(当前状态:{task.status})") + + s = get_settings() + deleted: list[str] = [] + + # 删字幕输出目录 + out_dir = s.output_dir() / f"task_{task.id}" + if out_dir.is_dir(): + shutil.rmtree(out_dir, ignore_errors=True) + deleted.append("outputs") + + # 删中间音频 + if task.wav_path: + wav = Path(task.wav_path) + if wav.is_file(): + wav.unlink(missing_ok=True) + deleted.append("audio") + + # 删保留的原始视频 + if task.source_path: + src = s.upload_dir() / task.source_path + if src.is_file(): + src.unlink(missing_ok=True) + deleted.append("video") + + # 删 DB 记录(先删关联的 UploadSession,再删 Task) + db.query(UploadSession).filter(UploadSession.task_id == task.id).delete() + db.delete(task) + db.commit() + + logger.info("删除任务 %d(%s):%s", task_id, task.filename, ", ".join(deleted) or "无产物") + return {"status": "deleted", "task_id": task_id, "cleaned": deleted} diff --git a/app/database.py b/app/database.py index 4dae34e..e9af836 100644 --- a/app/database.py +++ b/app/database.py @@ -58,6 +58,7 @@ def init_db_schema() -> None: """ from .models.task import Task # noqa: F401 from .models.upload_session import UploadSession # noqa: F401 + from .models.setting import Setting # noqa: F401 engine = get_engine() Base.metadata.create_all(engine) diff --git a/app/main.py b/app/main.py index 98eee80..49d1778 100644 --- a/app/main.py +++ b/app/main.py @@ -19,21 +19,21 @@ import logging import threading from contextlib import asynccontextmanager -from fastapi import Depends, FastAPI +from fastapi import FastAPI from fastapi.openapi.docs import get_redoc_html, get_swagger_ui_html from fastapi.responses import HTMLResponse, JSONResponse from sqlalchemy.orm import Session from .config import get_settings -from .controllers import log_router, task_router, upload_router +from .controllers import log_router, settings_router, task_router, upload_router from .database import get_db, init_db_schema -from .security import require_docs_auth from .services.cache_cleaner import purge_expired_cache, run_forever as run_cache_cleaner from .services.log_buffer import init_log_buffer from .services.reaper import reap_stale_sessions from .views.history_html import render as render_history_html from .views.home_html import render as render_home_html from .views.logs_html import render as render_logs_html +from .views.settings_html import render as render_settings_html # 日志分层: # - audio2text logger 始终设 DEBUG,确保所有记录(含子步骤)都能产生。 @@ -126,20 +126,21 @@ def create_app() -> FastAPI: app.include_router(upload_router) app.include_router(task_router) app.include_router(log_router) + app.include_router(settings_router) - # 受 Basic Auth 保护的文档接口 + # 文档接口(无认证,直接公开) @app.get("/openapi.json") - def protected_openapi(_: str = Depends(require_docs_auth)) -> JSONResponse: + def openapi_endpoint() -> JSONResponse: return JSONResponse(app.openapi()) @app.get("/docs") - def protected_docs(_: str = Depends(require_docs_auth)): + def docs_endpoint(): return get_swagger_ui_html( openapi_url="/openapi.json", title="audio2text docs", swagger_favicon_url="" ) @app.get("/redoc") - def protected_redoc(_: str = Depends(require_docs_auth)): + def redoc_endpoint(): return get_redoc_html( openapi_url="/openapi.json", title="audio2text docs", redoc_favicon_url="" ) @@ -187,6 +188,10 @@ def create_app() -> FastAPI: def logs_page() -> HTMLResponse: return HTMLResponse(render_logs_html()) + @app.get("/settings", response_class=HTMLResponse) + def settings_page() -> HTMLResponse: + return HTMLResponse(render_settings_html()) + return app diff --git a/app/models/setting.py b/app/models/setting.py new file mode 100644 index 0000000..2e0646f --- /dev/null +++ b/app/models/setting.py @@ -0,0 +1,36 @@ +"""运行时设置覆盖(键值存储)。 + +config.yaml 是只读挂载(镜像内不含配置),改完需重启容器才生效。 +本表持久化用户在「设置页」修改的参数,进程重启后自动加载, +无需改 config.yaml 或重建镜像。 + +当前支持的键见 _ALLOWED_KEYS(settings_router 维护),值为 JSON 字符串。 +""" + +from __future__ import annotations + +from datetime import datetime, timezone + +from sqlalchemy import String, DateTime, Text +from sqlalchemy.orm import Mapped, mapped_column + +from ..database import Base + + +def _now() -> datetime: + return datetime.now(timezone.utc) + + +class Setting(Base): + """单个配置项的覆盖值(key = 'asr.batch_size' 之类的点分路径)。""" + + __tablename__ = "setting" + + key: Mapped[str] = mapped_column(String(128), primary_key=True) + value: Mapped[str] = mapped_column(Text) # JSON 编码的值 + updated_at: Mapped[datetime] = mapped_column( + DateTime, default=_now, onupdate=_now, + ) + + def __repr__(self) -> str: + return f"Setting(key={self.key!r}, value={self.value!r})" diff --git a/app/views/_shared.py b/app/views/_shared.py index e6f1109..5cf0eb0 100644 --- a/app/views/_shared.py +++ b/app/views/_shared.py @@ -19,6 +19,7 @@ _NAV_ITEMS = [ ("/", "主页", "home"), ("/history", "历史", "history"), ("/logs", "日志", "logs"), + ("/settings", "设置", "settings"), ] diff --git a/app/views/history_html.py b/app/views/history_html.py index 1036692..2851e35 100644 --- a/app/views/history_html.py +++ b/app/views/history_html.py @@ -82,7 +82,12 @@ function renderTable(tasks) {{ }} else if (task.status === "failed") {{ action = `查看错误`; }} else {{ - action = `—`; + action = `-`; + }} + // done/failed 且非上传中:加删除按钮 + const canDelete = (task.status === "done" || task.status === "failed") && !task.is_upload; + if (canDelete) {{ + action += ` `; }} let progress; @@ -117,6 +122,17 @@ function renderPagination() {{ paginationEl.innerHTML = html; }} +async function deleteTask(taskId) {{ + if (!confirm("确认删除任务 #" + taskId + "?字幕和中间文件将被清除。")) return; + try {{ + const r = await fetch("/api/tasks/" + taskId, {{ method: "DELETE" }}); + if (!r.ok) {{ alert("删除失败:" + await r.text()); return; }} + load(currentOffset); + }} catch (e) {{ + alert("删除失败:" + e); + }} +}} + load(0); """ diff --git a/app/views/home_html.py b/app/views/home_html.py index c743354..f4f1afb 100644 --- a/app/views/home_html.py +++ b/app/views/home_html.py @@ -77,6 +77,12 @@ function renderTaskInner(task) {{ try {{ created = fmtDateTime24(new Date(task.created_at + "Z")); }} catch (e) {{ created = task.created_at; }} + // 删除按钮:仅 done/failed 且非上传中任务显示 + const canDelete = (task.status === "done" || task.status === "failed") && !task.is_upload; + const delBtn = canDelete + ? `` + : ""; + let body; if (task.status === "done") {{ body = `
`; @@ -91,11 +97,30 @@ function renderTaskInner(task) {{调整批处理大小等运行时参数。保存后对后续任务生效,已在运行的任务不受影响。
+ +