From 40aa3c07264252eb837926cdd029f5254debb151 Mon Sep 17 00:00:00 2001 From: xiaosheng <73678111+xiaoshengbao@users.noreply.github.com> Date: Fri, 2 Oct 2026 16:03:17 +0800 Subject: [PATCH] fix(voice): add domestic model download source switch --- README.md | 4 + frontend/components/SettingsDialog.vue | 45 +++++- frontend/composables/useApi.js | 1 + frontend/tests/voice-download-source.test.js | 26 ++++ src/wechat_decrypt_tool/routers/chat_media.py | 13 +- src/wechat_decrypt_tool/runtime_settings.py | 22 +++ .../voice_transcription.py | 37 +++++ tests/test_voice_model_download_source.py | 146 ++++++++++++++++++ 8 files changed, 289 insertions(+), 5 deletions(-) create mode 100644 frontend/tests/voice-download-source.test.js create mode 100644 tests/test_voice_model_download_source.py diff --git a/README.md b/README.md index ce25c002..03474788 100644 --- a/README.md +++ b/README.md @@ -243,6 +243,10 @@ npm run dev 可在「设置 → AI 服务 → 本地检索」按账号开启可选的语义检索,使用 Hugging Face 固定版本模型,支持 CPU 与 NVIDIA GPU 自动回退。使用方法、下载来源和兼容性实测见 [本地语义检索说明](docs/local-semantic-search.md)。 +## 语音模型下载源 + +在「设置 → 语音转文字 → 模型下载源」可切换官方源(Hugging Face)与国内镜像([HF-Mirror](https://hf-mirror.com))。默认使用官方源,选择会自动保存并用于后续从模型列表发起的下载;已开始的下载保持原源。连接超时后可切换来源,再点击模型的下载按钮重试。镜像站在部分网络下会重定向到官方站,应用会遵循这一永久跳转;镜像本身的可用性仍取决于所在网络。 + ## MCP 服务 设置页中的“AI 接入提示词”会包含 endpoint 和 Bearer token,可直接复制给客户端作为接入指令。 diff --git a/frontend/components/SettingsDialog.vue b/frontend/components/SettingsDialog.vue index 0240f11d..4737ade8 100644 --- a/frontend/components/SettingsDialog.vue +++ b/frontend/components/SettingsDialog.vue @@ -318,6 +318,26 @@ 当前:{{ voiceModelText }} +
+
+ +

仅影响新发起的模型下载。连接超时可切换后重试。

+
+ +
+ +
@@ -383,7 +403,7 @@ v-if="!model.downloaded && !isVoiceModelDeletePending(model.id)" type="button" class="voice-setting-focus whitespace-nowrap rounded-[5px] bg-[var(--app-accent)] px-2 py-1 text-[10px] font-medium text-white transition hover:bg-[var(--app-accent-hover)] disabled:cursor-not-allowed disabled:opacity-40" - :disabled="!model.downloadable || isVoiceModelDownloading(model) || isVoiceModelActionBusy(model.id)" + :disabled="voiceDownloadSourceBusy || !model.downloadable || isVoiceModelDownloading(model) || isVoiceModelActionBusy(model.id)" :title="model.downloadable ? `下载 ${model.name}` : (model.reason || '当前无法下载')" @click="startVoiceModelDownload(model)" > @@ -935,6 +955,9 @@ const voiceDeviceSource = ref('default') const voiceActiveDevice = ref('') const voiceModel = ref('zipformer-small-ctc-int8') const voiceModels = ref([]) +const voiceDownloadSource = ref('huggingface') +const voiceDownloadSourceBusy = ref(false) +const voiceDownloadSourceError = ref('') const voiceSupportedDevices = ref(['cpu', 'cuda']) const voiceModelSource = ref('default') const voiceModelAction = ref({ id: '', type: '' }) @@ -1240,6 +1263,7 @@ const canDeleteVoiceModel = (model) => !!model?.id && ( const applyVoiceTranscriptionStatus = (status) => { if (!status || typeof status !== 'object') return + voiceDownloadSource.value = status.modelDownloadSource === 'hf-mirror' ? 'hf-mirror' : 'huggingface' const requestedDevice = String(status.requestedDevice || status.device || 'cpu').trim().toLowerCase() voiceDevicePreference.value = requestedDevice === 'cuda' ? 'cuda' : 'cpu' voiceDeviceSource.value = String(status.deviceSource || 'default').trim() || 'default' @@ -1436,6 +1460,7 @@ function scheduleVoiceModelDownloadPolling() { } const startVoiceModelDownload = async (model) => { + if (voiceDownloadSourceBusy.value) return if (!model?.id || !model.downloadable || isVoiceModelDownloading(model) || isVoiceModelActionBusy(model.id)) return const generation = voiceModelDownloadGeneration(model.id) voiceModelAction.value = { id: model.id, type: 'download' } @@ -1582,6 +1607,24 @@ const refreshVoiceTranscriptionStatus = async () => { } } +const setVoiceDownloadSource = async (event) => { + const select = event.target + const next = select.value + if (voiceDownloadSourceBusy.value || next === voiceDownloadSource.value) return + voiceDownloadSourceBusy.value = true + voiceDownloadSourceError.value = '' + try { + const resp = await api.setVoiceTranscriptionSettings({ download_source: next }) + applyVoiceTranscriptionStatus(resp?.configuration || resp) + } catch (error) { + voiceDownloadSourceError.value = error?.message || '切换模型下载源失败,请重试' + } finally { + // 保存失败时恢复实际选项,避免界面与后端使用不同的下载源。 + select.value = voiceDownloadSource.value + voiceDownloadSourceBusy.value = false + } +} + const setVoiceDevice = async (device) => { const next = String(device || '').trim().toLowerCase() if (!['cpu', 'cuda'].includes(next) || voiceDeviceBusy.value || voiceDeviceLocked.value) return diff --git a/frontend/composables/useApi.js b/frontend/composables/useApi.js index 53a454e8..67595746 100644 --- a/frontend/composables/useApi.js +++ b/frontend/composables/useApi.js @@ -567,6 +567,7 @@ export const useApi = () => { const body = {} if (data.device != null) body.device = String(data.device || '').trim().toLowerCase() if (data.model != null) body.model = String(data.model || '').trim() + if (data.download_source != null) body.download_source = String(data.download_source || '').trim() return await request('/chat/media/voice/transcription/settings', { method: 'PUT', body diff --git a/frontend/tests/voice-download-source.test.js b/frontend/tests/voice-download-source.test.js new file mode 100644 index 00000000..eae79e92 --- /dev/null +++ b/frontend/tests/voice-download-source.test.js @@ -0,0 +1,26 @@ +import { afterEach, expect, it, vi } from 'vitest' +import { useApi } from '~/composables/useApi' + +vi.mock('~/stores/chatAccounts', () => ({ + useChatAccountsStore: () => ({ applySourceResponse: vi.fn() }) +})) + +afterEach(() => vi.unstubAllGlobals()) + +it('下载源切换通过现有设置接口发送,并保留模型与设备设置请求', async () => { + const fetch = vi.fn(async () => ({ status: 'success' })) + vi.stubGlobal('useApiBase', () => '/api') + vi.stubGlobal('$fetch', fetch) + const api = useApi() + for (const data of [ + { download_source: 'hf-mirror' }, + { download_source: 'huggingface' }, + { model: 'turbo' }, + { device: 'cpu' } + ]) { + await api.setVoiceTranscriptionSettings(data) + expect(fetch).toHaveBeenLastCalledWith('/chat/media/voice/transcription/settings', expect.objectContaining({ + method: 'PUT', body: data + })) + } +}) diff --git a/src/wechat_decrypt_tool/routers/chat_media.py b/src/wechat_decrypt_tool/routers/chat_media.py index 64448a40..cc3473ae 100644 --- a/src/wechat_decrypt_tool/routers/chat_media.py +++ b/src/wechat_decrypt_tool/routers/chat_media.py @@ -12,7 +12,7 @@ import time import re from pathlib import Path -from typing import Any, Optional +from typing import Any, Literal, Optional from urllib.parse import urlparse import requests @@ -89,6 +89,7 @@ load_voice_data, set_voice_transcription_device, set_voice_transcription_model, + set_voice_model_download_source, ) from ..native_voice_transcription import ( NativeVoiceTriggerError, @@ -126,6 +127,7 @@ class VoiceTranscriptionCacheLookupRequest(BaseModel): class VoiceTranscriptionSettingsRequest(BaseModel): device: Optional[str] = Field(None, description="推理设备:cpu 或 cuda") model: Optional[str] = Field(None, description="本地语音模型") + download_source: Optional[Literal["huggingface", "hf-mirror"]] = Field(None, description="语音模型下载源") class VoiceTranscriptionBatchRequest(BaseModel): @@ -3747,19 +3749,22 @@ async def get_chat_voice_transcription_status(): return await asyncio.to_thread(get_voice_transcription_service().status) -@router.put("/api/chat/media/voice/transcription/settings", summary="设置本地语音模型或推理设备") +@router.put("/api/chat/media/voice/transcription/settings", summary="设置本地语音模型、推理设备或下载源") async def set_chat_voice_transcription_settings(req: VoiceTranscriptionSettingsRequest, request: Request): _require_local_voice_mutation(request) device = str(req.device or "").strip() model = str(req.model or "").strip() - if int(bool(device)) + int(bool(model)) != 1: - raise HTTPException(status_code=400, detail="每次只能修改 device 或 model 中的一项。") + download_source = req.download_source + if sum(bool(value) for value in (device, model, download_source)) != 1: + raise HTTPException(status_code=400, detail="每次只能修改 device、model 或 download_source 中的一项。") try: configuration = None if model: configuration = await asyncio.to_thread(set_voice_transcription_model, model) if device: configuration = await asyncio.to_thread(set_voice_transcription_device, device) + if download_source: + configuration = await asyncio.to_thread(set_voice_model_download_source, download_source) except VoiceTranscriptionError as exc: status_code = 409 if exc.code in {"device_locked", "model_locked", "model_busy"} else 400 raise HTTPException( diff --git a/src/wechat_decrypt_tool/runtime_settings.py b/src/wechat_decrypt_tool/runtime_settings.py index 28e1b3c1..91d40efa 100644 --- a/src/wechat_decrypt_tool/runtime_settings.py +++ b/src/wechat_decrypt_tool/runtime_settings.py @@ -14,6 +14,11 @@ MCP_TOKEN_KEY = "mcp_token" VOICE_TRANSCRIPTION_DEVICE_KEY = "voice_transcription_device" VOICE_TRANSCRIPTION_MODEL_KEY = "voice_transcription_model" +VOICE_MODEL_DOWNLOAD_SOURCE_KEY = "voice_model_download_source" +VOICE_MODEL_DOWNLOAD_ENDPOINTS = { + "huggingface": "https://huggingface.co", + "hf-mirror": "https://hf-mirror.com", +} ENV_PORT_KEY = "WECHAT_TOOL_PORT" ENV_HOST_KEY = "WECHAT_TOOL_HOST" ENV_ALLOW_REMOTE_CALLS_KEY = "WECHAT_TOOL_ALLOW_REMOTE_CALLS" @@ -329,6 +334,23 @@ def read_effective_voice_transcription_model(default: str = "zipformer-small-ctc return _normalize_voice_transcription_model(default) or "zipformer-small-ctc-int8", "default" +def read_voice_model_download_source() -> str: + """读取语音模型下载源;旧配置或无效配置使用官方源。""" + source = _read_runtime_settings().get(VOICE_MODEL_DOWNLOAD_SOURCE_KEY) + return source if isinstance(source, str) and source in VOICE_MODEL_DOWNLOAD_ENDPOINTS else "huggingface" + + +def write_voice_model_download_source(source: str) -> None: + """仅持久化预设下载源,不接受任意网址。""" + if source not in VOICE_MODEL_DOWNLOAD_ENDPOINTS: + raise ValueError("不支持的语音模型下载源") + data = _read_runtime_settings() + data[VOICE_MODEL_DOWNLOAD_SOURCE_KEY] = source + _write_runtime_settings(data) + if read_voice_model_download_source() != source: + raise OSError("语音模型下载源保存失败") + + def ensure_mcp_token() -> tuple[str, str]: token, source = read_effective_mcp_token() if token: diff --git a/src/wechat_decrypt_tool/voice_transcription.py b/src/wechat_decrypt_tool/voice_transcription.py index 020dc3a0..16698adb 100644 --- a/src/wechat_decrypt_tool/voice_transcription.py +++ b/src/wechat_decrypt_tool/voice_transcription.py @@ -61,12 +61,15 @@ def _prepare_whisper_cuda_libraries() -> None: logger.debug("可选 CUDA 库目录不可用,继续使用系统运行库。", exc_info=True) from .runtime_settings import ( + VOICE_MODEL_DOWNLOAD_ENDPOINTS, VOICE_TRANSCRIPTION_DEVICE_CPU, VOICE_TRANSCRIPTION_DEVICE_CUDA, read_effective_voice_transcription_device, read_effective_voice_transcription_model, + read_voice_model_download_source, write_voice_transcription_device_setting, write_voice_transcription_model_setting, + write_voice_model_download_source, ) from .app_paths import get_data_dir, get_output_databases_dir, get_output_dir from .asr_models import ( @@ -1571,6 +1574,7 @@ def status(self) -> dict[str, Any]: "supportedDevices": spec["devices"] if spec else ["cpu", "cuda"], "modelSettingSource": self.config.model_source, "models": get_voice_model_catalog(selected_model=self.config.model), + "modelDownloadSource": read_voice_model_download_source(), "language": self.config.language, "device": self.config.device, "computeType": self.config.compute_type, @@ -2324,6 +2328,24 @@ def _write_cache( return True +def _voice_model_download_endpoint(source: str, repo_id: str) -> str: + """兼容镜像站在部分网络下永久跳转至官方站的行为。""" + endpoint = VOICE_MODEL_DOWNLOAD_ENDPOINTS[source] + if source == "hf-mirror": + path = f"/api/models/{repo_id}" + official = VOICE_MODEL_DOWNLOAD_ENDPOINTS["huggingface"] + try: + # SDK 不跟随跨站的 HEAD 跳转;仅接受镜像明确返回的同仓库官方地址。 + response = httpx.head(endpoint + path, follow_redirects=False, timeout=5.0) + if response.status_code in {301, 308} and response.headers.get("location") == official + path: + logger.info("[voice-model-download] mirror redirected to official host repo=%s", repo_id) + return official + except httpx.HTTPError: + # 探测失败仍交给原下载流程处理,不自动更换用户选择的源。 + pass + return endpoint + + def _download_voice_model_snapshot( model_id: str, *, @@ -2339,6 +2361,10 @@ def _download_voice_model_snapshot( repo_id = VOICE_MODEL_REPOSITORIES[model_id] common = { "local_dir": str(output_dir), + # 单次任务固定下载源,元数据与文件请求保持一致,不修改全局 HF_ENDPOINT。 + "endpoint": _voice_model_download_endpoint(read_voice_model_download_source(), repo_id), + # 内置模型均为公开仓库,避免将本机 Hugging Face 凭据发送给镜像站。 + "token": False, "allow_patterns": list(VOICE_MODEL_DOWNLOAD_ALLOW_PATTERNS), # A cancelled worker must not wait for other executor workers to drain. "max_workers": 1, @@ -3498,6 +3524,17 @@ def _reset_voice_transcription_service() -> VoiceTranscriptionService: _VOICE_TRANSCRIPTION_SERVICE_CONDITION.notify_all() +def set_voice_model_download_source(source: str) -> dict[str, Any]: + """切换后续下载的来源,无需重建推理服务或中断当前下载。""" + if source not in VOICE_MODEL_DOWNLOAD_ENDPOINTS: + raise VoiceTranscriptionError("invalid_download_source", "不支持该模型下载源。") + try: + write_voice_model_download_source(source) + except OSError as exc: + raise VoiceTranscriptionError("download_source_save_failed", "下载源保存失败,请重试。") from exc + return get_voice_transcription_service().status() + + def set_voice_transcription_device(device: str) -> dict[str, Any]: """Persist a user preference and make it effective for subsequent requests.""" diff --git a/tests/test_voice_model_download_source.py b/tests/test_voice_model_download_source.py new file mode 100644 index 00000000..83659202 --- /dev/null +++ b/tests/test_voice_model_download_source.py @@ -0,0 +1,146 @@ +"""语音模型下载源的持久化、接口边界与单次下载一致性。""" +import os +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest +import httpx +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from wechat_decrypt_tool import runtime_settings as settings +from wechat_decrypt_tool import voice_transcription as voice +from wechat_decrypt_tool.routers import chat_media + + +@pytest.fixture(autouse=True) +def isolated_settings(tmp_path, monkeypatch): + monkeypatch.setenv("WECHAT_TOOL_DATA_DIR", str(tmp_path)) + monkeypatch.setattr(voice.httpx, "head", Mock(return_value=httpx.Response(200))) + + +def test_download_source_persists_without_changing_model_or_device(): + settings.write_voice_transcription_model_setting("turbo") + settings.write_voice_transcription_device_setting("cuda") + assert settings.read_voice_model_download_source() == "huggingface" + for source in ("hf-mirror", "huggingface"): + settings.write_voice_model_download_source(source) + assert settings.read_voice_model_download_source() == source + assert settings._read_runtime_settings()[settings.VOICE_MODEL_DOWNLOAD_SOURCE_KEY] == source + assert settings.read_voice_transcription_model_setting() == "turbo" + assert settings.read_voice_transcription_device_setting() == "cuda" + + +@pytest.mark.parametrize("value", ["unknown", "https://example.com", [], None]) +def test_invalid_saved_source_falls_back_to_official(value): + settings._write_runtime_settings({settings.VOICE_MODEL_DOWNLOAD_SOURCE_KEY: value}) + assert settings.read_voice_model_download_source() == "huggingface" + + +@pytest.mark.parametrize("model", list(voice.VOICE_MODEL_REPOSITORIES)) +@pytest.mark.parametrize("source", ["huggingface", "hf-mirror"]) +def test_all_models_use_selected_source_for_metadata_and_files(model, source, tmp_path, monkeypatch): + from huggingface_hub import constants + + settings.write_voice_model_download_source(source) + monkeypatch.setenv("HF_ENDPOINT", "https://example.com") + original_endpoint = constants.ENDPOINT + calls = [] + + def snapshot(repo, **kwargs): + calls.append((repo, kwargs)) + if kwargs.get("dry_run"): + # 下载期间切换设置,当前文件请求仍应使用原先选定的源。 + settings.write_voice_model_download_source("huggingface" if source == "hf-mirror" else "hf-mirror") + return [SimpleNamespace(file_size=100)] + return kwargs["local_dir"] + + monkeypatch.setattr("huggingface_hub.snapshot_download", snapshot) + result = voice._download_voice_model_snapshot(model, output_dir=tmp_path, progress_callback=lambda **_: None) + assert result == tmp_path + assert len(calls) == 2 + for repo, kwargs in calls: + assert repo == voice.VOICE_MODEL_REPOSITORIES[model] + assert kwargs["endpoint"] == settings.VOICE_MODEL_DOWNLOAD_ENDPOINTS[source] + assert kwargs["token"] is False + if model in voice.ASR_MODEL_SPECS: + assert kwargs["revision"] == voice.ASR_MODEL_SPECS[model]["revision"] + assert set(kwargs["allow_patterns"]) == set(voice.ASR_MODEL_SPECS[model]["files"]) + assert constants.ENDPOINT == original_endpoint + assert os.environ["HF_ENDPOINT"] == "https://example.com" + + calls.clear() + voice._download_voice_model_snapshot(model, output_dir=tmp_path, progress_callback=lambda **_: None) + assert calls[0][1]["endpoint"] != settings.VOICE_MODEL_DOWNLOAD_ENDPOINTS[source] + + +def test_failed_metadata_probe_keeps_mirror_for_actual_download(tmp_path, monkeypatch): + settings.write_voice_model_download_source("hf-mirror") + snapshot = Mock(side_effect=[TimeoutError("metadata timeout"), str(tmp_path)]) + monkeypatch.setattr("huggingface_hub.snapshot_download", snapshot) + assert voice._download_voice_model_snapshot("turbo", output_dir=tmp_path, progress_callback=lambda **_: None) == tmp_path + assert all(call.kwargs["endpoint"] == "https://hf-mirror.com" for call in snapshot.call_args_list) + + +@pytest.mark.parametrize("status, location, expected", [ + (308, "https://huggingface.co/api/models/pkufool/zipformer-small", "https://huggingface.co"), + (301, "https://huggingface.co/api/models/pkufool/zipformer-small", "https://huggingface.co"), + (302, "https://huggingface.co/api/models/pkufool/zipformer-small", "https://hf-mirror.com"), + (308, "https://example.com/api/models/pkufool/zipformer-small", "https://hf-mirror.com"), + (308, "https://huggingface.co/api/models/other/repo", "https://hf-mirror.com"), +]) +def test_mirror_permanent_redirect_is_limited_to_same_official_repo(status, location, expected, monkeypatch): + head = Mock(return_value=httpx.Response(status, headers={"location": location})) + monkeypatch.setattr(voice.httpx, "head", head) + assert voice._voice_model_download_endpoint("hf-mirror", "pkufool/zipformer-small") == expected + head.assert_called_once_with("https://hf-mirror.com/api/models/pkufool/zipformer-small", follow_redirects=False, timeout=5.0) + + +def test_failed_probe_does_not_switch_sources(monkeypatch): + head = Mock(side_effect=httpx.ConnectTimeout("probe timeout")) + monkeypatch.setattr(voice.httpx, "head", head) + assert voice._voice_model_download_endpoint("hf-mirror", "pkufool/zipformer-small") == "https://hf-mirror.com" + head.reset_mock() + assert voice._voice_model_download_endpoint("huggingface", "pkufool/zipformer-small") == "https://huggingface.co" + head.assert_not_called() + + +@pytest.fixture +def client(monkeypatch): + service = SimpleNamespace(status=lambda: {"modelDownloadSource": settings.read_voice_model_download_source()}) + monkeypatch.setattr(voice, "get_voice_transcription_service", lambda: service) + reset = Mock(side_effect=AssertionError("切换下载源不应重建推理服务")) + monkeypatch.setattr(voice, "_reset_voice_transcription_service", reset) + app = FastAPI() + app.include_router(chat_media.router) + with TestClient(app, base_url="http://127.0.0.1:10392", client=("127.0.0.1", 50000)) as client: + yield client + + +def test_settings_api_saves_and_returns_download_source(client): + for source in ("hf-mirror", "huggingface"): + response = client.put("/api/chat/media/voice/transcription/settings", json={"download_source": source}) + assert response.status_code == 200 + assert response.json()["configuration"]["modelDownloadSource"] == source + assert settings.read_voice_model_download_source() == source + + +@pytest.mark.parametrize("body, status", [ + ({"download_source": "https://example.com"}, 422), + ({"download_source": ""}, 422), + ({"download_source": "hf-mirror", "device": "cpu"}, 400), + ({"download_source": "hf-mirror", "model": "turbo"}, 400), + ({}, 400), +]) +def test_settings_api_rejects_invalid_or_combined_source_updates(client, body, status): + response = client.put("/api/chat/media/voice/transcription/settings", json=body) + assert response.status_code == status + assert settings.read_voice_model_download_source() == "huggingface" + + +def test_save_failure_does_not_report_success(client, monkeypatch): + monkeypatch.setattr(settings, "_write_runtime_settings", lambda _: None) + response = client.put("/api/chat/media/voice/transcription/settings", json={"download_source": "hf-mirror"}) + assert response.status_code == 400 + assert response.json()["detail"]["code"] == "download_source_save_failed" + assert settings.read_voice_model_download_source() == "huggingface"