diff --git a/.env.example b/.env.example index fcc1cb9..f78c1cc 100644 --- a/.env.example +++ b/.env.example @@ -33,7 +33,7 @@ AI_PUBLISH_PROVIDER=codex AI_CODEX_PATH=codex AI_CODEX_HOME= AI_CODEX_MODEL=gpt-5.6-sol -AI_CODEX_TIMEOUT_SECONDS=300 +AI_CODEX_TIMEOUT_SECONDS=600 AI_REQUEST_TIMEOUT_SECONDS=120 # 远程 OpenAI-compatible / DeepSeek:文字稿分析 @@ -68,10 +68,22 @@ AI_LOCAL_HEALTH_TIMEOUT_SECONDS=30 # 语音转写 # ======================== -# volcengine:远程火山引擎;local:本地 faster-whisper。 -TRANSCRIPTION_PROVIDER=volcengine +# 默认完全离线,本地音频不会上传;火山配置仅保留用于手动回滚。 +TRANSCRIPTION_PROVIDER=local TRANSCRIPTION_FALLBACK_PROVIDER= +TRANSCRIPTION_OFFLINE_ONLY=true +TRANSCRIPTION_MODEL=large-v3 +TRANSCRIPTION_MODEL_REVISION=edaa852ec7e145841d8ffdb056a99866b5f0a478 +TRANSCRIPTION_MODEL_CACHE_DIR=E:\直播间切片工作流存储\_模型\faster-whisper +TRANSCRIPTION_LOCAL_FILES_ONLY=true +TRANSCRIPTION_LANGUAGE=zh +TRANSCRIPTION_DEVICE=cuda +TRANSCRIPTION_COMPUTE_TYPE=float16 +TRANSCRIPTION_CPU_FALLBACK_MODEL=medium +TRANSCRIPTION_CHUNK_SECONDS=120 +TRANSCRIPTION_CHUNK_OVERLAP_SECONDS=5 +# 以下值不会在完全离线模式下使用,也不会被自动删除。 VOLCENGINE_ASR_API_URL=https://openspeech.bytedance.com/api/v3/auc/bigmodel/recognize/flash VOLCENGINE_ASR_API_KEY= VOLCENGINE_ASR_APP_KEY= @@ -80,14 +92,6 @@ VOLCENGINE_ASR_RESOURCE_ID=volc.bigasr.auc_turbo VOLCENGINE_ASR_TIMEOUT_SECONDS=300 VOLCENGINE_ASR_AUDIO_FORMAT=mp3 -TRANSCRIPTION_MODEL=medium -TRANSCRIPTION_LANGUAGE=zh -TRANSCRIPTION_DEVICE=cpu -TRANSCRIPTION_COMPUTE_TYPE=int8 -TRANSCRIPTION_CPU_FALLBACK_MODEL=medium -TRANSCRIPTION_CHUNK_SECONDS=120 -TRANSCRIPTION_CHUNK_OVERLAP_SECONDS=5 - # ======================== # 排期与 Windows Chrome Worker # ======================== diff --git a/DEVELOPMENT_LOG.md b/DEVELOPMENT_LOG.md index 4786d1a..e0c8728 100644 --- a/DEVELOPMENT_LOG.md +++ b/DEVELOPMENT_LOG.md @@ -1,5 +1,33 @@ # Development Log +## 2026-09-01 最新任务恢复与本地转写简体化 + +- 本地 `faster-whisper` 的段落和逐词文本在写入分段 checkpoint 前统一经过 OpenCC `t2s.json`;只转换繁简字形,继续保留“计程车、软体”等台湾用词,英文、数字和时间戳不变。 +- 固定新增 `opencc==1.4.2`;本地转写 checkpoint 模型身份加入 `opencc-t2s-v1`,旧繁体分块不会被新任务静默复用。转写进度 JSON 同步记录实际文本标准化版本。 +- AI 完整性门禁前移到 `AI_ANALYZING`:覆盖率或单元不完整会直接进入 `FAILED_AI_ANALYZING`;`CLIP_SELECTING` 保留二次防护,不会继续生成切片或发布记录。 +- `/api/tasks/{task_id}/process/auto-retry` 在发现计费结果不确定的 AI 单元时先返回结构化 409。任务详情会明确说明风险并要求二次确认;取消时旧 Job 和全部证据保持不变。 +- 输入未变化时,确认后原位重置不确定单元并复用成功 AI checkpoint;transcript 已变化时,从 `AI_ANALYZING` 创建独立新 Job,旧失败 Job 不修改。有任何下游切片或发布记录时继续 fail-closed。 +- Codex 单窗口默认超时由 300 秒调整为 600 秒;新增 transcript 原位转换工具,转换前写入 `transcripts/backups/`,并校验时间戳、行数和 Markdown 结构。 +- 专项回归 `66 passed`、最终完整回归 `825 passed`;Ruff、Python compileall、JavaScript 语法和 `git diff --check` 均通过。`release_gate.ps1` 因当前分支没有正式版 `acceptance-results/latest.json` 按设计拒绝放行;正式服务部署和任务 `614fb38401e0` 的现场恢复继续单独验收。 +- 已把两项代码提交同步到运行工作树 `codex/runtime-offline-transcription-cutover`,项目环境确认安装 `opencc==1.4.2`;运行 `.env` 先备份,再只把 `AI_CODEX_TIMEOUT_SECONDS` 从 300 调整为 600。 +- 活动 SQLite 已通过 Online Backup 保存到 `data/backups/workflow-before-latest-ai-retry-20260901-165534-858701-d0d0492a.sqlite3`;备份及活动库均为 `quick_check=ok`、外键违规 0。恢复前任务无活动 Job、无切片、无发布记录。 +- 任务 `614fb38401e0` 的原 transcript 已备份到任务目录 `transcripts/backups/transcript.before-t2s-20260901-165555.md`,原 SHA-256 为 `d4f17ff7...d09f9d`;原位转换后的 SHA-256 为 `c666a5dc...43f18d`,1407 行、2781 个时间戳和 Markdown 结构保持不变。 +- 只通过 Alter 重启了 `Niuma-Studio` Web:8001 Listener 更新为 PID `151000`,祖先进程落在运行工作树;8765 发布 Worker Listener 仍为 PID `39056`。`/health`、深度 readiness、Scheduler 和 Worker 均健康。 +- 未确认重试先按设计返回 `409 / ai_retry_confirmation_required`;显式确认后创建独立 Job `dec329256a0b`,旧失败 Job `434da2a9bc97` 及证据保持不变。新 AI Run `6c497a6ac3f7` 完成 `18/18` 单元、覆盖率 100%,随后生成 5 个切片并停在 `PENDING_SUBTITLE_REVIEW`;发布记录仍为 0,未触发抖音或 B站投稿。 + +## 2026-08-31 完全离线长音频转写 + +- 默认转写 Provider 改为本地 `faster-whisper`,固定 `large-v3@edaa852e / cuda / float16`;`TRANSCRIPTION_OFFLINE_ONLY=true` 时,任务页、自动流水线、重试入口和显式火山请求都会在远程 HTTP 请求前拒绝执行。 +- 新增固定版本模型缓存和 Windows CUDA 运行时管理:模型只从 `E:\直播间切片工作流存储\_模型\faster-whisper` 读取,任务运行中不会偷偷联网下载;当前进程会加载项目 `.venv` 内的 cuBLAS 12 DLL,并复用 CTranslate2 自带的 cuDNN 9。 +- 新增一次性初始化脚本 `scripts/setup_local_transcription.ps1`,锁定下载 `large-v3` GPU 主模型和 `medium` CPU/int8 兜底模型,写入原子模型清单,并支持真实 20 秒音频的 GPU 推理冒烟。 +- 新增 `scripts/benchmark_local_transcription.py`,使用隔离数据库和任务目录执行 10/40 分钟本地验收;支持在指定分块保存 checkpoint 后主动停止,并用同一目录原地续跑,不接触正式任务数据。 +- 保留原 120 秒分块、5 秒重叠和 SQLite checkpoint;checkpoint 在实际 GPU/CPU 模型加载完成后才创建,`transcription_runs.model` 使用带 revision 的模型身份,音频指纹改为完整 SHA-256。 +- `/api/tasks/{id}/transcript-status` 增加离线锁、模型缓存、GPU 和 revision 状态;系统状态页展示完全离线、模型/DLL 就绪度、实际配置设备和外部转写费用 `¥0`,火山配置保留作人工回滚且不删除。 +- 新增离线边界专项回归,覆盖默认 Provider、409 阻断、HTTP 前置拦截、模型缺失、DLL 路径、CPU 兜底、模型 revision 与完整音频指纹;正式 `.env` 和运行中 Web/Worker 尚未切换或重启,等待真实素材验收后再启用。 +- 真实验收已验证主模型 `model.bin` 大小 `3,087,284,237` 字节及官方 SHA-256 `69f74147e3334731bc3a76048724833325d2ec74642fb52620eda87352e3d4f1`;20 秒中文音频 GPU 推理成功,没有退回 CPU。`medium` 兜底模型的 `1,527,906,378` 字节大文件也通过官方 SHA-256 `9b45e1009dcc4ab601eff815b61d80e60ce3fd8c74c1a14f4a282258286b51ae`,并实际以 `cpu / int8` 加载成功;就绪检查同时兼容官方模型使用的 `vocabulary.json` 或 `vocabulary.txt`。 +- 同一份 41 分 18 秒含多人对白和音乐素材,10 分钟样本用时 83.2 秒并输出 336 条;完整素材先在第 1/22 分块后主动中断,再复用该 checkpoint,用时 291.85 秒完成其余分块并输出 1085 条。SQLite 记录为 `large-v3@edaa852e / cuda / float16`、22/22 完成、完整 64 位 SHA-256;运行中转写进程 TCP 连接数为 0,GPU 采样为 65% 且约占 8.1 GB 显存。 +- 已在 `E:\直播间切片工作流存储\_验收\offline-transcription-20260831\41min\spot-checks` 均匀生成 10 个 15 秒试听片段和人工清单;人名、关键对白和剪辑时间点仍待用户试听确认,因此没有切换正式 `.env`,也没有重启 Web/Worker。 + ## 2026-08-28 2.1 集成 PR 与 Docker 冒烟修复 - 将当前线性领先 `master` 的 26 个提交完整保留到 `codex/integrate-v2.1-stable`,新增独立空白清理提交并创建顶层集成 PR #60;不 rebase、不 squash,也不自动合并。 diff --git a/NEXT_STEPS.md b/NEXT_STEPS.md index 54e6cc9..b8e8930 100644 --- a/NEXT_STEPS.md +++ b/NEXT_STEPS.md @@ -1,5 +1,20 @@ # Next Steps +## 2026-09-01 最新任务恢复验收 + +1. [x] 运行工作树已同步代码并固定安装 OpenCC;运行 `.env` 只把 Codex 单窗口超时调整为 600 秒,原配置已备份。 +2. [x] 活动 SQLite、当前 transcript 均已先备份;数据库完整性正常,transcript 已转为仅字形简体且结构保持不变。 +3. [x] 只通过 Alter 重启 `Niuma-Studio` Web;8001 新进程、祖先链、深度 readiness 和 Scheduler 正常,8765 发布 Worker 未重启且健康。 +4. [x] 独立 Job `dec329256a0b` 已完成,AI 覆盖率 `18/18(100%)`,生成 5 个切片并停在 `PENDING_SUBTITLE_REVIEW`;旧失败 Job 保留,发布记录为 0。 +5. 下一步由用户在任务详情页人工审核字幕和切片。不要点击立即发送,也不要创建或修改抖音/B站排期;PR #72 等 CI 通过后仍需用户确认才可合并。 + +## 2026-08-31 完全离线转写验收与启用 + +1. 固定版本 cuBLAS 12、`large-v3` 主模型和 `medium / CPU / int8` 兜底模型均已初始化并通过加载校验;20 秒、10 分钟和完整 41 分钟 GPU 验收均通过,完整素材续跑约 4 分 52 秒完成并成功复用第一个 checkpoint。 +2. 打开 `E:\直播间切片工作流存储\_验收\offline-transcription-20260831\41min\spot-checks\README.md`,依次试听 10 个 15 秒片段,核对人名、关键对白和剪辑起止点;第 6、10 段已标出需要重点确认的疑似人名/背景声识别。 +3. 只有人工试听确认后,才把正式 `.env` 切换为本地离线配置,并在记录进程、端口和健康状态后受控重启 Web/Worker。当前正式服务仍保持原状。 +4. 回滚时只需恢复旧 `TRANSCRIPTION_PROVIDER` 并设置 `TRANSCRIPTION_OFFLINE_ONLY=false`;不要删除火山配置、E 盘模型缓存、历史 transcript 或 SQLite checkpoint。 + ## 2026-08-28 2.1 集成收口 1. 等待 PR #60 的 Linux、Windows 与 Docker 三组 CI 全部通过;Docker 页面冒烟必须确认 `/`、`/tasks`、`/clips`、`/publish` 均返回成功。 diff --git a/THIRD_PARTY_NOTICES.md b/THIRD_PARTY_NOTICES.md index 367c8df..027a556 100644 --- a/THIRD_PARTY_NOTICES.md +++ b/THIRD_PARTY_NOTICES.md @@ -30,6 +30,12 @@ - 许可证:MIT License。 - 本项目锁定运行时版本 `pysubs2==1.9.0`;完整许可证保存在 `third_party_licenses/pysubs2-LICENSE.txt`。 +## OpenCC 1.4.2 + +- 项目:https://github.com/BYVoid/OpenCC +- 用途:把本地 faster-whisper 的繁体字形输出转换为简体字形;本项目使用 `t2s.json`,不做词汇本地化。 +- 许可证:Apache License 2.0;本项目通过 `opencc==1.4.2` 安装包使用,安装包内包含完整许可证。 + ## wavesurfer.js 7.12.11 - 项目:https://github.com/katspaugh/wavesurfer.js diff --git a/app/core/config.py b/app/core/config.py index bac3df5..e9018a8 100644 --- a/app/core/config.py +++ b/app/core/config.py @@ -2,6 +2,12 @@ import os from pathlib import Path +from app.core.transcription_defaults import ( + CPU_FALLBACK_TRANSCRIPTION_MODEL, + PRIMARY_TRANSCRIPTION_MODEL, + PRIMARY_TRANSCRIPTION_MODEL_REVISION, +) + PROJECT_ROOT = Path(__file__).resolve().parents[2] EXTERNAL_STORAGE_ROOT = Path(r"E:\直播间切片工作流存储") @@ -88,7 +94,7 @@ class Settings: ai_codex_path: str = _env("AI_CODEX_PATH", "codex") ai_codex_home: str = _env("AI_CODEX_HOME", "") ai_codex_model: str = _env("AI_CODEX_MODEL", "gpt-5.6-sol") - ai_codex_timeout_seconds: int = int(_env("AI_CODEX_TIMEOUT_SECONDS", "300")) + ai_codex_timeout_seconds: int = int(_env("AI_CODEX_TIMEOUT_SECONDS", "600")) ai_analysis_remote_base_url: str = _env_first( ("AI_ANALYSIS_REMOTE_BASE_URL", "AI_REMOTE_BASE_URL"), @@ -221,13 +227,26 @@ class Settings: ai_windows_wsl_setup_acknowledged: str = _env("AI_WINDOWS_WSL_SETUP_ACKNOWLEDGED", "true") ai_model_context_window: int = int(_env("AI_MODEL_CONTEXT_WINDOW", "1000000")) ai_model_auto_compact_token_limit: int = int(_env("AI_MODEL_AUTO_COMPACT_TOKEN_LIMIT", "900000")) - transcription_provider: str = _env("TRANSCRIPTION_PROVIDER", "volcengine") + transcription_provider: str = _env("TRANSCRIPTION_PROVIDER", "local") transcription_fallback_provider: str = _env("TRANSCRIPTION_FALLBACK_PROVIDER", "") - transcription_model: str = _env("TRANSCRIPTION_MODEL", "medium") + transcription_offline_only: bool = _env_bool("TRANSCRIPTION_OFFLINE_ONLY", True) + transcription_model: str = _env("TRANSCRIPTION_MODEL", PRIMARY_TRANSCRIPTION_MODEL) + transcription_model_revision: str = _env( + "TRANSCRIPTION_MODEL_REVISION", + PRIMARY_TRANSCRIPTION_MODEL_REVISION, + ) + transcription_model_cache_dir: Path = _env_path( + "TRANSCRIPTION_MODEL_CACHE_DIR", + _env_path("STORAGE_ROOT", EXTERNAL_STORAGE_ROOT) / "_模型" / "faster-whisper", + ) + transcription_local_files_only: bool = _env_bool("TRANSCRIPTION_LOCAL_FILES_ONLY", True) transcription_language: str = _env("TRANSCRIPTION_LANGUAGE", "zh") - transcription_device: str = _env("TRANSCRIPTION_DEVICE", "cpu") - transcription_compute_type: str = _env("TRANSCRIPTION_COMPUTE_TYPE", "int8") - transcription_cpu_fallback_model: str = _env("TRANSCRIPTION_CPU_FALLBACK_MODEL", "medium") + transcription_device: str = _env("TRANSCRIPTION_DEVICE", "cuda") + transcription_compute_type: str = _env("TRANSCRIPTION_COMPUTE_TYPE", "float16") + transcription_cpu_fallback_model: str = _env( + "TRANSCRIPTION_CPU_FALLBACK_MODEL", + CPU_FALLBACK_TRANSCRIPTION_MODEL, + ) transcription_chunk_seconds: int = int(_env("TRANSCRIPTION_CHUNK_SECONDS", "120")) transcription_chunk_overlap_seconds: int = int(_env("TRANSCRIPTION_CHUNK_OVERLAP_SECONDS", "5")) volcengine_asr_api_url: str = _env( diff --git a/app/core/transcription_defaults.py b/app/core/transcription_defaults.py new file mode 100644 index 0000000..ad53c6a --- /dev/null +++ b/app/core/transcription_defaults.py @@ -0,0 +1,15 @@ +"""离线语音转写的固定模型与版本。""" + +PRIMARY_TRANSCRIPTION_MODEL = "large-v3" +PRIMARY_TRANSCRIPTION_MODEL_REPOSITORY = "Systran/faster-whisper-large-v3" +PRIMARY_TRANSCRIPTION_MODEL_REVISION = "edaa852ec7e145841d8ffdb056a99866b5f0a478" +PRIMARY_TRANSCRIPTION_MODEL_BIN_SIZE = 3_087_284_237 +PRIMARY_TRANSCRIPTION_MODEL_BIN_SHA256 = "69f74147e3334731bc3a76048724833325d2ec74642fb52620eda87352e3d4f1" + +CPU_FALLBACK_TRANSCRIPTION_MODEL = "medium" +CPU_FALLBACK_TRANSCRIPTION_MODEL_REPOSITORY = "Systran/faster-whisper-medium" +CPU_FALLBACK_TRANSCRIPTION_MODEL_REVISION = "08e178d48790749d25932bbc082711ddcfdfbc4f" +CPU_FALLBACK_TRANSCRIPTION_MODEL_BIN_SIZE = 1_527_906_378 +CPU_FALLBACK_TRANSCRIPTION_MODEL_BIN_SHA256 = "9b45e1009dcc4ab601eff815b61d80e60ce3fd8c74c1a14f4a282258286b51ae" + +WINDOWS_CUBLAS_PACKAGE = "nvidia-cublas-cu12==12.9.2.10" diff --git a/app/models/settings.py b/app/models/settings.py index f7b8137..21a309c 100644 --- a/app/models/settings.py +++ b/app/models/settings.py @@ -1,6 +1,13 @@ from urllib.parse import urlsplit -from pydantic import BaseModel, Field, field_validator +from pathlib import Path, PureWindowsPath + +from pydantic import BaseModel, Field, field_validator, model_validator + +from app.core.transcription_defaults import ( + PRIMARY_TRANSCRIPTION_MODEL, + PRIMARY_TRANSCRIPTION_MODEL_REVISION, +) _LOCAL_AI_HOSTS = {"localhost", "127.0.0.1", "::1", "host.docker.internal", "host.containers.internal"} @@ -46,10 +53,32 @@ class AIConfigUpdate(BaseModel): ai_codex_path: str = Field(default="codex", max_length=1000) ai_codex_home: str = Field(default="", max_length=1000) ai_codex_model: str = Field(default="gpt-5.6-sol", max_length=200) - ai_codex_timeout_seconds: int = Field(default=300, ge=10, le=1800) + ai_codex_timeout_seconds: int = Field(default=600, ge=10, le=1800) - transcription_provider: str = Field(default="volcengine", pattern="^(volcengine|local)$", max_length=20) + transcription_provider: str = Field(default="local", pattern="^(volcengine|local)$", max_length=20) transcription_fallback_provider: str = Field(default="", pattern="^(|volcengine|local)$", max_length=20) + transcription_offline_only: bool = Field(default=True) + transcription_model: str = Field( + default=PRIMARY_TRANSCRIPTION_MODEL, + pattern=r"^[A-Za-z0-9._/-]+$", + max_length=200, + ) + transcription_model_revision: str = Field( + default=PRIMARY_TRANSCRIPTION_MODEL_REVISION, + pattern=r"^[0-9a-f]{40}$", + max_length=40, + ) + transcription_model_cache_dir: str = Field( + default=r"E:\直播间切片工作流存储\_模型\faster-whisper", + max_length=2000, + ) + transcription_local_files_only: bool = Field(default=True) + transcription_device: str = Field(default="cuda", pattern="^(cuda|cpu|auto)$", max_length=20) + transcription_compute_type: str = Field( + default="float16", + pattern="^(float16|int8_float16|int8|float32|auto)$", + max_length=30, + ) volcengine_asr_api_url: str = Field(default="", max_length=2000) volcengine_asr_api_key: str = Field(default="", max_length=4096) volcengine_asr_app_key: str = Field(default="", max_length=4096) @@ -113,6 +142,25 @@ def reject_control_characters(cls, value): def validate_volcengine_url(cls, value: str) -> str: return _validate_http_url(value, https_only=True) + @field_validator("transcription_model_cache_dir") + @classmethod + def validate_transcription_model_cache_dir(cls, value: str) -> str: + text = value.strip() + if not text or not (Path(text).is_absolute() or PureWindowsPath(text).is_absolute()): + raise ValueError("本地转写模型缓存目录必须是绝对路径") + return text + + @model_validator(mode="after") + def validate_offline_transcription_policy(self): + if self.transcription_offline_only: + if self.transcription_provider != "local": + raise ValueError("完全离线转写开启时,转写方式必须为 local") + if self.transcription_fallback_provider not in {"", "local"}: + raise ValueError("完全离线转写开启时,不能配置远程转写兜底") + if not self.transcription_local_files_only: + raise ValueError("完全离线转写开启时,必须只读取本地模型文件") + return self + @field_validator(*_REMOTE_URL_FIELDS) @classmethod def validate_remote_url(cls, value: str) -> str: diff --git a/app/routers/tasks.py b/app/routers/tasks.py index a788890..9f8cde4 100644 --- a/app/routers/tasks.py +++ b/app/routers/tasks.py @@ -18,6 +18,10 @@ ) from app.services import task_service from app.services.ai_prompt_preset_service import update_task_ai_prompt_preset +from app.services.ai_retry_service import ( + AIAnalysisRetryConfirmationRequired, + AutoPipelineRetryConflictError, +) from app.services.pipeline_engine import start_auto_pipeline from app.services.storage_service import ( allocate_task_dir_name, @@ -222,10 +226,17 @@ async def process_audio(task_id: str) -> dict: async def process_transcript( task_id: str, background_tasks: BackgroundTasks, - provider: str | None = Query(default=None, pattern="^(remote|local)$"), + provider: str | None = Query(default=None, pattern="^(remote|local|volcengine)$"), ) -> dict: try: - return task_service.process_task_transcript(task_id, background_tasks=background_tasks, provider=provider) + resolved_provider = task_service.validate_transcription_provider_choice(provider) + return task_service.process_task_transcript( + task_id, + background_tasks=background_tasks, + provider=resolved_provider, + ) + except task_service.TranscriptionOfflinePolicyError as exc: + raise HTTPException(status_code=409, detail=str(exc)) from exc except ValueError as exc: raise HTTPException(status_code=400, detail=str(exc)) from exc except RuntimeError as exc: @@ -237,9 +248,10 @@ async def process_transcript_workflow( task_id: str, background_tasks: BackgroundTasks, force: bool = Query(default=False), - provider: str | None = Query(default=None, pattern="^(remote|local)$"), + provider: str | None = Query(default=None, pattern="^(remote|local|volcengine)$"), ) -> dict: try: + resolved_provider = task_service.validate_transcription_provider_choice(provider) task = task_service.get_task(task_id, include_video_probe=False) if not task: raise ValueError("任务不存在") @@ -248,7 +260,7 @@ async def process_transcript_workflow( job, created = job_service.create_or_get_active_job( task_id=task_id, job_type=job_service.JOB_TYPE_TRANSCRIPT, - payload={"force": force, "provider": provider}, + payload={"force": force, "provider": resolved_provider}, ) return { "status": job["status"], @@ -257,6 +269,8 @@ async def process_transcript_workflow( "job": job, "task": task, } + except task_service.TranscriptionOfflinePolicyError as exc: + raise HTTPException(status_code=409, detail=str(exc)) from exc except ValueError as exc: raise HTTPException(status_code=400, detail=str(exc)) from exc except RuntimeError as exc: @@ -413,9 +427,22 @@ async def process_auto_pipeline( @router.post("/{task_id}/process/auto-retry") -async def retry_auto_pipeline(task_id: str, background_tasks: BackgroundTasks) -> dict: +async def retry_auto_pipeline( + task_id: str, + background_tasks: BackgroundTasks, + confirm_uncertain_ai: bool = Query(default=False), +) -> dict: try: - return start_auto_pipeline(task_id, background_tasks=background_tasks, retry=True) + return start_auto_pipeline( + task_id, + background_tasks=background_tasks, + retry=True, + confirm_uncertain_ai=confirm_uncertain_ai, + ) + except AIAnalysisRetryConfirmationRequired as exc: + raise HTTPException(status_code=409, detail=exc.detail) from exc + except AutoPipelineRetryConflictError as exc: + raise HTTPException(status_code=409, detail=exc.detail) from exc except ValueError as exc: raise HTTPException(status_code=400, detail=str(exc)) from exc diff --git a/app/services/ai_config_service.py b/app/services/ai_config_service.py index fb4c464..8afdea8 100644 --- a/app/services/ai_config_service.py +++ b/app/services/ai_config_service.py @@ -5,6 +5,7 @@ from app.models.settings import AIConfigUpdate from app.services.ai.codex_cli_provider import CodexCliConfig, CodexCliProvider from app.services.ai.diagnostics import fetch_ollama_models +from app.services.local_transcription_runtime import get_local_transcription_runtime_status SETTING_ATTRS = { @@ -17,6 +18,13 @@ "AI_CODEX_TIMEOUT_SECONDS": "ai_codex_timeout_seconds", "TRANSCRIPTION_PROVIDER": "transcription_provider", "TRANSCRIPTION_FALLBACK_PROVIDER": "transcription_fallback_provider", + "TRANSCRIPTION_OFFLINE_ONLY": "transcription_offline_only", + "TRANSCRIPTION_MODEL": "transcription_model", + "TRANSCRIPTION_MODEL_REVISION": "transcription_model_revision", + "TRANSCRIPTION_MODEL_CACHE_DIR": "transcription_model_cache_dir", + "TRANSCRIPTION_LOCAL_FILES_ONLY": "transcription_local_files_only", + "TRANSCRIPTION_DEVICE": "transcription_device", + "TRANSCRIPTION_COMPUTE_TYPE": "transcription_compute_type", "VOLCENGINE_ASR_API_URL": "volcengine_asr_api_url", "VOLCENGINE_ASR_API_KEY": "volcengine_asr_api_key", "VOLCENGINE_ASR_APP_KEY": "volcengine_asr_app_key", @@ -82,6 +90,15 @@ "ai_model_auto_compact_token_limit", } +BOOLEAN_ATTRS = { + "transcription_offline_only", + "transcription_local_files_only", +} + +PATH_ATTRS = { + "transcription_model_cache_dir", +} + LOCAL_MODEL_OPTIONS = ["qwen3:8b", "gemma3:12b", "qwen3:14b"] REMOTE_MODEL_OPTIONS = ["deepseek-v4-flash", "deepseek-v4-pro", "deepseek-chat", "deepseek-reasoner"] REMOTE_PROTOCOL_OPTIONS = ["chat_completions", "responses"] @@ -117,10 +134,17 @@ ], ), ( - "# 1. Audio transcription - Volcengine ASR", + "# 1. Audio transcription - offline local primary, Volcengine rollback only", [ "TRANSCRIPTION_PROVIDER", "TRANSCRIPTION_FALLBACK_PROVIDER", + "TRANSCRIPTION_OFFLINE_ONLY", + "TRANSCRIPTION_MODEL", + "TRANSCRIPTION_MODEL_REVISION", + "TRANSCRIPTION_MODEL_CACHE_DIR", + "TRANSCRIPTION_LOCAL_FILES_ONLY", + "TRANSCRIPTION_DEVICE", + "TRANSCRIPTION_COMPUTE_TYPE", "VOLCENGINE_ASR_API_URL", "VOLCENGINE_ASR_API_KEY", "VOLCENGINE_ASR_APP_KEY", @@ -238,6 +262,10 @@ def _apply_runtime_values(values: dict[str, str]) -> None: os.environ[env_key] = value if attr_name in INTEGER_ATTRS: object.__setattr__(settings, attr_name, int(value)) + elif attr_name in BOOLEAN_ATTRS: + object.__setattr__(settings, attr_name, value.strip().lower() in {"1", "true", "yes", "on"}) + elif attr_name in PATH_ATTRS: + object.__setattr__(settings, attr_name, Path(value)) else: object.__setattr__(settings, attr_name, value) _sync_legacy_runtime_aliases(values) @@ -302,16 +330,21 @@ def _public_config_values(values: dict[str, str]) -> dict[str, str]: def get_ai_config_context() -> dict: values = _current_config_values() public_values = _public_config_values(values) + transcription_runtime = get_local_transcription_runtime_status() secret_configured = { key: bool(str(values.get(key) or "").strip()) for key in SECRET_SETTING_KEYS } transcription_ready = bool( - values["TRANSCRIPTION_PROVIDER"] == "local" + ( + values["TRANSCRIPTION_PROVIDER"] == "local" + and transcription_runtime["ready"] + ) or ( values["VOLCENGINE_ASR_API_URL"] and values["VOLCENGINE_ASR_RESOURCE_ID"] and (values["VOLCENGINE_ASR_API_KEY"] or values["VOLCENGINE_ASR_APP_KEY"]) + and values["TRANSCRIPTION_OFFLINE_ONLY"].strip().lower() not in {"1", "true", "yes", "on"} ) ) analysis_ready = _remote_ready(values, "AI_ANALYSIS_REMOTE") @@ -344,6 +377,7 @@ def get_ai_config_context() -> dict: "values": public_values, "secret_configured": secret_configured, "transcription_ready": transcription_ready, + "transcription_runtime": transcription_runtime, "analysis_ready": analysis_ready, "analysis_key_valid": _key_valid(values["AI_ANALYSIS_REMOTE_API_KEY"]), "publish_ready": publish_ready, @@ -398,8 +432,15 @@ def save_ai_config(payload: AIConfigUpdate) -> dict: "AI_CODEX_HOME": payload.ai_codex_home.strip(), "AI_CODEX_MODEL": payload.ai_codex_model.strip() or "gpt-5.6-sol", "AI_CODEX_TIMEOUT_SECONDS": str(payload.ai_codex_timeout_seconds), - "TRANSCRIPTION_PROVIDER": payload.transcription_provider.strip() or "volcengine", + "TRANSCRIPTION_PROVIDER": payload.transcription_provider.strip() or "local", "TRANSCRIPTION_FALLBACK_PROVIDER": payload.transcription_fallback_provider.strip(), + "TRANSCRIPTION_OFFLINE_ONLY": str(payload.transcription_offline_only).lower(), + "TRANSCRIPTION_MODEL": payload.transcription_model.strip(), + "TRANSCRIPTION_MODEL_REVISION": payload.transcription_model_revision.strip(), + "TRANSCRIPTION_MODEL_CACHE_DIR": payload.transcription_model_cache_dir.strip(), + "TRANSCRIPTION_LOCAL_FILES_ONLY": str(payload.transcription_local_files_only).lower(), + "TRANSCRIPTION_DEVICE": payload.transcription_device.strip(), + "TRANSCRIPTION_COMPUTE_TYPE": payload.transcription_compute_type.strip(), "VOLCENGINE_ASR_API_URL": payload.volcengine_asr_api_url.strip() or "https://openspeech.bytedance.com/api/v3/auc/bigmodel/recognize/flash", "VOLCENGINE_ASR_API_KEY": payload.volcengine_asr_api_key.strip() diff --git a/app/services/ai_retry_service.py b/app/services/ai_retry_service.py new file mode 100644 index 0000000..2f8e78c --- /dev/null +++ b/app/services/ai_retry_service.py @@ -0,0 +1,392 @@ +"""不完整 AI 分析的显式确认与安全恢复。""" + +from __future__ import annotations + +from copy import deepcopy +from datetime import datetime, timezone +import hashlib +import json +from pathlib import Path +from typing import Any + +from app.db.database import get_connection +from app.models.task import TaskStatus +from app.services import job_service +from app.services.storage_service import get_artifact_paths + + +_AI_UNIT_CHECKPOINT_KEY = "_ai_analysis_units_v1" +_SAFE_RETAINED_STEPS = ( + TaskStatus.PREPARING_SOURCE.value, + TaskStatus.TRANSCRIBING.value, +) + + +class AIAnalysisRetryConfirmationRequired(ValueError): + """AI 单元可能已计费,必须由用户显式确认。""" + + def __init__(self, detail: dict[str, Any]) -> None: + super().__init__(str(detail.get("message") or "AI 重试需要确认")) + self.detail = detail + + +class AutoPipelineRetryConflictError(ValueError): + """当前任务不满足安全重试前提。""" + + def __init__(self, message: str, *, code: str = "auto_retry_conflict") -> None: + super().__init__(message) + self.detail = {"code": code, "message": message} + + +def get_ai_retry_confirmation(task_id: str) -> dict[str, Any] | None: + """返回待确认信息;没有计费不确定单元时返回 ``None``。""" + with get_connection() as connection: + active = _active_auto_job(connection, task_id) + if active: + return None + row = _latest_failed_auto_job(connection, task_id) + if not row: + return None + checkpoint = _decode_checkpoint(row["checkpoint_json"]) + uncertain_units = _uncertain_units(checkpoint) + if not uncertain_units: + return None + _require_no_downstream_records(connection, task_id) + retry_mode = _resolve_retry_mode(task_id, checkpoint) + return _confirmation_detail( + task_id=task_id, + job_id=str(row["id"]), + uncertain_count=len(uncertain_units), + retry_mode=retry_mode, + ) + + +def prepare_confirmed_ai_retry(task_id: str) -> tuple[dict, bool, dict[str, Any]]: + """在同一事务中重新验证前提并准备已确认的 AI 重试。""" + now = _now_iso() + with get_connection() as connection: + connection.execute("BEGIN IMMEDIATE") + active = _active_auto_job(connection, task_id) + if active: + connection.commit() + job = job_service.get_job(str(active["id"])) + return job, False, { + "retry_mode": "active", + "uncertain_unit_count": 0, + "restart_step": TaskStatus.AI_ANALYZING.value, + } + + row = _latest_failed_auto_job(connection, task_id) + if not row: + connection.rollback() + raise AutoPipelineRetryConflictError("没有可恢复的失败全自动 Job") + checkpoint = _decode_checkpoint(row["checkpoint_json"]) + uncertain_units = _uncertain_units(checkpoint) + if not uncertain_units: + connection.rollback() + raise AutoPipelineRetryConflictError("失败 Job 不包含需要确认的 AI 单元") + _require_no_downstream_records(connection, task_id) + retry_mode = _resolve_retry_mode(task_id, checkpoint) + _prepare_task_for_ai_retry(connection, task_id, now=now) + + if retry_mode == "fresh_ai": + new_job_id, created = job_service.create_or_get_active_job_with_connection( + connection, + task_id=task_id, + job_type=job_service.JOB_TYPE_AUTO_PIPELINE, + payload={ + "retry": True, + "start_step": TaskStatus.AI_ANALYZING.value, + "confirmed_uncertain_ai": True, + "retry_mode": retry_mode, + "retry_of_job_id": str(row["id"]), + }, + ) + connection.commit() + job = job_service.get_job(new_job_id) + return job, created, { + "retry_mode": retry_mode, + "uncertain_unit_count": len(uncertain_units), + "restart_step": TaskStatus.AI_ANALYZING.value, + "previous_job_id": str(row["id"]), + } + + updated_checkpoint = _checkpoint_for_uncertain_resume( + checkpoint, + authorized_at=now, + ) + payload = _decode_json_object(row["payload_json"], field="payload_json") + payload.update( + { + "retry": True, + "start_step": TaskStatus.PREPARING_SOURCE.value, + "confirmed_uncertain_ai": True, + "retry_mode": retry_mode, + } + ) + cursor = connection.execute( + """ + UPDATE workflow_jobs + SET status = ?, progress = 0, message = '已确认 AI 重试,任务已重新加入队列', + payload_json = ?, result_json = '{}', checkpoint_json = ?, + checkpoint_updated_at = ?, error_message = NULL, finished_at = NULL, + next_attempt_at = ?, cancel_requested = 0, lease_owner = NULL, + lease_token = NULL, lease_expires_at = NULL, heartbeat_at = NULL, + attempt_count = 0, updated_at = ? + WHERE id = ? AND status IN (?, ?) + """, + ( + job_service.JOB_STATUS_QUEUED, + json.dumps(payload, ensure_ascii=False), + json.dumps(updated_checkpoint, ensure_ascii=False), + now, + now, + now, + str(row["id"]), + job_service.JOB_STATUS_FAILED, + job_service.JOB_STATUS_CANCELLED, + ), + ) + if cursor.rowcount != 1: + connection.rollback() + raise AutoPipelineRetryConflictError("失败 Job 状态已经变化,请刷新页面后重试") + connection.commit() + job = job_service.get_job(str(row["id"])) + return job, True, { + "retry_mode": retry_mode, + "uncertain_unit_count": len(uncertain_units), + "restart_step": TaskStatus.AI_ANALYZING.value, + "previous_job_id": str(row["id"]), + } + + +def _active_auto_job(connection, task_id: str): + return connection.execute( + """ + SELECT id FROM workflow_jobs + WHERE task_id = ? AND job_type = ? AND status IN (?, ?) + ORDER BY created_at DESC LIMIT 1 + """, + ( + task_id, + job_service.JOB_TYPE_AUTO_PIPELINE, + job_service.JOB_STATUS_QUEUED, + job_service.JOB_STATUS_RUNNING, + ), + ).fetchone() + + +def _latest_failed_auto_job(connection, task_id: str): + return connection.execute( + """ + SELECT id, payload_json, checkpoint_json + FROM workflow_jobs + WHERE task_id = ? AND job_type = ? AND status IN (?, ?) + ORDER BY created_at DESC LIMIT 1 + """, + ( + task_id, + job_service.JOB_TYPE_AUTO_PIPELINE, + job_service.JOB_STATUS_FAILED, + job_service.JOB_STATUS_CANCELLED, + ), + ).fetchone() + + +def _require_no_downstream_records(connection, task_id: str) -> None: + output_count = int( + connection.execute( + "SELECT COUNT(*) AS value FROM output_clip WHERE task_id = ?", + (task_id,), + ).fetchone()["value"] + or 0 + ) + publish_count = int( + connection.execute( + "SELECT COUNT(*) AS value FROM publish_jobs WHERE task_id = ?", + (task_id,), + ).fetchone()["value"] + or 0 + ) + if output_count or publish_count: + raise AutoPipelineRetryConflictError( + "任务已经存在下游切片或发布记录,为避免覆盖现有结果,已拒绝 AI 覆盖式重试。", + code="ai_retry_downstream_conflict", + ) + + +def _prepare_task_for_ai_retry(connection, task_id: str, *, now: str) -> None: + """把旧的下游失败态收敛到允许进入 AI_ANALYZING 的失败态。""" + cursor = connection.execute( + """ + UPDATE tasks + SET status = ?, progress = 45, + error_message = '等待执行已确认的 AI 分析重试', + last_error = '等待执行已确认的 AI 分析重试', updated_at = ? + WHERE id = ? AND COALESCE(is_deleted, 0) = 0 AND auto_mode = 1 + AND status LIKE 'FAILED_%' + """, + (TaskStatus.FAILED_AI_ANALYZING.value, now, task_id), + ) + if cursor.rowcount != 1: + raise AutoPipelineRetryConflictError( + "任务已不处于可恢复的全自动失败状态,请刷新页面后重试。", + code="ai_retry_task_state_conflict", + ) + + +def _resolve_retry_mode(task_id: str, checkpoint: dict[str, Any]) -> str: + steps = checkpoint.get("steps") or {} + transcription = steps.get(TaskStatus.TRANSCRIBING.value) or {} + outputs = transcription.get("outputs") or {} + transcript_evidence = outputs.get("transcript") or {} + stored_sha256 = str(transcript_evidence.get("sha256") or "") + transcript_path = get_artifact_paths(task_id)["transcript_path"] + if not stored_sha256 or not transcript_path.is_file(): + raise AutoPipelineRetryConflictError( + "旧 Job 缺少可信的转写输入校验值,无法判断是否可安全复用 AI checkpoint。", + code="ai_retry_input_evidence_missing", + ) + return "resume_uncertain" if _sha256_file(transcript_path) == stored_sha256 else "fresh_ai" + + +def _checkpoint_for_uncertain_resume( + checkpoint: dict[str, Any], + *, + authorized_at: str, +) -> dict[str, Any]: + updated = deepcopy(checkpoint) + steps = updated.get("steps") + completed = updated.get("completed_steps") + if not isinstance(steps, dict) or not isinstance(completed, list): + raise AutoPipelineRetryConflictError( + "自动流水线 checkpoint 结构损坏,已拒绝重复调用 AI。", + code="ai_retry_checkpoint_invalid", + ) + if updated.get("start_step") != TaskStatus.PREPARING_SOURCE.value: + raise AutoPipelineRetryConflictError( + "旧 Job 不是从完整全自动流程启动,无法原位复用 AI 单元。", + code="ai_retry_checkpoint_invalid", + ) + retained_steps: dict[str, Any] = {} + for step in _SAFE_RETAINED_STEPS: + record = steps.get(step) + if not isinstance(record, dict) or record.get("state") != "succeeded": + raise AutoPipelineRetryConflictError( + "准备素材或转写 checkpoint 缺少成功证据,已拒绝原位重试。", + code="ai_retry_checkpoint_invalid", + ) + retained_steps[step] = deepcopy(record) + + ai_record = steps.get(TaskStatus.AI_ANALYZING.value) + attempts = int(ai_record.get("attempts") or 1) if isinstance(ai_record, dict) else 1 + baseline = deepcopy(ai_record.get("baseline") or {}) if isinstance(ai_record, dict) else {} + retained_steps[TaskStatus.AI_ANALYZING.value] = { + "state": "failed", + "attempts": max(1, attempts), + "started_at": str((ai_record or {}).get("started_at") or authorized_at), + "failed_at": authorized_at, + "baseline": baseline, + "outputs": {}, + "error": "用户已确认重试计费结果不确定的 AI 单元", + } + updated["completed_steps"] = list(_SAFE_RETAINED_STEPS) + updated["current_step"] = TaskStatus.AI_ANALYZING.value + updated["steps"] = retained_steps + updated["last_error"] = "用户已确认重试计费结果不确定的 AI 单元" + updated["updated_at"] = authorized_at + + reset_count = 0 + for _, unit in _uncertain_units(updated): + previous_error = str(unit.get("error") or "AI 请求结果不确定") + unit.clear() + unit.update( + { + "status": "retryable_failed", + "error": previous_error, + "retry_authorized_at": authorized_at, + "retry_authorized_reason": "explicit_user_confirmation", + } + ) + reset_count += 1 + if not reset_count: + raise AutoPipelineRetryConflictError("没有可重置的计费不确定 AI 单元") + return updated + + +def _uncertain_units(checkpoint: dict[str, Any]) -> list[tuple[str, dict[str, Any]]]: + root = checkpoint.get(_AI_UNIT_CHECKPOINT_KEY) + if not isinstance(root, dict): + return [] + namespaces = root.get("namespaces") + if not isinstance(namespaces, dict): + return [] + result: list[tuple[str, dict[str, Any]]] = [] + for namespace, state in namespaces.items(): + units = state.get("units") if isinstance(state, dict) else None + if not isinstance(units, dict): + continue + for unit_id, unit in units.items(): + if isinstance(unit, dict) and unit.get("status") in {"running", "uncertain"}: + result.append((f"{namespace}/{unit_id}", unit)) + return result + + +def _confirmation_detail( + *, + task_id: str, + job_id: str, + uncertain_count: int, + retry_mode: str, +) -> dict[str, Any]: + action = ( + "只重新请求结果不确定的 AI 单元,并复用其余成功 checkpoint" + if retry_mode == "resume_uncertain" + else "转写输入已变化,将从 AI 分析阶段创建全新 Job,不复用旧 AI 单元" + ) + return { + "code": "ai_retry_confirmation_required", + "message": ( + f"检测到 {uncertain_count} 个 AI 单元的请求可能已经产生费用,但结果未确认。" + f"确认后将{action}。" + ), + "task_id": task_id, + "previous_job_id": job_id, + "uncertain_unit_count": uncertain_count, + "retry_mode": retry_mode, + "restart_step": TaskStatus.AI_ANALYZING.value, + } + + +def _decode_checkpoint(value: Any) -> dict[str, Any]: + return _decode_json_object(value, field="checkpoint_json") + + +def _decode_json_object(value: Any, *, field: str) -> dict[str, Any]: + if isinstance(value, dict): + return deepcopy(value) + try: + decoded = json.loads(str(value or "{}")) + except json.JSONDecodeError as exc: + raise AutoPipelineRetryConflictError( + f"Workflow Job {field} 已损坏,已拒绝重试。", + code="ai_retry_checkpoint_invalid", + ) from exc + if not isinstance(decoded, dict): + raise AutoPipelineRetryConflictError( + f"Workflow Job {field} 不是对象,已拒绝重试。", + code="ai_retry_checkpoint_invalid", + ) + return decoded + + +def _sha256_file(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as source: + for block in iter(lambda: source.read(1024 * 1024), b""): + digest.update(block) + return digest.hexdigest() + + +def _now_iso() -> str: + return datetime.now(timezone.utc).isoformat(timespec="seconds") diff --git a/app/services/chinese_text_service.py b/app/services/chinese_text_service.py new file mode 100644 index 0000000..2b91fee --- /dev/null +++ b/app/services/chinese_text_service.py @@ -0,0 +1,24 @@ +"""本地转写中文本的字形标准化。""" + +from functools import lru_cache + + +SIMPLIFIED_CHINESE_NORMALIZATION_ID = "opencc-t2s-v1" + + +@lru_cache(maxsize=1) +def _t2s_converter(): + try: + from opencc import OpenCC + except ImportError as exc: + raise RuntimeError( + "未安装简体转换依赖 opencc。请先安装 requirements.txt,再重新生成本地转写。" + ) from exc + return OpenCC("t2s.json") + + +def simplify_chinese_text(text: str) -> str: + """只做繁体到简体的字形转换,不做大陆词汇本地化。""" + if not text: + return text + return _t2s_converter().convert(text) diff --git a/app/services/local_transcription_runtime.py b/app/services/local_transcription_runtime.py new file mode 100644 index 0000000..3cf0efb --- /dev/null +++ b/app/services/local_transcription_runtime.py @@ -0,0 +1,251 @@ +"""本地 faster-whisper 模型缓存和 Windows CUDA 运行时检查。""" + +from __future__ import annotations + +import ctypes +import importlib.util +import os +from pathlib import Path +import sys +from typing import Any + +from app.core.config import settings +from app.core.transcription_defaults import ( + CPU_FALLBACK_TRANSCRIPTION_MODEL, + CPU_FALLBACK_TRANSCRIPTION_MODEL_REVISION, + PRIMARY_TRANSCRIPTION_MODEL, + PRIMARY_TRANSCRIPTION_MODEL_REVISION, +) + + +REQUIRED_MODEL_FILES = ( + "config.json", + "model.bin", + "tokenizer.json", +) + +OPTIONAL_MODEL_FILES = ( + "preprocessor_config.json", +) + +MODEL_VOCABULARY_FILES = ( + "vocabulary.json", + "vocabulary.txt", +) + +_DLL_DIRECTORY_HANDLES: list[Any] = [] +_DLL_DIRECTORY_PATHS: set[str] = set() +_DLL_LIBRARY_HANDLES: list[Any] = [] +_DLL_LIBRARY_PATHS: set[str] = set() + + +class TranscriptionOfflinePolicyError(RuntimeError): + """完全离线模式禁止调用远程转写。""" + + +def ensure_transcription_provider_allowed(provider: str) -> str: + normalized = (provider or "").strip().lower() + if settings.transcription_offline_only and normalized != "local": + raise TranscriptionOfflinePolicyError( + "已启用完全离线转写,禁止调用火山引擎或其他远程转写服务。" + ) + return normalized + + +def model_revision_for(model_name: str) -> str: + normalized = (model_name or "").strip() + if normalized == CPU_FALLBACK_TRANSCRIPTION_MODEL: + return CPU_FALLBACK_TRANSCRIPTION_MODEL_REVISION + if normalized == PRIMARY_TRANSCRIPTION_MODEL: + return PRIMARY_TRANSCRIPTION_MODEL_REVISION + if normalized == settings.transcription_model: + return settings.transcription_model_revision + return settings.transcription_model_revision + + +def model_identity(model_name: str, revision: str | None = None) -> str: + resolved_revision = (revision or model_revision_for(model_name)).strip() + suffix = resolved_revision[:8] if resolved_revision else "unversioned" + return f"{model_name}@{suffix}" + + +def model_cache_directory(model_name: str, revision: str | None = None) -> Path: + resolved_revision = (revision or model_revision_for(model_name)).strip() + safe_name = (model_name or "model").replace("/", "--").replace("\\", "--") + suffix = resolved_revision[:8] if resolved_revision else "unversioned" + return Path(settings.transcription_model_cache_dir) / f"{safe_name}-{suffix}" + + +def model_files_ready(model_path: Path) -> bool: + return bool( + model_path.is_dir() + and all((model_path / name).is_file() for name in REQUIRED_MODEL_FILES) + and any((model_path / name).is_file() for name in MODEL_VOCABULARY_FILES) + ) + + +def _package_directory(package_name: str) -> Path | None: + spec = importlib.util.find_spec(package_name) + if spec is None: + return None + if spec.submodule_search_locations: + return Path(next(iter(spec.submodule_search_locations))) + if spec.origin: + return Path(spec.origin).parent + return None + + +def configure_windows_cuda_dll_directories() -> dict[str, Any]: + """仅修改当前 Python 进程的 DLL 搜索路径,不改 Windows 全局 PATH。""" + if os.name != "nt": + return { + "configured": True, + "cublas_ready": True, + "cudnn_ready": True, + "directories": [], + "errors": [], + } + + site_packages = Path(sys.prefix) / "Lib" / "site-packages" + cublas_dir = site_packages / "nvidia" / "cublas" / "bin" + ctranslate2_dir = _package_directory("ctranslate2") + candidates = [cublas_dir] + if ctranslate2_dir is not None: + candidates.append(ctranslate2_dir) + + errors: list[str] = [] + for directory in candidates: + directory_text = str(directory) + if not directory.is_dir() or directory_text in _DLL_DIRECTORY_PATHS: + continue + try: + _DLL_DIRECTORY_HANDLES.append(os.add_dll_directory(directory_text)) + _DLL_DIRECTORY_PATHS.add(directory_text) + except (FileNotFoundError, OSError) as exc: + errors.append(f"{directory}: {exc}") + + library_paths = [ + cublas_dir / "cublasLt64_12.dll", + cublas_dir / "cublas64_12.dll", + ] + if ctranslate2_dir is not None: + library_paths.append(ctranslate2_dir / "cudnn64_9.dll") + library_errors: dict[str, str] = {} + for library_path in library_paths: + library_text = str(library_path) + if not library_path.is_file(): + library_errors[library_path.name] = "文件不存在" + continue + if library_text in _DLL_LIBRARY_PATHS: + continue + try: + _DLL_LIBRARY_HANDLES.append(ctypes.WinDLL(library_text)) + _DLL_LIBRARY_PATHS.add(library_text) + except OSError as exc: + library_errors[library_path.name] = str(exc) + errors.append(f"{library_path}: {exc}") + + return { + "configured": not errors, + "cublas_ready": all( + path.is_file() and path.name not in library_errors + for path in library_paths[:2] + ), + "cudnn_ready": bool( + ctranslate2_dir + and library_paths[-1].is_file() + and library_paths[-1].name not in library_errors + ), + "directories": [str(path) for path in candidates if path.is_dir()], + "libraries": [str(path) for path in library_paths if path.is_file()], + "library_errors": library_errors, + "errors": errors, + } + + +def resolve_local_model_source(model_name: str, revision: str | None = None) -> str: + model_path = model_cache_directory(model_name, revision) + if model_files_ready(model_path): + return str(model_path) + raise RuntimeError( + "本地转写模型尚未初始化或缓存不完整:" + f"{model_path}。任务运行时不会联网下载,请先运行 scripts/setup_local_transcription.ps1。" + ) + + +def get_local_transcription_runtime_status() -> dict[str, Any]: + dll_status = configure_windows_cuda_dll_directories() + model_path = model_cache_directory( + settings.transcription_model, + settings.transcription_model_revision, + ) + model_ready = model_files_ready(model_path) + cuda_device_count = 0 + cuda_error = "" + try: + import ctranslate2 + + cuda_device_count = int(ctranslate2.get_cuda_device_count()) + except Exception as exc: + cuda_error = str(exc) + + gpu_ready = bool( + cuda_device_count > 0 + and dll_status["cublas_ready"] + and dll_status["cudnn_ready"] + ) + errors = list(dll_status["errors"]) + if cuda_error: + errors.append(cuda_error) + if not model_ready: + errors.append("固定版本模型尚未完整缓存") + + return { + "offline_only": bool(settings.transcription_offline_only), + "local_files_only": bool(settings.transcription_local_files_only), + "provider": settings.transcription_provider, + "model": settings.transcription_model, + "model_revision": settings.transcription_model_revision, + "model_identity": model_identity( + settings.transcription_model, + settings.transcription_model_revision, + ), + "model_cache_dir": str(settings.transcription_model_cache_dir), + "model_path": str(model_path), + "model_ready": model_ready, + "gpu_ready": gpu_ready, + "cuda_device_count": cuda_device_count, + "cublas_ready": bool(dll_status["cublas_ready"]), + "cudnn_ready": bool(dll_status["cudnn_ready"]), + "device": settings.transcription_device, + "compute_type": settings.transcription_compute_type, + "external_cost": "0 元", + "external_cost_yuan": 0, + "ready": bool( + settings.transcription_provider == "local" + and model_ready + and ( + settings.transcription_device != "cuda" + or gpu_ready + ) + ), + "errors": errors, + } + + +__all__ = [ + "PRIMARY_TRANSCRIPTION_MODEL", + "PRIMARY_TRANSCRIPTION_MODEL_REVISION", + "MODEL_VOCABULARY_FILES", + "OPTIONAL_MODEL_FILES", + "REQUIRED_MODEL_FILES", + "TranscriptionOfflinePolicyError", + "configure_windows_cuda_dll_directories", + "ensure_transcription_provider_allowed", + "get_local_transcription_runtime_status", + "model_cache_directory", + "model_files_ready", + "model_identity", + "model_revision_for", + "resolve_local_model_source", +] diff --git a/app/services/pipeline_engine.py b/app/services/pipeline_engine.py index 2d6c6c4..4365500 100644 --- a/app/services/pipeline_engine.py +++ b/app/services/pipeline_engine.py @@ -12,6 +12,11 @@ from app.db.database import get_connection from app.models.task import TaskStatus from app.services import job_service, task_service +from app.services.ai_retry_service import ( + AIAnalysisRetryConfirmationRequired, + get_ai_retry_confirmation, + prepare_confirmed_ai_retry, +) from app.services.auto_publish_service import create_auto_publish_jobs, platforms_for_task from app.services.metadata_generator import MetadataGenerator from app.services.pipeline_checkpoint_service import ( @@ -1375,43 +1380,55 @@ def _transcribe_or_read_text(self, task_id: str, context: dict) -> dict: def _run_ai_analysis(self, task_id: str, context: dict) -> dict: result = task_service.process_task_ai_analysis(task_id) clip_count = len(result.get("clips") or []) + analysis_meta = (result.get("analysis_run") or {}).get("analysis_meta") or {} + self._require_complete_ai_analysis(task_id, analysis_meta) append_task_log(task_id, f"全自动 AI 分析完成,候选片段:{clip_count} 条") return { "clip_count": clip_count, "analysis_run_id": result.get("analysis_run_id") or "", "analysis_path": result.get("analysis_path") or "", - "analysis_meta": (result.get("analysis_run") or {}).get("analysis_meta") or {}, + "analysis_meta": analysis_meta, } - def _select_clips(self, task_id: str, context: dict) -> dict: + def _require_complete_ai_analysis(self, task_id: str, meta: dict | None) -> dict: task = self._get_task(task_id) - config = context["config"] - meta = task_service.get_task_ai_analysis_meta(task_id) profile = str(task.get("selection_profile") or "general") if not meta: raise ValueError( "AI 分析缺少可信的完整性元数据;请重新分析或恢复可信历史," "当前不会进入自动切片或发送中心。" ) - meta = task_service.validate_ai_analysis_meta_for_cut(meta, profile) - coverage = float(meta.get("coverage_percent") or 0) + validated = task_service.validate_ai_analysis_meta_for_cut(meta, profile) + coverage = float(validated.get("coverage_percent") or 0) if profile == "long_live_talk" and ( - meta.get("analysis_incomplete") or float(meta.get("coverage_ratio") or 0) < 0.90 + validated.get("analysis_incomplete") + or float(validated.get("coverage_ratio") or 0) < 0.90 ): raise ValueError( f"长直播分析覆盖率仅 {coverage:.2f}%,低于 90%;" "请重试 AI 分析补齐缺失窗口,当前不会进入自动切片或发送中心。" ) - if meta.get("analysis_incomplete"): + if validated.get("analysis_incomplete"): raise ValueError( f"{profile} AI 分析存在未完成单元,当前覆盖率 {coverage:.2f}%;" "请重试 AI 分析补齐失败单元,当前不会进入自动切片或发送中心。" ) + return validated + + def _select_clips(self, task_id: str, context: dict) -> dict: + task = self._get_task(task_id) + config = context["config"] + meta = self._require_complete_ai_analysis( + task_id, + task_service.get_task_ai_analysis_meta(task_id), + ) + profile = str(task.get("selection_profile") or "general") if meta.get("quality_degraded"): raise ValueError( f"{profile} AI 分析质量评审未完整通过;" "候选可供人工检查,但当前不会进入自动切片或发送中心。" ) + candidates = self._list_raw_candidates(task_id) if not candidates: latest_run = task_service.get_latest_ai_analysis_run(task_id) @@ -1895,11 +1912,29 @@ def start_auto_pipeline( background_tasks: Any | None = None, retry: bool = False, start_step: TaskStatus | str | None = None, + confirm_uncertain_ai: bool = False, ) -> dict: if background_tasks is not None: from app.services import job_service if retry: + confirmation = get_ai_retry_confirmation(task_id) + if confirmation and not confirm_uncertain_ai: + raise AIAnalysisRetryConfirmationRequired(confirmation) + if confirmation: + job, requeued, retry_info = prepare_confirmed_ai_retry(task_id) + return { + "status": job["status"], + "message": ( + "转写输入已变化,已创建新的 AI 分析 Job。" + if retry_info["retry_mode"] == "fresh_ai" + else "已确认计费不确定单元,成功 checkpoint 将继续复用。" + ), + "task_id": task_id, + "job_id": job["id"], + "created": requeued, + **retry_info, + } job, requeued = job_service.retry_latest_or_get_active_job( task_id, job_service.JOB_TYPE_AUTO_PIPELINE, diff --git a/app/services/task_service.py b/app/services/task_service.py index 8805e82..6dc8f7b 100644 --- a/app/services/task_service.py +++ b/app/services/task_service.py @@ -83,6 +83,7 @@ update_default_subtitle_style, ) from app.services.transcript_service import read_transcript_progress, read_transcript_range +from app.services.local_transcription_runtime import TranscriptionOfflinePolicyError from app.services.transcript_workflow_service import ( TranscriptCancelledError, _can_retry_transcript_with_local, @@ -99,6 +100,7 @@ process_task_audio, process_task_transcript, process_task_transcript_workflow, + validate_transcription_provider_choice, ) from app.services.video_cut_workflow_service import process_task_video_cuts diff --git a/app/services/transcript_service.py b/app/services/transcript_service.py index 5743f1f..2d932b2 100644 --- a/app/services/transcript_service.py +++ b/app/services/transcript_service.py @@ -18,7 +18,18 @@ from app.core.config import settings from app.services import job_service +from app.services.chinese_text_service import ( + SIMPLIFIED_CHINESE_NORMALIZATION_ID, + simplify_chinese_text, +) from app.services.managed_process_service import popen_process_group, terminate_process_tree +from app.services.local_transcription_runtime import ( + configure_windows_cuda_dll_directories, + ensure_transcription_provider_allowed, + model_identity, + model_revision_for, + resolve_local_model_source, +) from app.services.transcription_checkpoint_service import ( RemoteTranscriptionResultUncertainError, TranscriptionCheckpoint, @@ -66,6 +77,7 @@ def duration_seconds(self) -> float: _WHISPER_MODEL = None _WHISPER_MODEL_KEY: tuple[str, str, str] | None = None +_EFFECTIVE_TRANSCRIPTION_MODEL_KEY: tuple[str, str, str] | None = None _CPU_FALLBACK_DEVICE = "cpu" _CPU_FALLBACK_COMPUTE_TYPE = "int8" _TRANSCRIPT_PROGRESS_FILE_NAME = "transcript_progress.json" @@ -218,6 +230,17 @@ def write_transcript_markdown( transcript_path.parent.mkdir(parents=True, exist_ok=True) progress_path = get_transcript_progress_path(transcript_path) _set_configured_transcription_runtime(provider) + _emit_transcript_progress( + progress_path, + progress_callback, + status="running", + current_chunk=0, + total_chunks=0, + percent=0, + message="正在验证固定版本本地模型和计算设备", + ) + if _ACTIVE_TRANSCRIPTION_PROVIDER == "local": + _prepare_local_transcription_model_for_run(audio_path) checkpoint = None task_id = str(task.get("id") or "").strip() if task_id: @@ -225,7 +248,7 @@ def write_transcript_markdown( task_id=task_id, source_path=audio_path, provider=_ACTIVE_TRANSCRIPTION_PROVIDER, - model=_ACTIVE_TRANSCRIPTION_MODEL, + model=_transcription_checkpoint_model(), device=_ACTIVE_TRANSCRIPTION_DEVICE, compute_type=_ACTIVE_TRANSCRIPTION_COMPUTE_TYPE, chunk_seconds=settings.transcription_chunk_seconds, @@ -335,17 +358,19 @@ def transcribe_audio_with_provider( ) -> list[TranscriptSegment]: provider = _normalize_provider_name(provider) if provider == "local": + model_key = _WHISPER_MODEL_KEY or _primary_model_key() _set_active_transcription_runtime( provider="local", provider_label="本地 faster-whisper", - model=(_WHISPER_MODEL_KEY or _primary_model_key())[0], - device=(_WHISPER_MODEL_KEY or _primary_model_key())[1], - compute_type=(_WHISPER_MODEL_KEY or _primary_model_key())[2], + model=model_identity(model_key[0]), + device=model_key[1], + compute_type=model_key[2], ) return transcribe_audio_in_chunks( audio_path, working_dir, progress_path, progress_callback, checkpoint=checkpoint ) if provider == "volcengine": + ensure_transcription_provider_allowed(provider) return transcribe_audio_with_volcengine( audio_path, working_dir, progress_path, progress_callback, checkpoint=checkpoint ) @@ -1104,10 +1129,16 @@ def _write_transcript_progress( "provider": _ACTIVE_TRANSCRIPTION_PROVIDER, "provider_label": _ACTIVE_TRANSCRIPTION_PROVIDER_LABEL, "model": _transcription_model_label(), + "model_revision": _transcription_model_revision_label(), "device": _transcription_device_label(), "compute_type": _transcription_compute_type_label(), "chunk_seconds": settings.transcription_chunk_seconds, "chunk_overlap_seconds": settings.transcription_chunk_overlap_seconds, + "text_normalization": ( + SIMPLIFIED_CHINESE_NORMALIZATION_ID + if _ACTIVE_TRANSCRIPTION_PROVIDER == "local" + else "" + ), "updated_at": datetime.now().isoformat(timespec="seconds"), } temp_path = progress_path.with_name(f"{progress_path.name}.tmp") @@ -1130,7 +1161,10 @@ def _chunk_start_percent(current_chunk: int, total_chunks: int) -> int: def transcribe_audio(audio_path: Path, allow_empty: bool = False) -> list[TranscriptSegment]: try: - segments = _transcribe_audio_with_model(_get_whisper_model(), audio_path) + segments = _transcribe_audio_with_model( + _get_whisper_model(_EFFECTIVE_TRANSCRIPTION_MODEL_KEY or _primary_model_key()), + audio_path, + ) except Exception as exc: if _should_retry_with_cpu(exc): try: @@ -1151,6 +1185,45 @@ def transcribe_audio(audio_path: Path, allow_empty: bool = False) -> list[Transc return segments +def _prepare_local_transcription_model_for_run(audio_path: Path | None = None) -> None: + """在创建 checkpoint 前确定本次实际使用 GPU 主模型还是 CPU 兜底。""" + global _EFFECTIVE_TRANSCRIPTION_MODEL_KEY + + _EFFECTIVE_TRANSCRIPTION_MODEL_KEY = None + model = _get_whisper_model(_primary_model_key()) + if audio_path is not None: + try: + _probe_whisper_model_inference(model, audio_path) + except Exception as exc: + if _WHISPER_MODEL_KEY != _cpu_fallback_model_key() and _should_retry_with_cpu(exc): + model = _get_whisper_model(_cpu_fallback_model_key()) + _probe_whisper_model_inference(model, audio_path) + else: + raise + _EFFECTIVE_TRANSCRIPTION_MODEL_KEY = _WHISPER_MODEL_KEY or _primary_model_key() + model, device, compute_type = _EFFECTIVE_TRANSCRIPTION_MODEL_KEY + _set_active_transcription_runtime( + provider="local", + provider_label="本地 faster-whisper", + model=model_identity(model), + device=device, + compute_type=compute_type, + ) + + +def _probe_whisper_model_inference(model, audio_path: Path) -> None: + """只解码开头 1 秒,确保 checkpoint 记录的是可实际推理的设备。""" + segments, _info = model.transcribe( + str(audio_path), + language=settings.transcription_language, + vad_filter=False, + beam_size=1, + word_timestamps=False, + clip_timestamps="0,1", + ) + list(segments) + + def _transcribe_audio_with_model(model, audio_path: Path) -> list[TranscriptSegment]: raw_segments, _info = model.transcribe( str(audio_path), @@ -1161,19 +1234,27 @@ def _transcribe_audio_with_model(model, audio_path: Path) -> list[TranscriptSegm ) results: list[TranscriptSegment] = [] for segment in raw_segments: - text = normalize_transcript_text(segment.text) + text = _normalize_local_transcript_text(segment.text) if not text: continue - words = tuple( - TranscriptWord( - start_ms=round(float(word.start) * 1000), - end_ms=round(float(word.end) * 1000), - text=normalize_transcript_text(word.word), - confidence=float(word.probability) if getattr(word, "probability", None) is not None else None, + normalized_words: list[TranscriptWord] = [] + for word in (getattr(segment, "words", None) or []): + word_text = _normalize_local_transcript_text(getattr(word, "word", "")) + if not word_text: + continue + normalized_words.append( + TranscriptWord( + start_ms=round(float(word.start) * 1000), + end_ms=round(float(word.end) * 1000), + text=word_text, + confidence=( + float(word.probability) + if getattr(word, "probability", None) is not None + else None + ), + ) ) - for word in (getattr(segment, "words", None) or []) - if normalize_transcript_text(getattr(word, "word", "")) - ) + words = tuple(normalized_words) avg_logprob = getattr(segment, "avg_logprob", None) confidence = max(0.0, min(1.0, 2.718281828 ** float(avg_logprob))) if avg_logprob is not None else None results.append( @@ -1196,6 +1277,7 @@ def _get_whisper_model(model_key: tuple[str, str, str] | None = None): return _WHISPER_MODEL try: + configure_windows_cuda_dll_directories() from faster_whisper import WhisperModel except ImportError as exc: raise RuntimeError( @@ -1203,17 +1285,21 @@ def _get_whisper_model(model_key: tuple[str, str, str] | None = None): ) from exc try: + revision = model_revision_for(model_key[0]) + model_source = resolve_local_model_source(model_key[0], revision) _WHISPER_MODEL = WhisperModel( - model_key[0], + model_source, device=model_key[1], compute_type=model_key[2], + download_root=str(settings.transcription_model_cache_dir), + local_files_only=settings.transcription_local_files_only, ) _WHISPER_MODEL_KEY = model_key if _ACTIVE_TRANSCRIPTION_PROVIDER == "local": _set_active_transcription_runtime( provider="local", provider_label="本地 faster-whisper", - model=model_key[0], + model=model_identity(model_key[0], revision), device=model_key[1], compute_type=model_key[2], ) @@ -1325,7 +1411,7 @@ def _set_configured_transcription_runtime(provider: str | None = None) -> None: _set_active_transcription_runtime( provider="local", provider_label="本地 faster-whisper", - model=model, + model=model_identity(model), device=device, compute_type=compute_type, ) @@ -1350,7 +1436,14 @@ def _transcription_model_label() -> str: if _ACTIVE_TRANSCRIPTION_MODEL: return _ACTIVE_TRANSCRIPTION_MODEL model, _device, _compute_type = _WHISPER_MODEL_KEY or _primary_model_key() - return model + return model_identity(model) + + +def _transcription_model_revision_label() -> str: + if _ACTIVE_TRANSCRIPTION_PROVIDER != "local": + return "" + model, _device, _compute_type = _WHISPER_MODEL_KEY or _primary_model_key() + return model_revision_for(model) def _transcription_device_label() -> str: @@ -1420,6 +1513,16 @@ def normalize_transcript_text(text: str) -> str: return re.sub(r"\s+", " ", text or "").strip() +def _normalize_local_transcript_text(text: str) -> str: + return simplify_chinese_text(normalize_transcript_text(text)) + + +def _transcription_checkpoint_model() -> str: + if _ACTIVE_TRANSCRIPTION_PROVIDER != "local": + return _ACTIVE_TRANSCRIPTION_MODEL + return f"{_ACTIVE_TRANSCRIPTION_MODEL}|text={SIMPLIFIED_CHINESE_NORMALIZATION_ID}" + + def escape_markdown_table_text(text: str) -> str: return normalize_transcript_text(text).replace("|", "\\|") diff --git a/app/services/transcript_workflow_service.py b/app/services/transcript_workflow_service.py index 25358aa..b2e83d8 100644 --- a/app/services/transcript_workflow_service.py +++ b/app/services/transcript_workflow_service.py @@ -10,6 +10,10 @@ from app.core.config import settings from app.models.task import TaskStatus from app.services import job_service +from app.services.local_transcription_runtime import ( + ensure_transcription_provider_allowed, + get_local_transcription_runtime_status, +) from app.services.storage_service import get_artifact_paths, get_source_video_path, validate_source_video_path from app.services.task_log_service import append_task_log from app.services.transcript_service import ( @@ -29,9 +33,6 @@ _TRANSCRIPT_STALE_AFTER = timedelta(minutes=10) -_DEFAULT_REMOTE_TRANSCRIPTION_PROVIDER = "volcengine" - - # ---------- 转写进度辅助 ---------- def _parse_progress_updated_at(progress: dict) -> datetime | None: @@ -64,13 +65,18 @@ def _is_transcript_progress_stale(progress: dict) -> bool: def _resolve_transcription_provider_choice(provider: str | None = None) -> str: - choice = (provider or "remote").strip().lower() - if choice == "local": - return "local" configured_provider = (settings.transcription_provider or "").strip().lower() - if configured_provider and configured_provider != "local": - return configured_provider - return _DEFAULT_REMOTE_TRANSCRIPTION_PROVIDER + choice = (provider or configured_provider or "local").strip().lower() + if choice == "remote": + choice = configured_provider if configured_provider and configured_provider != "local" else "volcengine" + if choice not in {"local", "volcengine", "aliyun", "tencent", "xunfei"}: + raise ValueError(f"未知转写服务商:{choice or '空'}") + return ensure_transcription_provider_allowed(choice) + + +def validate_transcription_provider_choice(provider: str | None = None) -> str: + """在创建持久化 Job 前验证 Provider 和完全离线边界。""" + return _resolve_transcription_provider_choice(provider) def _transcription_choice_label(provider: str) -> str: @@ -82,6 +88,8 @@ def _transcription_choice_label(provider: str) -> str: def _can_retry_transcript_with_local(progress: dict, transcript_exists: bool) -> bool: + if settings.transcription_offline_only: + return False provider = str(progress.get("provider") or "").strip().lower() return not transcript_exists and progress.get("status") == "failed" and provider != "local" @@ -133,6 +141,7 @@ def get_task_transcript_status(task_id: str) -> dict: } transcript_exists = paths["transcript_path"].exists() + runtime = get_local_transcription_runtime_status() return { "task_id": task_id, "task_status": task.get("status"), @@ -145,6 +154,10 @@ def get_task_transcript_status(task_id: str) -> dict: "preview": read_transcript_preview(paths["transcript_path"]), "error_message": task.get("error_message") or "", "local_retry_available": _can_retry_transcript_with_local(progress, transcript_exists), + "offline_only": runtime["offline_only"], + "model_ready": runtime["model_ready"], + "gpu_ready": runtime["gpu_ready"], + "model_revision": runtime["model_revision"], } diff --git a/app/services/transcription_checkpoint_service.py b/app/services/transcription_checkpoint_service.py index 98dbc93..ae86fcc 100644 --- a/app/services/transcription_checkpoint_service.py +++ b/app/services/transcription_checkpoint_service.py @@ -50,6 +50,16 @@ def fingerprint_file(path_value: str | Path) -> str: return digest.hexdigest() +def fingerprint_file_full(path_value: str | Path) -> str: + """完整读取转写音频,避免相同大小且头尾相同的文件误复用 checkpoint。""" + path = Path(path_value).resolve() + digest = hashlib.sha256() + with path.open("rb") as source: + for block in iter(lambda: source.read(4 * 1024 * 1024), b""): + digest.update(block) + return digest.hexdigest() + + class TranscriptionCheckpoint: def __init__( self, @@ -65,7 +75,7 @@ def __init__( allow_uncertain_retry: bool = False, ) -> None: self.task_id = task_id - self.source_fingerprint = fingerprint_file(source_path) + self.source_fingerprint = fingerprint_file_full(source_path) self.provider = provider self.model = model self.device = device diff --git a/app/static/js/app.js b/app/static/js/app.js index 8c3cfa4..2eee3fe 100644 --- a/app/static/js/app.js +++ b/app/static/js/app.js @@ -117,10 +117,25 @@ async function handleProcessAction(button) { if (result) result.textContent = "正在执行,请稍等..."; try { - const response = await fetch(button.dataset.endpoint, { method: "POST" }); - const data = await response.json(); - if (!response.ok) { - throw new Error(data.detail || "处理失败"); + let endpoint = button.dataset.endpoint; + let data; + try { + data = await apiFetch(endpoint, { method: "POST" }); + } catch (error) { + const needsAIConfirmation = error.status === 409 + && error.details?.code === "ai_retry_confirmation_required"; + if (!needsAIConfirmation) throw error; + const confirmed = window.confirm( + `${error.details.message}\n\n这一步可能产生新的 AI 调用费用,是否确认继续?`, + ); + if (!confirmed) { + if (result) result.textContent = "已取消 AI 重试,原失败 Job 和证据保持不变。"; + return; + } + const separator = endpoint.includes("?") ? "&" : "?"; + endpoint = `${endpoint}${separator}confirm_uncertain_ai=true`; + if (result) result.textContent = "已确认,正在安全创建或恢复 AI 分析 Job..."; + data = await apiFetch(endpoint, { method: "POST" }); } if (result) result.textContent = data.message || "处理完成,正在刷新页面..."; if (button.dataset.endpoint.includes("/process/transcript")) { @@ -260,6 +275,7 @@ function updateWorkflowButtons(data) { if (!startButton) return; const progressStatus = data.progress?.status || ""; const canRetryWithLocal = Boolean(data.local_retry_available); + const retryLabel = data.offline_only ? "重新本地转写" : "重新远程转写"; if (localTranscriptButton) { localTranscriptButton.hidden = !canRetryWithLocal; } @@ -277,7 +293,7 @@ function updateWorkflowButtons(data) { return; } if (progressStatus === "failed") { - startButton.textContent = "重新远程转写"; + startButton.textContent = retryLabel; startButton.disabled = false; startButton.classList.add("js-process-action"); startButton.dataset.endpoint = `/api/tasks/${transcriptPanel?.dataset.taskId}/process/transcript-workflow?force=true`; @@ -286,7 +302,7 @@ function updateWorkflowButtons(data) { return; } if (progressStatus === "cancelled" || progressStatus === "stale") { - startButton.textContent = progressStatus === "stale" ? "重新远程转写" : "重新生成转写"; + startButton.textContent = progressStatus === "stale" ? retryLabel : "重新生成转写"; startButton.disabled = false; startButton.classList.add("js-process-action"); startButton.dataset.endpoint = `/api/tasks/${transcriptPanel?.dataset.taskId}/process/transcript-workflow?force=true`; @@ -2438,7 +2454,7 @@ if (aiConfigForm) { const formData = new FormData(aiConfigForm); const payload = Object.fromEntries(formData.entries()); payload.ai_request_timeout_seconds = Number(payload.ai_request_timeout_seconds || 120); - payload.ai_codex_timeout_seconds = Number(payload.ai_codex_timeout_seconds || 300); + payload.ai_codex_timeout_seconds = Number(payload.ai_codex_timeout_seconds || 600); payload.volcengine_asr_timeout_seconds = Number(payload.volcengine_asr_timeout_seconds || 300); payload.ai_analysis_request_timeout_seconds = Number(payload.ai_analysis_request_timeout_seconds || 120); payload.ai_publish_request_timeout_seconds = Number(payload.ai_publish_request_timeout_seconds || 120); diff --git a/app/templates/system_status.html b/app/templates/system_status.html index 8cefa69..f982a45 100644 --- a/app/templates/system_status.html +++ b/app/templates/system_status.html @@ -50,6 +50,12 @@

系统状态

{{ ai_config.configured_count }} 转写 {{ "正常" if ai_config.transcription_ready else "待配置" }} · 分析 {{ "正常" if ai_config.analysis_ready else "待配置" }} · 文案 {{ "正常" if ai_config.publish_ready else "待配置" }} +
+
+ 完全离线转写 + ¥{{ ai_config.transcription_runtime.external_cost_yuan }} + {{ ai_config.transcription_runtime.model_identity }} · GPU {{ "可用" if ai_config.transcription_runtime.gpu_ready else "待初始化" }} · 模型 {{ "已缓存" if ai_config.transcription_runtime.model_ready else "待下载" }} +
@@ -145,14 +151,37 @@

通用开关

1. 音频转写

+

当前:{{ "完全离线" if ai_config.transcription_runtime.offline_only else "允许远程" }} · {{ ai_config.transcription_runtime.model_identity }} · {{ ai_config.transcription_runtime.device }}/{{ ai_config.transcription_runtime.compute_type }} · 外部费用 ¥{{ ai_config.transcription_runtime.external_cost_yuan }}

+
+ + +
+ + +
+ + +
+ +

模型 {{ "已缓存" if ai_config.transcription_runtime.model_ready else "未缓存" }} · cuBLAS {{ "正常" if ai_config.transcription_runtime.cublas_ready else "缺失" }} · cuDNN {{ "正常" if ai_config.transcription_runtime.cudnn_ready else "缺失" }}。缺失时请运行 scripts/setup_local_transcription.ps1。

+

火山引擎回滚配置

@@ -268,7 +297,8 @@

关键路径

  • 任务资料根目录:{{ storage_root }}
  • 数据库文件:{{ database_path }}({{ "已存在" if database_exists else "尚未创建" }})
  • AI 接口配置文件:项目本地 .env({{ "已存在" if ai_config.env_exists else "尚未创建" }})
  • -
  • 1. 音频转写:{{ ai_config['values']['VOLCENGINE_ASR_API_URL'] }} · {{ ai_config['values']['VOLCENGINE_ASR_RESOURCE_ID'] }} · Key {{ "已填写" if ai_config.secret_configured['VOLCENGINE_ASR_API_KEY'] or ai_config.secret_configured['VOLCENGINE_ASR_APP_KEY'] else "未填写" }}
  • +
  • 1. 音频转写:{{ ai_config.transcription_runtime.model_identity }} · {{ ai_config.transcription_runtime.device }}/{{ ai_config.transcription_runtime.compute_type }} · 模型 {{ "已缓存" if ai_config.transcription_runtime.model_ready else "待初始化" }} · GPU {{ "可用" if ai_config.transcription_runtime.gpu_ready else "待初始化" }} · 外部费用 ¥{{ ai_config.transcription_runtime.external_cost_yuan }}
  • +
  • 火山回滚配置:{{ ai_config['values']['VOLCENGINE_ASR_RESOURCE_ID'] }} · Key {{ "已保留" if ai_config.secret_configured['VOLCENGINE_ASR_API_KEY'] or ai_config.secret_configured['VOLCENGINE_ASR_APP_KEY'] else "未填写" }} · {{ "完全离线锁已阻断" if ai_config.transcription_runtime.offline_only else "允许手动回滚" }}
  • 2. 分析文字稿,生成候选切片:{{ ai_config['values']['AI_ANALYSIS_REMOTE_BASE_URL'] }} · {{ ai_config['values']['AI_ANALYSIS_REMOTE_PROTOCOL'] }} · {{ ai_config['values']['AI_ANALYSIS_REMOTE_MODEL'] }} · Key {{ "有效" if ai_config.analysis_key_valid else "异常" }}
  • 3. 发送中心生成发布文案:{{ ai_config['values']['AI_PUBLISH_REMOTE_BASE_URL'] }} · {{ ai_config['values']['AI_PUBLISH_REMOTE_PROTOCOL'] }} · {{ ai_config['values']['AI_PUBLISH_REMOTE_MODEL'] }} · Key {{ "有效" if ai_config.publish_key_valid else "异常" }}
  • 本地 AI:{{ ai_config['values']['AI_LOCAL_BASE_URL'] }} · {{ ai_config['values']['AI_LOCAL_MODEL'] }} · Ollama {{ "在线" if ai_config.local_ollama_online else "离线" }}
  • diff --git a/app/templates/task_detail.html b/app/templates/task_detail.html index 24fc284..0938264 100644 --- a/app/templates/task_detail.html +++ b/app/templates/task_detail.html @@ -34,7 +34,7 @@

    任务详情 · {{ task.title }}

    {% endif %} {% if not task.auto_mode %} - + 进入片段审核 {% endif %} {% if not task.auto_mode and task.output_clip_count > 0 %} diff --git a/docs/DEPLOYMENT.md b/docs/DEPLOYMENT.md index dea9936..95f2d75 100644 --- a/docs/DEPLOYMENT.md +++ b/docs/DEPLOYMENT.md @@ -60,6 +60,22 @@ python -m venv .venv pip install -r requirements.txt ``` +若要启用 Windows NVIDIA GPU 的完全离线转写,再运行一次: + +```powershell +.\scripts\setup_local_transcription.ps1 +``` + +脚本会安装固定版本 cuBLAS 12(及其固定 NVRTC 依赖),并把 `large-v3` 主模型与 `medium` CPU 兜底模型下载到 E 盘。模型大文件按 64 MiB 官方分片下载,可跨次运行续传;合并后必须同时通过固定大小和 SHA-256 校验才会启用。两套模型和运行库首次约需 5.2 GB;初始化完成后,任务运行只读取本地缓存,音频不会上传,也不会在任务过程中自动联网下载。可通过 `-AudioPath "真实音频路径"` 同时执行前 20 秒 GPU 推理冒烟。脚本默认让 PyPI 安装直连,避免本机代理拖慢 500 MB 以上的 Windows wheel;只有网络必须经代理时才增加 `-UseEnvironmentProxyForPip`。 + +长音频验收可使用隔离基准脚本,不会写入正式数据库或任务目录: + +```powershell +python .\scripts\benchmark_local_transcription.py "E:\素材\source.wav" "E:\直播间切片工作流存储\_验收\offline-asr" --seconds 600 +``` + +要验证 checkpoint 续跑,首次增加 `--stop-after-chunks 1`;脚本会在第一个 120 秒分块安全落盘后主动停止。随后使用相同音频、相同验收目录和相同任务 ID,去掉该参数再次运行,日志应出现“已复用第 1/... 段 checkpoint”。脚本会强制 `HF_HUB_OFFLINE=1`、本地 Provider 和固定模型缓存,即使误配置远程 Key 也不会上传音频。 + 看到 `Successfully installed ...` 即为成功。 ### 2.4 配置环境变量 @@ -77,7 +93,9 @@ copy .env.example .env - `AI_CODEX_PATH`:受控 Codex CLI 命令路径(默认 `codex`) - `AI_CODEX_MODEL`:Codex CLI 分析模型(默认 `gpt-5.6-sol`) - `AI_ANALYSIS_REMOTE_API_KEY`:DeepSeek API Key(可选,用远程 AI 分析时需要) -- `VOLCENGINE_ASR_API_KEY`:火山引擎转写 Key(可选,用远程转写时需要) +- `TRANSCRIPTION_MODEL_CACHE_DIR`:离线模型缓存目录,默认 `E:\直播间切片工作流存储\_模型\faster-whisper` +- `TRANSCRIPTION_OFFLINE_ONLY`:默认 `true`,阻断所有远程语音转写请求 +- `VOLCENGINE_ASR_API_KEY`:火山引擎回滚 Key(默认不使用,不会删除) - `LOCAL_ADMIN_TOKEN`:管理接口鉴权 Token(可留空或设随机字符串) ### 2.5 启动服务 @@ -273,7 +291,7 @@ Windows 主机(运行 FastAPI) | `AI_CODEX_PATH` | `codex` | 受控 Codex CLI 命令路径 | | `AI_CODEX_MODEL` | `gpt-5.6-sol` | Codex CLI 分析模型 | | `AI_ANALYSIS_REMOTE_API_KEY` | 空 | DeepSeek API Key | -| `TRANSCRIPTION_PROVIDER` | `volcengine` | 转写引擎:`volcengine` 或 `faster_whisper` | +| `TRANSCRIPTION_PROVIDER` | `local` | 转写引擎:`local` 或 `volcengine`;完全离线模式只允许 `local` | | `AI_PROVIDER` | `codex` | AI 分析引擎:`codex`、`remote` 或 `local` | --- diff --git a/docs/UI_REFERENCE.md b/docs/UI_REFERENCE.md index b4d4a62..af98d9b 100644 --- a/docs/UI_REFERENCE.md +++ b/docs/UI_REFERENCE.md @@ -1,5 +1,20 @@ # UI 参考说明 +## 2026-09-01 更新:简体转写与 AI 安全重试 + +- 本地 faster-whisper 的新转写统一显示简体字形;转换规则为 OpenCC `t2s`,只做繁简字形转换,因此台湾节目中的“计程车、软体”等用词仍会保留。段落、逐词时间戳和字幕复用同一份转换结果。 +- 全自动 AI 覆盖不完整时,任务会停在“AI 分析失败”,不会先显示成选片失败,也不会继续生成切片或发送中心记录。 +- “重试全自动流程”检测到可能已计费但结果不确定的 AI 单元时,先显示包含单元数量和恢复方式的确认框;只有用户再次确认才会发起 AI 调用,取消后原 Job 与证据不变。 +- transcript 未变化时,提示会说明只补跑不确定单元;transcript 已变化时,提示会说明创建全新 AI Job 且不复用旧 AI 单元。有下游切片或发布记录时,页面直接显示拒绝原因。 +- 系统状态中的 Codex 超时默认值调整为 600 秒,仍可在 10–1800 秒范围内人工配置。 + +## 2026-08-31 更新:完全离线本地转写 + +- 系统状态新增“完全离线转写”卡片,直接展示固定模型身份、模型缓存、GPU、cuBLAS、cuDNN 和外部转写费用 `¥0`;缺模型或 DLL 时明确提示运行初始化脚本,不在任务运行中联网下载。 +- “1. 音频转写”默认使用 `local`,固定显示 `large-v3@edaa852e / cuda / float16` 和 E 盘模型目录;火山地址与 Key 继续保留在“回滚配置”中,完全离线锁开启时不会调用。 +- 任务详情和自动流水线未显式选择 Provider 时都服从系统本地设置;失败、取消或卡住后的重试文案改为本地转写语义。显式选择火山且离线锁开启时,页面收到 409 中文提示且不会产生远程请求。 +- 转写进度、完整原文、字幕工作台和 AI 分析继续使用原有输出结构;本轮不增加说话人分离,也不改变后续文本 AI Provider。 + ## 2026-08-26 更新:全站克制动效系统 - 动效遵循“状态反馈优先、装饰最少”:使用原生 CSS/JavaScript,统一短时缓动,不引入 Motion、AutoAnimate、Animate.css、AOS、GSAP 或 View Transitions 运行时。 diff --git a/docs/agent_tasks/2026-08-31-offline-local-transcription.md b/docs/agent_tasks/2026-08-31-offline-local-transcription.md new file mode 100644 index 0000000..4b81eb8 --- /dev/null +++ b/docs/agent_tasks/2026-08-31-offline-local-transcription.md @@ -0,0 +1,57 @@ +# 完全离线长音频转写实施任务 + +## 背景 + +火山引擎长音频转写成本过高。本机具备 RTX 4070 Ti 12GB 和 64GB 内存,项目已包含 faster-whisper、分块 checkpoint、持久化任务和时间戳输出,但默认 Provider 仍偏向远程,且 Windows GPU 运行库与固定模型缓存没有形成可验证的离线闭环。 + +## 目标 + +1. 默认使用固定版本 `large-v3 / cuda / float16`,正式转写不上传音频、不产生云端 ASR 费用。 +2. 首次联网初始化后只从 E 盘固定缓存加载模型,缺模型时明确失败,不在任务中静默下载。 +3. 保留 120 秒分块、5 秒重叠、SQLite checkpoint、词级时间戳及 medium CPU 本地兜底。 +4. 系统状态页和任务状态 API 提供模型、GPU、离线锁和实际版本证据。 + +## 允许修改范围 + +- 转写配置、Provider 解析、运行时加载、checkpoint 指纹与转写状态 API。 +- 系统状态页、相关 JavaScript、一次性初始化/诊断脚本。 +- 转写定向测试、项目进度文档和 UI 参考文档。 +- Windows GPU 可选依赖清单。 + +## 禁止修改范围 + +- 不删除或回显火山 Key、`.env`、历史 transcript、数据库记录和模型缓存。 +- 不增加说话人分离、降噪、人声分离或新的云端 ASR Provider。 +- 不改变下游字幕、AI 分析、切片、排期和发布数据格式。 +- 不重启正式 Web/Worker,不触发真实平台发布或远程转写。 + +## 已确定实现要求 + +- 主模型锁定 `Systran/faster-whisper-large-v3@edaa852ec7e145841d8ffdb056a99866b5f0a478`。 +- CPU 兜底锁定 `Systran/faster-whisper-medium@08e178d48790749d25932bbc082711ddcfdfbc4f`。 +- Windows 固定安装 `nvidia-cublas-cu12==12.9.2.10` 及其 NVRTC 依赖;继续使用 CTranslate2 自带 cuDNN 9,不再安装第二套 cuDNN。 +- 完全离线锁开启时,远程 Provider 在请求前失败;HTTP 路由返回 409。 +- 转写 checkpoint 使用完整音频 SHA-256,其他视频/封面指纹行为不变。 +- 转写 run 的 model 字段记录 `模型名@revision前8位`,旧记录保持可读。 + +## 验收标准 + +- 默认、普通任务、自动流水线都选择 local;离线锁下无法创建远程转写任务。 +- 模型未缓存、cuBLAS 缺失、CUDA 不可用时返回准确中文诊断。 +- 已完成分块可恢复,模型 revision 或音频内容改变时不复用旧 checkpoint。 +- 系统状态页不泄露 Secret,显示离线、模型、GPU、设备和外部费用 ¥0。 +- 20 秒真实中文音频在 GPU 上完成;代表性长音频验证另行记录,未获确认前不切换正式运行配置。 + +## 测试命令 + +- `.venv\\Scripts\\python.exe -m pytest tests/test_offline_transcription.py tests/test_long_live_foundation.py tests/test_transcription_checkpoint_resilience.py -q` +- `.venv\\Scripts\\python.exe -m pytest -q` +- `.venv\\Scripts\\python.exe -m ruff check app tests scripts` +- `.venv\\Scripts\\python.exe -m compileall -q app scripts` +- `node --check app/static/js/app.js` +- `git diff --check` + +## 返回格式 + +- 返回分支、提交、远端 SHA、PR、CI、测试数量与真实 GPU 冒烟证据。 +- 明确区分代码/自动化验证、模型下载、真实音频验收和正式运行时切换状态。 diff --git a/docs/agent_tasks/2026-09-01-latest-task-transcription-recovery.md b/docs/agent_tasks/2026-09-01-latest-task-transcription-recovery.md new file mode 100644 index 0000000..9733393 --- /dev/null +++ b/docs/agent_tasks/2026-09-01-latest-task-transcription-recovery.md @@ -0,0 +1,67 @@ +# 最新任务失败与本地转写简体化修复任务 + +## 背景 + +- 最新任务 `614fb38401e0` 已完成本地转写,但 AI 分析的一个召回窗口超时,覆盖率停在 `17/18`。 +- 当前本地 faster-whisper 输出包含繁体字;用户确认今后统一转换为简体字,只转换字形并保留台湾用词。 +- 现有全自动重试入口无法正确补跑 AI 分析,也没有针对“可能已计费但结果未知”的二次确认。 + +## 目标 + +1. 本地转写的段落与逐词结果在写入 checkpoint 前统一执行 OpenCC `t2s` 转换。 +2. 将转换规则加入本地转写 checkpoint 身份,避免复用旧繁体分块。 +3. AI 覆盖不完整时在 AI 分析阶段失败,并保留选片阶段的防御性门禁。 +4. 全自动重试支持结构化确认、未变化输入的失败单元续跑,以及变化输入的全新 AI Job。 +5. 将 Codex 默认窗口超时调整为 600 秒,并更新页面交互、项目文档和测试。 +6. 安全恢复任务 `614fb38401e0`,不重复转写、不触发真实发布。 + +## 允许修改范围 + +- `app/` 中转写、AI 流水线、Job 重试、任务路由和任务详情交互相关文件。 +- `tests/` 中与本次行为对应的自动化测试。 +- `requirements.in`、`requirements.txt`、`.env.example`。 +- `DEVELOPMENT_LOG.md`、`NEXT_STEPS.md`、`docs/UI_REFERENCE.md` 和本任务文件。 +- 通过测试后,将代码提交同步到运行工作树并定向调整运行 `.env` 的超时配置。 + +## 禁止修改范围 + +- 主目录已有 `.codemap/*` 与 `PROJECT_REAUDIT.md`。 +- 发布 Worker、真实抖音/B站发布、排期、账号登录与平台验证。 +- PR 合并、强推、历史改写、数据库结构和既有失败 Job 证据。 +- API Key、Token、Cookie、密码及其他敏感配置内容。 + +## 已确定实现要求 + +- 固定依赖 `opencc==1.4.2`,使用 `t2s.json`,不使用词汇本地化转换。 +- 当前 transcript 变化后从 `AI_ANALYZING` 新建独立 Job,不继承旧 AI 单元。 +- 输入不变时,确认后只重置不确定单元并复用成功 checkpoint。 +- 未传 `confirm_uncertain_ai=true` 时返回 HTTP 409 和结构化确认信息。 +- 有下游切片或发布记录时拒绝覆盖式恢复。 +- Codex 超时默认值和运行值均设为 600 秒。 + +## 验收标准 + +- 段落和逐词文本均转为简体;台湾用词和非中文内容不被错误改写。 +- 旧繁体 checkpoint 不会被新规则复用。 +- AI 覆盖不完整落入 `FAILED_AI_ANALYZING`。 +- 409 确认、同输入单元续跑、变输入新 Job、下游保护均有测试覆盖。 +- 相关测试、完整 pytest、Ruff、compileall、JavaScript 语法检查、`git diff --check` 和 Windows smoke/release gate 通过。 +- 正式 transcript 有可恢复备份,转换后时间戳、行数和 Markdown 结构不变。 +- 新 AI Job 达到 100% 覆盖并进入正常后续阶段,且无发布记录。 + +## 测试命令 + +- `python -m pytest tests/test_offline_transcription.py tests/test_pipeline_checkpoint.py tests/test_auto_pipeline.py` +- `python -m pytest` +- `python -m ruff check app tests` +- `python -m compileall app tests` +- `node --check app/static/js/app.js` +- `git diff --check` +- 项目现有 Windows smoke/release gate。 + +## 返回格式 + +- 修改文件及核心行为。 +- 测试命令、结果与失败证据(如有)。 +- Git 分支、提交、Push、PR、CI 与未合并状态。 +- 运行服务重启验证、当前任务新 Job、AI 覆盖率和发布记录状态。 diff --git a/requirements-windows-gpu.txt b/requirements-windows-gpu.txt new file mode 100644 index 0000000..d49710e --- /dev/null +++ b/requirements-windows-gpu.txt @@ -0,0 +1,4 @@ +# Optional Windows GPU runtime for local faster-whisper transcription. +# Install after requirements.txt with scripts/setup_local_transcription.ps1. +nvidia-cublas-cu12==12.9.2.10 +nvidia-cuda-nvrtc-cu12==12.9.86 diff --git a/requirements.in b/requirements.in index a50c48c..5494a2b 100644 --- a/requirements.in +++ b/requirements.in @@ -10,4 +10,5 @@ tzdata>=2026.1 playwright>=1.58,<1.63 aiofiles>=25.1,<26 faster-whisper>=1.2,<1.3 +opencc==1.4.2 pysubs2==1.9.0 diff --git a/requirements.txt b/requirements.txt index 0b340de..b86e0af 100644 --- a/requirements.txt +++ b/requirements.txt @@ -13,4 +13,5 @@ tzdata==2026.3 playwright==1.62.0 aiofiles==25.1.0 faster-whisper==1.2.1 +opencc==1.4.2 pysubs2==1.9.0 diff --git a/scripts/benchmark_local_transcription.py b/scripts/benchmark_local_transcription.py new file mode 100644 index 0000000..cef615a --- /dev/null +++ b/scripts/benchmark_local_transcription.py @@ -0,0 +1,179 @@ +from __future__ import annotations + +import argparse +import json +import os +import sys +import time +from pathlib import Path + + +PROJECT_ROOT = Path(__file__).resolve().parents[1] +DEFAULT_MODEL_CACHE_DIR = Path(r"E:\直播间切片工作流存储\_模型\faster-whisper") + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="使用隔离数据库验证本地长音频转写、分块 checkpoint 和续跑能力。", + ) + parser.add_argument("audio_path", help="本地音频文件路径") + parser.add_argument("workspace_dir", help="隔离验收目录;续跑时必须使用同一目录") + parser.add_argument("--seconds", type=int, default=0, help="仅验证前 N 秒;0 表示完整音频") + parser.add_argument("--task-id", default="offline-local-benchmark", help="隔离数据库中的固定任务 ID") + parser.add_argument( + "--stop-after-chunks", + type=int, + default=0, + help="完成指定数量分块并保存 checkpoint 后主动停止;0 表示运行到结束", + ) + parser.add_argument( + "--model-cache-dir", + default=str(DEFAULT_MODEL_CACHE_DIR), + help="固定版本 faster-whisper 模型缓存根目录", + ) + return parser.parse_args() + + +def configure_isolated_environment(args: argparse.Namespace, workspace_dir: Path) -> None: + data_dir = workspace_dir / "data" + storage_dir = workspace_dir / "storage" + environment = { + "DATA_DIR": str(data_dir), + "DATABASE_PATH": str(data_dir / "benchmark.sqlite3"), + "STORAGE_ROOT": str(storage_dir), + "TASKS_DIR": str(storage_dir), + "UPLOAD_TEMP_DIR": str(storage_dir / "_uploads"), + "TRANSCRIPTION_PROVIDER": "local", + "TRANSCRIPTION_FALLBACK_PROVIDER": "", + "TRANSCRIPTION_OFFLINE_ONLY": "true", + "TRANSCRIPTION_MODEL": "large-v3", + "TRANSCRIPTION_MODEL_REVISION": "edaa852ec7e145841d8ffdb056a99866b5f0a478", + "TRANSCRIPTION_MODEL_CACHE_DIR": str(Path(args.model_cache_dir).resolve()), + "TRANSCRIPTION_LOCAL_FILES_ONLY": "true", + "TRANSCRIPTION_DEVICE": "cuda", + "TRANSCRIPTION_COMPUTE_TYPE": "float16", + "HF_HUB_OFFLINE": "1", + "TRANSFORMERS_OFFLINE": "1", + } + os.environ.update(environment) + + +def main() -> int: + args = parse_args() + audio_path = Path(args.audio_path).resolve() + workspace_dir = Path(args.workspace_dir).resolve() + if not audio_path.is_file(): + raise SystemExit(f"音频文件不存在:{audio_path}") + if args.seconds < 0: + raise SystemExit("--seconds 不能小于 0") + if args.stop_after_chunks < 0: + raise SystemExit("--stop-after-chunks 不能小于 0") + + workspace_dir.mkdir(parents=True, exist_ok=True) + configure_isolated_environment(args, workspace_dir) + if str(PROJECT_ROOT) not in sys.path: + sys.path.insert(0, str(PROJECT_ROOT)) + + from app.db.database import get_connection, init_db + from app.models.task import TaskCreate + from app.services.task_lifecycle_service import create_task_record + from app.services.transcript_service import ( + TranscriptChunk, + _extract_audio_chunk, + write_transcript_markdown, + ) + + init_db() + with get_connection() as connection: + task_exists = connection.execute("SELECT 1 FROM tasks WHERE id = ?", (args.task_id,)).fetchone() + if not task_exists: + create_task_record( + TaskCreate(task_name="完全离线长音频验收", selection_profile="general"), + task_id=args.task_id, + task_dir_name=args.task_id, + ) + + benchmark_audio_path = audio_path + if args.seconds: + benchmark_audio_path = workspace_dir / f"source-first-{args.seconds}s.wav" + if not benchmark_audio_path.exists(): + _extract_audio_chunk( + audio_path, + benchmark_audio_path, + TranscriptChunk(index=1, start_seconds=0, end_seconds=max(1, args.seconds)), + ) + + transcript_path = workspace_dir / "transcript.md" + started_at = time.perf_counter() + reuse_count = 0 + + def report_progress(progress: dict) -> None: + nonlocal reuse_count + message = str(progress.get("message") or "") + if "已复用第" in message and "checkpoint" in message: + reuse_count += 1 + print(json.dumps(progress, ensure_ascii=False), flush=True) + if ( + args.stop_after_chunks + and message.startswith(f"已完成第 {args.stop_after_chunks}/") + ): + raise RuntimeError( + f"验收脚本按计划在第 {args.stop_after_chunks} 个分块落盘后停止;" + "使用相同命令去掉 --stop-after-chunks 即可续跑。" + ) + + try: + result = write_transcript_markdown( + { + "id": args.task_id, + "task_name": "完全离线长音频验收", + "source": str(audio_path), + }, + benchmark_audio_path, + transcript_path, + progress_callback=report_progress, + provider="local", + ) + except RuntimeError as exc: + elapsed_seconds = round(time.perf_counter() - started_at, 2) + if "验收脚本按计划" in str(exc): + print( + json.dumps( + { + "status": "stopped_after_checkpoint", + "elapsed_seconds": elapsed_seconds, + "reason": str(exc), + }, + ensure_ascii=False, + ), + flush=True, + ) + return 75 + raise + + progress_path = transcript_path.with_name("transcript_progress.json") + progress = json.loads(progress_path.read_text(encoding="utf-8")) + elapsed_seconds = round(time.perf_counter() - started_at, 2) + print( + json.dumps( + { + "status": "completed", + "elapsed_seconds": elapsed_seconds, + "reuse_count": reuse_count, + "audio_path": str(benchmark_audio_path), + "transcript_path": str(transcript_path), + "provider": result.get("provider"), + "model": progress.get("model"), + "device": progress.get("device"), + "compute_type": progress.get("compute_type"), + "segment_count": result.get("segment_count"), + }, + ensure_ascii=False, + ), + flush=True, + ) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/scripts/diagnose_transcription_environment.py b/scripts/diagnose_transcription_environment.py index 50ff500..f351cdb 100644 --- a/scripts/diagnose_transcription_environment.py +++ b/scripts/diagnose_transcription_environment.py @@ -8,7 +8,8 @@ if str(PROJECT_ROOT) not in sys.path: sys.path.insert(0, str(PROJECT_ROOT)) -from app.core.config import settings +from app.core.config import settings # noqa: E402 +from app.services.local_transcription_runtime import get_local_transcription_runtime_status # noqa: E402 def check_command(name: str) -> tuple[bool, str]: @@ -43,7 +44,11 @@ def main() -> None: print("=== 当前转写配置 ===") print(f"TRANSCRIPTION_PROVIDER={settings.transcription_provider}") print(f"TRANSCRIPTION_FALLBACK_PROVIDER={settings.transcription_fallback_provider}") + print(f"TRANSCRIPTION_OFFLINE_ONLY={settings.transcription_offline_only}") print(f"TRANSCRIPTION_MODEL={settings.transcription_model}") + print(f"TRANSCRIPTION_MODEL_REVISION={settings.transcription_model_revision}") + print(f"TRANSCRIPTION_MODEL_CACHE_DIR={settings.transcription_model_cache_dir}") + print(f"TRANSCRIPTION_LOCAL_FILES_ONLY={settings.transcription_local_files_only}") print(f"TRANSCRIPTION_LANGUAGE={settings.transcription_language}") print(f"TRANSCRIPTION_DEVICE={settings.transcription_device}") print(f"TRANSCRIPTION_COMPUTE_TYPE={settings.transcription_compute_type}") @@ -55,9 +60,18 @@ def main() -> None: print(f"VOLCENGINE_ASR_API_KEY={'已填写' if settings.volcengine_asr_api_key else '未填写'}") print(f"VOLCENGINE_ASR_AUDIO_FORMAT={settings.volcengine_asr_audio_format}") - if settings.transcription_device.lower() == "cuda": - print() - print("提示:当前配置使用 CUDA。如果仍然看到 cublas64_12.dll 或 cudnn 错误,请先改用 CPU/int8 跑通。") + runtime = get_local_transcription_runtime_status() + print() + print("=== 本地离线转写就绪状态 ===") + print(f"模型固定身份:{runtime['model_identity']}") + print(f"模型已缓存:{'是' if runtime['model_ready'] else '否'}") + print(f"GPU 可用:{'是' if runtime['gpu_ready'] else '否'}") + print(f"CUDA 设备数:{runtime['cuda_device_count']}") + print(f"cuBLAS 12:{'正常' if runtime['cublas_ready'] else '缺失'}") + print(f"cuDNN 9:{'正常' if runtime['cudnn_ready'] else '缺失'}") + print(f"外部转写费用:{runtime['external_cost']}") + if runtime["errors"]: + print("待处理:" + ";".join(runtime["errors"])) if __name__ == "__main__": diff --git a/scripts/normalize_transcript_simplified.py b/scripts/normalize_transcript_simplified.py new file mode 100644 index 0000000..d4b2718 --- /dev/null +++ b/scripts/normalize_transcript_simplified.py @@ -0,0 +1,95 @@ +"""备份并原位把 transcript.md 转为 OpenCC t2s 简体字形。""" + +from __future__ import annotations + +import argparse +from datetime import datetime +import hashlib +import json +from pathlib import Path +import re +import shutil +import sys + + +PROJECT_ROOT = Path(__file__).resolve().parents[1] +if str(PROJECT_ROOT) not in sys.path: + sys.path.insert(0, str(PROJECT_ROOT)) + +from app.services.chinese_text_service import simplify_chinese_text # noqa: E402 + + +_TIMESTAMP_PATTERN = re.compile(r"\b\d{2}:\d{2}:\d{2}\b") + + +def normalize_transcript_file(transcript_path: Path, backup_dir: Path) -> dict: + transcript_path = transcript_path.resolve() + backup_dir = backup_dir.resolve() + if not transcript_path.is_file(): + raise FileNotFoundError(f"未找到 transcript:{transcript_path}") + original_bytes = transcript_path.read_bytes() + original = original_bytes.decode("utf-8") + converted = simplify_chinese_text(original) + _validate_structure(original, converted) + + backup_dir.mkdir(parents=True, exist_ok=True) + timestamp = datetime.now().strftime("%Y%m%d-%H%M%S") + backup_path = backup_dir / f"transcript.before-t2s-{timestamp}.md" + if backup_path.exists(): + raise FileExistsError(f"备份文件已存在:{backup_path}") + shutil.copy2(transcript_path, backup_path) + + temp_path = transcript_path.with_name(f"{transcript_path.name}.t2s.tmp") + try: + temp_path.write_bytes(converted.encode("utf-8")) + if temp_path.read_bytes().decode("utf-8") != converted: + raise RuntimeError("转换后的临时文件回读校验失败") + temp_path.replace(transcript_path) + finally: + if temp_path.exists(): + temp_path.unlink() + + final_bytes = transcript_path.read_bytes() + return { + "status": "converted" if final_bytes != original_bytes else "unchanged", + "transcript_path": str(transcript_path), + "backup_path": str(backup_path), + "before_sha256": hashlib.sha256(original_bytes).hexdigest(), + "after_sha256": hashlib.sha256(final_bytes).hexdigest(), + "line_count": len(original.splitlines()), + "timestamp_count": len(_TIMESTAMP_PATTERN.findall(original)), + } + + +def _validate_structure(original: str, converted: str) -> None: + original_lines = original.splitlines() + converted_lines = converted.splitlines() + if len(original_lines) != len(converted_lines): + raise RuntimeError("转换前后行数变化,已停止覆盖 transcript") + if _TIMESTAMP_PATTERN.findall(original) != _TIMESTAMP_PATTERN.findall(converted): + raise RuntimeError("转换前后时间戳变化,已停止覆盖 transcript") + if [line.count("|") for line in original_lines] != [line.count("|") for line in converted_lines]: + raise RuntimeError("转换前后 Markdown 表格结构变化,已停止覆盖 transcript") + if [len(line) - len(line.lstrip("#")) for line in original_lines] != [ + len(line) - len(line.lstrip("#")) for line in converted_lines + ]: + raise RuntimeError("转换前后 Markdown 标题结构变化,已停止覆盖 transcript") + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--transcript", type=Path, required=True) + parser.add_argument("--backup-dir", type=Path, required=True) + args = parser.parse_args() + print( + json.dumps( + normalize_transcript_file(args.transcript, args.backup_dir), + ensure_ascii=False, + indent=2, + ) + ) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/scripts/setup_local_transcription.ps1 b/scripts/setup_local_transcription.ps1 new file mode 100644 index 0000000..72eb800 --- /dev/null +++ b/scripts/setup_local_transcription.ps1 @@ -0,0 +1,92 @@ +param( + [string]$PythonPath = "", + [string]$ModelCacheDir = "", + [string]$AudioPath = "", + [ValidateRange(1, 120)] + [int]$Seconds = 20, + [switch]$SkipDependencyInstall, + [switch]$UseEnvironmentProxyForPip +) + +$ErrorActionPreference = "Stop" +$projectRoot = Split-Path -Parent $PSScriptRoot + +function Install-VerifiedWheel { + param( + [Parameter(Mandatory = $true)] + [string]$Url, + [Parameter(Mandatory = $true)] + [string]$Sha256, + [Parameter(Mandatory = $true)] + [string]$Destination + ) + + if (Test-Path -LiteralPath $Destination) { + $existingHash = (Get-FileHash -Algorithm SHA256 -LiteralPath $Destination).Hash.ToLowerInvariant() + if ($existingHash -ne $Sha256) { + Remove-Item -LiteralPath $Destination + } + } + if (-not (Test-Path -LiteralPath $Destination)) { + & curl.exe --noproxy "*" --location --fail --retry 10 --retry-delay 2 --continue-at - --output $Destination $Url + if ($LASTEXITCODE -ne 0) { + throw "GPU 运行库 wheel 下载失败:$Destination" + } + } + $actualHash = (Get-FileHash -Algorithm SHA256 -LiteralPath $Destination).Hash.ToLowerInvariant() + if ($actualHash -ne $Sha256) { + throw "GPU 运行库 wheel SHA-256 校验失败:$Destination" + } +} + +if (-not $PythonPath) { + $PythonPath = Join-Path $projectRoot ".venv\Scripts\python.exe" +} +$resolvedPython = Resolve-Path -LiteralPath $PythonPath -ErrorAction SilentlyContinue +if (-not $resolvedPython) { + throw "找不到项目 Python:$PythonPath。请通过 -PythonPath 指定 NiuMa Studio 的 .venv Python。" +} +$pythonExe = $resolvedPython.Path + +if (-not $SkipDependencyInstall) { + if ($UseEnvironmentProxyForPip) { + & $pythonExe -m pip install --disable-pip-version-check --timeout 600 --retries 5 -r (Join-Path $projectRoot "requirements-windows-gpu.txt") + if ($LASTEXITCODE -ne 0) { + throw "Windows GPU 运行库安装失败,未继续下载模型。" + } + } + else { + $wheelDirectory = Join-Path ([System.IO.Path]::GetTempPath()) "niuma-offline-transcription-wheels" + New-Item -ItemType Directory -Force -Path $wheelDirectory | Out-Null + $cublasWheel = Join-Path $wheelDirectory "nvidia_cublas_cu12-12.9.2.10-py3-none-win_amd64.whl" + $nvrtcWheel = Join-Path $wheelDirectory "nvidia_cuda_nvrtc_cu12-12.9.86-py3-none-win_amd64.whl" + Install-VerifiedWheel ` + -Url "https://files.pythonhosted.org/packages/20/e2/fc9a0e985249d873150276d5afb02e39a66817fedbf1a385724393e505ed/nvidia_cublas_cu12-12.9.2.10-py3-none-win_amd64.whl" ` + -Sha256 "623f43027d40d44ceadf0043f002bd25cf353e8f13ce90b9a87057019f560661" ` + -Destination $cublasWheel + Install-VerifiedWheel ` + -Url "https://files.pythonhosted.org/packages/52/de/823919be3b9d0ccbf1f784035423c5f18f4267fb0123558d58b813c6ec86/nvidia_cuda_nvrtc_cu12-12.9.86-py3-none-win_amd64.whl" ` + -Sha256 "72972ebdcf504d69462d3bcd67e7b81edd25d0fb85a2c46d3ea3517666636349" ` + -Destination $nvrtcWheel + & $pythonExe -m pip install --disable-pip-version-check --no-index --no-deps $nvrtcWheel $cublasWheel + if ($LASTEXITCODE -ne 0) { + throw "Windows GPU 运行库本地 wheel 安装失败,未继续下载模型。" + } + } +} + +$arguments = @( + "-X", "utf8", + (Join-Path $projectRoot "scripts\setup_local_transcription.py") +) +if ($ModelCacheDir) { + $arguments += @("--cache-dir", $ModelCacheDir) +} +if ($AudioPath) { + $arguments += @("--audio", $AudioPath, "--seconds", $Seconds) +} + +& $pythonExe @arguments +if ($LASTEXITCODE -ne 0) { + throw "本地离线转写初始化失败。" +} diff --git a/scripts/setup_local_transcription.py b/scripts/setup_local_transcription.py new file mode 100644 index 0000000..6d3fe6f --- /dev/null +++ b/scripts/setup_local_transcription.py @@ -0,0 +1,472 @@ +from __future__ import annotations + +import argparse +from concurrent.futures import ThreadPoolExecutor, as_completed +from datetime import datetime, timezone +import hashlib +import json +import os +from pathlib import Path +import shutil +import subprocess +import sys +from tempfile import TemporaryDirectory +import time +from urllib.request import Request, getproxies, urlopen + + +PROJECT_ROOT = Path(__file__).resolve().parents[1] +if str(PROJECT_ROOT) not in sys.path: + sys.path.insert(0, str(PROJECT_ROOT)) + +from app.core.config import settings # noqa: E402 +from app.core.transcription_defaults import ( # noqa: E402 + CPU_FALLBACK_TRANSCRIPTION_MODEL, + CPU_FALLBACK_TRANSCRIPTION_MODEL_BIN_SHA256, + CPU_FALLBACK_TRANSCRIPTION_MODEL_BIN_SIZE, + CPU_FALLBACK_TRANSCRIPTION_MODEL_REPOSITORY, + CPU_FALLBACK_TRANSCRIPTION_MODEL_REVISION, + PRIMARY_TRANSCRIPTION_MODEL, + PRIMARY_TRANSCRIPTION_MODEL_BIN_SHA256, + PRIMARY_TRANSCRIPTION_MODEL_BIN_SIZE, + PRIMARY_TRANSCRIPTION_MODEL_REPOSITORY, + PRIMARY_TRANSCRIPTION_MODEL_REVISION, +) +from app.services.local_transcription_runtime import ( # noqa: E402 + MODEL_VOCABULARY_FILES, + OPTIONAL_MODEL_FILES, + REQUIRED_MODEL_FILES, + configure_windows_cuda_dll_directories, + model_cache_directory, + model_files_ready, +) + + +MODEL_SPECS = ( + ( + PRIMARY_TRANSCRIPTION_MODEL, + PRIMARY_TRANSCRIPTION_MODEL_REPOSITORY, + PRIMARY_TRANSCRIPTION_MODEL_REVISION, + PRIMARY_TRANSCRIPTION_MODEL_BIN_SIZE, + PRIMARY_TRANSCRIPTION_MODEL_BIN_SHA256, + "GPU 主模型", + ), + ( + CPU_FALLBACK_TRANSCRIPTION_MODEL, + CPU_FALLBACK_TRANSCRIPTION_MODEL_REPOSITORY, + CPU_FALLBACK_TRANSCRIPTION_MODEL_REVISION, + CPU_FALLBACK_TRANSCRIPTION_MODEL_BIN_SIZE, + CPU_FALLBACK_TRANSCRIPTION_MODEL_BIN_SHA256, + "CPU 兜底模型", + ), +) +_MODEL_DOWNLOAD_ATTEMPTS = 20 +_CURL_DOWNLOAD_ATTEMPTS = 3 +_MODEL_BIN_PART_SIZE = 64 * 1024 * 1024 +_MODEL_BIN_DOWNLOAD_WORKERS = 8 + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="初始化 NiuMa Studio 完全离线转写模型。") + parser.add_argument("--cache-dir", default=str(settings.transcription_model_cache_dir)) + parser.add_argument("--audio", default="", help="可选:真实音频路径,用于 GPU 推理冒烟。") + parser.add_argument("--seconds", type=int, default=20, help="真实音频冒烟截取秒数。") + return parser.parse_args() + + +def _sha256_file(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as source: + for block in iter(lambda: source.read(4 * 1024 * 1024), b""): + digest.update(block) + return digest.hexdigest() + + +def _download_model_bin_part(url: str, target: Path, start: int, end: int) -> int: + if shutil.which("curl"): + try: + return _download_model_bin_part_with_curl(url, target, start, end) + except Exception as exc: + print( + f"官方模型分片 curl 连续失败,切换 Python Range 续传:{start}-{end};{exc}", + flush=True, + ) + return _download_model_bin_part_with_urllib(url, target, start, end) + + +def _append_download_tail(target: Path, tail: Path, expected_size: int) -> None: + if not tail.exists(): + return + current_size = target.stat().st_size if target.exists() else 0 + if current_size + tail.stat().st_size > expected_size: + tail.unlink() + raise RuntimeError(f"官方模型分片响应超过预期大小:{target.name}") + with target.open("ab" if current_size else "wb") as output, tail.open("rb") as source: + shutil.copyfileobj(source, output, length=1024 * 1024) + tail.unlink() + + +def _download_model_bin_part_with_curl(url: str, target: Path, start: int, end: int) -> int: + expected_size = end - start + 1 + tail = target.with_suffix(".download") + _append_download_tail(target, tail, expected_size) + proxy = getproxies().get("https") or getproxies().get("http") or "" + child_environment = os.environ.copy() + if proxy: + child_environment["HTTPS_PROXY"] = proxy + for attempt in range(1, _CURL_DOWNLOAD_ATTEMPTS + 1): + current_size = target.stat().st_size if target.exists() else 0 + if current_size == expected_size: + return expected_size + if current_size > expected_size: + target.unlink() + current_size = 0 + + request_start = start + current_size + command = [ + "curl", + "--location", + "--fail", + "--silent", + "--show-error", + "--connect-timeout", + "20", + "--max-time", + "60", + "--range", + f"{request_start}-{end}", + "--output", + str(tail), + ] + command.append(url) + result = subprocess.run( + command, + capture_output=True, + text=True, + check=False, + env=child_environment, + ) + _append_download_tail(target, tail, expected_size) + updated_size = target.stat().st_size if target.exists() else 0 + if updated_size == expected_size: + return expected_size + if attempt >= _CURL_DOWNLOAD_ATTEMPTS: + raise RuntimeError( + f"官方模型分片 curl 下载失败:{start}-{end};" + f"exit={result.returncode},error={result.stderr.strip()[:300]}" + ) + time.sleep(min(20, attempt * 3)) + raise RuntimeError(f"官方模型分片 curl 下载失败:{start}-{end}") + + +def _download_model_bin_part_with_urllib(url: str, target: Path, start: int, end: int) -> int: + expected_size = end - start + 1 + for attempt in range(1, _MODEL_DOWNLOAD_ATTEMPTS + 1): + current_size = target.stat().st_size if target.exists() else 0 + if current_size == expected_size: + return expected_size + if current_size > expected_size: + target.unlink() + current_size = 0 + + request_start = start + current_size + request = Request( + url, + headers={ + "Range": f"bytes={request_start}-{end}", + "User-Agent": "NiuMa-Studio-offline-transcription-setup/1.0", + }, + ) + try: + with urlopen(request, timeout=120) as response: + if getattr(response, "status", 0) != 206: + raise RuntimeError(f"官方模型分片接口未返回 206:{getattr(response, 'status', 0)}") + with target.open("ab" if current_size else "wb") as output: + while True: + block = response.read(1024 * 1024) + if not block: + break + output.write(block) + except Exception: + if attempt >= _MODEL_DOWNLOAD_ATTEMPTS: + raise + time.sleep(min(20, attempt * 3)) + continue + + if target.stat().st_size == expected_size: + return expected_size + raise RuntimeError(f"官方模型分片下载失败:{start}-{end}") + + +def _download_verified_model_bin( + target_dir: Path, + repository: str, + revision: str, + expected_size: int, + expected_sha256: str, + label: str, +) -> Path: + target = target_dir / "model.bin" + if target.is_file(): + if target.stat().st_size != expected_size: + raise RuntimeError(f"{label} 已有 model.bin 大小异常,请人工移走后重新初始化:{target}") + if _sha256_file(target) != expected_sha256: + raise RuntimeError(f"{label} 已有 model.bin SHA-256 不一致,请人工检查:{target}") + print(f"[{label}] model.bin 已存在且 SHA-256 校验通过") + return target + + part_dir = target_dir / ".model-bin-parts" + part_dir.mkdir(parents=True, exist_ok=True) + url = f"https://huggingface.co/{repository}/resolve/{revision}/model.bin?download=true" + ranges: list[tuple[int, int, Path]] = [] + for index, start in enumerate(range(0, expected_size, _MODEL_BIN_PART_SIZE)): + end = min(expected_size - 1, start + _MODEL_BIN_PART_SIZE - 1) + ranges.append((start, end, part_dir / f"part-{index:04d}.bin")) + + completed_bytes = sum( + end - start + 1 + for start, end, part in ranges + if part.is_file() and part.stat().st_size == end - start + 1 + ) + print( + f"[{label}] 从 Hugging Face 官方固定 revision 下载 model.bin:" + f"{completed_bytes}/{expected_size} 字节已缓存", + flush=True, + ) + failures: list[Exception] = [] + with ThreadPoolExecutor(max_workers=min(_MODEL_BIN_DOWNLOAD_WORKERS, len(ranges))) as executor: + future_to_range = { + executor.submit(_download_model_bin_part, url, part, start, end): (start, end) + for start, end, part in ranges + } + finished_parts = 0 + for future in as_completed(future_to_range): + try: + future.result() + except Exception as exc: + failures.append(exc) + print(f"[{label}] 一个分片仍未完成,其他分片继续:{exc}", flush=True) + else: + finished_parts += 1 + print( + f"[{label}] 官方 model.bin 分片完成:{finished_parts}/{len(ranges)}", + flush=True, + ) + if failures: + raise RuntimeError( + f"{label} 仍有 {len(failures)} 个分片未完成;已保留全部分片,下次运行只补缺失字节。" + ) from failures[0] + + temporary = target.with_suffix(".bin.tmp") + digest = hashlib.sha256() + total_size = 0 + with temporary.open("wb") as output: + for _start, _end, part in ranges: + with part.open("rb") as source: + while True: + block = source.read(4 * 1024 * 1024) + if not block: + break + output.write(block) + digest.update(block) + total_size += len(block) + actual_sha256 = digest.hexdigest() + if total_size != expected_size or actual_sha256 != expected_sha256: + temporary.unlink(missing_ok=True) + raise RuntimeError( + f"{label} model.bin 完整性校验失败:size={total_size}/{expected_size}," + f"sha256={actual_sha256}/{expected_sha256}" + ) + temporary.replace(target) + for _start, _end, part in ranges: + part.unlink(missing_ok=True) + try: + part_dir.rmdir() + except OSError: + pass + print(f"[{label}] model.bin 官方 SHA-256 校验通过:{actual_sha256}") + return target + + +def download_model( + model_name: str, + repository: str, + revision: str, + model_bin_size: int, + model_bin_sha256: str, + label: str, +) -> Path: + from huggingface_hub import snapshot_download + + target = model_cache_directory(model_name, revision) + target.mkdir(parents=True, exist_ok=True) + print(f"[{label}] 固定版本:{repository}@{revision}") + print(f"[{label}] 本地目录:{target}") + for attempt in range(1, _MODEL_DOWNLOAD_ATTEMPTS + 1): + try: + snapshot_download( + repo_id=repository, + revision=revision, + local_dir=target, + allow_patterns=[ + name + for name in (*REQUIRED_MODEL_FILES, *MODEL_VOCABULARY_FILES, *OPTIONAL_MODEL_FILES) + if name != "model.bin" + ], + max_workers=2, + ) + break + except Exception as exc: + if attempt >= _MODEL_DOWNLOAD_ATTEMPTS: + raise RuntimeError( + f"{label} 固定版本下载连续 {_MODEL_DOWNLOAD_ATTEMPTS} 次失败," + "已保留本地缓存,可稍后重新运行脚本继续。" + ) from exc + delay_seconds = min(30, attempt * 5) + print( + f"[{label}] 第 {attempt}/{_MODEL_DOWNLOAD_ATTEMPTS} 次下载中断:{exc}。" + f"保留缓存,{delay_seconds} 秒后重试。", + flush=True, + ) + time.sleep(delay_seconds) + _download_verified_model_bin( + target, + repository, + revision, + model_bin_size, + model_bin_sha256, + label, + ) + if not model_files_ready(target): + missing = [name for name in REQUIRED_MODEL_FILES if not (target / name).is_file()] + if not any((target / name).is_file() for name in MODEL_VOCABULARY_FILES): + missing.append("vocabulary.json 或 vocabulary.txt") + raise RuntimeError(f"{label} 缓存不完整,缺少:{', '.join(missing)}") + return target + + +def write_manifest(cache_dir: Path, model_paths: dict[str, Path]) -> Path: + payload = { + "schema_version": 1, + "created_at": datetime.now(timezone.utc).isoformat(timespec="seconds"), + "models": [ + { + "name": model_name, + "repository": repository, + "revision": revision, + "model_bin_size": model_bin_size, + "model_bin_sha256": model_bin_sha256, + "path": str(model_paths[model_name]), + } + for model_name, repository, revision, model_bin_size, model_bin_sha256, _label in MODEL_SPECS + ], + } + target = cache_dir / "offline-transcription-manifest.json" + temporary = target.with_suffix(".json.tmp") + temporary.write_text(json.dumps(payload, ensure_ascii=False, indent=2), encoding="utf-8") + temporary.replace(target) + return target + + +def extract_smoke_audio(source: Path, target: Path, seconds: int) -> None: + command = [ + "ffmpeg", + "-y", + "-nostdin", + "-loglevel", + "error", + "-i", + str(source), + "-t", + str(max(1, seconds)), + "-ar", + "16000", + "-ac", + "1", + str(target), + ] + result = subprocess.run( + command, + capture_output=True, + text=True, + encoding="utf-8", + errors="replace", + check=False, + ) + if result.returncode != 0: + raise RuntimeError(result.stderr.strip() or "FFmpeg 冒烟音频截取失败") + + +def verify_gpu_model(primary_path: Path, audio_path: Path | None, seconds: int) -> None: + dll_status = configure_windows_cuda_dll_directories() + if not dll_status["cublas_ready"] or not dll_status["cudnn_ready"]: + raise RuntimeError( + "Windows GPU 运行库未就绪:" + f"cuBLAS={dll_status['cublas_ready']},cuDNN={dll_status['cudnn_ready']}" + ) + + from faster_whisper import WhisperModel + + model = WhisperModel(str(primary_path), device="cuda", compute_type="float16") + print("GPU 主模型加载成功:large-v3 / cuda / float16") + if audio_path is None: + print("未提供 --audio,本次只验证模型加载,未执行真实语音推理。") + return + + if not audio_path.is_file(): + raise RuntimeError(f"真实冒烟音频不存在:{audio_path}") + with TemporaryDirectory(prefix="niuma_local_asr_smoke_") as temp_dir: + clip_path = Path(temp_dir) / "smoke.wav" + extract_smoke_audio(audio_path, clip_path, seconds) + segments, _info = model.transcribe( + str(clip_path), + language="zh", + vad_filter=True, + beam_size=5, + word_timestamps=True, + ) + results = list(segments) + if not results: + raise RuntimeError("真实 GPU 冒烟完成,但没有识别到语音") + print(f"真实 GPU 冒烟成功:识别到 {len(results)} 句") + for segment in results[:5]: + print(f"{segment.start:.1f}-{segment.end:.1f}: {segment.text.strip()}") + + +def verify_cpu_fallback_model(fallback_path: Path) -> None: + from faster_whisper import WhisperModel + + WhisperModel(str(fallback_path), device="cpu", compute_type="int8") + print("CPU 兜底模型加载成功:medium / cpu / int8") + + +def main() -> None: + args = parse_args() + cache_dir = Path(args.cache_dir).expanduser().resolve() + cache_dir.mkdir(parents=True, exist_ok=True) + object.__setattr__(settings, "transcription_model_cache_dir", cache_dir) + + model_paths: dict[str, Path] = {} + for model_name, repository, revision, model_bin_size, model_bin_sha256, label in MODEL_SPECS: + model_paths[model_name] = download_model( + model_name, + repository, + revision, + model_bin_size, + model_bin_sha256, + label, + ) + + manifest_path = write_manifest(cache_dir, model_paths) + verify_gpu_model( + model_paths[PRIMARY_TRANSCRIPTION_MODEL], + Path(args.audio).expanduser().resolve() if args.audio else None, + args.seconds, + ) + verify_cpu_fallback_model(model_paths[CPU_FALLBACK_TRANSCRIPTION_MODEL]) + print(f"离线模型清单:{manifest_path}") + print("初始化完成。运行时可断网加载,不会上传音频。") + + +if __name__ == "__main__": + main() diff --git a/tests/test_ai_retry.py b/tests/test_ai_retry.py new file mode 100644 index 0000000..9269975 --- /dev/null +++ b/tests/test_ai_retry.py @@ -0,0 +1,267 @@ +from __future__ import annotations + +import hashlib +import json +import shutil + +import pytest +from fastapi.testclient import TestClient + +from app.db.database import get_connection +from app.main import app +from app.models.settings import AIConfigUpdate +from app.models.task import TaskCreate, TaskStatus +from app.services import job_service, task_service +from app.services.pipeline_checkpoint_service import AUTO_PIPELINE_CHECKPOINT_KIND +from app.services.pipeline_engine import PipelineEngine +from app.services.storage_service import get_artifact_paths +from app.services.task_lifecycle_service import create_task_record + + +client = TestClient(app) + + +@pytest.fixture(autouse=True) +def cleanup_ai_retry_data(): + yield + pattern = "test-ai-retry-%" + with get_connection() as connection: + task_ids = [ + str(row["id"]) + for row in connection.execute( + "SELECT id FROM tasks WHERE id LIKE ?", + (pattern,), + ).fetchall() + ] + connection.execute("DELETE FROM publish_jobs WHERE task_id LIKE ?", (pattern,)) + connection.execute("DELETE FROM output_clip WHERE task_id LIKE ?", (pattern,)) + connection.execute("DELETE FROM clip_candidates WHERE task_id LIKE ?", (pattern,)) + connection.execute("DELETE FROM ai_analysis_windows WHERE task_id LIKE ?", (pattern,)) + connection.execute("DELETE FROM ai_analysis_runs WHERE task_id LIKE ?", (pattern,)) + connection.execute("DELETE FROM workflow_jobs WHERE task_id LIKE ?", (pattern,)) + connection.execute("DELETE FROM tasks WHERE id LIKE ?", (pattern,)) + connection.commit() + for task_id in task_ids: + shutil.rmtree(get_artifact_paths(task_id)["task_dir"], ignore_errors=True) + + +def _create_task(task_id: str) -> None: + create_task_record( + TaskCreate( + task_name=task_id, + source_type="upload", + platform="general", + selection_profile="variety_comedy", + auto_mode=True, + ), + task_id=task_id, + task_dir_name=task_id, + ) + + +def _sha256(value: bytes) -> str: + return hashlib.sha256(value).hexdigest() + + +def _failed_job(task_id: str, transcript: bytes = b"original transcript") -> dict: + paths = get_artifact_paths(task_id) + paths["transcript_path"].parent.mkdir(parents=True, exist_ok=True) + paths["transcript_path"].write_bytes(transcript) + job = job_service.create_job( + task_id, + job_service.JOB_TYPE_AUTO_PIPELINE, + payload={"retry": False, "start_step": None}, + ) + checkpoint = { + "kind": AUTO_PIPELINE_CHECKPOINT_KIND, + "task_id": task_id, + "run_key": "test-run-key", + "start_step": TaskStatus.PREPARING_SOURCE.value, + "current_step": TaskStatus.CLIP_SELECTING.value, + "completed_steps": [ + TaskStatus.PREPARING_SOURCE.value, + TaskStatus.TRANSCRIBING.value, + TaskStatus.AI_ANALYZING.value, + ], + "steps": { + TaskStatus.PREPARING_SOURCE.value: { + "state": "succeeded", + "attempts": 1, + "outputs": {}, + }, + TaskStatus.TRANSCRIBING.value: { + "state": "succeeded", + "attempts": 1, + "outputs": {"transcript": {"sha256": _sha256(transcript)}}, + }, + TaskStatus.AI_ANALYZING.value: { + "state": "succeeded", + "attempts": 1, + "outputs": {}, + }, + TaskStatus.CLIP_SELECTING.value: { + "state": "failed", + "attempts": 1, + "outputs": {}, + "error": "analysis incomplete", + }, + }, + "last_error": "analysis incomplete", + "_ai_analysis_units_v1": { + "namespaces": { + "variety_recall": { + "input_fingerprint": "input-v1", + "units": { + "window_001": { + "status": "completed", + "result_json": "{}", + "result_checksum": _sha256(b"{}"), + }, + "window_004": { + "status": "uncertain", + "error": "Codex CLI timeout; billing uncertain", + }, + }, + } + } + }, + } + with get_connection() as connection: + connection.execute( + """ + UPDATE workflow_jobs + SET status = ?, checkpoint_json = ?, error_message = 'analysis incomplete' + WHERE id = ? + """, + (job_service.JOB_STATUS_FAILED, json.dumps(checkpoint), job["id"]), + ) + connection.execute( + "UPDATE tasks SET status = ? WHERE id = ?", + (TaskStatus.FAILED_CLIP_SELECTING.value, task_id), + ) + connection.commit() + return job_service.get_job(job["id"]) + + +def test_retry_requires_structured_confirmation() -> None: + assert AIConfigUpdate().ai_codex_timeout_seconds == 600 + task_id = "test-ai-retry-confirm" + _create_task(task_id) + old_job = _failed_job(task_id) + + response = client.post(f"/api/tasks/{task_id}/process/auto-retry") + + assert response.status_code == 409 + detail = response.json()["detail"] + assert detail["code"] == "ai_retry_confirmation_required" + assert detail["uncertain_unit_count"] == 1 + assert detail["retry_mode"] == "resume_uncertain" + assert detail["previous_job_id"] == old_job["id"] + + +def test_confirmed_same_input_retries_only_uncertain_unit() -> None: + task_id = "test-ai-retry-resume" + _create_task(task_id) + old_job = _failed_job(task_id) + + response = client.post( + f"/api/tasks/{task_id}/process/auto-retry?confirm_uncertain_ai=true" + ) + + assert response.status_code == 200 + payload = response.json() + assert payload["job_id"] == old_job["id"] + assert payload["retry_mode"] == "resume_uncertain" + stored = job_service.get_job(old_job["id"]) + units = stored["checkpoint_json"]["_ai_analysis_units_v1"]["namespaces"]["variety_recall"]["units"] + assert units["window_001"]["status"] == "completed" + assert units["window_004"]["status"] == "retryable_failed" + assert units["window_004"]["retry_authorized_reason"] == "explicit_user_confirmation" + assert stored["checkpoint_json"]["completed_steps"] == [ + TaskStatus.PREPARING_SOURCE.value, + TaskStatus.TRANSCRIBING.value, + ] + assert stored["checkpoint_json"]["current_step"] == TaskStatus.AI_ANALYZING.value + assert task_service.get_task(task_id, include_video_probe=False)["status"] == TaskStatus.FAILED_AI_ANALYZING.value + + +def test_confirmed_changed_input_creates_fresh_ai_job_and_preserves_old_evidence() -> None: + task_id = "test-ai-retry-fresh" + _create_task(task_id) + old_job = _failed_job(task_id) + old_checkpoint = old_job["checkpoint_json"] + get_artifact_paths(task_id)["transcript_path"].write_bytes(b"simplified transcript") + + response = client.post( + f"/api/tasks/{task_id}/process/auto-retry?confirm_uncertain_ai=true" + ) + + assert response.status_code == 200 + payload = response.json() + assert payload["retry_mode"] == "fresh_ai" + assert payload["job_id"] != old_job["id"] + new_job = job_service.get_job(payload["job_id"]) + assert new_job["payload_json"]["start_step"] == TaskStatus.AI_ANALYZING.value + assert new_job["checkpoint_json"] == {} + preserved = job_service.get_job(old_job["id"]) + assert preserved["status"] == job_service.JOB_STATUS_FAILED + assert preserved["checkpoint_json"] == old_checkpoint + assert task_service.get_task(task_id, include_video_probe=False)["status"] == TaskStatus.FAILED_AI_ANALYZING.value + + +def test_retry_refuses_existing_downstream_records() -> None: + task_id = "test-ai-retry-downstream" + _create_task(task_id) + _failed_job(task_id) + now = task_service._now_iso() + with get_connection() as connection: + connection.execute( + """ + INSERT INTO output_clip ( + id, task_id, output_file_path, output_file_name, status, + created_at, updated_at + ) VALUES ('output-existing', ?, 'clip.mp4', 'clip.mp4', 'completed', ?, ?) + """, + (task_id, now, now), + ) + connection.commit() + + response = client.post(f"/api/tasks/{task_id}/process/auto-retry") + + assert response.status_code == 409 + assert response.json()["detail"]["code"] == "ai_retry_downstream_conflict" + + +def test_incomplete_analysis_fails_in_ai_stage(monkeypatch) -> None: + task_id = "test-ai-retry-incomplete-stage" + _create_task(task_id) + incomplete_meta = { + "analysis_incomplete": True, + "coverage_ratio": 0.9444, + "coverage_percent": 94.44, + } + monkeypatch.setattr( + task_service, + "process_task_ai_analysis", + lambda _task_id: { + "clips": [{"clip_id": "one"}], + "analysis_run_id": "run-one", + "analysis_path": "candidate_clips.json", + "analysis_run": {"analysis_meta": incomplete_meta}, + }, + ) + monkeypatch.setattr( + task_service, + "validate_ai_analysis_meta_for_cut", + lambda meta, _profile: meta, + ) + + result = PipelineEngine().run( + task_id, + retry=True, + start_step=TaskStatus.AI_ANALYZING, + ) + + assert result["failed_step"] == TaskStatus.AI_ANALYZING.value + assert result["failed_status"] == TaskStatus.FAILED_AI_ANALYZING.value + assert task_service.get_task(task_id, include_video_probe=False)["status"] == TaskStatus.FAILED_AI_ANALYZING.value diff --git a/tests/test_offline_transcription.py b/tests/test_offline_transcription.py new file mode 100644 index 0000000..48e3568 --- /dev/null +++ b/tests/test_offline_transcription.py @@ -0,0 +1,584 @@ +from __future__ import annotations + +from contextlib import contextmanager +import hashlib +import io +from pathlib import Path +import sys +from types import SimpleNamespace + +import pytest +from fastapi.testclient import TestClient +from pydantic import ValidationError + +from app.core.config import settings +from app.core.transcription_defaults import ( + CPU_FALLBACK_TRANSCRIPTION_MODEL_REVISION, + PRIMARY_TRANSCRIPTION_MODEL_REVISION, +) +from app.main import app +from app.models.settings import AIConfigUpdate +from app.services import local_transcription_runtime, transcript_service +from app.services.local_transcription_runtime import ( + MODEL_VOCABULARY_FILES, + REQUIRED_MODEL_FILES, + TranscriptionOfflinePolicyError, + get_local_transcription_runtime_status, + model_cache_directory, + model_files_ready, + model_identity, + resolve_local_model_source, +) +from app.services.transcript_service import TranscriptChunk +from app.services.transcript_workflow_service import ( + get_task_transcript_status, + validate_transcription_provider_choice, +) +from app.services.transcription_checkpoint_service import ( + TranscriptionCheckpoint, + fingerprint_file_full, +) +from scripts import setup_local_transcription +from scripts.normalize_transcript_simplified import normalize_transcript_file + + +@contextmanager +def _override_settings(**values): + previous = {name: getattr(settings, name) for name in values} + try: + for name, value in values.items(): + object.__setattr__(settings, name, value) + yield + finally: + for name, value in previous.items(): + object.__setattr__(settings, name, value) + + +def _write_model_files(target: Path) -> None: + target.mkdir(parents=True, exist_ok=True) + for name in REQUIRED_MODEL_FILES: + (target / name).write_bytes(b"test") + (target / MODEL_VOCABULARY_FILES[0]).write_bytes(b"test") + + +def test_default_transcription_configuration_is_fully_offline_local() -> None: + assert settings.transcription_provider == "local" + assert settings.transcription_fallback_provider == "" + assert settings.transcription_offline_only is True + assert settings.transcription_local_files_only is True + assert settings.transcription_model == "large-v3" + assert settings.transcription_model_revision == PRIMARY_TRANSCRIPTION_MODEL_REVISION + assert settings.transcription_device == "cuda" + assert settings.transcription_compute_type == "float16" + assert settings.transcription_chunk_seconds == 120 + assert settings.transcription_chunk_overlap_seconds == 5 + + +def test_local_transcription_converts_segments_and_words_to_simplified(tmp_path) -> None: + class FakeModel: + def transcribe(self, *_args, **_kwargs): + word_items = [ + SimpleNamespace(start=0.0, end=0.5, word="老師", probability=0.9), + SimpleNamespace(start=0.5, end=1.0, word="計程車", probability=0.8), + SimpleNamespace(start=1.0, end=1.5, word=" ABC-123 ", probability=0.7), + ] + segment = SimpleNamespace( + start=0.0, + end=1.5, + text="老師讓計程車載著軟體 ABC-123", + words=word_items, + avg_logprob=-0.1, + ) + return [segment], SimpleNamespace() + + result = transcript_service._transcribe_audio_with_model(FakeModel(), tmp_path / "audio.wav") + + assert result[0].text == "老师让计程车载著软体 ABC-123" + assert [word.text for word in result[0].words] == ["老师", "计程车", "ABC-123"] + assert "计程车" in result[0].text + assert "软体" in result[0].text + + +def test_transcript_file_conversion_backs_up_and_preserves_markdown_structure(tmp_path) -> None: + transcript_path = tmp_path / "transcript.md" + original = ( + "# 逐句時間戳原文\r\n\r\n" + "| 開始 | 結束 | 文本 |\r\n" + "| --- | --- | --- |\r\n" + "| 00:00:01 | 00:00:03 | 老師搭計程車,使用軟體。 |\r\n" + ) + transcript_path.write_bytes(original.encode("utf-8")) + + result = normalize_transcript_file(transcript_path, tmp_path / "backups") + + converted = transcript_path.read_text(encoding="utf-8") + assert "老师搭计程车,使用软体。" in converted + assert len(converted.splitlines()) == len(original.splitlines()) + assert Path(result["backup_path"]).read_bytes() == original.encode("utf-8") + assert result["before_sha256"] != result["after_sha256"] + + +def test_offline_config_rejects_remote_provider_and_network_model_loading() -> None: + with pytest.raises(ValidationError, match="转写方式必须为 local"): + AIConfigUpdate(transcription_provider="volcengine") + with pytest.raises(ValidationError, match="必须只读取本地模型文件"): + AIConfigUpdate(transcription_local_files_only=False) + + +def test_provider_resolver_uses_local_and_blocks_explicit_remote() -> None: + with _override_settings(transcription_provider="local", transcription_offline_only=True): + assert validate_transcription_provider_choice() == "local" + with pytest.raises(TranscriptionOfflinePolicyError, match="完全离线"): + validate_transcription_provider_choice("volcengine") + with pytest.raises(TranscriptionOfflinePolicyError, match="完全离线"): + validate_transcription_provider_choice("remote") + + +@pytest.mark.parametrize( + "endpoint", + ( + "/api/tasks/not-needed/process/transcript?provider=volcengine", + "/api/tasks/not-needed/process/transcript-workflow?provider=volcengine", + ), +) +def test_explicit_volcengine_request_returns_409_before_job_creation(monkeypatch, endpoint) -> None: + called = False + + def unexpected_job(*_args, **_kwargs): + nonlocal called + called = True + raise AssertionError("离线锁开启时不应创建远程转写 Job") + + monkeypatch.setattr("app.routers.tasks.job_service.create_or_get_active_job", unexpected_job) + with _override_settings(transcription_provider="local", transcription_offline_only=True): + response = TestClient(app).post(endpoint) + + assert response.status_code == 409 + assert "完全离线" in response.json()["detail"] + assert called is False + + +def test_remote_transcription_guard_runs_before_http(monkeypatch, tmp_path) -> None: + requested = False + + def unexpected_urlopen(*_args, **_kwargs): + nonlocal requested + requested = True + raise AssertionError("离线锁开启时不应发送 HTTP 请求") + + monkeypatch.setattr(transcript_service, "urlopen", unexpected_urlopen) + audio_path = tmp_path / "audio.mp3" + audio_path.write_bytes(b"audio") + with _override_settings(transcription_offline_only=True): + with pytest.raises(TranscriptionOfflinePolicyError, match="完全离线"): + transcript_service.transcribe_audio_with_provider( + audio_path, + tmp_path, + tmp_path / "progress.json", + "volcengine", + ) + + assert requested is False + + +def test_missing_model_never_downloads_during_task(tmp_path) -> None: + with _override_settings( + transcription_model_cache_dir=tmp_path, + transcription_offline_only=True, + transcription_local_files_only=True, + ): + with pytest.raises(RuntimeError, match="不会联网下载"): + resolve_local_model_source("large-v3", PRIMARY_TRANSCRIPTION_MODEL_REVISION) + + +def test_medium_repository_file_set_with_text_vocabulary_is_ready(tmp_path) -> None: + with _override_settings(transcription_model_cache_dir=tmp_path): + target = model_cache_directory("medium", CPU_FALLBACK_TRANSCRIPTION_MODEL_REVISION) + target.mkdir(parents=True, exist_ok=True) + for name in REQUIRED_MODEL_FILES: + (target / name).write_bytes(b"test") + (target / "vocabulary.txt").write_bytes(b"test") + + assert model_files_ready(target) is True + assert resolve_local_model_source("medium", CPU_FALLBACK_TRANSCRIPTION_MODEL_REVISION) == str(target) + + +def test_official_model_bin_range_download_resumes_partial_file(monkeypatch, tmp_path) -> None: + target = tmp_path / "part.bin" + target.write_bytes(b"abc") + requested_range = "" + + class FakeResponse: + status = 206 + + def __init__(self): + self.stream = io.BytesIO(b"defghij") + + def __enter__(self): + return self + + def __exit__(self, *_args): + return False + + def read(self, size): + return self.stream.read(size) + + def fake_urlopen(request, timeout): + nonlocal requested_range + assert timeout == 120 + requested_range = request.get_header("Range") + return FakeResponse() + + monkeypatch.setattr(setup_local_transcription, "urlopen", fake_urlopen) + + size = setup_local_transcription._download_model_bin_part_with_urllib( + "https://huggingface.co/example/model.bin", + target, + 100, + 109, + ) + + assert size == 10 + assert requested_range == "bytes=103-109" + assert target.read_bytes() == b"abcdefghij" + + +def test_model_bin_part_falls_back_to_urllib_after_curl_failure(monkeypatch, tmp_path) -> None: + target = tmp_path / "part.bin" + calls: list[str] = [] + + monkeypatch.setattr(setup_local_transcription.shutil, "which", lambda _name: "curl.exe") + + def failed_curl(*_args, **_kwargs): + calls.append("curl") + raise RuntimeError("curl failed") + + def successful_urllib(*_args, **_kwargs): + calls.append("urllib") + return 10 + + monkeypatch.setattr(setup_local_transcription, "_download_model_bin_part_with_curl", failed_curl) + monkeypatch.setattr(setup_local_transcription, "_download_model_bin_part_with_urllib", successful_urllib) + + assert setup_local_transcription._download_model_bin_part("https://example.test/model", target, 0, 9) == 10 + assert calls == ["curl", "urllib"] + + +def test_existing_model_bin_requires_matching_official_sha256(tmp_path) -> None: + payload = b"fixed official model" + (tmp_path / "model.bin").write_bytes(payload) + + result = setup_local_transcription._download_verified_model_bin( + tmp_path, + "Systran/test-model", + "a" * 40, + len(payload), + hashlib.sha256(payload).hexdigest(), + "测试模型", + ) + + assert result == tmp_path / "model.bin" + + +def test_runtime_status_reports_cached_model_and_gpu(monkeypatch, tmp_path) -> None: + with _override_settings( + transcription_provider="local", + transcription_model_cache_dir=tmp_path, + transcription_device="cuda", + ): + _write_model_files(model_cache_directory("large-v3", PRIMARY_TRANSCRIPTION_MODEL_REVISION)) + monkeypatch.setattr( + local_transcription_runtime, + "configure_windows_cuda_dll_directories", + lambda: { + "configured": True, + "cublas_ready": True, + "cudnn_ready": True, + "directories": [], + "errors": [], + }, + ) + import ctranslate2 + + monkeypatch.setattr(ctranslate2, "get_cuda_device_count", lambda: 1) + status = get_local_transcription_runtime_status() + + assert status["model_ready"] is True + assert status["gpu_ready"] is True + assert status["ready"] is True + assert status["model_identity"] == "large-v3@edaa852e" + assert status["external_cost_yuan"] == 0 + + +@pytest.mark.skipif(sys.platform != "win32", reason="Windows DLL 搜索路径仅在 Windows 验证") +def test_windows_dll_directories_include_venv_cublas_and_ctranslate2(monkeypatch, tmp_path) -> None: + fake_prefix = tmp_path / ".venv" + cublas_dir = fake_prefix / "Lib" / "site-packages" / "nvidia" / "cublas" / "bin" + ctranslate2_dir = tmp_path / "ctranslate2" + cublas_dir.mkdir(parents=True) + ctranslate2_dir.mkdir() + (cublas_dir / "cublasLt64_12.dll").write_bytes(b"dll") + (cublas_dir / "cublas64_12.dll").write_bytes(b"dll") + (ctranslate2_dir / "cudnn64_9.dll").write_bytes(b"dll") + added: list[str] = [] + + monkeypatch.setattr(local_transcription_runtime.sys, "prefix", str(fake_prefix)) + monkeypatch.setattr( + local_transcription_runtime, + "_package_directory", + lambda _name: ctranslate2_dir, + ) + monkeypatch.setattr( + local_transcription_runtime.os, + "add_dll_directory", + lambda path: added.append(path) or SimpleNamespace(close=lambda: None), + raising=False, + ) + monkeypatch.setattr(local_transcription_runtime, "_DLL_DIRECTORY_HANDLES", []) + monkeypatch.setattr(local_transcription_runtime, "_DLL_DIRECTORY_PATHS", set()) + monkeypatch.setattr(local_transcription_runtime, "_DLL_LIBRARY_HANDLES", []) + monkeypatch.setattr(local_transcription_runtime, "_DLL_LIBRARY_PATHS", set()) + loaded: list[str] = [] + monkeypatch.setattr( + local_transcription_runtime.ctypes, + "WinDLL", + lambda path: loaded.append(path) or SimpleNamespace(), + raising=False, + ) + + first = local_transcription_runtime.configure_windows_cuda_dll_directories() + second = local_transcription_runtime.configure_windows_cuda_dll_directories() + + assert first["cublas_ready"] is True + assert first["cudnn_ready"] is True + assert second["configured"] is True + assert added == [str(cublas_dir), str(ctranslate2_dir)] + assert loaded == [ + str(cublas_dir / "cublasLt64_12.dll"), + str(cublas_dir / "cublas64_12.dll"), + str(ctranslate2_dir / "cudnn64_9.dll"), + ] + + +def test_cuda_load_failure_selects_versioned_cpu_fallback(monkeypatch, tmp_path) -> None: + calls: list[tuple[str, str, str]] = [] + + class FakeWhisperModel: + def __init__(self, model_source, *, device, compute_type, **_kwargs): + calls.append((str(model_source), device, compute_type)) + if device == "cuda": + raise RuntimeError("cublas64_12.dll could not be loaded") + + fake_module = SimpleNamespace(WhisperModel=FakeWhisperModel) + monkeypatch.setitem(sys.modules, "faster_whisper", fake_module) + monkeypatch.setattr( + transcript_service, + "configure_windows_cuda_dll_directories", + lambda: {"configured": True}, + ) + monkeypatch.setattr(transcript_service, "_WHISPER_MODEL", None) + monkeypatch.setattr(transcript_service, "_WHISPER_MODEL_KEY", None) + monkeypatch.setattr(transcript_service, "_EFFECTIVE_TRANSCRIPTION_MODEL_KEY", None) + + with _override_settings( + transcription_model_cache_dir=tmp_path, + transcription_model="large-v3", + transcription_model_revision=PRIMARY_TRANSCRIPTION_MODEL_REVISION, + transcription_device="cuda", + transcription_compute_type="float16", + transcription_cpu_fallback_model="medium", + transcription_local_files_only=True, + ): + _write_model_files(model_cache_directory("large-v3", PRIMARY_TRANSCRIPTION_MODEL_REVISION)) + _write_model_files(model_cache_directory("medium", CPU_FALLBACK_TRANSCRIPTION_MODEL_REVISION)) + transcript_service._set_configured_transcription_runtime("local") + transcript_service._prepare_local_transcription_model_for_run() + + assert [call[1:] for call in calls] == [("cuda", "float16"), ("cpu", "int8")] + assert transcript_service._EFFECTIVE_TRANSCRIPTION_MODEL_KEY == ("medium", "cpu", "int8") + assert transcript_service._ACTIVE_TRANSCRIPTION_MODEL == model_identity( + "medium", + CPU_FALLBACK_TRANSCRIPTION_MODEL_REVISION, + ) + assert transcript_service._transcription_model_revision_label() == CPU_FALLBACK_TRANSCRIPTION_MODEL_REVISION + + +def test_cuda_inference_failure_selects_cpu_before_checkpoint(monkeypatch, tmp_path) -> None: + calls: list[tuple[str, str, str]] = [] + + class FakeWhisperModel: + def __init__(self, model_source, *, device, compute_type, **_kwargs): + self.device = device + calls.append((str(model_source), device, compute_type)) + + def transcribe(self, *_args, **_kwargs): + if self.device == "cuda": + raise RuntimeError("CUDA failed to initialize during inference") + return iter(()), None + + fake_module = SimpleNamespace(WhisperModel=FakeWhisperModel) + monkeypatch.setitem(sys.modules, "faster_whisper", fake_module) + monkeypatch.setattr( + transcript_service, + "configure_windows_cuda_dll_directories", + lambda: {"configured": True}, + ) + monkeypatch.setattr(transcript_service, "_WHISPER_MODEL", None) + monkeypatch.setattr(transcript_service, "_WHISPER_MODEL_KEY", None) + monkeypatch.setattr(transcript_service, "_EFFECTIVE_TRANSCRIPTION_MODEL_KEY", None) + audio_path = tmp_path / "audio.wav" + audio_path.write_bytes(b"audio") + + with _override_settings( + transcription_model_cache_dir=tmp_path, + transcription_model="large-v3", + transcription_model_revision=PRIMARY_TRANSCRIPTION_MODEL_REVISION, + transcription_device="cuda", + transcription_compute_type="float16", + transcription_cpu_fallback_model="medium", + transcription_local_files_only=True, + ): + _write_model_files(model_cache_directory("large-v3", PRIMARY_TRANSCRIPTION_MODEL_REVISION)) + _write_model_files(model_cache_directory("medium", CPU_FALLBACK_TRANSCRIPTION_MODEL_REVISION)) + transcript_service._set_configured_transcription_runtime("local") + transcript_service._prepare_local_transcription_model_for_run(audio_path) + + assert [call[1:] for call in calls] == [("cuda", "float16"), ("cpu", "int8")] + assert transcript_service._EFFECTIVE_TRANSCRIPTION_MODEL_KEY == ("medium", "cpu", "int8") + assert transcript_service._ACTIVE_TRANSCRIPTION_MODEL == model_identity( + "medium", + CPU_FALLBACK_TRANSCRIPTION_MODEL_REVISION, + ) + + +def test_transcript_status_exposes_offline_runtime_fields(monkeypatch, tmp_path) -> None: + transcript_path = tmp_path / "transcript.md" + monkeypatch.setattr( + "app.services.task_service.get_task", + lambda _task_id: { + "status": "pending_processing", + "status_label": "待处理", + "progress": 5, + "error_message": "", + }, + ) + monkeypatch.setattr( + "app.services.transcript_workflow_service.get_artifact_paths", + lambda _task_id: {"transcript_path": transcript_path}, + ) + monkeypatch.setattr( + "app.services.transcript_workflow_service.get_local_transcription_runtime_status", + lambda: { + "offline_only": True, + "model_ready": True, + "gpu_ready": True, + "model_revision": PRIMARY_TRANSCRIPTION_MODEL_REVISION, + }, + ) + + status = get_task_transcript_status("offline-status-test") + + assert status["offline_only"] is True + assert status["model_ready"] is True + assert status["gpu_ready"] is True + assert status["model_revision"] == PRIMARY_TRANSCRIPTION_MODEL_REVISION + + +def test_model_revision_and_full_audio_hash_invalidate_checkpoint(tmp_path) -> None: + from datetime import datetime, timezone + + from app.db.database import get_connection + + task_id = "offline-model-revision-checkpoint" + source = tmp_path / "same-size.wav" + head = b"h" * (1024 * 1024) + tail = b"t" * (1024 * 1024) + source.write_bytes(head + (b"a" * (1024 * 1024)) + tail) + first_hash = fingerprint_file_full(source) + now = datetime.now(timezone.utc).isoformat(timespec="seconds") + chunks = [TranscriptChunk(1, 0, 120)] + with get_connection() as connection: + connection.execute( + """ + INSERT INTO tasks ( + id, task_name, task_dir_name, source_type, platform, selection_profile, + status, progress, is_deleted, created_at, updated_at + ) VALUES (?, 'offline checkpoint', ?, 'upload', 'general', 'long_live_talk', + 'pending_processing', 0, 0, ?, ?) + """, + (task_id, task_id, now, now), + ) + connection.commit() + try: + first = TranscriptionCheckpoint( + task_id=task_id, + source_path=source, + provider="local", + model="large-v3@edaa852e", + device="cuda", + compute_type="float16", + chunk_seconds=120, + overlap_seconds=5, + ) + first.ensure_run(chunks) + revision_changed = TranscriptionCheckpoint( + task_id=task_id, + source_path=source, + provider="local", + model="large-v3@12345678", + device="cuda", + compute_type="float16", + chunk_seconds=120, + overlap_seconds=5, + ) + revision_changed.ensure_run(chunks) + + normalization_changed = TranscriptionCheckpoint( + task_id=task_id, + source_path=source, + provider="local", + model="large-v3@edaa852e|text=opencc-t2s-v1", + device="cuda", + compute_type="float16", + chunk_seconds=120, + overlap_seconds=5, + ) + normalization_changed.ensure_run(chunks) + + source.write_bytes(head + (b"b" * (1024 * 1024)) + tail) + second_hash = fingerprint_file_full(source) + content_changed = TranscriptionCheckpoint( + task_id=task_id, + source_path=source, + provider="local", + model="large-v3@edaa852e", + device="cuda", + compute_type="float16", + chunk_seconds=120, + overlap_seconds=5, + ) + content_changed.ensure_run(chunks) + + assert first_hash != second_hash + assert revision_changed.run_id != first.run_id + assert normalization_changed.run_id != first.run_id + assert content_changed.run_id != first.run_id + with get_connection() as connection: + models = { + row["model"] + for row in connection.execute( + "SELECT model FROM transcription_runs WHERE task_id = ?", + (task_id,), + ).fetchall() + } + assert { + "large-v3@edaa852e", + "large-v3@12345678", + "large-v3@edaa852e|text=opencc-t2s-v1", + } <= models + finally: + with get_connection() as connection: + connection.execute("DELETE FROM transcription_chunks WHERE task_id = ?", (task_id,)) + connection.execute("DELETE FROM transcription_runs WHERE task_id = ?", (task_id,)) + connection.execute("DELETE FROM tasks WHERE id = ?", (task_id,)) + connection.commit()