diff --git a/.github/workflows/chat-image-quality.yml b/.github/workflows/chat-image-quality.yml
index aefecf55..64c3e036 100644
--- a/.github/workflows/chat-image-quality.yml
+++ b/.github/workflows/chat-image-quality.yml
@@ -5,12 +5,16 @@ on:
paths:
- '.github/workflows/chat-image-quality.yml'
- 'src/wechat_decrypt_tool/media_helpers.py'
+ - 'src/wechat_decrypt_tool/key_store.py'
- 'src/wechat_decrypt_tool/chat_export_service.py'
- 'src/wechat_decrypt_tool/routers/media.py'
- 'src/wechat_decrypt_tool/routers/chat*.py'
- 'tests/test_chat_image_local_quality.py'
- 'tests/test_chat_media_image_cache_upgrade.py'
- 'tests/test_media_decrypt_stream_cancel.py'
+ - 'tests/test_media_xor_decode.py'
+ - 'tests/test_media_wxgf_ffmpeg_fallback.py'
+ - 'tests/test_key_file_permissions.py'
- 'frontend/lib/chat/chat-history.js'
- 'frontend/lib/chat/message-normalizer.js'
- 'frontend/tests/chat-image-quality.test.js'
@@ -54,6 +58,9 @@ jobs:
tests/test_image_key_resolver.py
tests/test_media_emoticon_catalog.py
tests/test_media_emoji_download_stream.py
+ tests/test_media_xor_decode.py
+ tests/test_media_wxgf_ffmpeg_fallback.py
+ tests/test_key_file_permissions.py
- name: Install frontend dependencies
working-directory: frontend
run: npm ci
diff --git a/docs/chat-ai.md b/docs/chat-ai.md
index 330c48f6..c8db9e10 100644
--- a/docs/chat-ai.md
+++ b/docs/chat-ai.md
@@ -22,6 +22,8 @@
选择 Ollama 或 LM Studio 前,请先启动其本地服务并准备好模型。默认地址分别为 `http://127.0.0.1:11434/v1` 和 `http://127.0.0.1:1234/v1`,未启用鉴权时密钥可留空;选择后会自动获取模型。启用鉴权时需填写密钥,端口不同则修改地址。连接失败时检查服务、地址、端口和鉴权设置,再点击「从上游获取」。也可勾选「手动输入(备用)」填写真实模型名称后保存;保存配置不代表连接已验证。
+服务运行在局域网内的另一台机器上时,把地址改为 `http://<局域网 IP>:端口/v1`。明文 HTTP 仅支持本机(`localhost`、`127.0.0.1`、`[::1]`)和局域网 IP(`10.x`、`172.16-31.x`、`192.168.x`、IPv6 `fc`/`fd` 开头),`nas.local` 这类主机名及其他地址必须使用 HTTPS。明文 HTTP 不加密也不校验对方身份,同一网络上的其他设备可能看到聊天内容和密钥,或冒充该服务返回内容;保存后自动任务和关注提醒也会持续向该地址发送内容,请仅在可信的局域网中使用。
+
## 内容与运行边界
- 支持文本、引用、卡片标题、已有语音转写、JPEG/PNG/WebP/GIF(首帧)、PDF(含扫描页)、DOCX、XLSX、PPTX、TXT、MD、CSV。
diff --git a/frontend/components/AiSettings.vue b/frontend/components/AiSettings.vue
index 9659748a..b9c1462f 100644
--- a/frontend/components/AiSettings.vue
+++ b/frontend/components/AiSettings.vue
@@ -62,6 +62,7 @@
请先启动本地服务并准备好模型;未启用鉴权时,API 密钥可留空。
+ 此地址使用明文 HTTP,聊天内容和密钥不加密传输,请仅在可信的局域网中使用。
@@ -247,10 +248,15 @@ const blank = () => ({ provider: 'deepseek', name: 'DeepSeek', protocol: 'openai
const form = reactive(blank())
// 切换预设后明确清空凭据,不让后端复用原配置的密钥。
const credentialsReset = ref(false)
-const isLocalService = computed(() => {
- if (!['ollama', 'lmstudio'].includes(form.provider)) return false
- try { return ['localhost', '127.0.0.1', '[::1]'].includes(new URL(form.base_url).hostname) } catch { return false }
+// 后端只对本机和局域网 IP 放行明文 HTTP,这里按协议判断,不重复网段规则。
+const serviceAddress = computed(() => {
+ try {
+ const url = new URL(form.base_url)
+ return { http: url.protocol === 'http:', loopback: ['localhost', '127.0.0.1', '[::1]'].includes(url.hostname) }
+ } catch { return {} }
})
+const isLocalService = computed(() => ['ollama', 'lmstudio'].includes(form.provider) && Boolean(serviceAddress.value.http || serviceAddress.value.loopback))
+const isLanHttp = computed(() => Boolean(serviceAddress.value.http && !serviceAddress.value.loopback))
const manualModel = ref(false), modelDetails = ref([]), modelError = ref(''), modelsLoading = ref(false)
const manualMetadata = ref(null)
const selectedMetadata = computed(() => {
diff --git a/frontend/components/GlobalExportDialog.vue b/frontend/components/GlobalExportDialog.vue
index 0306df61..b839373b 100644
--- a/frontend/components/GlobalExportDialog.vue
+++ b/frontend/components/GlobalExportDialog.vue
@@ -566,10 +566,7 @@ const startExport = async () => {
} else {
task.value.message = '正在保存到浏览器目录...'
task.value.progress = 98
- const zipPath = String(finalJob.zipPath || '').trim()
- const query = new URLSearchParams()
- query.set('path', zipPath)
- const downloadUrl = `${apiBase}/account/archive_export/download?${query.toString()}`
+ const downloadUrl = `${apiBase}/account/archive_export/${encodeURIComponent(currentExportId.value)}/download`
const downloadResponse = await fetch(downloadUrl)
if (!downloadResponse.ok) {
throw new Error(`下载导出文件失败(${downloadResponse.status})。`)
diff --git a/frontend/components/chat/ChatExportDialog.vue b/frontend/components/chat/ChatExportDialog.vue
index 957ae3c1..722584c4 100644
--- a/frontend/components/chat/ChatExportDialog.vue
+++ b/frontend/components/chat/ChatExportDialog.vue
@@ -451,7 +451,7 @@
@@ -497,6 +497,16 @@
重新探测缺失媒体
+
+
+
+
+
+
+ 本次未导出位置消息
+ 该目录的基线不含“位置”类型,已按基线的消息类型更新;需要时请重置增量基线或改用新目录。
+
+
diff --git a/frontend/composables/chat/useChatExport.js b/frontend/composables/chat/useChatExport.js
index b0461c41..646373e4 100644
--- a/frontend/composables/chat/useChatExport.js
+++ b/frontend/composables/chat/useChatExport.js
@@ -21,6 +21,7 @@ export const useChatExport = ({ api, apiBase, contacts, selectedAccount, selecte
{ value: 'emoji', label: '表情' },
{ value: 'video', label: '视频' },
{ value: 'voice', label: '语音' },
+ { value: 'location', label: '位置' },
{ value: 'chatHistory', label: '聊天记录' },
{ value: 'transfer', label: '转账' },
{ value: 'redPacket', label: '红包' },
diff --git a/frontend/tests/ai-presets.test.js b/frontend/tests/ai-presets.test.js
index a6af35a3..60dd051b 100644
--- a/frontend/tests/ai-presets.test.js
+++ b/frontend/tests/ai-presets.test.js
@@ -119,6 +119,29 @@ describe('AI 服务预设', () => {
expect(request.mock.calls.find(([path]) => path === '/profiles')[1].body).toMatchObject({ provider, model: 'local-model', api_key: '' })
})
+ it.each(['ollama', 'lmstudio'])('%s 改填局域网 HTTP 地址后仍免密钥自动获取,并提示明文传输', async provider => {
+ await open(); await choose(provider)
+ const dialogText = () => wrapper.find('[role=dialog]').text()
+ const fetches = () => request.mock.calls.filter(([path]) => path === '/models')
+ const fillAddress = async value => {
+ await wrapper.find('input[type=url]').setValue(value)
+ await wrapper.find('input[type=url]').trigger('blur'); await flushPromises()
+ }
+ expect(fetches()).toHaveLength(1)
+ expect(dialogText()).not.toContain('明文 HTTP')
+ await fillAddress('http://192.168.1.5:11434/v1')
+ expect(fetches()).toHaveLength(2)
+ expect(fetches().at(-1)[1].body).toMatchObject({ base_url: 'http://192.168.1.5:11434/v1', api_key: '' })
+ expect(dialogText()).toContain('API 密钥可留空')
+ expect(dialogText()).toContain('明文 HTTP')
+ // 同一预设改填 HTTPS 远程地址时,仍需先填写密钥。
+ await fillAddress('https://ollama.example.com/v1')
+ expect(wrapper.find('input[type=url]').element.value).toBe('https://ollama.example.com/v1')
+ expect(fetches()).toHaveLength(2)
+ expect(dialogText()).not.toContain('API 密钥可留空')
+ expect(dialogText()).not.toContain('明文 HTTP')
+ })
+
it('已有云端配置切换到本地时明确清空保存的密钥', async () => {
const original = request.getMockImplementation()
const profile = { ...presets[0], id: 'saved', has_key: true, model: 'cloud-model', vision: false }
diff --git a/frontend/tests/chat-export-model-options.test.js b/frontend/tests/chat-export-model-options.test.js
index 93ad570c..1ff7f199 100644
--- a/frontend/tests/chat-export-model-options.test.js
+++ b/frontend/tests/chat-export-model-options.test.js
@@ -16,7 +16,7 @@ function setup({ types = ['link'], privacy = false, transcribe = false, availabl
selectedAccount: ref('test-account'), selectedContact: ref(null), privacyMode: ref(privacy) }))
state.exportSelectedUsernames.value = ['friend']
state.exportFolderHandle.value = {}
- state.exportMessageTypes.value = types
+ if (types) state.exportMessageTypes.value = types
state.exportTranscribeVoice.value = transcribe
return { state, api }
}
@@ -59,4 +59,15 @@ describe('导出媒体与语音模型选项', () => {
expect(api.createChatExport).not.toHaveBeenCalled()
expect(state.exportError.value).toBe('请先下载模型')
})
+
+ it('位置消息默认勾选并随导出请求提交', async () => {
+ const { state, api } = setup({ types: null })
+ expect(state.exportMessageTypeOptions).toContainEqual({ value: 'location', label: '位置' })
+ expect(state.areAllExportMessageTypesSelected.value).toBe(true)
+ await state.startChatExport()
+ expect(api.createChatExport).toHaveBeenCalledWith(expect.objectContaining({
+ message_types: state.exportMessageTypeOptions.map(item => item.value),
+ }))
+ expect(api.createChatExport.mock.calls[0][0].message_types).toContain('location')
+ })
})
diff --git a/main.py b/main.py
index 63344a7c..a0a0807b 100644
--- a/main.py
+++ b/main.py
@@ -10,14 +10,25 @@
import multiprocessing
import os
+import sys
from pathlib import Path
# Keep standalone/frozen launches safe when scanner code uses multiprocessing.
if __name__ == "__main__":
multiprocessing.freeze_support()
+# Source launches must run this checkout's src/, not the copy that
+# `uv sync --no-editable` froze into site-packages. Keep this above every
+# project import.
+SRC_DIR = Path(__file__).resolve().parent / "src"
+if SRC_DIR.is_dir():
+ if str(SRC_DIR) in sys.path:
+ sys.path.remove(str(SRC_DIR))
+ sys.path.insert(0, str(SRC_DIR))
+
import uvicorn
+import wechat_decrypt_tool
from wechat_decrypt_tool.desktop_parent_watchdog import (
start_desktop_parent_watchdog_from_env,
)
@@ -55,6 +66,12 @@ def main():
else:
print("监听地址来源: 默认值")
print(f"监听地址: {host}")
+ code_source = str(Path(wechat_decrypt_tool.__file__).resolve())
+ try:
+ print(f"代码来源: {code_source}")
+ except UnicodeEncodeError:
+ # 标准输出的编码表示不了仓库路径时,退回转义形式,不让这行诊断信息中断启动。
+ print(f"代码来源: {ascii(code_source)}")
print(f"API文档: http://{access_host}:{port}/docs")
print(f"健康检查: http://{access_host}:{port}/api/health")
if lan_access_host != access_host:
@@ -62,7 +79,6 @@ def main():
print("按 Ctrl+C 停止服务")
print("=" * 60)
- repo_root = Path(__file__).resolve().parent
enable_reload = os.environ.get("WECHAT_TOOL_RELOAD", "0") == "1"
# 启动API服务
@@ -71,7 +87,7 @@ def main():
host=host,
port=port,
reload=enable_reload,
- reload_dirs=[str(repo_root / "src")] if enable_reload else None,
+ reload_dirs=[str(SRC_DIR)] if enable_reload else None,
reload_excludes=[
"output/*",
"output/**",
diff --git a/src/wechat_decrypt_tool/ai/providers.py b/src/wechat_decrypt_tool/ai/providers.py
index 9fa59fd2..efd7e078 100644
--- a/src/wechat_decrypt_tool/ai/providers.py
+++ b/src/wechat_decrypt_tool/ai/providers.py
@@ -3,6 +3,7 @@
import logging
import asyncio
+import ipaddress
import json
import re
import time
@@ -74,12 +75,28 @@ def public_profile(profile):
return {k: v for k, v in profile.items() if k != "api_key"} | {"has_key": bool(profile.get("api_key"))}
+# 明文 HTTP 的局域网例外只列 RFC 1918 私有网段和 IPv6 唯一本地地址。不用 is_private:
+# 它还包含 0.0.0.0/8、链路本地(含 169.254.169.254)和文档保留网段。
+LAN_NETWORKS = tuple(ipaddress.ip_network(value) for value in ("10.0.0.0/8", "172.16.0.0/12", "192.168.0.0/16", "fc00::/7"))
+
+
+def is_lan_address(host):
+ """只认 IP 字面量:不解析域名,不展开 IPv4 映射等嵌入地址,也不接受带 zone id 的写法。"""
+ if not isinstance(host, str) or "%" in host:
+ return False
+ try:
+ address = ipaddress.ip_address(host)
+ except ValueError:
+ return False
+ return any(address in network for network in LAN_NETWORKS)
+
+
def validate_url(value):
url = urlparse(value)
if url.scheme not in {"https", "http"} or not url.hostname or url.username or url.password or url.query or url.fragment:
raise ValueError("请输入有效的 HTTP(S) 服务地址,不要在地址中包含密钥")
- if url.scheme == "http" and url.hostname not in {"localhost", "127.0.0.1", "::1"}:
- raise ValueError("远程模型服务必须使用 HTTPS")
+ if url.scheme == "http" and url.hostname not in {"localhost", "127.0.0.1", "::1"} and not is_lan_address(url.hostname):
+ raise ValueError("明文 HTTP 仅支持本机(localhost、127.0.0.1、[::1])和局域网 IP(10.x、172.16-31.x、192.168.x、IPv6 fc/fd 开头),域名等其他地址必须使用 HTTPS")
def model_base_url(value):
diff --git a/src/wechat_decrypt_tool/chat_export_service.py b/src/wechat_decrypt_tool/chat_export_service.py
index 810ab22a..f768a0f4 100644
--- a/src/wechat_decrypt_tool/chat_export_service.py
+++ b/src/wechat_decrypt_tool/chat_export_service.py
@@ -2172,6 +2172,13 @@ def _run_job(self, job: ExportJob, account_dir: Path) -> None:
missing_files=list(opts.get("missingFiles") or []),
reset_baseline=bool(opts.get("resetBaseline")),
)
+ if folder_context.location_type_skipped:
+ # 探测、渲染都要和基线用同一份类型清单,否则已导出的历史会被误判为有差异。
+ want_types = set(folder_context.config.get("messageTypes") or [])
+ job.options["messageTypes"] = [
+ value for value in message_types_raw if _normalize_render_type_key(value) in want_types
+ ]
+ _safe_trace(trace, "incremental_location_type_skipped", messageTypes=sorted(want_types))
preferred_missing_owner_keys = {
incremental_conversation_key(salt=folder_context.salt, username=username)
for username in target_usernames
@@ -3655,6 +3662,10 @@ def esc_attr(v: Any) -> str:
warning_parts: list[str] = []
if folder_context.reset_baseline:
warning_parts.append("已重置基线并完整重建本次选择的会话。")
+ if folder_context.location_type_skipped:
+ warning_parts.append(
+ "该增量目录的基线不含“位置”类型,本次仍按基线的消息类型更新,未导出位置消息;需要时请重置增量基线或改用新目录。"
+ )
recovered_files = int(job.incremental.get("filesRecovered") or 0)
if recovered_files:
warning_parts.append(f"已补回 {recovered_files} 个缺失或异常的受管理文件。")
diff --git a/src/wechat_decrypt_tool/chat_incremental_export.py b/src/wechat_decrypt_tool/chat_incremental_export.py
index 9a433422..f670ef72 100644
--- a/src/wechat_decrypt_tool/chat_incremental_export.py
+++ b/src/wechat_decrypt_tool/chat_incremental_export.py
@@ -135,6 +135,17 @@ def config_fingerprint(config: dict[str, Any]) -> str:
return hashlib.sha256(_canonical_json(config)).hexdigest()
+def _config_without_location(config: dict[str, Any]) -> Optional[dict[str, Any]]:
+ """返回去掉“位置”后的配置;本次没有勾选位置,或只勾选了位置时返回 None。"""
+
+ requested = list(config.get("messageTypes") or [])
+ kept = [value for value in requested if value != "location"]
+ # 空清单在基线里表示“不过滤、导出全部类型”,不能把“只勾选位置”当成它。
+ if not kept or len(kept) == len(requested):
+ return None
+ return {**config, "messageTypes": kept}
+
+
def conversation_key(*, salt: str, username: str) -> str:
payload = f"{str(salt or '')}\0{str(username or '')}".encode("utf-8", errors="replace")
return hashlib.sha256(payload).hexdigest()
@@ -275,6 +286,7 @@ class ChatFolderContext:
unresolved_media_conversations: list[dict[str, Any]] = field(default_factory=list)
unresolved_missing_owner_keys: set[str] = field(default_factory=set)
metadata_changed: bool = False
+ location_type_skipped: bool = False
@property
def export_runtime_id(self) -> str:
@@ -361,6 +373,7 @@ def prepare_folder_context(
_validate_baseline_paths(old_state)
desired_hash = config_fingerprint(config)
+ location_type_skipped = False
if owned:
baseline_account = str(old_state.get("account") or "")
baseline_account_fingerprint = str(old_state.get("accountFingerprint") or "")
@@ -371,11 +384,19 @@ def prepare_folder_context(
)
if not account_matches:
raise ChatIncrementalError("incremental_account_mismatch", "该增量目录属于其他微信账号,请选择新目录。")
- if str(old_state.get("configFingerprint") or "") != desired_hash and not reset_baseline:
- raise ChatIncrementalError(
- "incremental_config_mismatch",
- "导出格式或筛选配置与该增量目录不一致,请选择新目录或重置后完整重建。",
- )
+ baseline_hash = str(old_state.get("configFingerprint") or "")
+ if baseline_hash != desired_hash and not reset_baseline:
+ # “位置”是导出面板后来补上的类型,而且默认勾选。基线只差这一项时沿用基线的类型清单,
+ # 已导出的历史与后续追加保持同一口径;需要位置消息时重置基线即可。
+ baseline_config = _config_without_location(config)
+ if baseline_config is None or config_fingerprint(baseline_config) != baseline_hash:
+ raise ChatIncrementalError(
+ "incremental_config_mismatch",
+ "导出格式或筛选配置与该增量目录不一致,请选择新目录或重置后完整重建。",
+ )
+ config = baseline_config
+ desired_hash = baseline_hash
+ location_type_skipped = True
if reset_baseline:
if old_state and not owned:
@@ -421,6 +442,7 @@ def prepare_folder_context(
salt=salt,
missing_files=missing,
reset_baseline=bool(reset_baseline),
+ location_type_skipped=location_type_skipped,
)
@@ -848,6 +870,7 @@ def materialize_folder_archive(
"filesReused": max(0, len(current_files) - len(staged_entries)),
"filesRemoved": len(stale),
"filesRecovered": recovered_count,
+ "locationTypeSkipped": bool(context.location_type_skipped),
}
job.repair_candidates = list(context.repair_candidates)
job.unresolved_media = {
diff --git a/src/wechat_decrypt_tool/chat_search_index.py b/src/wechat_decrypt_tool/chat_search_index.py
index 3e42f3fc..4acc1d2d 100644
--- a/src/wechat_decrypt_tool/chat_search_index.py
+++ b/src/wechat_decrypt_tool/chat_search_index.py
@@ -285,24 +285,28 @@ def start_chat_search_index_build(account_dir: Path, *, rebuild: bool = False, s
now = int(time.time())
with _BUILD_LOCK:
st = _BUILD_STATE.get(key)
- if st and st.get("status") == "building":
- return get_chat_search_index_status(account_dir, source=source_norm)
- _BUILD_STATE[key] = {
- "status": "building",
- "rebuild": bool(rebuild),
- "source": source_norm,
- "startedAt": now,
- "finishedAt": None,
- "indexedMessages": 0,
- "fetchedMessages": 0,
- "fetchCalls": 0,
- "totalConversations": 0,
- "completedConversations": 0,
- "messagesPerSec": 0,
- "currentDb": "",
- "currentConversation": "",
- "error": "",
- }
+ already_building = bool(st and st.get("status") == "building")
+ if not already_building:
+ _BUILD_STATE[key] = {
+ "status": "building",
+ "rebuild": bool(rebuild),
+ "source": source_norm,
+ "startedAt": now,
+ "finishedAt": None,
+ "indexedMessages": 0,
+ "fetchedMessages": 0,
+ "fetchCalls": 0,
+ "totalConversations": 0,
+ "completedConversations": 0,
+ "messagesPerSec": 0,
+ "currentDb": "",
+ "currentConversation": "",
+ "error": "",
+ }
+
+ # get_chat_search_index_status 会再次获取 _BUILD_LOCK(不可重入),只能在锁外调用。
+ if already_building:
+ return get_chat_search_index_status(account_dir, source=source_norm)
t = threading.Thread(
target=_build_worker,
diff --git a/src/wechat_decrypt_tool/key_store.py b/src/wechat_decrypt_tool/key_store.py
index 3a206c9f..dc9fd39c 100644
--- a/src/wechat_decrypt_tool/key_store.py
+++ b/src/wechat_decrypt_tool/key_store.py
@@ -1,5 +1,6 @@
import datetime
import json
+import os
import threading
from pathlib import Path
from typing import Any, Iterable, Optional
@@ -119,6 +120,14 @@ def _same_complete_image_key_pair(
def _atomic_write_json(path: Path, payload: Any) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
tmp = path.with_suffix(path.suffix + ".tmp")
+ # 密钥文件只允许属主读写:先把临时文件建成 0600(残留的旧临时文件也收紧)再写入内容,
+ # 替换后目标路径沿用这个权限,密钥不会以更宽的权限落盘。Windows 没有这套权限位,不处理。
+ if os.name == "posix":
+ try:
+ tmp.touch(mode=0o600, exist_ok=True)
+ os.chmod(tmp, 0o600)
+ except Exception:
+ pass
tmp.write_text(
json.dumps(payload, ensure_ascii=False, indent=2),
encoding="utf-8",
diff --git a/src/wechat_decrypt_tool/mcp/tools.py b/src/wechat_decrypt_tool/mcp/tools.py
index 10605cde..7bcd51cc 100644
--- a/src/wechat_decrypt_tool/mcp/tools.py
+++ b/src/wechat_decrypt_tool/mcp/tools.py
@@ -289,23 +289,27 @@ def _list_contacts(args: dict[str, Any], ctx: McpToolContext) -> dict[str, Any]:
return {**result, "contacts": _clip_deep(page), "offset": offset, "limit": limit, "hasMore": offset + limit < len(contacts)}
+def _match_confidence(query: str, item: dict[str, Any]) -> int:
+ q_lower = query.lower()
+ hay = " ".join(str(item.get(k) or "") for k in ("username", "remark", "nickname", "name", "displayName", "alias")).lower()
+ score = 0
+ if q_lower in hay:
+ score += 60
+ if hay.startswith(q_lower):
+ score += 20
+ if str(item.get("username") or "") == query:
+ score += 30
+ return min(100, score or 20)
+
+
def _resolve_contact(args: dict[str, Any], ctx: McpToolContext) -> dict[str, Any]:
query = _str(args, "query")
if not query:
raise ValueError("query is required.")
base = _list_contacts({**args, "keyword": query, "limit": _int(args, "limit", 10, minimum=1, maximum=50)}, ctx)
candidates = []
- q_lower = query.lower()
for item in list(base.get("contacts") or []):
- hay = " ".join(str(item.get(k) or "") for k in ("username", "remark", "nickname", "name", "displayName", "alias")).lower()
- score = 0
- if query in hay:
- score += 60
- if hay.startswith(q_lower):
- score += 20
- if str(item.get("username") or "") == query:
- score += 30
- candidates.append({**item, "confidence": min(100, score or 20)})
+ candidates.append({**item, "confidence": _match_confidence(query, item)})
candidates.sort(key=lambda x: int(x.get("confidence") or 0), reverse=True)
return {"status": "success", "query": query, "count": len(candidates), "candidates": _clip_deep(candidates, max_items=50)}
@@ -1087,12 +1091,13 @@ def _mobile_resolve_target(args: dict[str, Any], ctx: McpToolContext) -> dict[st
warnings: list[dict[str, Any]] = []
def extend(kind: str, result: dict[str, Any]) -> None:
- for idx, item in enumerate(_first_list(result, ("candidates", "users", "accounts", "sessions", "contacts", "items"))[:limit]):
+ for item in _first_list(result, ("candidates", "users", "accounts", "sessions", "contacts", "items"))[:limit]:
if not isinstance(item, dict):
continue
username = str(item.get("username") or item.get("id") or item.get("userName") or "").strip()
display = _candidate_display(item)
- confidence = int(item.get("confidence") or max(20, 80 - idx * 8))
+ # 朋友圈用户等结果不带 confidence:沿用联系人的打分规则,不按返回位置给分。
+ confidence = int(item.get("confidence") or _match_confidence(query, item))
candidates.append(
{
"kind": kind,
@@ -1133,7 +1138,9 @@ def extend(kind: str, result: dict[str, Any]) -> None:
candidates.sort(key=lambda x: int(x.get("confidence") or 0), reverse=True)
candidates = candidates[:limit]
best = candidates[0] if candidates else None
- ambiguous = len(candidates) > 1 and best is not None and int(best.get("confidence") or 0) - int(candidates[1].get("confidence") or 0) < 15
+ # 同一个人会同时以联系人、会话、朋友圈用户出现,歧义只和另一个目标比。
+ rival = next((c for c in candidates[1:] if c["id"] != best["id"]), None)
+ ambiguous = rival is not None and int(best.get("confidence") or 0) - int(rival.get("confidence") or 0) < 15
return {"status": "success", "ok": True, "query": query, "targetType": target_type, "count": len(candidates), "best": best, "ambiguous": ambiguous, "candidates": candidates, "warnings": warnings}
@@ -1409,6 +1416,14 @@ def _tools_catalog(args: dict[str, Any], _: McpToolContext) -> dict[str, Any]:
"limit": int_schema("Maximum records to return.", minimum=1, maximum=200),
"offset": int_schema("Pagination offset.", minimum=0),
}
+MESSAGE_SERVER_ID = {
+ "server_id": string_schema("消息的服务端 ID,使用精确的十进制字符串。"),
+ "msg_svr_id": string_schema("server_id 的兼容别名。"),
+}
+EMOJI_REMOTE_SOURCE = {
+ "emoji_url": string_schema("消息返回的表情远程地址,本地缺失时使用。"),
+ "aes_key": string_schema("消息返回的表情解密 key,与 emoji_url 配合使用。"),
+}
def _install_tools() -> None:
@@ -1436,9 +1451,40 @@ def _install_tools() -> None:
_register("wechat.moments.list_timeline", "List Moments timeline by users, keyword, and pagination.", object_schema({**COMMON_ACCOUNT, **PAGING, "usernames": array_schema("Optional poster usernames.", string_schema("Username.")), "keyword": string_schema("Optional content keyword.")}), _sns_timeline, package="wechat.moments")
_register("wechat.moments.search_moments", "Alias for timeline keyword/user search.", object_schema({**COMMON_ACCOUNT, **PAGING, "usernames": array_schema("Optional poster usernames.", string_schema("Username.")), "query": string_schema("Content keyword.")}), _sns_timeline, package="wechat.moments")
_register("wechat.moments.list_users", "List Moments posters with post counts.", object_schema({**COMMON_ACCOUNT, "keyword": string_schema("Optional poster keyword."), "limit": int_schema("Maximum users.", minimum=1, maximum=500)}), _sns_users, package="wechat.moments")
- _register("wechat.moments.get_media_url", "Build a URL for a Moments image resource.", object_schema(additional_properties=True), _sns_media_url, package="wechat.media")
+ _register(
+ "wechat.moments.get_media_url",
+ "获取朋友圈图片链接。",
+ object_schema({
+ **COMMON_ACCOUNT,
+ "post_id": string_schema("朋友圈动态 ID。"),
+ "tid": string_schema("post_id 的兼容别名。"),
+ "media_id": string_schema("动态内的媒体 ID。"),
+ "create_time": int_schema("动态发布时间戳。"),
+ "width": int_schema("图片宽度。", minimum=0),
+ "height": int_schema("图片高度。", minimum=0),
+ "total_size": int_schema("图片文件大小。", minimum=0),
+ "idx": int_schema("同一动态内相同尺寸图片中的序号。", minimum=0),
+ "post_type": int_schema("时间线返回的动态类型。"),
+ "media_type": int_schema("时间线返回的媒体类型。"),
+ "md5": string_schema("图片 MD5。"),
+ "token": string_schema("时间线返回的图片 token。"),
+ "url": string_schema("时间线返回的远程图片地址。"),
+ "key": string_schema("时间线返回的图片解密 key。"),
+ }, additional_properties=True),
+ _sns_media_url, package="wechat.media",
+ )
_register("wechat.moments.get_article_thumb_url", "Build a URL for an official-article thumbnail image.", object_schema({"url": string_schema("Article URL.")}, required=["url"]), _sns_article_thumb_url, package="wechat.media")
- _register("wechat.moments.get_remote_video_url", "Build a URL for a remote Moments video/live-photo resource.", object_schema(additional_properties=True), _sns_video_remote_url, package="wechat.media")
+ _register(
+ "wechat.moments.get_remote_video_url",
+ "获取朋友圈远程视频或实况链接。",
+ object_schema({
+ **COMMON_ACCOUNT,
+ "url": string_schema("时间线返回的远程视频或实况地址。"),
+ "token": string_schema("时间线返回的视频 token。"),
+ "key": string_schema("时间线返回的视频解密 key。"),
+ }, additional_properties=True),
+ _sns_video_remote_url, package="wechat.media",
+ )
_register("wechat.moments.get_local_video_url", "Build a URL for a local cached Moments video resource.", object_schema({**COMMON_ACCOUNT, "post_id": string_schema("Moments post id."), "media_id": string_schema("Media id.")}, required=["post_id", "media_id"]), _sns_video_url, package="wechat.media")
_register("wechat.biz.list_accounts", "List official account/service account message sources.", object_schema(COMMON_ACCOUNT), _biz_accounts, package="wechat.biz")
@@ -1471,10 +1517,50 @@ def _install_tools() -> None:
}, additional_properties=True),
_chat_image_url, package="wechat.media",
)
- _register("wechat.media.get_chat_emoji_url", "Build a URL for a chat emoji message resource.", object_schema(additional_properties=True), _chat_emoji_url, package="wechat.media")
- _register("wechat.media.get_chat_video_thumb_url", "Build a URL for a chat video thumbnail.", object_schema(additional_properties=True), _chat_video_thumb_url, package="wechat.media")
- _register("wechat.media.get_chat_video_url", "Build a URL for a chat video resource.", object_schema(additional_properties=True), _chat_video_url, package="wechat.media")
- _register("wechat.media.get_chat_voice_url", "Build a URL for a chat voice file. This does not transcribe audio.", object_schema(additional_properties=True), _chat_voice_url, package="wechat.media")
+ _register(
+ "wechat.media.get_chat_emoji_url",
+ "获取聊天表情链接。",
+ object_schema({
+ **COMMON_ACCOUNT,
+ "md5": string_schema("表情 MD5。"),
+ "username": string_schema("表情所属会话。"),
+ **EMOJI_REMOTE_SOURCE,
+ }, additional_properties=True),
+ _chat_emoji_url, package="wechat.media",
+ )
+ _register(
+ "wechat.media.get_chat_video_thumb_url",
+ "获取聊天视频缩略图链接。",
+ object_schema({
+ **COMMON_ACCOUNT,
+ "md5": string_schema("视频缩略图 MD5。"),
+ "file_id": string_schema("视频缩略图文件标识。"),
+ "username": string_schema("视频所属会话。"),
+ "deep_scan": bool_schema("允许扩大本地文件搜索范围。", default=False),
+ }, additional_properties=True),
+ _chat_video_thumb_url, package="wechat.media",
+ )
+ _register(
+ "wechat.media.get_chat_video_url",
+ "获取聊天视频链接。",
+ object_schema({
+ **COMMON_ACCOUNT,
+ "md5": string_schema("视频 MD5。"),
+ "file_id": string_schema("视频文件标识。"),
+ "username": string_schema("视频所属会话。"),
+ "deep_scan": bool_schema("允许扩大本地文件搜索范围。", default=False),
+ }, additional_properties=True),
+ _chat_video_url, package="wechat.media",
+ )
+ _register(
+ "wechat.media.get_chat_voice_url",
+ "获取聊天语音文件链接,不做语音转写。",
+ object_schema({
+ **COMMON_ACCOUNT,
+ **MESSAGE_SERVER_ID,
+ }, additional_properties=True),
+ _chat_voice_url, package="wechat.media",
+ )
_register("wechat.media.get_decrypted_resource_url", "Build a URL for a previously decrypted resource by MD5.", object_schema({**COMMON_ACCOUNT, "md5": string_schema("32-character resource md5.")}, required=["md5"]), _decrypted_media_resource_url, package="wechat.media")
_register("wechat.media.get_proxy_image_url", "Build a backend proxy URL for a remote chat image.", object_schema({"url": string_schema("Remote image URL.")}, required=["url"]), _chat_proxy_image_url, package="wechat.media")
_register("wechat.media.get_favicon_url", "Build a backend URL for a web page favicon.", object_schema({"url": string_schema("Page URL.")}, required=["url"]), _chat_favicon_url, package="wechat.media")
@@ -1488,8 +1574,42 @@ def _install_tools() -> None:
_register("wechat.mobile.get_chat_context", "Return a compact chat context by recent page, anchor, or day. Recent mode defaults to live WeChat data when available.", object_schema({**COMMON_ACCOUNT, **CHAT_SOURCE, "username": string_schema("Session username."), "target": string_schema("Optional fuzzy session clue."), "mode": string_schema("recent, around, or day."), "anchor_id": string_schema("Message anchor id."), "message_id": string_schema("Alias for anchor_id."), "date": string_schema("YYYY-MM-DD for day mode."), "limit": int_schema("Message count.", minimum=1, maximum=100), "offset": int_schema("Message offset.", minimum=0), "order": string_schema("asc or desc."), "render_types": string_schema("Optional render type filter."), "before": int_schema("Messages before anchor.", minimum=0, maximum=30), "after": int_schema("Messages after anchor.", minimum=0, maximum=30)}), _mobile_get_chat_context, package="wechat.mobile")
_register("wechat.mobile.get_session_bundle", "Return one session's metadata, messages, and optional calendar counts for mobile UI. Messages default to live WeChat data when available.", object_schema({**COMMON_ACCOUNT, **CHAT_SOURCE, "username": string_schema("Session username."), "limit": int_schema("Message count.", minimum=1, maximum=100), "offset": int_schema("Message offset.", minimum=0), "order": string_schema("asc or desc."), "render_types": string_schema("Optional render type filter."), "year": int_schema("Optional year for daily counts."), "month": int_schema("Optional month for daily counts.", minimum=1, maximum=12)}, required=["username"]), _mobile_session_bundle, package="wechat.mobile")
_register("wechat.mobile.search_moments", "Search Moments posts with compact media references.", object_schema({**COMMON_ACCOUNT, "query": string_schema("Content keyword."), "poster": string_schema("Optional poster clue."), "usernames": array_schema("Poster usernames.", string_schema("Username.")), "limit": int_schema("Post count.", minimum=1, maximum=30), "offset": int_schema("Offset cursor.", minimum=0)}), _mobile_search_moments, package="wechat.mobile")
- _register("wechat.mobile.get_media_links", "Return URL resources for chat, Moments, avatar, link, or emoji media.", object_schema(additional_properties=True), _mobile_get_media_links, package="wechat.mobile")
- _register("wechat.mobile.get_message_media_bundle", "Return likely media URLs for a message or link without fetching binary content.", object_schema(additional_properties=True), _mobile_message_media_bundle, package="wechat.mobile")
+ _register(
+ "wechat.mobile.get_media_links",
+ "返回聊天、朋友圈、头像、链接或表情媒体的链接。这里只列常用参数;需要聊天大图或更多定位参数时,改用 wechat.media.get_chat_image_url、wechat.moments.get_media_url 等专用工具。",
+ object_schema({
+ **COMMON_ACCOUNT,
+ "kind": string_schema("资源类型。auto 按已提供的 username、md5、file_id、server_id、emoji_url 返回头像和聊天媒体链接;也可指定 avatar、chat_image、chat_emoji、chat_video_thumb、chat_video、chat_voice、moments_image、moments_video、favicon 或 proxy_image。", default="auto"),
+ "max_items": int_schema("最多返回的链接数,默认 20。", minimum=1, maximum=20),
+ "username": string_schema("头像或聊天媒体所属会话。"),
+ "md5": string_schema("聊天图片、表情、视频或朋友圈图片的 MD5。"),
+ "file_id": string_schema("聊天图片或视频的文件标识。"),
+ **MESSAGE_SERVER_ID,
+ **EMOJI_REMOTE_SOURCE,
+ "post_id": string_schema("朋友圈动态 ID。"),
+ "media_id": string_schema("动态内的媒体 ID。"),
+ "token": string_schema("时间线返回的媒体 token。"),
+ "key": string_schema("时间线返回的媒体解密 key。"),
+ "url": string_schema("朋友圈远程图片或视频地址,或 favicon、proxy_image 的目标地址。"),
+ }, additional_properties=True),
+ _mobile_get_media_links, package="wechat.mobile",
+ )
+ _register(
+ "wechat.mobile.get_message_media_bundle",
+ "返回一条消息或链接可能用到的媒体链接,不读取二进制内容。",
+ object_schema({
+ **COMMON_ACCOUNT,
+ "username": string_schema("消息所属会话。"),
+ "session_id": string_schema("username 的兼容别名。"),
+ **MESSAGE_SERVER_ID,
+ "md5": string_schema("消息返回的图片、视频或表情 MD5。"),
+ "file_id": string_schema("消息返回的图片或视频文件标识。"),
+ **EMOJI_REMOTE_SOURCE,
+ "url": string_schema("链接消息的网页或图片地址。"),
+ "link_url": string_schema("url 的兼容别名。"),
+ }, additional_properties=True),
+ _mobile_message_media_bundle, package="wechat.mobile",
+ )
_register("wechat.mobile.get_analytics", "Return compact analytics data by metric without loading full annual payloads. Chat daily-count analytics default to live WeChat data when available.", object_schema({**CHAT_SOURCE}, additional_properties=True), _mobile_get_analytics, package="wechat.mobile")
diff --git a/src/wechat_decrypt_tool/media_helpers.py b/src/wechat_decrypt_tool/media_helpers.py
index 38bfbb7d..147d177d 100644
--- a/src/wechat_decrypt_tool/media_helpers.py
+++ b/src/wechat_decrypt_tool/media_helpers.py
@@ -13,6 +13,7 @@
import tempfile
import threading
import time
+from collections import OrderedDict
from dataclasses import dataclass
from functools import lru_cache
from pathlib import Path
@@ -1277,12 +1278,167 @@ def _get_wxam_decoder():
return None
+# wxgf 容器里的每个分区都是独立的 Annex-B HEVC 码流,以「起始码 + VPS」开头。
+_WXGF_HEVC_PARTITION_START = b"\x00\x00\x00\x01\x40\x01"
+_WXGF_FFMPEG_TIMEOUT_SECONDS = 20
+# 数据可能来自远程地址,体积很小的文件也能声明任意多的分区、NAL 或极大的画面,所以都设上限:
+# 真实文件至多两个分区([alpha, 画面]),单帧图片只有少量 NAL,画面尺寸也远小于这个像素数。
+_WXGF_MAX_PARTITIONS = 2
+_WXGF_MAX_NAL_UNITS = 4096
+_WXGF_FFMPEG_MAX_PIXELS = 50_000_000
+# 同一份 wxgf 会被反复送进来(一次读取内的多次重试、每次图片请求重新比较各个变体),
+# 按内容记住转换结果(包括失败),不为同样的数据再启动 ffmpeg。
+_WXGF_FFMPEG_CACHE_BYTES = 16 * 1024 * 1024
+_WXGF_FFMPEG_CACHE_ENTRIES = 512
+_WXGF_FFMPEG_CACHE: "OrderedDict[bytes, Optional[bytes]]" = OrderedDict()
+_WXGF_FFMPEG_CACHE_LOCK = threading.Lock()
+
+
+def _wxgf_hevc_partitions(data: bytes) -> list[bytes]:
+ """按「起始码 + VPS」切出 wxgf 里的各个 HEVC 分区。
+
+ 头部与码流之间的元数据(如 ICC)里会出现形似起始码的字节,所以不能直接取第一个起始码。
+ """
+ starts: list[int] = []
+ pos = data.find(_WXGF_HEVC_PARTITION_START, 4)
+ while pos >= 0:
+ if len(starts) >= _WXGF_MAX_PARTITIONS:
+ return []
+ starts.append(pos)
+ pos = data.find(_WXGF_HEVC_PARTITION_START, pos + len(_WXGF_HEVC_PARTITION_START))
+
+ partitions: list[bytes] = []
+ for start, limit in zip(starts, starts[1:] + [len(data)]):
+ # 分区前是 4 字节大端长度,按它截断(后面还跟着下一个长度前缀或容器尾部数据);
+ # 长度对不上说明文件被截断或不是这种布局,整个放弃,不把残缺的码流交给解码器。
+ size = int.from_bytes(data[start - 4 : start], "big")
+ if not 0 < size <= limit - start:
+ return []
+ partitions.append(data[start : start + size])
+ return partitions
+
+
+def _hevc_picture_count(stream: bytes) -> int:
+ """数出码流里的编码图像:VCL NAL 且 first_slice_segment_in_pic_flag 为 1。
+
+ NAL 多到不像一张图片时不再往下数,按多帧处理。
+ """
+ count = 0
+ nal_units = 0
+ pos = stream.find(b"\x00\x00\x01")
+ while 0 <= pos < len(stream) - 5:
+ nal_units += 1
+ if nal_units > _WXGF_MAX_NAL_UNITS:
+ return max(count, 2)
+ if (stream[pos + 3] >> 1) & 0x3F < 32 and stream[pos + 5] & 0x80:
+ count += 1
+ pos = stream.find(b"\x00\x00\x01", pos + 3)
+ return count
+
+
+def _wxgf_ffmpeg_first_frame(ffmpeg_exe: str, stream: bytes, output_args: list[str], deadline: float) -> bytes:
+ """把一段 HEVC 码流经 stdin 交给 ffmpeg,从 stdout 取回首帧;失败返回 b""。"""
+ import subprocess
+
+ # 各分区共用一个截止时间,一份文件的总耗时不超过 _WXGF_FFMPEG_TIMEOUT_SECONDS。
+ timeout = deadline - time.monotonic()
+ if timeout <= 0:
+ logger.warning(f"wxgf ffmpeg decode timed out after {_WXGF_FFMPEG_TIMEOUT_SECONDS}s")
+ return b""
+ try:
+ proc = subprocess.run(
+ [
+ ffmpeg_exe,
+ "-hide_banner",
+ "-loglevel",
+ "error",
+ # 解码器一报错就整体失败,尽量不把残缺的画面当成有效图片交给调用方缓存。
+ "-xerror",
+ "-err_detect",
+ "explode",
+ # 超出像素上限的画面在分配解码缓冲之前就被拒绝。
+ "-max_pixels",
+ str(_WXGF_FFMPEG_MAX_PIXELS),
+ "-f",
+ "hevc",
+ "-i",
+ "pipe:0",
+ "-frames:v",
+ "1",
+ *output_args,
+ "pipe:1",
+ ],
+ input=stream,
+ check=False,
+ capture_output=True,
+ timeout=timeout,
+ creationflags=getattr(subprocess, "CREATE_NO_WINDOW", 0) if os.name == "nt" else 0,
+ )
+ except subprocess.TimeoutExpired:
+ logger.warning(f"wxgf ffmpeg decode timed out after {_WXGF_FFMPEG_TIMEOUT_SECONDS}s")
+ return b""
+ if proc.returncode == 0 and proc.stdout:
+ return proc.stdout
+ # 损坏的码流能让 ffmpeg 输出上千行错误,只记录结尾一小段。
+ err = (proc.stderr or b"")[-400:].decode("utf-8", errors="ignore").strip()
+ logger.warning(f"wxgf ffmpeg decode failed (rc={proc.returncode}): {err or 'no output'}")
+ return b""
+
+
+def _wxgf_partitions_to_jpeg(ffmpeg_exe: str, partitions: list[bytes]) -> Optional[bytes]:
+ deadline = time.monotonic() + _WXGF_FFMPEG_TIMEOUT_SECONDS
+ for alpha in partitions[:-1]:
+ plane = _wxgf_ffmpeg_first_frame(ffmpeg_exe, alpha, ["-f", "rawvideo", "-pix_fmt", "gray"], deadline)
+ if not plane or plane.count(b"\xff") != len(plane):
+ return None
+ out = _wxgf_ffmpeg_first_frame(
+ ffmpeg_exe, partitions[-1], ["-f", "image2pipe", "-codec:v", "mjpeg", "-q:v", "3"], deadline
+ )
+ return out if _detect_image_media_type(out[:32]) == "image/jpeg" else None
+
+
+def _wxgf_to_jpeg_with_ffmpeg(data: bytes) -> Optional[bytes]:
+ """没有 WxAM 解码器时(非 Windows,或 DLL 缺失),用 ffmpeg 把静态、不透明的 wxgf 转成 JPEG。
+
+ 带透明度的 wxgf 有两个分区,实测顺序固定为 [alpha, 画面]:取最后一个分区为画面,
+ 它之前的分区必须解出来全不透明。单帧 JPEG 表示不了的情况(多帧动图、alpha 不是
+ 全不透明)返回 None,调用方照旧走原有的回退(如远程表情),不会缓存降级后的画面。
+ 输出不带容器里的 ICC;边长为奇数时按 HEVC 编码尺寸多出 1 像素。
+ """
+ try:
+ ffmpeg_exe = _find_ffmpeg_executable()
+ if not ffmpeg_exe:
+ return None
+ # 调用方只要在数据里扫到 "wxgf" 就会进来,没有可解码的单帧码流时不启动进程。
+ partitions = _wxgf_hevc_partitions(data)
+ if not partitions or _hevc_picture_count(partitions[-1]) != 1:
+ return None
+
+ key = hashlib.sha256(data).digest()
+ with _WXGF_FFMPEG_CACHE_LOCK:
+ if key in _WXGF_FFMPEG_CACHE:
+ _WXGF_FFMPEG_CACHE.move_to_end(key)
+ return _WXGF_FFMPEG_CACHE[key]
+ converted = _wxgf_partitions_to_jpeg(ffmpeg_exe, partitions)
+ with _WXGF_FFMPEG_CACHE_LOCK:
+ _WXGF_FFMPEG_CACHE[key] = converted
+ while (
+ len(_WXGF_FFMPEG_CACHE) > _WXGF_FFMPEG_CACHE_ENTRIES
+ or sum(len(value or b"") for value in _WXGF_FFMPEG_CACHE.values()) > _WXGF_FFMPEG_CACHE_BYTES
+ ):
+ _WXGF_FFMPEG_CACHE.popitem(last=False)
+ return converted
+ except Exception as e:
+ logger.warning(f"wxgf to JPEG conversion failed: {e}")
+ return None
+
+
def _wxgf_to_image_bytes(data: bytes) -> Optional[bytes]:
if not data or not data.startswith(b"wxgf"):
return None
fn = _get_wxam_decoder()
if fn is None:
- return None
+ return _wxgf_to_jpeg_with_ffmpeg(data)
max_output_size = 52 * 1024 * 1024
for mode in (0, 3):
@@ -3157,6 +3313,16 @@ def _guess_media_type_by_path(path: Path, fallback: str = "application/octet-str
return fallback
+@lru_cache(maxsize=256)
+def _xor_table(key: int) -> bytes:
+ return bytes(b ^ key for b in range(256))
+
+
+def _xor_bytes(data: bytes, key: int) -> bytes:
+ """单字节 XOR:查表交给 bytes.translate,避免对整个文件逐字节跑 Python 循环。"""
+ return bytes(data).translate(_xor_table(key))
+
+
def _try_xor_decrypt_by_magic(data: bytes) -> tuple[Optional[bytes], Optional[str]]:
if not data:
return None, None
@@ -3200,7 +3366,7 @@ def _try_xor_decrypt_by_magic(data: bytes) -> tuple[Optional[bytes], Optional[st
if not ok:
continue
- decoded = bytes(b ^ key for b in data)
+ decoded = _xor_bytes(data, key)
if magic == b"wxgf":
try:
@@ -3244,7 +3410,7 @@ def _try_xor_decrypt_by_magic(data: bytes) -> tuple[Optional[bytes], Optional[st
if preview_len > 0:
for key in range(256):
try:
- pv = bytes(b ^ key for b in data[:preview_len])
+ pv = _xor_bytes(data[:preview_len], key)
except Exception:
continue
try:
@@ -3258,7 +3424,7 @@ def _try_xor_decrypt_by_magic(data: bytes) -> tuple[Optional[bytes], Optional[st
or (scan.find(b"RIFF") >= 0)
or (scan.find(b"ftyp") >= 0)
):
- decoded = bytes(b ^ key for b in data)
+ decoded = _xor_bytes(data, key)
dec2, mt2 = _try_strip_media_prefix(decoded)
if mt2 != "application/octet-stream":
if mt2.startswith("image/") and (not _is_probably_valid_image(dec2, mt2)):
@@ -3437,7 +3603,15 @@ def _save_media_keys(account_dir: Path, xor_key: int, aes_key16: Optional[bytes]
"xor": int(xor_key),
"aes": aes_str,
}
- (account_dir / "_media_keys.json").write_text(
+ path = account_dir / "_media_keys.json"
+ # 密钥文件只允许属主读写:先建成/收紧到 0600 再写入内容。Windows 没有这套权限位,不处理。
+ if os.name == "posix":
+ try:
+ path.touch(mode=0o600, exist_ok=True)
+ os.chmod(path, 0o600)
+ except Exception:
+ pass
+ path.write_text(
json.dumps(payload, ensure_ascii=False, indent=2),
encoding="utf-8",
)
@@ -3446,7 +3620,7 @@ def _save_media_keys(account_dir: Path, xor_key: int, aes_key16: Optional[bytes]
def _decrypt_wechat_dat_v3(data: bytes, xor_key: int) -> bytes:
- return bytes(b ^ xor_key for b in data)
+ return _xor_bytes(data, xor_key)
def _decrypt_wechat_dat_v4(data: bytes, xor_key: int, aes_key: bytes) -> bytes:
@@ -3466,7 +3640,7 @@ def _decrypt_wechat_dat_v4(data: bytes, xor_key: int, aes_key: bytes) -> bytes:
if xor_size > 0:
raw_data = rest[aes_size:-xor_size]
xor_data = rest[-xor_size:]
- xored_data = bytes(b ^ xor_key for b in xor_data)
+ xored_data = _xor_bytes(xor_data, xor_key)
else:
xored_data = b""
diff --git a/src/wechat_decrypt_tool/routers/account_archive_export.py b/src/wechat_decrypt_tool/routers/account_archive_export.py
index 28e21bb1..331d6fce 100644
--- a/src/wechat_decrypt_tool/routers/account_archive_export.py
+++ b/src/wechat_decrypt_tool/routers/account_archive_export.py
@@ -538,9 +538,16 @@ async def export_account_archive(req: AccountArchiveExportRequest):
return {"status": "success", "job": job.to_public_dict()}
-@router.get("/api/account/archive_export/download", summary="Download account archive by file path")
-async def download_account_archive(path: str):
- zip_path = Path(str(path or "").strip()).expanduser().resolve()
+@router.get("/api/account/archive_export/{export_id}/download", summary="Download account archive export file")
+async def download_account_archive(export_id: str):
+ job = _get_job(export_id)
+ if not job:
+ raise HTTPException(status_code=404, detail="Export not found.")
+ # zip_path 由请求里的 output_dir/file_name 拼出,任务一开始就已写入;只有任务完成后它才是本次导出的产物,
+ # 这里的 done 判断不能放宽成“文件存在即可”。
+ if job.status != "done" or not job.zip_path:
+ raise HTTPException(status_code=409, detail="Export not ready.")
+ zip_path = Path(job.zip_path)
if not zip_path.exists() or not zip_path.is_file():
raise HTTPException(status_code=404, detail="Export file not found.")
if zip_path.suffix.lower() not in {".zip", ".wec"}:
diff --git a/src/wechat_decrypt_tool/routers/chat_contacts.py b/src/wechat_decrypt_tool/routers/chat_contacts.py
index 9dfcd5d1..c60f6d3c 100644
--- a/src/wechat_decrypt_tool/routers/chat_contacts.py
+++ b/src/wechat_decrypt_tool/routers/chat_contacts.py
@@ -519,6 +519,65 @@ def _build_contact_pinyin_initial(name: str) -> str:
return "#"
+# pypinyin 把 ü 记作 v。lüe / nüe 另接受输入法同样支持的 lue / nue;lü / nü 与输入法一致,只有 lv / nv。
+_PINYIN_UE_SPELLINGS = {"lve": "lue", "nve": "nue"}
+
+
+# 关键词筛选每次都按相同顺序查询全部联系人的备注/昵称,名称数一旦超过容量,缓存就会整体失效
+# (每个请求都全部重新转换),所以容量要比上面两个缓存大得多。
+@lru_cache(maxsize=65536)
+def _build_contact_pinyin_search_forms(name: str) -> tuple[tuple[str, str, tuple[int, ...]], ...]:
+ # 返回若干组 (首字母串, 全拼串, 各音节在全拼串中的起始下标),供关键词按拼音匹配。
+ text = _normalize_text(name)
+ if not text:
+ return ()
+
+ # errors=list 让非汉字逐字符原样返回。一个汉字都没转换出来的名称(如 "Li Wei🌸")不做拼音匹配,
+ # 否则去掉分隔符后会比普通子串匹配更宽松。
+ parts = lazy_pinyin(text, style=Style.NORMAL, errors=list)
+ if "".join(parts) == text:
+ return ()
+
+ # 汉字取拼音,ASCII 字母各算一个音节;数字、空格、标点、表情、全角字母等都丢弃(关键词只含 ASCII 字母)。
+ syllables = ["".join(_PINYIN_ALPHA_RE.findall(part)).lower() for part in parts]
+ syllables = [syllable for syllable in syllables if syllable]
+ if not syllables:
+ return ()
+
+ # 首字是多音字姓氏时同时保留姓氏读音与默认读音(“曾国藩”按 zeng,“乐乐”仍可按 le 搜到)。
+ readings = [syllables]
+ override = _SURNAME_PINYIN_OVERRIDES.get(text[0])
+ if override and override != syllables[0]:
+ readings.append([override, *syllables[1:]])
+ for reading in list(readings):
+ respelled = [_PINYIN_UE_SPELLINGS.get(syllable, syllable) for syllable in reading]
+ if respelled != reading:
+ readings.append(respelled)
+
+ forms: list[tuple[str, str, tuple[int, ...]]] = []
+ for reading in readings:
+ starts: list[int] = []
+ offset = 0
+ for syllable in reading:
+ starts.append(offset)
+ offset += len(syllable)
+ forms.append(("".join(syllable[0] for syllable in reading), "".join(reading), tuple(starts)))
+ return tuple(forms)
+
+
+def _matches_pinyin_keyword(name: str, keyword: str) -> bool:
+ # 纯 ASCII 名称不含汉字,直接跳过,也不占用缓存。
+ if (not name) or name.isascii():
+ return False
+ for initials, full, starts in _build_contact_pinyin_search_forms(name):
+ if keyword in initials:
+ return True
+ # 全拼只从音节边界起匹配,避免 "an" 命中 "zhangwei"。
+ if any(full.startswith(keyword, start) for start in starts):
+ return True
+ return False
+
+
def _decode_varint(raw: bytes, offset: int) -> tuple[Optional[int], int]:
value = 0
shift = 0
@@ -1417,7 +1476,7 @@ def _infer_contact_type(username: str, row: dict[str, Any]) -> Optional[str]:
return None
-def _matches_keyword(contact: dict[str, Any], keyword: str) -> bool:
+def _matches_keyword(contact: dict[str, Any], keyword: str, *, pinyin: bool = True) -> bool:
kw = _normalize_text(keyword).lower()
if not kw:
return True
@@ -1437,6 +1496,11 @@ def _matches_keyword(contact: dict[str, Any], keyword: str) -> bool:
for field in fields:
if kw in _normalize_text(field).lower():
return True
+
+ # 纯字母关键词再按拼音匹配名称:首字母(zw -> 张伟)或从任一音节起的全拼前缀(zhangw / wei)。
+ if pinyin and kw.isascii() and kw.isalpha():
+ names = {_normalize_text(contact.get(key)) for key in ("displayName", "remark", "nickname")}
+ return any(_matches_pinyin_keyword(name, kw) for name in names)
return False
@@ -2141,6 +2205,8 @@ def _collect_contacts_for_account_realtime(
contacts.sort(
key=lambda x: (
+ # 字段子串命中的排在仅拼音命中的之前,按条数截取结果的调用方(如 MCP)不会因拼音命中而丢掉原有结果。
+ not _matches_keyword(x, keyword or "", pinyin=False),
-_to_int(x.get("_sortTs", 0)),
_normalize_text(x.get("displayName", "")).lower(),
_normalize_text(x.get("username", "")).lower(),
@@ -2344,6 +2410,8 @@ def _collect_contacts_for_account(
contacts.sort(
key=lambda x: (
+ # 字段子串命中的排在仅拼音命中的之前,按条数截取结果的调用方(如 MCP)不会因拼音命中而丢掉原有结果。
+ not _matches_keyword(x, keyword or "", pinyin=False),
-_to_int(x.get("_sortTs", 0)),
_normalize_text(x.get("displayName", "")).lower(),
_normalize_text(x.get("username", "")).lower(),
diff --git a/src/wechat_decrypt_tool/routers/chat_export.py b/src/wechat_decrypt_tool/routers/chat_export.py
index db355699..cd905e0d 100644
--- a/src/wechat_decrypt_tool/routers/chat_export.py
+++ b/src/wechat_decrypt_tool/routers/chat_export.py
@@ -31,6 +31,7 @@
"link",
"transfer",
"redPacket",
+ "location",
"system",
"quote",
"voip",
diff --git a/src/wechat_decrypt_tool/routers/chat_media.py b/src/wechat_decrypt_tool/routers/chat_media.py
index 46e98492..18abe1e1 100644
--- a/src/wechat_decrypt_tool/routers/chat_media.py
+++ b/src/wechat_decrypt_tool/routers/chat_media.py
@@ -3086,11 +3086,14 @@ async def get_chat_emoji(
data = b""
media_type = "application/octet-stream"
if p:
- data, media_type = _read_and_maybe_decrypt_media(p, account_dir=account_dir, weixin_root=wxid_dir)
+ # 本地 wxgf 在非 Windows 上要启动 ffmpeg 解码,放到线程里,不阻塞事件循环。
+ data, media_type = await asyncio.to_thread(
+ _read_and_maybe_decrypt_media, p, account_dir=account_dir, weixin_root=wxid_dir
+ )
if media_type == "application/octet-stream":
# Some emojis are stored encrypted (see emoticon.db); try remote fetch as fallback.
- data2, mt2 = _try_fetch_emoticon_from_remote(account_dir, str(md5).lower())
+ data2, mt2 = await asyncio.to_thread(_try_fetch_emoticon_from_remote, account_dir, str(md5).lower())
if data2 is not None and mt2:
data, media_type = data2, mt2
@@ -3113,7 +3116,7 @@ async def get_chat_emoji(
if not blob:
continue
try:
- data2, mt = _try_strip_media_prefix(blob)
+ data2, mt = await asyncio.to_thread(_try_strip_media_prefix, blob)
except Exception:
data2, mt = blob, "application/octet-stream"
diff --git a/src/wechat_decrypt_tool/sns_media.py b/src/wechat_decrypt_tool/sns_media.py
index a24bc2d8..e9220287 100644
--- a/src/wechat_decrypt_tool/sns_media.py
+++ b/src/wechat_decrypt_tool/sns_media.py
@@ -179,6 +179,16 @@ def _sns_cdn_media_source(url: str) -> str:
return "video-or-unknown"
+# These hosts resolve via CNAME to socwxsns.video.qq.com and serve its
+# *.video.qq.com certificate, so https only passes verification under the CNAME
+# target. Checked 2026-10 (SAN has no *.tc.qq.com); re-check with
+# `openssl s_client -connect :443 -verify_hostname ` before removing.
+_SNS_CDN_TLS_HOST_ALIASES = {
+ "vweixinthumb.tc.qq.com": "socwxsns.video.qq.com",
+ "vweixinf.tc.qq.com": "socwxsns.video.qq.com",
+}
+
+
def fix_sns_cdn_url(
url: str,
*,
@@ -189,6 +199,7 @@ def fix_sns_cdn_url(
"""WeFlow-compatible SNS CDN URL normalization.
- Force https for Tencent CDNs.
+ - Swap hosts whose certificate does not cover them for their CNAME target.
- Preserve image size variants by default because Tencent binds `/60`, `/150`,
`/200`, `/480`, and `/0` to their matching credentials.
- Only an explicit original-image request may replace a size suffix with `/0`.
@@ -210,6 +221,10 @@ def fix_sns_cdn_url(
# http -> https
u = re.sub(r"^http://", "https://", u, flags=re.I)
+ tls_alias = _SNS_CDN_TLS_HOST_ALIASES.get(host)
+ if tls_alias:
+ u = re.sub(r"^https://[^/?#]+", f"https://{tls_alias}", u, flags=re.I)
+
if force_original and not is_video:
u = re.sub(r"/(?:60|150|200|480)(?=($|\?))", "/0", u)
diff --git a/src/wechat_decrypt_tool/wechat_detection.py b/src/wechat_decrypt_tool/wechat_detection.py
index bda63d25..356c3550 100644
--- a/src/wechat_decrypt_tool/wechat_detection.py
+++ b/src/wechat_decrypt_tool/wechat_detection.py
@@ -614,11 +614,6 @@ def auto_detect_wechat_data_dirs():
_append_detected_dir(detected_dirs, item_path)
logger.debug("目录扫描检测成功: %s", item_path)
- # macOS default candidates can already point at the data root even when
- # its version name is unfamiliar to this release.
- if sys.platform == "darwin" and _contains_wechat_account_dirs(Path(scan_path)):
- _append_detected_dir(detected_dirs, scan_path)
-
# 策略2:进程内存分析(简化版)
try:
process_list = get_process_list()
diff --git a/tests/test_account_archive_export_download.py b/tests/test_account_archive_export_download.py
new file mode 100644
index 00000000..4cc58978
--- /dev/null
+++ b/tests/test_account_archive_export_download.py
@@ -0,0 +1,110 @@
+"""账号归档下载接口只能按服务端任务的 export_id 取文件,不能由调用方指定路径。"""
+from pathlib import Path
+
+import pytest
+from fastapi import FastAPI
+from fastapi.testclient import TestClient
+
+from wechat_decrypt_tool.routers import account_archive_export
+
+
+ROOT = Path(__file__).resolve().parents[1]
+
+
+@pytest.fixture
+def client():
+ app = FastAPI()
+ app.include_router(account_archive_export.router)
+ with TestClient(app) as test_client:
+ yield test_client
+
+
+@pytest.fixture
+def register_job():
+ registered: list[str] = []
+
+ def _register(**fields):
+ job = account_archive_export.AccountArchiveExportJob(**fields)
+ with account_archive_export._JOBS_LOCK:
+ account_archive_export._JOBS[job.export_id] = job
+ registered.append(job.export_id)
+ return job
+
+ yield _register
+ with account_archive_export._JOBS_LOCK:
+ for export_id in registered:
+ account_archive_export._JOBS.pop(export_id, None)
+
+
+def test_unknown_export_id_returns_404(client):
+ response = client.get("/api/account/archive_export/missing-export/download")
+
+ assert response.status_code == 404
+ assert response.json() == {"detail": "Export not found."}
+
+
+@pytest.mark.parametrize("status", ["queued", "running", "error", "cancelled"])
+def test_unfinished_job_is_not_downloadable(client, register_job, tmp_path, status):
+ # 任务一开始就会写入 zip_path;同名旧文件可能还在磁盘上,不能在任务完成前被取走。
+ stale = tmp_path / "wechat_archive_account.zip"
+ stale.write_bytes(b"stale-archive")
+ register_job(export_id=f"archive-{status}", status=status, zip_path=str(stale), file_name=stale.name)
+
+ response = client.get(f"/api/account/archive_export/archive-{status}/download")
+
+ assert response.status_code == 409
+ assert response.json() == {"detail": "Export not ready."}
+
+
+@pytest.mark.parametrize(
+ ("file_name", "media_type"),
+ [
+ ("wechat_archive_account.zip", "application/zip"),
+ ("wechat_archive_account.zip.wec", "application/octet-stream"),
+ ],
+)
+def test_finished_job_serves_its_own_file(client, register_job, tmp_path, file_name, media_type):
+ archive = tmp_path / file_name
+ archive.write_bytes(b"archive-bytes")
+ register_job(export_id="archive-done", status="done", zip_path=str(archive), file_name=file_name)
+
+ response = client.get("/api/account/archive_export/archive-done/download")
+
+ assert response.status_code == 200
+ assert response.content == b"archive-bytes"
+ assert response.headers["content-type"] == media_type
+ assert f'filename="{file_name}"' in response.headers["content-disposition"]
+
+
+def test_finished_job_with_removed_file_returns_404(client, register_job, tmp_path):
+ removed = tmp_path / "wechat_archive_account.zip"
+ register_job(export_id="archive-removed", status="done", zip_path=str(removed), file_name=removed.name)
+
+ response = client.get("/api/account/archive_export/archive-removed/download")
+
+ assert response.status_code == 404
+ assert response.json() == {"detail": "Export file not found."}
+
+
+def test_path_query_cannot_read_other_files(client, register_job, tmp_path):
+ other = tmp_path / "other" / "private.zip"
+ other.parent.mkdir()
+ other.write_bytes(b"not-an-export")
+ archive = tmp_path / "wechat_archive_account.zip"
+ archive.write_bytes(b"archive-bytes")
+ register_job(export_id="archive-done", status="done", zip_path=str(archive), file_name=archive.name)
+
+ legacy = client.get("/api/account/archive_export/download", params={"path": str(other)})
+ assert legacy.status_code == 404
+ assert b"not-an-export" not in legacy.content
+
+ overridden = client.get("/api/account/archive_export/archive-done/download", params={"path": str(other)})
+ assert overridden.status_code == 200
+ assert overridden.content == b"archive-bytes"
+
+
+def test_frontend_downloads_archive_by_export_id():
+ source = (ROOT / "frontend" / "components" / "GlobalExportDialog.vue").read_text(encoding="utf-8")
+
+ assert "/account/archive_export/${encodeURIComponent(currentExportId.value)}/download" in source
+ assert "archive_export/download?" not in source
diff --git a/tests/test_ai_services.py b/tests/test_ai_services.py
index 42f11f84..e18d7f7a 100644
--- a/tests/test_ai_services.py
+++ b/tests/test_ai_services.py
@@ -128,6 +128,71 @@ def test_profile_mask_and_url_validation(service):
validate_url("http://localhost:11434/v1")
+@pytest.mark.parametrize("value", [
+ "http://127.0.0.1:11434/v1",
+ "http://[::1]:11434/v1",
+ # RFC 1918 私有网段及 IPv6 唯一本地地址的首尾。
+ "http://10.0.0.1:11434/v1",
+ "http://10.255.255.254:8000/v1",
+ "http://172.16.0.1:1234/v1",
+ "http://172.31.255.254:1234/v1",
+ "http://192.168.0.1:11434/v1",
+ "http://192.168.255.254:4646/v1",
+ "http://[fc00::1]:11434/v1",
+ "http://[fdff:ffff::1]:11434/v1",
+ # HTTPS 不受网段限制。
+ "https://192.168.1.5:8443/v1",
+])
+def test_url_accepted_for_loopback_lan_ip_literals_and_https(value):
+ validate_url(value)
+
+
+@pytest.mark.parametrize("value", [
+ # 域名不解析,即使看起来指向局域网。
+ "http://example.com/v1",
+ "http://nas.local:11434/v1",
+ "http://192.168.1.5.nip.io/v1",
+ # 公网地址及紧邻私有网段的地址。
+ "http://8.8.8.8/v1",
+ "http://9.255.255.255/v1",
+ "http://11.0.0.1/v1",
+ "http://172.15.255.255/v1",
+ "http://172.32.0.1/v1",
+ "http://192.167.255.255/v1",
+ "http://192.169.0.1/v1",
+ "http://[2606:4700:4700::1111]/v1",
+ "http://[fe00::1]/v1",
+ "http://[fd::1]/v1",
+ # 本机只认三个固定写法,其余回环地址不放行。
+ "http://127.0.0.2/v1",
+ "http://[::ffff:127.0.0.1]/v1",
+ # 未指定、链路本地(含云元数据地址)、运营商级 NAT、文档与测试保留网段。
+ "http://0.0.0.0:11434/v1",
+ "http://169.254.169.254/latest",
+ "http://100.64.0.1/v1",
+ "http://192.0.2.1/v1",
+ "http://198.18.0.1/v1",
+ "http://[::]/v1",
+ "http://[fe80::1]/v1",
+ "http://[2001:db8::1]/v1",
+ # 嵌入 IPv4 的 IPv6 地址不展开,带 zone id 的写法不接受。
+ "http://[::ffff:192.168.1.5]/v1",
+ "http://[::ffff:c0a8:105]/v1",
+ "http://[::ffff:8.8.8.8]/v1",
+ "http://[64:ff9b::a00:1]/v1",
+ "http://[fd00::1%25eth0]:11434/v1",
+ "http://[fd00::1%eth0]:11434/v1",
+ # 非规范 IPv4 写法在不同解析器下含义不一致。
+ "http://010.0.0.1/v1",
+ "http://10.1/v1",
+ "http://167772161/v1",
+])
+def test_http_rejected_outside_loopback_and_lan_ip_literals(value):
+ # 同时确认报错来自这条规则,并向用户说明了局域网例外。
+ with pytest.raises(ValueError, match="局域网 IP.*必须使用 HTTPS"):
+ validate_url(value)
+
+
def test_summary_graph_and_checkpoint_replay(service):
async def run():
task = service.create_task(task_options())
diff --git a/tests/test_chat_export_message_types_semantics.py b/tests/test_chat_export_message_types_semantics.py
index fbaaa7a6..de904c1e 100644
--- a/tests/test_chat_export_message_types_semantics.py
+++ b/tests/test_chat_export_message_types_semantics.py
@@ -1,6 +1,7 @@
import os
import json
import hashlib
+import re
import sqlite3
import sys
import unittest
@@ -14,6 +15,14 @@
sys.path.insert(0, str(ROOT / "src"))
+def _dialog_message_types() -> list[str]:
+ """导出面板默认勾选并提交的消息类型。"""
+
+ source = (ROOT / "frontend" / "composables" / "chat" / "useChatExport.js").read_text(encoding="utf-8")
+ options = source.split("const exportMessageTypeOptions = [", 1)[1].split("]", 1)[0]
+ return re.findall(r"value: '([^']+)'", options)
+
+
class TestChatExportMessageTypesSemantics(unittest.TestCase):
def _reload_export_modules(self):
import wechat_decrypt_tool.app_paths as app_paths
@@ -403,6 +412,47 @@ def test_checked_location_exports_location_fields(self):
else:
os.environ["WECHAT_TOOL_DATA_DIR"] = prev_data
+ def test_export_request_accepts_every_dialog_message_type(self):
+ from wechat_decrypt_tool.routers.chat_export import ChatExportCreateRequest
+
+ dialog_types = _dialog_message_types()
+ self.assertIn("location", dialog_types)
+ self.assertEqual(len(dialog_types), len(set(dialog_types)))
+ self.assertEqual(ChatExportCreateRequest(message_types=dialog_types).message_types, dialog_types)
+
+ def test_dialog_default_types_export_location_message(self):
+ with TemporaryDirectory() as td:
+ root = Path(td)
+ account = "wxid_test"
+ username = "wxid_friend"
+ self._prepare_account(root, account=account, username=username)
+
+ prev_data = os.environ.get("WECHAT_TOOL_DATA_DIR")
+ try:
+ os.environ["WECHAT_TOOL_DATA_DIR"] = str(root)
+ svc = self._reload_export_modules()
+ job = self._create_job(
+ svc.CHAT_EXPORT_MANAGER,
+ account=account,
+ username=username,
+ message_types=_dialog_message_types(),
+ include_media=False,
+ )
+ self.assertEqual(job.status, "done", msg=job.error)
+
+ payload, manifest, _ = self._load_export_payload(job.zip_path)
+ location_msg = next((m for m in payload.get("messages", []) if int(m.get("type") or 0) == 48), None)
+ self.assertIsNotNone(location_msg)
+ self.assertEqual(str(location_msg.get("renderType") or ""), "location")
+ self.assertEqual(str(location_msg.get("locationPoiname") or ""), "天安门")
+ self.assertEqual(len(payload.get("messages", [])), 7)
+ self.assertIn("location", manifest.get("filters", {}).get("messageTypes") or [])
+ finally:
+ if prev_data is None:
+ os.environ.pop("WECHAT_TOOL_DATA_DIR", None)
+ else:
+ os.environ["WECHAT_TOOL_DATA_DIR"] = prev_data
+
def test_privacy_mode_never_exports_media(self):
with TemporaryDirectory() as td:
root = Path(td)
diff --git a/tests/test_chat_export_panel_frontend.py b/tests/test_chat_export_panel_frontend.py
index c42665e2..c6701ff6 100644
--- a/tests/test_chat_export_panel_frontend.py
+++ b/tests/test_chat_export_panel_frontend.py
@@ -79,6 +79,19 @@ def test_incremental_result_separates_success_repair_and_unavailable_media(self)
self.assertIn("源端暂不可用,重复修复不会改变结果。", dialog)
self.assertIn("查看完整任务说明", dialog)
+ def test_location_type_is_offered_and_folder_result_shows_when_it_is_skipped(self):
+ dialog = (ROOT / "frontend" / "components" / "chat" / "ChatExportDialog.vue").read_text(encoding="utf-8")
+ export_state = (ROOT / "frontend" / "composables" / "chat" / "useChatExport.js").read_text(encoding="utf-8")
+
+ self.assertIn("{ value: 'location', label: '位置' }", export_state)
+ # 目录沿用不含位置的基线时,提示要放在折叠的任务说明之外。
+ followups = dialog.index('class="chat-export-folder-result__followups"')
+ notice = dialog.index("本次未导出位置消息")
+ details = dialog.index('class="chat-export-folder-result__details"')
+ self.assertLess(followups, notice)
+ self.assertLess(notice, details)
+ self.assertIn('v-if="exportJob.incremental?.locationTypeSkipped"', dialog)
+
def test_incremental_baseline_card_uses_compact_status_and_custom_checkbox(self):
dialog = (ROOT / "frontend" / "components" / "chat" / "ChatExportDialog.vue").read_text(encoding="utf-8")
diff --git a/tests/test_chat_incremental_export.py b/tests/test_chat_incremental_export.py
index 6f882507..245b8280 100644
--- a/tests/test_chat_incremental_export.py
+++ b/tests/test_chat_incremental_export.py
@@ -24,6 +24,36 @@ class TestChatIncrementalExport(unittest.TestCase):
_seed_wxid_media_files = _BaseChatExportTest._seed_wxid_media_files
_seed_source_info = _BaseChatExportTest._seed_source_info
+ # 导出面板加入“位置”之前,“全部选择”提交的消息类型。
+ _LEGACY_DIALOG_MESSAGE_TYPES = (
+ "text",
+ "image",
+ "emoji",
+ "video",
+ "voice",
+ "chatHistory",
+ "transfer",
+ "redPacket",
+ "file",
+ "link",
+ "quote",
+ "system",
+ "voip",
+ )
+ _NEW_LOCATION_ROWS = (
+ (8, 1008, 1, 8, 2, 1735689608, "升级后的新消息", None),
+ (
+ 9,
+ 1009,
+ 48,
+ 9,
+ 2,
+ 1735689609,
+ '',
+ None,
+ ),
+ )
+
def _wait_for_job(self, manager, export_id: str):
for _ in range(400):
job = manager.get_job(export_id)
@@ -132,6 +162,17 @@ def _managed_message_file(folder: Path, suffix: str) -> Path:
assert len(matches) == 1
return matches[0]
+ def _insert_message_rows(self, account_dir: Path, username: str, rows) -> None:
+ connection = sqlite3.connect(str(account_dir / "message_0.db"))
+ try:
+ connection.executemany(
+ f"INSERT INTO {self._message_table(username)} VALUES (?, ?, ?, ?, ?, ?, ?, ?)",
+ list(rows),
+ )
+ connection.commit()
+ finally:
+ connection.close()
+
def test_unavailable_pending_media_is_deduplicated_without_repair_prompt(self):
with TemporaryDirectory() as td:
root = Path(td)
@@ -886,6 +927,292 @@ def test_config_conflict_and_privacy_baseline(self):
else:
os.environ["WECHAT_TOOL_DATA_DIR"] = previous
+ def test_folder_without_location_keeps_updating_when_location_is_requested(self):
+ legacy_types = list(self._LEGACY_DIALOG_MESSAGE_TYPES)
+ current_types = [*legacy_types, "location"]
+ with TemporaryDirectory() as td:
+ root = Path(td)
+ account = "wxid_incremental"
+ username = "wxid_friend"
+ account_dir = self._prepare_account(root, account=account, username=username)
+ output_dir = root / "exports"
+
+ previous = os.environ.get("WECHAT_TOOL_DATA_DIR")
+ try:
+ os.environ["WECHAT_TOOL_DATA_DIR"] = str(root)
+ service = self._reload_export_modules()
+ first = self._create_folder_job(
+ service.CHAT_EXPORT_MANAGER,
+ account=account,
+ username=username,
+ output_dir=output_dir,
+ message_types=legacy_types,
+ )
+ self.assertEqual(first.status, "done", msg=first.error)
+ self.assertEqual(first.warning, "")
+ self.assertFalse(first.incremental.get("locationTypeSkipped"))
+ folder = output_dir / "聊天增量测试"
+ state_path = folder / ".wechat-chat-export.json"
+ baseline = json.loads(state_path.read_text(encoding="utf-8"))
+ self.assertNotIn("location", baseline["config"]["messageTypes"])
+ message_file = self._managed_message_file(folder, "json")
+
+ def exported_messages():
+ return json.loads(message_file.read_text(encoding="utf-8")).get("messages", [])
+
+ self.assertEqual(len(exported_messages()), 6)
+
+ self._insert_message_rows(account_dir, username, self._NEW_LOCATION_ROWS)
+ updated = self._create_folder_job(
+ service.CHAT_EXPORT_MANAGER,
+ account=account,
+ username=username,
+ output_dir=output_dir,
+ message_types=current_types,
+ )
+ self.assertEqual(updated.status, "done", msg=updated.error)
+ self.assertEqual(updated.incremental.get("messagesAdded"), 1)
+ self.assertFalse(updated.repair_candidates)
+ self.assertTrue(updated.incremental.get("locationTypeSkipped"))
+ self.assertIn("基线不含“位置”类型", updated.warning)
+ self.assertEqual(updated.options["messageTypes"], legacy_types)
+ messages = exported_messages()
+ self.assertEqual(len(messages), 7)
+ self.assertIn("升级后的新消息", [str(item.get("content") or "") for item in messages])
+ self.assertNotIn("location", [str(item.get("renderType") or "") for item in messages])
+ state = json.loads(state_path.read_text(encoding="utf-8"))
+ self.assertEqual(state["config"], baseline["config"])
+ self.assertEqual(state["configFingerprint"], baseline["configFingerprint"])
+
+ for different_types in (
+ ["text", "location"],
+ [value for value in current_types if value != "voip"],
+ [value for value in legacy_types if value != "voip"],
+ ):
+ with self.subTest(message_types=different_types):
+ with self.assertRaises(service.ChatIncrementalError) as captured:
+ self._create_folder_job(
+ service.CHAT_EXPORT_MANAGER,
+ account=account,
+ username=username,
+ output_dir=output_dir,
+ message_types=different_types,
+ )
+ self.assertEqual(captured.exception.code, "incremental_config_mismatch")
+
+ rebuilt = self._create_folder_job(
+ service.CHAT_EXPORT_MANAGER,
+ account=account,
+ username=username,
+ output_dir=output_dir,
+ message_types=current_types,
+ reset_baseline=True,
+ )
+ self.assertEqual(rebuilt.status, "done", msg=rebuilt.error)
+ self.assertFalse(rebuilt.incremental.get("locationTypeSkipped"))
+ self.assertNotIn("位置", rebuilt.warning)
+ self.assertEqual(rebuilt.options["messageTypes"], current_types)
+ messages = exported_messages()
+ self.assertEqual(len(messages), 9)
+ self.assertEqual(
+ [str(item.get("locationPoiname") or "") for item in messages if item.get("renderType") == "location"],
+ ["天安门", "外滩"],
+ )
+ state = json.loads(state_path.read_text(encoding="utf-8"))
+ self.assertIn("location", state["config"]["messageTypes"])
+
+ # 基线已经包含位置后,不勾选位置就是一次真实的配置变化。
+ with self.assertRaises(service.ChatIncrementalError) as captured:
+ self._create_folder_job(
+ service.CHAT_EXPORT_MANAGER,
+ account=account,
+ username=username,
+ output_dir=output_dir,
+ message_types=legacy_types,
+ )
+ self.assertEqual(captured.exception.code, "incremental_config_mismatch")
+ finally:
+ if previous is None:
+ os.environ.pop("WECHAT_TOOL_DATA_DIR", None)
+ else:
+ os.environ["WECHAT_TOOL_DATA_DIR"] = previous
+
+ def test_html_folder_without_location_appends_only_baseline_types(self):
+ legacy_types = list(self._LEGACY_DIALOG_MESSAGE_TYPES)
+ current_types = [*legacy_types, "location"]
+ with TemporaryDirectory() as td:
+ root = Path(td)
+ account = "wxid_incremental"
+ username = "wxid_friend"
+ account_dir = self._prepare_account(root, account=account, username=username)
+ output_dir = root / "exports"
+
+ previous = os.environ.get("WECHAT_TOOL_DATA_DIR")
+ try:
+ os.environ["WECHAT_TOOL_DATA_DIR"] = str(root)
+ service = self._reload_export_modules()
+ first = self._create_folder_job(
+ service.CHAT_EXPORT_MANAGER,
+ account=account,
+ username=username,
+ output_dir=output_dir,
+ export_format="html",
+ message_types=legacy_types,
+ )
+ self.assertEqual(first.status, "done", msg=first.error)
+ folder = output_dir / "聊天增量测试"
+ state_path = folder / ".wechat-chat-export.json"
+ baseline = json.loads(state_path.read_text(encoding="utf-8"))
+ message_file = self._managed_message_file(folder, "html")
+ # 会话列表的预览不受类型筛选影响,这里只看消息本身的 renderType 标记。
+ location_marker = 'data-render-type="location"'
+ first_html = message_file.read_text(encoding="utf-8")
+ self.assertIn("普通文本消息", first_html)
+ self.assertNotIn(location_marker, first_html)
+
+ original_full_probe = service._probe_incremental_conversation
+
+ def reject_full_probe(**_kwargs):
+ raise AssertionError("沿用基线类型的分页会话不应重新扫描完整历史")
+
+ service._probe_incremental_conversation = reject_full_probe
+ try:
+ # 源数据里已有的位置消息排在水位之后,不能被当成新消息追加。
+ no_change = self._create_folder_job(
+ service.CHAT_EXPORT_MANAGER,
+ account=account,
+ username=username,
+ output_dir=output_dir,
+ export_format="html",
+ message_types=current_types,
+ )
+ self.assertEqual(no_change.status, "done", msg=no_change.error)
+ self.assertEqual(no_change.incremental.get("messagesAdded"), 0)
+ self.assertEqual(no_change.incremental.get("filesChanged"), 0)
+ self.assertTrue(no_change.incremental.get("locationTypeSkipped"))
+
+ self._insert_message_rows(account_dir, username, self._NEW_LOCATION_ROWS)
+ appended = self._create_folder_job(
+ service.CHAT_EXPORT_MANAGER,
+ account=account,
+ username=username,
+ output_dir=output_dir,
+ export_format="html",
+ message_types=current_types,
+ )
+ self.assertEqual(appended.status, "done", msg=appended.error)
+ self.assertEqual(appended.incremental.get("messagesAdded"), 1)
+ self.assertEqual(appended.progress.messages_exported, 1)
+ self.assertFalse(appended.repair_candidates)
+ self.assertTrue(appended.incremental.get("locationTypeSkipped"))
+ self.assertIn("基线不含“位置”类型", appended.warning)
+ current_html = message_file.read_text(encoding="utf-8")
+ self.assertIn("升级后的新消息", current_html)
+ self.assertNotIn("普通文本消息", current_html)
+ self.assertNotIn(location_marker, current_html)
+ state = json.loads(state_path.read_text(encoding="utf-8"))
+ self.assertEqual(state["config"], baseline["config"])
+ self.assertEqual(state["configFingerprint"], baseline["configFingerprint"])
+ finally:
+ service._probe_incremental_conversation = original_full_probe
+ finally:
+ if previous is None:
+ os.environ.pop("WECHAT_TOOL_DATA_DIR", None)
+ else:
+ os.environ["WECHAT_TOOL_DATA_DIR"] = previous
+
+ def test_only_a_missing_location_type_keeps_the_baseline_config(self):
+ import wechat_decrypt_tool.chat_incremental_export as incremental
+
+ # 写进基线的是归一化后的 renderType。
+ legacy_types = sorted(value.lower() for value in self._LEGACY_DIALOG_MESSAGE_TYPES)
+ current_types = [*legacy_types, "location"]
+ partial_types = [value for value in legacy_types if value != "system"]
+
+ def build(message_types, **overrides):
+ values = {
+ "export_format": "html",
+ "start_time": None,
+ "end_time": None,
+ "message_types": message_types,
+ "include_media": True,
+ "media_kinds": ["image", "emoji", "video", "video_thumb", "voice", "file"],
+ "download_remote_media": False,
+ "html_page_size": 1000,
+ "privacy_mode": False,
+ "transcribe_voice": False,
+ }
+ values.update(overrides)
+ return incremental.build_config(**values)
+
+ fingerprint = incremental.config_fingerprint
+ # v2.4.0 起(远程缩略图改为默认关闭之后)、面板还没有“位置”时,
+ # 保持默认选项建立增量目录写进基线的配置指纹。
+ legacy_fingerprint = "cf8c33b421da906e8f144a30c9504867264488b4ab94d373e5e4d52813f681ea"
+ self.assertEqual(fingerprint(build(legacy_types)), legacy_fingerprint)
+
+ with TemporaryDirectory() as td:
+
+ def prepare(config, *, baseline_fingerprint=legacy_fingerprint, reset_baseline=False):
+ # 走浏览器回传基线的路径,不依赖磁盘上的目录。
+ return incremental.prepare_folder_context(
+ account="wxid_incremental",
+ exports_root=Path(td),
+ requested_folder_name="聊天增量测试",
+ config=config,
+ privacy_mode=False,
+ desktop_output=False,
+ supplied_baseline={
+ "schemaVersion": 1,
+ "artifactType": "wechat-chat-incremental-folder",
+ "account": "wxid_incremental",
+ "folderName": "聊天增量测试",
+ "conversationSalt": "salt",
+ "configFingerprint": baseline_fingerprint,
+ "conversations": {},
+ "files": {},
+ },
+ missing_files=[],
+ reset_baseline=reset_baseline,
+ )
+
+ kept = prepare(build(current_types))
+ self.assertTrue(kept.location_type_skipped)
+ self.assertEqual(kept.config, build(legacy_types))
+ self.assertEqual(kept.config_hash, legacy_fingerprint)
+
+ # 旧目录只勾选了部分类型时,重复同样的勾选也会多带一个默认勾选的位置。
+ kept_partial = prepare(
+ build([*partial_types, "location"]),
+ baseline_fingerprint=fingerprint(build(partial_types)),
+ )
+ self.assertTrue(kept_partial.location_type_skipped)
+ self.assertEqual(kept_partial.config, build(partial_types))
+
+ unchanged = prepare(build(legacy_types))
+ self.assertFalse(unchanged.location_type_skipped)
+ self.assertEqual(unchanged.config_hash, legacy_fingerprint)
+
+ reset = prepare(build(current_types), reset_baseline=True)
+ self.assertFalse(reset.location_type_skipped)
+ self.assertEqual(reset.config, build(current_types))
+ self.assertEqual(reset.config_hash, fingerprint(build(current_types)))
+
+ mismatches = {
+ "其它配置不同": (build(current_types, export_format="json"), legacy_fingerprint),
+ "分页不同": (build(current_types, html_page_size=500), legacy_fingerprint),
+ "少勾了其它类型": (build([*partial_types, "location"]), legacy_fingerprint),
+ "多勾了其它类型": (build(current_types), fingerprint(build(partial_types))),
+ # 空清单表示导出全部类型,本来就包含位置。
+ "只勾选位置": (build(["location"]), fingerprint(build([]))),
+ "基线已含位置": (build(legacy_types), fingerprint(build(current_types))),
+ }
+ for label, (config, baseline_fingerprint) in mismatches.items():
+ with self.subTest(label):
+ with self.assertRaises(incremental.ChatIncrementalError) as captured:
+ prepare(config, baseline_fingerprint=baseline_fingerprint)
+ self.assertEqual(captured.exception.code, "incremental_config_mismatch")
+
def test_unselected_conversation_is_preserved_and_missing_file_is_restored(self):
with TemporaryDirectory() as td:
root = Path(td)
diff --git a/tests/test_chat_search_index_targets.py b/tests/test_chat_search_index_targets.py
index a6b86fd8..b1e7fde9 100644
--- a/tests/test_chat_search_index_targets.py
+++ b/tests/test_chat_search_index_targets.py
@@ -283,6 +283,59 @@ def test_auto_index_requires_local_decrypted_databases_and_does_not_call_wcdb(se
self.assertEqual(status["index"]["build"].get("status"), "error")
self.assertIn("No sessions found", str(status["index"]["build"].get("error") or ""))
+ def test_start_build_returns_status_without_deadlock_when_already_building(self):
+ import wechat_decrypt_tool.chat_search_index as idx
+
+ with TemporaryDirectory() as td:
+ account_dir = Path(td) / "wxid_account"
+ account_dir.mkdir(parents=True, exist_ok=True)
+ running_build = {"status": "building", "source": "decrypted", "startedAt": 123}
+ result = {}
+
+ def request_build():
+ try:
+ result["status"] = idx.start_chat_search_index_build(account_dir, source="decrypted")
+ except Exception as exc:
+ result["error"] = exc
+
+ # Private lock/state: a regression leaves the caller stuck on these, not on the module's real lock.
+ with (
+ patch.object(idx, "_BUILD_LOCK", threading.Lock()),
+ patch.object(idx, "_BUILD_STATE", {idx._account_key(account_dir): dict(running_build)}),
+ patch.object(idx, "_build_worker") as build_worker,
+ ):
+ caller = threading.Thread(target=request_build, daemon=True)
+ caller.start()
+ caller.join(timeout=5)
+
+ self.assertFalse(caller.is_alive(), "start_chat_search_index_build deadlocked on _BUILD_LOCK")
+ self.assertIsNone(result.get("error"))
+ self.assertEqual(result["status"]["index"]["build"], running_build)
+ self.assertEqual(idx._BUILD_STATE[idx._account_key(account_dir)], running_build)
+ build_worker.assert_not_called()
+
+ def test_start_build_registers_state_and_starts_one_worker(self):
+ import wechat_decrypt_tool.chat_search_index as idx
+
+ with TemporaryDirectory() as td:
+ account_dir = Path(td) / "wxid_account"
+ account_dir.mkdir(parents=True, exist_ok=True)
+
+ worker_started = threading.Event()
+ with (
+ patch.object(idx, "_BUILD_LOCK", threading.Lock()),
+ patch.object(idx, "_BUILD_STATE", {}),
+ patch.object(idx, "_build_worker", side_effect=lambda *_args: worker_started.set()) as build_worker,
+ ):
+ status = idx.start_chat_search_index_build(account_dir, rebuild=True, source="decrypted")
+ registered = dict(idx._BUILD_STATE[idx._account_key(account_dir)])
+ self.assertTrue(worker_started.wait(5), "build worker thread was not started")
+
+ self.assertEqual(status["index"]["build"].get("status"), "building")
+ self.assertEqual(registered.get("status"), "building")
+ self.assertTrue(registered.get("rebuild"))
+ build_worker.assert_called_once_with(account_dir, True, "decrypted")
+
def test_auto_search_uses_decrypted_index_for_single_character(self):
import wechat_decrypt_tool.chat_search_index as idx
from wechat_decrypt_tool.routers import chat as chat_router
diff --git a/tests/test_contacts_keyword_pinyin.py b/tests/test_contacts_keyword_pinyin.py
new file mode 100644
index 00000000..a2bfd388
--- /dev/null
+++ b/tests/test_contacts_keyword_pinyin.py
@@ -0,0 +1,391 @@
+import importlib.util
+import sqlite3
+import sys
+import unittest
+import uuid
+from pathlib import Path
+from tempfile import TemporaryDirectory
+from types import SimpleNamespace
+from unittest.mock import patch
+
+
+ROOT = Path(__file__).resolve().parents[1]
+sys.path.insert(0, str(ROOT / "src"))
+
+
+def _contact(**fields):
+ # username 用纯数字,避免字母关键词被 username 的普通子串匹配命中。
+ contact = {
+ "username": "10001",
+ "displayName": "",
+ "remark": "",
+ "nickname": "",
+ "alias": "",
+ "region": "",
+ "source": "",
+ "country": "",
+ "province": "",
+ "city": "",
+ }
+ contact.update(fields)
+ return contact
+
+
+# (username, remark, nick_name, sort_timestamp):字面命中的会话最旧,仅拼音命中的会话更新且多于 MCP 默认的 10 条。
+_RANKED_ROWS = [
+ ("wxid_n01", "", "Lily", 20),
+ ("wxid_n02", "HR 小王", "王小明", 10),
+ ("wxid_n03", "", "Bob", 0),
+ *[
+ (f"wxid_p{index:02d}", "", name, 1000 + index)
+ for index, name in enumerate(
+ ("李娜", "丽丽", "黎明", "立华", "林涛", "刘洋", "梁静", "廖凡", "凌风", "连胜", "力宏", "莉莉")
+ )
+ ],
+ ("wxid_q01", "", "胡蓉", 2001),
+ ("wxid_q02", "", "何瑞", 2000),
+]
+_RANKED_PINYIN_LI = [f"wxid_p{index:02d}" for index in range(11, -1, -1)]
+
+
+def _write_account(account_dir, rows):
+ account_dir.mkdir()
+ conn = sqlite3.connect(str(account_dir / "contact.db"))
+ try:
+ conn.execute(
+ """
+ CREATE TABLE contact (
+ username TEXT,
+ remark TEXT,
+ nick_name TEXT,
+ alias TEXT,
+ local_type INTEGER,
+ verify_flag INTEGER,
+ big_head_url TEXT,
+ small_head_url TEXT
+ )
+ """
+ )
+ conn.executemany(
+ "INSERT INTO contact VALUES (?, ?, ?, '', 1, 0, '', '')",
+ [row[:3] for row in rows],
+ )
+ conn.commit()
+ finally:
+ conn.close()
+
+ conn = sqlite3.connect(str(account_dir / "session.db"))
+ try:
+ conn.execute("CREATE TABLE SessionTable (username TEXT, sort_timestamp INTEGER)")
+ conn.executemany("INSERT INTO SessionTable VALUES (?, ?)", [(row[0], row[3]) for row in rows])
+ conn.commit()
+ finally:
+ conn.close()
+
+
+@unittest.skipUnless(importlib.util.find_spec("pypinyin"), "pypinyin is not installed")
+class TestContactsKeywordPinyin(unittest.TestCase):
+ def assertMatches(self, contact, *keywords):
+ from wechat_decrypt_tool.routers.chat_contacts import _matches_keyword
+
+ for keyword in keywords:
+ with self.subTest(keyword=keyword):
+ self.assertTrue(_matches_keyword(contact, keyword))
+
+ def assertNotMatches(self, contact, *keywords):
+ from wechat_decrypt_tool.routers.chat_contacts import _matches_keyword
+
+ for keyword in keywords:
+ with self.subTest(keyword=keyword):
+ self.assertFalse(_matches_keyword(contact, keyword))
+
+ def test_initials_match(self):
+ contact = _contact(displayName="张伟", nickname="张伟")
+ self.assertMatches(contact, "zw", "z", "w", "ZW", " zw ")
+ self.assertNotMatches(contact, "wz", "zz", "zwx", "h", "a", "g", "e")
+
+ def test_full_pinyin_matches_only_at_syllable_boundaries(self):
+ contact = _contact(displayName="张伟", nickname="张伟")
+ self.assertMatches(contact, "zhang", "zhangw", "zhangwei", "zha", "wei", "we", "ZhangWei")
+ self.assertNotMatches(contact, "an", "hang", "ang", "angwei", "gw", "ei", "zhangweii", "weizhang")
+
+ def test_remark_and_nickname_are_both_searchable(self):
+ contact = _contact(displayName="老板", remark="老板", nickname="张伟")
+ self.assertMatches(contact, "lb", "laoban", "zw", "zhangwei")
+
+ only_display_name = _contact(displayName="相亲相爱一家人")
+ self.assertMatches(only_display_name, "xqxa", "yjr", "xiangqin", "qinxiangai", "yijiaren")
+ self.assertNotMatches(only_display_name, "iang", "inxiangai")
+
+ def test_mixed_chinese_and_latin_names(self):
+ self.assertMatches(_contact(displayName="A张伟"), "azw", "azhangwei", "zw", "zhangw")
+ self.assertNotMatches(_contact(displayName="A张伟"), "azhw", "aw", "ab")
+
+ self.assertMatches(_contact(displayName="张伟Bob"), "zwb", "zwbob", "weibob", "zhangweib")
+ self.assertMatches(_contact(displayName="Lily妈妈"), "lilymm", "lilymama", "mm", "mama")
+ self.assertNotMatches(_contact(displayName="Lily妈妈"), "lilyam", "am")
+
+ # 空格、标点、表情、数字不参与拼音匹配,也不会隔断前后的字。
+ self.assertMatches(_contact(displayName="张伟(公司)"), "zwgs", "gs", "gongsi", "weigong")
+ self.assertMatches(_contact(displayName="张 伟🎉"), "zw", "zhangwei")
+ self.assertMatches(_contact(displayName="王5哥"), "wg", "wangge", "5")
+ self.assertMatches(_contact(displayName="A1张伟"), "azw", "a1")
+
+ def test_names_without_chinese_keep_plain_substring_matching(self):
+ # 不含汉字的名称即使带表情、标点或重音字母,也不会因为去掉分隔符而变宽松。
+ for name in ("Li Wei", "Li Wei🌸", "John Smith 🎉", "José García", "Mr.Li🌙"):
+ contact = _contact(displayName=name, nickname=name)
+ self.assertNotMatches(contact, "liwei", "lw", "iw", "hns", "ns", "js", "garca", "sg", "rl")
+ self.assertMatches(_contact(displayName="Li Wei🌸"), "li wei", "wei")
+ self.assertMatches(_contact(displayName="José García"), "garc", "jos")
+
+ def test_polyphone_readings(self):
+ # 词语读音由 pypinyin 按词典消歧:重庆是 chong qing 而不是 zhong qing。
+ chongqing = _contact(displayName="重庆客户")
+ self.assertMatches(chongqing, "cq", "cqkh", "chongqing", "chongqingkehu")
+ self.assertNotMatches(chongqing, "zq", "zhongqing")
+
+ # 多音字姓氏按姓氏读音命中,与列表里的拼音分组一致。
+ self.assertMatches(_contact(displayName="曾国藩"), "zgf", "zengguofan", "zeng")
+ self.assertMatches(_contact(displayName="单田芳"), "stf", "shantianfang")
+
+ # 首字不是姓氏用法时,默认读音仍可命中。
+ self.assertMatches(_contact(displayName="乐乐"), "ll", "lele")
+ self.assertMatches(_contact(displayName="单车少年"), "dcsn", "danche")
+
+ def test_u_umlaut_syllables(self):
+ # pypinyin 把 ü 记作 v:lüe / nüe 同时接受 lue / nue 拼法。
+ self.assertMatches(_contact(displayName="侵略者"), "qlz", "qinlve", "qinlue", "lvezhe", "luezhe")
+ self.assertMatches(_contact(displayName="虐心"), "nx", "nvexin", "nuexin")
+ self.assertMatches(_contact(displayName="曾略"), "zl", "zenglue", "zenglve", "cenglue")
+ # lü / nü 与输入法一致,只接受 lv / nv。
+ self.assertMatches(_contact(displayName="吕布"), "lb", "lv", "lvbu")
+ self.assertNotMatches(_contact(displayName="吕布"), "lu", "lubu")
+
+ def test_non_letter_keywords_are_not_pinyin_matched(self):
+ contact = _contact(displayName="张伟", nickname="张伟")
+ self.assertNotMatches(contact, "zw1", "z w", "zhang wei", "zw_", "z.w", "张w", "张wei", "9", "_")
+ self.assertMatches(contact, "张", "伟", "张伟")
+ self.assertNotMatches(contact, "李", "伟张")
+
+ def test_only_name_fields_are_pinyin_matched(self):
+ contact = _contact(
+ displayName="Bob",
+ nickname="Bob",
+ region="中国大陆·北京",
+ country="中国大陆",
+ province="北京",
+ city="海淀",
+ source="通过扫一扫添加",
+ )
+ self.assertMatches(contact, "北京", "海淀", "扫一扫", "bob")
+ self.assertNotMatches(contact, "bj", "beijing", "hd", "haidian", "sys", "zgdl")
+
+ def test_existing_substring_matching_is_unchanged(self):
+ contact = _contact(
+ username="wxid_Abc01",
+ displayName="三哥",
+ remark="三哥",
+ nickname="Zhang San",
+ alias="ZS_Alias",
+ region="中国大陆·四川·成都",
+ source="通过搜索微信号添加",
+ )
+ self.assertMatches(
+ contact,
+ "",
+ None,
+ " ",
+ "wxid_abc",
+ "ABC01",
+ "三",
+ "zhang san",
+ "ang s",
+ "zs_alias",
+ "s_al",
+ "成都",
+ "微信号",
+ )
+ # 纯 ASCII 名称仍只做普通子串匹配,不会因拼音逻辑变宽松。
+ self.assertNotMatches(contact, "zhangsan", "zs1", "wxid_san", "重庆")
+
+ def test_pinyin_conversion_is_cached_per_name(self):
+ from wechat_decrypt_tool.routers import chat_contacts
+
+ # 每次运行使用新的名称,保证断言不受进程内已有缓存影响。
+ suffix = uuid.uuid4().hex
+ remark, nickname = f"缓存备注{suffix}", f"缓存昵称{suffix}"
+ contact = _contact(displayName=remark, remark=remark, nickname=nickname)
+ with patch.object(chat_contacts, "lazy_pinyin", wraps=chat_contacts.lazy_pinyin) as spy:
+ for keyword in ("hcbz", "huancun", "nicheng", "qq", "zzz"):
+ chat_contacts._matches_keyword(contact, keyword)
+ chat_contacts._matches_keyword(dict(contact), keyword)
+
+ converted = sorted(call.args[0] for call in spy.call_args_list)
+ self.assertEqual(converted, sorted([remark, nickname]))
+
+ def test_pinyin_conversion_is_skipped_when_it_cannot_help(self):
+ from wechat_decrypt_tool.routers import chat_contacts
+
+ contact = _contact(displayName="免转换专用名", remark="免转换专用名", nickname="Plain Ascii", alias="mzh_alias")
+ with patch.object(
+ chat_contacts,
+ "lazy_pinyin",
+ side_effect=AssertionError("unexpected pinyin conversion"),
+ ) as spy:
+ # 非 ASCII / 含非字母字符的关键词。
+ self.assertTrue(chat_contacts._matches_keyword(contact, "专用"))
+ self.assertFalse(chat_contacts._matches_keyword(contact, "李"))
+ self.assertFalse(chat_contacts._matches_keyword(contact, "mzh1"))
+ self.assertFalse(chat_contacts._matches_keyword(contact, "m z"))
+ # 普通字段已命中的字母关键词。
+ self.assertTrue(chat_contacts._matches_keyword(contact, "mzh"))
+ self.assertTrue(chat_contacts._matches_keyword(contact, "ascii"))
+ # 名称全是 ASCII 的联系人。
+ ascii_contact = _contact(displayName="Plain Ascii", nickname="Plain Ascii")
+ self.assertFalse(chat_contacts._matches_keyword(ascii_contact, "pa"))
+
+ spy.assert_not_called()
+
+ def test_contacts_api_filters_by_pinyin_keyword(self):
+ from wechat_decrypt_tool.routers import chat_contacts
+
+ with TemporaryDirectory() as td:
+ account_dir = Path(td) / "wxid_account"
+ _write_account(
+ account_dir,
+ [
+ ("wxid_a1", "", "张伟", 0),
+ ("wxid_b2", "重庆客户", "李娜", 0),
+ ("wxid_c3", "", "Bob", 0),
+ ("11@chatroom", "", "相亲相爱一家人", 0),
+ ],
+ )
+
+ def search(keyword):
+ with patch.object(chat_contacts, "_resolve_account_dir", return_value=account_dir):
+ payload = chat_contacts.list_chat_contacts(
+ SimpleNamespace(base_url="http://test/"),
+ account="wxid_account",
+ source="decrypted",
+ keyword=keyword,
+ )
+ self.assertEqual(payload["total"], len(payload["contacts"]))
+ return sorted(item["username"] for item in payload["contacts"])
+
+ self.assertEqual(search(None), ["11@chatroom", "wxid_a1", "wxid_b2", "wxid_c3"])
+ self.assertEqual(search("zw"), ["wxid_a1"])
+ self.assertEqual(search("zhangw"), ["wxid_a1"])
+ self.assertEqual(search("wei"), ["wxid_a1"])
+ self.assertEqual(search("cq"), ["wxid_b2"])
+ self.assertEqual(search("ln"), ["wxid_b2"])
+ self.assertEqual(search("yijiaren"), ["11@chatroom"])
+ self.assertEqual(search("bob"), ["wxid_c3"])
+ self.assertEqual(search("张"), ["wxid_a1"])
+ self.assertEqual(search("an"), [])
+ self.assertEqual(search("hang"), [])
+ self.assertEqual(search("zw1"), [])
+
+ def test_literal_matches_are_listed_before_pinyin_only_matches(self):
+ from wechat_decrypt_tool.routers import chat_contacts
+
+ with TemporaryDirectory() as td:
+ account_dir = Path(td) / "wxid_account"
+ _write_account(account_dir, _RANKED_ROWS)
+
+ def search(keyword):
+ with patch.object(chat_contacts, "_resolve_account_dir", return_value=account_dir):
+ payload = chat_contacts.list_chat_contacts(
+ SimpleNamespace(base_url="http://test/"),
+ account="wxid_account",
+ source="decrypted",
+ keyword=keyword,
+ )
+ return [item["username"] for item in payload["contacts"]]
+
+ # 字面命中的联系人会话再旧也排在仅拼音命中的之前,其余仍按最近会话排序。
+ self.assertEqual(search("li"), ["wxid_n01", *_RANKED_PINYIN_LI])
+ self.assertEqual(search("hr"), ["wxid_n02", "wxid_q01", "wxid_q02"])
+ # 没有拼音命中参与时,顺序与原来一致(只按最近会话)。
+ self.assertEqual(search("小王"), ["wxid_n02"])
+ self.assertEqual(
+ search(None),
+ ["wxid_q01", "wxid_q02", *_RANKED_PINYIN_LI, "wxid_n01", "wxid_n02", "wxid_n03"],
+ )
+
+ def test_realtime_literal_matches_are_listed_before_pinyin_only_matches(self):
+ from wechat_decrypt_tool.routers import chat_contacts
+
+ class DummyLock:
+ def __enter__(self):
+ return self
+
+ def __exit__(self, exc_type, exc, tb):
+ return False
+
+ contact_rows = [
+ {"username": username, "remark": remark, "nick_name": nick_name, "local_type": 1, "flag": 0}
+ for username, remark, nick_name, _ in _RANKED_ROWS
+ ]
+ sessions = [{"username": row[0], "sort_timestamp": row[3]} for row in _RANKED_ROWS]
+
+ def search(keyword):
+ with (
+ patch.object(
+ chat_contacts.WCDB_REALTIME,
+ "ensure_connected",
+ return_value=SimpleNamespace(handle=1, lock=DummyLock()),
+ ),
+ patch.object(chat_contacts, "_wcdb_get_sessions", return_value=sessions),
+ patch.object(chat_contacts, "_query_realtime_contact_rows", return_value=contact_rows),
+ patch.object(chat_contacts, "_query_realtime_official_account_type_map", return_value={}),
+ patch.object(chat_contacts, "_query_realtime_enterprise_group_usernames", return_value=set()),
+ patch.object(chat_contacts, "_wcdb_get_display_names", return_value={}),
+ patch.object(chat_contacts, "_wcdb_get_avatar_urls", return_value={}),
+ ):
+ contacts = chat_contacts._collect_contacts_for_account_realtime(
+ account_dir=Path("account"),
+ base_url="http://test",
+ keyword=keyword,
+ include_friends=True,
+ include_groups=True,
+ include_officials=True,
+ )
+ return [item["username"] for item in contacts]
+
+ self.assertEqual(search("li"), ["wxid_n01", *_RANKED_PINYIN_LI])
+ self.assertEqual(search("hr"), ["wxid_n02", "wxid_q01", "wxid_q02"])
+ self.assertEqual(
+ search(None),
+ ["wxid_q01", "wxid_q02", *_RANKED_PINYIN_LI, "wxid_n01", "wxid_n02", "wxid_n03"],
+ )
+
+ def test_mcp_resolve_contact_keeps_literal_match_ahead_of_pinyin_matches(self):
+ from wechat_decrypt_tool.mcp import tools
+ from wechat_decrypt_tool.mcp.registry import McpToolContext
+ from wechat_decrypt_tool.routers import chat_contacts
+
+ ctx = McpToolContext(request=SimpleNamespace(base_url="http://test/"))
+ with TemporaryDirectory() as td:
+ account_dir = Path(td) / "wxid_account"
+ _write_account(account_dir, _RANKED_ROWS)
+
+ def resolve(query, **extra):
+ args = {"account": "wxid_account", "source": "decrypted", "query": query, **extra}
+ with patch.object(chat_contacts, "_resolve_account_dir", return_value=account_dir):
+ return [item["username"] for item in tools._resolve_contact(args, ctx)["candidates"]]
+
+ # 默认只取 10 条:比字面命中更新的 12 个拼音命中不能把 Lily 挤出结果。
+ usernames = resolve("li")
+ self.assertEqual(len(usernames), 10)
+ self.assertEqual(usernames[0], "wxid_n01")
+ self.assertLessEqual(set(usernames[1:]), set(_RANKED_PINYIN_LI))
+ self.assertEqual(resolve("li", limit=1), ["wxid_n01"])
+
+ usernames = resolve("hr")
+ self.assertEqual(usernames[0], "wxid_n02")
+ self.assertEqual(set(usernames[1:]), {"wxid_q01", "wxid_q02"})
+
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/tests/test_key_file_permissions.py b/tests/test_key_file_permissions.py
new file mode 100644
index 00000000..a13fcd0b
--- /dev/null
+++ b/tests/test_key_file_permissions.py
@@ -0,0 +1,130 @@
+"""数据库密钥与图片密钥文件在 POSIX 上只允许属主读写。"""
+
+import json
+import os
+import stat
+from pathlib import Path
+
+import pytest
+
+from wechat_decrypt_tool import key_store, media_helpers
+
+
+pytestmark = pytest.mark.skipif(os.name == "nt", reason="Windows 没有 POSIX 权限位")
+
+DB_KEY = "ab" * 32
+MEDIA_KEYS = {"xor": 0xA5, "aes": "1234567890abcdef"}
+
+
+def mode(path):
+ return stat.S_IMODE(Path(path).stat().st_mode)
+
+
+def seed(path, permissions=0o644):
+ path.parent.mkdir(parents=True, exist_ok=True)
+ path.write_bytes(b"{}")
+ path.chmod(permissions)
+
+
+def save_db_key():
+ return key_store.upsert_account_keys_in_store("wxid_demo", db_key=DB_KEY, raise_on_write_error=True)
+
+
+def save_media_keys(account_dir):
+ media_helpers._save_media_keys(account_dir, MEDIA_KEYS["xor"], MEDIA_KEYS["aes"].encode("ascii"))
+
+
+@pytest.fixture(autouse=True)
+def default_umask():
+ previous = os.umask(0o022)
+ try:
+ yield
+ finally:
+ os.umask(previous)
+
+
+@pytest.fixture
+def store_path(tmp_path, monkeypatch):
+ path = tmp_path / "output" / "account_keys.json"
+ monkeypatch.setattr(key_store, "_KEY_STORE_PATH", path)
+ return path
+
+
+@pytest.fixture
+def modes_at_write(monkeypatch):
+ """记录密钥内容写入那一刻文件已有的权限;文件还不存在时记 None。"""
+ seen = []
+ write_text = Path.write_text
+
+ def recording_write_text(self, *args, **kwargs):
+ seen.append((self.name, mode(self) if self.exists() else None))
+ return write_text(self, *args, **kwargs)
+
+ monkeypatch.setattr(Path, "write_text", recording_write_text)
+ return seen
+
+
+def test_account_keys_store_is_created_owner_only(store_path):
+ save_db_key()
+
+ assert mode(store_path) == 0o600
+
+
+def test_account_keys_store_overwrite_narrows_existing_mode(store_path):
+ seed(store_path)
+
+ save_db_key()
+
+ assert mode(store_path) == 0o600
+ assert json.loads(store_path.read_text(encoding="utf-8"))["wxid_demo"]["db_key"] == DB_KEY
+
+
+@pytest.mark.parametrize("stale_temp_file", [False, True], ids=["new", "stale-0644-temp-file"])
+def test_account_keys_are_only_written_to_an_owner_only_file(store_path, modes_at_write, stale_temp_file):
+ if stale_temp_file:
+ # 上次异常退出留下的临时文件权限更宽,也必须在写入密钥前收紧。
+ seed(store_path.with_name("account_keys.json.tmp"))
+
+ save_db_key()
+
+ assert modes_at_write == [("account_keys.json.tmp", 0o600)]
+ assert mode(store_path) == 0o600
+
+
+def test_media_keys_file_is_created_owner_only(tmp_path):
+ save_media_keys(tmp_path)
+
+ assert mode(tmp_path / "_media_keys.json") == 0o600
+
+
+def test_media_keys_overwrite_narrows_existing_mode(tmp_path):
+ path = tmp_path / "_media_keys.json"
+ seed(path)
+
+ save_media_keys(tmp_path)
+
+ assert mode(path) == 0o600
+ assert json.loads(path.read_text(encoding="utf-8")) == MEDIA_KEYS
+
+
+@pytest.mark.parametrize("existing", [False, True], ids=["new", "existing-0644"])
+def test_media_keys_are_only_written_to_an_owner_only_file(tmp_path, modes_at_write, existing):
+ if existing:
+ seed(tmp_path / "_media_keys.json")
+
+ save_media_keys(tmp_path)
+
+ assert modes_at_write == [("_media_keys.json", 0o600)]
+
+
+def test_keys_are_still_saved_when_chmod_is_refused(store_path, tmp_path, monkeypatch):
+ def refuse(*_args, **_kwargs):
+ raise PermissionError("chmod refused")
+
+ monkeypatch.setattr(os, "chmod", refuse)
+
+ save_db_key()
+ save_media_keys(tmp_path)
+
+ assert json.loads(store_path.read_text(encoding="utf-8"))["wxid_demo"]["db_key"] == DB_KEY
+ assert json.loads((tmp_path / "_media_keys.json").read_text(encoding="utf-8")) == MEDIA_KEYS
diff --git a/tests/test_main_source_bootstrap.py b/tests/test_main_source_bootstrap.py
new file mode 100644
index 00000000..fb17dd78
--- /dev/null
+++ b/tests/test_main_source_bootstrap.py
@@ -0,0 +1,48 @@
+from __future__ import annotations
+
+import os
+import subprocess
+import sys
+from pathlib import Path
+
+
+ROOT = Path(__file__).resolve().parents[1]
+
+
+def test_main_prefers_checkout_src_over_stale_installed_copy(tmp_path: Path) -> None:
+ # 模拟 `uv sync --no-editable` 留在 site-packages 的旧副本。PYTHONPATH 排在
+ # site-packages 与 editable 安装的 .pth 之前,因此 editable 环境下也不会碰巧通过。
+ stale = tmp_path / "stale-site-packages"
+ package = stale / "wechat_decrypt_tool"
+ package.mkdir(parents=True)
+ (package / "__init__.py").write_text(
+ 'raise ImportError("stale installed copy of wechat_decrypt_tool was imported")\n',
+ encoding="utf-8",
+ )
+ # main.py 顶层会 import uvicorn;用空桩代替,测试既不加载也不启动真实服务。
+ (stale / "uvicorn.py").write_text("", encoding="utf-8")
+
+ env = dict(os.environ)
+ inherited = env.get("PYTHONPATH", "")
+ env["PYTHONPATH"] = f"{stale}{os.pathsep}{inherited}" if inherited else str(stale)
+ env["PYTHONIOENCODING"] = "utf-8"
+ env["PYTHONDONTWRITEBYTECODE"] = "1"
+
+ result = subprocess.run(
+ [
+ sys.executable,
+ "-c",
+ "import main, wechat_decrypt_tool; print(wechat_decrypt_tool.__file__)",
+ ],
+ cwd=ROOT,
+ env=env,
+ capture_output=True,
+ text=True,
+ encoding="utf-8",
+ check=False,
+ timeout=120,
+ )
+
+ assert result.returncode == 0, result.stderr
+ resolved = Path(result.stdout.strip().splitlines()[-1]).resolve()
+ assert resolved == (ROOT / "src" / "wechat_decrypt_tool" / "__init__.py").resolve()
diff --git a/tests/test_mcp_router.py b/tests/test_mcp_router.py
index 3ff59e71..b5c04261 100644
--- a/tests/test_mcp_router.py
+++ b/tests/test_mcp_router.py
@@ -85,6 +85,30 @@ class TestMcpRouter(unittest.TestCase):
"wechat.media.download_chat_emoji",
"wechat.media.open_chat_media_folder",
}
+ UNSAFE_INTEGER_ID = "9007199254740993"
+ MEDIA_LINK_KINDS = (
+ "avatar", "chat_image", "chat_emoji", "chat_video_thumb", "chat_video",
+ "chat_voice", "moments_image", "moments_video", "favicon", "proxy_image",
+ )
+ MEDIA_TOOL_ARGUMENTS = {
+ "wechat.moments.get_media_url": {
+ "account", "post_id", "tid", "media_id", "create_time", "width", "height", "total_size",
+ "idx", "post_type", "media_type", "md5", "token", "url", "key",
+ },
+ "wechat.moments.get_remote_video_url": {"account", "url", "token", "key"},
+ "wechat.media.get_chat_emoji_url": {"account", "username", "md5", "emoji_url", "aes_key"},
+ "wechat.media.get_chat_video_thumb_url": {"account", "username", "md5", "file_id", "deep_scan"},
+ "wechat.media.get_chat_video_url": {"account", "username", "md5", "file_id", "deep_scan"},
+ "wechat.media.get_chat_voice_url": {"account", "server_id", "msg_svr_id"},
+ "wechat.mobile.get_media_links": {
+ "account", "kind", "max_items", "username", "md5", "file_id", "server_id", "msg_svr_id",
+ "emoji_url", "aes_key", "post_id", "media_id", "token", "key", "url",
+ },
+ "wechat.mobile.get_message_media_bundle": {
+ "account", "username", "session_id", "server_id", "msg_svr_id", "md5", "file_id",
+ "emoji_url", "aes_key", "url", "link_url",
+ },
+ }
def setUp(self):
self._old_mcp_token = os.environ.get("WECHAT_TOOL_MCP_TOKEN")
@@ -766,6 +790,59 @@ def test_image_tool_exposes_large_image_options_and_defaults(self):
else:
self.assertNotIn("fetch_remote", query)
+ def test_media_tools_advertise_arguments_their_handlers_read(self):
+ client = self._client()
+ tools = {tool["name"]: tool for tool in client.post("/mcp", json=self._rpc("tools/list")).json()["result"]["tools"]}
+
+ def call(name, arguments):
+ result = client.post("/mcp", json=self._rpc("tools/call", {"name": name, "arguments": arguments})).json()["result"]
+ self.assertFalse(result["isError"])
+ return result["structuredContent"]
+
+ for name, expected in self.MEDIA_TOOL_ARGUMENTS.items():
+ schema = tools[name]["inputSchema"]
+ with self.subTest(tool=name):
+ self.assertTrue(schema["additionalProperties"])
+ self.assertEqual(set(schema["properties"]), expected)
+ contexts = [{}, {"md5": "0" * 32}]
+ if name == "wechat.mobile.get_media_links":
+ contexts.append({"kind": "moments_image"})
+ for key, prop in schema["properties"].items():
+ # server_id 会按整数解析,字符串参数也用十进制数字探测。
+ value = {"boolean": True, "integer": 1}.get(prop["type"], self.UNSAFE_INTEGER_ID)
+ with self.subTest(tool=name, argument=key):
+ self.assertTrue(
+ any(call(name, {**context, key: value}) != call(name, context) for context in contexts),
+ f"{name} advertises {key}, but it changes nothing in {contexts}",
+ )
+
+ kind_description = tools["wechat.mobile.get_media_links"]["inputSchema"]["properties"]["kind"]["description"]
+ for kind in self.MEDIA_LINK_KINDS:
+ with self.subTest(kind=kind):
+ self.assertIn(kind, kind_description)
+ resources = call("wechat.mobile.get_media_links", {"kind": kind, "username": "wxid_a", "url": "https://example.com/a"})["resources"]
+ self.assertEqual([item["kind"] for item in resources], [kind])
+
+ def test_media_tools_take_server_ids_as_exact_strings(self):
+ client = self._client()
+ tools = {tool["name"]: tool for tool in client.post("/mcp", json=self._rpc("tools/list")).json()["result"]["tools"]}
+ server_id = self.UNSAFE_INTEGER_ID
+
+ for name in ("wechat.media.get_chat_voice_url", "wechat.mobile.get_media_links", "wechat.mobile.get_message_media_bundle"):
+ properties = tools[name]["inputSchema"]["properties"]
+ for key in ("server_id", "msg_svr_id"):
+ with self.subTest(tool=name, argument=key):
+ self.assertEqual(properties.get(key, {}).get("type"), "string")
+ structured = client.post("/mcp", json=self._rpc(name, {key: server_id})).json()["result"]["structuredContent"]
+ if name == "wechat.media.get_chat_voice_url":
+ voice = structured
+ elif name == "wechat.mobile.get_media_links":
+ voice = next(item for item in structured["resources"] if item["kind"] == "chat_voice")
+ else:
+ self.assertEqual(structured["serverId"], server_id)
+ voice = structured["urls"]["voice"]
+ self.assertEqual(parse_qs(urlsplit(voice["url"]).query)["server_id"], [server_id])
+
def test_completed_mcp_packages_and_mobile_facade_are_listed(self):
client = self._client()
@@ -961,6 +1038,96 @@ def test_mobile_resolve_target_normalizes_candidates(self):
self.assertEqual(structured["best"]["username"], "wxid_friend")
self.assertEqual(structured["best"]["kind"], "contact")
+ def test_resolve_contact_scores_query_case_insensitively(self):
+ client = self._client()
+
+ class FakeContactsRouter:
+ def list_chat_contacts(self, _request, **_kwargs):
+ return {
+ "status": "success",
+ "contacts": [
+ {"username": "wxid_aaron", "displayName": "Aaron", "remark": "", "nickname": "Aaron", "alias": "", "region": "Alice Springs"},
+ {"username": "wxid_friend", "displayName": "Alice", "remark": "Alice", "nickname": "ali", "alias": ""},
+ ],
+ }
+
+ with patch("wechat_decrypt_tool.mcp.tools._contacts_router", return_value=FakeContactsRouter()):
+ for query in ("alice", "Alice", "ALICE"):
+ with self.subTest(query=query):
+ resp = client.post("/mcp", json=self._rpc("wechat.contacts.resolve_contact", {"query": query}))
+ self.assertEqual(resp.status_code, 200)
+ candidates = resp.json()["result"]["structuredContent"]["candidates"]
+ self.assertEqual(
+ [(c["username"], c["confidence"]) for c in candidates],
+ [("wxid_friend", 60), ("wxid_aaron", 20)],
+ )
+
+ def test_mobile_resolve_target_scores_unscored_candidates_by_match(self):
+ client = self._client()
+ sns_users = []
+
+ class FakeContactsRouter:
+ def list_chat_contacts(self, _request, **_kwargs):
+ return {
+ "status": "success",
+ "contacts": [{"username": "wxid_friend", "displayName": "Alice", "remark": "Alice", "nickname": "ali", "alias": ""}],
+ }
+
+ class FakeChatRouter:
+ def list_chat_sessions(self, _request, **_kwargs):
+ return {"status": "success", "sessions": []}
+
+ class FakeSnsRouter:
+ def list_sns_users(self, **_kwargs):
+ return {"items": list(sns_users), "count": len(sns_users), "limit": 5}
+
+ class FakeBizRouter:
+ def get_biz_account_list(self, **_kwargs):
+ return {"status": "success", "total": 0, "data": []}
+
+ def resolve(arguments):
+ resp = client.post("/mcp", json=self._rpc("wechat.mobile.resolve_target", {"limit": 5, **arguments}))
+ self.assertEqual(resp.status_code, 200)
+ return resp.json()["result"]["structuredContent"]
+
+ with patch("wechat_decrypt_tool.mcp.tools._contacts_router", return_value=FakeContactsRouter()), patch(
+ "wechat_decrypt_tool.mcp.tools._chat_router", return_value=FakeChatRouter()
+ ), patch("wechat_decrypt_tool.mcp.tools._sns_router", return_value=FakeSnsRouter()), patch(
+ "wechat_decrypt_tool.mcp.tools._biz_router", return_value=FakeBizRouter()
+ ):
+ sns_users[:] = [
+ {"username": "wxid_poster", "displayName": "Malice Daily", "postCount": 900},
+ {"username": "wxid_reader", "displayName": "Palace Alice Tea", "postCount": 300},
+ ]
+ for query in ("alice", "Alice"):
+ with self.subTest(query=query):
+ structured = resolve({"query": query})
+ self.assertEqual(structured["warnings"], [])
+ self.assertEqual(
+ [(c["kind"], c["username"], c["confidence"]) for c in structured["candidates"]],
+ [("contact", "wxid_friend", 60), ("moments_user", "wxid_poster", 60), ("moments_user", "wxid_reader", 60)],
+ )
+ self.assertTrue(structured["ambiguous"])
+
+ sns_users[:] = [{"username": "wxid_friend", "displayName": "Alice", "postCount": 3}]
+ for query in ("alice", "wxid_friend"):
+ with self.subTest(same_person=query):
+ structured = resolve({"query": query})
+ self.assertEqual([c["kind"] for c in structured["candidates"]], ["contact", "moments_user"])
+ self.assertEqual(structured["best"]["kind"], "contact")
+ self.assertFalse(structured["ambiguous"])
+
+ sns_users[:] = [
+ {"username": "wxid_friend_fan", "displayName": "Fan", "postCount": 900},
+ {"username": "wxid_friend", "displayName": "Alice", "postCount": 3},
+ ]
+ structured = resolve({"query": "wxid_friend", "target_type": "moments_user"})
+ self.assertEqual(
+ [(c["username"], c["confidence"]) for c in structured["candidates"]],
+ [("wxid_friend", 100), ("wxid_friend_fan", 80)],
+ )
+ self.assertFalse(structured["ambiguous"])
+
def test_mobile_media_links_does_not_fetch_binary_content(self):
client = self._client()
diff --git a/tests/test_media_wxgf_ffmpeg_fallback.py b/tests/test_media_wxgf_ffmpeg_fallback.py
new file mode 100644
index 00000000..a71c3a7b
--- /dev/null
+++ b/tests/test_media_wxgf_ffmpeg_fallback.py
@@ -0,0 +1,360 @@
+"""wxgf(HEVC)图片在没有 WxAM 解码器时(非 Windows)改由 ffmpeg 解码。"""
+
+import asyncio
+import ctypes
+import functools
+import io
+import logging
+import os
+import subprocess
+
+import pytest
+from PIL import Image
+
+from wechat_decrypt_tool import media_helpers
+
+
+JPEG = b"\xff\xd8\xff\xe0" + bytes(16) + b"\xff\xd9"
+VPS = b"\x00\x00\x00\x01\x40\x01"
+# IDR 条带的起始码 + NAL 头,后一个字节的最高位是 first_slice_segment_in_pic_flag。
+IDR = b"\x00\x00\x01\x26\x01\x80"
+# 真实文件在头部和码流之间夹着 ICC 等元数据,里面会出现形似起始码的字节。
+METADATA = b"\x03\x02\x02\x00\x01" + b"\x00\x00\x01\x2a" * 4 + b"\x00\x00\x00\x01\x00\x00"
+# 单分区文件在码流之后还有 24 字节容器数据,里面同样可能出现起始码。
+TAIL = IDR * 4
+FAKE_FFMPEG = "/fake/bin/ffmpeg"
+
+
+def hevc(body=b"picture", pictures=1):
+ return VPS + (IDR + body) * pictures
+
+
+def wxgf(*partitions, tail=TAIL):
+ """按真实文件的布局拼容器:19 字节头、元数据,然后是带 4 字节大端长度前缀的各分区。"""
+ header = b"wxgf\x13\x00\x02\x00\xa0\x00\x78" + bytes(5) + b"\x40\x01\xa2"
+ return header + METADATA + b"".join(len(part).to_bytes(4, "big") + part for part in partitions) + tail
+
+
+class FfmpegSpy:
+ def __init__(self, returncode=0, stdout=JPEG, stderr=b"decode error", error=None, alpha=b"\xff" * 16):
+ self.calls = []
+ self.returncode, self.stdout, self.stderr, self.error, self.alpha = returncode, stdout, stderr, error, alpha
+
+ def __call__(self, command, **kwargs):
+ self.calls.append((command, kwargs))
+ if self.error is not None:
+ raise self.error
+ stdout = self.alpha if "rawvideo" in command else self.stdout
+ return subprocess.CompletedProcess(command, self.returncode, stdout=stdout, stderr=self.stderr)
+
+
+@pytest.fixture(autouse=True)
+def empty_conversion_cache():
+ media_helpers._WXGF_FFMPEG_CACHE.clear()
+ yield
+ media_helpers._WXGF_FFMPEG_CACHE.clear()
+
+
+@pytest.fixture
+def use_ffmpeg(monkeypatch):
+ """默认按非 Windows 处理:没有 WxAM 解码器,ffmpeg 可用但被替身接管。"""
+ monkeypatch.setattr(media_helpers, "_get_wxam_decoder", lambda: None)
+ monkeypatch.setattr(media_helpers, "_find_ffmpeg_executable", lambda: FAKE_FFMPEG)
+
+ def install(**kwargs):
+ spy = FfmpegSpy(**kwargs)
+ monkeypatch.setattr(subprocess, "run", spy)
+ return spy
+
+ return install
+
+
+def test_fallback_pipes_only_the_picture_stream_through_ffmpeg(use_ffmpeg):
+ ffmpeg = use_ffmpeg()
+ picture = hevc(b"main picture" * 8)
+
+ assert media_helpers._wxgf_to_image_bytes(wxgf(picture)) == JPEG
+
+ (command, kwargs), = ffmpeg.calls
+ assert command[0] == FAKE_FFMPEG
+ assert command[command.index("-i") + 1] == "pipe:0" and command[-1] == "pipe:1"
+ assert command[command.index("-frames:v") + 1] == "1"
+ # 解码出错即失败,不输出残缺的画面。
+ assert "-xerror" in command
+ assert command[command.index("-err_detect") + 1] == "explode"
+ assert command.index("-err_detect") < command.index("-i")
+ # 解码尺寸有上限,且作为输入选项放在 -i 之前。
+ assert command[command.index("-max_pixels") + 1] == str(media_helpers._WXGF_FFMPEG_MAX_PIXELS)
+ assert command.index("-max_pixels") < command.index("-i")
+ # 只喂码流本身:不含容器头、元数据、长度前缀和尾部数据。
+ assert kwargs["input"] == picture
+ assert 0 < kwargs["timeout"] <= 60
+ assert kwargs["creationflags"] == (subprocess.CREATE_NO_WINDOW if os.name == "nt" else 0)
+
+
+def test_fallback_returns_none_without_ffmpeg(use_ffmpeg, monkeypatch):
+ ffmpeg = use_ffmpeg()
+ monkeypatch.setattr(media_helpers, "_find_ffmpeg_executable", lambda: "")
+
+ assert media_helpers._wxgf_to_image_bytes(wxgf(hevc())) is None
+ assert ffmpeg.calls == []
+
+
+@pytest.mark.parametrize(
+ "failure",
+ [
+ {"returncode": 1, "stdout": b""},
+ {"returncode": 1},
+ {"stdout": b""},
+ {"stdout": b"not a jpeg"},
+ {"error": subprocess.TimeoutExpired(FAKE_FFMPEG, 1)},
+ {"error": OSError("exec format error")},
+ ],
+ ids=["non-zero-exit", "non-zero-exit-with-output", "no-output", "garbage-output", "timeout", "spawn-error"],
+)
+def test_fallback_never_raises_on_ffmpeg_failure(use_ffmpeg, failure):
+ ffmpeg = use_ffmpeg(**failure)
+
+ assert media_helpers._wxgf_to_image_bytes(wxgf(hevc())) is None
+ assert len(ffmpeg.calls) == 1
+
+
+@pytest.mark.parametrize(
+ "payload",
+ [
+ b"wxgf",
+ b"wxgf\x13" + bytes(range(256)) * 4,
+ b"wxgf\x13" + METADATA + IDR + b"slice only",
+ wxgf(VPS + b"parameter sets only"),
+ # 长度前缀比剩余数据长:文件被截断。
+ wxgf(hevc(b"main picture" * 8))[:-40],
+ # 多帧(动图):单帧 JPEG 表示不了,留给调用方原有的回退。
+ wxgf(hevc(pictures=3)),
+ # 真实文件至多 [alpha, 画面] 两个分区,更多的一律不解码。
+ wxgf(hevc(b"alpha"), hevc(b"alpha"), hevc()),
+ # 一张图片不会有这么多 NAL。
+ wxgf(VPS + b"\x00\x00\x01\x4e\x01" * (media_helpers._WXGF_MAX_NAL_UNITS + 1) + IDR + b"picture"),
+ ],
+ ids=[
+ "header-only",
+ "no-start-code",
+ "no-parameter-sets",
+ "no-picture",
+ "truncated-file",
+ "animation",
+ "too-many-partitions",
+ "too-many-nal-units",
+ ],
+)
+def test_fallback_does_not_spawn_without_a_single_decodable_picture(use_ffmpeg, payload):
+ ffmpeg = use_ffmpeg()
+
+ assert media_helpers._wxgf_to_image_bytes(payload) is None
+ assert ffmpeg.calls == []
+
+
+def test_prefix_scan_does_not_spawn_for_a_stray_wxgf_marker(use_ffmpeg):
+ ffmpeg = use_ffmpeg()
+ payload = bytes(64) + b"wxgf" + bytes(range(256)) * 64
+
+ assert media_helpers._try_strip_media_prefix(payload) == (payload, "application/octet-stream")
+ assert ffmpeg.calls == []
+
+
+def test_fallback_decodes_the_picture_behind_an_opaque_alpha_partition(use_ffmpeg):
+ ffmpeg = use_ffmpeg()
+ # 画面分区 300 字节,它的长度前缀 00 00 01 2c 紧跟在 alpha 分区后面,本身就像一个起始码。
+ alpha, picture = hevc(b"alpha"), hevc(b"p" * 288)
+ assert len(picture) == 300
+
+ assert media_helpers._wxgf_to_image_bytes(wxgf(alpha, picture, tail=b"")) == JPEG
+
+ (alpha_command, alpha_kwargs), (picture_command, picture_kwargs) = ffmpeg.calls
+ assert alpha_command[alpha_command.index("-pix_fmt") + 1] == "gray" and "rawvideo" in alpha_command
+ assert alpha_kwargs["input"] == alpha
+ assert "mjpeg" in picture_command
+ assert picture_kwargs["input"] == picture
+
+
+def test_partitions_share_one_deadline(use_ffmpeg, monkeypatch):
+ ffmpeg = use_ffmpeg()
+ clock = iter([100.0, 100.0, 112.0])
+ monkeypatch.setattr(media_helpers.time, "monotonic", lambda: next(clock))
+
+ assert media_helpers._wxgf_to_image_bytes(wxgf(hevc(b"alpha"), hevc())) == JPEG
+
+ (_, alpha_kwargs), (_, picture_kwargs) = ffmpeg.calls
+ budget = media_helpers._WXGF_FFMPEG_TIMEOUT_SECONDS
+ assert alpha_kwargs["timeout"] == budget
+ # alpha 分区用掉的 12 秒从画面分区的时限里扣除。
+ assert picture_kwargs["timeout"] == budget - 12
+
+
+def test_picture_is_not_decoded_once_the_deadline_has_passed(use_ffmpeg, monkeypatch):
+ ffmpeg = use_ffmpeg()
+ clock = iter([100.0, 100.0, 100.0 + media_helpers._WXGF_FFMPEG_TIMEOUT_SECONDS])
+ monkeypatch.setattr(media_helpers.time, "monotonic", lambda: next(clock))
+
+ assert media_helpers._wxgf_to_image_bytes(wxgf(hevc(b"alpha"), hevc())) is None
+ assert len(ffmpeg.calls) == 1
+
+
+@pytest.mark.parametrize("alpha_plane", [b"\xff" * 15 + b"\x80", b""], ids=["translucent", "undecodable"])
+def test_fallback_declines_a_picture_with_transparency(use_ffmpeg, alpha_plane):
+ ffmpeg = use_ffmpeg(alpha=alpha_plane)
+
+ assert media_helpers._wxgf_to_image_bytes(wxgf(hevc(b"alpha"), hevc())) is None
+ # 画面分区不再解码:结果只会是丢了透明度的 JPEG。
+ assert len(ffmpeg.calls) == 1
+
+
+def test_emoji_route_still_fetches_the_remote_gif_for_an_animated_wxgf(use_ffmpeg, monkeypatch, tmp_path):
+ from wechat_decrypt_tool.routers import chat_media
+
+ ffmpeg = use_ffmpeg()
+ local = tmp_path / "sticker.dat"
+ local.write_bytes(wxgf(hevc(pictures=3)))
+ gif = b"GIF89a" + bytes(16) + b"\x3b"
+ monkeypatch.setattr(chat_media, "_resolve_account_dir", lambda _account: tmp_path)
+ monkeypatch.setattr(chat_media, "_resolve_account_wxid_dir", lambda _account_dir: None)
+ monkeypatch.setattr(chat_media, "_resolve_media_path_for_kind", lambda *_args, **_kwargs: local)
+ monkeypatch.setattr(chat_media, "_try_fetch_emoticon_from_remote", lambda _account_dir, _md5: (gif, "image/gif"))
+
+ response = asyncio.run(chat_media.get_chat_emoji(md5="0123456789abcdef0123456789abcdef", account="wxid_demo"))
+
+ assert (response.body, response.media_type) == (gif, "image/gif")
+ assert ffmpeg.calls == []
+
+
+def test_same_payload_is_converted_only_once(use_ffmpeg):
+ ffmpeg = use_ffmpeg()
+ payload = wxgf(hevc())
+
+ assert media_helpers._wxgf_to_image_bytes(payload) == JPEG
+ assert media_helpers._wxgf_to_image_bytes(bytes(payload)) == JPEG
+ assert len(ffmpeg.calls) == 1
+
+
+def test_undecodable_file_spawns_ffmpeg_once_per_read(use_ffmpeg, tmp_path):
+ # 读取路径会从多个入口把同一份 wxgf 交给转换函数,失败结果也要记住。
+ ffmpeg = use_ffmpeg(returncode=1, stdout=b"")
+ path = tmp_path / "0123456789abcdef0123456789abcdef.dat"
+ path.write_bytes(wxgf(hevc()))
+
+ for _ in range(2):
+ data, media_type = media_helpers._read_and_maybe_decrypt_media(path)
+ assert (data, media_type) == (path.read_bytes(), "application/octet-stream")
+ assert len(ffmpeg.calls) == 1
+
+
+def test_conversion_cache_stays_within_its_byte_budget(use_ffmpeg, monkeypatch):
+ use_ffmpeg()
+ monkeypatch.setattr(media_helpers, "_WXGF_FFMPEG_CACHE_BYTES", 2 * len(JPEG))
+
+ for index in range(5):
+ assert media_helpers._wxgf_to_image_bytes(wxgf(hevc(b"picture %d" % index))) == JPEG
+
+ assert len(media_helpers._WXGF_FFMPEG_CACHE) == 2
+
+
+def test_failure_log_keeps_only_the_end_of_ffmpeg_stderr(use_ffmpeg, caplog):
+ use_ffmpeg(returncode=1, stdout=b"", stderr=b"decode error\n" * 20_000 + b"last line")
+
+ with caplog.at_level(logging.WARNING, logger=media_helpers.logger.name):
+ assert media_helpers._wxgf_to_image_bytes(wxgf(hevc())) is None
+
+ (record,) = caplog.records
+ assert "last line" in record.getMessage() and "rc=1" in record.getMessage()
+ assert len(record.getMessage()) < 1000
+
+
+def _wxam_decoder(result, output=JPEG):
+ def decode(_input, _input_size, output_address, output_size, _config):
+ ctypes.memmove(output_address, output, len(output))
+ output_size._obj.value = len(output)
+ return result
+
+ return decode
+
+
+def test_wxam_decoder_success_does_not_call_ffmpeg(use_ffmpeg, monkeypatch):
+ ffmpeg = use_ffmpeg()
+ monkeypatch.setattr(media_helpers, "_get_wxam_decoder", lambda: _wxam_decoder(0))
+
+ assert media_helpers._wxgf_to_image_bytes(wxgf(hevc())) == JPEG
+ assert ffmpeg.calls == []
+
+
+def test_wxam_decoder_failure_does_not_fall_back_to_ffmpeg(use_ffmpeg, monkeypatch):
+ # Windows 上解码器存在时行为不变:它拒绝的数据不再交给 ffmpeg。
+ ffmpeg = use_ffmpeg()
+ monkeypatch.setattr(media_helpers, "_get_wxam_decoder", lambda: _wxam_decoder(-1))
+
+ assert media_helpers._wxgf_to_image_bytes(wxgf(hevc())) is None
+ assert ffmpeg.calls == []
+
+
+@pytest.fixture
+def real_ffmpeg(monkeypatch):
+ # 绕过模块级缓存做一次真实查找,不改动其他测试会用到的缓存。
+ ffmpeg_exe = media_helpers._find_ffmpeg_executable.__wrapped__()
+ if not ffmpeg_exe:
+ pytest.skip("未找到 ffmpeg(可用 WECHAT_TOOL_FFMPEG 指定)")
+ monkeypatch.setattr(media_helpers, "_get_wxam_decoder", lambda: None)
+ monkeypatch.setattr(media_helpers, "_find_ffmpeg_executable", lambda: ffmpeg_exe)
+ return functools.partial(_encode_hevc, ffmpeg_exe)
+
+
+@functools.lru_cache(maxsize=None)
+def _encode_hevc(ffmpeg_exe, source, frames=1, pix_fmt="yuv420p", lossless=False):
+ proc = subprocess.run(
+ [
+ ffmpeg_exe, "-hide_banner", "-loglevel", "error",
+ "-f", "lavfi", "-i", source,
+ "-frames:v", str(frames), "-pix_fmt", pix_fmt,
+ "-c:v", "libx265", "-x265-params", "log-level=none" + (":lossless=1" if lossless else ""),
+ "-f", "hevc", "pipe:1",
+ ],
+ capture_output=True,
+ timeout=120,
+ )
+ if proc.returncode != 0 or not proc.stdout.startswith(VPS):
+ pytest.skip("ffmpeg 的 libx265 不可用,无法生成 HEVC 测试图")
+ return proc.stdout
+
+
+@pytest.mark.parametrize("with_alpha", [False, True], ids=["opaque", "opaque-alpha-partition"])
+def test_real_ffmpeg_decodes_a_static_picture(real_ffmpeg, with_alpha):
+ picture = real_ffmpeg("testsrc=size=160x120")
+ partitions = (picture,)
+ if with_alpha:
+ # 真实文件先是 alpha 分区,后面才是画面;alpha 是全范围灰度,全不透明时每个像素都是 255。
+ partitions = (real_ffmpeg("color=c=white:size=160x120", pix_fmt="gray", lossless=True), picture)
+
+ converted = media_helpers._wxgf_to_image_bytes(wxgf(*partitions))
+
+ assert converted is not None and converted.startswith(b"\xff\xd8\xff")
+ with Image.open(io.BytesIO(converted)) as image:
+ image.load()
+ assert (image.format, image.size) == ("JPEG", (160, 120))
+
+
+def test_real_ffmpeg_declines_a_picture_with_transparency(real_ffmpeg):
+ picture = real_ffmpeg("testsrc=size=160x120")
+
+ # alpha 分区不是全白,即图片带透明度。
+ assert media_helpers._wxgf_to_image_bytes(wxgf(picture, picture)) is None
+
+
+def test_real_ffmpeg_declines_a_picture_above_the_pixel_cap(real_ffmpeg, monkeypatch):
+ picture = real_ffmpeg("testsrc=size=160x120")
+ monkeypatch.setattr(media_helpers, "_WXGF_FFMPEG_MAX_PIXELS", 160 * 120 - 1)
+
+ assert media_helpers._wxgf_to_image_bytes(wxgf(picture)) is None
+
+
+def test_real_ffmpeg_declines_an_animation(real_ffmpeg):
+ animation = real_ffmpeg("testsrc=size=160x120:rate=5", frames=3)
+
+ assert media_helpers._hevc_picture_count(animation) == 3
+ assert media_helpers._wxgf_to_image_bytes(wxgf(animation)) is None
diff --git a/tests/test_media_xor_decode.py b/tests/test_media_xor_decode.py
new file mode 100644
index 00000000..d0df50c2
--- /dev/null
+++ b/tests/test_media_xor_decode.py
@@ -0,0 +1,97 @@
+"""单字节 XOR 解码改为查表后,结果必须与逐字节实现完全一致。"""
+
+import struct
+
+import pytest
+from Crypto.Cipher import AES
+from Crypto.Util import Padding
+
+from wechat_decrypt_tool import media_helpers
+
+
+PNG = b"\x89PNG\r\n\x1a\n" + bytes(range(256)) + b"\x00\x00\x00\x00IEND\xaeB`\x82"
+BUFFERS = {
+ "empty": b"",
+ "one": b"\x5a",
+ "short": b"wxgf\x00\xff",
+ "all-bytes": bytes(range(256)),
+ "long": bytes((i * 131 + 7) & 0xFF for i in range(70_001)),
+}
+
+
+def naive_xor(data, key):
+ return bytes(b ^ key for b in data)
+
+
+class NoIterBytes(bytes):
+ """一旦被 Python 层逐字节迭代就失败,用来守住整文件解码不回退成解释器循环。"""
+
+ def __iter__(self):
+ raise AssertionError("逐字节 Python 循环")
+
+
+@pytest.mark.parametrize("name", BUFFERS)
+def test_xor_bytes_matches_naive_for_every_key(name):
+ data = BUFFERS[name]
+ # 逐字节的参照实现很慢:256 个 key 由短缓冲区(含全部 256 个字节值)覆盖,长缓冲区只抽几个。
+ for key in (0x00, 0x01, 0x5A, 0xA5, 0xFF) if name == "long" else range(256):
+ assert media_helpers._xor_bytes(data, key) == naive_xor(data, key)
+
+
+@pytest.mark.parametrize("buffer_type", [bytearray, memoryview])
+def test_xor_bytes_accepts_any_buffer_like_naive(buffer_type):
+ data = BUFFERS["all-bytes"]
+
+ decoded = media_helpers._xor_bytes(buffer_type(data), 0xA5)
+
+ assert type(decoded) is bytes and decoded == naive_xor(buffer_type(data), 0xA5)
+
+
+@pytest.mark.parametrize("key", [-1, 256])
+def test_xor_bytes_rejects_out_of_range_key_like_naive(key):
+ with pytest.raises(ValueError):
+ naive_xor(b"\x00", key)
+ with pytest.raises(ValueError):
+ media_helpers._xor_bytes(b"\x00", key)
+
+
+@pytest.mark.parametrize("name", BUFFERS)
+def test_dat_v3_matches_naive(name):
+ data = BUFFERS[name]
+ for key in (0x00, 0x01, 0xA5, 0xFF):
+ assert media_helpers._decrypt_wechat_dat_v3(data, key) == naive_xor(data, key)
+
+
+@pytest.mark.parametrize("tail", [b"", b"\x01", BUFFERS["long"]], ids=["no-tail", "one", "long"])
+def test_dat_v4_xor_tail_matches_naive(tail):
+ aes_key = b"cfcd208495d565ef"
+ head, raw, key = b"aes protected head", b"raw middle", 0xA5
+ encrypted_head = AES.new(aes_key, AES.MODE_ECB).encrypt(Padding.pad(head, AES.block_size))
+ data = (
+ struct.pack("<6sLLx", b"\x07\x08V1\x08\x07", len(head), len(tail))
+ + encrypted_head
+ + raw
+ + naive_xor(tail, key)
+ )
+
+ assert media_helpers._decrypt_wechat_dat_v4(data, key, aes_key) == head + raw + tail
+
+
+@pytest.mark.parametrize("key", [0x00, 0x37, 0xA5, 0xFF])
+def test_magic_guess_decodes_whole_payload(key):
+ assert media_helpers._try_xor_decrypt_by_magic(naive_xor(PNG, key)) == (PNG, "image/png")
+
+
+@pytest.mark.parametrize("key", [0x00, 0x37, 0xA5, 0xFF])
+def test_magic_guess_bruteforce_strips_prefix(key):
+ # 魔数不在固定偏移时走 256 个 key 的预览穷举,再剥掉前缀。
+ data = naive_xor(b"junk!" + PNG, key)
+
+ assert media_helpers._try_xor_decrypt_by_magic(data) == (PNG, "image/png")
+
+
+def test_whole_buffer_decode_does_not_iterate_in_python():
+ data = NoIterBytes(naive_xor(PNG, 0xA5))
+
+ assert media_helpers._decrypt_wechat_dat_v3(data, 0xA5) == PNG
+ assert media_helpers._try_xor_decrypt_by_magic(data) == (PNG, "image/png")
diff --git a/tests/test_sns_media.py b/tests/test_sns_media.py
index 91880cac..f100116d 100644
--- a/tests/test_sns_media.py
+++ b/tests/test_sns_media.py
@@ -864,6 +864,82 @@ def test_fix_sns_cdn_url_non_tencent_host_passthrough(self):
out = sns_media.fix_sns_cdn_url(u, token="tkn", is_video=False)
self.assertEqual(out, u)
+ def test_fix_sns_cdn_url_uses_tls_alias_for_cert_mismatched_hosts(self):
+ # Both hosts are CNAMEs of socwxsns.video.qq.com and serve its *.video.qq.com
+ # certificate, so https only verifies under the CNAME target.
+ for host in ("vweixinthumb.tc.qq.com", "vweixinf.tc.qq.com"):
+ for scheme in ("http", "https"):
+ with self.subTest(host=host, scheme=scheme):
+ out = sns_media.fix_sns_cdn_url(
+ f"{scheme}://{host}/150/20250/snsvideodownload?filekey=abc&bizid=1023",
+ token="tkn",
+ )
+ self.assertEqual(
+ out,
+ "https://socwxsns.video.qq.com/150/20250/snsvideodownload"
+ "?filekey=abc&bizid=1023&token=tkn&idx=1",
+ )
+
+ video = sns_media.fix_sns_cdn_url(
+ "http://vweixinf.tc.qq.com/102/20202/snsvideodownload?filekey=abc&bizid=1023",
+ token="tkn",
+ is_video=True,
+ )
+ self.assertEqual(
+ video,
+ "https://socwxsns.video.qq.com/102/20202/snsvideodownload"
+ "?token=tkn&idx=1&filekey=abc&bizid=1023",
+ )
+ for alias_host in sns_media._SNS_CDN_TLS_HOST_ALIASES.values():
+ self.assertTrue(sns_media.is_allowed_sns_media_host(alias_host), alias_host)
+
+ def test_remote_fetch_uses_tls_alias_host_and_keeps_cache_consistent(self):
+ requests: list[str] = []
+
+ def handler(request: httpx.Request) -> httpx.Response:
+ requests.append(str(request.url.host))
+ if "/102/" in request.url.path:
+ return httpx.Response(200, content=b"\x00\x00\x00\x18ftypmp42", request=request)
+ return httpx.Response(200, content=b"\xff\xd8\xff\x00jpeg", request=request)
+
+ async def run(account_dir: Path):
+ async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client:
+ images = [
+ await sns_media.try_fetch_and_decrypt_sns_image_remote(
+ account_dir=account_dir,
+ url="http://vweixinthumb.tc.qq.com/150/20250/snsvideodownload?filekey=abc",
+ key="",
+ token="thumb-token",
+ use_cache=True,
+ client=client,
+ )
+ for _ in range(2)
+ ]
+ video = await sns_media.materialize_sns_remote_video(
+ account_dir=account_dir,
+ url="http://vweixinf.tc.qq.com/102/20202/snsvideodownload?filekey=abc",
+ key="",
+ token="video-token",
+ use_cache=True,
+ client=client,
+ )
+ return images, video
+
+ with TemporaryDirectory() as td:
+ account_dir = Path(td)
+ images, video = asyncio.run(run(account_dir))
+ cached_video = sns_media.get_cached_sns_remote_video(
+ account_dir=account_dir,
+ url="http://vweixinf.tc.qq.com/102/20202/snsvideodownload?filekey=abc",
+ key="",
+ token="rotated-token",
+ )
+
+ self.assertEqual(requests, ["socwxsns.video.qq.com", "socwxsns.video.qq.com"])
+ self.assertEqual([image.source for image in images], ["remote", "remote-cache"])
+ self.assertIsNotNone(video)
+ self.assertEqual(cached_video, video)
+
def test_cdn_capture_keeps_thumbnail_and_original_credentials_paired(self):
requests: list[str] = []
diff --git a/tests/test_wechat_detection_auto_detect.py b/tests/test_wechat_detection_auto_detect.py
index 271db937..e253e676 100644
--- a/tests/test_wechat_detection_auto_detect.py
+++ b/tests/test_wechat_detection_auto_detect.py
@@ -207,6 +207,47 @@ def test_macos_detects_accounts_nested_under_xwechat_files(self):
self.assertEqual(accounts[0]["data_dir"], str(account_dir))
self.assertEqual(accounts[0]["database_count"], 1)
+ def test_macos_generic_user_roots_are_not_reported_as_data_roots(self):
+ from wechat_decrypt_tool import wechat_detection as wd
+
+ with TemporaryDirectory() as td:
+ home = Path(td) / "home"
+ # 与微信无关、但直接包含 *.db 的目录(例如 ~/.hermes/state.db)
+ for unrelated_dir in (
+ home / ".hermes",
+ home / "Documents" / "notes",
+ home / "Desktop" / "project",
+ home / "Downloads" / "tool",
+ ):
+ unrelated_dir.mkdir(parents=True)
+ (unrelated_dir / "state.db").write_bytes(b"not-wechat")
+
+ container_root = home / "Library" / "Containers" / "com.tencent.xinWeChat" / "Data"
+ xwechat_root = container_root / "Documents" / "xwechat_files"
+ container_db_storage = xwechat_root / "wxid_demo_abcd" / "db_storage"
+ container_db_storage.mkdir(parents=True)
+ (container_db_storage / "contact.db").write_bytes(b"demo")
+
+ # 通用目录下名称匹配的子目录仍然要能识别
+ copied_root = home / "Documents" / "WeChat Files"
+ copied_db_storage = copied_root / "wxid_copied" / "db_storage"
+ copied_db_storage.mkdir(parents=True)
+ (copied_db_storage / "contact.db").write_bytes(b"demo")
+
+ with (
+ patch.object(wd.sys, "platform", "darwin"),
+ patch.object(wd.Path, "home", return_value=home),
+ patch.object(wd, "get_process_list", return_value=[]),
+ ):
+ detected_dirs = wd.auto_detect_wechat_data_dirs()
+ accounts = wd.detect_wechat_accounts_from_data_root()
+
+ self.assertEqual(detected_dirs, [str(copied_root), str(xwechat_root)])
+ self.assertEqual(
+ sorted(item["account_name"] for item in accounts),
+ ["wxid_copied", "wxid_demo_abcd"],
+ )
+
def test_xwechat_config_ini_real_path_returns_data_root(self):
import hashlib