Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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,可直接复制给客户端作为接入指令。
Expand Down
45 changes: 44 additions & 1 deletion frontend/components/SettingsDialog.vue
Original file line number Diff line number Diff line change
Expand Up @@ -318,6 +318,26 @@
<span class="shrink-0 rounded-full bg-[var(--app-surface-muted)] px-2 py-1 text-[10px] text-[var(--app-text-secondary)]">当前:{{ voiceModelText }}</span>
</div>

<div class="mt-3 flex flex-col gap-2 sm:flex-row sm:items-center sm:justify-between">
<div class="min-w-0">
<label for="voice-model-download-source" class="text-[12px] font-medium text-[var(--app-text-primary)]">模型下载源</label>
<p id="voice-model-download-source-hint" class="mt-0.5 text-[11px] leading-relaxed text-[var(--app-text-muted)]">仅影响新发起的模型下载。连接超时可切换后重试。</p>
</div>
<select
id="voice-model-download-source"
:value="voiceDownloadSource"
:disabled="voiceStatusLoading || voiceDownloadSourceBusy"
:aria-busy="voiceDownloadSourceBusy"
aria-describedby="voice-model-download-source-hint"
class="voice-setting-focus w-full rounded-[6px] border border-[var(--app-border)] bg-[var(--app-surface-bg)] px-2.5 py-1.5 text-[12px] text-[var(--app-text-primary)] disabled:cursor-not-allowed disabled:opacity-50 sm:w-auto"
@change="setVoiceDownloadSource"
>
<option value="huggingface">官方源(Hugging Face)</option>
<option value="hf-mirror">国内镜像(HF-Mirror)</option>
</select>
</div>
<ErrorNotice v-if="voiceDownloadSourceError" :message="voiceDownloadSourceError" compact manual class="mt-1.5 text-[11px] text-[var(--danger-color)]" />

<div v-if="voiceStatusLoading" class="mt-3 grid gap-2 sm:grid-cols-2" aria-label="正在读取模型列表">
<div v-for="index in 4" :key="index" class="h-[134px] rounded-[9px] bg-[var(--app-surface-muted)]" />
</div>
Expand Down Expand Up @@ -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)"
>
Expand Down Expand Up @@ -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: '' })
Expand Down Expand Up @@ -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'
Expand Down Expand Up @@ -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' }
Expand Down Expand Up @@ -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
Expand Down
1 change: 1 addition & 0 deletions frontend/composables/useApi.js
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
26 changes: 26 additions & 0 deletions frontend/tests/voice-download-source.test.js
Original file line number Diff line number Diff line change
@@ -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
}))
}
})
13 changes: 9 additions & 4 deletions src/wechat_decrypt_tool/routers/chat_media.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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(
Expand Down
22 changes: 22 additions & 0 deletions src/wechat_decrypt_tool/runtime_settings.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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:
Expand Down
37 changes: 37 additions & 0 deletions src/wechat_decrypt_tool/voice_transcription.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
*,
Expand All @@ -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,
Expand Down Expand Up @@ -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."""

Expand Down
Loading
Loading