Files
zWorkFlow/tests/test_voice_transcription_cache.py
btc-z 66eddaff0e feat: transcribe_voice 新增 OpenAI Whisper API 后端 (#66)
默认 local,零行为变化。opt-in 双因素:transcription_backend=openai
且 openai_api_key 都齐才生效;任一缺失静默回退 local + stderr 一行警告。
首次进入云路径会 stderr 警告"语音将上传至 OpenAI 服务器"。

新增 config.json 字段:
- transcription_backend: "local" (默认) | "openai"
- local_whisper_model: "base" (替换 mcp_server.py 里硬编码 DEFAULT_WHISPER_MODEL)
- openai_api_key: "" (默认空;openai 包为 optional,按需 pip install)

关键技术选择:
- _transcribe(wav, backend) 单一 if/else 分发,不引入插件/工厂层
  (Rule of Three —— 只有一个云后端时不值得抽象)
- 文件 > 25MB 在 OpenAI() 实例化之前提前拒绝,避免无谓上传
- 错误分类清晰: 缺 key / 缺 openai 包 / 401 / 429 / APIError 各自的提示
- PR #58 缓存 schema 自然扩展: 条目加 backend 字段,命中需 backend+model_size 都匹配
- 旧条目缺 backend 字段视为 "local",向前兼容 PR #58 已落盘的所有数据
- transcribe_chat.py 批量 CLI 与 MCP 工具共享同一份配置,保持一致

新增 2 个测试 (tests/test_openai_backend.py),只覆盖回归风险最高的两条:
- 文件 > 25MB 必须在 SDK 实例化前拒绝(隐私契约的防线)
- backend 不匹配的旧条目不命中(避免切后端时返回错后端结果)

其余路径要么琐碎(默认值读取)、要么坏掉时声音很大(SDK 错误、ImportError),
要么已被 PR #58 现有测试隐式覆盖(缺 backend 字段的旧条目),不再单独写测试。

顺手把 README 里 PR #53 漏掉的 voice 三件套(get_voice_messages /
decode_voice / transcribe_voice)补进 MCP 工具表,并新增"⚠️ 语音转录隐私"
章节说清数据流向、成本(约 \$0.006/分钟)、25MB 上限、回退行为。

Closes ylytdeng/wechat-decrypt#59
2026-05-01 13:56:32 +08:00

275 lines
12 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

import json
import os
import tempfile
import threading
import unittest
from unittest.mock import patch
import mcp_server
class _CacheIsolationMixin:
"""所有测试共享:隔离 module-level 缓存状态 + 指向 tempdir 的 cache 文件。"""
def setUp(self):
self._saved_cache = mcp_server._voice_transcription_cache
self._saved_path = mcp_server.VOICE_TRANSCRIPTION_CACHE_FILE
self._saved_warned = mcp_server._voice_transcription_save_warned
mcp_server._voice_transcription_cache = None
mcp_server._voice_transcription_save_warned = False
self._tmp = tempfile.TemporaryDirectory()
self.addCleanup(self._tmp.cleanup)
mcp_server.VOICE_TRANSCRIPTION_CACHE_FILE = os.path.join(
self._tmp.name, "voice_transcriptions.json"
)
def tearDown(self):
mcp_server._voice_transcription_cache = self._saved_cache
mcp_server.VOICE_TRANSCRIPTION_CACHE_FILE = self._saved_path
mcp_server._voice_transcription_save_warned = self._saved_warned
class VoiceTranscriptionCachePersistenceTests(_CacheIsolationMixin, unittest.TestCase):
"""_load_voice_transcription_cache / _save_voice_transcription_cache 的持久化行为。"""
def test_load_missing_file_returns_empty_dict(self):
self.assertEqual(mcp_server._load_voice_transcription_cache(), {})
def test_save_and_reload_roundtrip(self):
cache = mcp_server._load_voice_transcription_cache()
cache["wxid_foo:42"] = {
"text": "你好",
"language": "zh",
"create_time": 1700000000,
"model_size": "base",
}
mcp_server._save_voice_transcription_cache()
# 强制下一次 load 从磁盘读
mcp_server._voice_transcription_cache = None
reloaded = mcp_server._load_voice_transcription_cache()
self.assertEqual(reloaded["wxid_foo:42"]["text"], "你好")
self.assertEqual(reloaded["wxid_foo:42"]["language"], "zh")
def test_corrupt_file_returns_empty_dict(self):
with open(mcp_server.VOICE_TRANSCRIPTION_CACHE_FILE, "w", encoding="utf-8") as f:
f.write("{{ not valid json")
self.assertEqual(mcp_server._load_voice_transcription_cache(), {})
def test_non_dict_payload_returns_empty_dict(self):
with open(mcp_server.VOICE_TRANSCRIPTION_CACHE_FILE, "w", encoding="utf-8") as f:
json.dump(["not", "a", "dict"], f)
self.assertEqual(mcp_server._load_voice_transcription_cache(), {})
def test_utf8_preserved_on_disk(self):
# ensure_ascii=False 必须生效,否则中文会被转义成 \uXXXX
cache = mcp_server._load_voice_transcription_cache()
cache["wxid_bar:1"] = {"text": "中文测试", "language": "zh"}
mcp_server._save_voice_transcription_cache()
with open(mcp_server.VOICE_TRANSCRIPTION_CACHE_FILE, "rb") as f:
raw = f.read()
self.assertIn("中文测试".encode("utf-8"), raw)
def test_save_without_prior_load_persists_empty_dict(self):
# 从未 load 过就直接 save应落盘一个空 dict而不是静默丢弃。
mcp_server._voice_transcription_cache = None
mcp_server._save_voice_transcription_cache()
self.assertTrue(os.path.exists(mcp_server.VOICE_TRANSCRIPTION_CACHE_FILE))
with open(mcp_server.VOICE_TRANSCRIPTION_CACHE_FILE, encoding="utf-8") as f:
self.assertEqual(json.load(f), {})
class VoiceTranscriptionCacheAtomicityTests(_CacheIsolationMixin, unittest.TestCase):
"""原子写 + crash-during-save 行为。"""
def test_write_is_atomic_via_rename(self):
# 先写入一份已有缓存
cache = mcp_server._load_voice_transcription_cache()
cache["wxid_x:1"] = {"text": "initial", "language": "zh", "model_size": "base"}
mcp_server._save_voice_transcription_cache()
# 模拟:写 .tmp 正常但 os.replace 阶段失败
original_replace = os.replace
def flaky_replace(src, dst):
raise OSError("disk full during rename")
cache["wxid_x:1"] = {"text": "MUTATED", "language": "zh", "model_size": "base"}
with patch.object(os, "replace", side_effect=flaky_replace):
mcp_server._save_voice_transcription_cache() # 不应抛
# 磁盘上应仍然是 initial不是 MUTATED也不是损坏的半截文件
with open(mcp_server.VOICE_TRANSCRIPTION_CACHE_FILE, encoding="utf-8") as f:
disk = json.load(f)
self.assertEqual(disk["wxid_x:1"]["text"], "initial")
# .tmp 应该被清理,避免污染目录
tmp_path = mcp_server.VOICE_TRANSCRIPTION_CACHE_FILE + ".tmp"
# 注patch 生效期间 os.replace 失败finally 里会尝试 unlink
_ = original_replace # 防 lint 警告
self.assertFalse(os.path.exists(tmp_path))
def test_early_save_error_preserves_existing_file(self):
# json.dump 在 .tmp 上抛异常时(模拟磁盘满 / 权限问题),主文件应保持原样;
# 注意此测试不是"写到一半中断"而是"写前就失败"的场景。
cache = mcp_server._load_voice_transcription_cache()
cache["wxid_y:1"] = {"text": "survives", "language": "zh", "model_size": "base"}
mcp_server._save_voice_transcription_cache()
cache["wxid_y:1"] = {"text": "DO NOT SEE", "language": "zh", "model_size": "base"}
def boom(*args, **kwargs):
raise OSError("disk full")
with patch.object(mcp_server.json, "dump", side_effect=boom):
mcp_server._save_voice_transcription_cache() # 静默降级,不抛
# 主文件没被破坏:仍然可 json.load 出原先内容
mcp_server._voice_transcription_cache = None
reloaded = mcp_server._load_voice_transcription_cache()
self.assertEqual(reloaded["wxid_y:1"]["text"], "survives")
class VoiceTranscriptionCacheConcurrencyTests(_CacheIsolationMixin, unittest.TestCase):
"""多线程下的 load/save 行为。"""
def test_concurrent_load_returns_same_dict_instance(self):
# 16 个线程同时触发首次 load应当只实际化一份 dictlock 生效)
barrier = threading.Barrier(16)
results = []
results_lock = threading.Lock()
def worker():
barrier.wait()
d = mcp_server._load_voice_transcription_cache()
with results_lock:
results.append(id(d))
threads = [threading.Thread(target=worker) for _ in range(16)]
for t in threads:
t.start()
for t in threads:
t.join()
self.assertEqual(len(set(results)), 1, "并发 load 应返回同一个 dict 对象")
def test_concurrent_save_does_not_corrupt(self):
# 多个线程同时 save磁盘上最终文件必须是合法 JSON原子写 + lock 保障)
cache = mcp_server._load_voice_transcription_cache()
for i in range(100):
cache[f"wxid_z:{i}"] = {
"text": f"msg-{i}",
"language": "zh",
"model_size": "base",
}
def worker():
mcp_server._save_voice_transcription_cache()
threads = [threading.Thread(target=worker) for _ in range(8)]
for t in threads:
t.start()
for t in threads:
t.join()
with open(mcp_server.VOICE_TRANSCRIPTION_CACHE_FILE, encoding="utf-8") as f:
disk = json.load(f) # 必须能解析
self.assertEqual(len(disk), 100)
class TranscribeVoiceCacheHitTests(_CacheIsolationMixin, unittest.TestCase):
"""transcribe_voice 的缓存命中 / 失效路径。"""
def _seed(self, key, entry):
cache = mcp_server._load_voice_transcription_cache()
cache[key] = entry
mcp_server._save_voice_transcription_cache()
def test_cache_hit_skips_fetch_and_transcribe(self):
key = mcp_server._voice_transcription_cache_key("wxid_test", 7)
self._seed(key, {
"text": "缓存命中文本",
"language": "zh",
"create_time": 1700000000,
"model_size": mcp_server.LOCAL_WHISPER_MODEL,
})
with patch.object(mcp_server, "resolve_username", return_value="wxid_test") as mock_resolve, \
patch.object(mcp_server, "_fetch_voice_row") as mock_fetch, \
patch.object(mcp_server, "_silk_to_wav") as mock_silk, \
patch.object(mcp_server, "_get_whisper_model") as mock_model:
result = mcp_server.transcribe_voice("test_contact", 7)
mock_resolve.assert_called_once_with("test_contact")
mock_fetch.assert_not_called()
mock_silk.assert_not_called()
mock_model.assert_not_called()
self.assertIn("缓存命中文本", result)
self.assertIn("(zh)", result)
def test_cache_hit_uses_placeholder_when_create_time_missing(self):
# 旧条目若没有 create_time 字段,不应崩溃
key = mcp_server._voice_transcription_cache_key("wxid_test", 8)
self._seed(key, {
"text": "历史条目",
"language": "zh",
"model_size": mcp_server.LOCAL_WHISPER_MODEL,
})
with patch.object(mcp_server, "resolve_username", return_value="wxid_test"), \
patch.object(mcp_server, "_fetch_voice_row") as mock_fetch:
result = mcp_server.transcribe_voice("test_contact", 8)
mock_fetch.assert_not_called()
self.assertIn("历史条目", result)
def test_cache_hit_returns_empty_text_without_retranscribing(self):
# Whisper 返回空也要缓存;再次调用应直接返回空,不进入 miss 路径
key = mcp_server._voice_transcription_cache_key("wxid_test", 9)
self._seed(key, {
"text": "",
"language": "zh",
"create_time": 1700000000,
"model_size": mcp_server.LOCAL_WHISPER_MODEL,
})
with patch.object(mcp_server, "resolve_username", return_value="wxid_test"), \
patch.object(mcp_server, "_fetch_voice_row") as mock_fetch:
result = mcp_server.transcribe_voice("test_contact", 9)
mock_fetch.assert_not_called()
self.assertIn("(zh)", result)
def test_model_mismatch_is_treated_as_miss(self):
# 缓存条目的 model_size 和当前 LOCAL_WHISPER_MODEL 不一致时,
# 不应命中;进入 miss 路径(这里无 whisper 依赖,应落到"缺少依赖"分支)。
key = mcp_server._voice_transcription_cache_key("wxid_test", 10)
self._seed(key, {
"text": "旧模型结果",
"language": "zh",
"create_time": 1700000000,
"model_size": "OUTDATED_MODEL",
})
with patch.object(mcp_server, "resolve_username", return_value="wxid_test"), \
patch.dict("sys.modules", {"whisper": None}):
# whisper=None 时 `import whisper` 触发 ImportError
result = mcp_server.transcribe_voice("test_contact", 10)
# 走了 miss 路径 → 返回缺依赖提示,而不是返回旧缓存文本
self.assertNotIn("旧模型结果", result)
self.assertIn("缺少依赖", result)
def test_cache_key_handles_colon_in_username(self):
# 若上游未来的 resolve_username 放出带 ':' 的 username也不会和其他条目冲突
key_a = mcp_server._voice_transcription_cache_key("wxid:foo", 1)
key_b = mcp_server._voice_transcription_cache_key("wxid", 1)
self.assertNotEqual(key_a, key_b)
if __name__ == "__main__":
unittest.main()