commit 96ccbbdc3dc73138f90edb26918b37aec117d69d Author: pandorafuture <157389297+pandorafuture@users.noreply.github.com> Date: Thu Jun 4 22:45:34 2026 +0800 feat: initial release of wx-cli v0.7.2 WeChat macOS database decryption and query tool, supporting key extraction, message querying, full-text search, conversation export, real-time monitoring, media handling, and HTTP server mode with REST API and SSE. diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..0592392 --- /dev/null +++ b/.gitignore @@ -0,0 +1,2 @@ +/target +.DS_Store diff --git a/CHANGELOG.md b/CHANGELOG.md new file mode 100644 index 0000000..f3515a2 --- /dev/null +++ b/CHANGELOG.md @@ -0,0 +1,21 @@ +# Changelog + +All notable changes to this project will be documented in this file. + +## [0.7.2] - 2026-04-06 + +### Features + +- **Decrypt WeChat databases** — Automatically decrypt WeChat macOS (4.1.7.x / 4.1.8.x) encrypted databases +- **Extract encryption keys** — Two methods available: `key extract` (recommended, uses LLDB) and `key scan` (memory scan, requires sudo) +- **Browse contacts** — Search and view your WeChat contacts with details like phone, signature, region, labels, and memo +- **Browse conversations** — List recent conversations with unread counts and last message preview +- **Query messages** — Filter messages by contact, group, date range, or message type, with pagination support +- **Full-text search** — Search across all conversations by keyword, with automatic index building +- **Export conversations** — Export chats to TXT or JSON, including images, voice messages, videos, and file attachments +- **Real-time monitoring** — Watch for new messages as they arrive with `watch` command +- **Media handling** — Decrypt images, convert WeChat voice messages to standard audio, decode WeChat-format images (WXGF), and decrypt video channel videos +- **HTTP server mode** — Run as a local HTTP service with REST API and real-time event stream (SSE) for integration with other apps +- **Privacy filtering** — Hide specific contacts or group members from query and server results +- **Environment check** — `doctor` command verifies all prerequisites are met +- **Parallel processing** — Large exports process images, voice, and video in parallel for faster completion diff --git a/CHANGELOG.zh-CN.md b/CHANGELOG.zh-CN.md new file mode 100644 index 0000000..fe6f581 --- /dev/null +++ b/CHANGELOG.zh-CN.md @@ -0,0 +1,21 @@ +# 更新日志 + +本文件记录项目的所有重要变更。 + +## [0.7.2] - 2026-04-06 + +### 功能 + +- **解密微信数据库** — 自动解密 macOS 微信(4.1.7.x / 4.1.8.x)的加密数据库 +- **提取加密密钥** — 提供两种方式:`key extract`(推荐,通过 LLDB)和 `key scan`(内存扫描,需要 sudo) +- **浏览联系人** — 搜索和查看微信联系人,支持手机号、签名、地区、标签、备注等详情 +- **浏览会话** — 列出近期会话,显示未读数和最新消息预览 +- **查询消息** — 按联系人、群聊、时间范围或消息类型筛选消息,支持分页浏览 +- **全文搜索** — 按关键词搜索所有会话的聊天记录,自动建立搜索索引 +- **导出会话** — 将聊天记录导出为 TXT 或 JSON 格式,包含图片、语音、视频和文件附件 +- **实时监听** — 通过 `watch` 命令实时监听新消息 +- **媒体处理** — 解密图片、将微信语音转为标准音频格式、解码微信专有图片格式(WXGF)、解密视频号视频 +- **HTTP 服务模式** — 以本地 HTTP 服务运行,提供 REST API 和实时事件推送(SSE),方便与其他应用集成 +- **隐私过滤** — 可在查询和服务结果中隐藏指定联系人或群成员 +- **环境检查** — `doctor` 命令一键检查所有运行前置条件 +- **并行处理** — 大型导出任务并行处理图片、语音和视频,显著提升导出速度 diff --git a/Cargo.lock b/Cargo.lock new file mode 100644 index 0000000..1adf620 --- /dev/null +++ b/Cargo.lock @@ -0,0 +1,3362 @@ +# This file is automatically @generated by Cargo. +# It is not intended for manual editing. +version = 4 + +[[package]] +name = "adler2" +version = "2.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "320119579fcad9c21884f5c4861d16174d0e06250625266f50fe6898340abefa" + +[[package]] +name = "aes" +version = "0.8.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b169f7a6d4742236a0a00c541b845991d0ac43e546831af1249753ab4c3aa3a0" +dependencies = [ + "cfg-if", + "cipher", + "cpufeatures", +] + +[[package]] +name = "ahash" +version = "0.8.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5a15f179cd60c4584b8a8c596927aadc462e27f2ca70c04e0071964a73ba7a75" +dependencies = [ + "cfg-if", + "once_cell", + "version_check", + "zerocopy", +] + +[[package]] +name = "aho-corasick" +version = "1.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ddd31a130427c27518df266943a5308ed92d4b226cc639f5a8f1002816174301" +dependencies = [ + "memchr", +] + +[[package]] +name = "allocator-api2" +version = "0.2.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "683d7910e743518b0e34f1186f92494becacb047c7b6bf616c96772180fef923" + +[[package]] +name = "android_system_properties" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "819e7219dbd41043ac279b19830f2efc897156490d7fd6ea916720117ee66311" +dependencies = [ + "libc", +] + +[[package]] +name = "ansi_term" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d52a9bb7ec0cf484c551830a7ce27bd20d67eac647e1befb56b0be4ee39a55d2" +dependencies = [ + "winapi", +] + +[[package]] +name = "anstream" +version = "1.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "824a212faf96e9acacdbd09febd34438f8f711fb84e09a8916013cd7815ca28d" +dependencies = [ + "anstyle", + "anstyle-parse", + "anstyle-query", + "anstyle-wincon", + "colorchoice", + "is_terminal_polyfill", + "utf8parse", +] + +[[package]] +name = "anstyle" +version = "1.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "940b3a0ca603d1eade50a4846a2afffd5ef57a9feac2c0e2ec2e14f9ead76000" + +[[package]] +name = "anstyle-parse" +version = "1.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "52ce7f38b242319f7cabaa6813055467063ecdc9d355bbb4ce0c68908cd8130e" +dependencies = [ + "utf8parse", +] + +[[package]] +name = "anstyle-query" +version = "1.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "40c48f72fd53cd289104fc64099abca73db4166ad86ea0b4341abe65af83dadc" +dependencies = [ + "windows-sys 0.61.2", +] + +[[package]] +name = "anstyle-wincon" +version = "3.0.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "291e6a250ff86cd4a820112fb8898808a366d8f9f58ce16d1f538353ad55747d" +dependencies = [ + "anstyle", + "once_cell_polyfill", + "windows-sys 0.61.2", +] + +[[package]] +name = "anyhow" +version = "1.0.102" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7f202df86484c868dbad7eaa557ef785d5c66295e41b460ef922eca0723b842c" + +[[package]] +name = "atomic-waker" +version = "1.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1505bd5d3d116872e7271a6d4e16d81d0c8570876c8de68093a09ac269d8aac0" + +[[package]] +name = "atty" +version = "0.2.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d9b39be18770d11421cdb1b9947a45dd3f37e93092cbf377614828a319d5fee8" +dependencies = [ + "hermit-abi", + "libc", + "winapi", +] + +[[package]] +name = "autocfg" +version = "1.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c08606f8c3cbf4ce6ec8e28fb0014a2c086708fe954eaa885384a6165172e7e8" + +[[package]] +name = "axum" +version = "0.8.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b52af3cb4058c895d37317bb27508dccc8e5f2d39454016b297bf4a400597b8" +dependencies = [ + "axum-core", + "axum-macros", + "bytes", + "form_urlencoded", + "futures-util", + "http", + "http-body", + "http-body-util", + "hyper", + "hyper-util", + "itoa", + "matchit", + "memchr", + "mime", + "percent-encoding", + "pin-project-lite", + "serde_core", + "serde_json", + "serde_path_to_error", + "serde_urlencoded", + "sync_wrapper", + "tokio", + "tower", + "tower-layer", + "tower-service", + "tracing", +] + +[[package]] +name = "axum-core" +version = "0.5.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "08c78f31d7b1291f7ee735c1c6780ccde7785daae9a9206026862dab7d8792d1" +dependencies = [ + "bytes", + "futures-core", + "http", + "http-body", + "http-body-util", + "mime", + "pin-project-lite", + "sync_wrapper", + "tower-layer", + "tower-service", + "tracing", +] + +[[package]] +name = "axum-macros" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "604fde5e028fea851ce1d8570bbdc034bec850d157f7569d10f347d06808c05c" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "base64" +version = "0.22.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "72b3254f16251a8381aa12e40e3c4d2f0199f8c6508fbecb9d91f575e0fbb8c6" + +[[package]] +name = "bindgen" +version = "0.59.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2bd2a9a458e8f4304c52c43ebb0cfbd520289f8379a52e329a38afda99bf8eb8" +dependencies = [ + "bitflags 1.3.2", + "cexpr", + "clang-sys", + "clap 2.34.0", + "env_logger", + "lazy_static", + "lazycell", + "log", + "peeking_take_while", + "proc-macro2", + "quote", + "regex", + "rustc-hash", + "shlex", + "which", +] + +[[package]] +name = "bitflags" +version = "1.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bef38d45163c2f1dde094a7dfd33ccf595c92905c8f8f4fdc18d06fb1037718a" + +[[package]] +name = "bitflags" +version = "2.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "843867be96c8daad0d758b57df9392b6d8d271134fce549de6ce169ff98a92af" + +[[package]] +name = "block-buffer" +version = "0.10.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3078c7629b62d3f0439517fa394996acacc5cbc91c5a20d8c658e77abd503a71" +dependencies = [ + "generic-array", +] + +[[package]] +name = "block-padding" +version = "0.3.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a8894febbff9f758034a5b8e12d87918f56dfc64a8e1fe757d65e29041538d93" +dependencies = [ + "generic-array", +] + +[[package]] +name = "bumpalo" +version = "3.20.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5d20789868f4b01b2f2caec9f5c4e0213b41e3e5702a50157d699ae31ced2fcb" + +[[package]] +name = "bytes" +version = "1.11.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1e748733b7cbc798e1434b6ac524f0c1ff2ab456fe201501e6497c8417a4fc33" + +[[package]] +name = "cbc" +version = "0.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "26b52a9543ae338f279b96b0b9fed9c8093744685043739079ce85cd58f289a6" +dependencies = [ + "cipher", +] + +[[package]] +name = "cc" +version = "1.2.57" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7a0dd1ca384932ff3641c8718a02769f1698e7563dc6974ffd03346116310423" +dependencies = [ + "find-msvc-tools", + "jobserver", + "libc", + "shlex", +] + +[[package]] +name = "cexpr" +version = "0.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6fac387a98bb7c37292057cffc56d62ecb629900026402633ae9160df93a8766" +dependencies = [ + "nom", +] + +[[package]] +name = "cfg-if" +version = "1.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9330f8b2ff13f34540b44e946ef35111825727b38d33286ef986142615121801" + +[[package]] +name = "chrono" +version = "0.4.44" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c673075a2e0e5f4a1dde27ce9dee1ea4558c7ffe648f576438a20ca1d2acc4b0" +dependencies = [ + "iana-time-zone", + "js-sys", + "num-traits", + "serde", + "wasm-bindgen", + "windows-link", +] + +[[package]] +name = "cipher" +version = "0.4.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "773f3b9af64447d2ce9850330c473515014aa235e6a783b02db81ff39e4a3dad" +dependencies = [ + "crypto-common", + "inout", +] + +[[package]] +name = "clang-sys" +version = "1.8.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b023947811758c97c59bf9d1c188fd619ad4718dcaa767947df1cadb14f39f4" +dependencies = [ + "glob", + "libc", + "libloading", +] + +[[package]] +name = "clap" +version = "2.34.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a0610544180c38b88101fecf2dd634b174a62eef6946f84dfc6a7127512b381c" +dependencies = [ + "ansi_term", + "atty", + "bitflags 1.3.2", + "strsim 0.8.0", + "textwrap", + "unicode-width", + "vec_map", +] + +[[package]] +name = "clap" +version = "4.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b193af5b67834b676abd72466a96c1024e6a6ad978a1f484bd90b85c94041351" +dependencies = [ + "clap_builder", + "clap_derive", +] + +[[package]] +name = "clap_builder" +version = "4.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "714a53001bf66416adb0e2ef5ac857140e7dc3a0c48fb28b2f10762fc4b5069f" +dependencies = [ + "anstream", + "anstyle", + "clap_lex", + "strsim 0.11.1", +] + +[[package]] +name = "clap_derive" +version = "4.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1110bd8a634a1ab8cb04345d8d878267d57c3cf1b38d91b71af6686408bbca6a" +dependencies = [ + "heck", + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "clap_lex" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c8d4a3bb8b1e0c1050499d1815f5ab16d04f0959b233085fb31653fbfc9d98f9" + +[[package]] +name = "colorchoice" +version = "1.0.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1d07550c9036bf2ae0c684c4297d503f838287c83c53686d05370d0e139ae570" + +[[package]] +name = "console" +version = "0.15.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "054ccb5b10f9f2cbf51eb355ca1d05c2d279ce1804688d0db74b4733a5aeafd8" +dependencies = [ + "encode_unicode", + "libc", + "once_cell", + "windows-sys 0.59.0", +] + +[[package]] +name = "core-foundation-sys" +version = "0.8.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "773648b94d0e5d620f64f280777445740e61fe701025087ec8b57f45c791888b" + +[[package]] +name = "cpufeatures" +version = "0.2.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "59ed5838eebb26a2bb2e58f6d5b5316989ae9d08bab10e0e6d103e656d1b0280" +dependencies = [ + "libc", +] + +[[package]] +name = "crc32fast" +version = "1.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9481c1c90cbf2ac953f07c8d4a58aa3945c425b7185c9154d67a65e4230da511" +dependencies = [ + "cfg-if", +] + +[[package]] +name = "crossbeam-deque" +version = "0.8.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9dd111b7b7f7d55b72c0a6ae361660ee5853c9af73f70c3c2ef6858b950e2e51" +dependencies = [ + "crossbeam-epoch", + "crossbeam-utils", +] + +[[package]] +name = "crossbeam-epoch" +version = "0.9.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5b82ac4a3c2ca9c3460964f020e1402edd5753411d7737aa39c3714ad1b5420e" +dependencies = [ + "crossbeam-utils", +] + +[[package]] +name = "crossbeam-utils" +version = "0.8.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d0a5c400df2834b80a4c3327b3aad3a4c4cd4de0629063962b03235697506a28" + +[[package]] +name = "crypto-common" +version = "0.1.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "78c8292055d1c1df0cce5d180393dc8cce0abec0a7102adb6c7b1eef6016d60a" +dependencies = [ + "generic-array", + "typenum", +] + +[[package]] +name = "darling" +version = "0.20.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc7f46116c46ff9ab3eb1597a45688b6715c6e628b5c133e288e709a29bcb4ee" +dependencies = [ + "darling_core", + "darling_macro", +] + +[[package]] +name = "darling_core" +version = "0.20.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0d00b9596d185e565c2207a0b01f8bd1a135483d02d9b7b0a54b11da8d53412e" +dependencies = [ + "fnv", + "ident_case", + "proc-macro2", + "quote", + "strsim 0.11.1", + "syn", +] + +[[package]] +name = "darling_macro" +version = "0.20.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc34b93ccb385b40dc71c6fceac4b2ad23662c7eeb248cf10d529b7e055b6ead" +dependencies = [ + "darling_core", + "quote", + "syn", +] + +[[package]] +name = "dashmap" +version = "6.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5041cc499144891f3790297212f32a74fb938e5136a14943f338ef9e0ae276cf" +dependencies = [ + "cfg-if", + "crossbeam-utils", + "hashbrown 0.14.5", + "lock_api", + "once_cell", + "parking_lot_core", +] + +[[package]] +name = "deranged" +version = "0.5.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7cd812cc2bc1d69d4764bd80df88b4317eaef9e773c75226407d9bc0876b211c" +dependencies = [ + "powerfmt", +] + +[[package]] +name = "derive_builder" +version = "0.20.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "507dfb09ea8b7fa618fcf76e953f4f5e192547945816d5358edffe39f6f94947" +dependencies = [ + "derive_builder_macro", +] + +[[package]] +name = "derive_builder_core" +version = "0.20.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2d5bcf7b024d6835cfb3d473887cd966994907effbe9227e8c8219824d06c4e8" +dependencies = [ + "darling", + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "derive_builder_macro" +version = "0.20.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ab63b0e2bf4d5928aff72e83a7dace85d7bba5fe12dcc3c5a572d78caffd3f3c" +dependencies = [ + "derive_builder_core", + "syn", +] + +[[package]] +name = "digest" +version = "0.10.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" +dependencies = [ + "block-buffer", + "crypto-common", + "subtle", +] + +[[package]] +name = "dirs" +version = "6.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c3e8aa94d75141228480295a7d0e7feb620b1a5ad9f12bc40be62411e38cce4e" +dependencies = [ + "dirs-sys", +] + +[[package]] +name = "dirs-sys" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e01a3366d27ee9890022452ee61b2b63a67e6f13f58900b651ff5665f0bb1fab" +dependencies = [ + "libc", + "option-ext", + "redox_users", + "windows-sys 0.61.2", +] + +[[package]] +name = "displaydoc" +version = "0.2.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "97369cbbc041bc366949bc74d34658d6cda5621039731c6310521892a3a20ae0" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "ecb" +version = "0.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1a8bfa975b1aec2145850fcaa1c6fe269a16578c44705a532ae3edc92b8881c7" +dependencies = [ + "cipher", +] + +[[package]] +name = "either" +version = "1.15.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "48c757948c5ede0e46177b7add2e67155f70e33c07fea8284df6576da70b3719" + +[[package]] +name = "encode_unicode" +version = "1.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "34aa73646ffb006b8f5147f3dc182bd4bcb190227ce861fc4a4844bf8e3cb2c0" + +[[package]] +name = "env_logger" +version = "0.9.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a12e6657c4c97ebab115a42dcee77225f7f482cdd841cf7088c657a42e9e00e7" +dependencies = [ + "atty", + "humantime", + "log", + "regex", + "termcolor", +] + +[[package]] +name = "equivalent" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f" + +[[package]] +name = "errno" +version = "0.3.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" +dependencies = [ + "libc", + "windows-sys 0.61.2", +] + +[[package]] +name = "fallible-iterator" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2acce4a10f12dc2fb14a218589d4f1f62ef011b2d0cc4b3cb1bba8e94da14649" + +[[package]] +name = "fallible-streaming-iterator" +version = "0.1.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7360491ce676a36bf9bb3c56c1aa791658183a54d2744120f27285738d90465a" + +[[package]] +name = "fastrand" +version = "2.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "37909eebbb50d72f9059c3b6d82c0463f2ff062c9e95845c43a6c9c0355411be" + +[[package]] +name = "filetime" +version = "0.2.27" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f98844151eee8917efc50bd9e8318cb963ae8b297431495d3f758616ea5c57db" +dependencies = [ + "cfg-if", + "libc", + "libredox", +] + +[[package]] +name = "find-msvc-tools" +version = "0.1.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5baebc0774151f905a1a2cc41989300b1e6fbb29aff0ceffa1064fdd3088d582" + +[[package]] +name = "flate2" +version = "1.1.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "843fba2746e448b37e26a819579957415c8cef339bf08564fe8b7ddbd959573c" +dependencies = [ + "crc32fast", + "miniz_oxide", +] + +[[package]] +name = "fnv" +version = "1.0.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3f9eec918d3f24069decb9af1554cad7c880e2da24a9afd88aca000531ab82c1" + +[[package]] +name = "foldhash" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d9c4f5dac5e15c24eb999c26181a6ca40b39fe946cbe4c263c7209467bc83af2" + +[[package]] +name = "form_urlencoded" +version = "1.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cb4cb245038516f5f85277875cdaa4f7d2c9a0fa0468de06ed190163b1581fcf" +dependencies = [ + "percent-encoding", +] + +[[package]] +name = "fsevent-sys" +version = "4.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "76ee7a02da4d231650c7cea31349b889be2f45ddb3ef3032d2ec8185f6313fd2" +dependencies = [ + "libc", +] + +[[package]] +name = "futures-channel" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "07bbe89c50d7a535e539b8c17bc0b49bdb77747034daa8087407d655f3f7cc1d" +dependencies = [ + "futures-core", +] + +[[package]] +name = "futures-core" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7e3450815272ef58cec6d564423f6e755e25379b217b0bc688e295ba24df6b1d" + +[[package]] +name = "futures-macro" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e835b70203e41293343137df5c0664546da5745f82ec9b84d40be8336958447b" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "futures-sink" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c39754e157331b013978ec91992bde1ac089843443c49cbc7f46150b0fad0893" + +[[package]] +name = "futures-task" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "037711b3d59c33004d3856fbdc83b99d4ff37a24768fa1be9ce3538a1cde4393" + +[[package]] +name = "futures-util" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "389ca41296e6190b48053de0321d02a77f32f8a5d2461dd38762c0593805c6d6" +dependencies = [ + "futures-core", + "futures-macro", + "futures-task", + "pin-project-lite", + "slab", +] + +[[package]] +name = "generic-array" +version = "0.14.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "85649ca51fd72272d7821adaf274ad91c288277713d9c18820d8499a7ff69e9a" +dependencies = [ + "typenum", + "version_check", +] + +[[package]] +name = "getrandom" +version = "0.2.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ff2abc00be7fca6ebc474524697ae276ad847ad0a6b3faa4bcb027e9a4614ad0" +dependencies = [ + "cfg-if", + "libc", + "wasi", +] + +[[package]] +name = "getrandom" +version = "0.3.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "899def5c37c4fd7b2664648c28120ecec138e4d395b459e5ca34f9cce2dd77fd" +dependencies = [ + "cfg-if", + "libc", + "r-efi 5.3.0", + "wasip2", +] + +[[package]] +name = "getrandom" +version = "0.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0de51e6874e94e7bf76d726fc5d13ba782deca734ff60d5bb2fb2607c7406555" +dependencies = [ + "cfg-if", + "libc", + "r-efi 6.0.0", + "wasip2", + "wasip3", +] + +[[package]] +name = "glob" +version = "0.3.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0cc23270f6e1808e30a928bdc84dea0b9b4136a8bc82338574f23baf47bbd280" + +[[package]] +name = "hashbrown" +version = "0.14.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e5274423e17b7c9fc20b6e7e208532f9b19825d82dfd615708b70edd83df41f1" +dependencies = [ + "ahash", +] + +[[package]] +name = "hashbrown" +version = "0.15.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9229cfe53dfd69f0609a49f65461bd93001ea1ef889cd5529dd176593f5338a1" +dependencies = [ + "allocator-api2", + "equivalent", + "foldhash", +] + +[[package]] +name = "hashbrown" +version = "0.16.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "841d1cc9bed7f9236f321df977030373f4a4163ae1a7dbfe1a51a2c1a51d9100" + +[[package]] +name = "hashlink" +version = "0.9.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6ba4ff7128dee98c7dc9794b6a411377e1404dba1c97deb8d1a55297bd25d8af" +dependencies = [ + "hashbrown 0.14.5", +] + +[[package]] +name = "heck" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" + +[[package]] +name = "hermit-abi" +version = "0.1.19" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "62b467343b94ba476dcb2500d242dadbb39557df889310ac77c5d99100aaac33" +dependencies = [ + "libc", +] + +[[package]] +name = "hex" +version = "0.4.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7f24254aa9a54b5c858eaee2f5bccdb46aaf0e486a595ed5fd8f86ba55232a70" + +[[package]] +name = "hmac" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6c49c37c09c17a53d937dfbb742eb3a961d65a994e6bcdcf37e7399d0cc8ab5e" +dependencies = [ + "digest", +] + +[[package]] +name = "home" +version = "0.5.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cc627f471c528ff0c4a49e1d5e60450c8f6461dd6d10ba9dcd3a61d3dff7728d" +dependencies = [ + "windows-sys 0.61.2", +] + +[[package]] +name = "http" +version = "1.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e3ba2a386d7f85a81f119ad7498ebe444d2e22c2af0b86b069416ace48b3311a" +dependencies = [ + "bytes", + "itoa", +] + +[[package]] +name = "http-body" +version = "1.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1efedce1fb8e6913f23e0c92de8e62cd5b772a67e7b3946df930a62566c93184" +dependencies = [ + "bytes", + "http", +] + +[[package]] +name = "http-body-util" +version = "0.1.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b021d93e26becf5dc7e1b75b1bed1fd93124b374ceb73f43d4d4eafec896a64a" +dependencies = [ + "bytes", + "futures-core", + "http", + "http-body", + "pin-project-lite", +] + +[[package]] +name = "http-range-header" +version = "0.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9171a2ea8a68358193d15dd5d70c1c10a2afc3e7e4c5bc92bc9f025cebd7359c" + +[[package]] +name = "httparse" +version = "1.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6dbf3de79e51f3d586ab4cb9d5c3e2c14aa28ed23d180cf89b4df0454a69cc87" + +[[package]] +name = "httpdate" +version = "1.0.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "df3b46402a9d5adb4c86a0cf463f42e19994e3ee891101b1841f30a545cb49a9" + +[[package]] +name = "humantime" +version = "2.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "135b12329e5e3ce057a9f972339ea52bc954fe1e9358ef27f95e89716fbc5424" + +[[package]] +name = "hyper" +version = "1.8.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2ab2d4f250c3d7b1c9fcdff1cece94ea4e2dfbec68614f7b87cb205f24ca9d11" +dependencies = [ + "atomic-waker", + "bytes", + "futures-channel", + "futures-core", + "http", + "http-body", + "httparse", + "httpdate", + "itoa", + "pin-project-lite", + "pin-utils", + "smallvec", + "tokio", +] + +[[package]] +name = "hyper-util" +version = "0.1.20" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "96547c2556ec9d12fb1578c4eaf448b04993e7fb79cbaad930a656880a6bdfa0" +dependencies = [ + "bytes", + "http", + "http-body", + "hyper", + "pin-project-lite", + "tokio", + "tower-service", +] + +[[package]] +name = "iana-time-zone" +version = "0.1.65" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e31bc9ad994ba00e440a8aa5c9ef0ec67d5cb5e5cb0cc7f8b744a35b389cc470" +dependencies = [ + "android_system_properties", + "core-foundation-sys", + "iana-time-zone-haiku", + "js-sys", + "log", + "wasm-bindgen", + "windows-core", +] + +[[package]] +name = "iana-time-zone-haiku" +version = "0.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f31827a206f56af32e590ba56d5d2d085f558508192593743f16b2306495269f" +dependencies = [ + "cc", +] + +[[package]] +name = "icu_collections" +version = "2.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4c6b649701667bbe825c3b7e6388cb521c23d88644678e83c0c4d0a621a34b43" +dependencies = [ + "displaydoc", + "potential_utf", + "yoke", + "zerofrom", + "zerovec", +] + +[[package]] +name = "icu_locale_core" +version = "2.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "edba7861004dd3714265b4db54a3c390e880ab658fec5f7db895fae2046b5bb6" +dependencies = [ + "displaydoc", + "litemap", + "tinystr", + "writeable", + "zerovec", +] + +[[package]] +name = "icu_normalizer" +version = "2.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5f6c8828b67bf8908d82127b2054ea1b4427ff0230ee9141c54251934ab1b599" +dependencies = [ + "icu_collections", + "icu_normalizer_data", + "icu_properties", + "icu_provider", + "smallvec", + "zerovec", +] + +[[package]] +name = "icu_normalizer_data" +version = "2.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7aedcccd01fc5fe81e6b489c15b247b8b0690feb23304303a9e560f37efc560a" + +[[package]] +name = "icu_properties" +version = "2.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "020bfc02fe870ec3a66d93e677ccca0562506e5872c650f893269e08615d74ec" +dependencies = [ + "icu_collections", + "icu_locale_core", + "icu_properties_data", + "icu_provider", + "zerotrie", + "zerovec", +] + +[[package]] +name = "icu_properties_data" +version = "2.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "616c294cf8d725c6afcd8f55abc17c56464ef6211f9ed59cccffe534129c77af" + +[[package]] +name = "icu_provider" +version = "2.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "85962cf0ce02e1e0a629cc34e7ca3e373ce20dda4c4d7294bbd0bf1fdb59e614" +dependencies = [ + "displaydoc", + "icu_locale_core", + "writeable", + "yoke", + "zerofrom", + "zerotrie", + "zerovec", +] + +[[package]] +name = "id-arena" +version = "2.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3d3067d79b975e8844ca9eb072e16b31c3c1c36928edf9c6789548c524d0d954" + +[[package]] +name = "ident_case" +version = "1.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b9e0384b61958566e926dc50660321d12159025e767c18e043daf26b70104c39" + +[[package]] +name = "idna" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3b0875f23caa03898994f6ddc501886a45c7d3d62d04d2d90788d47be1b1e4de" +dependencies = [ + "idna_adapter", + "smallvec", + "utf8_iter", +] + +[[package]] +name = "idna_adapter" +version = "1.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3acae9609540aa318d1bc588455225fb2085b9ed0c4f6bd0d9d5bcd86f1a0344" +dependencies = [ + "icu_normalizer", + "icu_properties", +] + +[[package]] +name = "indexmap" +version = "2.13.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7714e70437a7dc3ac8eb7e6f8df75fd8eb422675fc7678aff7364301092b1017" +dependencies = [ + "equivalent", + "hashbrown 0.16.1", + "serde", + "serde_core", +] + +[[package]] +name = "inotify" +version = "0.11.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bd5b3eaf1a28b758ac0faa5a4254e8ab2705605496f1b1f3fbbc3988ad73d199" +dependencies = [ + "bitflags 2.11.0", + "inotify-sys", + "libc", +] + +[[package]] +name = "inotify-sys" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e05c02b5e89bff3b946cedeca278abc628fe811e604f027c45a8aa3cf793d0eb" +dependencies = [ + "libc", +] + +[[package]] +name = "inout" +version = "0.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "879f10e63c20629ecabbb64a8010319738c66a5cd0c29b02d63d272b03751d01" +dependencies = [ + "block-padding", + "generic-array", +] + +[[package]] +name = "insta" +version = "1.46.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e82db8c87c7f1ccecb34ce0c24399b8a73081427f3c7c50a5d597925356115e4" +dependencies = [ + "console", + "once_cell", + "serde", + "similar", + "tempfile", +] + +[[package]] +name = "is_terminal_polyfill" +version = "1.70.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a6cb138bb79a146c1bd460005623e142ef0181e3d0219cb493e02f7d08a35695" + +[[package]] +name = "itertools" +version = "0.14.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2b192c782037fadd9cfa75548310488aabdbf3d2da73885b31bd0abd03351285" +dependencies = [ + "either", +] + +[[package]] +name = "itoa" +version = "1.0.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "92ecc6618181def0457392ccd0ee51198e065e016d1d527a7ac1b6dc7c1f09d2" + +[[package]] +name = "jobserver" +version = "0.1.34" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9afb3de4395d6b3e67a780b6de64b51c978ecf11cb9a462c66be7d4ca9039d33" +dependencies = [ + "getrandom 0.3.4", + "libc", +] + +[[package]] +name = "js-sys" +version = "0.3.91" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b49715b7073f385ba4bc528e5747d02e66cb39c6146efb66b781f131f0fb399c" +dependencies = [ + "once_cell", + "wasm-bindgen", +] + +[[package]] +name = "kqueue" +version = "1.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "eac30106d7dce88daf4a3fcb4879ea939476d5074a9b7ddd0fb97fa4bed5596a" +dependencies = [ + "kqueue-sys", + "libc", +] + +[[package]] +name = "kqueue-sys" +version = "1.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ed9625ffda8729b85e45cf04090035ac368927b8cebc34898e7c120f52e4838b" +dependencies = [ + "bitflags 1.3.2", + "libc", +] + +[[package]] +name = "lazy_static" +version = "1.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bbd2bcb4c963f2ddae06a2efc7e9f3591312473c50c6685e1f298068316e66fe" + +[[package]] +name = "lazycell" +version = "1.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "830d08ce1d1d941e6b30645f1a0eb5643013d835ce3779a5fc208261dbe10f55" + +[[package]] +name = "leb128fmt" +version = "0.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09edd9e8b54e49e587e4f6295a7d29c3ea94d469cb40ab8ca70b288248a81db2" + +[[package]] +name = "libc" +version = "0.2.183" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b5b646652bf6661599e1da8901b3b9522896f01e736bad5f723fe7a3a27f899d" + +[[package]] +name = "libloading" +version = "0.8.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d7c4b02199fee7c5d21a5ae7d8cfa79a6ef5bb2fc834d6e9058e89c825efdc55" +dependencies = [ + "cfg-if", + "windows-link", +] + +[[package]] +name = "libredox" +version = "0.1.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1744e39d1d6a9948f4f388969627434e31128196de472883b39f148769bfe30a" +dependencies = [ + "bitflags 2.11.0", + "libc", + "plain", + "redox_syscall 0.7.3", +] + +[[package]] +name = "libsqlite3-sys" +version = "0.30.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2e99fb7a497b1e3339bc746195567ed8d3e24945ecd636e3619d20b9de9e9149" +dependencies = [ + "cc", + "pkg-config", + "vcpkg", +] + +[[package]] +name = "linux-raw-sys" +version = "0.4.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d26c52dbd32dccf2d10cac7725f8eae5296885fb5703b261f7d0a0739ec807ab" + +[[package]] +name = "linux-raw-sys" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32a66949e030da00e8c7d4434b251670a91556f4144941d37452769c25d58a53" + +[[package]] +name = "litemap" +version = "0.8.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6373607a59f0be73a39b6fe456b8192fcc3585f602af20751600e974dd455e77" + +[[package]] +name = "lock_api" +version = "0.4.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "224399e74b87b5f3557511d98dff8b14089b3dadafcab6bb93eab67d3aace965" +dependencies = [ + "scopeguard", +] + +[[package]] +name = "log" +version = "0.4.29" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5e5032e24019045c762d3c0f28f5b6b8bbf38563a65908389bf7978758920897" + +[[package]] +name = "lru" +version = "0.12.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "234cf4f4a04dc1f57e24b96cc0cd600cf2af460d4161ac5ecdd0af8e1f3b2a38" +dependencies = [ + "hashbrown 0.15.5", +] + +[[package]] +name = "mach2" +version = "0.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dae608c151f68243f2b000364e1f7b186d9c29845f7d2d85bd31b9ad77ad552b" + +[[package]] +name = "matchers" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d1525a2a28c7f4fa0fc98bb91ae755d1e2d1505079e05539e35bc876b5d65ae9" +dependencies = [ + "regex-automata", +] + +[[package]] +name = "matchit" +version = "0.8.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "47e1ffaa40ddd1f3ed91f717a33c8c0ee23fff369e3aa8772b9605cc1d22f4c3" + +[[package]] +name = "md5" +version = "0.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "490cc448043f947bae3cbee9c203358d62dbee0db12107a74be5c30ccfd09771" + +[[package]] +name = "memchr" +version = "2.8.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f8ca58f447f06ed17d5fc4043ce1b10dd205e060fb3ce5b979b8ed8e59ff3f79" + +[[package]] +name = "mime" +version = "0.3.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6877bb514081ee2a7ff5ef9de3281f14a4dd4bceac4c09388074a6b5df8a139a" + +[[package]] +name = "mime_guess" +version = "2.0.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f7c44f8e672c00fe5308fa235f821cb4198414e1c77935c1ab6948d3fd78550e" +dependencies = [ + "mime", + "unicase", +] + +[[package]] +name = "minimal-lexical" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "68354c5c6bd36d73ff3feceb05efa59b6acb7626617f4962be322a825e61f79a" + +[[package]] +name = "miniz_oxide" +version = "0.8.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1fa76a2c86f704bdb222d66965fb3d63269ce38518b83cb0575fca855ebb6316" +dependencies = [ + "adler2", + "simd-adler32", +] + +[[package]] +name = "mio" +version = "1.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a69bcab0ad47271a0234d9422b131806bf3968021e5dc9328caf2d4cd58557fc" +dependencies = [ + "libc", + "log", + "wasi", + "windows-sys 0.61.2", +] + +[[package]] +name = "nom" +version = "7.1.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d273983c5a657a70a3e8f2a01329822f3b8c8172b73826411a55751e404a0a4a" +dependencies = [ + "memchr", + "minimal-lexical", +] + +[[package]] +name = "notify" +version = "8.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4d3d07927151ff8575b7087f245456e549fea62edf0ec4e565a5ee50c8402bc3" +dependencies = [ + "bitflags 2.11.0", + "fsevent-sys", + "inotify", + "kqueue", + "libc", + "log", + "mio", + "notify-types", + "walkdir", + "windows-sys 0.60.2", +] + +[[package]] +name = "notify-types" +version = "2.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "42b8cfee0e339a0337359f3c88165702ac6e600dc01c0cc9579a92d62b08477a" +dependencies = [ + "bitflags 2.11.0", +] + +[[package]] +name = "nu-ansi-term" +version = "0.50.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7957b9740744892f114936ab4a57b3f487491bbeafaf8083688b16841a4240e5" +dependencies = [ + "windows-sys 0.61.2", +] + +[[package]] +name = "num-conv" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cf97ec579c3c42f953ef76dbf8d55ac91fb219dde70e49aa4a6b7d74e9919050" + +[[package]] +name = "num-traits" +version = "0.2.19" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "071dfc062690e90b734c0b2273ce72ad0ffa95f0c74596bc250dcfd960262841" +dependencies = [ + "autocfg", +] + +[[package]] +name = "num_threads" +version = "0.1.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5c7398b9c8b70908f6371f47ed36737907c87c52af34c268fed0bf0ceb92ead9" +dependencies = [ + "libc", +] + +[[package]] +name = "once_cell" +version = "1.21.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50" + +[[package]] +name = "once_cell_polyfill" +version = "1.70.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "384b8ab6d37215f3c5301a95a4accb5d64aa607f1fcb26a11b5303878451b4fe" + +[[package]] +name = "option-ext" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "04744f49eae99ab78e0d5c0b603ab218f515ea8cfe5a456d7629ad883a3b6e7d" + +[[package]] +name = "parking_lot_core" +version = "0.9.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2621685985a2ebf1c516881c026032ac7deafcda1a2c9b7850dc81e3dfcb64c1" +dependencies = [ + "cfg-if", + "libc", + "redox_syscall 0.5.18", + "smallvec", + "windows-link", +] + +[[package]] +name = "pbkdf2" +version = "0.12.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f8ed6a7761f76e3b9f92dfb0a60a6a6477c61024b775147ff0973a02653abaf2" +dependencies = [ + "digest", + "hmac", + "sha2", +] + +[[package]] +name = "peeking_take_while" +version = "0.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "19b17cddbe7ec3f8bc800887bab5e717348c95ea2ca0b1bf0837fb964dc67099" + +[[package]] +name = "percent-encoding" +version = "2.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9b4f627cb1b25917193a259e49bdad08f671f8d9708acfd5fe0a8c1455d87220" + +[[package]] +name = "pin-project-lite" +version = "0.2.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a89322df9ebe1c1578d689c92318e070967d1042b512afbe49518723f4e6d5cd" + +[[package]] +name = "pin-utils" +version = "0.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b870d8c151b6f2fb93e84a13146138f05d02ed11c7e7c54f8826aaaf7c9f184" + +[[package]] +name = "pkg-config" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7edddbd0b52d732b21ad9a5fab5c704c14cd949e5e9a1ec5929a24fded1b904c" + +[[package]] +name = "plain" +version = "0.2.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b4596b6d070b27117e987119b4dac604f3c58cfb0b191112e24771b2faeac1a6" + +[[package]] +name = "potential_utf" +version = "0.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b73949432f5e2a09657003c25bca5e19a0e9c84f8058ca374f49e0ebe605af77" +dependencies = [ + "zerovec", +] + +[[package]] +name = "powerfmt" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "439ee305def115ba05938db6eb1644ff94165c5ab5e9420d1c1bcedbba909391" + +[[package]] +name = "prettyplease" +version = "0.2.37" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "479ca8adacdd7ce8f1fb39ce9ecccbfe93a3f1344b3d0d97f20bc0196208f62b" +dependencies = [ + "proc-macro2", + "syn", +] + +[[package]] +name = "proc-macro2" +version = "1.0.106" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8fd00f0bb2e90d81d1044c2b32617f68fcb9fa3bb7640c23e9c748e53fb30934" +dependencies = [ + "unicode-ident", +] + +[[package]] +name = "prost" +version = "0.13.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2796faa41db3ec313a31f7624d9286acf277b52de526150b7e69f3debf891ee5" +dependencies = [ + "bytes", + "prost-derive", +] + +[[package]] +name = "prost-derive" +version = "0.13.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8a56d757972c98b346a9b766e3f02746cde6dd1cd1d1d563472929fdd74bec4d" +dependencies = [ + "anyhow", + "itertools", + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "quote" +version = "1.0.45" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "41f2619966050689382d2b44f664f4bc593e129785a36d6ee376ddf37259b924" +dependencies = [ + "proc-macro2", +] + +[[package]] +name = "r-efi" +version = "5.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "69cdb34c158ceb288df11e18b4bd39de994f6657d83847bdffdbd7f346754b0f" + +[[package]] +name = "r-efi" +version = "6.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f8dcc9c7d52a811697d2151c701e0d08956f92b0e24136cf4cf27b57a6a0d9bf" + +[[package]] +name = "rayon" +version = "1.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "368f01d005bf8fd9b1206fb6fa653e6c4a81ceb1466406b81792d87c5677a58f" +dependencies = [ + "either", + "rayon-core", +] + +[[package]] +name = "rayon-core" +version = "1.13.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "22e18b0f0062d30d4230b2e85ff77fdfe4326feb054b9783a3460d8435c8ab91" +dependencies = [ + "crossbeam-deque", + "crossbeam-utils", +] + +[[package]] +name = "redox_syscall" +version = "0.5.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ed2bf2547551a7053d6fdfafda3f938979645c44812fbfcda098faae3f1a362d" +dependencies = [ + "bitflags 2.11.0", +] + +[[package]] +name = "redox_syscall" +version = "0.7.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6ce70a74e890531977d37e532c34d45e9055d2409ed08ddba14529471ed0be16" +dependencies = [ + "bitflags 2.11.0", +] + +[[package]] +name = "redox_users" +version = "0.5.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a4e608c6638b9c18977b00b475ac1f28d14e84b27d8d42f70e0bf1e3dec127ac" +dependencies = [ + "getrandom 0.2.17", + "libredox", + "thiserror 2.0.18", +] + +[[package]] +name = "regex" +version = "1.12.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e10754a14b9137dd7b1e3e5b0493cc9171fdd105e0ab477f51b72e7f3ac0e276" +dependencies = [ + "aho-corasick", + "memchr", + "regex-automata", + "regex-syntax", +] + +[[package]] +name = "regex-automata" +version = "0.4.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6e1dd4122fc1595e8162618945476892eefca7b88c52820e74af6262213cae8f" +dependencies = [ + "aho-corasick", + "memchr", + "regex-syntax", +] + +[[package]] +name = "regex-syntax" +version = "0.8.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dc897dd8d9e8bd1ed8cdad82b5966c3e0ecae09fb1907d58efaa013543185d0a" + +[[package]] +name = "ring" +version = "0.17.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a4689e6c2294d81e88dc6261c768b63bc4fcdb852be6d1352498b114f61383b7" +dependencies = [ + "cc", + "cfg-if", + "getrandom 0.2.17", + "libc", + "untrusted", + "windows-sys 0.52.0", +] + +[[package]] +name = "rusqlite" +version = "0.32.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7753b721174eb8ff87a9a0e799e2d7bc3749323e773db92e0984debb00019d6e" +dependencies = [ + "bitflags 2.11.0", + "fallible-iterator", + "fallible-streaming-iterator", + "hashlink", + "libsqlite3-sys", + "smallvec", +] + +[[package]] +name = "rust-stemmers" +version = "1.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e46a2036019fdb888131db7a4c847a1063a7493f971ed94ea82c67eada63ca54" +dependencies = [ + "serde", + "serde_derive", +] + +[[package]] +name = "rustc-hash" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "08d43f7aa6b08d49f382cde6a7982047c3426db949b1424bc4b7ec9ae12c6ce2" + +[[package]] +name = "rustix" +version = "0.38.44" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fdb5bc1ae2baa591800df16c9ca78619bf65c0488b41b96ccec5d11220d8c154" +dependencies = [ + "bitflags 2.11.0", + "errno", + "libc", + "linux-raw-sys 0.4.15", + "windows-sys 0.59.0", +] + +[[package]] +name = "rustix" +version = "1.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6fe4565b9518b83ef4f91bb47ce29620ca828bd32cb7e408f0062e9930ba190" +dependencies = [ + "bitflags 2.11.0", + "errno", + "libc", + "linux-raw-sys 0.12.1", + "windows-sys 0.61.2", +] + +[[package]] +name = "rustls" +version = "0.23.37" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "758025cb5fccfd3bc2fd74708fd4682be41d99e5dff73c377c0646c6012c73a4" +dependencies = [ + "log", + "once_cell", + "ring", + "rustls-pki-types", + "rustls-webpki", + "subtle", + "zeroize", +] + +[[package]] +name = "rustls-pki-types" +version = "1.14.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "be040f8b0a225e40375822a563fa9524378b9d63112f53e19ffff34df5d33fdd" +dependencies = [ + "zeroize", +] + +[[package]] +name = "rustls-webpki" +version = "0.103.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "df33b2b81ac578cabaf06b89b0631153a3f416b0a886e8a7a1707fb51abbd1ef" +dependencies = [ + "ring", + "rustls-pki-types", + "untrusted", +] + +[[package]] +name = "rustversion" +version = "1.0.22" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b39cdef0fa800fc44525c84ccb54a029961a8215f9619753635a9c0d2538d46d" + +[[package]] +name = "ryu" +version = "1.0.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9774ba4a74de5f7b1c1451ed6cd5285a32eddb5cccb8cc655a4e50009e06477f" + +[[package]] +name = "same-file" +version = "1.0.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "93fc1dc3aaa9bfed95e02e6eadabb4baf7e3078b0bd1b4d7b6b0b68378900502" +dependencies = [ + "winapi-util", +] + +[[package]] +name = "scopeguard" +version = "1.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "94143f37725109f92c262ed2cf5e59bce7498c01bcc1502d7b9afe439a4e9f49" + +[[package]] +name = "semver" +version = "1.0.27" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d767eb0aabc880b29956c35734170f26ed551a859dbd361d140cdbeca61ab1e2" + +[[package]] +name = "serde" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9a8e94ea7f378bd32cbbd37198a4a91436180c5bb472411e48b5ec2e2124ae9e" +dependencies = [ + "serde_core", + "serde_derive", +] + +[[package]] +name = "serde_core" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "41d385c7d4ca58e59fc732af25c3983b67ac852c1a25000afe1175de458b67ad" +dependencies = [ + "serde_derive", +] + +[[package]] +name = "serde_derive" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d540f220d3187173da220f885ab66608367b6574e925011a9353e4badda91d79" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "serde_json" +version = "1.0.149" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "83fc039473c5595ace860d8c4fafa220ff474b3fc6bfdb4293327f1a37e94d86" +dependencies = [ + "itoa", + "memchr", + "serde", + "serde_core", + "zmij", +] + +[[package]] +name = "serde_path_to_error" +version = "0.1.20" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "10a9ff822e371bb5403e391ecd83e182e0e77ba7f6fe0160b795797109d1b457" +dependencies = [ + "itoa", + "serde", + "serde_core", +] + +[[package]] +name = "serde_spanned" +version = "0.6.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bf41e0cfaf7226dca15e8197172c295a782857fcb97fad1808a166870dee75a3" +dependencies = [ + "serde", +] + +[[package]] +name = "serde_urlencoded" +version = "0.7.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d3491c14715ca2294c4d6a88f15e84739788c1d030eed8c110436aafdaa2f3fd" +dependencies = [ + "form_urlencoded", + "itoa", + "ryu", + "serde", +] + +[[package]] +name = "sha2" +version = "0.10.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283" +dependencies = [ + "cfg-if", + "cpufeatures", + "digest", +] + +[[package]] +name = "sharded-slab" +version = "0.1.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f40ca3c46823713e0d4209592e8d6e826aa57e928f09752619fc696c499637f6" +dependencies = [ + "lazy_static", +] + +[[package]] +name = "shlex" +version = "1.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0fda2ff0d084019ba4d7c6f371c95d8fd75ce3524c3cb8fb653a3023f6323e64" + +[[package]] +name = "signal-hook-registry" +version = "1.4.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c4db69cba1110affc0e9f7bcd48bbf87b3f4fc7c61fc9155afd4c469eb3d6c1b" +dependencies = [ + "errno", + "libc", +] + +[[package]] +name = "silk-rs" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "014e6619f35a385ff848570e73a0b8c36b31031e0ee11cac70192b50097b1cfe" +dependencies = [ + "bindgen", + "bytes", + "cc", + "thiserror 1.0.69", +] + +[[package]] +name = "simd-adler32" +version = "0.3.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e320a6c5ad31d271ad523dcf3ad13e2767ad8b1cb8f047f75a8aeaf8da139da2" + +[[package]] +name = "similar" +version = "2.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bbbb5d9659141646ae647b42fe094daf6c6192d1620870b449d9557f748b2daa" + +[[package]] +name = "slab" +version = "0.4.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0c790de23124f9ab44544d7ac05d60440adc586479ce501c1d6d7da3cd8c9cf5" + +[[package]] +name = "smallvec" +version = "1.15.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "67b1b7a3b5fe4f1376887184045fcf45c69e92af734b7aaddc05fb777b6fbd03" + +[[package]] +name = "socket2" +version = "0.6.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3a766e1110788c36f4fa1c2b71b387a7815aa65f88ce0229841826633d93723e" +dependencies = [ + "libc", + "windows-sys 0.61.2", +] + +[[package]] +name = "stable_deref_trait" +version = "1.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6ce2be8dc25455e1f91df71bfa12ad37d7af1092ae736f3a6cd0e37bc7810596" + +[[package]] +name = "strsim" +version = "0.8.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8ea5119cdb4c55b55d432abb513a0429384878c15dde60cc77b1c99de1a95a6a" + +[[package]] +name = "strsim" +version = "0.11.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7da8b5736845d9f2fcb837ea5d9e2628564b3b043a70948a3f0b778838c5fb4f" + +[[package]] +name = "subtle" +version = "2.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "13c2bddecc57b384dee18652358fb23172facb8a2c51ccc10d74c157bdea3292" + +[[package]] +name = "syn" +version = "2.0.117" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e665b8803e7b1d2a727f4023456bbbbe74da67099c585258af0ad9c5013b9b99" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + +[[package]] +name = "sync_wrapper" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0bf256ce5efdfa370213c1dabab5935a12e49f2c58d15e9eac2870d3b4f27263" + +[[package]] +name = "synstructure" +version = "0.13.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "728a70f3dbaf5bab7f0c4b1ac8d7ae5ea60a4b5549c8a5914361c99147a709d2" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "tempfile" +version = "3.27.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32497e9a4c7b38532efcdebeef879707aa9f794296a4f0244f6f69e9bc8574bd" +dependencies = [ + "fastrand", + "getrandom 0.4.2", + "once_cell", + "rustix 1.1.4", + "windows-sys 0.61.2", +] + +[[package]] +name = "termcolor" +version = "1.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "06794f8f6c5c898b3275aebefa6b8a1cb24cd2c6c79397ab15774837a0bc5755" +dependencies = [ + "winapi-util", +] + +[[package]] +name = "textwrap" +version = "0.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d326610f408c7a4eb6f51c37c330e496b08506c9457c9d34287ecc38809fb060" +dependencies = [ + "unicode-width", +] + +[[package]] +name = "thiserror" +version = "1.0.69" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6aaf5339b578ea85b50e080feb250a3e8ae8cfcdff9a461c9ec2904bc923f52" +dependencies = [ + "thiserror-impl 1.0.69", +] + +[[package]] +name = "thiserror" +version = "2.0.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4288b5bcbc7920c07a1149a35cf9590a2aa808e0bc1eafaade0b80947865fbc4" +dependencies = [ + "thiserror-impl 2.0.18", +] + +[[package]] +name = "thiserror-impl" +version = "1.0.69" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4fee6c4efc90059e10f81e6d42c60a18f76588c3d74cb83a0b242a2b6c7504c1" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "thiserror-impl" +version = "2.0.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ebc4ee7f67670e9b64d05fa4253e753e016c6c95ff35b89b7941d6b856dec1d5" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "thread_local" +version = "1.1.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f60246a4944f24f6e018aa17cdeffb7818b76356965d03b07d6a9886e8962185" +dependencies = [ + "cfg-if", +] + +[[package]] +name = "time" +version = "0.3.47" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "743bd48c283afc0388f9b8827b976905fb217ad9e647fae3a379a9283c4def2c" +dependencies = [ + "deranged", + "itoa", + "libc", + "num-conv", + "num_threads", + "powerfmt", + "serde_core", + "time-core", + "time-macros", +] + +[[package]] +name = "time-core" +version = "0.1.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7694e1cfe791f8d31026952abf09c69ca6f6fa4e1a1229e18988f06a04a12dca" + +[[package]] +name = "time-macros" +version = "0.2.27" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2e70e4c5a0e0a8a4823ad65dfe1a6930e4f4d756dcd9dd7939022b5e8c501215" +dependencies = [ + "num-conv", + "time-core", +] + +[[package]] +name = "tinystr" +version = "0.8.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "42d3e9c45c09de15d06dd8acf5f4e0e399e85927b7f00711024eb7ae10fa4869" +dependencies = [ + "displaydoc", + "zerovec", +] + +[[package]] +name = "tokio" +version = "1.50.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "27ad5e34374e03cfffefc301becb44e9dc3c17584f414349ebe29ed26661822d" +dependencies = [ + "bytes", + "libc", + "mio", + "pin-project-lite", + "signal-hook-registry", + "socket2", + "tokio-macros", + "windows-sys 0.61.2", +] + +[[package]] +name = "tokio-macros" +version = "2.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5c55a2eff8b69ce66c84f85e1da1c233edc36ceb85a2058d11b0d6a3c7e7569c" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "tokio-stream" +version = "0.1.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32da49809aab5c3bc678af03902d4ccddea2a87d028d86392a4b1560c6906c70" +dependencies = [ + "futures-core", + "pin-project-lite", + "tokio", + "tokio-util", +] + +[[package]] +name = "tokio-util" +version = "0.7.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9ae9cec805b01e8fc3fd2fe289f89149a9b66dd16786abd8b19cfa7b48cb0098" +dependencies = [ + "bytes", + "futures-core", + "futures-sink", + "futures-util", + "pin-project-lite", + "tokio", +] + +[[package]] +name = "toml" +version = "0.8.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dc1beb996b9d83529a9e75c17a1686767d148d70663143c7854d8b4a09ced362" +dependencies = [ + "serde", + "serde_spanned", + "toml_datetime", + "toml_edit", +] + +[[package]] +name = "toml_datetime" +version = "0.6.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "22cddaf88f4fbc13c51aebbf5f8eceb5c7c5a9da2ac40a13519eb5b0a0e8f11c" +dependencies = [ + "serde", +] + +[[package]] +name = "toml_edit" +version = "0.22.27" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "41fe8c660ae4257887cf66394862d21dbca4a6ddd26f04a3560410406a2f819a" +dependencies = [ + "indexmap", + "serde", + "serde_spanned", + "toml_datetime", + "toml_write", + "winnow", +] + +[[package]] +name = "toml_write" +version = "0.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5d99f8c9a7727884afe522e9bd5edbfc91a3312b36a77b5fb8926e4c31a41801" + +[[package]] +name = "tower" +version = "0.5.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ebe5ef63511595f1344e2d5cfa636d973292adc0eec1f0ad45fae9f0851ab1d4" +dependencies = [ + "futures-core", + "futures-util", + "pin-project-lite", + "sync_wrapper", + "tokio", + "tower-layer", + "tower-service", + "tracing", +] + +[[package]] +name = "tower-http" +version = "0.6.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d4e6559d53cc268e5031cd8429d05415bc4cb4aefc4aa5d6cc35fbf5b924a1f8" +dependencies = [ + "bitflags 2.11.0", + "bytes", + "futures-core", + "futures-util", + "http", + "http-body", + "http-body-util", + "http-range-header", + "httpdate", + "mime", + "mime_guess", + "percent-encoding", + "pin-project-lite", + "tokio", + "tokio-util", + "tower-layer", + "tower-service", + "tracing", +] + +[[package]] +name = "tower-layer" +version = "0.3.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "121c2a6cda46980bb0fcd1647ffaf6cd3fc79a013de288782836f6df9c48780e" + +[[package]] +name = "tower-service" +version = "0.3.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8df9b6e13f2d32c91b9bd719c00d1958837bc7dec474d94952798cc8e69eeec3" + +[[package]] +name = "tracing" +version = "0.1.44" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "63e71662fa4b2a2c3a26f570f037eb95bb1f85397f3cd8076caed2f026a6d100" +dependencies = [ + "log", + "pin-project-lite", + "tracing-attributes", + "tracing-core", +] + +[[package]] +name = "tracing-attributes" +version = "0.1.31" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7490cfa5ec963746568740651ac6781f701c9c5ea257c58e057f3ba8cf69e8da" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "tracing-core" +version = "0.1.36" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "db97caf9d906fbde555dd62fa95ddba9eecfd14cb388e4f491a66d74cd5fb79a" +dependencies = [ + "once_cell", + "valuable", +] + +[[package]] +name = "tracing-log" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ee855f1f400bd0e5c02d150ae5de3840039a3f54b025156404e34c23c03f47c3" +dependencies = [ + "log", + "once_cell", + "tracing-core", +] + +[[package]] +name = "tracing-subscriber" +version = "0.3.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cb7f578e5945fb242538965c2d0b04418d38ec25c79d160cd279bf0731c8d319" +dependencies = [ + "matchers", + "nu-ansi-term", + "once_cell", + "regex-automata", + "sharded-slab", + "smallvec", + "thread_local", + "tracing", + "tracing-core", + "tracing-log", +] + +[[package]] +name = "typenum" +version = "1.19.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "562d481066bde0658276a35467c4af00bdc6ee726305698a55b86e61d7ad82bb" + +[[package]] +name = "unicase" +version = "2.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dbc4bc3a9f746d862c45cb89d705aa10f187bb96c76001afab07a0d35ce60142" + +[[package]] +name = "unicode-ident" +version = "1.0.24" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" + +[[package]] +name = "unicode-width" +version = "0.1.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7dd6e30e90baa6f72411720665d41d89b9a3d039dc45b8faea1ddd07f617f6af" + +[[package]] +name = "unicode-xid" +version = "0.2.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ebc1c04c71510c7f702b52b7c350734c9ff1295c464a03335b00bb84fc54f853" + +[[package]] +name = "untrusted" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8ecb6da28b8a351d773b68d5825ac39017e680750f980f3a1a85cd8dd28a47c1" + +[[package]] +name = "ureq" +version = "2.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "02d1a66277ed75f640d608235660df48c8e3c19f3b4edb6a263315626cc3c01d" +dependencies = [ + "base64", + "flate2", + "log", + "once_cell", + "rustls", + "rustls-pki-types", + "serde", + "serde_json", + "url", + "webpki-roots 0.26.11", +] + +[[package]] +name = "url" +version = "2.5.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ff67a8a4397373c3ef660812acab3268222035010ab8680ec4215f38ba3d0eed" +dependencies = [ + "form_urlencoded", + "idna", + "percent-encoding", + "serde", +] + +[[package]] +name = "utf8_iter" +version = "1.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6c140620e7ffbb22c2dee59cafe6084a59b5ffc27a8859a5f0d494b5d52b6be" + +[[package]] +name = "utf8parse" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "06abde3611657adf66d383f00b093d7faecc7fa57071cce2578660c9f1010821" + +[[package]] +name = "valuable" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ba73ea9cf16a25df0c8caa16c51acb937d5712a8429db78a3ee29d5dcacd3a65" + +[[package]] +name = "vcpkg" +version = "0.2.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "accd4ea62f7bb7a82fe23066fb0957d48ef677f6eeb8215f372f52e48bb32426" + +[[package]] +name = "vec_map" +version = "0.8.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f1bddf1187be692e79c5ffeab891132dfb0f236ed36a43c7ed39f1165ee20191" + +[[package]] +name = "vergen" +version = "9.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b849a1f6d8639e8de261e81ee0fc881e3e3620db1af9f2e0da015d4382ceaf75" +dependencies = [ + "anyhow", + "derive_builder", + "rustversion", + "time", + "vergen-lib", +] + +[[package]] +name = "vergen-gitcl" +version = "9.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "77ff3b5300a085d6bcd8fc96a507f706a28ae3814693236c9b409db71a1d15b9" +dependencies = [ + "anyhow", + "derive_builder", + "rustversion", + "time", + "vergen", + "vergen-lib", +] + +[[package]] +name = "vergen-lib" +version = "9.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b34a29ba7e9c59e62f229ae1932fb1b8fb8a6fdcc99215a641913f5f5a59a569" +dependencies = [ + "anyhow", + "derive_builder", + "rustversion", +] + +[[package]] +name = "version_check" +version = "0.9.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b928f33d975fc6ad9f86c8f283853ad26bdd5b10b7f1542aa2fa15e2289105a" + +[[package]] +name = "walkdir" +version = "2.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "29790946404f91d9c5d06f9874efddea1dc06c5efe94541a7d6863108e3a5e4b" +dependencies = [ + "same-file", + "winapi-util", +] + +[[package]] +name = "wasi" +version = "0.11.1+wasi-snapshot-preview1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ccf3ec651a847eb01de73ccad15eb7d99f80485de043efb2f370cd654f4ea44b" + +[[package]] +name = "wasip2" +version = "1.0.2+wasi-0.2.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9517f9239f02c069db75e65f174b3da828fe5f5b945c4dd26bd25d89c03ebcf5" +dependencies = [ + "wit-bindgen", +] + +[[package]] +name = "wasip3" +version = "0.4.0+wasi-0.3.0-rc-2026-01-06" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5428f8bf88ea5ddc08faddef2ac4a67e390b88186c703ce6dbd955e1c145aca5" +dependencies = [ + "wit-bindgen", +] + +[[package]] +name = "wasm-bindgen" +version = "0.2.114" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6532f9a5c1ece3798cb1c2cfdba640b9b3ba884f5db45973a6f442510a87d38e" +dependencies = [ + "cfg-if", + "once_cell", + "rustversion", + "wasm-bindgen-macro", + "wasm-bindgen-shared", +] + +[[package]] +name = "wasm-bindgen-macro" +version = "0.2.114" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "18a2d50fcf105fb33bb15f00e7a77b772945a2ee45dcf454961fd843e74c18e6" +dependencies = [ + "quote", + "wasm-bindgen-macro-support", +] + +[[package]] +name = "wasm-bindgen-macro-support" +version = "0.2.114" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "03ce4caeaac547cdf713d280eda22a730824dd11e6b8c3ca9e42247b25c631e3" +dependencies = [ + "bumpalo", + "proc-macro2", + "quote", + "syn", + "wasm-bindgen-shared", +] + +[[package]] +name = "wasm-bindgen-shared" +version = "0.2.114" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "75a326b8c223ee17883a4251907455a2431acc2791c98c26279376490c378c16" +dependencies = [ + "unicode-ident", +] + +[[package]] +name = "wasm-encoder" +version = "0.244.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "990065f2fe63003fe337b932cfb5e3b80e0b4d0f5ff650e6985b1048f62c8319" +dependencies = [ + "leb128fmt", + "wasmparser", +] + +[[package]] +name = "wasm-metadata" +version = "0.244.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bb0e353e6a2fbdc176932bbaab493762eb1255a7900fe0fea1a2f96c296cc909" +dependencies = [ + "anyhow", + "indexmap", + "wasm-encoder", + "wasmparser", +] + +[[package]] +name = "wasmparser" +version = "0.244.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "47b807c72e1bac69382b3a6fb3dbe8ea4c0ed87ff5629b8685ae6b9a611028fe" +dependencies = [ + "bitflags 2.11.0", + "hashbrown 0.15.5", + "indexmap", + "semver", +] + +[[package]] +name = "webpki-roots" +version = "0.26.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "521bc38abb08001b01866da9f51eb7c5d647a19260e00054a8c7fd5f9e57f7a9" +dependencies = [ + "webpki-roots 1.0.6", +] + +[[package]] +name = "webpki-roots" +version = "1.0.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "22cfaf3c063993ff62e73cb4311efde4db1efb31ab78a3e5c457939ad5cc0bed" +dependencies = [ + "rustls-pki-types", +] + +[[package]] +name = "which" +version = "4.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "87ba24419a2078cd2b0f2ede2691b6c66d8e47836da3b6db8265ebad47afbfc7" +dependencies = [ + "either", + "home", + "once_cell", + "rustix 0.38.44", +] + +[[package]] +name = "winapi" +version = "0.3.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5c839a674fcd7a98952e593242ea400abe93992746761e38641405d28b00f419" +dependencies = [ + "winapi-i686-pc-windows-gnu", + "winapi-x86_64-pc-windows-gnu", +] + +[[package]] +name = "winapi-i686-pc-windows-gnu" +version = "0.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ac3b87c63620426dd9b991e5ce0329eff545bccbbb34f3be09ff6fb6ab51b7b6" + +[[package]] +name = "winapi-util" +version = "0.1.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22" +dependencies = [ + "windows-sys 0.61.2", +] + +[[package]] +name = "winapi-x86_64-pc-windows-gnu" +version = "0.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "712e227841d057c1ee1cd2fb22fa7e5a5461ae8e48fa2ca79ec42cfc1931183f" + +[[package]] +name = "windows-core" +version = "0.62.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b8e83a14d34d0623b51dce9581199302a221863196a1dde71a7663a4c2be9deb" +dependencies = [ + "windows-implement", + "windows-interface", + "windows-link", + "windows-result", + "windows-strings", +] + +[[package]] +name = "windows-implement" +version = "0.60.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "053e2e040ab57b9dc951b72c264860db7eb3b0200ba345b4e4c3b14f67855ddf" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "windows-interface" +version = "0.59.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3f316c4a2570ba26bbec722032c4099d8c8bc095efccdc15688708623367e358" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "windows-link" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5" + +[[package]] +name = "windows-result" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7781fa89eaf60850ac3d2da7af8e5242a5ea78d1a11c49bf2910bb5a73853eb5" +dependencies = [ + "windows-link", +] + +[[package]] +name = "windows-strings" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7837d08f69c77cf6b07689544538e017c1bfcf57e34b4c0ff58e6c2cd3b37091" +dependencies = [ + "windows-link", +] + +[[package]] +name = "windows-sys" +version = "0.52.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "282be5f36a8ce781fad8c8ae18fa3f9beff57ec1b52cb3de0789201425d9a33d" +dependencies = [ + "windows-targets 0.52.6", +] + +[[package]] +name = "windows-sys" +version = "0.59.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1e38bc4d79ed67fd075bcc251a1c39b32a1776bbe92e5bef1f0bf1f8c531853b" +dependencies = [ + "windows-targets 0.52.6", +] + +[[package]] +name = "windows-sys" +version = "0.60.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f2f500e4d28234f72040990ec9d39e3a6b950f9f22d3dba18416c35882612bcb" +dependencies = [ + "windows-targets 0.53.5", +] + +[[package]] +name = "windows-sys" +version = "0.61.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ae137229bcbd6cdf0f7b80a31df61766145077ddf49416a728b02cb3921ff3fc" +dependencies = [ + "windows-link", +] + +[[package]] +name = "windows-targets" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9b724f72796e036ab90c1021d4780d4d3d648aca59e491e6b98e725b84e99973" +dependencies = [ + "windows_aarch64_gnullvm 0.52.6", + "windows_aarch64_msvc 0.52.6", + "windows_i686_gnu 0.52.6", + "windows_i686_gnullvm 0.52.6", + "windows_i686_msvc 0.52.6", + "windows_x86_64_gnu 0.52.6", + "windows_x86_64_gnullvm 0.52.6", + "windows_x86_64_msvc 0.52.6", +] + +[[package]] +name = "windows-targets" +version = "0.53.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4945f9f551b88e0d65f3db0bc25c33b8acea4d9e41163edf90dcd0b19f9069f3" +dependencies = [ + "windows-link", + "windows_aarch64_gnullvm 0.53.1", + "windows_aarch64_msvc 0.53.1", + "windows_i686_gnu 0.53.1", + "windows_i686_gnullvm 0.53.1", + "windows_i686_msvc 0.53.1", + "windows_x86_64_gnu 0.53.1", + "windows_x86_64_gnullvm 0.53.1", + "windows_x86_64_msvc 0.53.1", +] + +[[package]] +name = "windows_aarch64_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32a4622180e7a0ec044bb555404c800bc9fd9ec262ec147edd5989ccd0c02cd3" + +[[package]] +name = "windows_aarch64_gnullvm" +version = "0.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a9d8416fa8b42f5c947f8482c43e7d89e73a173cead56d044f6a56104a6d1b53" + +[[package]] +name = "windows_aarch64_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09ec2a7bb152e2252b53fa7803150007879548bc709c039df7627cabbd05d469" + +[[package]] +name = "windows_aarch64_msvc" +version = "0.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b9d782e804c2f632e395708e99a94275910eb9100b2114651e04744e9b125006" + +[[package]] +name = "windows_i686_gnu" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8e9b5ad5ab802e97eb8e295ac6720e509ee4c243f69d781394014ebfe8bbfa0b" + +[[package]] +name = "windows_i686_gnu" +version = "0.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "960e6da069d81e09becb0ca57a65220ddff016ff2d6af6a223cf372a506593a3" + +[[package]] +name = "windows_i686_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0eee52d38c090b3caa76c563b86c3a4bd71ef1a819287c19d586d7334ae8ed66" + +[[package]] +name = "windows_i686_gnullvm" +version = "0.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fa7359d10048f68ab8b09fa71c3daccfb0e9b559aed648a8f95469c27057180c" + +[[package]] +name = "windows_i686_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "240948bc05c5e7c6dabba28bf89d89ffce3e303022809e73deaefe4f6ec56c66" + +[[package]] +name = "windows_i686_msvc" +version = "0.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1e7ac75179f18232fe9c285163565a57ef8d3c89254a30685b57d83a38d326c2" + +[[package]] +name = "windows_x86_64_gnu" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "147a5c80aabfbf0c7d901cb5895d1de30ef2907eb21fbbab29ca94c5b08b1a78" + +[[package]] +name = "windows_x86_64_gnu" +version = "0.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9c3842cdd74a865a8066ab39c8a7a473c0778a3f29370b5fd6b4b9aa7df4a499" + +[[package]] +name = "windows_x86_64_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "24d5b23dc417412679681396f2b49f3de8c1473deb516bd34410872eff51ed0d" + +[[package]] +name = "windows_x86_64_gnullvm" +version = "0.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0ffa179e2d07eee8ad8f57493436566c7cc30ac536a3379fdf008f47f6bb7ae1" + +[[package]] +name = "windows_x86_64_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec" + +[[package]] +name = "windows_x86_64_msvc" +version = "0.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d6bbff5f0aada427a1e5a6da5f1f98158182f26556f345ac9e04d36d0ebed650" + +[[package]] +name = "winnow" +version = "0.7.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "df79d97927682d2fd8adb29682d1140b343be4ac0f08fd68b7765d9c059d3945" +dependencies = [ + "memchr", +] + +[[package]] +name = "wit-bindgen" +version = "0.51.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d7249219f66ced02969388cf2bb044a09756a083d0fab1e566056b04d9fbcaa5" +dependencies = [ + "wit-bindgen-rust-macro", +] + +[[package]] +name = "wit-bindgen-core" +version = "0.51.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ea61de684c3ea68cb082b7a88508a8b27fcc8b797d738bfc99a82facf1d752dc" +dependencies = [ + "anyhow", + "heck", + "wit-parser", +] + +[[package]] +name = "wit-bindgen-rust" +version = "0.51.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b7c566e0f4b284dd6561c786d9cb0142da491f46a9fbed79ea69cdad5db17f21" +dependencies = [ + "anyhow", + "heck", + "indexmap", + "prettyplease", + "syn", + "wasm-metadata", + "wit-bindgen-core", + "wit-component", +] + +[[package]] +name = "wit-bindgen-rust-macro" +version = "0.51.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0c0f9bfd77e6a48eccf51359e3ae77140a7f50b1e2ebfe62422d8afdaffab17a" +dependencies = [ + "anyhow", + "prettyplease", + "proc-macro2", + "quote", + "syn", + "wit-bindgen-core", + "wit-bindgen-rust", +] + +[[package]] +name = "wit-component" +version = "0.244.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9d66ea20e9553b30172b5e831994e35fbde2d165325bec84fc43dbf6f4eb9cb2" +dependencies = [ + "anyhow", + "bitflags 2.11.0", + "indexmap", + "log", + "serde", + "serde_derive", + "serde_json", + "wasm-encoder", + "wasm-metadata", + "wasmparser", + "wit-parser", +] + +[[package]] +name = "wit-parser" +version = "0.244.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ecc8ac4bc1dc3381b7f59c34f00b67e18f910c2c0f50015669dde7def656a736" +dependencies = [ + "anyhow", + "id-arena", + "indexmap", + "log", + "semver", + "serde", + "serde_derive", + "serde_json", + "unicode-xid", + "wasmparser", +] + +[[package]] +name = "writeable" +version = "0.6.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9edde0db4769d2dc68579893f2306b26c6ecfbe0ef499b013d731b7b9247e0b9" + +[[package]] +name = "wx-cli" +version = "0.7.2" +dependencies = [ + "axum", + "chrono", + "clap 4.6.0", + "futures-util", + "hex", + "libc", + "lru", + "rayon", + "rusqlite", + "serde", + "serde_json", + "silk-rs", + "tempfile", + "tokio", + "tokio-stream", + "tokio-util", + "toml", + "tower", + "tower-http", + "tracing-subscriber", + "ureq", + "url", + "vergen-gitcl", + "wx-context", + "wx-db", + "wx-decrypt", + "wx-keychain", + "wx-media", + "wx-monitor", + "wx-paths", +] + +[[package]] +name = "wx-context" +version = "0.7.2" +dependencies = [ + "aes", + "cbc", + "dashmap", + "filetime", + "hex", + "hmac", + "pbkdf2", + "rayon", + "rusqlite", + "rust-stemmers", + "serde", + "sha2", + "tempfile", + "thiserror 2.0.18", + "wx-db", + "wx-decrypt", + "wx-keychain", + "wx-paths", +] + +[[package]] +name = "wx-db" +version = "0.7.2" +dependencies = [ + "hex", + "insta", + "md5", + "prost", + "rusqlite", + "serde", + "serde_json", + "tempfile", + "thiserror 2.0.18", + "zstd", +] + +[[package]] +name = "wx-decrypt" +version = "0.7.2" +dependencies = [ + "aes", + "cbc", + "hmac", + "pbkdf2", + "sha2", + "tempfile", + "thiserror 2.0.18", +] + +[[package]] +name = "wx-keychain" +version = "0.7.2" +dependencies = [ + "aes", + "cbc", + "chrono", + "hex", + "hmac", + "libc", + "mach2", + "regex", + "serde", + "sha2", + "tempfile", + "thiserror 2.0.18", + "tokio", + "toml", + "wx-decrypt", + "wx-paths", +] + +[[package]] +name = "wx-media" +version = "0.7.2" +dependencies = [ + "aes", + "base64", + "cipher", + "ecb", + "hex", + "md5", + "rusqlite", + "serde", + "silk-rs", + "tempfile", + "thiserror 2.0.18", + "wx-keychain", +] + +[[package]] +name = "wx-monitor" +version = "0.7.2" +dependencies = [ + "aes", + "cbc", + "futures-core", + "hmac", + "notify", + "pbkdf2", + "rusqlite", + "serde", + "sha2", + "tempfile", + "thiserror 2.0.18", + "tokio", + "tracing", + "wx-db", + "wx-decrypt", +] + +[[package]] +name = "wx-paths" +version = "0.7.2" +dependencies = [ + "dirs", + "libc", + "serde", + "tempfile", + "thiserror 2.0.18", +] + +[[package]] +name = "yoke" +version = "0.8.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "72d6e5c6afb84d73944e5cedb052c4680d5657337201555f9f2a16b7406d4954" +dependencies = [ + "stable_deref_trait", + "yoke-derive", + "zerofrom", +] + +[[package]] +name = "yoke-derive" +version = "0.8.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b659052874eb698efe5b9e8cf382204678a0086ebf46982b79d6ca3182927e5d" +dependencies = [ + "proc-macro2", + "quote", + "syn", + "synstructure", +] + +[[package]] +name = "zerocopy" +version = "0.8.47" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "efbb2a062be311f2ba113ce66f697a4dc589f85e78a4aea276200804cea0ed87" +dependencies = [ + "zerocopy-derive", +] + +[[package]] +name = "zerocopy-derive" +version = "0.8.47" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0e8bc7269b54418e7aeeef514aa68f8690b8c0489a06b0136e5f57c4c5ccab89" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "zerofrom" +version = "0.1.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "50cc42e0333e05660c3587f3bf9d0478688e15d870fab3346451ce7f8c9fbea5" +dependencies = [ + "zerofrom-derive", +] + +[[package]] +name = "zerofrom-derive" +version = "0.1.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d71e5d6e06ab090c67b5e44993ec16b72dcbaabc526db883a360057678b48502" +dependencies = [ + "proc-macro2", + "quote", + "syn", + "synstructure", +] + +[[package]] +name = "zeroize" +version = "1.8.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b97154e67e32c85465826e8bcc1c59429aaaf107c1e4a9e53c8d8ccd5eff88d0" + +[[package]] +name = "zerotrie" +version = "0.2.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2a59c17a5562d507e4b54960e8569ebee33bee890c70aa3fe7b97e85a9fd7851" +dependencies = [ + "displaydoc", + "yoke", + "zerofrom", +] + +[[package]] +name = "zerovec" +version = "0.11.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6c28719294829477f525be0186d13efa9a3c602f7ec202ca9e353d310fb9a002" +dependencies = [ + "yoke", + "zerofrom", + "zerovec-derive", +] + +[[package]] +name = "zerovec-derive" +version = "0.11.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "eadce39539ca5cb3985590102671f2567e659fca9666581ad3411d59207951f3" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "zmij" +version = "1.0.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b8848ee67ecc8aedbaf3e4122217aff892639231befc6a1b58d29fff4c2cabaa" + +[[package]] +name = "zstd" +version = "0.13.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e91ee311a569c327171651566e07972200e76fcfe2242a4fa446149a3881c08a" +dependencies = [ + "zstd-safe", +] + +[[package]] +name = "zstd-safe" +version = "7.2.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8f49c4d5f0abb602a93fb8736af2a4f4dd9512e36f7f570d66e65ff867ed3b9d" +dependencies = [ + "zstd-sys", +] + +[[package]] +name = "zstd-sys" +version = "2.0.16+zstd.1.5.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91e19ebc2adc8f83e43039e79776e3fda8ca919132d68a1fed6a5faca2683748" +dependencies = [ + "cc", + "pkg-config", +] diff --git a/Cargo.toml b/Cargo.toml new file mode 100644 index 0000000..8912cbf --- /dev/null +++ b/Cargo.toml @@ -0,0 +1,7 @@ +[workspace] +members = ["crates/wx-decrypt", "crates/wx-keychain", "crates/wx-cli", "crates/wx-db", "crates/wx-media", "crates/wx-monitor", "crates/wx-context", "crates/wx-paths"] +resolver = "2" + +[workspace.package] +version = "0.7.2" +edition = "2021" diff --git a/README.md b/README.md new file mode 100644 index 0000000..e6a0238 --- /dev/null +++ b/README.md @@ -0,0 +1,218 @@ +# wx-cli + +WeChat macOS 数据库解密与查询工具。支持通过 `key scan`(Mach VM 内存扫描)或 `key extract`(LLDB hook)提取密钥,解密并查询 WeChat 4.1.7.x / 4.1.8.x 的 Apple SEE 加密 SQLite 数据库。 + +## 支持范围 + +- **平台**:macOS(arm64 / Apple Silicon) +- **WeChat 版本**:4.1.7.x / 4.1.8.x + +## 前置条件 + +密钥提取**需要 SIP 关闭**(SIP enabled 时 `task_for_pid` 被内核拒绝,即使 root 也不行)。如果你已有密钥,可以跳过 SIP 要求,直接用 `key set` 手动录入。 + +| 条件 | key scan | key extract | +|------|----------|-------------| +| SIP disabled + sudo | **可用** | **可用,但通常不需要 sudo** | +| SIP disabled + 无 sudo | 不可用 | **可用(推荐)** | +| SIP enabled | 不可用 | 不可用 | + +`key extract`(LLDB 方式)还需要: + +1. `sudo DevToolsSecurity -enable` +2. `sudo dscl . append /Groups/_developer GroupMembership $USER` +3. `xcode-select --install`(提供 `lldb` 和 `python3`) + +## 安装 + +### 从源码构建 + +```bash +# 需要 Rust 工具链(rustup 安装即可) +cargo build --release +``` + +编译产物位于 `target/release/wx-cli`。 + +### 部署二进制 + +```bash +mkdir -p ~/.local/bin +cp target/release/wx-cli ~/.local/bin/wx-cli +chmod +x ~/.local/bin/wx-cli + +# 确保 PATH 包含该目录 +echo 'export PATH="$HOME/.local/bin:$PATH"' >> ~/.zshrc +source ~/.zshrc + +wx-cli --version +``` + +如果目标机器在远程(例如 VM),本机构建后传输即可,远程不需要 Rust 工具链: + +```bash +scp target/release/wx-cli user@remote:~/.local/bin/wx-cli +``` + +## 使用 + +### 1. 检查环境 + +```bash +wx-cli doctor # 检查 SIP、DevToolsSecurity、_developer 组、LLDB/python3 +wx-cli status # 查看 WeChat 运行状态和所有账号密钥/缓存状态 +``` + +### 2. 提取密钥 + +```bash +# 方式 A(推荐):内存扫描 — 不重启 WeChat,需要 sudo +sudo wx-cli key scan + +# 方式 B:LLDB hook — 会重启 WeChat,通常不需要 sudo +wx-cli key extract --timeout 120 + +# 查看已保存的密钥 +wx-cli key list +``` + +`key scan` 从 WeChat 进程内存中提取已缓存的数据库密钥;`key extract` 通过 LLDB hook 捕获 PBKDF2 调用获取原始密钥,覆盖范围更广。 + +手动设置密钥: + +```bash +wx-cli key set <64-hex-key> # 数据库密钥 +wx-cli key set-image # 图片密钥 +``` + +### 3. 解密数据库 + +```bash +wx-cli decrypt # 自动解密到缓存目录 +wx-cli decrypt --incremental # 增量解密(只处理变化的文件) + +# 手动指定路径和密钥 +wx-cli decrypt -k <64位hex密钥> -d /path/to/xwechat_files/ -o /tmp/decrypted +``` + +### 4. 查询聊天记录 + +```bash +wx-cli sessions --limit 10 # 最近会话 +wx-cli contacts --search 张三 # 搜索联系人 +wx-cli query 张三 --limit 20 # 查某人的消息 +wx-cli search 周末 --limit 20 # 全局关键词搜索 +wx-cli query 张三 --type text # 按消息类型过滤 +wx-cli query 周末爬山群 # 群聊消息 +wx-cli export 张三 -o /tmp/export --format json # 导出会话 +wx-cli watch --poll --poll-ms 3000 # 实时监听新消息 +``` + +如果本机已启动 `server run` 服务,查询命令会自动复用 REST API(默认探测 `http://127.0.0.1:9100`)。可用 `--no-server` 强制本地查询,或 `--server-only` 强制远程。 + +### 5. 媒体解密 + +```bash +wx-cli decode-image input.dat -d -o output.png # 解密图片 +wx-cli decode-image /path/to/dat_dir/ -d -o /tmp/ # 批量解密 +wx-cli media extract-voice --media-dir -o voice.mp3 # 提取语音(需 ffmpeg) +wx-cli media decrypt-video encrypted.bin --seed 2105122989 -o video.mp4 # 解密视频号视频 +``` + +### 6. HTTP API 服务 + +```bash +wx-cli server run # 启动(默认 127.0.0.1:9100) +wx-cli server run --host 0.0.0.0 --token mysecret # 远程访问(必须设 token) +wx-cli server status # 查看状态 +wx-cli server stop # 停止 +wx-cli server restart # 重启 +``` + +REST 端点:`/api/v1/health`、`/api/v1/sessions`、`/api/v1/contacts`、`/api/v1/messages`、`/api/v1/search`、`/api/v1/media`、`/api/v1/events`(SSE)。 + +所有查询命令加 `--format json` 可获取 JSON 格式输出。 + +## 命令一览 + +| 命令 | 说明 | +|------|------| +| `wx-cli status` | 查看 WeChat 运行状态 | +| `wx-cli doctor` | 检查环境(SIP 等) | +| `sudo wx-cli key scan` | 内存扫描提取密钥(推荐) | +| `wx-cli key extract` | LLDB hook 提取密钥 | +| `wx-cli key list` | 查看已保存密钥 | +| `wx-cli key set ` | 手动设置密钥 | +| `wx-cli decrypt` | 解密数据库 | +| `wx-cli sessions` | 最近会话列表 | +| `wx-cli contacts --search <名字>` | 搜索联系人 | +| `wx-cli query <联系人>` | 查询消息 | +| `wx-cli search <关键词>` | 全局搜索 | +| `wx-cli export <联系人>` | 导出会话 | +| `wx-cli watch` | 实时监听新消息 | +| `wx-cli decode-image <路径>` | 解密图片 | +| `wx-cli media extract-voice` | 提取语音 | +| `wx-cli media decrypt-video` | 解密视频号视频 | +| `wx-cli server run` | 启动 HTTP API 服务 | +| `wx-cli server status/stop/restart` | 管理服务 | +| `wx-cli paths` | 查看所有数据路径 | +| `wx-cli info ` | 查看数据库加密状态 | + +## Contact Hiding + +按账号隐藏指定联系人、群聊或带特定标签的联系人。启用后,查询、导出、监控和服务接口默认应用隐藏规则。 + +配置文件:`~/Library/Application Support/wx-cli/config/settings.toml` + +```toml +[accounts.""] +ignore_contacts = ["wxid_xxx", "12345@chatroom"] +ignore_tags = ["同事", "客户"] +``` + +本地命令支持 `--show-hidden` 忽略隐藏规则查看完整结果。`search` 当前不会自动应用隐藏配置。 + +## 文件路径 + +| 类别 | 路径(macOS) | 用途 | 可删除? | +|------|---------------|------|----------| +| Config | `~/Library/Application Support/wx-cli/config/` | 密钥、设置 | 否(先备份) | +| Cache | `~/Library/Caches/wx-cli/` | 解密后数据库 | 可(重新 decrypt) | +| State | `~/Library/Application Support/wx-cli/state/` | 服务运行时元数据 | 可 | +| Logs | `~/Library/Logs/wx-cli/` | 服务日志 | 可 | +| Temp | `$TMPDIR/wx-cli/` | 密钥提取临时文件 | 可 | + +使用 `wx-cli paths` 查看所有路径。清理缓存:`rm -rf ~/Library/Caches/wx-cli/`。 + +## 项目结构 + +``` +wx-cli/ +├── crates/ +│ ├── wx-decrypt/ # 核心解密库(KDF、逐页解密、整库解密) +│ ├── wx-keychain/ # 密钥提取(LLDB / Mach VM)与本地存储 +│ ├── wx-cli/ # CLI 入口 +│ ├── wx-db/ # 数据库查询(联系人、消息、会话、群聊) +│ ├── wx-media/ # 媒体解密(图片、语音、视频) +│ ├── wx-monitor/ # 实时消息监听与增量监控 +│ ├── wx-context/ # 账号解析、解密缓存、联系人解析 +│ └── wx-paths/ # 平台路径管理 +``` + +## 常见问题 + +### `key extract` 超时 + +- 确认 WeChat 已弹出登录界面并完成登录 +- 增加超时:`--timeout 300` +- 检查日志:`$TMPDIR/wx-cli/lldb/wechat_lldb_output.txt` + +### SIP / DevToolsSecurity 报错 + +密钥提取需要 SIP 关闭。重启进入恢复模式执行 `csrutil disable`,然后运行 `wx-cli doctor` 逐项检查。 + +### 解密后数据库无法打开 + +- `wx-cli key list` 确认密钥正确 +- `wx-cli info ` 检查文件是否为加密状态 +- 确认 WeChat 版本在 4.1.7.x / 4.1.8.x 范围内 diff --git a/crates/wx-cli/Cargo.toml b/crates/wx-cli/Cargo.toml new file mode 100644 index 0000000..dbef4b6 --- /dev/null +++ b/crates/wx-cli/Cargo.toml @@ -0,0 +1,48 @@ +[package] +name = "wx-cli" +version.workspace = true +edition.workspace = true + +[[bin]] +name = "wx-cli" +path = "src/main.rs" + +[features] +default = ["audio"] +audio = ["wx-media/audio"] + +[dependencies] +wx-decrypt = { path = "../wx-decrypt" } +wx-keychain = { path = "../wx-keychain" } +wx-db = { path = "../wx-db" } +wx-media = { path = "../wx-media" } +wx-monitor = { path = "../wx-monitor" } +wx-context = { path = "../wx-context" } +wx-paths = { path = "../wx-paths" } +rusqlite = { version = "0.32", features = ["bundled"] } +clap = { version = "4", features = ["derive"] } +futures-util = "0.3" +hex = "0.4" +serde = { version = "1", features = ["derive"] } +serde_json = "1" +toml = "0.8" +ureq = { version = "2", features = ["json"] } +url = "2" +tokio = { version = "1", features = ["rt-multi-thread", "macros", "signal"] } +tracing-subscriber = { version = "0.3", features = ["env-filter"] } +chrono = { version = "0.4", default-features = false, features = ["clock"] } +axum = { version = "0.8", features = ["macros"] } +tokio-stream = { version = "0.1", features = ["sync"] } +tokio-util = { version = "0.7", features = ["rt"] } +tower-http = { version = "0.6", features = ["cors", "fs"] } +tower = { version = "0.5", features = ["util"] } +libc = "0.2" +lru = "0.12" +rayon = "1" + +[dev-dependencies] +tempfile = "3" +silk-rs = "0.2" + +[build-dependencies] +vergen-gitcl = { version = "9", features = ["build"] } diff --git a/crates/wx-cli/build.rs b/crates/wx-cli/build.rs new file mode 100644 index 0000000..178960f --- /dev/null +++ b/crates/wx-cli/build.rs @@ -0,0 +1,16 @@ +use vergen_gitcl::{BuildBuilder, Emitter, GitclBuilder}; + +fn main() -> Result<(), Box> { + let build = BuildBuilder::default() + .build_date(true) + .use_local(true) + .build()?; + let gitcl = GitclBuilder::default().sha(true).build()?; + + Emitter::default() + .add_instructions(&build)? + .add_instructions(&gitcl)? + .emit()?; + + Ok(()) +} diff --git a/crates/wx-cli/src/cmd/contacts.rs b/crates/wx-cli/src/cmd/contacts.rs new file mode 100644 index 0000000..fe03e06 --- /dev/null +++ b/crates/wx-cli/src/cmd/contacts.rs @@ -0,0 +1,210 @@ +use std::path::PathBuf; + +use wx_context::{AccountContext, ContactResolver, ResolveParams, VisibilityIndex}; +use wx_db::Contact; + +use super::thin_client::{ThinClientCliArgs, ThinClientOptions}; +use crate::output::JsonEnvelope; +use crate::settings::Settings; +use crate::util::{effective_limit_all, open_db_core, print_cache_stats, print_detection_note, try_remote_or_local}; +use crate::visibility_projection::project_contacts_envelope; +use crate::OutputFormat; + +#[allow(clippy::too_many_arguments)] +pub fn cmd_contacts( + data_dir: Option, + account: Option, + key: Option, + search: Option, + limit: usize, + offset: usize, + all: bool, + format: OutputFormat, + show_hidden: bool, + server: ThinClientCliArgs, +) -> Result<(), Box> { + let options = ThinClientOptions::resolve_from_process_env(server); + let effective_limit = effective_limit_all(all, limit); + + let envelope = try_remote_or_local( + &options, + |client| { + let mut query = vec![ + ("limit".to_string(), effective_limit.to_string()), + ("offset".to_string(), offset.to_string()), + ]; + if let Some(ref search) = search { + query.push(("search".to_string(), search.to_string())); + } + if show_hidden { + query.push(("show_hidden".to_string(), "1".to_string())); + } + client.get_json("/api/v1/contacts", &query) + }, + || { + load_local_contacts( + data_dir, + account, + key, + search.clone(), + effective_limit, + offset, + show_hidden, + ) + }, + "contacts", + )?; + print_contacts_output(&envelope, format) +} + +pub(crate) fn build_visibility( + acct: &AccountContext, + resolver: &ContactResolver, +) -> VisibilityIndex { + let settings = Settings::load_default().unwrap_or_default(); + let account_settings = settings.for_account(&acct.account_id); + VisibilityIndex::build( + &account_settings.ignore_contacts, + &account_settings.ignore_tags, + resolver, + ) +} + +fn load_local_contacts( + data_dir: Option, + account: Option, + key: Option, + search: Option, + effective_limit: usize, + offset: usize, + show_hidden: bool, +) -> Result, Box> { + let acct = AccountContext::resolve(&ResolveParams { + account: account.as_deref(), + data_dir: data_dir.as_deref(), + key_hex: key.as_deref(), + })?; + print_detection_note(&acct); + + let (db, stats) = open_db_core(&acct, crate::util::decrypt_progress_callback)?; + if let Some(ref s) = stats { + print_cache_stats(s); + } + + let resolver = ContactResolver::build(&db)?; + let visibility = build_visibility(&acct, &resolver); + + let mut query = wx_db::ContactQuery::new() + .limit(wx_db::MAX_QUERY_LIMIT) + .offset(0); + if let Some(ref kw) = search { + query = query.keyword(kw); + } + + let result = db.query_contacts(&query)?; + let envelope = JsonEnvelope::from_query_result(result, wx_db::MAX_QUERY_LIMIT, 0, |c| c); + Ok(project_contacts_envelope( + envelope.items, + &visibility, + effective_limit, + offset, + &envelope.stats, + show_hidden, + )) +} + +fn print_contacts_output( + envelope: &JsonEnvelope, + format: OutputFormat, +) -> Result<(), Box> { + match format { + OutputFormat::Json => println!("{}", serde_json::to_string_pretty(envelope)?), + OutputFormat::Text => { + render_contacts_text(&envelope.items); + eprintln!("{} contacts found.", envelope.items.len()); + } + } + Ok(()) +} + +fn render_contacts_text(items: &[Contact]) { + for c in items { + let alias_part = if c.alias.is_empty() { + String::new() + } else { + format!(" / {}", c.alias) + }; + + let display = if !c.remark.is_empty() { + format!("{}{}", c.remark, alias_part) + } else if !c.nick_name.is_empty() { + format!("{}{}", c.nick_name, alias_part) + } else { + alias_part.trim_start_matches(" / ").to_string() + }; + + if display.is_empty() { + println!(" ({})", c.user_name); + } else { + println!(" {:<30} ({})", display, c.user_name); + } + + print_contact_tree(c); + } +} + +fn gender_label(g: u32) -> &'static str { + match g { + 1 => "male", + 2 => "female", + _ => "unknown", + } +} + +fn source_scene_label(s: u32) -> String { + match s { + 1 => "通过QQ号添加".to_string(), + 3 => "通过微信号添加".to_string(), + 6 => "通过手机号添加".to_string(), + 10 => "通过名片添加".to_string(), + 14 => "通过群聊添加".to_string(), + 30 => "通过扫一扫添加".to_string(), + _ => format!("场景码 {s}"), + } +} + +fn print_contact_tree(c: &Contact) { + let mut lines: Vec<(String, String)> = Vec::new(); + + if let Some(ref phone) = c.phone { + lines.push(("Phone".to_string(), phone.clone())); + } + if let Some(ref sig) = c.signature { + lines.push(("Signature".to_string(), sig.clone())); + } + if let Some(ref region) = c.region { + lines.push(("Region".to_string(), region.clone())); + } + if let Some(g) = c.gender { + lines.push(("Gender".to_string(), gender_label(g).to_string())); + } + if let Some(s) = c.source_scene { + lines.push(("Source".to_string(), source_scene_label(s))); + } + if let Some(ref memo) = c.memo { + lines.push(("Memo".to_string(), memo.clone())); + } + if !c.labels.is_empty() { + lines.push(("Labels".to_string(), c.labels.join(", "))); + } + + if lines.is_empty() { + return; + } + + let last_idx = lines.len() - 1; + for (i, (key, val)) in lines.iter().enumerate() { + let prefix = if i == last_idx { "└─" } else { "├─" }; + println!(" {prefix} {key}: {val}"); + } +} diff --git a/crates/wx-cli/src/cmd/db_dev.rs b/crates/wx-cli/src/cmd/db_dev.rs new file mode 100644 index 0000000..12c5e1f --- /dev/null +++ b/crates/wx-cli/src/cmd/db_dev.rs @@ -0,0 +1,65 @@ +use crate::DbDevAction; + +pub fn cmd_db_dev( + path: &std::path::Path, + action: DbDevAction, +) -> Result<(), Box> { + let db = wx_db::WechatDb::open(path)?; + + match action { + DbDevAction::Contacts { + keyword, + limit, + offset, + } => { + let mut q = wx_db::ContactQuery::new().limit(limit).offset(offset); + if let Some(kw) = keyword { + q = q.keyword(kw); + } + let result = db.query_contacts(&q)?; + println!("{}", serde_json::to_string_pretty(&result)?); + } + DbDevAction::Sessions { limit, offset } => { + let result = + db.query_sessions(&wx_db::SessionQuery::new().limit(limit).offset(offset))?; + println!("{}", serde_json::to_string_pretty(&result)?); + } + DbDevAction::Messages { + talker, + start, + end, + keyword, + limit, + offset, + } => { + let mut q = wx_db::MessageQuery::for_talker(talker) + .limit(limit) + .offset(offset); + if let Some(s) = start { + q = q.since(s); + } + if let Some(e) = end { + q = q.until(e); + } + if let Some(kw) = keyword { + q = q.keyword(kw); + } + let result = db.query_messages(&q)?; + println!("{}", serde_json::to_string_pretty(&result)?); + } + DbDevAction::Chatrooms { + username, + limit, + offset, + } => { + let mut q = wx_db::ChatRoomQuery::new().limit(limit).offset(offset); + if let Some(name) = username { + q = q.username(name); + } + let result = db.query_chatrooms(&q)?; + println!("{}", serde_json::to_string_pretty(&result)?); + } + } + + Ok(()) +} diff --git a/crates/wx-cli/src/cmd/decode_image.rs b/crates/wx-cli/src/cmd/decode_image.rs new file mode 100644 index 0000000..3b3257e --- /dev/null +++ b/crates/wx-cli/src/cmd/decode_image.rs @@ -0,0 +1,20 @@ +use std::path::PathBuf; + +use crate::MediaAction; + +pub fn cmd_decode_image( + input: PathBuf, + output: Option, + account: Option, + data_dir: Option, +) -> Result<(), Box> { + let action = MediaAction::DecryptDat { + input, + output, + v2_key: None, + account, + data_dir, + xor_key: None, + }; + super::media::cmd_media(action) +} diff --git a/crates/wx-cli/src/cmd/decrypt.rs b/crates/wx-cli/src/cmd/decrypt.rs new file mode 100644 index 0000000..d83d9e4 --- /dev/null +++ b/crates/wx-cli/src/cmd/decrypt.rs @@ -0,0 +1,158 @@ +use std::path::PathBuf; + +use wx_context::{AccountContext, DecryptRequest, PersistentCache, ResolveParams}; +use wx_decrypt::KeyMaterial; + +use crate::util::{find_db_files, print_cache_stats, print_detection_note}; + +pub fn cmd_decrypt( + key_hex: Option, + data_dir: Option, + account: Option, + output: Option, + incremental: bool, +) -> Result<(), Box> { + let params = &wx_decrypt::MACOS_4_1_7_31; + + // Use PersistentCache path when no explicit output is given + if output.is_none() || incremental { + let acct = AccountContext::resolve(&ResolveParams { + account: account.as_deref(), + data_dir: data_dir.as_deref(), + key_hex: key_hex.as_deref(), + })?; + print_detection_note(&acct); + + let cache = PersistentCache::new(&acct, params)?; + let stats = DecryptRequest::new() + .all() + .execute_with_progress(&cache, crate::util::decrypt_progress_callback)?; + print_cache_stats(&stats); + + eprintln!( + "Done: {} decrypted, {} cached, {} errors, {} WAL patched", + stats.decrypted, stats.skipped, stats.errors, stats.wal_patched + ); + eprintln!("Output: {}", cache.decrypted_root().display()); + return Ok(()); + } + + // Explicit output directory — full decrypt (legacy behavior) + let output_dir = output.unwrap(); + + let acct = AccountContext::resolve(&ResolveParams { + account: account.as_deref(), + data_dir: data_dir.as_deref(), + key_hex: key_hex.as_deref(), + })?; + print_detection_note(&acct); + + let db_storage = acct.data_dir.join("db_storage"); + if !db_storage.exists() { + return Err(format!("db_storage not found in {}", acct.data_dir.display()).into()); + } + + let db_files = find_db_files(&db_storage)?; + if db_files.is_empty() { + return Err("no .db files found in db_storage/".into()); + } + + eprintln!("Found {} database files.", db_files.len()); + + let mut ok_count = 0u32; + let mut skip_count = 0u32; + let mut err_count = 0u32; + let mut wal_ok_count = 0u32; + let mut wal_err_count = 0u32; + + for db_path in &db_files { + let rel = db_path.strip_prefix(&acct.data_dir).unwrap_or(db_path); + let out_path = output_dir.join(rel); + + let db_result = match &acct.key_material { + KeyMaterial::RawKey(key) => wx_decrypt::decrypt_db(db_path, &out_path, key, params), + KeyMaterial::EncKey { key, salt } => { + wx_decrypt::decrypt_db_direct(db_path, &out_path, key, salt, params) + } + KeyMaterial::EncKeys(pairs) => { + match wx_decrypt::read_main_db_salt_for_path(db_path) { + Ok(db_salt) => match pairs.iter().find(|p| p.salt == db_salt) { + Some(pair) => wx_decrypt::decrypt_db_direct( + db_path, &out_path, &pair.key, &pair.salt, params, + ), + None => Err(wx_decrypt::DecryptError::NoMatchingEncKey), + }, + Err(e) => Err(e), + } + } + }; + + match db_result { + Ok(()) => { + eprintln!(" OK {}", rel.display()); + ok_count += 1; + + let wal_path = db_path.with_extension("db-wal"); + if wal_path.exists() { + let wal_result = match &acct.key_material { + KeyMaterial::RawKey(key) => { + wx_decrypt::decrypt_wal(&wal_path, &out_path, key, params) + } + KeyMaterial::EncKey { key, salt } => wx_decrypt::decrypt_wal_direct( + &wal_path, &out_path, key, salt, params, + ), + KeyMaterial::EncKeys(pairs) => { + match wx_decrypt::read_main_db_salt_for_path(&wal_path) { + Ok(db_salt) => match pairs.iter().find(|p| p.salt == db_salt) { + Some(pair) => wx_decrypt::decrypt_wal_direct( + &wal_path, &out_path, &pair.key, &pair.salt, params, + ), + None => Err(wx_decrypt::DecryptError::NoMatchingEncKey), + }, + Err(e) => Err(e), + } + } + }; + match wal_result { + Ok(n) if n > 0 => { + eprintln!(" WAL {} — {n} frames patched", rel.display()); + wal_ok_count += 1; + } + Ok(_) => {} + Err(e) => { + eprintln!(" WAL ERR {} — {e}", rel.display()); + wal_err_count += 1; + } + } + } + } + Err(wx_decrypt::DecryptError::AlreadyDecrypted) => { + eprintln!(" SKIP {}", rel.display()); + skip_count += 1; + } + Err(e) => { + eprintln!(" ERR {} — {e}", rel.display()); + err_count += 1; + } + } + } + + eprint!("\nDone: {ok_count} decrypted, {skip_count} skipped, {err_count} errors"); + if wal_ok_count > 0 || wal_err_count > 0 { + eprint!(", WAL: {wal_ok_count} patched / {wal_err_count} failed"); + } + eprintln!("."); + + if err_count > 0 || wal_err_count > 0 { + let mut parts = Vec::new(); + if err_count > 0 { + parts.push(format!("{err_count} databases failed")); + } + if wal_err_count > 0 { + parts.push(format!("{wal_err_count} WAL patches failed")); + } + return Err(parts.join(", ").into()); + } + + Ok(()) +} diff --git a/crates/wx-cli/src/cmd/doctor.rs b/crates/wx-cli/src/cmd/doctor.rs new file mode 100644 index 0000000..62f4477 --- /dev/null +++ b/crates/wx-cli/src/cmd/doctor.rs @@ -0,0 +1,33 @@ +pub fn cmd_doctor(fix: bool) -> Result<(), Box> { + let checks = wx_keychain::all_preflight_checks(); + + let all_passed = checks.iter().all(|c| c.passed); + + for c in &checks { + let icon = if c.passed { "\u{2705}" } else { "\u{2717}" }; + println!("{icon} {:<24} {}", c.name, c.detail); + } + + if fix && !all_passed { + println!("\n--- Fix commands ---"); + let mut has_fix = false; + for c in &checks { + if !c.passed { + if let Some(ref cmd) = c.fix_cmd { + println!("\n# Fix: {}", c.name); + println!("{cmd}"); + has_fix = true; + } + } + } + if !has_fix { + println!("(no automatic fixes available)"); + } + } + + if all_passed { + println!("\nAll checks passed."); + } + + Ok(()) +} diff --git a/crates/wx-cli/src/cmd/export.rs b/crates/wx-cli/src/cmd/export.rs new file mode 100644 index 0000000..84a5a31 --- /dev/null +++ b/crates/wx-cli/src/cmd/export.rs @@ -0,0 +1,856 @@ +use std::io::Write; +use std::path::PathBuf; + +use serde::Serialize; +use wx_context::{ + AccountContext, ContactResolver, DecryptRequest, Direction, PersistentCache, ResolveParams, +}; +use wx_db::{is_group_chat, MessageContent, MessageQuery, SortOrder, MAX_QUERY_LIMIT}; + +use crate::cmd::export_media::{MediaKind, MediaStats}; +use crate::cmd::query::resolve_talker; +use crate::cmd::contacts::build_visibility; +use crate::output::{JsonEnvelope, PagingMeta, StatsMeta}; +use crate::schema::{enrich_message, project_message_items, EnrichedMessage}; +use crate::util::{ + decrypt_progress_callback, open_db_all, print_cache_stats, print_detection_note, + sanitize_filename, +}; +use crate::{ExportFormat, SortOrderArg}; + +// ── JSON export types ─────────────────────────────────────────────── + +#[derive(Serialize)] +struct ExportEnvelope { + export_info: ExportInfo, + conversation: ConversationMeta, + #[serde(flatten)] + envelope: JsonEnvelope, +} + +#[derive(Serialize)] +struct ExportInfo { + version: &'static str, + exported_at: String, + generator: String, +} + +#[derive(Serialize)] +struct ConversationMeta { + talker: String, + display_name: String, + #[serde(rename = "type")] + conv_type: &'static str, + message_count: usize, + #[serde(skip_serializing_if = "Option::is_none")] + time_range_start: Option, + #[serde(skip_serializing_if = "Option::is_none")] + time_range_end: Option, +} + +#[derive(Serialize)] +struct ExportedMessage { + #[serde(flatten)] + enriched: EnrichedMessage, + #[serde(skip_serializing_if = "Vec::is_empty")] + media_files: Vec, +} + +// ── Main entry ────────────────────────────────────────────────────── + +#[allow(clippy::too_many_arguments)] +pub fn cmd_export( + contact: &str, + output_dir: PathBuf, + data_dir: Option, + account: Option, + key: Option, + since: Option, + until: Option, + limit: usize, + offset: usize, + order: SortOrderArg, + all: bool, + format: ExportFormat, + no_media: bool, + show_emoji: bool, + show_hidden: bool, + parallel: Option, +) -> Result<(), Box> { + let params = &wx_decrypt::MACOS_4_1_7_31; + + // Bootstrap + let acct = AccountContext::resolve(&ResolveParams { + account: account.as_deref(), + data_dir: data_dir.as_deref(), + key_hex: key.as_deref(), + })?; + print_detection_note(&acct); + + // For --no-media with raw_key: fully direct. Otherwise need cache for media. + let (db, cache) = if acct.raw_key.is_some() && no_media { + let (db, _, _) = open_db_all(&acct, decrypt_progress_callback)?; + (db, None::) + } else if acct.raw_key.is_some() { + // Messages via direct encrypted open; media files (dat/video/file) live outside + // db_storage and need decrypted hardlink.db + message_resource DB for resolution. + // DecryptScope doesn't yet have a MediaOnly variant, so we decrypt all DBs. + // The message *query* still goes through the direct encrypted WechatDb. + eprintln!( + "Direct encrypted open (SQLCipher) for messages; decrypt cache for media resolution" + ); + let db = wx_context::open_encrypted_db(&acct)?; + let cache = PersistentCache::new(&acct, params)?; + let stats = DecryptRequest::new() + .all() + .execute_with_progress(&cache, decrypt_progress_callback)?; + print_cache_stats(&stats); + (db, Some(cache)) + } else { + let cache = PersistentCache::new(&acct, params)?; + let stats = DecryptRequest::new() + .all() + .execute_with_progress(&cache, decrypt_progress_callback)?; + print_cache_stats(&stats); + let db = wx_db::WechatDb::open(cache.decrypted_root())?; + (db, Some(cache)) + }; + let resolver = ContactResolver::build(&db)?; + let visibility = build_visibility(&acct, &resolver); + let self_wxid = &acct.base_wxid; + + let talker = resolve_talker(contact, &resolver, &db, Some(&visibility), show_hidden)?; + let is_group = is_group_chat(&talker); + + // Query params + let effective_order: SortOrder = order.into(); + let effective_limit = if all { + usize::MAX + } else { + wx_db::effective_limit(limit) + }; + + // Batch query + let mut all_messages = Vec::new(); + let mut cursor = offset; + let mut total_count: Option = None; + let mut total_scanned = 0usize; + let mut total_skipped = 0usize; + let mut shard_warnings = Vec::new(); + + if all { + loop { + let mut query = MessageQuery::for_talker(&talker) + .limit(MAX_QUERY_LIMIT) + .offset(cursor) + .order(effective_order) + .with_filtered_count(cursor == offset); + if let Some(s) = since { + query = query.since(s); + } + if let Some(u) = until { + query = query.until(u); + } + let result = db.query_messages(&query)?; + let count = result.items.len(); + if cursor == offset { + total_count = Some( + result + .stats + .filtered_count + .unwrap_or(result.stats.total_rows), + ); + } + // Only capture stats from the first batch (they cover the full query scope) + if cursor == offset { + total_scanned = result.stats.total_rows; + total_skipped = result.stats.skipped; + shard_warnings = result.shard_warnings; + } + all_messages.extend(result.items); + if count < MAX_QUERY_LIMIT { + break; + } + cursor += count; + } + } else { + loop { + let remaining = effective_limit.saturating_sub(all_messages.len()); + if remaining == 0 { + break; + } + let batch_size = remaining.min(MAX_QUERY_LIMIT); + let mut query = MessageQuery::for_talker(&talker) + .limit(batch_size) + .offset(cursor) + .order(effective_order) + .with_filtered_count(cursor == offset); + if let Some(s) = since { + query = query.since(s); + } + if let Some(u) = until { + query = query.until(u); + } + let result = db.query_messages(&query)?; + let count = result.items.len(); + if cursor == offset { + total_count = Some( + result + .stats + .filtered_count + .unwrap_or(result.stats.total_rows), + ); + } + if cursor == offset { + total_scanned = result.stats.total_rows; + total_skipped = result.stats.skipped; + shard_warnings = result.shard_warnings; + } + all_messages.extend(result.items); + if count < batch_size { + break; + } + cursor += count; + } + } + + // Print warnings + for w in &shard_warnings { + eprintln!("warning: shard {}: {}", w.path, w.reason); + } + if total_skipped > 0 { + eprintln!("warning: {total_skipped} messages skipped (decode error)"); + } + + if all_messages.is_empty() { + eprintln!("No messages found for export."); + return Ok(()); + } + + // Phase 2: Enrich messages, then apply sender projection + let enriched: Vec = all_messages + .into_iter() + .map(|m| enrich_message(m, self_wxid, &resolver)) + .collect(); + let projected = project_message_items(enriched, &talker, &visibility, show_hidden); + + if projected.is_empty() { + eprintln!("No visible messages found for export (all filtered by sender hiding)."); + return Ok(()); + } + + // Create output directories + std::fs::create_dir_all(&output_dir)?; + let media_dir = output_dir.join("media"); + if !no_media { + std::fs::create_dir_all(&media_dir)?; + } + + // Resolve media via parallel pipeline (or skip) + let (media_map, media_stats, _media_errors) = if no_media || cache.is_none() { + (vec![vec![]; projected.len()], MediaStats::default(), None) + } else { + let c = cache.as_ref().unwrap(); + let attach_dir = acct.data_dir.join("msg").join("attach"); + let decrypted_media = c.decrypted_root().join("message"); + let hardlink_db = c.decrypted_root().join("hardlink").join("hardlink.db"); + + let v2_aes_key = wx_media::derive_v2_key_from_dir(&acct.data_dir).ok(); + let dat_opts = wx_media::DatDecryptOptions { + v2_aes_key, + xor_key: None, + }; + + let file_dir = acct.data_dir.join("msg").join("file"); + let video_dir = acct.data_dir.join("msg").join("video"); + + let ctx = crate::cmd::export_task::build_shared_context( + attach_dir, + decrypted_media, + file_dir, + video_dir, + hardlink_db, + media_dir.clone(), + &talker, + dat_opts, + ); + + let tasks = crate::cmd::export_task::classify(&projected); + let (unique_tasks, dup_map) = crate::cmd::export_task::dedup(tasks); + let (results, resolve_errors) = crate::cmd::export_task::resolve_parallel( + unique_tasks, + std::sync::Arc::new(ctx), + parallel, + ); + let (media_map, stats, collect_errors) = + crate::cmd::export_task::collect(results, &dup_map, projected.len()); + + let combined = crate::cmd::export_task::ErrorSummary { + errors: resolve_errors + .errors + .into_iter() + .chain(collect_errors.errors) + .collect(), + }; + combined.print_report(); + (media_map, stats, Some(combined)) + }; + + let total_media: usize = media_map.iter().map(Vec::len).sum(); + + // Output + let talker_display = resolver.display_name(&talker).to_string(); + let safe_name = sanitize_filename(&talker_display); + let date_str = chrono::Local::now().format("%Y-%m-%d").to_string(); + + match format { + ExportFormat::Txt => { + let filename = format!("{safe_name}_{date_str}.txt"); + let out_path = output_dir.join(&filename); + write_txt( + &out_path, + &talker, + &talker_display, + self_wxid, + is_group, + &projected, + &media_map, + &resolver, + show_emoji, + )?; + eprintln!( + "Exported {} messages, {} media files → {}", + projected.len(), + total_media, + out_path.display() + ); + } + ExportFormat::Json => { + let filename = format!("{safe_name}_{date_str}.json"); + let out_path = output_dir.join(&filename); + let exported_count = projected.len(); + write_json( + &out_path, + &talker, + &talker_display, + is_group, + projected, + &media_map, + effective_limit, + offset, + total_count.unwrap_or(0), + total_scanned, + total_skipped, + shard_warnings, + )?; + eprintln!( + "Exported {} messages, {} media files → {}", + exported_count, + total_media, + out_path.display() + ); + } + } + + // Media quality hints + for hint in media_quality_hints(&media_stats) { + eprintln!("{hint}"); + } + + Ok(()) +} + +// ── TXT writer ────────────────────────────────────────────────────── + +#[allow(clippy::too_many_arguments)] +fn write_txt( + out_path: &std::path::Path, + _talker: &str, + talker_display: &str, + self_wxid: &str, + is_group: bool, + messages: &[EnrichedMessage], + media_map: &[Vec], + resolver: &ContactResolver, + show_emoji: bool, +) -> Result<(), Box> { + let mut f = std::io::BufWriter::new(std::fs::File::create(out_path)?); + + let self_name = { + let name = resolver.display_name(self_wxid); + if name == self_wxid { + self_wxid.to_string() + } else { + name.to_string() + } + }; + + // Header + if is_group { + writeln!(f, "\"{}\"的聊天记录如下:", talker_display)?; + } else { + writeln!( + f, + "\"{}\"和\"{}\"的聊天记录如下:", + self_name, talker_display + )?; + } + + // Meta info + let export_time = chrono::Local::now().format("%Y-%m-%d %H:%M:%S").to_string(); + let first_time = messages.iter().map(|em| em.message.create_time).min(); + let last_time = messages.iter().map(|em| em.message.create_time).max(); + + writeln!(f)?; + writeln!(f, "导出时间:{export_time}")?; + writeln!(f, "消息数:{}", messages.len())?; + if let (Some(first), Some(last)) = (first_time, last_time) { + let first_date = format_date_short(first); + let last_date = format_date_short(last); + writeln!(f, "时间范围:{first_date} ~ {last_date}")?; + } + + // Messages + let mut last_date: Option = None; + let mut image_counter = 0usize; + let mut voice_counter = 0usize; + let mut video_counter = 0usize; + let mut file_counter = 0usize; + let mut attachments: Vec<(String, String)> = Vec::new(); // (label, filename) + + for (i, em) in messages.iter().enumerate() { + let msg = &em.message; + let dt = chrono::DateTime::from_timestamp(msg.create_time, 0) + .map(|dt| dt.with_timezone(&chrono::Local)); + + // Date separator + if let Some(dt) = &dt { + let date = dt.date_naive(); + if last_date != Some(date) { + let (y, m, d) = (date.year(), date.month(), date.day()); + write!(f, "\n\n————— {y}-{m}-{d} —————\n")?; + last_date = Some(date); + } + } + + // Sender name + let sender_name = if em.direction == Direction::Outgoing { + self_name.clone() + } else { + em.sender_display_name.clone() + }; + + // Time + let time_str = dt + .map(|dt| dt.format("%H:%M").to_string()) + .unwrap_or_default(); + + // Content — use enriched snippet (already includes quote redaction) + let assets = &media_map[i]; + let content = if assets.is_empty() { + if !show_emoji && matches!(&msg.content, MessageContent::Emoji(_)) { + "[动画表情]".to_string() + } else { + em.snippet.clone() + } + } else { + let mut parts = Vec::new(); + for asset in assets { + let label = match asset.kind { + MediaKind::Image => { + image_counter += 1; + format!("图片{image_counter}") + } + MediaKind::Voice => { + voice_counter += 1; + format!("语音{voice_counter}") + } + MediaKind::Video => { + video_counter += 1; + format!("视频{video_counter}") + } + MediaKind::File => { + file_counter += 1; + format!("文件{file_counter}") + } + }; + parts.push(format!("{label}(可在附件中查看)")); + attachments.push((label, media_asset_path(&asset.filename))); + } + parts.join("\n") + }; + + write!(f, "\n{sender_name} {time_str}\n{content}\n")?; + } + + // Attachment list + if !attachments.is_empty() { + write!(f, "\n附件:\n")?; + for (label, path) in &attachments { + writeln!(f, "\n[{label}] {path}")?; + } + } + + f.flush()?; + Ok(()) +} + +// ── JSON writer ───────────────────────────────────────────────────── + +#[allow(clippy::too_many_arguments)] +fn write_json( + out_path: &std::path::Path, + talker: &str, + talker_display: &str, + is_group: bool, + messages: Vec, + media_map: &[Vec], + effective_limit: usize, + user_offset: usize, + total: usize, + scanned: usize, + skipped: usize, + shard_warnings: Vec, +) -> Result<(), Box> { + let message_count = messages.len(); + let time_range_start = messages.iter().map(|em| em.message.create_time).min(); + let time_range_end = messages.iter().map(|em| em.message.create_time).max(); + + let exported_items: Vec = messages + .into_iter() + .enumerate() + .map(|(i, enriched)| { + let media_files = media_asset_paths(&media_map[i]); + ExportedMessage { + enriched, + media_files, + } + }) + .collect(); + + let returned = exported_items.len(); + let has_more = user_offset + returned < total; + + let envelope = ExportEnvelope { + export_info: ExportInfo { + version: "1", + exported_at: chrono::Local::now().to_rfc3339(), + generator: format!("wx-cli {}", env!("CARGO_PKG_VERSION")), + }, + conversation: ConversationMeta { + talker: talker.to_string(), + display_name: talker_display.to_string(), + conv_type: if is_group { "group" } else { "private" }, + message_count, + time_range_start, + time_range_end, + }, + envelope: JsonEnvelope { + items: exported_items, + paging: PagingMeta { + limit: effective_limit, + offset: user_offset, + returned, + has_more, + total, + }, + stats: StatsMeta { + scanned, + skipped, + elapsed_ms: None, + shard_warnings, + }, + }, + }; + + let json = serde_json::to_string_pretty(&envelope)?; + std::fs::write(out_path, json)?; + Ok(()) +} + +// ── Helpers ───────────────────────────────────────────────────────── + +fn media_asset_path(filename: &str) -> String { + format!("media/{filename}") +} + +fn media_asset_paths(assets: &[crate::cmd::export_media::MediaAsset]) -> Vec { + assets + .iter() + .map(|asset| media_asset_path(&asset.filename)) + .collect() +} + +fn media_quality_hints(stats: &MediaStats) -> Vec { + let mut hints = Vec::new(); + + if stats.fallback_videos > 0 { + hints.push(format!( + "hint: {} video(s) resolved via directory scan fallback", + stats.fallback_videos + )); + } + if stats.fallback_files > 0 { + hints.push(format!( + "hint: {} file(s) resolved via directory scan fallback", + stats.fallback_files + )); + } + if stats.skipped_videos > 0 { + hints.push(format!( + "hint: {} video(s) skipped — not found after hardlink + directory scan", + stats.skipped_videos + )); + } + if stats.thumbnail_images > 0 { + hints.push(format!( + "hint: {} image(s) exported as thumbnail only — open them in WeChat to download full resolution", + stats.thumbnail_images + )); + } + if stats.silk_voices > 0 { + hints.push(format!( + "hint: {} voice(s) exported as raw SILK — install ffmpeg for MP3 transcoding", + stats.silk_voices + )); + } + if stats.skipped_files > 0 { + hints.push(format!( + "hint: {} file(s) skipped — not found after hardlink + directory scan", + stats.skipped_files + )); + } + if stats.wxgf_transcoded > 0 { + hints.push(format!( + "hint: {} wxgf image(s) transcoded to standard image format", + stats.wxgf_transcoded + )); + } + if stats.wxgf_fallback > 0 { + hints.push(format!( + "hint: {} wxgf image(s) kept as .wxgf - install ffmpeg for PNG/GIF export", + stats.wxgf_fallback + )); + } + + hints +} + +fn format_date_short(ts: i64) -> String { + chrono::DateTime::from_timestamp(ts, 0) + .map(|dt| { + let local = dt.with_timezone(&chrono::Local); + let d = local.date_naive(); + format!("{}-{}-{}", d.year(), d.month(), d.day()) + }) + .unwrap_or_else(|| ts.to_string()) +} + +use chrono::Datelike; + +#[cfg(test)] +mod tests { + use super::*; + use crate::cmd::export_media::{MediaAsset, MediaStats}; + use crate::schema::enrich_message; + use rusqlite::Connection; + + #[test] + fn media_asset_paths_follow_asset_filenames() { + let paths = media_asset_paths(&[ + MediaAsset { + kind: MediaKind::Image, + filename: "4865625c4e99e4d3b0959a0fe84f41cd.png".into(), + }, + MediaAsset { + kind: MediaKind::Image, + filename: "cdb2f853d5e1cdebbdc66bb8c80e1714.wxgf".into(), + }, + ]); + + assert_eq!( + paths, + vec![ + "media/4865625c4e99e4d3b0959a0fe84f41cd.png".to_string(), + "media/cdb2f853d5e1cdebbdc66bb8c80e1714.wxgf".to_string(), + ] + ); + } + + #[test] + fn media_quality_hints_include_wxgf_summary_lines() { + let hints = media_quality_hints(&MediaStats { + wxgf_transcoded: 2, + wxgf_fallback: 1, + ..MediaStats::default() + }); + + assert!(hints.contains(&"hint: 2 wxgf image(s) transcoded to standard image format".into())); + assert!(hints.contains( + &"hint: 1 wxgf image(s) kept as .wxgf - install ffmpeg for PNG/GIF export".into() + )); + } + + #[test] + fn write_json_uses_asset_filenames_for_media_files() { + let tmp = tempfile::TempDir::new().unwrap(); + let resolver = test_resolver(tmp.path()); + let out_path = tmp.path().join("export.json"); + + write_json( + &out_path, + "wxid_other", + "Alice", + false, + vec![sample_enriched_message(&resolver)], + &[vec![MediaAsset { + kind: MediaKind::Image, + filename: "4865625c4e99e4d3b0959a0fe84f41cd.png".into(), + }]], + 100, + 0, + 1, + 1, + 0, + vec![], + ) + .unwrap(); + + let json = std::fs::read_to_string(out_path).unwrap(); + assert!(json.contains("\"media_files\": [")); + assert!(json.contains("\"media/4865625c4e99e4d3b0959a0fe84f41cd.png\"")); + } + + #[test] + fn write_json_keeps_wxgf_media_paths_when_asset_filename_is_wxgf() { + let tmp = tempfile::TempDir::new().unwrap(); + let resolver = test_resolver(tmp.path()); + let out_path = tmp.path().join("export-wxgf.json"); + + write_json( + &out_path, + "wxid_other", + "Alice", + false, + vec![sample_enriched_message(&resolver)], + &[vec![MediaAsset { + kind: MediaKind::Image, + filename: "cdb2f853d5e1cdebbdc66bb8c80e1714.wxgf".into(), + }]], + 100, + 0, + 1, + 1, + 0, + vec![], + ) + .unwrap(); + + let json = std::fs::read_to_string(out_path).unwrap(); + assert!(json.contains("\"media/cdb2f853d5e1cdebbdc66bb8c80e1714.wxgf\"")); + } + + #[test] + fn write_txt_uses_asset_filenames_for_attachment_list() { + let tmp = tempfile::TempDir::new().unwrap(); + let resolver = test_resolver(tmp.path()); + let out_path = tmp.path().join("export.txt"); + + write_txt( + &out_path, + "wxid_other", + "Alice", + "wxid_me", + false, + &[sample_enriched_message(&resolver)], + &[vec![MediaAsset { + kind: MediaKind::Image, + filename: "cdb2f853d5e1cdebbdc66bb8c80e1714.wxgf".into(), + }]], + &resolver, + true, + ) + .unwrap(); + + let txt = std::fs::read_to_string(out_path).unwrap(); + assert!(txt.contains("图片1(可在附件中查看)")); + assert!(txt.contains("[图片1] media/cdb2f853d5e1cdebbdc66bb8c80e1714.wxgf")); + } + + fn sample_image_message() -> wx_db::Message { + wx_db::Message { + sort_seq: 1, + server_id: 1, + msg_type: wx_db::MSG_TYPE_IMAGE, + sub_type: 0, + sender: "wxid_other".into(), + talker: "wxid_other".into(), + create_time: 1_710_504_000, + content: wx_db::MessageContent::Image { + md5: Some("deadbeef".into()), + }, + status: 0, + } + } + + fn sample_enriched_message(resolver: &ContactResolver) -> EnrichedMessage { + enrich_message(sample_image_message(), "wxid_me", resolver) + } + + fn test_resolver(base: &std::path::Path) -> ContactResolver { + let contact_dir = base.join("contact"); + let session_dir = base.join("session"); + let message_dir = base.join("message"); + std::fs::create_dir_all(&contact_dir).unwrap(); + std::fs::create_dir_all(&session_dir).unwrap(); + std::fs::create_dir_all(&message_dir).unwrap(); + + let contact_conn = Connection::open(contact_dir.join("contact.db")).unwrap(); + contact_conn + .execute_batch( + "CREATE TABLE contact ( + username TEXT PRIMARY KEY, + alias TEXT DEFAULT '', + remark TEXT DEFAULT '', + nick_name TEXT DEFAULT '', + description TEXT DEFAULT NULL, + extra_buffer BLOB DEFAULT NULL + ); + CREATE TABLE contact_label ( + label_id_ TEXT, + label_name_ TEXT, + sort_order_ INTEGER + );", + ) + .unwrap(); + contact_conn + .execute( + "INSERT INTO contact (username, remark, nick_name) VALUES (?1, ?2, ?3)", + rusqlite::params!["wxid_me", "Me", ""], + ) + .unwrap(); + contact_conn + .execute( + "INSERT INTO contact (username, remark, nick_name) VALUES (?1, ?2, ?3)", + rusqlite::params!["wxid_other", "Alice", ""], + ) + .unwrap(); + + let session_conn = Connection::open(session_dir.join("session.db")).unwrap(); + session_conn + .execute_batch( + "CREATE TABLE SessionTable ( + username TEXT, + sort_timestamp INTEGER, + summary TEXT, + last_msg_type INTEGER, + last_msg_sender TEXT, + last_sender_display_name TEXT + );", + ) + .unwrap(); + + let db = wx_db::WechatDb::open(base).unwrap(); + ContactResolver::build(&db).unwrap() + } +} diff --git a/crates/wx-cli/src/cmd/export_media.rs b/crates/wx-cli/src/cmd/export_media.rs new file mode 100644 index 0000000..daef437 --- /dev/null +++ b/crates/wx-cli/src/cmd/export_media.rs @@ -0,0 +1,661 @@ +#[cfg(test)] +use std::collections::HashSet; +#[cfg(test)] +use std::path::PathBuf; + +#[cfg(test)] +use crate::util::{format_month, sanitize_filename}; +#[cfg(test)] +use wx_db::{Message, MessageContent}; +#[cfg(test)] +use wx_media::DatDecryptOptions; + +#[derive(Debug, Clone)] +pub struct MediaAsset { + pub kind: MediaKind, + pub filename: String, +} + +#[derive(Debug, Default)] +pub struct MediaStats { + pub skipped_videos: usize, + pub skipped_files: usize, + pub fallback_videos: usize, + pub fallback_files: usize, + pub thumbnail_images: usize, + pub silk_voices: usize, + pub wxgf_transcoded: usize, + pub wxgf_fallback: usize, +} + +#[derive(Debug, Clone, Copy)] +pub enum MediaKind { + Image, + Voice, + Video, + File, +} + +#[cfg(test)] +pub struct MediaBridge { + attach_dir: PathBuf, + media_dir: PathBuf, + file_dir: PathBuf, + video_dir: PathBuf, + hardlink_db: PathBuf, + output_media_dir: PathBuf, + dat_opts: DatDecryptOptions, + exported: HashSet, + xor_key_detected: bool, + pub stats: MediaStats, +} + +#[cfg(test)] +impl MediaBridge { + pub fn new( + attach_dir: PathBuf, + media_dir: PathBuf, + file_dir: PathBuf, + video_dir: PathBuf, + hardlink_db: PathBuf, + output_media_dir: PathBuf, + dat_opts: DatDecryptOptions, + ) -> Self { + Self { + attach_dir, + media_dir, + file_dir, + video_dir, + hardlink_db, + output_media_dir, + dat_opts, + exported: HashSet::new(), + xor_key_detected: false, + stats: MediaStats::default(), + } + } + + pub fn resolve(&mut self, msg: &Message, talker: &str) -> Vec { + match &msg.content { + MessageContent::Image { md5: Some(md5) } => self.resolve_image(md5, talker), + MessageContent::Voice => self.resolve_voice(msg.server_id), + MessageContent::Video { md5: Some(md5) } => self.resolve_video(md5, msg.create_time), + MessageContent::File { + md5: Some(md5), + title, + .. + } => self.resolve_file(md5, msg.create_time, title.as_deref()), + _ => vec![], + } + } + + fn resolve_image(&mut self, md5: &str, talker: &str) -> Vec { + // Lazily detect XOR key from talker's attach subdirectory + if !self.xor_key_detected { + let username_hash = format!("{:x}", wx_media::md5_hash(talker.as_bytes())); + let talker_attach = self.attach_dir.join(&username_hash); + if let Some(key) = wx_media::detect_xor_key(&talker_attach) { + self.dat_opts.xor_key = Some(key); + } + self.xor_key_detected = true; + } + + let lookup = match wx_media::resolve_image_by_md5(talker, &self.attach_dir, md5) { + Ok(r) => r, + Err(e) => { + eprintln!("warning: image resolve failed for md5={md5}: {e}"); + return vec![]; + } + }; + + let dat_path = match lookup.recommended { + Some(p) => p, + None => { + eprintln!("warning: no recommended .dat for md5={md5}"); + return vec![]; + } + }; + + // Detect if this is a thumbnail (_t.dat = ~9KB thumbnail, not the compressed _h or original) + let is_thumbnail = dat_path + .file_name() + .map(|n| n.to_string_lossy().contains("_t.")) + .unwrap_or(false); + + let data = match std::fs::read(&dat_path) { + Ok(d) => d, + Err(e) => { + eprintln!("warning: cannot read {}: {e}", dat_path.display()); + return vec![]; + } + }; + + let decoded = match wx_media::decrypt_dat(&data, &self.dat_opts) { + Ok(d) => d, + Err(e) => { + eprintln!("warning: decrypt_dat failed for md5={md5}: {e}"); + return vec![]; + } + }; + + let (image_data, image_ext, wxgf_transcoded, wxgf_fallback) = + export_image_bytes(decoded.data, &decoded.ext); + + let filename = format!("{}.{}", md5, image_ext); + if !self.exported.insert(filename.clone()) { + return vec![MediaAsset { + kind: MediaKind::Image, + filename, + }]; + } + + let out_path = self.output_media_dir.join(&filename); + if let Err(e) = std::fs::write(&out_path, &image_data) { + eprintln!("warning: cannot write {}: {e}", out_path.display()); + return vec![]; + } + + if is_thumbnail { + self.stats.thumbnail_images += 1; + } + if wxgf_transcoded { + self.stats.wxgf_transcoded += 1; + } + if wxgf_fallback { + self.stats.wxgf_fallback += 1; + } + + vec![MediaAsset { + kind: MediaKind::Image, + filename, + }] + } + + fn resolve_voice(&mut self, server_id: i64) -> Vec { + let svr_id = server_id.to_string(); + + let blob = match wx_media::extract_voice(&self.media_dir, &svr_id) { + Ok(b) => b, + Err(e) => { + eprintln!("warning: voice extract failed for svr_id={svr_id}: {e}"); + return vec![]; + } + }; + + let (data, ext) = match wx_media::transcode_silk_to_mp3(&blob.data) { + Ok(result) => { + if !result.transcoded { + self.stats.silk_voices += 1; + } + (result.data, result.ext.to_string()) + } + Err(e) => { + eprintln!("warning: voice transcode failed for svr_id={svr_id}: {e}"); + (blob.data, "silk".to_string()) + } + }; + + let filename = format!("{svr_id}.{ext}"); + if !self.exported.insert(filename.clone()) { + return vec![MediaAsset { + kind: MediaKind::Voice, + filename, + }]; + } + + let out_path = self.output_media_dir.join(&filename); + if let Err(e) = std::fs::write(&out_path, &data) { + eprintln!("warning: cannot write {}: {e}", out_path.display()); + return vec![]; + } + + vec![MediaAsset { + kind: MediaKind::Voice, + filename, + }] + } + + fn resolve_video(&mut self, md5: &str, create_time: i64) -> Vec { + let entries = match wx_media::query_hardlink(&self.hardlink_db, "video", md5) { + Ok(e) => e, + Err(e) => { + if !matches!(&e, wx_media::MediaError::NotFound(_)) { + eprintln!("warning: video hardlink query failed for md5={md5}: {e}"); + } + // Try directory scan fallback + let month = format_month(create_time); + return match wx_media::find_video_by_md5(&self.video_dir, md5, &month) { + Some(source) => self.copy_fallback_video(md5, &source), + None => { + self.stats.skipped_videos += 1; + vec![] + } + }; + } + }; + + let entry = match entries.first() { + Some(e) => e, + None => return vec![], + }; + + // Try candidate paths to find the physical video file + let candidates = [ + self.attach_dir + .join(&entry.dir1) + .join(&entry.dir2) + .join("Video") + .join(&entry.file_name), + self.attach_dir + .join(&entry.dir1) + .join(&entry.dir2) + .join(&entry.file_name), + self.attach_dir + .join(&entry.dir1) + .join("Video") + .join(&entry.file_name), + ]; + + let source = match candidates.iter().find(|p| p.exists()) { + Some(p) => p.clone(), + None => { + // Hardlink entry exists but physical file missing; try directory scan + let month = format_month(create_time); + return match wx_media::find_video_by_md5(&self.video_dir, md5, &month) { + Some(source) => self.copy_fallback_video(md5, &source), + None => { + self.stats.skipped_videos += 1; + vec![] + } + }; + } + }; + + let filename = entry.file_name.clone(); + if !self.exported.insert(filename.clone()) { + return vec![MediaAsset { + kind: MediaKind::Video, + filename, + }]; + } + + let out_path = self.output_media_dir.join(&filename); + if let Err(e) = std::fs::copy(&source, &out_path) { + eprintln!("warning: cannot copy video {}: {e}", source.display()); + return vec![]; + } + + vec![MediaAsset { + kind: MediaKind::Video, + filename, + }] + } + + fn copy_fallback_video(&mut self, md5: &str, source: &std::path::Path) -> Vec { + let filename = format!("{md5}.mp4"); + if !self.exported.insert(filename.clone()) { + self.stats.fallback_videos += 1; + return vec![MediaAsset { + kind: MediaKind::Video, + filename, + }]; + } + + let out_path = self.output_media_dir.join(&filename); + if let Err(e) = std::fs::copy(source, &out_path) { + eprintln!( + "warning: cannot copy fallback video {}: {e}", + source.display() + ); + return vec![]; + } + + self.stats.fallback_videos += 1; + vec![MediaAsset { + kind: MediaKind::Video, + filename, + }] + } + + fn resolve_file( + &mut self, + md5: &str, + create_time: i64, + title: Option<&str>, + ) -> Vec { + let entries = match wx_media::query_hardlink(&self.hardlink_db, "file", md5) { + Ok(e) => e, + Err(e) => { + if !matches!(&e, wx_media::MediaError::NotFound(_)) { + eprintln!("warning: file hardlink query failed for md5={md5}: {e}"); + } + return self.try_file_fallback(md5, create_time, title); + } + }; + + let entry = match entries.first() { + Some(e) => e, + None => { + return self.try_file_fallback(md5, create_time, title); + } + }; + + // Try candidate paths to find the physical file + let candidates = [ + self.file_dir + .join(&entry.dir1) + .join(&entry.dir2) + .join(&entry.file_name), + self.file_dir.join(&entry.dir1).join(&entry.file_name), + ]; + + let source = match candidates.iter().find(|p| p.exists()) { + Some(p) => p.clone(), + None => { + // Hardlink entry exists but physical file missing; try directory scan + return self.try_file_fallback(md5, create_time, title); + } + }; + + let filename = format!("{}_{}", md5, entry.file_name); + if !self.exported.insert(filename.clone()) { + return vec![MediaAsset { + kind: MediaKind::File, + filename, + }]; + } + + let out_path = self.output_media_dir.join(&filename); + if let Err(e) = std::fs::copy(&source, &out_path) { + eprintln!("warning: cannot copy file {}: {e}", source.display()); + return vec![]; + } + + vec![MediaAsset { + kind: MediaKind::File, + filename, + }] + } + + fn try_file_fallback( + &mut self, + md5: &str, + create_time: i64, + title: Option<&str>, + ) -> Vec { + if let Some(t) = title { + let month = format_month(create_time); + if let Some(source) = wx_media::find_file_by_name(&self.file_dir, t, &month) { + let basename = std::path::Path::new(t) + .file_name() + .map(|n| n.to_string_lossy().to_string()) + .unwrap_or_else(|| t.to_string()); + let safe_name = sanitize_filename(&basename); + let filename = format!("{md5}_{safe_name}"); + if !self.exported.insert(filename.clone()) { + self.stats.fallback_files += 1; + return vec![MediaAsset { + kind: MediaKind::File, + filename, + }]; + } + + let out_path = self.output_media_dir.join(&filename); + if let Err(e) = std::fs::copy(&source, &out_path) { + eprintln!( + "warning: cannot copy fallback file {}: {e}", + source.display() + ); + return vec![]; + } + + self.stats.fallback_files += 1; + return vec![MediaAsset { + kind: MediaKind::File, + filename, + }]; + } + } + self.stats.skipped_files += 1; + vec![] + } +} + +pub fn export_image_bytes(decoded_data: Vec, decoded_ext: &str) -> (Vec, String, bool, bool) { + if decoded_ext != "wxgf" { + return (decoded_data, decoded_ext.to_string(), false, false); + } + + match wx_media::transcode_wxgf(&decoded_data) { + Ok(transcoded) if transcoded.transcoded => { + (transcoded.data, transcoded.ext.to_string(), true, false) + } + Ok(_) => (decoded_data, "wxgf".to_string(), false, true), + Err(e) => { + eprintln!("warning: wxgf image export kept as .wxgf due to transcode error: {e}"); + (decoded_data, "wxgf".to_string(), false, true) + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::Mutex; + + static FFMPEG_ENV_LOCK: Mutex<()> = Mutex::new(()); + + #[test] + fn test_resolve_video_fallback_when_no_hardlink_db() { + let tmp = tempfile::TempDir::new().unwrap(); + let root = tmp.path(); + + // Setup video_dir with a video file in 2024-03 + let video_dir = root.join("video"); + let month_dir = video_dir.join("2024-03"); + std::fs::create_dir_all(&month_dir).unwrap(); + std::fs::write(month_dir.join("deadbeef.mp4"), b"fake-video").unwrap(); + + // Setup output media dir + let output_media = root.join("output"); + std::fs::create_dir_all(&output_media).unwrap(); + + // Use a nonexistent hardlink DB so query_hardlink will fail → triggers fallback + let mut bridge = MediaBridge::new( + root.join("attach"), // attach_dir (unused for this test) + root.join("media"), // media_dir (unused) + root.join("file"), // file_dir (unused) + video_dir, // video_dir + root.join("nonexistent.db"), // hardlink_db (will fail) + output_media.clone(), + wx_media::DatDecryptOptions::default(), + ); + + // create_time = 2024-03-15T12:00:00Z = 1710504000 + let msg = Message { + sort_seq: 0, + server_id: 1, + msg_type: 43, + sub_type: 0, + sender: "wxid_test".into(), + talker: "wxid_other".into(), + create_time: 1710504000, + content: MessageContent::Video { + md5: Some("deadbeef".into()), + }, + status: 0, + }; + + let assets = bridge.resolve(&msg, "wxid_other"); + assert_eq!(assets.len(), 1); + assert!(matches!(assets[0].kind, MediaKind::Video)); + assert_eq!(assets[0].filename, "deadbeef.mp4"); + assert_eq!(bridge.stats.fallback_videos, 1); + assert!(output_media.join("deadbeef.mp4").exists()); + } + + #[test] + fn test_resolve_image_transcodes_embedded_wxgf_to_standard_image() { + let tmp = tempfile::TempDir::new().unwrap(); + let root = tmp.path(); + let talker = "wxid_other"; + let md5 = "4865625c4e99e4d3b0959a0fe84f41cd"; + let xor_key = 0xa5; + let wxgf = sample_embedded_png_wxgf(); + + write_xor_dat(root, talker, md5, &wxgf, xor_key); + + let output_media = root.join("output"); + std::fs::create_dir_all(&output_media).unwrap(); + + let mut bridge = MediaBridge::new( + root.join("attach"), + root.join("media"), + root.join("file"), + root.join("video"), + root.join("hardlink.db"), + output_media.clone(), + wx_media::DatDecryptOptions { + v2_aes_key: None, + xor_key: Some(xor_key), + }, + ); + + let assets = bridge.resolve_image(md5, talker); + assert_eq!(assets.len(), 1); + assert_eq!(assets[0].filename, format!("{md5}.png")); + assert_eq!(bridge.stats.wxgf_transcoded, 1); + assert_eq!(bridge.stats.wxgf_fallback, 0); + assert_eq!( + std::fs::read(output_media.join(format!("{md5}.png"))).unwrap(), + sample_png() + ); + } + + #[test] + fn test_resolve_image_keeps_wxgf_when_hevc_cannot_be_transcoded() { + let _guard = FFMPEG_ENV_LOCK.lock().unwrap(); + unsafe { + std::env::set_var("FFMPEG_PATH", "/definitely-missing-ffmpeg"); + } + wx_media::reset_ffmpeg_cache(); + + let tmp = tempfile::TempDir::new().unwrap(); + let root = tmp.path(); + let talker = "wxid_other"; + let md5 = "cdb2f853d5e1cdebbdc66bb8c80e1714"; + let xor_key = 0xa5; + let wxgf = sample_hevc_wxgf(); + + write_xor_dat(root, talker, md5, &wxgf, xor_key); + + let output_media = root.join("output"); + std::fs::create_dir_all(&output_media).unwrap(); + + let mut bridge = MediaBridge::new( + root.join("attach"), + root.join("media"), + root.join("file"), + root.join("video"), + root.join("hardlink.db"), + output_media.clone(), + wx_media::DatDecryptOptions { + v2_aes_key: None, + xor_key: Some(xor_key), + }, + ); + + let assets = bridge.resolve_image(md5, talker); + assert_eq!(assets.len(), 1); + assert_eq!(assets[0].filename, format!("{md5}.wxgf")); + assert_eq!(bridge.stats.wxgf_transcoded, 0); + assert_eq!(bridge.stats.wxgf_fallback, 1); + assert_eq!( + std::fs::read(output_media.join(format!("{md5}.wxgf"))).unwrap(), + wxgf + ); + + unsafe { + std::env::remove_var("FFMPEG_PATH"); + } + wx_media::reset_ffmpeg_cache(); + } + + #[test] + fn test_resolve_image_keeps_wxgf_when_transcode_errors() { + let tmp = tempfile::TempDir::new().unwrap(); + let root = tmp.path(); + let talker = "wxid_other"; + let md5 = "0badf00d0badf00d0badf00d0badf00d"; + let xor_key = 0xa5; + let wxgf = b"wxgf".to_vec(); + + write_xor_dat(root, talker, md5, &wxgf, xor_key); + + let output_media = root.join("output"); + std::fs::create_dir_all(&output_media).unwrap(); + + let mut bridge = MediaBridge::new( + root.join("attach"), + root.join("media"), + root.join("file"), + root.join("video"), + root.join("hardlink.db"), + output_media.clone(), + wx_media::DatDecryptOptions { + v2_aes_key: None, + xor_key: Some(xor_key), + }, + ); + + let assets = bridge.resolve_image(md5, talker); + assert_eq!(assets.len(), 1); + assert_eq!(assets[0].filename, format!("{md5}.wxgf")); + assert_eq!(bridge.stats.wxgf_transcoded, 0); + assert_eq!(bridge.stats.wxgf_fallback, 1); + assert_eq!( + std::fs::read(output_media.join(format!("{md5}.wxgf"))).unwrap(), + wxgf + ); + } + + fn write_xor_dat( + root: &std::path::Path, + talker: &str, + md5: &str, + plaintext: &[u8], + xor_key: u8, + ) { + let username_hash = format!("{:x}", wx_media::md5_hash(talker.as_bytes())); + let img_dir = root + .join("attach") + .join(username_hash) + .join("2026-03") + .join("Img"); + std::fs::create_dir_all(&img_dir).unwrap(); + let encrypted: Vec = plaintext.iter().map(|b| b ^ xor_key).collect(); + std::fs::write(img_dir.join(format!("{md5}.dat")), encrypted).unwrap(); + } + + fn sample_embedded_png_wxgf() -> Vec { + let mut wxgf = b"wxgfmetadata".to_vec(); + wxgf.extend_from_slice(&sample_png()); + wxgf + } + + fn sample_hevc_wxgf() -> Vec { + let mut wxgf = b"wxgfmetadata".to_vec(); + wxgf.extend_from_slice(&[0x00, 0x00, 0x00, 0x01, 0x26, 0x01, 0x02, 0x03, 0x04]); + wxgf + } + + fn sample_png() -> Vec { + vec![ + 0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A, 0x00, 0x00, 0x00, 0x0D, 0x49, 0x48, + 0x44, 0x52, 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01, 0x08, 0x06, 0x00, 0x00, + 0x00, 0x1F, 0x15, 0xC4, 0x89, 0x00, 0x00, 0x00, 0x0D, 0x49, 0x44, 0x41, 0x54, 0x78, + 0x9C, 0x63, 0xF8, 0xCF, 0xC0, 0xF0, 0x1F, 0x00, 0x05, 0x00, 0x01, 0xFF, 0x89, 0x99, + 0x3D, 0x1D, 0x00, 0x00, 0x00, 0x00, 0x49, 0x45, 0x4E, 0x44, 0xAE, 0x42, 0x60, 0x82, + ] + } +} diff --git a/crates/wx-cli/src/cmd/export_task.rs b/crates/wx-cli/src/cmd/export_task.rs new file mode 100644 index 0000000..de3a3a1 --- /dev/null +++ b/crates/wx-cli/src/cmd/export_task.rs @@ -0,0 +1,1602 @@ +use std::cell::RefCell; +use std::collections::{HashMap, HashSet}; +use std::path::{Path, PathBuf}; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::{Arc, Mutex}; + +use rayon::prelude::*; +use rusqlite::Connection; + +use crate::cmd::export_media::{export_image_bytes, MediaAsset, MediaKind, MediaStats}; +use crate::schema::EnrichedMessage; +use crate::util::{format_month, sanitize_filename}; +use wx_db::MessageContent; +use wx_media::DatDecryptOptions; + +// --------------------------------------------------------------------------- +// Core types +// --------------------------------------------------------------------------- + +/// Pre-scanned typed task descriptor produced by the classify stage. +#[derive(Debug, Clone)] +pub enum MediaTask { + Image { + md5: String, + msg_index: usize, + }, + Voice { + server_id: i64, + msg_index: usize, + }, + Video { + md5: String, + create_time: i64, + msg_index: usize, + }, + File { + md5: String, + create_time: i64, + title: Option, + msg_index: usize, + }, +} + +impl MediaTask { + pub fn kind(&self) -> TaskKind { + match self { + MediaTask::Image { .. } => TaskKind::Image, + MediaTask::Voice { .. } => TaskKind::Voice, + MediaTask::Video { .. } => TaskKind::Video, + MediaTask::File { .. } => TaskKind::File, + } + } + + pub fn msg_index(&self) -> usize { + match self { + MediaTask::Image { msg_index, .. } => *msg_index, + MediaTask::Voice { msg_index, .. } => *msg_index, + MediaTask::Video { msg_index, .. } => *msg_index, + MediaTask::File { msg_index, .. } => *msg_index, + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub enum TaskKind { + Image, + Voice, + Video, + File, +} + +/// Result of resolving a single task. +#[derive(Debug)] +pub struct ResolvedAsset { + pub msg_index: usize, + pub asset: Option, + pub tags: Vec, + pub error: Option, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum TaskTag { + ThumbnailImage, + SilkVoice, + WxgfTranscoded, + WxgfFallback, + FallbackVideo, + FallbackFile, + SkippedVideo, + SkippedFile, +} + +impl TaskTag { + /// Whether this tag should be counted for duplicate messages. + /// + /// Matches old `MediaBridge` behavior: + /// - Image tags (Thumbnail, WxgfTranscoded, WxgfFallback): NOT counted for duplicates + /// because old code returned early when `exported.insert()` failed (before counting). + /// - All other tags: counted for duplicates because old code counted them + /// before or regardless of the `exported.insert()` check. + pub fn counts_on_duplicate(self) -> bool { + match self { + TaskTag::ThumbnailImage => false, + TaskTag::WxgfTranscoded => false, + TaskTag::WxgfFallback => false, + TaskTag::SilkVoice => true, + TaskTag::FallbackVideo => true, + TaskTag::FallbackFile => true, + TaskTag::SkippedVideo => true, + TaskTag::SkippedFile => true, + } + } +} + +/// Structured error for a single failed task. +#[derive(Debug)] +pub struct ExportError { + pub task_kind: &'static str, + pub key: String, + pub reason: String, +} + +/// Aggregated error summary with grouped reporting. +#[derive(Debug, Default)] +pub struct ErrorSummary { + pub errors: Vec, +} + +impl ErrorSummary { + pub fn print_report(&self) { + if self.errors.is_empty() { + return; + } + let mut groups: HashMap<&str, Vec<&ExportError>> = HashMap::new(); + for e in &self.errors { + groups.entry(e.task_kind).or_default().push(e); + } + for (kind, errs) in groups { + eprintln!("media errors [{kind}]: {} failure(s)", errs.len()); + for e in errs { + eprintln!(" - {}: {}", e.key, e.reason); + } + } + } +} + +// --------------------------------------------------------------------------- +// Write gate — thread-safe output filename dedup +// --------------------------------------------------------------------------- + +pub struct WriteGate { + written: Mutex>, +} + +impl WriteGate { + pub fn new() -> Self { + Self { + written: Mutex::new(HashSet::new()), + } + } + + /// Try to claim a filename. Returns `true` if this thread should write. + pub fn claim(&self, filename: &str) -> bool { + self.written.lock().unwrap().insert(filename.to_string()) + } +} + +// --------------------------------------------------------------------------- +// Thread-local connection pools +// --------------------------------------------------------------------------- + +/// Per-thread connection pool for voice media_*.db files. +pub struct VoiceConnectionPool { + db_paths: Vec, + path_key: u64, +} + +impl VoiceConnectionPool { + pub fn new(media_dir: &Path) -> Self { + let db_paths = wx_media::find_media_dbs(media_dir).unwrap_or_default(); + let path_key = Self::compute_path_key(&db_paths); + Self { db_paths, path_key } + } + + fn compute_path_key(paths: &[PathBuf]) -> u64 { + use std::collections::hash_map::DefaultHasher; + use std::hash::{Hash, Hasher}; + let mut hasher = DefaultHasher::new(); + let mut sorted: Vec<&Path> = paths.iter().map(|p| p.as_path()).collect(); + sorted.sort(); + for p in &sorted { + p.hash(&mut hasher); + } + hasher.finish() + } + + fn open_all(&self) -> Vec { + let mut conns = Vec::new(); + for path in &self.db_paths { + if let Ok(conn) = Connection::open_with_flags( + path, + rusqlite::OpenFlags::SQLITE_OPEN_READ_ONLY, + ) { + conns.push(conn); + } + } + conns + } + + pub fn with_connections(&self, f: impl FnOnce(&[Connection]) -> R) -> R { + thread_local! { + static CONNS: RefCell)>> = RefCell::new(None); + } + CONNS.with(|cell| { + let mut borrow = cell.borrow_mut(); + if let Some((key, conns)) = borrow.as_ref() { + if *key == self.path_key { + return f(conns); + } + } + let conns = self.open_all(); + *borrow = Some((self.path_key, conns)); + f(borrow.as_ref().unwrap().1.as_slice()) + }) + } +} + +/// Per-thread connection pool for hardlink.db. +pub struct HardlinkConnectionPool { + db_path: PathBuf, + path_key: u64, +} + +impl HardlinkConnectionPool { + pub fn new(db_path: PathBuf) -> Self { + use std::collections::hash_map::DefaultHasher; + use std::hash::{Hash, Hasher}; + let mut hasher = DefaultHasher::new(); + db_path.hash(&mut hasher); + let path_key = hasher.finish(); + Self { db_path, path_key } + } + + fn open(&self) -> Option { + Connection::open_with_flags( + &self.db_path, + rusqlite::OpenFlags::SQLITE_OPEN_READ_ONLY, + ) + .ok() + } + + pub fn with_connection(&self, f: impl FnOnce(&Connection) -> R) -> Option { + thread_local! { + static CONN: RefCell> = RefCell::new(None); + } + CONN.with(|cell| { + let mut borrow = cell.borrow_mut(); + if let Some((key, conn)) = borrow.as_ref() { + if *key == self.path_key { + return Some(f(conn)); + } + } + let conn = self.open()?; + *borrow = Some((self.path_key, conn)); + Some(f(&borrow.as_ref().unwrap().1)) + }) + } +} + +// --------------------------------------------------------------------------- +// Shared context — immutable, Arc-shared across rayon threads +// --------------------------------------------------------------------------- + +pub struct SharedContext { + pub attach_dir: PathBuf, + pub media_dir: PathBuf, + pub file_dir: PathBuf, + pub video_dir: PathBuf, + pub output_media_dir: PathBuf, + pub dat_opts: DatDecryptOptions, + pub talker: String, + pub voice_chat_name_id_hint: Arc>>, + pub voice_pool: VoiceConnectionPool, + pub hardlink_pool: HardlinkConnectionPool, + pub write_gate: WriteGate, +} + +// --------------------------------------------------------------------------- +// DupMap — dedup tracking +// --------------------------------------------------------------------------- + +pub struct DupMap { + /// (duplicate msg_index, canonical msg_index) + pub duplicates: Vec<(usize, usize)>, +} + +// --------------------------------------------------------------------------- +// Pipeline functions +// --------------------------------------------------------------------------- + +/// Build shared context from account/session info (pre-compute stage). +pub fn build_shared_context( + attach_dir: PathBuf, + media_dir: PathBuf, + file_dir: PathBuf, + video_dir: PathBuf, + hardlink_db: PathBuf, + output_media_dir: PathBuf, + talker: &str, + dat_opts: DatDecryptOptions, +) -> SharedContext { + // Pre-detect XOR key + let mut dat_opts = dat_opts; + let username_hash = format!("{:x}", wx_media::md5_hash(talker.as_bytes())); + let talker_attach = attach_dir.join(&username_hash); + if let Some(key) = wx_media::detect_xor_key(&talker_attach) { + dat_opts.xor_key = Some(key); + } + + // Pre-cache ffmpeg availability (OnceLock, one-time check) + let _ = wx_media::ffmpeg_available(); + + SharedContext { + voice_chat_name_id_hint: Arc::new(Mutex::new(None)), + voice_pool: VoiceConnectionPool::new(&media_dir), + hardlink_pool: HardlinkConnectionPool::new(hardlink_db), + write_gate: WriteGate::new(), + attach_dir, + media_dir, + file_dir, + video_dir, + output_media_dir, + dat_opts, + talker: talker.to_string(), + } +} + +fn update_voice_chat_name_id_hint(ctx: &SharedContext, blob: &wx_media::VoiceBlob) { + if let Some(chat_name_id) = blob.chat_name_id { + if let Ok(mut hint) = ctx.voice_chat_name_id_hint.lock() { + *hint = Some(chat_name_id); + } + } +} + +/// Stage 1: Classify messages into typed tasks. +pub fn classify(messages: &[EnrichedMessage]) -> Vec { + let mut tasks = Vec::new(); + for (idx, em) in messages.iter().enumerate() { + match &em.message.content { + MessageContent::Image { md5: Some(md5) } => { + tasks.push(MediaTask::Image { + md5: md5.clone(), + msg_index: idx, + }); + } + MessageContent::Voice => { + tasks.push(MediaTask::Voice { + server_id: em.message.server_id, + msg_index: idx, + }); + } + MessageContent::Video { md5: Some(md5) } => { + tasks.push(MediaTask::Video { + md5: md5.clone(), + create_time: em.message.create_time, + msg_index: idx, + }); + } + MessageContent::File { + md5: Some(md5), + title, + .. + } => { + tasks.push(MediaTask::File { + md5: md5.clone(), + create_time: em.message.create_time, + title: title.clone(), + msg_index: idx, + }); + } + _ => {} + } + } + tasks +} + +/// Stage 2: Deduplicate tasks by content key. +/// +/// Image/voice are safely deduped (md5/server_id fully determines output). +/// Video/file are NOT deduped when they have different fallback parameters +/// (create_time/title), since different parameters may hit different source files. +pub fn dedup(tasks: Vec) -> (Vec, DupMap) { + let mut canonical: HashMap = HashMap::new(); + let mut duplicates: Vec<(usize, usize)> = Vec::new(); + let mut unique: Vec = Vec::new(); + + for task in tasks { + let msg_idx = task.msg_index(); + let (key, can_dedup) = match &task { + MediaTask::Image { md5, .. } => (format!("img:{md5}"), true), + MediaTask::Voice { server_id, .. } => (format!("voi:{server_id}"), true), + MediaTask::Video { + md5, create_time, .. + } => (format!("vid:{md5}:{create_time}"), false), + MediaTask::File { + md5, + create_time, + title, + .. + } => { + let t = title.as_deref().unwrap_or(""); + (format!("fil:{md5}:{create_time}:{t}"), false) + } + }; + + if can_dedup { + if let Some(&canonical_msg_idx) = canonical.get(&key) { + duplicates.push((msg_idx, canonical_msg_idx)); + continue; + } + } else { + // For video/file: only dedup if the key is identical + // (same md5 AND same fallback params) + if let Some(&canonical_msg_idx) = canonical.get(&key) { + duplicates.push((msg_idx, canonical_msg_idx)); + continue; + } + } + + canonical.insert(key, msg_idx); + unique.push(task); + } + + ( + unique, + DupMap { + duplicates, + }, + ) +} + +/// Default rayon thread pool size: min(num_cpus, 4). +fn rayon_default_threads() -> usize { + let cpus = std::thread::available_parallelism() + .map(|n| n.get()) + .unwrap_or(1); + cpus.min(4) +} + +/// Stage 3-4: Parallel resolve. +/// +/// Tasks are batched by type, each batch runs in parallel within a rayon thread pool. +/// Progress is reported per-type at ~10% intervals. +pub fn resolve_parallel( + tasks: Vec, + ctx: Arc, + parallel: Option, +) -> (Vec, ErrorSummary) { + let num_threads = parallel.unwrap_or_else(rayon_default_threads).max(1); + let pool = rayon::ThreadPoolBuilder::new() + .num_threads(num_threads) + .build() + .unwrap(); + + // Group by kind + let mut batches: HashMap> = HashMap::new(); + for task in tasks { + batches.entry(task.kind()).or_default().push(task); + } + + let order = [TaskKind::Image, TaskKind::Voice, TaskKind::Video, TaskKind::File]; + let mut all_results = Vec::new(); + let mut all_errors = ErrorSummary::default(); + + for kind in order { + let batch = match batches.remove(&kind) { + Some(b) => b, + None => continue, + }; + let total = batch.len(); + if total == 0 { + continue; + } + + let counter = AtomicUsize::new(0); + let kind_label = match kind { + TaskKind::Image => "image", + TaskKind::Voice => "voice", + TaskKind::Video => "video", + TaskKind::File => "file", + }; + + let results: Vec = pool.install(|| { + batch + .par_iter() + .map(|task| { + let result = resolve_one(task, &ctx); + let done = counter.fetch_add(1, Ordering::Relaxed) + 1; + let prev_threshold = (done - 1) * 10 / total; + let cur_threshold = done * 10 / total; + if cur_threshold != prev_threshold || done == total { + eprintln!("media: {kind_label} {done}/{total}"); + } + result + }) + .collect() + }); + + for r in results { + if let Some(e) = &r.error { + all_errors.errors.push(ExportError { + task_kind: kind_label, + key: e.key.clone(), + reason: e.reason.clone(), + }); + } + all_results.push(ResolvedAsset { + msg_index: r.msg_index, + asset: r.asset, + tags: r.tags, + error: None, // errors collected separately + }); + } + } + + (all_results, all_errors) +} + +/// Resolve a single task. +fn resolve_one(task: &MediaTask, ctx: &SharedContext) -> ResolvedAsset { + match task { + MediaTask::Image { md5, msg_index } => resolve_image(md5, *msg_index, ctx), + MediaTask::Voice { server_id, msg_index } => resolve_voice(*server_id, *msg_index, ctx), + MediaTask::Video { + md5, + create_time, + msg_index, + } => resolve_video(md5, *create_time, *msg_index, ctx), + MediaTask::File { + md5, + create_time, + title, + msg_index, + } => resolve_file(md5, *create_time, title.as_deref(), *msg_index, ctx), + } +} + +fn resolve_image(md5: &str, msg_index: usize, ctx: &SharedContext) -> ResolvedAsset { + let lookup = match wx_media::resolve_image_by_md5(&ctx.talker, &ctx.attach_dir, md5) { + Ok(r) => r, + Err(e) => { + return ResolvedAsset { + msg_index, + asset: None, + tags: vec![], + error: Some(ExportError { + task_kind: "image", + key: md5.to_string(), + reason: format!("resolve failed: {e}"), + }), + }; + } + }; + + let dat_path = match lookup.recommended { + Some(p) => p, + None => { + return ResolvedAsset { + msg_index, + asset: None, + tags: vec![], + error: Some(ExportError { + task_kind: "image", + key: md5.to_string(), + reason: "no recommended .dat".to_string(), + }), + }; + } + }; + + let is_thumbnail = dat_path + .file_name() + .map(|n| n.to_string_lossy().contains("_t.")) + .unwrap_or(false); + + let data = match std::fs::read(&dat_path) { + Ok(d) => d, + Err(e) => { + return ResolvedAsset { + msg_index, + asset: None, + tags: vec![], + error: Some(ExportError { + task_kind: "image", + key: md5.to_string(), + reason: format!("read {}: {e}", dat_path.display()), + }), + }; + } + }; + + let decoded = match wx_media::decrypt_dat(&data, &ctx.dat_opts) { + Ok(d) => d, + Err(e) => { + return ResolvedAsset { + msg_index, + asset: None, + tags: vec![], + error: Some(ExportError { + task_kind: "image", + key: md5.to_string(), + reason: format!("decrypt: {e}"), + }), + }; + } + }; + + let (image_data, image_ext, wxgf_transcoded, wxgf_fallback) = + export_image_bytes(decoded.data, &decoded.ext); + + let filename = format!("{md5}.{image_ext}"); + if ctx.write_gate.claim(&filename) { + let out_path = ctx.output_media_dir.join(&filename); + if let Err(e) = std::fs::write(&out_path, &image_data) { + return ResolvedAsset { + msg_index, + asset: None, + tags: vec![], + error: Some(ExportError { + task_kind: "image", + key: md5.to_string(), + reason: format!("write {}: {e}", out_path.display()), + }), + }; + } + } + + let mut tags = vec![]; + if is_thumbnail { + tags.push(TaskTag::ThumbnailImage); + } + if wxgf_transcoded { + tags.push(TaskTag::WxgfTranscoded); + } + if wxgf_fallback { + tags.push(TaskTag::WxgfFallback); + } + + ResolvedAsset { + msg_index, + asset: Some(MediaAsset { + kind: MediaKind::Image, + filename, + }), + tags, + error: None, + } +} + +fn resolve_voice(server_id: i64, msg_index: usize, ctx: &SharedContext) -> ResolvedAsset { + let svr_id = server_id.to_string(); + let chat_name_id_hint = ctx + .voice_chat_name_id_hint + .lock() + .ok() + .and_then(|hint| *hint); + + let blob = ctx.voice_pool.with_connections(|conns| { + for conn in conns { + if let Ok(b) = wx_media::extract_voice_with_conn_hint(conn, &svr_id, chat_name_id_hint) { + return Some(b); + } + } + None + }); + + let blob = match blob { + Some(b) => b, + None => { + // Fallback to opening fresh connections + match wx_media::extract_voice(&ctx.media_dir, &svr_id) { + Ok(b) => b, + Err(e) => { + return ResolvedAsset { + msg_index, + asset: None, + tags: vec![], + error: Some(ExportError { + task_kind: "voice", + key: svr_id, + reason: format!("extract failed: {e}"), + }), + }; + } + } + } + }; + + update_voice_chat_name_id_hint(ctx, &blob); + + let (data, ext, is_silk) = match wx_media::transcode_silk_to_mp3(&blob.data) { + Ok(result) => { + let is_silk = !result.transcoded; + (result.data, result.ext.to_string(), is_silk) + } + Err(e) => { + eprintln!("warning: voice transcode failed for svr_id={svr_id}: {e}"); + (blob.data, "silk".to_string(), true) + } + }; + + let filename = format!("{svr_id}.{ext}"); + if ctx.write_gate.claim(&filename) { + let out_path = ctx.output_media_dir.join(&filename); + if let Err(e) = std::fs::write(&out_path, &data) { + return ResolvedAsset { + msg_index, + asset: None, + tags: vec![], + error: Some(ExportError { + task_kind: "voice", + key: svr_id, + reason: format!("write {}: {e}", out_path.display()), + }), + }; + } + } + + let mut tags = vec![]; + if is_silk { + tags.push(TaskTag::SilkVoice); + } + + ResolvedAsset { + msg_index, + asset: Some(MediaAsset { + kind: MediaKind::Voice, + filename, + }), + tags, + error: None, + } +} + +fn resolve_video( + md5: &str, + create_time: i64, + msg_index: usize, + ctx: &SharedContext, +) -> ResolvedAsset { + // Try hardlink DB first + let hardlink_result = ctx.hardlink_pool.with_connection(|conn| { + wx_media::query_hardlink_with_conn(conn, "video", md5) + }); + + let entries = match hardlink_result { + Some(Ok(e)) => Some(e), + Some(Err(e)) => { + if !matches!(&e, wx_media::MediaError::NotFound(_)) { + eprintln!("warning: video hardlink query failed for md5={md5}: {e}"); + } + None + } + None => None, + }; + + if let Some(entries) = entries { + if let Some(entry) = entries.first() { + let candidates = [ + ctx.attach_dir + .join(&entry.dir1) + .join(&entry.dir2) + .join("Video") + .join(&entry.file_name), + ctx.attach_dir + .join(&entry.dir1) + .join(&entry.dir2) + .join(&entry.file_name), + ctx.attach_dir + .join(&entry.dir1) + .join("Video") + .join(&entry.file_name), + ]; + + if let Some(source) = candidates.iter().find(|p| p.exists()) { + let filename = entry.file_name.clone(); + if ctx.write_gate.claim(&filename) { + let out_path = ctx.output_media_dir.join(&filename); + if let Err(e) = std::fs::copy(source, &out_path) { + return ResolvedAsset { + msg_index, + asset: None, + tags: vec![], + error: Some(ExportError { + task_kind: "video", + key: md5.to_string(), + reason: format!("copy {}: {e}", source.display()), + }), + }; + } + } + return ResolvedAsset { + msg_index, + asset: Some(MediaAsset { + kind: MediaKind::Video, + filename, + }), + tags: vec![], + error: None, + }; + } + } + } + + // Fallback: directory scan + let month = format_month(create_time); + match wx_media::find_video_by_md5(&ctx.video_dir, md5, &month) { + Some(source) => { + let filename = format!("{md5}.mp4"); + if ctx.write_gate.claim(&filename) { + let out_path = ctx.output_media_dir.join(&filename); + if let Err(e) = std::fs::copy(&source, &out_path) { + return ResolvedAsset { + msg_index, + asset: None, + tags: vec![], + error: Some(ExportError { + task_kind: "video", + key: md5.to_string(), + reason: format!("copy fallback {}: {e}", source.display()), + }), + }; + } + } + ResolvedAsset { + msg_index, + asset: Some(MediaAsset { + kind: MediaKind::Video, + filename, + }), + tags: vec![TaskTag::FallbackVideo], + error: None, + } + } + None => ResolvedAsset { + msg_index, + asset: None, + tags: vec![TaskTag::SkippedVideo], + error: None, + }, + } +} + +fn resolve_file( + md5: &str, + create_time: i64, + title: Option<&str>, + msg_index: usize, + ctx: &SharedContext, +) -> ResolvedAsset { + // Try hardlink DB first + let hardlink_result = ctx.hardlink_pool.with_connection(|conn| { + wx_media::query_hardlink_with_conn(conn, "file", md5) + }); + + let entries = match hardlink_result { + Some(Ok(e)) => Some(e), + Some(Err(e)) => { + if !matches!(&e, wx_media::MediaError::NotFound(_)) { + eprintln!("warning: file hardlink query failed for md5={md5}: {e}"); + } + None + } + None => None, + }; + + if let Some(entries) = entries { + if let Some(entry) = entries.first() { + let candidates = [ + ctx.file_dir + .join(&entry.dir1) + .join(&entry.dir2) + .join(&entry.file_name), + ctx.file_dir.join(&entry.dir1).join(&entry.file_name), + ]; + + if let Some(source) = candidates.iter().find(|p| p.exists()) { + let filename = format!("{}_{}", md5, entry.file_name); + if ctx.write_gate.claim(&filename) { + let out_path = ctx.output_media_dir.join(&filename); + if let Err(e) = std::fs::copy(&source, &out_path) { + return ResolvedAsset { + msg_index, + asset: None, + tags: vec![], + error: Some(ExportError { + task_kind: "file", + key: md5.to_string(), + reason: format!("copy {}: {e}", source.display()), + }), + }; + } + } + return ResolvedAsset { + msg_index, + asset: Some(MediaAsset { + kind: MediaKind::File, + filename, + }), + tags: vec![], + error: None, + }; + } + } + } + + // Fallback: directory scan by title + if let Some(t) = title { + let month = format_month(create_time); + if let Some(source) = wx_media::find_file_by_name(&ctx.file_dir, t, &month) { + let basename = std::path::Path::new(t) + .file_name() + .map(|n| n.to_string_lossy().to_string()) + .unwrap_or_else(|| t.to_string()); + let safe_name = sanitize_filename(&basename); + let filename = format!("{md5}_{safe_name}"); + if ctx.write_gate.claim(&filename) { + let out_path = ctx.output_media_dir.join(&filename); + if let Err(e) = std::fs::copy(&source, &out_path) { + return ResolvedAsset { + msg_index, + asset: None, + tags: vec![], + error: Some(ExportError { + task_kind: "file", + key: md5.to_string(), + reason: format!("copy fallback {}: {e}", source.display()), + }), + }; + } + } + return ResolvedAsset { + msg_index, + asset: Some(MediaAsset { + kind: MediaKind::File, + filename, + }), + tags: vec![TaskTag::FallbackFile], + error: None, + }; + } + } + + ResolvedAsset { + msg_index, + asset: None, + tags: vec![TaskTag::SkippedFile], + error: None, + } +} + +/// Stage 5: Collect resolved assets back into a media_map indexed by message position. +/// +/// Also populates MediaStats from TaskTags. +pub fn collect( + results: Vec, + dup_map: &DupMap, + total_messages: usize, +) -> (Vec>, MediaStats, ErrorSummary) { + let mut media_map: Vec> = vec![vec![]; total_messages]; + let mut stats = MediaStats::default(); + let mut errors = ErrorSummary::default(); + + // Build index from results by msg_index + let mut by_index: HashMap, Vec, Option)> = + HashMap::new(); + for r in results { + if let Some(e) = r.error { + errors.errors.push(e); + } + by_index.insert( + r.msg_index, + (r.asset, r.tags, None), + ); + } + + // Place canonical results — count tags always, copy asset only when present. + // Matches old MediaBridge: SkippedVideo/SkippedFile stats counted unconditionally; + // image stats (ThumbnailImage, WxgfTranscoded, WxgfFallback) also counted + // because canonical always does the full resolve. + for (msg_idx, (asset, tags, _)) in &by_index { + apply_tags(&mut stats, tags); + if let Some(a) = asset { + media_map[*msg_idx].push(a.clone()); + } + } + + // Resolve duplicates — copy the canonical task's asset to duplicate msg positions. + // Step 1: Count tags that should be counted per-message (SilkVoice, FallbackVideo, + // FallbackFile, SkippedVideo, SkippedFile) — always, even when asset is None. + // Step 2: Copy the asset to duplicate msg positions (only when asset exists). + // This two-step approach matches old MediaBridge behavior where skipped/fallback + // stats were counted regardless of dedup, but image stats only counted once. + for (dup_msg_idx, canonical_msg_idx) in &dup_map.duplicates { + if let Some((asset, tags, _)) = by_index.get(canonical_msg_idx) { + let dup_tags: Vec = tags.iter().copied().filter(|t| t.counts_on_duplicate()).collect(); + apply_tags(&mut stats, &dup_tags); + if let Some(a) = asset { + media_map[*dup_msg_idx].push(a.clone()); + } + } + } + + (media_map, stats, errors) +} + +fn apply_tags(stats: &mut MediaStats, tags: &[TaskTag]) { + for tag in tags { + match tag { + TaskTag::ThumbnailImage => stats.thumbnail_images += 1, + TaskTag::SilkVoice => stats.silk_voices += 1, + TaskTag::WxgfTranscoded => stats.wxgf_transcoded += 1, + TaskTag::WxgfFallback => stats.wxgf_fallback += 1, + TaskTag::FallbackVideo => stats.fallback_videos += 1, + TaskTag::FallbackFile => stats.fallback_files += 1, + TaskTag::SkippedVideo => stats.skipped_videos += 1, + TaskTag::SkippedFile => stats.skipped_files += 1, + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::Arc; + + #[test] + fn test_classify_empty_messages() { + let tasks = classify(&[]); + assert!(tasks.is_empty()); + } + + #[test] + fn test_dedup_no_duplicates() { + let tasks = vec![ + MediaTask::Image { + md5: "aaa".to_string(), + msg_index: 0, + }, + MediaTask::Image { + md5: "bbb".to_string(), + msg_index: 1, + }, + ]; + let (unique, dup_map) = dedup(tasks); + assert_eq!(unique.len(), 2); + assert!(dup_map.duplicates.is_empty()); + } + + #[test] + fn test_dedup_duplicate_md5_image() { + let tasks = vec![ + MediaTask::Image { + md5: "same".to_string(), + msg_index: 0, + }, + MediaTask::Image { + md5: "same".to_string(), + msg_index: 1, + }, + ]; + let (unique, dup_map) = dedup(tasks); + assert_eq!(unique.len(), 1); + assert_eq!(dup_map.duplicates.len(), 1); + assert_eq!(dup_map.duplicates[0].0, 1); // duplicate msg_index + } + + #[test] + fn test_dedup_video_different_create_time_not_deduped() { + let tasks = vec![ + MediaTask::Video { + md5: "same".to_string(), + create_time: 1000, + msg_index: 0, + }, + MediaTask::Video { + md5: "same".to_string(), + create_time: 2000, + msg_index: 1, + }, + ]; + let (unique, dup_map) = dedup(tasks); + assert_eq!(unique.len(), 2); + assert!(dup_map.duplicates.is_empty()); + } + + #[test] + fn test_dedup_file_different_title_not_deduped() { + let tasks = vec![ + MediaTask::File { + md5: "same".to_string(), + create_time: 1000, + title: Some("file_a.pdf".to_string()), + msg_index: 0, + }, + MediaTask::File { + md5: "same".to_string(), + create_time: 1000, + title: Some("file_b.pdf".to_string()), + msg_index: 1, + }, + ]; + let (unique, dup_map) = dedup(tasks); + assert_eq!(unique.len(), 2); + assert!(dup_map.duplicates.is_empty()); + } + + #[test] + fn test_dedup_voice_same_server_id() { + let tasks = vec![ + MediaTask::Voice { + server_id: 123, + msg_index: 0, + }, + MediaTask::Voice { + server_id: 123, + msg_index: 1, + }, + ]; + let (unique, dup_map) = dedup(tasks); + assert_eq!(unique.len(), 1); + assert_eq!(dup_map.duplicates.len(), 1); + } + + #[test] + fn test_write_gate_prevents_duplicate_writes() { + let gate = WriteGate::new(); + assert!(gate.claim("file1.jpg")); + assert!(!gate.claim("file1.jpg")); // second claim returns false + assert!(gate.claim("file2.jpg")); + } + + // --- Parity / integration tests --- + + /// Verify that duplicate messages get the same asset but tags are NOT double-counted. + /// This matches the old MediaBridge behavior where `exported.insert()` returned early + /// for duplicates, skipping stat increments. + #[test] + fn test_collect_duplicate_image_no_double_tag_count() { + // Two image messages with same md5 → dedup removes one, collect copies asset + let tasks = vec![ + MediaTask::Image { + md5: "abc123".to_string(), + msg_index: 0, + }, + MediaTask::Image { + md5: "abc123".to_string(), + msg_index: 1, + }, + ]; + let (unique, dup_map) = dedup(tasks); + assert_eq!(unique.len(), 1); + assert_eq!(dup_map.duplicates.len(), 1); + + // Simulate resolve producing a thumbnail + wxgf_transcoded asset + let results = vec![ResolvedAsset { + msg_index: 0, + asset: Some(MediaAsset { + kind: MediaKind::Image, + filename: "abc123.png".to_string(), + }), + tags: vec![TaskTag::ThumbnailImage, TaskTag::WxgfTranscoded], + error: None, + }]; + + let (media_map, stats, _errors) = collect(results, &dup_map, 2); + + // Both messages should have the asset + assert_eq!(media_map[0].len(), 1); + assert_eq!(media_map[1].len(), 1); + assert_eq!(media_map[0][0].filename, "abc123.png"); + assert_eq!(media_map[1][0].filename, "abc123.png"); + + // Tags should only be counted once (for the canonical task), not for the duplicate + assert_eq!(stats.thumbnail_images, 1); + assert_eq!(stats.wxgf_transcoded, 1); + } + + /// Verify that duplicate voice messages DO count silk_voices for each duplicate. + /// Matches old MediaBridge where silk was counted BEFORE the export check. + #[test] + fn test_collect_duplicate_voice_counts_silk() { + let tasks = vec![ + MediaTask::Voice { + server_id: 42, + msg_index: 0, + }, + MediaTask::Voice { + server_id: 42, + msg_index: 1, + }, + ]; + let (unique, dup_map) = dedup(tasks); + assert_eq!(unique.len(), 1); + + let results = vec![ResolvedAsset { + msg_index: 0, + asset: Some(MediaAsset { + kind: MediaKind::Voice, + filename: "42.mp3".to_string(), + }), + tags: vec![TaskTag::SilkVoice], + error: None, + }]; + + let (media_map, stats, _errors) = collect(results, &dup_map, 2); + + assert_eq!(media_map[0].len(), 1); + assert_eq!(media_map[1].len(), 1); + // silk_voices counted for each message (old behavior: counted before export check) + assert_eq!(stats.silk_voices, 2); + } + + /// Verify that duplicate video fallback counts fallback_videos for each duplicate. + #[test] + fn test_collect_duplicate_video_fallback_counts() { + let tasks = vec![ + MediaTask::Video { + md5: "v1".to_string(), + create_time: 1000, + msg_index: 0, + }, + MediaTask::Video { + md5: "v1".to_string(), + create_time: 1000, + msg_index: 1, + }, + ]; + let (unique, dup_map) = dedup(tasks); + assert_eq!(unique.len(), 1); + + let results = vec![ResolvedAsset { + msg_index: 0, + asset: Some(MediaAsset { + kind: MediaKind::Video, + filename: "v1.mp4".to_string(), + }), + tags: vec![TaskTag::FallbackVideo], + error: None, + }]; + + let (_, stats, _) = collect(results, &dup_map, 2); + + // fallback_videos counted for each message (old behavior) + assert_eq!(stats.fallback_videos, 2); + } + + /// Verify that duplicate skipped videos DO count skipped_videos for each duplicate. + /// Matches old MediaBridge where skipped_videos was counted unconditionally. + #[test] + fn test_collect_duplicate_skipped_video_counts() { + let tasks = vec![ + MediaTask::Video { + md5: "missing".to_string(), + create_time: 1000, + msg_index: 0, + }, + MediaTask::Video { + md5: "missing".to_string(), + create_time: 1000, + msg_index: 1, + }, + ]; + let (unique, dup_map) = dedup(tasks); + assert_eq!(unique.len(), 1); + + let results = vec![ResolvedAsset { + msg_index: 0, + asset: None, + tags: vec![TaskTag::SkippedVideo], + error: None, + }]; + + let (_, stats, _) = collect(results, &dup_map, 2); + + // skipped_videos counted for each message (old behavior: unconditional count) + assert_eq!(stats.skipped_videos, 2); + } + + /// Verify dedup produces consistent output: two identical image messages + /// result in the same asset being placed at both msg positions. + #[test] + fn test_dedup_produces_consistent_output() { + let tasks = vec![ + MediaTask::Image { + md5: "img1".to_string(), + msg_index: 0, + }, + MediaTask::Image { + md5: "img2".to_string(), + msg_index: 1, + }, + MediaTask::Image { + md5: "img1".to_string(), // duplicate of msg 0 + msg_index: 2, + }, + ]; + let (unique, dup_map) = dedup(tasks); + assert_eq!(unique.len(), 2); + assert_eq!(dup_map.duplicates.len(), 1); + assert_eq!(dup_map.duplicates[0], (2, 0)); // msg 2 is dup of canonical msg 0 + + // Simulate resolve for the 2 unique tasks + let results = vec![ + ResolvedAsset { + msg_index: 0, + asset: Some(MediaAsset { + kind: MediaKind::Image, + filename: "img1.jpg".to_string(), + }), + tags: vec![], + error: None, + }, + ResolvedAsset { + msg_index: 1, + asset: Some(MediaAsset { + kind: MediaKind::Image, + filename: "img2.jpg".to_string(), + }), + tags: vec![], + error: None, + }, + ]; + + let (media_map, _stats, _errors) = collect(results, &dup_map, 3); + + // All 3 messages should have their assets + assert_eq!(media_map[0].len(), 1); + assert_eq!(media_map[0][0].filename, "img1.jpg"); + assert_eq!(media_map[1].len(), 1); + assert_eq!(media_map[1][0].filename, "img2.jpg"); + assert_eq!(media_map[2].len(), 1); + assert_eq!(media_map[2][0].filename, "img1.jpg"); // duplicate gets canonical's asset + } + + /// Verify classify correctly maps message types to tasks. + #[test] + fn test_classify_mixed_messages() { + use crate::schema::EnrichedMessage; + use wx_context::Direction; + use wx_db::Message; + + let msgs = vec![ + EnrichedMessage { + message: Message { + sort_seq: 0, + server_id: 1, + msg_type: 3, + sub_type: 0, + sender: "a".into(), + talker: "b".into(), + create_time: 1000, + content: MessageContent::Image { + md5: Some("md5_a".into()), + }, + status: 0, + }, + sender_display_name: "A".into(), + direction: Direction::Incoming, + snippet: String::new(), + }, + EnrichedMessage { + message: Message { + sort_seq: 1, + server_id: 2, + msg_type: 34, + sub_type: 0, + sender: "a".into(), + talker: "b".into(), + create_time: 1001, + content: MessageContent::Voice, + status: 0, + }, + sender_display_name: "A".into(), + direction: Direction::Incoming, + snippet: String::new(), + }, + EnrichedMessage { + message: Message { + sort_seq: 2, + server_id: 3, + msg_type: 43, + sub_type: 0, + sender: "a".into(), + talker: "b".into(), + create_time: 1002, + content: MessageContent::Text("hello".into()), + status: 0, + }, + sender_display_name: "A".into(), + direction: Direction::Incoming, + snippet: "hello".into(), + }, + ]; + + let tasks = classify(&msgs); + assert_eq!(tasks.len(), 2); // Text message skipped + + // First task: Image + assert!(matches!( + &tasks[0], + MediaTask::Image { md5, msg_index: 0 } if md5 == "md5_a" + )); + + // Second task: Voice + assert!(matches!( + &tasks[1], + MediaTask::Voice { server_id: 2, msg_index: 1 } + )); + } + + /// End-to-end parity test: resolve an image through the new parallel pipeline + /// and compare with the old MediaBridge serial oracle. + /// Uses real .dat file I/O with mock XOR-encrypted data. + #[test] + fn test_parallel_equals_serial_image_resolve() { + use crate::cmd::export_media::MediaBridge; + + let tmp = tempfile::TempDir::new().unwrap(); + let root = tmp.path(); + let talker = "wxid_testuser"; + let md5 = "4865625c4e99e4d3b0959a0fe84f41cd"; + let xor_key = 0xa5u8; + + // Create a sample .dat file: XOR-encrypted embedded PNG WXGF + let png = vec![ + 0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A, 0x00, 0x00, 0x00, 0x0D, 0x49, 0x48, + 0x44, 0x52, 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01, 0x08, 0x06, 0x00, 0x00, + 0x00, 0x1F, 0x15, 0xC4, 0x89, 0x00, 0x00, 0x00, 0x0D, 0x49, 0x44, 0x41, 0x54, 0x78, + 0x9C, 0x63, 0xF8, 0xCF, 0xC0, 0xF0, 0x1F, 0x00, 0x05, 0x00, 0x01, 0xFF, 0x89, 0x99, + 0x3D, 0x1D, 0x00, 0x00, 0x00, 0x00, 0x49, 0x45, 0x4E, 0x44, 0xAE, 0x42, 0x60, 0x82, + ]; + let mut wxgf = b"wxgfmetadata".to_vec(); + wxgf.extend_from_slice(&png); + let encrypted: Vec = wxgf.iter().map(|b| b ^ xor_key).collect(); + + let username_hash = format!("{:x}", wx_media::md5_hash(talker.as_bytes())); + let img_dir = root.join("attach").join(&username_hash).join("2026-03").join("Img"); + std::fs::create_dir_all(&img_dir).unwrap(); + std::fs::write(img_dir.join(format!("{md5}.dat")), &encrypted).unwrap(); + + let output_media = root.join("output"); + std::fs::create_dir_all(&output_media).unwrap(); + + let dat_opts = wx_media::DatDecryptOptions { + v2_aes_key: None, + xor_key: Some(xor_key), + }; + + // --- Serial oracle (old MediaBridge) --- + let mut bridge = MediaBridge::new( + root.join("attach"), + root.join("media"), + root.join("file"), + root.join("video"), + root.join("hardlink.db"), + output_media.clone(), + dat_opts.clone(), + ); + let serial_assets = bridge.resolve( + &wx_db::Message { + sort_seq: 0, + server_id: 1, + msg_type: 3, + sub_type: 0, + sender: "sender".into(), + talker: talker.into(), + create_time: 1_700_000_000, + content: wx_db::MessageContent::Image { + md5: Some(md5.to_string()), + }, + status: 0, + }, + talker, + ); + + // --- New parallel pipeline --- + let ctx = build_shared_context( + root.join("attach"), + root.join("media"), + root.join("file"), + root.join("video"), + root.join("hardlink.db"), + output_media.clone(), + talker, + dat_opts, + ); + + // Clean output dir so the new pipeline can write + let _ = std::fs::remove_dir_all(&output_media); + std::fs::create_dir_all(&output_media).unwrap(); + + let tasks = vec![MediaTask::Image { + md5: md5.to_string(), + msg_index: 0, + }]; + let (results, _) = resolve_parallel(tasks, std::sync::Arc::new(ctx), Some(1)); + let (media_map, stats, _) = collect(results, &DupMap { duplicates: vec![] }, 1); + + // Compare: both should produce the same filename + assert_eq!(serial_assets.len(), 1); + assert_eq!(media_map[0].len(), 1); + assert_eq!(serial_assets[0].filename, media_map[0][0].filename); + + // Both should have written the same file + let serial_bytes = std::fs::read(output_media.join(&serial_assets[0].filename)).unwrap(); + let parallel_bytes = std::fs::read(output_media.join(&media_map[0][0].filename)).unwrap(); + assert_eq!(serial_bytes, parallel_bytes); + + // Stats should match (wxgf_transcoded = 1 in both) + assert_eq!(bridge.stats.wxgf_transcoded, stats.wxgf_transcoded); + } + + #[cfg(feature = "audio")] + fn sample_silk() -> Vec { + silk_rs::encode_silk(vec![0_u8; 24_000 / 1_000 * 40 * 2], 24_000, 24_000, true).unwrap() + } + + #[cfg(feature = "audio")] + fn create_voice_media_db(path: &Path, rows: &[(i64, i64, i64, i64, &[u8])]) { + let conn = rusqlite::Connection::open(path).unwrap(); + conn.execute_batch( + "CREATE TABLE VoiceInfo ( + chat_name_id INTEGER, + create_time INTEGER, + local_id INTEGER, + svr_id INTEGER, + voice_data BLOB, + data_index TEXT DEFAULT '0' + ); + CREATE INDEX VoiceInfo_INDEX ON VoiceInfo(chat_name_id, svr_id);", + ) + .unwrap(); + let mut stmt = conn + .prepare( + "INSERT INTO VoiceInfo (chat_name_id, create_time, local_id, svr_id, voice_data) + VALUES (?, ?, ?, ?, ?)", + ) + .unwrap(); + for &(chat_name_id, create_time, local_id, svr_id, data) in rows { + stmt.execute(rusqlite::params![ + chat_name_id, + create_time, + local_id, + svr_id, + data + ]) + .unwrap(); + } + } + + #[cfg(feature = "audio")] + #[test] + fn test_resolve_voice_caches_chat_name_id_hint_after_first_lookup() { + let tmp = tempfile::TempDir::new().unwrap(); + let media_dir = tmp.path().join("media"); + let output_media = tmp.path().join("output"); + std::fs::create_dir_all(&media_dir).unwrap(); + std::fs::create_dir_all(&output_media).unwrap(); + + let silk = sample_silk(); + create_voice_media_db( + &media_dir.join("media_0.db"), + &[ + (55, 1000, 1, 101, &silk), + (55, 1001, 2, 102, &silk), + ], + ); + + let ctx = Arc::new(build_shared_context( + tmp.path().join("attach"), + media_dir, + tmp.path().join("file"), + tmp.path().join("video"), + tmp.path().join("hardlink.db"), + output_media, + "he593121260", + wx_media::DatDecryptOptions::default(), + )); + + assert_eq!(*ctx.voice_chat_name_id_hint.lock().unwrap(), None); + + let first = resolve_voice(101, 0, &ctx); + assert!(first.error.is_none()); + assert_eq!(*ctx.voice_chat_name_id_hint.lock().unwrap(), Some(55)); + + let second = resolve_voice(102, 1, &ctx); + assert!(second.error.is_none()); + assert_eq!(*ctx.voice_chat_name_id_hint.lock().unwrap(), Some(55)); + } +} diff --git a/crates/wx-cli/src/cmd/info.rs b/crates/wx-cli/src/cmd/info.rs new file mode 100644 index 0000000..1e97258 --- /dev/null +++ b/crates/wx-cli/src/cmd/info.rs @@ -0,0 +1,27 @@ +pub fn cmd_info(db_file: &std::path::Path) -> Result<(), Box> { + use std::io::Read; + + let metadata = std::fs::metadata(db_file)?; + let mut f = std::fs::File::open(db_file)?; + let mut header = [0u8; 16]; + f.read_exact(&mut header)?; + + let is_sqlite = &header[..] == b"SQLite format 3\0"; + let salt_hex = hex::encode(header); + let page_count = metadata.len() / 4096; + + println!("File: {}", db_file.display()); + println!( + "Size: {} bytes ({} pages)", + metadata.len(), + page_count + ); + if is_sqlite { + println!("Status: decrypted (SQLite header detected)"); + } else { + println!("Status: encrypted"); + println!("Salt: {salt_hex}"); + } + + Ok(()) +} diff --git a/crates/wx-cli/src/cmd/key.rs b/crates/wx-cli/src/cmd/key.rs new file mode 100644 index 0000000..59681c5 --- /dev/null +++ b/crates/wx-cli/src/cmd/key.rs @@ -0,0 +1,287 @@ +use std::time::Duration; + +use crate::util::{lookup_or_resolve_nickname, parse_hex_key_32}; + +pub async fn cmd_key_extract(timeout_secs: u64) -> Result<(), Box> { + eprintln!("Running pre-flight checks..."); + wx_keychain::preflight_checks()?; + eprintln!(" All checks passed."); + + let version = wx_keychain::ensure_supported_wechat_version()?; + eprintln!(" WeChat version: {version}"); + + let accounts = wx_keychain::find_account_dirs()?; + if accounts.is_empty() { + return Err("no WeChat account directories found".into()); + } + + let mut store = wx_keychain::KeyStore::load_default()?; + let mut store_dirty = false; + + eprintln!("Detected accounts:"); + for a in &accounts { + let nick = lookup_or_resolve_nickname(&mut store, a); + if nick.is_some() { + store_dirty = true; + } + eprintln!( + " {} ({})", + a.account_id, + nick.unwrap_or_else(|| "昵称未知".to_string()) + ); + } + if store_dirty { + store.save_default()?; + } + + let result = wx_keychain::capture_key(&accounts, Duration::from_secs(timeout_secs)).await?; + + let matched = &result.matched_account; + let hex_key = hex::encode(result.raw_key); + eprintln!("Key captured after {} PBKDF2 calls.", result.call_count); + eprintln!("Matched account: {}", matched.account_id); + println!("{hex_key}"); + + let nickname = wx_keychain::resolve_nickname( + &matched.data_dir, + &wx_decrypt::KeyMaterial::RawKey(result.raw_key), + &matched.base_wxid, + ) + .unwrap_or_else(|e| { + eprintln!(" Warning: nickname resolution failed: {e}"); + None + }); + + if let Some(ref n) = nickname { + eprintln!("Account nickname: {n}"); + } + + store.set( + &matched.account_id, + &hex_key, + &version, + nickname, + Some(matched.base_wxid.clone()), + ); + store.save_default()?; + eprintln!( + "Key saved to {:?}", + wx_keychain::KeyStore::default_path()? + ); + + Ok(()) +} + +#[cfg(target_os = "macos")] +pub fn cmd_key_scan() -> Result<(), Box> { + // SIP check — task_for_pid fails with kern_return=5 when SIP is enabled, + // even as root. This is a hard requirement (tested 2026-03-08). + let sip = wx_keychain::check_sip(); + if !sip.passed { + return Err(format!( + "{} — task_for_pid requires SIP disabled. Disable in Recovery Mode: csrutil disable", + sip.detail + ) + .into()); + } + + // Find WeChat process (PID + version only). + let (pid, version) = wx_keychain::find_wechat_pid()?; + eprintln!("Found WeChat PID {} (v{})", pid, version); + + // Load account directories. + let accounts = wx_keychain::find_account_dirs()?; + if accounts.is_empty() { + return Err("no WeChat account directories found".into()); + } + eprintln!( + "Found {} account director{}", + accounts.len(), + if accounts.len() == 1 { "y" } else { "ies" } + ); + + // Scan process memory. + eprintln!("Scanning WeChat process memory..."); + let results = + wx_keychain::capture_key_mach(pid, &accounts, &wx_decrypt::MACOS_4_1_7_31)?; + + // Count total pairs across all results + let total_pairs: usize = results + .iter() + .map(|r| match &r.key_material { + wx_decrypt::KeyMaterial::EncKeys(pairs) => pairs.len(), + _ => unreachable!("capture_key_mach always returns EncKeys"), + }) + .sum(); + eprintln!( + "Found {} valid key{} for {} account{}", + total_pairs, + if total_pairs == 1 { "" } else { "s" }, + results.len(), + if results.len() == 1 { "" } else { "s" }, + ); + + // Store results. + let mut store = wx_keychain::KeyStore::load_default()?; + for r in &results { + let matched = &r.matched_account; + + let nickname = wx_keychain::resolve_nickname( + &matched.data_dir, + &r.key_material, + &matched.base_wxid, + ) + .unwrap_or_else(|e| { + eprintln!( + " Warning: nickname resolution failed for {}: {e}", + matched.account_id + ); + None + }); + + let pairs = match &r.key_material { + wx_decrypt::KeyMaterial::EncKeys(pairs) => pairs, + _ => unreachable!("capture_key_mach always returns EncKeys"), + }; + + store.set_enc_keys( + &matched.account_id, + pairs, + &version, + nickname.clone(), + Some(matched.base_wxid.clone()), + ); + + let display = nickname + .as_ref() + .map(|n| format!("{} ({})", matched.account_id, n)) + .unwrap_or_else(|| matched.account_id.clone()); + eprintln!( + " {} — {} enc_key{} stored", + display, + pairs.len(), + if pairs.len() == 1 { "" } else { "s" } + ); + for pair in pairs { + println!( + "{}\t{}\t{}", + matched.account_id, + hex::encode(pair.key), + hex::encode(pair.salt) + ); + } + } + + store.save_default()?; + eprintln!( + "Keys saved to {:?}", + wx_keychain::KeyStore::default_path()? + ); + + Ok(()) +} + +#[cfg(not(target_os = "macos"))] +pub fn cmd_key_scan() -> Result<(), Box> { + Err("key scan is only supported on macOS".into()) +} + +pub fn cmd_key_list() -> Result<(), Box> { + let mut store = wx_keychain::KeyStore::load_default()?; + if store.accounts.is_empty() { + eprintln!("No keys stored."); + return Ok(()); + } + + let accounts = wx_keychain::find_account_dirs().unwrap_or_default(); + let mut store_dirty = false; + for account in &accounts { + if lookup_or_resolve_nickname(&mut store, account).is_some() { + store_dirty = true; + } + } + if store_dirty { + store.save_default()?; + } + + let mut ids: Vec<_> = store.accounts.keys().cloned().collect(); + ids.sort(); + for id in ids { + let key = store + .get(&id) + .ok_or_else(|| format!("missing key entry for account {id}"))?; + + let raw_status = if key.data_key.is_empty() { "no" } else { "yes" }; + let enc_count = key.enc_keys.len(); + let has_legacy_enc = key.enc_key.as_ref().is_some_and(|k| !k.is_empty()); + let enc_status = if enc_count > 0 { + format!("{enc_count}") + } else if has_legacy_enc { + "1".to_string() + } else { + "no".to_string() + }; + let img_status = if key.image_aes_key.is_some() { + "yes" + } else { + "no" + }; + + let display_key = if !key.data_key.is_empty() { + key.data_key.clone() + } else if enc_count > 0 { + let first = &key.enc_keys[0].enc_key; + if enc_count > 1 { + format!("{} (+{} more)", first, enc_count - 1) + } else { + first.clone() + } + } else if let Some(ref ek) = key.enc_key { + ek.clone() + } else { + "(no key)".to_string() + }; + + println!( + "{} {} (v{}, {}, raw={} enc={} img={})", + key.display_name(), + display_key, + key.wechat_version, + key.extracted_at.format("%Y-%m-%d %H:%M:%S UTC"), + raw_status, + enc_status, + img_status, + ); + } + Ok(()) +} + +pub fn cmd_key_set(account: &str, hex_key: &str) -> Result<(), Box> { + parse_hex_key_32(hex_key, "manual key")?; + + let mut store = wx_keychain::KeyStore::load_default()?; + store.set(account, hex_key, "manual", None, None); + store.save_default()?; + eprintln!("Key saved for {account}."); + Ok(()) +} + +pub fn cmd_key_set_image(account: &str, image_key: &str) -> Result<(), Box> { + let key_hex = if image_key.len() == 32 && image_key.chars().all(|c| c.is_ascii_hexdigit()) { + image_key.to_string() + } else if image_key.len() == 16 && image_key.is_ascii() { + hex::encode(image_key.as_bytes()) + } else { + return Err(format!( + "image key must be 16-byte ASCII string or 32-char hex, got {} chars", + image_key.len() + ) + .into()); + }; + + let mut store = wx_keychain::KeyStore::load_default()?; + store.set_image_key(account, &key_hex); + store.save_default()?; + eprintln!("Image AES key saved for {account}."); + Ok(()) +} diff --git a/crates/wx-cli/src/cmd/media.rs b/crates/wx-cli/src/cmd/media.rs new file mode 100644 index 0000000..72c03c6 --- /dev/null +++ b/crates/wx-cli/src/cmd/media.rs @@ -0,0 +1,335 @@ +use std::path::PathBuf; + +use crate::util::walkdir_dat_files; +use crate::MediaAction; + +pub fn cmd_media(action: MediaAction) -> Result<(), Box> { + match action { + MediaAction::DecryptDat { + input, + output, + v2_key, + account, + data_dir, + xor_key, + } => { + let v2_aes_key = if let Some(ref k) = v2_key { + let bytes = k.as_bytes(); + if bytes.len() != 16 { + return Err(format!("V2 key must be 16 bytes, got {}", bytes.len()).into()); + } + let mut arr = [0u8; 16]; + arr.copy_from_slice(bytes); + Some(arr) + } else if let Some(ref acct) = account { + let store = wx_keychain::KeyStore::load_default()?; + let entry = store + .get(acct) + .ok_or_else(|| format!("no key entry for account '{acct}' in KeyStore"))?; + let hex = entry + .image_aes_key + .as_ref() + .ok_or_else(|| format!("no V2 image key stored for account '{acct}' — use `key set-image` or `--data-dir` instead"))?; + let raw = hex::decode(hex) + .map_err(|e| format!("invalid image_aes_key hex in KeyStore: {e}"))?; + if raw.len() != 16 { + return Err(format!( + "stored image_aes_key is {} bytes, expected 16", + raw.len() + ) + .into()); + } + let mut arr = [0u8; 16]; + arr.copy_from_slice(&raw); + eprintln!("Using V2 key from KeyStore for account '{acct}'"); + Some(arr) + } else if let Some(ref dir) = data_dir { + let key = wx_media::derive_v2_key_from_dir(dir) + .map_err(|e| format!("V2 key derivation from --data-dir failed: {e}"))?; + let key_preview = String::from_utf8_lossy(&key[..8]); + eprintln!("Derived V2 key from UIN+WXID: {key_preview}..."); + Some(key) + } else { + None + }; + + let explicit_xor = xor_key + .map(|s| u8::from_str_radix(&s, 16)) + .transpose() + .map_err(|e| format!("invalid xor_key hex: {e}"))?; + + if input.is_dir() { + let out_dir = output.unwrap_or_else(|| input.join("decrypted")); + std::fs::create_dir_all(&out_dir)?; + + let xor = explicit_xor.or_else(|| { + let detected = wx_media::detect_xor_key(&input); + if let Some(k) = detected { + eprintln!("Auto-detected XOR key: 0x{k:02x}"); + } + detected + }); + + let opts = wx_media::DatDecryptOptions { + v2_aes_key, + xor_key: xor, + }; + + let mut ok_count = 0usize; + let mut err_count = 0usize; + let mut skip_count = 0usize; + + for entry in walkdir_dat_files(&input) { + let path = entry.path(); + let name = entry.file_name().to_string_lossy().to_string(); + + if name.ends_with("_t.dat") { + skip_count += 1; + continue; + } + + let data = match std::fs::read(&path) { + Ok(d) => d, + Err(e) => { + eprintln!(" Skip {}: {e}", path.display()); + err_count += 1; + continue; + } + }; + + match wx_media::decrypt_dat(&data, &opts) { + Ok(result) => { + let (final_data, final_ext) = + maybe_transcode_wxgf(result.data, &result.ext, &name); + let stem = path.file_stem().unwrap_or_default().to_string_lossy(); + let out_path = out_dir.join(format!("{}.{}", stem, final_ext)); + std::fs::write(&out_path, &final_data)?; + ok_count += 1; + } + Err(e) => { + eprintln!(" Failed {}: {e}", name); + err_count += 1; + } + } + } + + if skip_count > 0 { + eprintln!("Batch decrypt: {ok_count} succeeded, {err_count} failed, {skip_count} skipped"); + } else { + eprintln!("Batch decrypt: {ok_count} succeeded, {err_count} failed"); + } + } else { + let data = std::fs::read(&input)?; + + let xor = explicit_xor.or_else(|| { + input.parent().and_then(|dir| { + let detected = wx_media::detect_xor_key(dir); + if let Some(k) = detected { + eprintln!("Auto-detected XOR key: 0x{k:02x}"); + } + detected + }) + }); + + let opts = wx_media::DatDecryptOptions { + v2_aes_key, + xor_key: xor, + }; + + let result = wx_media::decrypt_dat(&data, &opts)?; + let display_name = input.display().to_string(); + let (final_data, final_ext) = + maybe_transcode_wxgf(result.data, &result.ext, &display_name); + + let out_path = if let Some(user_path) = output { + // Auto-correct suffix if user-specified extension doesn't match actual format + let user_ext = user_path + .extension() + .map(|e| e.to_string_lossy().to_string()) + .unwrap_or_default(); + if !user_ext.is_empty() && user_ext != final_ext { + let corrected = user_path.with_extension(&final_ext); + eprintln!( + "warning: requested .{} but actual format is .{}, writing to {}", + user_ext, + final_ext, + corrected.display(), + ); + corrected + } else { + user_path + } + } else { + let stem = input.file_stem().unwrap_or_default().to_string_lossy(); + input.with_file_name(format!("{}.{}", stem, final_ext)) + }; + + if let Some(parent) = out_path.parent() { + std::fs::create_dir_all(parent)?; + } + std::fs::write(&out_path, &final_data)?; + + eprintln!( + "Decrypted: {} -> {} ({:?}, {} bytes, {})", + input.display(), + out_path.display(), + result.format, + final_data.len(), + final_ext, + ); + } + } + MediaAction::ResolvePath { + db, + media_type, + key, + } => { + let entries = wx_media::query_hardlink(&db, &media_type, &key)?; + println!("{}", serde_json::to_string_pretty(&entries)?); + } + MediaAction::ExtractVoice { + media_dir, + svr_id, + output, + raw, + } => { + let blob = wx_media::extract_voice(&media_dir, &svr_id)?; + + if raw { + // Raw SILK output + let out_path = output.unwrap_or_else(|| PathBuf::from(format!("{}.silk", svr_id))); + if let Some(parent) = out_path.parent() { + std::fs::create_dir_all(parent)?; + } + std::fs::write(&out_path, &blob.data)?; + eprintln!( + "Extracted voice (raw SILK): svr_id={}, {} bytes -> {}", + blob.svr_id, + blob.data.len(), + out_path.display(), + ); + } else { + // Transcode to MP3 + match wx_media::transcode_silk_to_mp3(&blob.data) { + Ok(result) => { + if !result.transcoded { + return Err(voice_extract_requires_ffmpeg_message(&svr_id).into()); + } + + let out_path = if let Some(user_path) = output { + // Auto-correct suffix if user output does not match the actual format. + let user_ext = user_path + .extension() + .map(|e| e.to_string_lossy().to_string()) + .unwrap_or_default(); + if !user_ext.is_empty() && user_ext != result.ext { + let corrected = user_path.with_extension(result.ext); + eprintln!( + "warning: requested .{} but actual format is .{}, writing to {}", + user_ext, result.ext, corrected.display(), + ); + corrected + } else { + user_path + } + } else { + PathBuf::from(format!("{}.{}", svr_id, result.ext)) + }; + if let Some(parent) = out_path.parent() { + std::fs::create_dir_all(parent)?; + } + std::fs::write(&out_path, &result.data)?; + eprintln!( + "Extracted voice: svr_id={}, {} bytes -> {} ({})", + blob.svr_id, + result.data.len(), + out_path.display(), + result.ext, + ); + } + Err(wx_media::MediaError::AudioFeatureDisabled) => { + return Err(voice_extract_audio_feature_message(&svr_id).into()); + } + Err(e) => return Err(e.into()), + } + } + } + MediaAction::DecryptVideo { + input, + seed, + output, + } => { + let seed_val = parse_seed(&seed)?; + let ciphertext = std::fs::read(&input)?; + let result = wx_media::decrypt_video(&ciphertext, seed_val); + + let out_path = output.unwrap_or_else(|| { + let stem = input.file_stem().unwrap_or_default().to_string_lossy(); + input.with_file_name(format!("{}.mp4", stem)) + }); + + if let Some(parent) = out_path.parent() { + std::fs::create_dir_all(parent)?; + } + std::fs::write(&out_path, &result.data)?; + + if !result.is_valid_mp4 { + eprintln!("warning: 'ftyp' signature not found, output may not be a valid MP4"); + } + eprintln!( + "Decrypted video: {} -> {} ({} bytes)", + input.display(), + out_path.display(), + result.data.len(), + ); + } + } + Ok(()) +} + +/// Try to transcode WXGF data to a standard image format. +/// Returns the (possibly transcoded) data and extension. +fn maybe_transcode_wxgf(data: Vec, ext: &str, display_name: &str) -> (Vec, String) { + if ext != "wxgf" { + return (data, ext.to_string()); + } + + match wx_media::transcode_wxgf(&data) { + Ok(result) => { + if !result.transcoded { + eprintln!( + "warning: ffmpeg missing, writing raw HEVC for {} ({})", + display_name, + wx_media::MediaError::ffmpeg_install_hint() + ); + } + (result.data, result.ext.to_string()) + } + Err(e) => { + eprintln!("warning: WXGF transcode failed for {}: {}", display_name, e); + (data, "wxgf".to_string()) + } + } +} + +fn voice_extract_requires_ffmpeg_message(svr_id: &str) -> String { + format!( + "voice extraction for svr_id={svr_id} requires ffmpeg to produce MP3; {}. To export raw SILK instead, rerun with --raw", + wx_media::MediaError::ffmpeg_install_hint() + ) +} + +fn voice_extract_audio_feature_message(svr_id: &str) -> String { + format!( + "voice extraction for svr_id={svr_id} requires audio transcoding support; rebuild wx-cli with the 'audio' feature, or rerun with --raw" + ) +} + +/// Parse a seed string as either decimal or hex (with 0x prefix). +fn parse_seed(s: &str) -> Result> { + if let Some(hex) = s.strip_prefix("0x").or_else(|| s.strip_prefix("0X")) { + Ok(u64::from_str_radix(hex, 16)?) + } else { + Ok(s.parse::()?) + } +} diff --git a/crates/wx-cli/src/cmd/mod.rs b/crates/wx-cli/src/cmd/mod.rs new file mode 100644 index 0000000..71a6a0c --- /dev/null +++ b/crates/wx-cli/src/cmd/mod.rs @@ -0,0 +1,23 @@ +pub mod contacts; +pub mod db_dev; +pub mod decode_image; +pub mod decrypt; +pub mod doctor; +pub mod export; +pub mod export_media; +pub mod export_task; +pub mod info; +pub mod key; +pub mod media; +pub mod paths; +pub mod query; +pub mod search; +pub mod serve; +pub mod server; +pub mod sessions; +pub mod status; +pub mod thin_client; +pub mod watch; + +#[cfg(test)] +mod thin_client_tests; diff --git a/crates/wx-cli/src/cmd/paths.rs b/crates/wx-cli/src/cmd/paths.rs new file mode 100644 index 0000000..bcbf2c0 --- /dev/null +++ b/crates/wx-cli/src/cmd/paths.rs @@ -0,0 +1,54 @@ +use std::path::Path; + +use wx_paths::PathsSummary; + +pub fn cmd_paths(json: bool) -> Result<(), Box> { + let ap = wx_paths::AppPaths::new()?; + let summary = ap.summary(); + if json { + println!("{}", serde_json::to_string_pretty(&summary)?); + } else { + print_paths_table(&summary); + } + Ok(()) +} + +fn print_paths_table(summary: &PathsSummary) { + println!("Platform: {}", summary.platform); + println!(); + + let rows: Vec<(&str, &Path)> = vec![ + ("config_dir", &summary.config_dir), + ("keys_file", &summary.keys_file), + ("settings_file", &summary.settings_file), + ("cache_root", &summary.cache_root), + ("state_root", &summary.state_root), + ("logs_dir", &summary.logs_dir), + ("server_state_dir", &summary.server_state_dir), + ("server_stdout_log", &summary.server_stdout_log), + ("server_stderr_log", &summary.server_stderr_log), + ("temp_root", &summary.temp_root), + ]; + + let max_label = rows.iter().map(|(l, _)| l.len()).max().unwrap_or(0); + + for (label, path) in &rows { + let display = tilde_path(path); + let status = if path.exists() { + "[exists]" + } else { + "[missing]" + }; + println!("{: String { + if let Ok(home) = std::env::var("HOME") { + let home_path = Path::new(&home); + if let Ok(suffix) = path.strip_prefix(home_path) { + return format!("~/{}", suffix.display()); + } + } + path.display().to_string() +} diff --git a/crates/wx-cli/src/cmd/query.rs b/crates/wx-cli/src/cmd/query.rs new file mode 100644 index 0000000..1999bd5 --- /dev/null +++ b/crates/wx-cli/src/cmd/query.rs @@ -0,0 +1,423 @@ +use std::path::PathBuf; + +use wx_context::{ + route_shards_for_query, write_shard_metadata_sidecar, AccountContext, ContactResolver, + DecryptRequest, Direction, PersistentCache, ResolveParams, VisibilityIndex, +}; + +use super::contacts::build_visibility; +use super::thin_client::{ThinClient, ThinClientCliArgs, ThinClientOptions}; +use crate::output::JsonEnvelope; +use crate::schema::{enrich_message, project_message_items, EnrichedMessage}; +use crate::util::{open_db_all, print_cache_stats, print_detection_note}; +use crate::{OutputFormat, SortOrderArg}; + +/// Resolve a contact identifier to a talker wxid. +/// +/// Delegates to the shared `contact_id::resolve_contact_id` resolver and prints +/// resolution notes to stderr for CLI feedback. +pub(crate) fn resolve_talker( + contact: &str, + resolver: &ContactResolver, + db: &wx_db::WechatDb, + visibility: Option<&VisibilityIndex>, + show_hidden: bool, +) -> Result> { + use crate::contact_id::{resolve_contact_id, ContactResolveError}; + + match resolve_contact_id(contact, resolver, db, visibility, show_hidden) { + Ok(resolved) => { + if let Some(ref name) = resolved.display_name { + eprintln!("Resolved \"{contact}\" → {name}({})", resolved.wxid); + } + Ok(resolved.wxid) + } + Err( + ContactResolveError::NotFound(msg) + | ContactResolveError::Ambiguous(msg) + | ContactResolveError::Hidden(msg), + ) => Err(msg.into()), + } +} + +#[allow(clippy::too_many_arguments)] +pub fn cmd_query( + contact: &str, + data_dir: Option, + account: Option, + key: Option, + since: Option, + until: Option, + msg_type: Option, + limit: usize, + offset: usize, + order: SortOrderArg, + all: bool, + format: OutputFormat, + around_sort_seq: Option, + around_server_id: Option, + context: Option, + after_sort_seq: Option, + show_hidden: bool, + server: ThinClientCliArgs, +) -> Result<(), Box> { + let has_around = around_sort_seq.is_some() || around_server_id.is_some(); + let has_anchor = + around_sort_seq.is_some() || around_server_id.is_some() || after_sort_seq.is_some(); + let effective_limit = if has_anchor { + limit + } else if all { + wx_db::MAX_QUERY_LIMIT + } else { + wx_db::effective_limit(limit) + }; + + let options = ThinClientOptions::resolve_from_process_env(server); + let preserve_local_warning = context.is_some() && !has_around; + + if options.is_enabled() && !preserve_local_warning { + let client = ThinClient::new(options.clone()); + match client.probe_health().and_then(|_| { + fetch_remote_query( + &client, + contact, + since, + until, + msg_type.clone(), + effective_limit, + offset, + order.clone(), + around_sort_seq, + around_server_id, + context, + after_sort_seq, + show_hidden, + ) + }) { + Ok(envelope) => { + let is_group = envelope + .items + .first() + .map(|item| wx_db::is_group_chat(&item.message.talker)) + .unwrap_or_else(|| wx_db::is_group_chat(contact)); + return print_query_output(&envelope, format, is_group); + } + Err(err) if err.should_fallback(options.mode) => { + eprintln!( + "note: remote server unavailable, falling back to local query ({})", + err.fallback_detail() + ); + } + Err(err) => return Err(err.into()), + } + } + + let (envelope, is_group) = load_local_query( + contact, + data_dir, + account, + key, + since, + until, + msg_type, + limit, + offset, + order, + all, + around_sort_seq, + around_server_id, + context, + after_sort_seq, + show_hidden, + )?; + print_query_output(&envelope, format, is_group) +} + +#[allow(clippy::too_many_arguments)] +fn load_local_query( + contact: &str, + data_dir: Option, + account: Option, + key: Option, + since: Option, + until: Option, + msg_type: Option, + limit: usize, + offset: usize, + order: SortOrderArg, + all: bool, + around_sort_seq: Option, + around_server_id: Option, + context: Option, + after_sort_seq: Option, + show_hidden: bool, +) -> Result<(JsonEnvelope, bool), Box> { + let acct = AccountContext::resolve(&ResolveParams { + account: account.as_deref(), + data_dir: data_dir.as_deref(), + key_hex: key.as_deref(), + })?; + print_detection_note(&acct); + + let (db, _cache, _stats) = if acct.raw_key.is_some() { + // Direct encrypted open — no shard routing needed + let (db, _, stats) = open_db_all(&acct, crate::util::decrypt_progress_callback)?; + (db, None::, stats) + } else { + let params = &wx_decrypt::MACOS_4_1_7_31; + let cache = PersistentCache::new(&acct, params)?; + + // Two-phase decrypt: try shard routing when time bounds are present. + DecryptRequest::new() + .core() + .execute_with_progress(&cache, crate::util::decrypt_progress_callback)?; + + let shard_ids = route_shards_for_query(cache.decrypted_root(), contact, since, until); + + let stats = match shard_ids { + Some(ref ids) => { + eprintln!( + "shard routing: decrypting {} of N shards for time-bounded query", + ids.len() + ); + DecryptRequest::new() + .shards(ids) + .execute_with_progress(&cache, crate::util::decrypt_progress_callback)? + } + None => DecryptRequest::new() + .all() + .execute_with_progress(&cache, crate::util::decrypt_progress_callback)?, + }; + print_cache_stats(&stats); + + let db = wx_db::WechatDb::open(cache.decrypted_root())?; + + if let Err(e) = write_shard_metadata_sidecar(&db, cache.decrypted_root()) { + eprintln!("warning: failed to write shard metadata sidecar: {e}"); + } + + (db, Some(cache), Some(stats)) + }; + + let resolver = ContactResolver::build(&db)?; + let visibility = build_visibility(&acct, &resolver); + let self_wxid = &acct.base_wxid; + + let talker = resolve_talker(contact, &resolver, &db, Some(&visibility), show_hidden)?; + + // Determine if we're using anchor mode + let has_anchor = + around_sort_seq.is_some() || around_server_id.is_some() || after_sort_seq.is_some(); + + // Validate: offset must be 0 in anchor mode + if has_anchor && offset != 0 { + return Err("--offset is not supported with anchor queries (--around-sort-seq, --around-server-id, --after-sort-seq)".into()); + } + + // Warn if --context is used without --around-* + let has_around = around_sort_seq.is_some() || around_server_id.is_some(); + if context.is_some() && !has_around { + eprintln!("warning: --context is only meaningful with --around-sort-seq or --around-server-id, ignoring"); + } + + let result = if has_anchor { + let mut query = wx_db::MessageQuery::for_talker(&talker); + + if let Some(seq) = around_sort_seq { + query = query.around_sort_seq(seq); + } else if let Some(id) = around_server_id { + query = query.around_server_id(id); + } else if let Some(seq) = after_sort_seq { + query = query.after_sort_seq(seq).limit(limit); + } + + if has_around { + if let Some(ctx) = context { + query = query.context(ctx); + } + } + + if let Some(ref mt_str) = msg_type { + if let Some(mt) = wx_db::parse_msg_type(mt_str) { + query = query.msg_type(mt); + } + } + + db.query_messages_anchor(&query)? + } else { + let effective_limit = if all { + wx_db::MAX_QUERY_LIMIT + } else { + wx_db::effective_limit(limit) + }; + + let mut query = wx_db::MessageQuery::for_talker(&talker) + .limit(effective_limit) + .offset(offset) + .order(order.into()); + + if let Some(s) = since { + query = query.since(s); + } + if let Some(u) = until { + query = query.until(u); + } + + if let Some(ref mt_str) = msg_type { + if let Some(mt) = wx_db::parse_msg_type(mt_str) { + query = query.msg_type(mt); + } + } + + db.query_messages(&query)? + }; + + for w in &result.shard_warnings { + eprintln!("warning: shard {}: {}", w.path, w.reason); + } + if result.stats.skipped > 0 { + eprintln!( + "warning: {} messages skipped (decode error)", + result.stats.skipped + ); + } + + let is_group = wx_db::is_group_chat(&talker); + let effective_limit = if has_anchor { + result.items.len() + } else if all { + wx_db::MAX_QUERY_LIMIT + } else { + wx_db::effective_limit(limit) + }; + + let mut envelope = + JsonEnvelope::from_message_query_result(result, effective_limit, offset, |m| { + enrich_message(m, self_wxid, &resolver) + }); + + // When limit pushdown was used (non-anchor, non-all), total_rows only reflects the + // scanned window. Use a lightweight COUNT(*) query to get the actual DB-level total. + if !has_anchor && !all { + let mt_filter = msg_type + .as_ref() + .and_then(|s| wx_db::parse_msg_type(s)); + let db_total = db.count_messages( + &talker, + since.unwrap_or(0), + until.unwrap_or(i64::MAX), + mt_filter, + ); + envelope.paging.total = db_total; + envelope.paging.has_more = offset + envelope.paging.returned < db_total; + } + + // Phase 2: sender-level projection + let projected = project_message_items(envelope.items, &talker, &visibility, show_hidden); + envelope.paging.returned = projected.len(); + envelope.items = projected; + + Ok((envelope, is_group)) +} + +#[allow(clippy::too_many_arguments)] +fn fetch_remote_query( + client: &ThinClient, + contact: &str, + since: Option, + until: Option, + msg_type: Option, + limit: usize, + offset: usize, + order: SortOrderArg, + around_sort_seq: Option, + around_server_id: Option, + context: Option, + after_sort_seq: Option, + show_hidden: bool, +) -> Result, super::thin_client::ThinClientError> { + let mut query = vec![ + ("contact".to_string(), contact.to_string()), + ("limit".to_string(), limit.to_string()), + ("offset".to_string(), offset.to_string()), + ( + "order".to_string(), + match order { + SortOrderArg::Asc => "asc".to_string(), + SortOrderArg::Desc => "desc".to_string(), + }, + ), + ]; + if let Some(since) = since { + query.push(("since".to_string(), since.to_string())); + } + if let Some(until) = until { + query.push(("until".to_string(), until.to_string())); + } + if let Some(msg_type) = msg_type { + query.push(("type".to_string(), msg_type)); + } + if let Some(seq) = around_sort_seq { + query.push(("around_sort_seq".to_string(), seq.to_string())); + } + if let Some(id) = around_server_id { + query.push(("around_server_id".to_string(), id.to_string())); + } + if let Some(context) = context { + query.push(("context".to_string(), context.to_string())); + } + if let Some(seq) = after_sort_seq { + query.push(("after_sort_seq".to_string(), seq.to_string())); + } + if show_hidden { + query.push(("show_hidden".to_string(), "1".to_string())); + } + client.get_json("/api/v1/messages", &query) +} + +fn print_query_output( + envelope: &JsonEnvelope, + format: OutputFormat, + is_group: bool, +) -> Result<(), Box> { + match format { + OutputFormat::Json => println!("{}", serde_json::to_string_pretty(envelope)?), + OutputFormat::Text => render_query_text(&envelope.items, is_group), + } + Ok(()) +} + +fn render_query_text(items: &[EnrichedMessage], is_group: bool) { + for item in items { + let ts = chrono::DateTime::from_timestamp(item.message.create_time, 0) + .map(|dt| { + dt.with_timezone(&chrono::Local) + .format("%m-%d %H:%M") + .to_string() + }) + .unwrap_or_default(); + + if is_group { + let sender_name = if item.direction == Direction::Outgoing { + "我" + } else { + display_name_only(&item.sender_display_name, &item.message.sender) + }; + println!("{ts} [{sender_name}] {}", item.snippet); + } else { + let arrow = match item.direction { + Direction::Incoming => "<<", + Direction::Outgoing => ">>", + }; + println!("{ts} {arrow} {}", item.snippet); + } + } +} + +fn display_name_only<'a>(display: &'a str, wxid: &str) -> &'a str { + let suffix = format!("({wxid})"); + display + .strip_suffix(&suffix) + .filter(|name| !name.is_empty()) + .unwrap_or(display) +} diff --git a/crates/wx-cli/src/cmd/search.rs b/crates/wx-cli/src/cmd/search.rs new file mode 100644 index 0000000..b5db28e --- /dev/null +++ b/crates/wx-cli/src/cmd/search.rs @@ -0,0 +1,305 @@ +use std::path::PathBuf; + +use wx_context::{ + open_fts_connection_with_key, AccountContext, ContactResolver, ResolveParams, +}; + +use super::thin_client::{ThinClient, ThinClientCliArgs, ThinClientOptions}; +use crate::output::{JsonEnvelope, PagingMeta, StatsMeta}; +use crate::schema::{enrich_message_as_hit, enrich_native_fts_hit, SearchHit}; +use crate::util::{effective_limit_all, open_db_all, print_cache_stats, print_detection_note, try_remote_or_local}; +use crate::OutputFormat; + +// Unused imports kept for Task 6 cleanup reference: +// use crate::util::print_fts_stats; +// use wx_db::{FtsSearchResult}; + +#[allow(clippy::too_many_arguments)] +pub fn cmd_search( + keyword: &str, + data_dir: Option, + account: Option, + key: Option, + limit: usize, + offset: usize, + all: bool, + format: OutputFormat, + server: ThinClientCliArgs, +) -> Result<(), Box> { + let options = ThinClientOptions::resolve_from_process_env(server); + let effective_limit = effective_limit_all(all, limit); + + let envelope = try_remote_or_local( + &options, + |client| fetch_remote_search(client, keyword, effective_limit, offset), + || load_local_search(keyword, data_dir, account, key, effective_limit, offset), + "search", + )?; + print_search_output(&envelope, format) +} + +fn load_local_search( + keyword: &str, + data_dir: Option, + account: Option, + key: Option, + effective_limit: usize, + offset: usize, +) -> Result, Box> { + let acct = AccountContext::resolve(&ResolveParams { + account: account.as_deref(), + data_dir: data_dir.as_deref(), + key_hex: key.as_deref(), + })?; + print_detection_note(&acct); + + let (db, _cache, stats) = open_db_all(&acct, crate::util::decrypt_progress_callback)?; + if let Some(ref s) = stats { + print_cache_stats(s); + } + let resolver = ContactResolver::build(&db)?; + let self_wxid = &acct.base_wxid; + + let search_start = std::time::Instant::now(); + + // --- Native FTS search → fallback to scan --- + let use_fallback = match db.message_fts_path.as_deref() { + Some(fts_path) => match open_fts_connection_with_key(fts_path, acct.raw_key.as_ref()) { + Ok(conn) => { + match wx_db::native_fts::search_message_fts( + &conn, + keyword, + effective_limit, + offset, + ) { + Ok(result) => { + return native_fts_envelope( + result, + &resolver, + self_wxid, + effective_limit, + offset, + &search_start, + ); + } + Err(e) => { + eprintln!("Native FTS query failed, falling back to scan: {e}"); + true + } + } + } + Err(e) => { + eprintln!("Cannot open FTS DB, falling back to scan: {e}"); + true + } + }, + None => { + // No native FTS DB available — silent fallback + true + } + }; + + debug_assert!(use_fallback); + scan_envelope( + keyword, + &db, + &resolver, + self_wxid, + effective_limit, + offset, + &search_start, + ) +} + +fn fetch_remote_search( + client: &ThinClient, + keyword: &str, + limit: usize, + offset: usize, +) -> Result, super::thin_client::ThinClientError> { + let query = vec![ + ("q".to_string(), keyword.to_string()), + ("limit".to_string(), limit.to_string()), + ("offset".to_string(), offset.to_string()), + ]; + client.get_json("/api/v1/search", &query) +} + +fn print_search_output( + envelope: &JsonEnvelope, + format: OutputFormat, +) -> Result<(), Box> { + match format { + OutputFormat::Json => println!("{}", serde_json::to_string_pretty(envelope)?), + OutputFormat::Text => render_search_text(&envelope.items), + } + Ok(()) +} + +fn render_search_text(items: &[SearchHit]) { + for hit in items { + let ts = chrono::DateTime::from_timestamp(hit.create_time, 0) + .map(|dt| { + dt.with_timezone(&chrono::Local) + .format("%m-%d %H:%M") + .to_string() + }) + .unwrap_or_default(); + + if wx_db::is_group_chat(&hit.talker) { + println!( + "{ts} [群「{}」] {}: {}", + display_name_only(&hit.talker_display_name, &hit.talker), + display_name_only(&hit.sender_display_name, &hit.sender), + hit.snippet + ); + } else { + println!( + "{ts} [{}] {}", + display_name_only(&hit.talker_display_name, &hit.talker), + hit.snippet + ); + } + } +} + +fn display_name_only<'a>(display: &'a str, wxid: &str) -> &'a str { + let suffix = format!("({wxid})"); + display + .strip_suffix(&suffix) + .filter(|name| !name.is_empty()) + .unwrap_or(display) +} + +/// Native FTS path: return the existing JSON envelope contract. +fn native_fts_envelope( + result: wx_db::NativeFtsResult, + resolver: &ContactResolver, + self_wxid: &str, + limit: usize, + offset: usize, + search_start: &std::time::Instant, +) -> Result, Box> { + let total_hits = result.total_hits; + let returned = result.hits.len(); + + let enriched: Vec = result + .hits + .into_iter() + .map(|hit| enrich_native_fts_hit(hit, self_wxid, resolver)) + .collect(); + Ok(JsonEnvelope { + items: enriched, + paging: PagingMeta { + limit, + offset, + returned, + has_more: offset + returned < total_hits, + total: total_hits, + }, + stats: StatsMeta { + scanned: 0, + skipped: 0, + elapsed_ms: Some(search_start.elapsed().as_millis() as u64), + shard_warnings: Vec::new(), + }, + }) +} + +/// Existing scan path: iterate all sessions, query_messages per session, collect + sort + paginate. +#[allow(clippy::too_many_arguments)] +fn scan_envelope( + keyword: &str, + db: &wx_db::WechatDb, + resolver: &ContactResolver, + self_wxid: &str, + effective_limit: usize, + offset: usize, + search_start: &std::time::Instant, +) -> Result, Box> { + // Paginate through ALL sessions so we never miss conversations. + let mut all_sessions = Vec::new(); + let page_size = wx_db::MAX_QUERY_LIMIT; + let mut sess_offset = 0; + loop { + let page = db.query_sessions( + &wx_db::SessionQuery::new() + .limit(page_size) + .offset(sess_offset), + )?; + if page.items.is_empty() { + break; + } + sess_offset += page.items.len(); + all_sessions.extend(page.items); + if sess_offset >= page.stats.total_rows { + break; + } + } + + let mut all_hits: Vec<(wx_db::Message, String)> = Vec::new(); + let mut total_scanned: usize = 0; + let mut total_skipped: usize = 0; + let mut all_shard_warnings: Vec = Vec::new(); + + for session in &all_sessions { + // BUG FIX: per-session query uses MAX_QUERY_LIMIT (not user limit) + let query = wx_db::MessageQuery::for_talker(&session.username) + .keyword(keyword) + .limit(wx_db::MAX_QUERY_LIMIT); + + let result = db.query_messages(&query)?; + total_scanned += result.stats.total_rows; + total_skipped += result.stats.skipped; + all_shard_warnings.extend(result.shard_warnings); + for msg in result.items { + all_hits.push((msg, session.username.clone())); + } + } + + for w in &all_shard_warnings { + eprintln!("warning: shard {}: {}", w.path, w.reason); + } + if total_skipped > 0 { + eprintln!("warning: {} messages skipped (decode error)", total_skipped); + } + + // Sort by (sort_seq DESC, create_time DESC, server_id DESC) to align with + // query_messages() stable ordering. server_id provides a unique tie-breaker. + all_hits.sort_by(|a, b| { + b.0.sort_seq + .cmp(&a.0.sort_seq) + .then_with(|| b.0.create_time.cmp(&a.0.create_time)) + .then_with(|| b.0.server_id.cmp(&a.0.server_id)) + }); + + let total_hits = all_hits.len(); + // Apply offset + limit + let page: Vec<_> = all_hits + .into_iter() + .skip(offset) + .take(effective_limit) + .collect(); + + let enriched: Vec<_> = page + .into_iter() + .map(|(m, talker)| enrich_message_as_hit(m, talker, self_wxid, resolver)) + .collect(); + let returned = enriched.len(); + Ok(JsonEnvelope { + items: enriched, + paging: PagingMeta { + limit: effective_limit, + offset, + returned, + has_more: offset + returned < total_hits, + total: total_hits, + }, + stats: StatsMeta { + scanned: total_scanned, + skipped: total_skipped, + elapsed_ms: Some(search_start.elapsed().as_millis() as u64), + shard_warnings: all_shard_warnings, + }, + }) +} diff --git a/crates/wx-cli/src/cmd/serve/auth.rs b/crates/wx-cli/src/cmd/serve/auth.rs new file mode 100644 index 0000000..5d2c39b --- /dev/null +++ b/crates/wx-cli/src/cmd/serve/auth.rs @@ -0,0 +1,37 @@ +use std::sync::Arc; + +use axum::extract::State; +use axum::http::{Request, StatusCode}; +use axum::middleware::Next; +use axum::response::{IntoResponse, Response}; +use serde_json::json; + +use super::state::AppState; + +pub async fn bearer_auth( + State(state): State>, + request: Request, + next: Next, +) -> Response { + let Some(expected) = &state.auth_token else { + return next.run(request).await; + }; + + let authorized = request + .headers() + .get("authorization") + .and_then(|v| v.to_str().ok()) + .and_then(|v| v.strip_prefix("Bearer ")) + .map(|token| token == expected) + .unwrap_or(false); + + if authorized { + next.run(request).await + } else { + ( + StatusCode::UNAUTHORIZED, + axum::Json(json!({ "error": "unauthorized" })), + ) + .into_response() + } +} diff --git a/crates/wx-cli/src/cmd/serve/bridge.rs b/crates/wx-cli/src/cmd/serve/bridge.rs new file mode 100644 index 0000000..389e971 --- /dev/null +++ b/crates/wx-cli/src/cmd/serve/bridge.rs @@ -0,0 +1,458 @@ +use std::collections::HashMap; +use std::sync::Arc; +use std::time::Duration; + +use tokio::sync::{mpsc, watch}; +use tokio_util::sync::CancellationToken; +use wx_context::ContactResolver; +use wx_monitor::SessionEvent; + +use crate::schema::{ + enrich_message, enrich_session_event, project_message_items, project_session_sender, + EnrichedMessage, +}; + +use super::event::{MessagePayload, SessionPayload, SseEvent}; +use super::refresh::RefreshTrigger; +use super::state::AppState; + +struct BridgeState { + cursors: HashMap, + /// Global fallback baseline for talkers not seen at startup. + /// Set to max(sort_seq) across all talkers at serve startup time. + startup_watermark: i64, +} + +fn should_broadcast_talker( + visibility: &wx_context::VisibilityIndex, + talker: &str, +) -> bool { + !visibility.is_hidden_talker(talker) +} + +/// Initialize per-talker baselines and a global startup_watermark. +/// +/// Must be called **before** `WechatMonitor::start()` to avoid a race where +/// a message arriving between monitor start and baseline completion gets +/// absorbed into the baseline and never pushed to SSE clients. +pub fn init_baselines(db: &wx_db::WechatDb) -> Result<(HashMap, i64), String> { + let sessions = db + .query_sessions(&wx_db::SessionQuery::new().limit(10000)) + .map_err(|e| format!("query_sessions failed: {e}"))?; + + let usernames: Vec = sessions.items.iter().map(|s| s.username.clone()).collect(); + let cursors = db.bulk_max_sort_seq(&usernames); + let global_max = cursors.values().copied().max().unwrap_or(0); + + eprintln!( + "bridge: initialized baselines for {} talkers, startup_watermark={}", + cursors.len(), + global_max + ); + Ok((cursors, global_max)) +} + +/// Result from the blocking DB query, returned to the async context for +/// state updates and event broadcasting. +struct BridgeUpdate { + /// Effective cursor value (always `Some` in the unified path; `None` not produced). + last_sort_seq: Option, + /// Enriched messages to broadcast (empty when no new messages or on failure). + messages: Vec, +} + +pub async fn run_bridge( + mut receiver: mpsc::Receiver, + state: Arc, + cursors: HashMap, + startup_watermark: i64, + refresh_tx: mpsc::Sender, + mut refresh_watch: watch::Receiver, + shutdown: CancellationToken, +) { + let mut bridge_state = BridgeState { + cursors, + startup_watermark, + }; + let mut heartbeat = tokio::time::interval(Duration::from_secs(15)); + loop { + tokio::select! { + event = receiver.recv() => { + match event { + Some(ev) => handle_session_event( + ev, + &state, + &mut bridge_state, + &refresh_tx, + &mut refresh_watch, + &shutdown, + ).await, + None => break, + } + } + _ = heartbeat.tick() => { + let _ = state.broadcast_tx.send(Arc::new(SseEvent::Heartbeat)); + } + _ = shutdown.cancelled() => break, + } + } +} + +async fn handle_session_event( + ev: SessionEvent, + state: &AppState, + bridge_state: &mut BridgeState, + refresh_tx: &mpsc::Sender, + refresh_watch: &mut watch::Receiver, + shutdown: &CancellationToken, +) { + if !should_broadcast_talker(&state.visibility, &ev.username) { + return; + } + + // Read cursor snapshot BEFORE triggering refresh (async context only). + let baseline: i64 = bridge_state + .cursors + .get(&ev.username) + .copied() + .unwrap_or(bridge_state.startup_watermark); + + // Snapshot current epoch, then trigger refresh + let epoch_before = *refresh_watch.borrow(); + if refresh_tx.send(RefreshTrigger::Refresh).await.is_err() { + eprintln!("warn: bridge refresh_tx send failed (channel closed)"); + return; + } + + // Wait for refresh to complete (epoch advances past our snapshot). + // Timeout after 30s to avoid blocking forever if refresh keeps failing. + // Also abort if shutdown fires during the wait. + let wait_result = tokio::time::timeout(Duration::from_secs(30), async { + loop { + if *refresh_watch.borrow() > epoch_before { + return true; + } + tokio::select! { + result = refresh_watch.changed() => { + if result.is_err() { + return false; + } + } + _ = shutdown.cancelled() => return false, + } + } + }) + .await; + + match wait_result { + Ok(true) => {} // refresh succeeded + Ok(false) => { + if shutdown.is_cancelled() { + return; + } + eprintln!("warn: bridge refresh_watch closed"); + return; + } + Err(_) => { + eprintln!("warn: bridge refresh wait timed out (30s), skipping event"); + return; + } + } + + // DB query in blocking context; returns None on DB lock failure + let db = Arc::clone(&state.db); + let resolver = Arc::clone(&state.resolver); + let visibility = Arc::clone(&state.visibility); + let self_wxid = state.self_wxid.clone(); + let username = ev.username.clone(); + + let update = tokio::task::spawn_blocking(move || { + let guard = match db.lock() { + Ok(g) => g, + Err(e) => { + eprintln!("warn: bridge db lock failed: {e}"); + return None; + } + }; + + // Unified path: always query after baseline (works for both first and subsequent events) + let query = wx_db::MessageQuery::for_talker(&username).after_sort_seq(baseline); + match guard.query_messages_anchor(&query) { + Ok(result) => { + for w in &result.shard_warnings { + eprintln!("warn: bridge shard {}: {}", w.path, w.reason); + } + let max_seq = result.items.iter().map(|m| m.sort_seq).max(); + let effective_seq = max_seq.unwrap_or(baseline); + let messages = + enrich_messages(result.items, &self_wxid, &resolver, &username, &visibility); + Some(BridgeUpdate { + last_sort_seq: Some(effective_seq), + messages, + }) + } + Err(e) => { + eprintln!("warn: bridge query_messages_anchor failed: {e}"); + Some(BridgeUpdate { + last_sort_seq: Some(baseline), + messages: vec![], + }) + } + } + }) + .await; + + // Async context: update per-talker cursors and broadcast events + let username = ev.username.clone(); + + let (last_sort_seq, messages) = apply_update(&mut bridge_state.cursors, &username, update); + + // Session event (always sent) + let mut enriched = enrich_session_event(ev, &state.self_wxid, &state.resolver); + // Phase 2: redact hidden sender in session summary + project_session_sender(&mut enriched, &state.visibility); + let session_payload = SessionPayload { + enriched, + last_sort_seq, + }; + let _ = state + .broadcast_tx + .send(Arc::new(SseEvent::Session(session_payload))); + + // Message event (when there are new messages after baseline) + if !messages.is_empty() { + let payload = MessagePayload { + talker: username.clone(), + talker_display_name: state.resolver.display_with_id(&username), + messages, + anchor_sort_seq: last_sort_seq, + }; + let _ = state + .broadcast_tx + .send(Arc::new(SseEvent::Message(payload))); + } +} + +fn enrich_messages( + items: Vec, + self_wxid: &str, + resolver: &ContactResolver, + talker: &str, + visibility: &wx_context::VisibilityIndex, +) -> Vec { + let enriched: Vec = items + .into_iter() + .map(|m| enrich_message(m, self_wxid, resolver)) + .collect(); + // Phase 2: sender-level projection (SSE has no bypass) + project_message_items(enriched, talker, visibility, false) +} + +/// Apply a BridgeUpdate to the cursor map and return (last_sort_seq, messages). +/// +/// Cursor update rules (unified path — no first/subsequent distinction): +/// - Query ok + new messages: cursor advanced to max(sort_seq) of new messages +/// - Query ok + no new messages: cursor set to baseline (unchanged) +/// - Query failed: cursor set to baseline (unchanged) +/// - DB lock failed / spawn panic: cursor unchanged, fallback to existing value +fn apply_update( + cursors: &mut HashMap, + username: &str, + update: Result, tokio::task::JoinError>, +) -> (Option, Vec) { + match update { + Ok(Some(u)) => { + if let Some(seq) = u.last_sort_seq { + cursors.insert(username.to_string(), seq); + } + (u.last_sort_seq, u.messages) + } + _ => { + // spawn_blocking panicked or DB lock failed - use existing cursor if any + let fallback_seq = cursors.get(username).copied(); + (fallback_seq, vec![]) + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use wx_context::{ContactResolver, VisibilityIndex}; + + fn make_msg(sort_seq: i64) -> EnrichedMessage { + EnrichedMessage { + message: wx_db::Message { + sort_seq, + server_id: 0, + msg_type: 1, + sub_type: 0, + sender: String::new(), + talker: String::new(), + create_time: 0, + content: wx_db::MessageContent::Text(String::new()), + status: 0, + }, + sender_display_name: String::new(), + direction: wx_context::Direction::Incoming, + snippet: String::new(), + } + } + + #[test] + fn first_event_initializes_cursor_with_messages() { + let mut cursors = HashMap::new(); + let update = Ok(Some(BridgeUpdate { + last_sort_seq: Some(600), + messages: vec![make_msg(600)], + })); + let (last_sort_seq, messages) = apply_update(&mut cursors, "wxid_alice", update); + + assert_eq!(last_sort_seq, Some(600)); + assert_eq!(messages.len(), 1); + assert_eq!(cursors.get("wxid_alice"), Some(&600)); + } + + #[test] + fn first_event_no_new_messages_cursor_at_baseline() { + let mut cursors = HashMap::new(); + let update = Ok(Some(BridgeUpdate { + last_sort_seq: Some(500), + messages: vec![], + })); + let (last_sort_seq, messages) = apply_update(&mut cursors, "wxid_alice", update); + + assert_eq!(last_sort_seq, Some(500)); + assert!(messages.is_empty()); + assert_eq!(cursors.get("wxid_alice"), Some(&500)); + } + + #[test] + fn subsequent_event_with_new_messages_advances_cursor() { + let mut cursors = HashMap::new(); + cursors.insert("wxid_alice".to_string(), 500); + let update = Ok(Some(BridgeUpdate { + last_sort_seq: Some(800), + messages: vec![make_msg(600), make_msg(800)], + })); + let (last_sort_seq, messages) = apply_update(&mut cursors, "wxid_alice", update); + + assert_eq!(last_sort_seq, Some(800)); + assert_eq!(messages.len(), 2); + assert_eq!(cursors.get("wxid_alice"), Some(&800)); + } + + #[test] + fn subsequent_event_no_new_messages_keeps_cursor() { + let mut cursors = HashMap::new(); + cursors.insert("wxid_alice".to_string(), 500); + let update = Ok(Some(BridgeUpdate { + last_sort_seq: Some(500), + messages: vec![], + })); + let (last_sort_seq, messages) = apply_update(&mut cursors, "wxid_alice", update); + + assert_eq!(last_sort_seq, Some(500)); + assert!(messages.is_empty()); + assert_eq!(cursors.get("wxid_alice"), Some(&500)); + } + + #[test] + fn query_failed_reuses_baseline() { + let mut cursors = HashMap::new(); + let update = Ok(Some(BridgeUpdate { + last_sort_seq: Some(500), + messages: vec![], + })); + let (last_sort_seq, messages) = apply_update(&mut cursors, "wxid_alice", update); + + assert_eq!(last_sort_seq, Some(500)); + assert!(messages.is_empty()); + assert_eq!(cursors.get("wxid_alice"), Some(&500)); + } + + #[test] + fn db_lock_failed_uses_existing_cursor() { + let mut cursors = HashMap::new(); + cursors.insert("wxid_alice".to_string(), 500); + let update: Result, _> = Ok(None); + let (last_sort_seq, messages) = apply_update(&mut cursors, "wxid_alice", update); + + assert_eq!(last_sort_seq, Some(500)); + assert!(messages.is_empty()); + assert_eq!(cursors.get("wxid_alice"), Some(&500)); + } + + #[test] + fn db_lock_failed_no_prior_cursor() { + let mut cursors = HashMap::new(); + let update: Result, _> = Ok(None); + let (last_sort_seq, messages) = apply_update(&mut cursors, "wxid_alice", update); + + assert_eq!(last_sort_seq, None); + assert!(messages.is_empty()); + assert!(!cursors.contains_key("wxid_alice")); + } + + #[test] + fn hidden_talker_is_not_broadcast() { + let visibility = + VisibilityIndex::build(&["wxid_secret".to_string()], &[], &ContactResolver::empty()); + + assert!(!should_broadcast_talker(&visibility, "wxid_secret")); + assert!(should_broadcast_talker(&visibility, "wxid_visible")); + } + + // --- Phase 2: sender-level tests --- + + #[test] + fn enrich_messages_filters_hidden_sender_in_group() { + let visibility = VisibilityIndex::build( + &["wxid_spam".to_string()], + &[], + &ContactResolver::empty(), + ); + let msgs = vec![ + wx_db::Message { + sort_seq: 1, server_id: 1, msg_type: 1, sub_type: 0, + sender: "wxid_spam".to_string(), talker: "group@chatroom".to_string(), + create_time: 100, content: wx_db::MessageContent::Text("spam".into()), status: 0, + }, + wx_db::Message { + sort_seq: 2, server_id: 2, msg_type: 1, sub_type: 0, + sender: "wxid_normal".to_string(), talker: "group@chatroom".to_string(), + create_time: 101, content: wx_db::MessageContent::Text("hello".into()), status: 0, + }, + ]; + + let result = enrich_messages(msgs, "wxid_me", &ContactResolver::empty(), "group@chatroom", &visibility); + assert_eq!(result.len(), 1, "hidden sender message should be filtered"); + assert_eq!(result[0].message.sender, "wxid_normal"); + } + + #[test] + fn session_sender_redaction_in_bridge() { + use crate::schema::project_session_sender; + let visibility = VisibilityIndex::build( + &["wxid_spam".to_string()], + &[], + &ContactResolver::empty(), + ); + let ev = wx_monitor::SessionEvent { + username: "group@chatroom".to_string(), + sort_timestamp: 1, + detected_at: 2, + kind: wx_monitor::SessionEventKind::Updated, + summary: "spam message".to_string(), + last_msg_type: Some(1), + last_msg_sender: Some("wxid_spam".to_string()), + last_sender_display_name: Some("Spammer".to_string()), + }; + let mut enriched = crate::schema::enrich_session_event(ev, "wxid_me", &ContactResolver::empty()); + project_session_sender(&mut enriched, &visibility); + + assert_eq!(enriched.session.summary, "[消息已隐藏]"); + assert!(enriched.session.last_msg_sender.is_none()); + assert!(enriched.direction.is_none()); + } +} diff --git a/crates/wx-cli/src/cmd/serve/error.rs b/crates/wx-cli/src/cmd/serve/error.rs new file mode 100644 index 0000000..be649ae --- /dev/null +++ b/crates/wx-cli/src/cmd/serve/error.rs @@ -0,0 +1,27 @@ +use axum::http::StatusCode; +use axum::response::{IntoResponse, Response}; +use serde_json::json; + +pub enum ServeError { + Db(String), + Internal(String), + InvalidParam(String), + NotFound(String), + UnsupportedMedia(String), + Upstream(String), +} + +impl IntoResponse for ServeError { + fn into_response(self) -> Response { + let (status, message) = match self { + ServeError::Db(msg) => (StatusCode::INTERNAL_SERVER_ERROR, msg), + ServeError::Internal(msg) => (StatusCode::INTERNAL_SERVER_ERROR, msg), + ServeError::InvalidParam(msg) => (StatusCode::BAD_REQUEST, msg), + ServeError::NotFound(msg) => (StatusCode::NOT_FOUND, msg), + ServeError::UnsupportedMedia(msg) => (StatusCode::UNSUPPORTED_MEDIA_TYPE, msg), + ServeError::Upstream(msg) => (StatusCode::BAD_GATEWAY, msg), + }; + let body = axum::Json(json!({ "error": message })); + (status, body).into_response() + } +} diff --git a/crates/wx-cli/src/cmd/serve/event.rs b/crates/wx-cli/src/cmd/serve/event.rs new file mode 100644 index 0000000..2cb960b --- /dev/null +++ b/crates/wx-cli/src/cmd/serve/event.rs @@ -0,0 +1,121 @@ +use serde::Serialize; + +use crate::schema::{EnrichedMessage, EnrichedSession}; + +/// SSE event types broadcast to connected clients. +#[derive(Clone, Serialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum SseEvent { + Session(SessionPayload), + Message(MessagePayload), + Heartbeat, +} + +#[derive(Clone, Serialize)] +pub struct SessionPayload { + #[serde(flatten)] + pub enriched: EnrichedSession, + #[serde(skip_serializing_if = "Option::is_none")] + pub last_sort_seq: Option, +} + +#[derive(Clone, Serialize)] +pub struct MessagePayload { + pub talker: String, + pub talker_display_name: String, + pub messages: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + pub anchor_sort_seq: Option, +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::Value; + use wx_context::Direction; + + #[test] + fn session_payload_omits_direction_when_unknown() { + let payload = SessionPayload { + enriched: EnrichedSession { + session: wx_db::Session { + username: "wxid_friend".to_string(), + summary: "hello".to_string(), + sort_timestamp: 1, + last_msg_type: Some(1), + last_msg_sender: None, + last_sender_display_name: None, + }, + display_name: "wxid_friend".to_string(), + direction: None, + detected_at: None, + }, + last_sort_seq: None, + }; + + let json = serde_json::to_value(&payload).unwrap(); + assert!(json.get("direction").is_none()); + } + + #[test] + fn session_sse_event_omits_direction_when_unknown() { + let event = SseEvent::Session(SessionPayload { + enriched: EnrichedSession { + session: wx_db::Session { + username: "wxid_friend".to_string(), + summary: "hello".to_string(), + sort_timestamp: 1, + last_msg_type: Some(1), + last_msg_sender: None, + last_sender_display_name: None, + }, + display_name: "wxid_friend".to_string(), + direction: None, + detected_at: Some(2), + }, + last_sort_seq: Some(3), + }); + + let json = serde_json::to_value(&event).unwrap(); + assert_eq!( + json.get("type"), + Some(&Value::String("session".to_string())) + ); + assert_eq!(json.get("detected_at"), Some(&Value::Number(2.into()))); + assert!(json.get("direction").is_none()); + } + + #[test] + fn message_payload_keeps_message_direction() { + let payload = MessagePayload { + talker: "wxid_friend".to_string(), + talker_display_name: "wxid_friend".to_string(), + messages: vec![EnrichedMessage { + message: wx_db::Message { + sort_seq: 1, + server_id: 2, + msg_type: 1, + sub_type: 0, + sender: "wxid_me".to_string(), + talker: "wxid_friend".to_string(), + create_time: 3, + content: wx_db::MessageContent::Text("hello".to_string()), + status: 0, + }, + sender_display_name: "wxid_me".to_string(), + direction: Direction::Outgoing, + snippet: "hello".to_string(), + }], + anchor_sort_seq: Some(1), + }; + + let json = serde_json::to_value(&payload).unwrap(); + assert_eq!( + json.get("messages") + .and_then(Value::as_array) + .and_then(|items| items.first()) + .and_then(|item| item.get("direction")), + Some(&Value::String("outgoing".to_string())) + ); + } +} diff --git a/crates/wx-cli/src/cmd/serve/handlers.rs b/crates/wx-cli/src/cmd/serve/handlers.rs new file mode 100644 index 0000000..dcc89f1 --- /dev/null +++ b/crates/wx-cli/src/cmd/serve/handlers.rs @@ -0,0 +1,735 @@ +use std::convert::Infallible; +use std::sync::atomic::Ordering; +use std::sync::Arc; +use std::time::Duration; + +use axum::extract::{Query, Request, State}; +use axum::http::StatusCode; +use axum::response::sse::{Event, KeepAlive, Sse}; +use axum::response::IntoResponse; +use axum::Json; +use serde::Deserialize; +use serde_json::json; +use tokio_stream::wrappers::BroadcastStream; +use tokio_stream::StreamExt; + +use crate::cmd::server::types::{RuntimeAccountState, ServerHealthPayload}; +use crate::output::JsonEnvelope; +use crate::schema::{ + enrich_message, enrich_message_as_hit, enrich_native_fts_hit, enrich_session, + project_message_items, +}; +use crate::visibility_projection::{ + project_contacts_envelope, project_sessions_envelope_enriched, +}; + +use super::error::ServeError; +use super::event::SseEvent; +use super::media::{self, MediaRequest}; +use super::state::AppState; + +fn default_limit() -> usize { + 20 +} + +fn default_order() -> String { + "desc".to_string() +} + +fn parse_order(s: &str) -> wx_db::SortOrder { + match s.to_lowercase().as_str() { + "asc" => wx_db::SortOrder::Asc, + _ => wx_db::SortOrder::Desc, + } +} + +pub async fn handler_health( + State(state): State>, +) -> Result { + Ok(Json(ServerHealthPayload { + ready: state.ready.load(Ordering::Acquire), + worker_id: state.worker_id.clone(), + cli_version: state.cli_version.clone(), + current_account: RuntimeAccountState { + wxid: state.current_account.wxid.clone(), + name: state.current_account.name.clone(), + }, + })) +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::Value; + + #[test] + fn health_payload_can_include_current_account_metadata() { + let payload = serde_json::to_value(ServerHealthPayload { + ready: true, + worker_id: "worker-123".to_string(), + cli_version: "1.2.3 (abc1234 2026-03-25)".to_string(), + current_account: RuntimeAccountState { + wxid: "wxid_me".to_string(), + name: "Me".to_string(), + }, + }) + .unwrap(); + + assert_eq!(payload.get("ready"), Some(&Value::Bool(true))); + assert_eq!( + payload.get("worker_id"), + Some(&Value::String("worker-123".to_string())) + ); + assert_eq!( + payload.get("cli_version"), + Some(&Value::String("1.2.3 (abc1234 2026-03-25)".to_string())) + ); + assert_eq!( + payload.get("current_account").and_then(Value::as_object), + Some(&serde_json::Map::from_iter([ + ("wxid".to_string(), Value::String("wxid_me".to_string())), + ("name".to_string(), Value::String("Me".to_string())), + ])) + ); + } +} + +#[derive(Deserialize)] +pub struct MediaParams { + server_id: Option, + talker: Option, + format: Option, +} + +pub async fn handler_media( + State(state): State>, + Query(params): Query, + request: Request, +) -> Result { + let server_id = params + .server_id + .ok_or_else(|| ServeError::InvalidParam("missing required parameter: server_id".into()))?; + let talker = params + .talker + .filter(|value| !value.trim().is_empty()) + .ok_or_else(|| ServeError::InvalidParam("missing required parameter: talker".into()))?; + let format = media::MediaFormat::parse(params.format.as_deref())?; + + media::serve_media( + state, + MediaRequest { + server_id, + talker, + format, + }, + request, + ) + .await +} + +// --------------------------------------------------------------------------- +// Sessions +// --------------------------------------------------------------------------- + +#[derive(Deserialize)] +pub struct SessionParams { + #[serde(default = "default_limit")] + limit: usize, + #[serde(default)] + offset: usize, + #[serde(default = "default_order")] + order: String, + show_hidden: Option, +} + +pub async fn handler_sessions( + State(state): State>, + Query(params): Query, +) -> Result { + let db = Arc::clone(&state.db); + let resolver = Arc::clone(&state.resolver); + let visibility = Arc::clone(&state.visibility); + let self_wxid = state.self_wxid.clone(); + let limit = wx_db::effective_limit(params.limit); + let offset = params.offset; + let order = parse_order(¶ms.order); + let show_hidden = matches!(params.show_hidden.as_deref(), Some("1") | Some("true")); + + let result = tokio::task::spawn_blocking(move || { + let guard = db.lock().map_err(|e| ServeError::Internal(e.to_string()))?; + let result = guard + .query_sessions( + &wx_db::SessionQuery::new() + .limit(wx_db::MAX_QUERY_LIMIT) + .offset(0) + .order(order), + ) + .map_err(|e| ServeError::Db(e.to_string()))?; + + let envelope = JsonEnvelope::from_query_result(result, limit, offset, |s| { + enrich_session(s, &self_wxid, &resolver, None) + }); + Ok::<_, ServeError>(project_sessions_envelope_enriched( + envelope.items, + &visibility, + limit, + offset, + &envelope.stats, + show_hidden, + )) + }) + .await + .map_err(|e| ServeError::Internal(e.to_string()))??; + + Ok(Json(result)) +} + +// --------------------------------------------------------------------------- +// Contacts +// --------------------------------------------------------------------------- + +#[derive(Deserialize)] +pub struct ContactParams { + search: Option, + #[serde(default = "default_limit")] + limit: usize, + #[serde(default)] + offset: usize, + show_hidden: Option, +} + +pub async fn handler_contacts( + State(state): State>, + Query(params): Query, +) -> Result { + let db = Arc::clone(&state.db); + let visibility = Arc::clone(&state.visibility); + let limit = wx_db::effective_limit(params.limit); + let offset = params.offset; + let search = params.search; + let show_hidden = matches!(params.show_hidden.as_deref(), Some("1") | Some("true")); + + let result = tokio::task::spawn_blocking(move || { + let guard = db.lock().map_err(|e| ServeError::Internal(e.to_string()))?; + let mut query = wx_db::ContactQuery::new() + .limit(wx_db::MAX_QUERY_LIMIT) + .offset(0); + if let Some(kw) = &search { + query = query.keyword(kw); + } + let result = guard + .query_contacts(&query) + .map_err(|e| ServeError::Db(e.to_string()))?; + + let envelope = + JsonEnvelope::from_query_result(result, wx_db::MAX_QUERY_LIMIT, 0, |c| c); + Ok::<_, ServeError>(project_contacts_envelope( + envelope.items, + &visibility, + limit, + offset, + &envelope.stats, + show_hidden, + )) + }) + .await + .map_err(|e| ServeError::Internal(e.to_string()))??; + + Ok(Json(result)) +} + +// --------------------------------------------------------------------------- +// Messages +// --------------------------------------------------------------------------- + +#[derive(Deserialize)] +pub struct MessageParams { + pub contact: Option, + #[serde(default = "default_limit")] + pub limit: usize, + #[serde(default)] + pub offset: usize, + pub since: Option, + pub until: Option, + #[serde(rename = "type")] + pub msg_type: Option, + #[serde(default = "default_order")] + pub order: String, + pub around_sort_seq: Option, + pub around_server_id: Option, + pub context: Option, + pub after_sort_seq: Option, + pub show_hidden: Option, +} + +pub async fn handler_messages( + State(state): State>, + Query(params): Query, +) -> Result { + let contact = params + .contact + .ok_or_else(|| ServeError::InvalidParam("missing required parameter: contact".into()))?; + + // Anchor mutual exclusion validation + let anchor_count = [ + params.around_sort_seq.is_some(), + params.around_server_id.is_some(), + params.after_sort_seq.is_some(), + ] + .iter() + .filter(|&&b| b) + .count(); + if anchor_count > 1 { + return Err(ServeError::InvalidParam( + "around_sort_seq, around_server_id, and after_sort_seq are mutually exclusive".into(), + )); + } + let has_anchor = anchor_count == 1; + + // Anchor vs time-range mutual exclusion + if has_anchor && (params.since.is_some() || params.until.is_some()) { + return Err(ServeError::InvalidParam( + "anchor parameters cannot be combined with since/until".into(), + )); + } + + // Offset must be 0 in anchor mode + if has_anchor && params.offset != 0 { + return Err(ServeError::InvalidParam( + "offset is not supported with anchor queries".into(), + )); + } + + let db = Arc::clone(&state.db); + let resolver = Arc::clone(&state.resolver); + let visibility = Arc::clone(&state.visibility); + let self_wxid = state.self_wxid.clone(); + let limit = wx_db::effective_limit(params.limit); + let offset = params.offset; + let order = parse_order(¶ms.order); + let contact = contact.clone(); + let since = params.since; + let until = params.until; + let msg_type = params.msg_type.clone(); + let around_sort_seq = params.around_sort_seq; + let around_server_id = params.around_server_id; + let after_sort_seq = params.after_sort_seq; + let context = params.context; + let show_hidden = matches!(params.show_hidden.as_deref(), Some("1") | Some("true")); + + let result = tokio::task::spawn_blocking(move || { + let guard = db.lock().map_err(|e| ServeError::Internal(e.to_string()))?; + + // Resolve contact — support wxid direct pass-through and fuzzy match + let talker = resolve_contact(&contact, &resolver, &guard, Some(&visibility), show_hidden)?; + + let result = if has_anchor { + let mut query = wx_db::MessageQuery::for_talker(&talker); + + if let Some(seq) = around_sort_seq { + query = query.around_sort_seq(seq); + } else if let Some(id) = around_server_id { + query = query.around_server_id(id); + } else if let Some(seq) = after_sort_seq { + query = query.after_sort_seq(seq).limit(limit); + } + + let has_around = around_sort_seq.is_some() || around_server_id.is_some(); + if has_around { + if let Some(ctx) = context { + query = query.context(ctx); + } + } + + if let Some(ref t) = msg_type { + if let Some(type_val) = wx_db::parse_msg_type(t) { + query = query.msg_type(type_val); + } + } + + guard + .query_messages_anchor(&query) + .map_err(|e| ServeError::Db(e.to_string()))? + } else { + let mut query = wx_db::MessageQuery::for_talker(&talker) + .limit(limit) + .offset(offset) + .order(order); + if let Some(s) = since { + query = query.since(s); + } + if let Some(u) = until { + query = query.until(u); + } + if let Some(ref t) = msg_type { + if let Some(type_val) = wx_db::parse_msg_type(t) { + query = query.msg_type(type_val); + } + } + + guard + .query_messages(&query) + .map_err(|e| ServeError::Db(e.to_string()))? + }; + + let mut envelope = JsonEnvelope::from_message_query_result(result, limit, offset, |m| { + enrich_message(m, &self_wxid, &resolver) + }); + + // When limit pushdown was used (non-anchor), total_rows only reflects the + // scanned window. Use a lightweight COUNT(*) query for accurate DB-level total. + if !has_anchor { + let mt_filter = msg_type + .as_ref() + .and_then(|s| wx_db::parse_msg_type(s)); + let db_total = guard.count_messages( + &talker, + since.unwrap_or(0), + until.unwrap_or(i64::MAX), + mt_filter, + ); + envelope.paging.total = db_total; + envelope.paging.has_more = offset + envelope.paging.returned < db_total; + } + + // Phase 2: sender-level projection + let projected = project_message_items(envelope.items, &talker, &visibility, show_hidden); + envelope.paging.returned = projected.len(); + envelope.items = projected; + + Ok::<_, ServeError>(envelope) + }) + .await + .map_err(|e| ServeError::Internal(e.to_string()))??; + + Ok(Json(result)) +} + +// --------------------------------------------------------------------------- +// Search +// --------------------------------------------------------------------------- + +#[derive(Deserialize)] +pub struct SearchParams { + q: Option, + #[serde(default = "default_limit")] + limit: usize, + #[serde(default)] + offset: usize, +} + +pub async fn handler_search( + State(state): State>, + Query(params): Query, +) -> Result { + let search_start = std::time::Instant::now(); + let q = params + .q + .ok_or_else(|| ServeError::InvalidParam("missing required parameter: q".into()))?; + let resolver = Arc::clone(&state.resolver); + let self_wxid = state.self_wxid.clone(); + let limit = wx_db::effective_limit(params.limit); + let offset = params.offset; + + // Phase 1: Try native FTS using the independent connection (NO main DB lock). + let fts_conn = state.fts_conn.clone(); + let q_clone = q.clone(); + let resolver_clone = Arc::clone(&resolver); + let self_wxid_clone = self_wxid.clone(); + let state_arc = Arc::clone(&state); + + let native_result: Option, ServeError>> = if let Some(fts_mutex) = + fts_conn + { + let result = tokio::task::spawn_blocking(move || { + let fts_guard = fts_mutex + .lock() + .map_err(|e: std::sync::PoisonError<_>| ServeError::Internal(e.to_string()))?; + // Lazy-init name2id cache with proper error propagation. + let name2id = { + let mut cache_guard = state_arc.name2id_cache.lock() + .map_err(|e: std::sync::PoisonError<_>| ServeError::Internal(e.to_string()))?; + match cache_guard.as_ref() { + Some(map) => map.clone(), + None => { + let loaded = wx_db::native_fts::load_name2id(&fts_guard) + .map_err(|e| ServeError::Db(e.to_string()))?; + *cache_guard = Some(loaded.clone()); + loaded + } + } + }; + match wx_db::native_fts::search_message_fts_with_cache( + &fts_guard, + &q_clone, + limit, + offset, + Some(&name2id), + ) { + Ok(result) => { + let total = result.total_hits; + let items: Vec<_> = result + .hits + .into_iter() + .map(|hit| enrich_native_fts_hit(hit, &self_wxid_clone, &resolver_clone)) + .collect(); + let returned = items.len(); + let has_more = offset + returned < total; + Ok(JsonEnvelope { + items, + paging: crate::output::PagingMeta { + limit, + offset, + returned, + has_more, + total, + }, + stats: crate::output::StatsMeta { + scanned: 0, + skipped: 0, + elapsed_ms: Some(search_start.elapsed().as_millis() as u64), + shard_warnings: Vec::new(), + }, + }) + } + Err(e) => { + eprintln!("warn: native FTS search failed, falling back to scan: {e}"); + Err(ServeError::Internal("fts_failed".into())) + } + } + }) + .await + .map_err(|e| ServeError::Internal(e.to_string()))?; + + match result { + Ok(envelope) => Some(Ok(envelope)), + Err(_) => None, // Fall through to scan + } + } else { + None + }; + + if let Some(Ok(envelope)) = native_result { + return Ok(Json(envelope)); + } + + // Phase 2: Scan fallback — requires main DB lock. + let db = Arc::clone(&state.db); + let result = tokio::task::spawn_blocking(move || { + let mut guard = db.lock().map_err(|e| ServeError::Internal(e.to_string()))?; + + // Try native FTS first using pooled connection. + // On failure: reopen FTS connection and retry once before falling back to scan. + let native_result: Option = { + let first_attempt = guard + .pool() + .and_then(|pool| pool.fts_conn()) + .map(|fts_conn| { + wx_db::native_fts::search_message_fts(fts_conn, &q, limit, offset) + }); + + match first_attempt { + Some(Ok(r)) => Some(r), + Some(Err(e)) => { + // FTS query failed — attempt reopen + retry + eprintln!("warn: native FTS search failed, attempting reopen: {e}"); + match guard.reopen_fts() { + Ok(()) => { + // Reopen succeeded — retry query with fresh connection + guard + .pool() + .and_then(|pool| pool.fts_conn()) + .and_then( + |fts_conn| match wx_db::native_fts::search_message_fts( + fts_conn, &q, limit, offset, + ) { + Ok(r) => Some(r), + Err(e2) => { + eprintln!( + "warn: native FTS search failed after reopen, \ + falling back to scan: {e2}" + ); + None + } + }, + ) + } + Err(reopen_err) => { + eprintln!( + "warn: FTS reopen failed, falling back to scan: {reopen_err}" + ); + None + } + } + } + // No pool or no fts_conn — skip directly to scan fallback + None => None, + } + }; + + let (hits, total_hits, scan_scanned, scan_skipped, shard_warnings): ( + Vec<_>, + usize, + usize, + usize, + Vec, + ) = match native_result { + Some(result) => { + let total = result.total_hits; + let items = result + .hits + .into_iter() + .map(|hit| enrich_native_fts_hit(hit, &self_wxid, &resolver)) + .collect(); + (items, total, 0, 0, Vec::new()) + } + None => { + // Scan fallback: iterate all sessions and search by keyword. + let mut all_sessions = Vec::new(); + let page_size = wx_db::MAX_QUERY_LIMIT; + let mut sess_offset = 0; + loop { + let page = guard + .query_sessions( + &wx_db::SessionQuery::new() + .limit(page_size) + .offset(sess_offset), + ) + .map_err(|e| ServeError::Db(e.to_string()))?; + if page.items.is_empty() { + break; + } + let done = sess_offset + page.items.len() >= page.stats.total_rows; + sess_offset += page.items.len(); + all_sessions.extend(page.items); + if done { + break; + } + } + + let mut all_hits: Vec<(wx_db::Message, String)> = Vec::new(); + let mut total_scanned: usize = 0; + let mut total_skipped: usize = 0; + let mut all_shard_warnings: Vec = Vec::new(); + for session in &all_sessions { + let result = guard + .query_messages( + &wx_db::MessageQuery::for_talker(&session.username) + .keyword(&q) + .limit(wx_db::MAX_QUERY_LIMIT), + ) + .map_err(|e| ServeError::Db(e.to_string()))?; + total_scanned += result.stats.total_rows; + total_skipped += result.stats.skipped; + all_shard_warnings.extend(result.shard_warnings); + for msg in result.items { + all_hits.push((msg, session.username.clone())); + } + } + + // Sort by (sort_seq DESC, create_time DESC, server_id DESC) + all_hits.sort_by(|a, b| { + b.0.sort_seq + .cmp(&a.0.sort_seq) + .then_with(|| b.0.create_time.cmp(&a.0.create_time)) + .then_with(|| b.0.server_id.cmp(&a.0.server_id)) + }); + + let total = all_hits.len(); + let page: Vec<_> = all_hits.into_iter().skip(offset).take(limit).collect(); + let enriched: Vec<_> = page + .into_iter() + .map(|(m, talker)| enrich_message_as_hit(m, talker, &self_wxid, &resolver)) + .collect(); + ( + enriched, + total, + total_scanned, + total_skipped, + all_shard_warnings, + ) + } + }; + + let returned = hits.len(); + let has_more = offset + returned < total_hits; + let envelope = JsonEnvelope { + items: hits, + paging: crate::output::PagingMeta { + limit, + offset, + returned, + has_more, + total: total_hits, + }, + stats: crate::output::StatsMeta { + scanned: scan_scanned, + skipped: scan_skipped, + elapsed_ms: Some(search_start.elapsed().as_millis() as u64), + shard_warnings, + }, + }; + Ok::<_, ServeError>(envelope) + }) + .await + .map_err(|e| ServeError::Internal(e.to_string()))??; + + Ok(Json(result)) +} + +// --------------------------------------------------------------------------- +// Helpers +// --------------------------------------------------------------------------- + +fn resolve_contact( + contact: &str, + resolver: &wx_context::ContactResolver, + db: &wx_db::WechatDb, + visibility: Option<&wx_context::VisibilityIndex>, + show_hidden: bool, +) -> Result { + use crate::contact_id::{resolve_contact_id, ContactResolveError}; + + match resolve_contact_id(contact, resolver, db, visibility, show_hidden) { + Ok(resolved) => Ok(resolved.wxid), + Err(ContactResolveError::NotFound(msg)) | Err(ContactResolveError::Hidden(msg)) => { + Err(ServeError::InvalidParam(msg)) + } + Err(ContactResolveError::Ambiguous(msg)) => Err(ServeError::InvalidParam(msg)), + } +} + +// --------------------------------------------------------------------------- +// SSE +// --------------------------------------------------------------------------- + +pub async fn handler_sse(State(state): State>) -> impl IntoResponse { + if !state.ready.load(Ordering::Acquire) { + return ( + StatusCode::SERVICE_UNAVAILABLE, + Json(json!({"error": "bridge initializing, SSE not ready"})), + ) + .into_response(); + } + + let rx = state.broadcast_tx.subscribe(); + let shutdown = state.shutdown.clone(); + let base = BroadcastStream::new(rx).filter_map(|result| match result { + Ok(event) => { + let json = serde_json::to_string(event.as_ref()).ok()?; + let event_type = match event.as_ref() { + SseEvent::Session(_) => "session", + SseEvent::Message(_) => "message", + SseEvent::Heartbeat => "heartbeat", + }; + Some(Ok::<_, Infallible>( + Event::default().event(event_type).data(json), + )) + } + Err(_) => None, // Lagged — skip + }); + let stream = futures_util::StreamExt::take_until(base, shutdown.cancelled_owned()); + Sse::new(stream) + .keep_alive(KeepAlive::new().interval(Duration::from_secs(30))) + .into_response() +} diff --git a/crates/wx-cli/src/cmd/serve/media.rs b/crates/wx-cli/src/cmd/serve/media.rs new file mode 100644 index 0000000..ed2038a --- /dev/null +++ b/crates/wx-cli/src/cmd/serve/media.rs @@ -0,0 +1,671 @@ +use std::path::{Path, PathBuf}; +use std::sync::Arc; + +use axum::body::Body; +use axum::extract::Request; +use axum::http::header::{CONTENT_DISPOSITION, CONTENT_TYPE}; +use axum::http::HeaderValue; +use axum::response::{IntoResponse, Response}; +use tower::ServiceExt; +use tower_http::services::ServeFile; +use wx_db::{ + open_readonly_connection, Message, MessageContent, MessageQuery, SortOrder, WechatDb, +}; + +use crate::util::{format_month, sanitize_filename}; +use super::error::ServeError; +use super::state::{AppState, CachedVoicePayload}; + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum MediaFormat { + Ogg, + Mp3, +} + +impl MediaFormat { + pub fn parse(value: Option<&str>) -> Result { + match value.unwrap_or("ogg").to_ascii_lowercase().as_str() { + "ogg" => Ok(Self::Ogg), + "mp3" => Ok(Self::Mp3), + other => Err(ServeError::InvalidParam(format!( + "unsupported media format: {other}" + ))), + } + } +} + +pub struct MediaRequest { + pub server_id: i64, + pub talker: String, + pub format: MediaFormat, +} + +enum MediaPayload { + InlineBytes { + bytes: Vec, + content_type: &'static str, + }, + ServePath { + path: PathBuf, + content_type: &'static str, + disposition: Option, + }, +} + +impl MediaPayload { + async fn into_response(self, request: Request) -> Result { + match self { + Self::InlineBytes { + bytes, + content_type, + } => { + let mut response = bytes.into_response(); + response + .headers_mut() + .insert(CONTENT_TYPE, HeaderValue::from_static(content_type)); + Ok(response) + } + Self::ServePath { + path, + content_type, + disposition, + } => { + let mut response = ServeFile::new(&path) + .oneshot(request) + .await + .map_err(|never| match never {})? + .map(Body::new); + response + .headers_mut() + .insert(CONTENT_TYPE, HeaderValue::from_static(content_type)); + if let Some(disposition) = disposition { + let value = HeaderValue::from_str(&disposition).map_err(|e| { + ServeError::Internal(format!( + "invalid content disposition for {}: {e}", + path.display() + )) + })?; + response.headers_mut().insert(CONTENT_DISPOSITION, value); + } + Ok(response) + } + } + } +} + +pub async fn serve_media( + state: Arc, + media_request: MediaRequest, + request: Request, +) -> Result { + let payload = resolve_media(state, media_request).await?; + payload.into_response(request).await +} + +async fn resolve_media( + state: Arc, + media_request: MediaRequest, +) -> Result { + let db = Arc::clone(&state.db); + let attach_dir = state.attach_dir.clone(); + let file_dir = state.file_dir.clone(); + let video_dir = state.video_dir.clone(); + let hardlink_db_path = state.hardlink_db_path.clone(); + let hardlink_db_conn = state.hardlink_db_conn.clone(); + let raw_key = state.raw_key; + let dat_decrypt = state.dat_decrypt.clone(); + let voice_cache = Arc::clone(&state.voice_cache); + let image_xor_cache = Arc::clone(&state.image_xor_cache); + let visibility = Arc::clone(&state.visibility); + let state_for_cache = Arc::clone(&state); + + tokio::task::spawn_blocking(move || { + let message = { + let mut guard = db.lock().map_err(|e| ServeError::Internal(e.to_string()))?; + lookup_message(&mut guard, &media_request.talker, media_request.server_id)? + }; + let canonical_talker = message.talker.clone(); + + if !visibility.allows_media_for_sender(&canonical_talker, &message.sender) { + return Err(ServeError::NotFound(format!( + "message not found for talker={}, server_id={}", + media_request.talker, media_request.server_id + ))); + } + + match message.content { + MessageContent::Image { md5: Some(md5) } => resolve_image( + &canonical_talker, + &md5, + &attach_dir, + dat_decrypt, + &image_xor_cache, + ), + MessageContent::Image { md5: None } => Err(ServeError::NotFound(format!( + "image metadata missing for server_id={}", + media_request.server_id + ))), + MessageContent::Voice => { + let db_paths = { + let mut cache_guard = state_for_cache.media_db_paths.lock() + .map_err(|e: std::sync::PoisonError<_>| ServeError::Internal(e.to_string()))?; + match cache_guard.as_ref() { + Some(paths) => paths.clone(), + None => { + let loaded = find_media_db_paths(&state_for_cache.media_db_dir)?; + *cache_guard = Some(loaded.clone()); + loaded + } + } + }; + resolve_voice( + media_request.server_id, + media_request.format, + &db_paths, + raw_key, + &voice_cache, + ) + } + MessageContent::Video { md5: Some(md5) } => resolve_video( + &md5, + message.create_time, + &attach_dir, + &video_dir, + &hardlink_db_conn, + &hardlink_db_path, + raw_key, + ), + MessageContent::Video { md5: None } => Err(ServeError::NotFound(format!( + "video metadata missing for server_id={}", + media_request.server_id + ))), + MessageContent::File { + md5: Some(md5), + title, + .. + } => resolve_file( + &md5, + title.as_deref(), + message.create_time, + &file_dir, + &hardlink_db_conn, + &hardlink_db_path, + raw_key, + ), + MessageContent::File { md5: None, .. } => Err(ServeError::NotFound(format!( + "file metadata missing for server_id={}", + media_request.server_id + ))), + _ => Err(ServeError::UnsupportedMedia(format!( + "message type {} is not supported by /api/v1/media", + message.msg_type + ))), + } + }) + .await + .map_err(|e| ServeError::Internal(e.to_string()))? +} + +fn lookup_message(db: &mut WechatDb, talker: &str, server_id: i64) -> Result { + let query = MessageQuery::for_talker(talker) + .around_server_id(server_id) + .context(0) + .limit(1) + .order(SortOrder::Desc); + let result = db + .query_messages_anchor(&query) + .map_err(|e| ServeError::Db(e.to_string()))?; + + result + .items + .into_iter() + .find(|message| message.server_id == server_id) + .ok_or_else(|| { + ServeError::NotFound(format!( + "message not found for talker={talker}, server_id={server_id}" + )) + }) +} + +fn resolve_image( + talker: &str, + md5: &str, + attach_dir: &Path, + mut dat_decrypt: wx_media::DatDecryptOptions, + image_xor_cache: &Arc>>>, +) -> Result { + if dat_decrypt.xor_key.is_none() { + dat_decrypt.xor_key = cached_xor_key(talker, attach_dir, image_xor_cache); + } + + let lookup = wx_media::resolve_image_by_md5(talker, attach_dir, md5) + .map_err(|e| map_media_lookup_err(md5, e))?; + let dat_path = lookup.recommended.ok_or_else(|| { + ServeError::NotFound(format!("no candidate image file found for md5={md5}")) + })?; + let data = std::fs::read(&dat_path) + .map_err(|e| ServeError::Internal(format!("failed to read {}: {e}", dat_path.display())))?; + let decoded = wx_media::decrypt_dat(&data, &dat_decrypt) + .map_err(|e| ServeError::Internal(format!("decrypt_dat failed for {md5}: {e}")))?; + + if decoded.ext == "wxgf" { + let transcoded = + wx_media::transcode_wxgf(&decoded.data).map_err(|e| map_wxgf_err(md5, e))?; + if !transcoded.transcoded { + return Err(ServeError::UnsupportedMedia(format!( + "wxgf image for md5={md5} contains HEVC content that cannot be served directly; {}", + wx_media::MediaError::ffmpeg_install_hint() + ))); + } + return Ok(MediaPayload::InlineBytes { + bytes: transcoded.data, + content_type: image_content_type(transcoded.ext), + }); + } + + Ok(MediaPayload::InlineBytes { + bytes: decoded.data, + content_type: image_content_type(&decoded.ext), + }) +} + +fn resolve_voice( + server_id: i64, + format: MediaFormat, + db_paths: &[PathBuf], + raw_key: Option<[u8; 32]>, + voice_cache: &Arc>>, +) -> Result { + let cache_key = format!( + "{}:{}", + server_id, + match format { + MediaFormat::Ogg => "ogg", + MediaFormat::Mp3 => "mp3", + } + ); + if let Ok(mut cache) = voice_cache.lock() { + if let Some(cached) = cache.get(&cache_key) { + return Ok(MediaPayload::InlineBytes { + bytes: cached.bytes.clone(), + content_type: cached.content_type, + }); + } + } + + let svr_id = server_id.to_string(); + let mut first_db_error: Option = None; + + for db_path in db_paths { + let conn = match open_readonly_connection(&db_path, raw_key.as_ref()) { + Ok(conn) => conn, + Err(err) => { + if first_db_error.is_none() { + first_db_error = Some(err.to_string()); + } + continue; + } + }; + + match wx_media::extract_voice_with_conn(&conn, &svr_id) { + Ok(blob) => { + let result = match format { + MediaFormat::Ogg => wx_media::transcode_silk_to_ogg_opus(&blob.data), + MediaFormat::Mp3 => wx_media::transcode_silk_to_mp3(&blob.data), + } + .map_err(|err| map_audio_err(server_id, format, err))?; + + if !result.transcoded { + return Err(voice_transcode_unavailable(server_id, format)); + } + + if let Ok(mut cache) = voice_cache.lock() { + cache.put( + cache_key.clone(), + CachedVoicePayload { + bytes: result.data.clone(), + content_type: result.mime, + }, + ); + } + + return Ok(MediaPayload::InlineBytes { + bytes: result.data, + content_type: result.mime, + }); + } + Err(wx_media::MediaError::LookupMiss(_)) + | Err(wx_media::MediaError::SchemaMissing(_)) => continue, + Err(err) => { + if first_db_error.is_none() { + first_db_error = Some(err.to_string()); + } + } + } + } + + if let Some(err) = first_db_error { + return Err(ServeError::Db(err)); + } + + Err(ServeError::NotFound(format!( + "voice asset not found for server_id={server_id}" + ))) +} + +fn resolve_video( + md5: &str, + create_time: i64, + attach_dir: &Path, + video_dir: &Path, + hardlink_db_conn: &Arc>>, + hardlink_db_path: &Path, + raw_key: Option<[u8; 32]>, +) -> Result { + let mut candidates = Vec::new(); + match query_hardlink_entries(hardlink_db_conn, hardlink_db_path, raw_key, "video", md5) { + Ok(entries) => { + if let Some(entry) = entries.first() { + candidates.push( + attach_dir + .join(&entry.dir1) + .join(&entry.dir2) + .join("Video") + .join(&entry.file_name), + ); + candidates.push( + attach_dir + .join(&entry.dir1) + .join(&entry.dir2) + .join(&entry.file_name), + ); + candidates.push( + attach_dir + .join(&entry.dir1) + .join("Video") + .join(&entry.file_name), + ); + candidates.push( + video_dir + .join(&entry.dir1) + .join(&entry.dir2) + .join(&entry.file_name), + ); + candidates.push(video_dir.join(&entry.dir1).join(&entry.file_name)); + candidates.push(video_dir.join(&entry.file_name)); + + if let Some(path) = candidates.iter().find(|path| path.exists()) { + return Ok(MediaPayload::ServePath { + path: path.clone(), + content_type: video_content_type(path), + disposition: path.file_name().map(|name| { + format!( + "inline; filename=\"{}\"", + sanitize_header_filename(&name.to_string_lossy()) + ) + }), + }); + } + } + } + Err(ServeError::NotFound(_)) => {} + Err(err) => return Err(err), + } + + let month = format_month(create_time); + let fallback = wx_media::find_video_by_md5(video_dir, md5, &month).ok_or_else(|| { + ServeError::NotFound(format!( + "video asset not found for md5={md5}; the video may not be downloaded locally in WeChat yet" + )) + })?; + Ok(MediaPayload::ServePath { + path: fallback.clone(), + content_type: video_content_type(&fallback), + disposition: fallback.file_name().map(|name| { + format!( + "inline; filename=\"{}\"", + sanitize_header_filename(&name.to_string_lossy()) + ) + }), + }) +} + +fn resolve_file( + md5: &str, + title: Option<&str>, + create_time: i64, + file_dir: &Path, + hardlink_db_conn: &Arc>>, + hardlink_db_path: &Path, + raw_key: Option<[u8; 32]>, +) -> Result { + match query_hardlink_entries(hardlink_db_conn, hardlink_db_path, raw_key, "file", md5) { + Ok(entries) => { + if let Some(entry) = entries.first() { + let candidates = [ + file_dir + .join(&entry.dir1) + .join(&entry.dir2) + .join(&entry.file_name), + file_dir.join(&entry.dir1).join(&entry.file_name), + ]; + if let Some(path) = candidates.iter().find(|path| path.exists()) { + return Ok(MediaPayload::ServePath { + path: path.clone(), + content_type: file_content_type(path), + disposition: Some(format!( + "attachment; filename=\"{}\"", + sanitize_header_filename(&entry.file_name) + )), + }); + } + } + } + Err(ServeError::NotFound(_)) => {} + Err(err) => return Err(err), + } + + let title = title.ok_or_else(|| { + ServeError::NotFound(format!( + "file asset not found for md5={md5} (missing title)" + )) + })?; + let month = format_month(create_time); + let fallback = wx_media::find_file_by_name(file_dir, title, &month) + .ok_or_else(|| ServeError::NotFound(format!("file asset not found for md5={md5}")))?; + let file_name = fallback + .file_name() + .map(|name| name.to_string_lossy().to_string()) + .unwrap_or_else(|| sanitize_filename(title)); + Ok(MediaPayload::ServePath { + path: fallback, + content_type: file_content_type(Path::new(&file_name)), + disposition: Some(format!( + "attachment; filename=\"{}\"", + sanitize_header_filename(&file_name) + )), + }) +} + +fn query_hardlink_entries( + hardlink_db_conn: &Arc>>, + hardlink_db_path: &Path, + raw_key: Option<[u8; 32]>, + media_type: &str, + key: &str, +) -> Result, ServeError> { + let mut guard = hardlink_db_conn + .lock() + .map_err(|e| ServeError::Internal(format!("hardlink db lock failed: {e}")))?; + + // Lazily (re)open connection if absent (e.g. after refresh cleared it). + if guard.is_none() && hardlink_db_path.exists() { + match open_readonly_connection(hardlink_db_path, raw_key.as_ref()) { + Ok(conn) => { + eprintln!("server/hardlink: reopened pooled connection"); + *guard = Some(conn); + } + Err(e) => { + eprintln!("warn: cannot reopen hardlink.db: {e}"); + } + } + } + + let result = if let Some(ref conn) = *guard { + wx_media::query_hardlink_with_conn(conn, media_type, key) + .map_err(|e| map_media_lookup_err(key, e)) + } else { + // Fallback: open a new connection (file may not exist or open failed) + let conn = open_readonly_connection(hardlink_db_path, raw_key.as_ref()) + .map_err(|e| ServeError::Db(e.to_string()))?; + wx_media::query_hardlink_with_conn(&conn, media_type, key) + .map_err(|e| map_media_lookup_err(key, e)) + }; + result +} + +fn find_media_db_paths(media_db_dir: &Path) -> Result, ServeError> { + let entries = std::fs::read_dir(media_db_dir).map_err(|_| { + ServeError::NotFound(format!( + "media database directory not found: {}", + media_db_dir.display() + )) + })?; + + let mut paths: Vec = entries + .filter_map(|entry| entry.ok()) + .map(|entry| entry.path()) + .filter(|path| { + path.file_name() + .and_then(|name| name.to_str()) + .is_some_and(|name| { + (name == "media.db" || name.starts_with("media_")) && name.ends_with(".db") + }) + }) + .collect(); + paths.sort(); + + if paths.is_empty() { + return Err(ServeError::NotFound(format!( + "no media databases found under {}", + media_db_dir.display() + ))); + } + + Ok(paths) +} + +fn image_content_type(ext: &str) -> &'static str { + match ext { + "jpg" | "jpeg" => "image/jpeg", + "png" => "image/png", + "gif" => "image/gif", + "bmp" => "image/bmp", + "webp" => "image/webp", + "tif" | "tiff" => "image/tiff", + _ => "application/octet-stream", + } +} + +fn video_content_type(path: &Path) -> &'static str { + match path + .extension() + .and_then(|ext| ext.to_str()) + .unwrap_or_default() + { + "mp4" => "video/mp4", + "mov" => "video/quicktime", + _ => "application/octet-stream", + } +} + +fn file_content_type(path: &Path) -> &'static str { + match path + .extension() + .and_then(|ext| ext.to_str()) + .unwrap_or_default() + { + "txt" => "text/plain; charset=utf-8", + "pdf" => "application/pdf", + _ => "application/octet-stream", + } +} + +fn map_media_lookup_err(key: &str, err: wx_media::MediaError) -> ServeError { + match err { + wx_media::MediaError::NotFound(_) + | wx_media::MediaError::LookupMiss(_) + | wx_media::MediaError::NoDatFiles { .. } + | wx_media::MediaError::NoMediaDbs(_) => { + ServeError::NotFound(format!("media asset not found for key={key}")) + } + other => ServeError::Internal(other.to_string()), + } +} + +fn map_audio_err(server_id: i64, format: MediaFormat, err: wx_media::MediaError) -> ServeError { + match err { + wx_media::MediaError::FfmpegNotFound => voice_transcode_unavailable(server_id, format), + wx_media::MediaError::AudioFeatureDisabled => ServeError::Internal(format!( + "voice transcoding for server_id={server_id} is unavailable because wx-cli was built without the 'audio' feature" + )), + wx_media::MediaError::FfmpegFailed { .. } + | wx_media::MediaError::SilkDecodeFailed { .. } => ServeError::Upstream(format!( + "voice transcode failed for server_id={server_id}: {err}" + )), + other => ServeError::Internal(other.to_string()), + } +} + +fn map_wxgf_err(md5: &str, err: wx_media::MediaError) -> ServeError { + match err { + wx_media::MediaError::InvalidWxgf => ServeError::UnsupportedMedia(format!( + "wxgf image for md5={md5} is invalid or unreadable" + )), + wx_media::MediaError::FfmpegNotFound => ServeError::UnsupportedMedia(format!( + "wxgf image for md5={md5} requires ffmpeg for HEVC decode; {}", + wx_media::MediaError::ffmpeg_install_hint() + )), + wx_media::MediaError::FfmpegFailed { .. } => { + ServeError::Upstream(format!("wxgf transcode failed for {md5}: {err}")) + } + other => ServeError::Internal(format!("wxgf transcode failed for {md5}: {other}")), + } +} + +fn voice_transcode_unavailable(server_id: i64, format: MediaFormat) -> ServeError { + let format_name = match format { + MediaFormat::Ogg => "ogg", + MediaFormat::Mp3 => "mp3", + }; + ServeError::UnsupportedMedia(format!( + "voice media for server_id={server_id} requires ffmpeg to produce {format_name}; {}", + wx_media::MediaError::ffmpeg_install_hint() + )) +} + +fn sanitize_header_filename(name: &str) -> String { + sanitize_filename(name).replace('"', "_") +} + +fn cached_xor_key( + talker: &str, + attach_dir: &Path, + image_xor_cache: &Arc>>>, +) -> Option { + if let Ok(mut cache) = image_xor_cache.lock() { + if let Some(value) = cache.get(talker) { + return *value; + } + } + + let username_hash = format!("{:x}", wx_media::md5_hash(talker.as_bytes())); + let talker_attach = attach_dir.join(username_hash); + let detected = wx_media::detect_xor_key(&talker_attach); + + if let Ok(mut cache) = image_xor_cache.lock() { + cache.put(talker.to_string(), detected); + } + + detected +} diff --git a/crates/wx-cli/src/cmd/serve/mod.rs b/crates/wx-cli/src/cmd/serve/mod.rs new file mode 100644 index 0000000..e812cdc --- /dev/null +++ b/crates/wx-cli/src/cmd/serve/mod.rs @@ -0,0 +1,506 @@ +mod auth; +mod bridge; +mod error; +mod event; +mod handlers; +mod media; +pub(crate) mod refresh; +mod routes; +mod state; + +use std::num::NonZeroUsize; +use std::path::PathBuf; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::sync::Arc; +use std::time::{Duration, Instant}; + +use lru::LruCache; +use tokio::signal::unix::SignalKind; +use tokio::sync::{broadcast, mpsc, watch}; +use tokio_util::sync::CancellationToken; +use wx_context::{ + open_fts_connection_with_key, register_mm_fts_tokenizer, write_shard_metadata_sidecar, + AccountContext, ContactResolver, DecryptRequest, PersistentCache, ResolveParams, +}; + +use crate::util::{print_cache_stats, print_detection_note}; +use crate::version; + +use self::refresh::{RefreshTask, RefreshTrigger}; +use self::state::{AppState, CurrentAccount}; +use super::contacts::build_visibility; +use super::server::runtime::{base_url, RuntimeReporter}; +use super::server::types::{RuntimeAccountState, ServerRuntimeState, WorkerLifecycle}; + +pub(crate) fn is_loopback(host: &str) -> bool { + matches!(host, "127.0.0.1" | "::1" | "localhost") +} + +fn resolve_watch_mode(poll: bool, fsnotify: bool) -> wx_monitor::WatchMode { + if poll { + wx_monitor::WatchMode::Poll + } else if fsnotify { + wx_monitor::WatchMode::Fsnotify + } else { + wx_monitor::WatchMode::Auto + } +} + +#[allow(clippy::too_many_arguments)] +pub async fn cmd_serve( + key_hex: Option, + data_dir: Option, + account: Option, + poll: bool, + fsnotify: bool, + poll_ms: u64, + host: String, + port: u16, + token: Option, + worker_id: Option, + runtime_reporter: Option, +) -> Result<(), Box> { + // Security check: require token when binding to non-loopback + if !is_loopback(&host) && token.is_none() { + return Err("--token is required when --host is not loopback".into()); + } + + let t_total = Instant::now(); + let params = &wx_decrypt::MACOS_4_1_7_31; + + // 1. Account resolution + cache + let t = Instant::now(); + let acct = AccountContext::resolve(&ResolveParams { + account: account.as_deref(), + data_dir: data_dir.as_deref(), + key_hex: key_hex.as_deref(), + })?; + print_detection_note(&acct); + eprintln!( + "server/timing: account_resolve {:.0}ms", + t.elapsed().as_secs_f64() * 1000.0 + ); + + // 2 & 3. Open DB — direct encrypted or decrypt+cache depending on raw_key + let direct_mode = acct.raw_key.is_some(); + let (db, cache): (wx_db::WechatDb, Option) = if direct_mode { + let t = Instant::now(); + eprintln!("Direct encrypted open with pool (SQLCipher)"); + let db = wx_context::open_encrypted_db_with_pool(&acct)?; + eprintln!( + "server/timing: db_open_encrypted_with_pool {:.0}ms", + t.elapsed().as_secs_f64() * 1000.0 + ); + (db, None) + } else { + let t = Instant::now(); + let cache = PersistentCache::new(&acct, params)?; + let stats = DecryptRequest::new() + .all() + .execute_with_progress(&cache, crate::util::decrypt_progress_callback)?; + print_cache_stats(&stats); + eprintln!( + "server/timing: decrypt {:.0}ms", + t.elapsed().as_secs_f64() * 1000.0 + ); + + let t = Instant::now(); + let db = + wx_db::WechatDb::open_with_pool(cache.decrypted_root(), register_mm_fts_tokenizer)?; + eprintln!( + "server/timing: db_open_with_pool {:.0}ms", + t.elapsed().as_secs_f64() * 1000.0 + ); + + // Write shard metadata sidecar for future routing + if let Err(e) = write_shard_metadata_sidecar(&db, cache.decrypted_root()) { + eprintln!("warn: failed to write shard metadata sidecar: {e}"); + } + + (db, Some(cache)) + }; + + // 3b. Open independent FTS connection (outside WechatDb Mutex) + let fts_conn = db.message_fts_path.as_deref().and_then(|fts_path| { + match open_fts_connection_with_key(fts_path, acct.raw_key.as_ref()) { + Ok(conn) => { + if let Ok(mode) = + conn.query_row("PRAGMA journal_mode", [], |r| r.get::<_, String>(0)) + { + eprintln!("server/fts: journal_mode={mode}"); + } + Some(Arc::new(std::sync::Mutex::new(conn))) + } + Err(e) => { + eprintln!("warn: cannot open independent FTS connection: {e}"); + None + } + } + }); + + // 4. Build contact resolver + let t = Instant::now(); + let resolver = ContactResolver::build(&db)?; + let visibility = build_visibility(&acct, &resolver); + eprintln!( + "server/timing: resolver_build {:.0}ms", + t.elapsed().as_secs_f64() * 1000.0 + ); + + let self_wxid = acct.base_wxid.clone(); + let current_account = CurrentAccount { + wxid: self_wxid.clone(), + name: resolver.display_name(&self_wxid).to_string(), + }; + let worker_id = worker_id.unwrap_or_else(|| { + format!( + "worker-{}-{}", + std::process::id(), + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .map(|d| d.as_nanos()) + .unwrap_or_default() + ) + }); + let attach_dir = acct.data_dir.join("msg").join("attach"); + let file_dir = acct.data_dir.join("msg").join("file"); + let video_dir = acct.data_dir.join("msg").join("video"); + let dat_decrypt = wx_media::DatDecryptOptions { + v2_aes_key: wx_media::derive_v2_key_from_dir(&acct.data_dir).ok(), + xor_key: None, + }; + let media_db_dir = if direct_mode { + acct.data_dir.join("db_storage").join("message") + } else { + cache + .as_ref() + .map(|c| c.decrypted_root().join("message")) + .unwrap_or_else(|| acct.data_dir.join("db_storage").join("message")) + }; + let hardlink_db_path = if direct_mode { + acct.data_dir + .join("db_storage") + .join("hardlink") + .join("hardlink.db") + } else { + cache + .as_ref() + .map(|c| c.decrypted_root().join("hardlink").join("hardlink.db")) + .unwrap_or_else(|| { + acct.data_dir + .join("db_storage") + .join("hardlink") + .join("hardlink.db") + }) + }; + + // 3c. Open hardlink.db connection (pooled, outside WechatDb Mutex) + let hardlink_db_conn = if hardlink_db_path.exists() { + match wx_db::open_readonly_connection(&hardlink_db_path, acct.raw_key.as_ref()) { + Ok(conn) => { + eprintln!("server/hardlink: opened pooled connection"); + Some(conn) + } + Err(e) => { + eprintln!("warn: cannot open hardlink.db connection: {e}"); + None + } + } + } else { + None + }; + let hardlink_db_conn = Arc::new(std::sync::Mutex::new(hardlink_db_conn)); + + // Check encrypted session dir before constructing monitor config + let encrypted_session_dir = acct.data_dir.join("db_storage").join("session"); + if !encrypted_session_dir.exists() { + return Err(format!( + "session directory not found: {}", + encrypted_session_dir.display() + ) + .into()); + } + + let watch_mode = resolve_watch_mode(poll, fsnotify); + let config = wx_monitor::MonitorConfig { + encrypted_session_dir, + key_material: acct.key_material.clone(), + params, + watch_mode: watch_mode.clone(), + poll_interval: Duration::from_millis(poll_ms), + channel_capacity: 1000, + raw_key: acct.raw_key, + encrypted_root: if acct.raw_key.is_some() { + Some(acct.data_dir.join("db_storage")) + } else { + None + }, + }; + + // Capture values before moving db into Mutex + let fts_path_for_refresh = db.message_fts_path.clone(); + let raw_key_for_refresh = acct.raw_key; + + // 5. Create refresh task channels + let (refresh_tx, refresh_rx) = mpsc::channel::(64); + let (epoch_tx, _epoch_rx) = watch::channel(0u64); + + // 6. Build AppState + let shutdown = CancellationToken::new(); + let (broadcast_tx, _) = broadcast::channel(512); + let db_arc = Arc::new(std::sync::Mutex::new(db)); + let cache_arc = cache.map(Arc::new); + + let app_state = Arc::new(AppState { + db: Arc::clone(&db_arc), + self_wxid, + current_account, + worker_id: worker_id.clone(), + cli_version: version::cli_version_string(), + resolver: Arc::new(resolver), + visibility: Arc::new(visibility), + broadcast_tx, + auth_token: token.clone(), + ready: AtomicBool::new(false), + refresh_tx: refresh_tx.clone(), + shutdown: shutdown.clone(), + fts_conn, + attach_dir, + media_db_dir, + file_dir, + video_dir, + hardlink_db_path, + hardlink_db_conn, + raw_key: acct.raw_key, + dat_decrypt, + voice_cache: Arc::new(std::sync::Mutex::new(LruCache::new(NonZeroUsize::new(256).unwrap()))), + image_xor_cache: Arc::new(std::sync::Mutex::new(LruCache::new(NonZeroUsize::new(1024).unwrap()))), + name2id_cache: Arc::new(std::sync::Mutex::new(None)), + media_db_paths: Arc::new(std::sync::Mutex::new(None)), + }); + + // 7. Bind TCP listener BEFORE background init (port available immediately) + let router = routes::build_router(Arc::clone(&app_state)); + + let bind_addr = if host.contains(':') { + format!("[{host}]:{port}") + } else { + format!("{host}:{port}") + }; + let t = Instant::now(); + let listener = tokio::net::TcpListener::bind(&bind_addr).await?; + eprintln!( + "server/timing: tcp_bind {:.0}ms", + t.elapsed().as_secs_f64() * 1000.0 + ); + + let auth_status = if token.is_some() { + "Bearer token required" + } else { + "disabled" + }; + eprintln!("wx-cli server worker: listening on http://{bind_addr}"); + eprintln!(" SSE endpoint: GET /api/v1/events"); + eprintln!(" REST endpoints:"); + eprintln!(" GET /api/v1/health"); + eprintln!(" GET /api/v1/sessions"); + eprintln!(" GET /api/v1/contacts"); + eprintln!(" GET /api/v1/messages?contact="); + eprintln!(" GET /api/v1/media?server_id=&talker=[&format=ogg|mp3]"); + eprintln!(" GET /api/v1/search?q="); + eprintln!(" Auth: {auth_status}"); + let resolved = wx_monitor::resolve_watch_mode(&watch_mode); + eprintln!(" Monitor: mode={watch_mode:?} -> {resolved:?}, interval={poll_ms}ms"); + + if let Some(reporter) = &runtime_reporter { + reporter.write_state(ServerRuntimeState { + pid: std::process::id(), + worker_id: worker_id.clone(), + lifecycle: WorkerLifecycle::Starting, + ready: false, + host: host.clone(), + port, + base_url: base_url(&host, port), + token_configured: token.is_some(), + cli_version: app_state.cli_version.clone(), + current_account: Some(RuntimeAccountState { + wxid: app_state.current_account.wxid.clone(), + name: app_state.current_account.name.clone(), + }), + stdout_log: reporter.ap().server_stdout_log(), + stderr_log: reporter.ap().server_stderr_log(), + })?; + } + + // 8. Background task: init baselines → start monitor → spawn refresh → spawn bridge → set ready + // The background task retains monitor ownership and is responsible for cleanup on shutdown. + let shutdown_bg = shutdown.clone(); + let bg_state = Arc::clone(&app_state); + let bg_t_total = t_total; + let runtime_reporter_bg = runtime_reporter.clone(); + let bg_host = host.clone(); + let bg_handle = tokio::spawn(async move { + // init_baselines in spawn_blocking (holds db lock briefly) + let db = Arc::clone(&bg_state.db); + let baselines = tokio::task::spawn_blocking(move || { + let t = Instant::now(); + let guard = match db.lock() { + Ok(g) => g, + Err(e) => { + return Err(format!("db lock failed: {e}")); + } + }; + let result = bridge::init_baselines(&guard); + eprintln!( + "server/timing: init_baselines {:.0}ms", + t.elapsed().as_secs_f64() * 1000.0 + ); + result + }) + .await; + + let (bridge_cursors, startup_watermark) = match baselines { + Ok(Ok(r)) => r, + Ok(Err(e)) => { + eprintln!("error: bridge init_baselines: {e}, SSE will remain unavailable"); + return; + } + Err(e) => { + eprintln!( + "error: bridge init_baselines panicked: {e}, SSE will remain unavailable" + ); + return; + } + }; + + // Start monitor + let t = Instant::now(); + let mut monitor = match wx_monitor::WechatMonitor::start(config) { + Ok(m) => m, + Err(e) => { + eprintln!("error: monitor start failed: {e}, SSE will remain unavailable"); + return; + } + }; + eprintln!( + "server/timing: monitor_start {:.0}ms", + t.elapsed().as_secs_f64() * 1000.0 + ); + + let receiver = monitor.take_receiver().expect("receiver already taken"); + + // Spawn refresh task BEFORE bridge (so refresh loop is running when first events arrive) + let refresh_task = RefreshTask::new( + refresh_rx, + epoch_tx.clone(), + Arc::clone(&db_arc), + cache_arc.clone(), + shutdown_bg.clone(), + ) + .with_fts(bg_state.fts_conn.clone(), fts_path_for_refresh) + .with_raw_key(raw_key_for_refresh) + .with_caches( + Some(Arc::clone(&bg_state.name2id_cache)), + Some(Arc::clone(&bg_state.media_db_paths)), + Some(Arc::clone(&bg_state.hardlink_db_conn)), + ); + let refresh_handle = tokio::spawn(refresh_task.run()); + + // Spawn bridge with refresh channels + let bridge_refresh_watch = epoch_tx.subscribe(); + let bridge_handle = tokio::spawn(bridge::run_bridge( + receiver, + Arc::clone(&bg_state), + bridge_cursors, + startup_watermark, + refresh_tx, + bridge_refresh_watch, + shutdown_bg.clone(), + )); + + // Mark ready + bg_state.ready.store(true, Ordering::Release); + if let Some(reporter) = &runtime_reporter_bg { + let _ = reporter.write_state(ServerRuntimeState { + pid: std::process::id(), + worker_id: bg_state.worker_id.clone(), + lifecycle: WorkerLifecycle::Running, + ready: true, + host: bg_host.clone(), + port, + base_url: base_url(&bg_host, port), + token_configured: bg_state.auth_token.is_some(), + cli_version: bg_state.cli_version.clone(), + current_account: Some(RuntimeAccountState { + wxid: bg_state.current_account.wxid.clone(), + name: bg_state.current_account.name.clone(), + }), + stdout_log: reporter.ap().server_stdout_log(), + stderr_log: reporter.ap().server_stderr_log(), + }); + } + eprintln!( + "server/timing: TOTAL startup {:.0}ms", + bg_t_total.elapsed().as_secs_f64() * 1000.0 + ); + + // Wait for shutdown signal, then clean up + shutdown_bg.cancelled().await; + if let Some(reporter) = &runtime_reporter_bg { + let _ = reporter.write_state(ServerRuntimeState { + pid: std::process::id(), + worker_id: bg_state.worker_id.clone(), + lifecycle: WorkerLifecycle::Stopping, + ready: false, + host: bg_host.clone(), + port, + base_url: base_url(&bg_host, port), + token_configured: bg_state.auth_token.is_some(), + cli_version: bg_state.cli_version.clone(), + current_account: Some(RuntimeAccountState { + wxid: bg_state.current_account.wxid.clone(), + name: bg_state.current_account.name.clone(), + }), + stdout_log: reporter.ap().server_stdout_log(), + stderr_log: reporter.ap().server_stderr_log(), + }); + } + monitor.stop(); + let _ = bridge_handle.await; + let _ = refresh_handle.await; + }); + + // Signal handler: cancel the shutdown token on SIGTERM or SIGINT + let shutdown_signal = shutdown.clone(); + axum::serve(listener, router) + .with_graceful_shutdown(async move { + let mut sigterm = tokio::signal::unix::signal(SignalKind::terminate()) + .expect("failed to register SIGTERM handler"); + let mut sigint = tokio::signal::unix::signal(SignalKind::interrupt()) + .expect("failed to register SIGINT handler"); + tokio::select! { + _ = sigterm.recv() => {}, + _ = sigint.recv() => {}, + } + eprintln!("\nShutting down..."); + shutdown_signal.cancel(); + }) + .await?; + + // Wait for background supervisor to finish cleanup (monitor.stop, join bridge/refresh). + // On timeout, hard-exit to avoid Runtime::drop blocking on in-flight spawn_blocking tasks. + match tokio::time::timeout(Duration::from_secs(5), bg_handle).await { + Ok(Ok(())) => eprintln!("Server stopped."), + Ok(Err(e)) => eprintln!("warn: background task panicked: {e}"), + Err(_) => { + eprintln!("warn: shutdown timed out after 5s, exiting anyway"); + let _ = std::io::Write::flush(&mut std::io::stderr()); + std::process::exit(1); + } + } + if let Some(reporter) = runtime_reporter { + let _ = reporter.clear_state(); + } + let _ = std::io::Write::flush(&mut std::io::stderr()); + Ok(()) +} diff --git a/crates/wx-cli/src/cmd/serve/refresh.rs b/crates/wx-cli/src/cmd/serve/refresh.rs new file mode 100644 index 0000000..5acac90 --- /dev/null +++ b/crates/wx-cli/src/cmd/serve/refresh.rs @@ -0,0 +1,516 @@ +use std::collections::HashMap; +use std::path::PathBuf; +use std::sync::Arc; + +use rusqlite::Connection; +use tokio::sync::{mpsc, watch}; +use tokio_util::sync::CancellationToken; +use wx_context::{ + open_fts_connection, open_fts_connection_with_key, DecryptProgress, DecryptRequest, + PersistentCache, +}; +use wx_db::WechatDb; + +/// Signal sent to the refresh task. +pub enum RefreshTrigger { + Refresh, + #[allow(dead_code)] + Shutdown, +} + +/// Background task that runs DecryptRequest when triggered, then refreshes +/// pool connections. Multiple triggers are coalesced: if N signals queue up +/// while a refresh is in progress, only one subsequent refresh runs. +/// +/// Epoch is only advanced on successful refresh (decrypt + reopen all succeed). +/// Failed refreshes are logged but do NOT advance the epoch, so bridge waiters +/// will not proceed with stale data. +pub struct RefreshTask { + trigger_rx: mpsc::Receiver, + epoch_tx: watch::Sender, + db: Arc>, + cache: Option>, + shutdown: CancellationToken, + /// Independent FTS connection to reopen on refresh. + fts_conn: Option>>, + /// Path to FTS DB for reopening. + fts_path: Option, + /// Raw key for encrypted FTS reopen. + raw_key: Option<[u8; 32]>, + /// Cache of name2id mapping — cleared when FTS is reopened. + name2id_cache: Option>>>>, + /// Cache of media DB paths — cleared on every refresh. + media_db_paths: Option>>>>, + /// Cached hardlink.db connection — cleared on refresh so it is reopened lazily. + hardlink_db_conn: Option>>>, +} + +impl RefreshTask { + pub fn new( + trigger_rx: mpsc::Receiver, + epoch_tx: watch::Sender, + db: Arc>, + cache: Option>, + shutdown: CancellationToken, + ) -> Self { + RefreshTask { + trigger_rx, + epoch_tx, + db, + cache, + shutdown, + fts_conn: None, + fts_path: None, + raw_key: None, + name2id_cache: None, + media_db_paths: None, + hardlink_db_conn: None, + } + } + + pub fn with_raw_key(mut self, raw_key: Option<[u8; 32]>) -> Self { + self.raw_key = raw_key; + self + } + + /// Set the independent FTS connection and path for refresh reopening. + pub fn with_fts( + mut self, + fts_conn: Option>>, + fts_path: Option, + ) -> Self { + self.fts_conn = fts_conn; + self.fts_path = fts_path; + self + } + + /// Set the caches that should be invalidated on refresh. + pub fn with_caches( + mut self, + name2id_cache: Option>>>>, + media_db_paths: Option>>>>, + hardlink_db_conn: Option>>>, + ) -> Self { + self.name2id_cache = name2id_cache; + self.media_db_paths = media_db_paths; + self.hardlink_db_conn = hardlink_db_conn; + self + } + + pub async fn run(mut self) { + let mut epoch: u64 = 0; + + loop { + // Wait for next trigger or shutdown + let trigger = tokio::select! { + t = self.trigger_rx.recv() => t, + _ = self.shutdown.cancelled() => break, + }; + match trigger { + Some(RefreshTrigger::Refresh) => {} + Some(RefreshTrigger::Shutdown) | None => break, + } + + // Drain/coalesce any queued Refresh signals + loop { + match self.trigger_rx.try_recv() { + Ok(RefreshTrigger::Refresh) => continue, + Ok(RefreshTrigger::Shutdown) => { + // Shutdown takes priority — exit immediately + return; + } + Err(_) => break, + } + } + + // Run refresh in spawn_blocking. + let db = Arc::clone(&self.db); + let cache = self.cache.clone(); + let fts_conn = self.fts_conn.clone(); + let fts_path = self.fts_path.clone(); + let raw_key = self.raw_key; + let success = tokio::task::spawn_blocking(move || { + if let Some(cache) = cache { + // Decrypt-cache mode: decrypt then selective reopen + let modified_paths: Arc>> = + Arc::new(std::sync::Mutex::new(Vec::new())); + let modified_clone = Arc::clone(&modified_paths); + + let progress_cb = move |event: DecryptProgress| { + match &event { + DecryptProgress::Decrypted { path, .. } => { + modified_clone.lock().unwrap().push(path.clone()); + } + DecryptProgress::Skipped { + path, + wal_patched: true, + } => { + modified_clone.lock().unwrap().push(path.clone()); + } + _ => {} + } + crate::util::decrypt_progress_callback(event); + }; + + if let Err(e) = DecryptRequest::new() + .all() + .execute_with_progress(&cache, progress_cb) + { + eprintln!("warn: refresh decrypt failed: {e}"); + return false; + } + + let modified = modified_paths.lock().unwrap(); + let decrypted_root = cache.decrypted_root(); + + let mut guard = match db.lock() { + Ok(g) => g, + Err(e) => { + eprintln!("warn: refresh db lock failed: {e}"); + return false; + } + }; + + if let Err(e) = guard.reopen_sessions() { + eprintln!("warn: refresh reopen_sessions failed: {e}"); + return false; + } + + let mut contact_reopened = false; + let mut shards_reopened = 0u32; + let mut fts_changed = false; + + for rel_path in modified.iter() { + if rel_path.ends_with("contact/contact.db") && !contact_reopened { + if let Err(e) = guard.reopen_contacts() { + eprintln!("warn: refresh reopen_contacts failed: {e}"); + return false; + } + contact_reopened = true; + } else if is_message_shard_path(rel_path) { + let abs_path = decrypted_root.join(rel_path); + match guard.reopen_pooled_shard(&abs_path) { + Ok(true) => shards_reopened += 1, + Ok(false) => { + eprintln!( + "warn: unknown shard path (topology change?): {rel_path} — \ + restart server to pick up new shards" + ); + } + Err(e) => { + eprintln!("warn: refresh reopen_pooled_shard failed for {rel_path}: {e}"); + return false; + } + } + } else if rel_path.ends_with("message/message_fts.db") { + fts_changed = true; + } + } + + if fts_changed { + if let Err(e) = guard.reopen_fts() { + eprintln!("warn: refresh reopen_fts failed: {e}"); + return false; + } + if let (Some(fts_mutex), Some(path)) = (&fts_conn, &fts_path) { + match open_fts_connection(path) { + Ok(new_conn) => { + if let Ok(mut fts_guard) = fts_mutex.lock() + as Result, _> + { + *fts_guard = new_conn; + } + } + Err(e) => { + eprintln!("warn: refresh reopen FTS connection failed: {e}"); + } + } + } + } + + if !modified.is_empty() { + eprintln!( + "info: refresh: {} modified file(s), {} shard(s) reopened{}{}", + modified.len(), + shards_reopened, + if contact_reopened { ", contact reopened" } else { "" }, + if fts_changed { ", FTS reopened" } else { "" }, + ); + } + + true + } else { + // Direct encrypted mode: just reopen all connections + let mut guard = match db.lock() { + Ok(g) => g, + Err(e) => { + eprintln!("warn: refresh db lock failed: {e}"); + return false; + } + }; + + if let Err(e) = guard.reopen_sessions() { + eprintln!("warn: refresh reopen_sessions failed: {e}"); + return false; + } + if let Err(e) = guard.reopen_contacts() { + eprintln!("warn: refresh reopen_contacts failed: {e}"); + return false; + } + if let Err(e) = guard.reopen_all_pooled() { + eprintln!("warn: refresh reopen_all_pooled failed: {e}"); + return false; + } + + // Reopen independent FTS connection + if let (Some(fts_mutex), Some(path)) = (&fts_conn, &fts_path) { + match open_fts_connection_with_key(path, raw_key.as_ref()) { + Ok(new_conn) => { + if let Ok(mut fts_guard) = fts_mutex.lock() + as Result, _> + { + *fts_guard = new_conn; + } + } + Err(e) => { + eprintln!("warn: refresh reopen FTS connection failed: {e}"); + } + } + } + + eprintln!("info: refresh (direct mode): all connections reopened"); + true + } + }) + .await; + + let ok = match success { + Ok(v) => v, + Err(e) => { + eprintln!("warn: refresh task panicked: {e}"); + false + } + }; + + // Only advance epoch on successful refresh + if ok { + epoch += 1; + let _ = self.epoch_tx.send(epoch); + + // Invalidate caches that depend on reopened connections. + if let Some(cache) = &self.name2id_cache { + *cache.lock().unwrap() = None; + } + if let Some(cache) = &self.media_db_paths { + *cache.lock().unwrap() = None; + } + if let Some(cache) = &self.hardlink_db_conn { + // Drop the old connection; next media query will reopen lazily. + *cache.lock().unwrap() = None; + } + } + } + } +} + +/// Check if a relative path looks like a numbered message shard (e.g. `message/message_N.db`). +fn is_message_shard_path(rel_path: &str) -> bool { + let Some(filename) = rel_path.rsplit('/').next() else { + return false; + }; + if !rel_path.contains("message/") { + return false; + } + let Some(stem) = filename.strip_suffix(".db") else { + return false; + }; + let Some(suffix) = stem.strip_prefix("message_") else { + return false; + }; + !suffix.is_empty() && suffix.bytes().all(|b| b.is_ascii_digit()) +} + +#[cfg(test)] +mod tests { + use super::*; + + /// Helper: create a RefreshTask that will fail on decrypt (no real cache/db). + /// Used to test channel behavior without needing real WeChat infrastructure. + fn make_test_task() -> ( + mpsc::Sender, + watch::Receiver, + CancellationToken, + RefreshTask, + ) { + let (trigger_tx, trigger_rx) = mpsc::channel(64); + let (epoch_tx, epoch_rx) = watch::channel(0u64); + + // Create a minimal WechatDb fixture that will make decrypt fail + // (no PersistentCache/AccountContext). We use a temp dir with + // bare minimum structure so WechatDb::open succeeds. + let dir = tempfile::TempDir::new().unwrap(); + let base = dir.path(); + std::fs::create_dir_all(base.join("contact")).unwrap(); + rusqlite::Connection::open(base.join("contact/contact.db")) + .unwrap() + .execute_batch( + "CREATE TABLE contact (username TEXT PRIMARY KEY, alias TEXT, remark TEXT, nick_name TEXT, description TEXT, extra_buffer BLOB);", + ) + .unwrap(); + std::fs::create_dir_all(base.join("session")).unwrap(); + rusqlite::Connection::open(base.join("session/session.db")) + .unwrap() + .execute_batch( + "CREATE TABLE SessionTable (username TEXT, sort_timestamp INTEGER, summary TEXT);", + ) + .unwrap(); + std::fs::create_dir_all(base.join("message")).unwrap(); + let db = wx_db::WechatDb::open(base).unwrap(); + let db_arc = Arc::new(std::sync::Mutex::new(db)); + + // PersistentCache requires AccountContext — we can't easily create one. + // Instead, create a cache pointing to the temp dir with valid structure. + let params = &wx_decrypt::MACOS_4_1_7_31; + let acct = wx_context::AccountContext { + account_id: "test".to_string(), + base_wxid: "wxid_test".to_string(), + data_dir: base.to_path_buf(), + key_material: wx_decrypt::KeyMaterial::RawKey([0u8; 32]), + raw_key: Some([0u8; 32]), + writeback_enabled: false, + detection_note: None, + }; + // PersistentCache::new needs encrypted dirs to exist + std::fs::create_dir_all(base.join("db_storage/contact")).unwrap(); + std::fs::create_dir_all(base.join("db_storage/session")).unwrap(); + std::fs::create_dir_all(base.join("db_storage/message")).unwrap(); + let cache = wx_context::PersistentCache::new(&acct, params).unwrap(); + let cache_arc = Some(Arc::new(cache)); + + // Leak the TempDir to keep files alive for the duration of the test + std::mem::forget(dir); + + let shutdown = CancellationToken::new(); + let task = RefreshTask::new(trigger_rx, epoch_tx, db_arc, cache_arc, shutdown.clone()); + (trigger_tx, epoch_rx, shutdown, task) + } + + #[tokio::test] + async fn shutdown_signal_exits_task() { + let (tx, _rx, _shutdown, task) = make_test_task(); + let handle = tokio::spawn(task.run()); + + tx.send(RefreshTrigger::Shutdown).await.unwrap(); + // Task should exit promptly + tokio::time::timeout(std::time::Duration::from_secs(2), handle) + .await + .expect("task should exit on Shutdown") + .expect("task should not panic"); + } + + #[tokio::test] + async fn channel_close_exits_task() { + let (tx, _rx, _shutdown, task) = make_test_task(); + let handle = tokio::spawn(task.run()); + + drop(tx); + // Task should exit when channel is closed + tokio::time::timeout(std::time::Duration::from_secs(2), handle) + .await + .expect("task should exit on channel close") + .expect("task should not panic"); + } + + #[tokio::test] + async fn refresh_failure_does_not_advance_epoch() { + let (tx, rx, _shutdown, task) = make_test_task(); + let handle = tokio::spawn(task.run()); + + // Send a Refresh signal — decrypt will fail (no encrypted .db files) + tx.send(RefreshTrigger::Refresh).await.unwrap(); + + // Give refresh task time to process + tokio::time::sleep(std::time::Duration::from_millis(500)).await; + + // Epoch should NOT have advanced because decrypt failed + assert_eq!(*rx.borrow(), 0, "epoch must not advance on decrypt failure"); + + tx.send(RefreshTrigger::Shutdown).await.unwrap(); + let _ = handle.await; + } + + #[tokio::test] + async fn multiple_refresh_failures_still_dont_advance_epoch() { + let (tx, rx, _shutdown, task) = make_test_task(); + let handle = tokio::spawn(task.run()); + + // Send 3 Refresh signals rapidly — all will fail + for _ in 0..3 { + tx.send(RefreshTrigger::Refresh).await.unwrap(); + } + + // Give refresh task time to process + tokio::time::sleep(std::time::Duration::from_millis(500)).await; + + // Epoch should still be 0 — no successful refreshes + assert_eq!( + *rx.borrow(), + 0, + "epoch must not advance on repeated failures" + ); + + tx.send(RefreshTrigger::Shutdown).await.unwrap(); + let _ = handle.await; + } + + #[tokio::test] + async fn cancellation_token_exits_task() { + let (_tx, _rx, shutdown, task) = make_test_task(); + let handle = tokio::spawn(task.run()); + + // Cancel the token — task should exit promptly + shutdown.cancel(); + tokio::time::timeout(std::time::Duration::from_secs(2), handle) + .await + .expect("task should exit on CancellationToken cancel") + .expect("task should not panic"); + } + + #[test] + fn is_message_shard_path_recognizes_numbered_shards() { + assert!(is_message_shard_path("message/message_0.db")); + assert!(is_message_shard_path("message/message_1.db")); + assert!(is_message_shard_path("message/message_12.db")); + assert!(is_message_shard_path("message/message_999.db")); + } + + #[test] + fn is_message_shard_path_rejects_non_shards() { + // FTS database is not a numbered shard + assert!(!is_message_shard_path("message/message_fts.db")); + // contact.db is not a message shard + assert!(!is_message_shard_path("contact/contact.db")); + // session.db is not a message shard + assert!(!is_message_shard_path("session/session.db")); + // No message/ prefix + assert!(!is_message_shard_path("message_0.db")); + // Not a .db file + assert!(!is_message_shard_path("message/message_0.txt")); + // Empty suffix + assert!(!is_message_shard_path("message/message_.db")); + // Non-numeric suffix + assert!(!is_message_shard_path("message/message_abc.db")); + } + + #[test] + fn path_categorization_contact() { + assert!("contact/contact.db".ends_with("contact/contact.db")); + assert!(!"message/message_0.db".ends_with("contact/contact.db")); + } + + #[test] + fn path_categorization_fts() { + assert!("message/message_fts.db".ends_with("message/message_fts.db")); + assert!(!"message/message_0.db".ends_with("message/message_fts.db")); + } +} diff --git a/crates/wx-cli/src/cmd/serve/routes.rs b/crates/wx-cli/src/cmd/serve/routes.rs new file mode 100644 index 0000000..ab68380 --- /dev/null +++ b/crates/wx-cli/src/cmd/serve/routes.rs @@ -0,0 +1,27 @@ +use std::sync::Arc; + +use axum::middleware; +use axum::routing::get; +use axum::Router; +use tower_http::cors::CorsLayer; + +use super::auth; +use super::handlers; +use super::state::AppState; + +pub fn build_router(state: Arc) -> Router { + Router::new() + .route("/api/v1/health", get(handlers::handler_health)) + .route("/api/v1/sessions", get(handlers::handler_sessions)) + .route("/api/v1/contacts", get(handlers::handler_contacts)) + .route("/api/v1/messages", get(handlers::handler_messages)) + .route("/api/v1/media", get(handlers::handler_media)) + .route("/api/v1/search", get(handlers::handler_search)) + .route("/api/v1/events", get(handlers::handler_sse)) + .layer(middleware::from_fn_with_state( + state.clone(), + auth::bearer_auth, + )) + .layer(CorsLayer::permissive()) + .with_state(state) +} diff --git a/crates/wx-cli/src/cmd/serve/state.rs b/crates/wx-cli/src/cmd/serve/state.rs new file mode 100644 index 0000000..09fbdd3 --- /dev/null +++ b/crates/wx-cli/src/cmd/serve/state.rs @@ -0,0 +1,86 @@ +use std::collections::HashMap; +use std::path::PathBuf; +use std::sync::atomic::AtomicBool; +use std::sync::Arc; + +use lru::LruCache; +use rusqlite::Connection; +use tokio::sync::{broadcast, mpsc}; +use tokio_util::sync::CancellationToken; +use wx_context::{ContactResolver, VisibilityIndex}; +use wx_db::WechatDb; +use wx_media::DatDecryptOptions; + +use super::event::SseEvent; +use super::refresh::RefreshTrigger; + +#[derive(Clone)] +pub struct CachedVoicePayload { + pub bytes: Vec, + pub content_type: &'static str, +} + +#[derive(Clone)] +pub struct CurrentAccount { + pub wxid: String, + pub name: String, +} + +pub struct AppState { + /// WechatDb behind std::sync::Mutex — accessed only via spawn_blocking. + pub db: Arc>, + /// Self wxid for direction enrichment. + pub self_wxid: String, + /// Active account metadata for the currently served dataset. + pub current_account: CurrentAccount, + /// Stable per-worker identity persisted in runtime state and returned by health probes. + pub worker_id: String, + /// Shared CLI version string used by health/status surfaces. + pub cli_version: String, + /// Contact name resolver (read-only after construction). + pub resolver: Arc, + /// Compiled talker-level visibility rules for this worker. + pub visibility: Arc, + /// Broadcast channel for SSE events. + pub broadcast_tx: broadcast::Sender>, + /// Optional Bearer token for auth. + pub auth_token: Option, + /// Bridge initialization complete flag. SSE returns 503 until true. + pub ready: AtomicBool, + /// Channel for bridge to signal refresh task. + /// Held here to keep the sender alive; bridge receives its own clone directly. + #[allow(dead_code)] + pub refresh_tx: mpsc::Sender, + /// Shutdown coordination token — cancelled on SIGTERM/SIGINT. + pub shutdown: CancellationToken, + /// Independent FTS connection outside the main WechatDb Mutex. + /// Used by handler_search to avoid holding the main lock during FTS queries. + /// Wrapped in its own Mutex for thread-safe reopen. + pub fts_conn: Option>>, + /// Root attach directory for `.dat` image lookup. + pub attach_dir: PathBuf, + /// Directory containing `media*.db` voice shards for the active mode. + pub media_db_dir: PathBuf, + /// Root file directory for file fallback lookup. + pub file_dir: PathBuf, + /// Root video directory for video fallback lookup. + pub video_dir: PathBuf, + /// `hardlink.db` path for the active mode. + pub hardlink_db_path: PathBuf, + /// Cached connection to `hardlink.db` — lazily opened and cleared on refresh. + pub hardlink_db_conn: Arc>>, + /// Optional raw key for direct encrypted media access. + pub raw_key: Option<[u8; 32]>, + /// Image `.dat` decryption options shared by media handlers. + pub dat_decrypt: DatDecryptOptions, + /// Small in-memory cache for transcoded voice responses, keyed by `server_id:format`. + pub voice_cache: Arc>>, + /// Per-talker cached XOR keys for image `.dat` lookup. + pub image_xor_cache: Arc>>>, + /// Lazily initialized name2id cache for FTS search resolution. + /// Cleared on FTS reopen so stale data doesn't persist. + pub name2id_cache: Arc>>>, + /// Lazily initialized list of media database paths for voice lookup. + /// Cleared on refresh so new media DBs are discovered. + pub media_db_paths: Arc>>>, +} diff --git a/crates/wx-cli/src/cmd/server/manager.rs b/crates/wx-cli/src/cmd/server/manager.rs new file mode 100644 index 0000000..e584ddb --- /dev/null +++ b/crates/wx-cli/src/cmd/server/manager.rs @@ -0,0 +1,434 @@ +use std::fs::OpenOptions; +use std::process::{Child, Command, Stdio}; +use std::time::{Duration, Instant}; + +use wx_paths::AppPaths; + +use super::runtime::{ + acquire_management_lock, base_url, load_launch_config, load_runtime_state, pid_is_running, + pid_matches_managed_worker, probe_health, remove_runtime_state, save_launch_config, + save_runtime_state, terminate_pid, wait_for_pid_exit, +}; +use super::types::{ + RuntimeAccountState, ServerHealthState, ServerLaunchConfig, ServerRestartArgs, ServerRunArgs, + ServerRuntimeState, ServerStatusArgs, ServerStatusKind, ServerStatusReport, ServerStopArgs, + WorkerLifecycle, +}; +use crate::version; +use crate::OutputFormat; + +const START_TIMEOUT: Duration = Duration::from_secs(10); +const STOP_TIMEOUT: Duration = Duration::from_secs(5); + +fn resolve_app_paths(runtime_root: Option) -> Result> { + match runtime_root { + Some(root) => Ok(AppPaths::with_runtime_root(root)?), + None => Ok(AppPaths::new()?), + } +} + +use std::path::PathBuf; + +pub async fn cmd_server_run(args: ServerRunArgs) -> Result<(), Box> { + let args_runtime_root = args.runtime_root.clone(); + let ap = resolve_app_paths(args_runtime_root.clone())?; + let _lock = acquire_management_lock(&ap)?; + let config: ServerLaunchConfig = args.into(); + + validate_launch_config(&config)?; + ap.ensure_server_dirs()?; + + let existing_state = load_runtime_state(&ap)?; + let report = build_status_report(&ap)?; + if existing_state + .as_ref() + .is_some_and(|state| pid_is_running(state.pid)) + { + if let Some(existing_config) = load_launch_config(&ap)? { + if existing_config != config { + return Err( + "server already running with different launch configuration; stop or restart it before changing host/port/token" + .into(), + ); + } + } + return match report.status { + ServerStatusKind::Running | ServerStatusKind::Starting => { + print_run_summary("server already running", &report); + Ok(()) + } + _ => Err( + "managed server worker is still running but unhealthy; use `wx-cli server stop` or `wx-cli server restart`" + .into(), + ), + }; + } + + save_launch_config(&ap, &config)?; + remove_runtime_state(&ap)?; + let worker_id = generate_worker_id(); + let mut child = spawn_worker(&ap, &config, &worker_id, args_runtime_root.as_ref())?; + let state = starting_state(&ap, &config, child.id(), worker_id); + save_runtime_state(&ap, &state)?; + + wait_for_worker_ready(&ap, &config, &mut child)?; + let report = build_status_report(&ap)?; + print_run_summary("server started", &report); + Ok(()) +} + +pub fn cmd_server_status(args: ServerStatusArgs) -> Result<(), Box> { + let ap = resolve_app_paths(args.runtime_root)?; + let report = build_status_report(&ap)?; + match args.format { + OutputFormat::Json => println!("{}", serde_json::to_string_pretty(&report)?), + OutputFormat::Text => print_status_text(&report), + } + Ok(()) +} + +pub async fn cmd_server_stop(args: ServerStopArgs) -> Result<(), Box> { + let ap = resolve_app_paths(args.runtime_root)?; + let _lock = acquire_management_lock(&ap)?; + let Some(state) = load_runtime_state(&ap)? else { + println!("server not running"); + return Ok(()); + }; + + if !pid_is_running(state.pid) { + remove_runtime_state(&ap)?; + println!( + "removed stale server state (pid {} is not running)", + state.pid + ); + return Ok(()); + } + + terminate_pid(state.pid)?; + if !wait_for_pid_exit(state.pid, STOP_TIMEOUT) { + return Err(format!( + "server pid {} did not exit within {:?}", + state.pid, STOP_TIMEOUT + ) + .into()); + } + + remove_runtime_state(&ap)?; + println!("server stopped"); + Ok(()) +} + +pub async fn cmd_server_restart(args: ServerRestartArgs) -> Result<(), Box> { + let ap = resolve_app_paths(args.runtime_root.clone())?; + let config = load_launch_config(&ap)?.ok_or( + "no persisted server launch configuration found; run `wx-cli server run` first", + )?; + + let stop_args = ServerStopArgs { + runtime_root: args.runtime_root.clone(), + }; + let status = build_status_report(&ap)?; + if !matches!(status.status, ServerStatusKind::NotRunning) { + cmd_server_stop(stop_args).await?; + } + + cmd_server_run(ServerRunArgs { + key: config.key, + data_dir: config.data_dir, + account: config.account, + poll: config.poll, + fsnotify: config.fsnotify, + poll_ms: config.poll_ms, + host: config.host, + port: config.port, + token: config.token, + runtime_root: args.runtime_root, + }) + .await +} + +pub fn build_status_report( + ap: &AppPaths, +) -> Result> { + let state = load_runtime_state(ap)?; + let config = load_launch_config(ap)?; + + let mut notes = Vec::new(); + + let mut report = ServerStatusReport { + status: ServerStatusKind::NotRunning, + runtime_root: ap.server_state_dir(), + state_file: ap.server_state_file(), + config_file: ap.server_config_file(), + stdout_log: ap.server_stdout_log(), + stderr_log: ap.server_stderr_log(), + pid: state.as_ref().map(|s| s.pid), + base_url: state + .as_ref() + .map(|s| s.base_url.clone()) + .or_else(|| config.as_ref().map(|c| base_url(&c.host, c.port))), + ready: state.as_ref().is_some_and(|s| s.ready), + health: ServerHealthState::Skipped, + cli_version: state.as_ref().map(|s| s.cli_version.clone()), + current_account: state.as_ref().and_then(|s| s.current_account.clone()), + notes: Vec::new(), + }; + + match state { + None => { + if config.is_some() { + notes.push( + "persisted launch configuration exists, but no active runtime state" + .to_string(), + ); + } + } + Some(state) if !pid_is_running(state.pid) => { + report.status = ServerStatusKind::Stale; + notes.push(format!( + "runtime state references pid {} but that process is not running", + state.pid + )); + } + Some(state) => match probe_health( + &state.base_url, + config.as_ref().and_then(|c| c.token.as_deref()), + ) { + Ok(health) if health.ready => { + report.status = ServerStatusKind::Running; + report.health = ServerHealthState::Healthy; + report.ready = true; + report.cli_version = Some(health.cli_version); + report.current_account = Some(health.current_account); + } + Ok(health) if health.worker_id != state.worker_id => { + report.status = ServerStatusKind::Broken; + report.health = ServerHealthState::Unreachable; + notes.push(format!( + "health probe returned worker_id {} but runtime state expected {}", + health.worker_id, state.worker_id + )); + if !pid_matches_managed_worker(state.pid, &state.worker_id) { + notes.push(format!( + "pid {} also does not match the stored managed worker identity", + state.pid + )); + } + } + Ok(health) => { + report.status = match state.lifecycle { + WorkerLifecycle::Starting => ServerStatusKind::Starting, + WorkerLifecycle::Stopping => ServerStatusKind::Stopping, + WorkerLifecycle::Running => ServerStatusKind::Broken, + }; + report.health = ServerHealthState::NotReady; + report.ready = false; + report.cli_version = Some(health.cli_version); + report.current_account = Some(health.current_account); + notes.push("health probe returned ready=false".to_string()); + } + Err(err) => { + report.status = match state.lifecycle { + WorkerLifecycle::Starting => ServerStatusKind::Starting, + WorkerLifecycle::Stopping => ServerStatusKind::Stopping, + WorkerLifecycle::Running => ServerStatusKind::Broken, + }; + report.health = ServerHealthState::Unreachable; + notes.push(format!("health probe failed: {err}")); + if !pid_matches_managed_worker(state.pid, &state.worker_id) { + notes.push(format!( + "pid {} does not match the stored managed worker identity", + state.pid + )); + } + } + }, + } + + report.notes = notes; + Ok(report) +} + +fn validate_launch_config(config: &ServerLaunchConfig) -> Result<(), Box> { + if !crate::cmd::serve::is_loopback(&config.host) && config.token.is_none() { + return Err("--token is required when --host is not loopback".into()); + } + Ok(()) +} + +fn spawn_worker( + ap: &AppPaths, + config: &ServerLaunchConfig, + worker_id: &str, + runtime_root: Option<&PathBuf>, +) -> Result> { + let stdout = OpenOptions::new() + .create(true) + .append(true) + .open(ap.server_stdout_log())?; + let stderr = OpenOptions::new() + .create(true) + .append(true) + .open(ap.server_stderr_log())?; + + let mut command = Command::new(std::env::current_exe()?); + command + .arg("server") + .arg("_worker") + .stdin(Stdio::null()) + .stdout(Stdio::from(stdout)) + .stderr(Stdio::from(stderr)); + + // Only pass --runtime-root when the user explicitly specified it. + // In default mode, the worker calls AppPaths::new() and uses + // platform-native paths (separate state and logs directories). + if let Some(root) = runtime_root { + command.arg("--runtime-root").arg(root); + } + + command + .arg("--worker-id") + .arg(worker_id); + + if let Some(key) = &config.key { + command.arg("--key").arg(key); + } + if let Some(data_dir) = &config.data_dir { + command.arg("--data-dir").arg(data_dir); + } + if let Some(account) = &config.account { + command.arg("--account").arg(account); + } + if config.poll { + command.arg("--poll"); + } + if config.fsnotify { + command.arg("--fsnotify"); + } + command.arg("--poll-ms").arg(config.poll_ms.to_string()); + command.arg("--host").arg(&config.host); + command.arg("--port").arg(config.port.to_string()); + if let Some(token) = &config.token { + command.arg("--token").arg(token); + } + + #[cfg(unix)] + { + use std::os::unix::process::CommandExt; + command.process_group(0); + } + + Ok(command.spawn()?) +} + +fn wait_for_worker_ready( + ap: &AppPaths, + _config: &ServerLaunchConfig, + child: &mut Child, +) -> Result<(), Box> { + let deadline = Instant::now() + START_TIMEOUT; + while Instant::now() < deadline { + if let Some(status) = child.try_wait()? { + return Err(format!( + "server worker exited early with {status}; inspect {}", + ap.server_stderr_log().display() + ) + .into()); + } + + if let Ok(report) = build_status_report(ap) { + if matches!(report.status, ServerStatusKind::Running) { + return Ok(()); + } + } + + std::thread::sleep(Duration::from_millis(100)); + } + + let _ = child.kill(); + Err(format!( + "server worker did not become ready within {:?}; inspect {}", + START_TIMEOUT, + ap.server_stderr_log().display() + ) + .into()) +} + +fn starting_state( + ap: &AppPaths, + config: &ServerLaunchConfig, + pid: u32, + worker_id: String, +) -> ServerRuntimeState { + ServerRuntimeState { + pid, + worker_id, + lifecycle: WorkerLifecycle::Starting, + ready: false, + host: config.host.clone(), + port: config.port, + base_url: base_url(&config.host, config.port), + token_configured: config.token.is_some(), + cli_version: version::cli_version_string(), + current_account: None, + stdout_log: ap.server_stdout_log(), + stderr_log: ap.server_stderr_log(), + } +} + +fn generate_worker_id() -> String { + format!( + "worker-{}-{}", + std::process::id(), + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .map(|d| d.as_nanos()) + .unwrap_or_default() + ) +} + +fn print_run_summary(prefix: &str, report: &ServerStatusReport) { + println!("{prefix}"); + print_status_text(report); +} + +fn print_status_text(report: &ServerStatusReport) { + let status = match report.status { + ServerStatusKind::NotRunning => "not running", + ServerStatusKind::Starting => "starting", + ServerStatusKind::Running => "running", + ServerStatusKind::Stopping => "stopping", + ServerStatusKind::Stale => "stale", + ServerStatusKind::Broken => "broken", + }; + println!("Server: {status}"); + if let Some(pid) = report.pid { + println!("PID: {pid}"); + } + if let Some(base_url) = &report.base_url { + println!("Base URL: {base_url}"); + } + println!("Ready: {}", if report.ready { "yes" } else { "no" }); + println!( + "Health: {}", + match report.health { + ServerHealthState::Healthy => "healthy", + ServerHealthState::NotReady => "not ready", + ServerHealthState::Unreachable => "unreachable", + ServerHealthState::Skipped => "skipped", + } + ); + if let Some(version) = &report.cli_version { + println!("CLI: {version}"); + } + if let Some(RuntimeAccountState { wxid, name }) = &report.current_account { + println!("Account: {wxid} ({name})"); + } + println!("Runtime root: {}", report.runtime_root.display()); + println!("State file: {}", report.state_file.display()); + println!("Stdout log: {}", report.stdout_log.display()); + println!("Stderr log: {}", report.stderr_log.display()); + for note in &report.notes { + println!("Note: {note}"); + } +} diff --git a/crates/wx-cli/src/cmd/server/mod.rs b/crates/wx-cli/src/cmd/server/mod.rs new file mode 100644 index 0000000..0baa3d1 --- /dev/null +++ b/crates/wx-cli/src/cmd/server/mod.rs @@ -0,0 +1,43 @@ +pub mod manager; +pub mod runtime; +pub mod types; + +use runtime::RuntimeReporter; +pub use types::{ServerAction, ServerWorkerArgs}; + +use crate::cmd::serve; + +pub async fn cmd_server(action: ServerAction) -> Result<(), Box> { + match action { + ServerAction::Run(args) => manager::cmd_server_run(args).await, + ServerAction::Status(args) => manager::cmd_server_status(args), + ServerAction::Stop(args) => manager::cmd_server_stop(args).await, + ServerAction::Restart(args) => manager::cmd_server_restart(args).await, + ServerAction::Worker(args) => cmd_server_worker(args).await, + } +} + +async fn cmd_server_worker(args: ServerWorkerArgs) -> Result<(), Box> { + let ap = match args.runtime_root.clone() { + Some(root) => wx_paths::AppPaths::with_runtime_root(root)?, + None => wx_paths::AppPaths::new()?, + }; + ap.ensure_server_dirs()?; + + let config = types::ServerLaunchConfig::from(args.clone()); + let reporter = RuntimeReporter::new(ap, config.clone()); + serve::cmd_serve( + config.key, + config.data_dir, + config.account, + config.poll, + config.fsnotify, + config.poll_ms, + config.host, + config.port, + config.token, + args.worker_id, + Some(reporter), + ) + .await +} diff --git a/crates/wx-cli/src/cmd/server/runtime.rs b/crates/wx-cli/src/cmd/server/runtime.rs new file mode 100644 index 0000000..1df713f --- /dev/null +++ b/crates/wx-cli/src/cmd/server/runtime.rs @@ -0,0 +1,205 @@ +use std::fs::{self, OpenOptions}; +use std::io::{ErrorKind, Write}; +use std::path::{Path, PathBuf}; +use std::process::Command; +use std::time::{Duration, Instant}; + +use serde::de::DeserializeOwned; +use serde::Serialize; +use wx_paths::AppPaths; + +use super::types::{ServerHealthPayload, ServerLaunchConfig, ServerRuntimeState}; + +pub struct ManagementLockGuard { + lock_file: PathBuf, +} + +impl Drop for ManagementLockGuard { + fn drop(&mut self) { + let _ = fs::remove_file(&self.lock_file); + } +} + +pub fn acquire_management_lock( + ap: &AppPaths, +) -> Result> { + ap.ensure_server_dirs()?; + let lock_file = ap.server_lock_file(); + + loop { + match OpenOptions::new() + .create_new(true) + .write(true) + .open(&lock_file) + { + Ok(mut file) => { + writeln!(file, "{}", std::process::id())?; + return Ok(ManagementLockGuard { lock_file }); + } + Err(err) if err.kind() == ErrorKind::AlreadyExists => { + let owner = fs::read_to_string(&lock_file) + .ok() + .and_then(|s| s.trim().parse::().ok()); + if owner.is_some_and(pid_is_running) { + return Err("another server management operation is in progress".into()); + } + fs::remove_file(&lock_file)?; + } + Err(err) => return Err(err.into()), + } + } +} + +pub fn save_json(path: &Path, value: &T) -> Result<(), Box> { + if let Some(parent) = path.parent() { + fs::create_dir_all(parent)?; + } + + let temp = path.with_extension("json.tmp"); + fs::write(&temp, serde_json::to_vec_pretty(value)?)?; + fs::rename(temp, path)?; + Ok(()) +} + +pub fn load_json( + path: &Path, +) -> Result, Box> { + match fs::read(path) { + Ok(bytes) => Ok(Some(serde_json::from_slice(&bytes)?)), + Err(err) if err.kind() == ErrorKind::NotFound => Ok(None), + Err(err) => Err(err.into()), + } +} + +pub fn load_runtime_state( + ap: &AppPaths, +) -> Result, Box> { + load_json(&ap.server_state_file()) +} + +pub fn save_runtime_state( + ap: &AppPaths, + state: &ServerRuntimeState, +) -> Result<(), Box> { + save_json(&ap.server_state_file(), state) +} + +pub fn load_launch_config( + ap: &AppPaths, +) -> Result, Box> { + load_json(&ap.server_config_file()) +} + +pub fn save_launch_config( + ap: &AppPaths, + config: &ServerLaunchConfig, +) -> Result<(), Box> { + save_json(&ap.server_config_file(), config) +} + +pub fn remove_runtime_state(ap: &AppPaths) -> Result<(), Box> { + match fs::remove_file(ap.server_state_file()) { + Ok(()) => Ok(()), + Err(err) if err.kind() == ErrorKind::NotFound => Ok(()), + Err(err) => Err(err.into()), + } +} + +pub fn pid_is_running(pid: u32) -> bool { + let result = unsafe { libc::kill(pid as i32, 0) }; + if result == 0 { + return true; + } + std::io::Error::last_os_error().raw_os_error() == Some(libc::EPERM) +} + +pub fn pid_matches_managed_worker(pid: u32, worker_id: &str) -> bool { + process_command_line(pid) + .map(|command| command.contains("server _worker") && command.contains(worker_id)) + .unwrap_or(false) +} + +pub fn terminate_pid(pid: u32) -> Result<(), Box> { + let result = unsafe { libc::kill(pid as i32, libc::SIGTERM) }; + if result == 0 { + Ok(()) + } else { + Err(std::io::Error::last_os_error().into()) + } +} + +fn process_command_line(pid: u32) -> Option { + let output = Command::new("ps") + .args(["-o", "command=", "-p", &pid.to_string()]) + .output() + .ok()?; + if !output.status.success() { + return None; + } + let command = String::from_utf8(output.stdout).ok()?; + let trimmed = command.trim(); + if trimmed.is_empty() { + None + } else { + Some(trimmed.to_string()) + } +} + +pub fn wait_for_pid_exit(pid: u32, timeout: Duration) -> bool { + let deadline = Instant::now() + timeout; + while Instant::now() < deadline { + if !pid_is_running(pid) { + return true; + } + std::thread::sleep(Duration::from_millis(50)); + } + !pid_is_running(pid) +} + +pub fn base_url(host: &str, port: u16) -> String { + if host.contains(':') { + format!("http://[{host}]:{port}") + } else { + format!("http://{host}:{port}") + } +} + +pub fn probe_health( + base_url: &str, + token: Option<&str>, +) -> Result> { + let agent = ureq::AgentBuilder::new() + .timeout_connect(Duration::from_millis(250)) + .timeout_read(Duration::from_millis(750)) + .build(); + let url = format!("{}/api/v1/health", base_url.trim_end_matches('/')); + let mut request = agent.get(&url); + if let Some(token) = token { + request = request.set("Authorization", &format!("Bearer {token}")); + } + let response = request.call()?; + Ok(response.into_json::()?) +} + +#[derive(Clone, Debug)] +pub struct RuntimeReporter { + ap: AppPaths, +} + +impl RuntimeReporter { + pub fn new(ap: AppPaths, _config: ServerLaunchConfig) -> Self { + Self { ap } + } + + pub fn write_state(&self, state: ServerRuntimeState) -> Result<(), Box> { + save_runtime_state(&self.ap, &state) + } + + pub fn clear_state(&self) -> Result<(), Box> { + remove_runtime_state(&self.ap) + } + + pub fn ap(&self) -> &AppPaths { + &self.ap + } +} diff --git a/crates/wx-cli/src/cmd/server/types.rs b/crates/wx-cli/src/cmd/server/types.rs new file mode 100644 index 0000000..7e04659 --- /dev/null +++ b/crates/wx-cli/src/cmd/server/types.rs @@ -0,0 +1,259 @@ +use std::path::PathBuf; + +use clap::{Args, Subcommand}; +use serde::{Deserialize, Serialize}; + +use crate::OutputFormat; + +#[derive(Args, Clone, Debug)] +pub struct ServerRunArgs { + /// 32-byte hex key (overrides KeyStore lookup) + #[arg(short, long)] + pub key: Option, + + /// WeChat data directory (auto-detect if omitted) + #[arg(short, long)] + pub data_dir: Option, + + /// Account directory name or base account ID + #[arg(long)] + pub account: Option, + + /// Force mtime polling instead of fsnotify + #[arg(long, conflicts_with = "fsnotify")] + pub poll: bool, + + /// Force fsnotify backend (opt-in on macOS, where polling is the default) + #[arg(long, conflicts_with = "poll")] + pub fsnotify: bool, + + /// Polling interval in milliseconds + #[arg(long, default_value = "2000")] + pub poll_ms: u64, + + /// Listen host address + #[arg(long, default_value = "127.0.0.1")] + pub host: String, + + /// Listen port + #[arg(long, default_value = "9100")] + pub port: u16, + + /// Bearer token for authentication (required when --host is not loopback) + #[arg(long)] + pub token: Option, + + /// Internal runtime root override for tests/non-public plumbing + #[arg(long, hide = true)] + pub runtime_root: Option, +} + +#[derive(Args, Clone, Debug)] +pub struct ServerStatusArgs { + /// Output format + #[arg(long, default_value = "text", value_enum)] + pub format: OutputFormat, + + /// Internal runtime root override for tests/non-public plumbing + #[arg(long, hide = true)] + pub runtime_root: Option, +} + +#[derive(Args, Clone, Debug)] +pub struct ServerStopArgs { + /// Internal runtime root override for tests/non-public plumbing + #[arg(long, hide = true)] + pub runtime_root: Option, +} + +#[derive(Args, Clone, Debug)] +pub struct ServerRestartArgs { + /// Internal runtime root override for tests/non-public plumbing + #[arg(long, hide = true)] + pub runtime_root: Option, +} + +#[derive(Args, Clone, Debug)] +pub struct ServerWorkerArgs { + /// 32-byte hex key (overrides KeyStore lookup) + #[arg(short, long)] + pub key: Option, + + /// WeChat data directory (auto-detect if omitted) + #[arg(short, long)] + pub data_dir: Option, + + /// Account directory name or base account ID + #[arg(long)] + pub account: Option, + + /// Force mtime polling instead of fsnotify + #[arg(long, conflicts_with = "fsnotify")] + pub poll: bool, + + /// Force fsnotify backend (opt-in on macOS, where polling is the default) + #[arg(long, conflicts_with = "poll")] + pub fsnotify: bool, + + /// Polling interval in milliseconds + #[arg(long, default_value = "2000")] + pub poll_ms: u64, + + /// Listen host address + #[arg(long, default_value = "127.0.0.1")] + pub host: String, + + /// Listen port + #[arg(long, default_value = "9100")] + pub port: u16, + + /// Bearer token for authentication (required when --host is not loopback) + #[arg(long)] + pub token: Option, + + /// Internal runtime root override for tests/non-public plumbing + #[arg(long, hide = true)] + pub runtime_root: Option, + + /// Internal worker identity used to verify the managed process + #[arg(long, hide = true)] + pub worker_id: Option, +} + +#[derive(Subcommand, Clone, Debug)] +pub enum ServerAction { + /// Start the managed HTTP service + Run(ServerRunArgs), + /// Show managed service status + Status(ServerStatusArgs), + /// Stop the managed HTTP service + Stop(ServerStopArgs), + /// Restart the managed HTTP service + Restart(ServerRestartArgs), + /// Hidden foreground worker for the service manager and integration tests + #[command(name = "_worker", hide = true)] + Worker(ServerWorkerArgs), +} + +#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)] +pub struct ServerLaunchConfig { + pub key: Option, + pub data_dir: Option, + pub account: Option, + pub poll: bool, + pub fsnotify: bool, + pub poll_ms: u64, + pub host: String, + pub port: u16, + pub token: Option, +} + +impl From for ServerLaunchConfig { + fn from(value: ServerRunArgs) -> Self { + Self { + key: value.key, + data_dir: value.data_dir, + account: value.account, + poll: value.poll, + fsnotify: value.fsnotify, + poll_ms: value.poll_ms, + host: value.host, + port: value.port, + token: value.token, + } + } +} + +impl From for ServerLaunchConfig { + fn from(value: ServerWorkerArgs) -> Self { + Self { + key: value.key, + data_dir: value.data_dir, + account: value.account, + poll: value.poll, + fsnotify: value.fsnotify, + poll_ms: value.poll_ms, + host: value.host, + port: value.port, + token: value.token, + } + } +} + +#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)] +pub struct RuntimeAccountState { + pub wxid: String, + pub name: String, +} + +#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)] +#[serde(rename_all = "snake_case")] +pub enum WorkerLifecycle { + Starting, + Running, + Stopping, +} + +#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)] +pub struct ServerRuntimeState { + pub pid: u32, + pub worker_id: String, + pub lifecycle: WorkerLifecycle, + pub ready: bool, + pub host: String, + pub port: u16, + pub base_url: String, + pub token_configured: bool, + pub cli_version: String, + pub current_account: Option, + pub stdout_log: PathBuf, + pub stderr_log: PathBuf, +} + +#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)] +pub struct ServerHealthPayload { + pub ready: bool, + pub worker_id: String, + pub cli_version: String, + pub current_account: RuntimeAccountState, +} + +#[derive(Clone, Debug, Serialize)] +#[serde(rename_all = "snake_case")] +pub enum ServerHealthState { + Healthy, + NotReady, + Unreachable, + Skipped, +} + +#[derive(Clone, Debug, Serialize)] +#[serde(rename_all = "snake_case")] +pub enum ServerStatusKind { + NotRunning, + Starting, + Running, + Stopping, + Stale, + Broken, +} + +#[derive(Clone, Debug, Serialize)] +pub struct ServerStatusReport { + pub status: ServerStatusKind, + /// Default: `AppPaths::server_state_dir()`. When `--runtime-root` is + /// specified, ALL runtime files (state + config + lock + logs) go under + /// this root. In default mode, logs go to `AppPaths::logs_dir()/server/`. + pub runtime_root: PathBuf, + pub state_file: PathBuf, + pub config_file: PathBuf, + pub stdout_log: PathBuf, + pub stderr_log: PathBuf, + pub pid: Option, + pub base_url: Option, + pub ready: bool, + pub health: ServerHealthState, + pub cli_version: Option, + pub current_account: Option, + pub notes: Vec, +} diff --git a/crates/wx-cli/src/cmd/sessions.rs b/crates/wx-cli/src/cmd/sessions.rs new file mode 100644 index 0000000..fc71b34 --- /dev/null +++ b/crates/wx-cli/src/cmd/sessions.rs @@ -0,0 +1,143 @@ +use std::path::PathBuf; + +use wx_context::{AccountContext, ContactResolver, ResolveParams}; + +use super::contacts::build_visibility; +use super::thin_client::{ThinClientCliArgs, ThinClientOptions}; +use crate::output::JsonEnvelope; +use crate::schema::{enrich_session, EnrichedSession}; +use crate::util::{effective_limit_all, open_db_core, print_cache_stats, print_detection_note, try_remote_or_local}; +use crate::visibility_projection::project_sessions_envelope_enriched; +use crate::{OutputFormat, SortOrderArg}; + +#[allow(clippy::too_many_arguments)] +pub fn cmd_sessions( + data_dir: Option, + account: Option, + key: Option, + limit: usize, + offset: usize, + order: SortOrderArg, + all: bool, + format: OutputFormat, + show_hidden: bool, + server: ThinClientCliArgs, +) -> Result<(), Box> { + let options = ThinClientOptions::resolve_from_process_env(server); + let effective_limit = effective_limit_all(all, limit); + + let envelope = try_remote_or_local( + &options, + |client| { + let mut query = vec![ + ("limit".to_string(), effective_limit.to_string()), + ("offset".to_string(), offset.to_string()), + ( + "order".to_string(), + match order { + SortOrderArg::Asc => "asc".to_string(), + SortOrderArg::Desc => "desc".to_string(), + }, + ), + ]; + if show_hidden { + query.push(("show_hidden".to_string(), "1".to_string())); + } + client.get_json("/api/v1/sessions", &query) + }, + || { + load_local_sessions( + data_dir, + account, + key, + effective_limit, + offset, + order.clone(), + show_hidden, + ) + }, + "sessions", + )?; + print_sessions_output(&envelope, format) +} + +fn load_local_sessions( + data_dir: Option, + account: Option, + key: Option, + effective_limit: usize, + offset: usize, + order: SortOrderArg, + show_hidden: bool, +) -> Result, Box> { + let acct = AccountContext::resolve(&ResolveParams { + account: account.as_deref(), + data_dir: data_dir.as_deref(), + key_hex: key.as_deref(), + })?; + print_detection_note(&acct); + + let (db, stats) = open_db_core(&acct, crate::util::decrypt_progress_callback)?; + if let Some(ref s) = stats { + print_cache_stats(s); + } + let resolver = ContactResolver::build(&db)?; + let self_wxid = &acct.base_wxid; + let visibility = build_visibility(&acct, &resolver); + + let result = db.query_sessions( + &wx_db::SessionQuery::new() + .limit(wx_db::MAX_QUERY_LIMIT) + .offset(0) + .order(order.into()), + )?; + + let envelope = JsonEnvelope::from_query_result(result, wx_db::MAX_QUERY_LIMIT, 0, |s| { + enrich_session(s, self_wxid, &resolver, None) + }); + + Ok(project_sessions_envelope_enriched( + envelope.items, + &visibility, + effective_limit, + offset, + &envelope.stats, + show_hidden, + )) +} + +fn print_sessions_output( + envelope: &JsonEnvelope, + format: OutputFormat, +) -> Result<(), Box> { + match format { + OutputFormat::Json => println!("{}", serde_json::to_string_pretty(envelope)?), + OutputFormat::Text => render_sessions_text(&envelope.items), + } + Ok(()) +} + +fn render_sessions_text(items: &[EnrichedSession]) { + for s in items { + let ts = chrono::DateTime::from_timestamp(s.session.sort_timestamp, 0) + .map(|dt| { + dt.with_timezone(&chrono::Local) + .format("%m-%d %H:%M") + .to_string() + }) + .unwrap_or_default(); + + let is_group = wx_db::is_group_chat(&s.session.username); + let summary = if is_group { + if let Some(ref sender) = s.session.last_sender_display_name { + format!("{sender}: {}", s.session.summary) + } else { + s.session.summary.clone() + } + } else { + s.session.summary.clone() + }; + let summary: String = summary.chars().take(60).collect(); + println!("{ts} {} {summary}", s.display_name); + } +} diff --git a/crates/wx-cli/src/cmd/status.rs b/crates/wx-cli/src/cmd/status.rs new file mode 100644 index 0000000..a08e7d2 --- /dev/null +++ b/crates/wx-cli/src/cmd/status.rs @@ -0,0 +1,125 @@ +pub fn cmd_status() -> Result<(), Box> { + // WeChat process status (pgrep-only, no lsof) + match wx_keychain::find_wechat_pid() { + Ok((pid, version)) => { + println!("WeChat: running (pid {pid}, v{version})"); + } + Err(_) => { + println!("WeChat: not running"); + } + } + + // Account directories + let accounts = wx_keychain::find_account_dirs().unwrap_or_default(); + let store = wx_keychain::KeyStore::load_default().unwrap_or_default(); + + if accounts.is_empty() { + println!("Accounts: (none found)"); + } else { + println!("Accounts:"); + for a in &accounts { + let key_entry = store.get(&a.account_id); + let has_key = key_entry.is_some(); + + // Display name from KeyStore + let display = key_entry + .and_then(|k| k.nickname.as_deref()) + .map(|n| format!(" ({n})")) + .unwrap_or_default(); + + let key_icon = if has_key { "\u{2705}" } else { "\u{2717}" }; + + // Cache status + let cache_info = cache_status(&a.account_id); + + println!(" {}{display} key {key_icon} {cache_info}", a.account_id,); + } + } + + // Paths summary + if let Ok(ap) = wx_paths::AppPaths::new() { + let config = tilde_path(&ap.config_dir()); + let cache = tilde_path(ap.cache_root()); + println!("Paths: config {config} cache {cache}"); + } + + Ok(()) +} + +fn tilde_path(path: &std::path::Path) -> String { + if let Ok(home) = std::env::var("HOME") { + let home_path = std::path::Path::new(&home); + if let Ok(suffix) = path.strip_prefix(home_path) { + return format!("~/{}", suffix.display()); + } + } + path.display().to_string() +} + +fn cache_status(account_id: &str) -> String { + let cache_dir = match wx_paths::AppPaths::new() { + Ok(ap) => ap.account_db_cache_dir(account_id), + Err(_) => return "cache: unknown".into(), + }; + + if !cache_dir.exists() { + return "no cache".into(); + } + + // Count .db files and find most recent mtime + let mut db_count = 0usize; + let mut newest = std::time::SystemTime::UNIX_EPOCH; + + if let Ok(entries) = std::fs::read_dir(&cache_dir) { + for entry in entries.flatten() { + count_db_files_recursive(&entry.path(), &mut db_count, &mut newest); + } + } + + if db_count == 0 { + return "cache empty".into(); + } + + let age = newest + .elapsed() + .map(format_duration) + .unwrap_or_else(|_| "?".into()); + + format!("{db_count} DBs last decrypt {age} ago") +} + +fn count_db_files_recursive( + path: &std::path::Path, + count: &mut usize, + newest: &mut std::time::SystemTime, +) { + if path.is_dir() { + if let Ok(entries) = std::fs::read_dir(path) { + for entry in entries.flatten() { + count_db_files_recursive(&entry.path(), count, newest); + } + } + } else if path.extension().is_some_and(|e| e == "db") { + *count += 1; + if let Ok(meta) = path.metadata() { + if let Ok(mtime) = meta.modified() { + if mtime > *newest { + *newest = mtime; + } + } + } + } +} + +fn format_duration(d: std::time::Duration) -> String { + let secs = d.as_secs(); + if secs < 60 { + format!("{secs}s") + } else if secs < 3600 { + format!("{}m", secs / 60) + } else if secs < 86400 { + format!("{}h", secs / 3600) + } else { + format!("{}d", secs / 86400) + } +} diff --git a/crates/wx-cli/src/cmd/thin_client.rs b/crates/wx-cli/src/cmd/thin_client.rs new file mode 100644 index 0000000..689026a --- /dev/null +++ b/crates/wx-cli/src/cmd/thin_client.rs @@ -0,0 +1,246 @@ +use std::fmt; +use std::time::Duration; + +use serde::de::DeserializeOwned; +use url::Url; + +pub const DEFAULT_SERVER_URL: &str = "http://127.0.0.1:9100"; +const CONNECT_TIMEOUT_MS: u64 = 250; +const READ_TIMEOUT_MS: u64 = 1_500; + +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub struct ThinClientCliArgs { + pub server_url: Option, + pub server_token: Option, + pub server_only: bool, + pub no_server: bool, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum ServerMode { + Auto, + ServerOnly, + Disabled, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct ThinClientOptions { + pub base_url: String, + pub token: Option, + pub mode: ServerMode, +} + +impl ThinClientOptions { + pub fn resolve( + cli: ThinClientCliArgs, + env_url: Option, + env_token: Option, + ) -> Self { + let base_url = cli + .server_url + .or(env_url) + .unwrap_or_else(|| DEFAULT_SERVER_URL.to_string()); + let token = cli.server_token.or(env_token); + let mode = if cli.no_server { + ServerMode::Disabled + } else if cli.server_only { + ServerMode::ServerOnly + } else { + ServerMode::Auto + }; + Self { + base_url, + token, + mode, + } + } + + pub fn resolve_from_process_env(cli: ThinClientCliArgs) -> Self { + Self::resolve( + cli, + std::env::var("WECHAT_CLI_SERVER_URL").ok(), + std::env::var("WECHAT_CLI_SERVER_TOKEN").ok(), + ) + } + + pub fn is_enabled(&self) -> bool { + self.mode != ServerMode::Disabled + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum ThinClientErrorKind { + Unavailable, + Unauthorized, + BadRequest, + Server, + Decode, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum ThinClientDecision { + FallbackToLocal, + Fail, +} + +impl ThinClientDecision { + pub fn from_error(mode: ServerMode, kind: ThinClientErrorKind) -> Self { + if mode == ServerMode::ServerOnly { + return Self::Fail; + } + match kind { + ThinClientErrorKind::Unavailable | ThinClientErrorKind::Unauthorized => { + Self::FallbackToLocal + } + ThinClientErrorKind::BadRequest + | ThinClientErrorKind::Server + | ThinClientErrorKind::Decode => Self::Fail, + } + } +} + +#[derive(Debug)] +pub struct ThinClientError { + pub kind: ThinClientErrorKind, + message: String, +} + +impl ThinClientError { + fn new(kind: ThinClientErrorKind, message: impl Into) -> Self { + Self { + kind, + message: message.into(), + } + } + + pub fn should_fallback(&self, mode: ServerMode) -> bool { + ThinClientDecision::from_error(mode, self.kind) == ThinClientDecision::FallbackToLocal + } + + pub fn fallback_detail(&self) -> &str { + match self.kind { + ThinClientErrorKind::Unavailable => simplify_unavailable_message(&self.message), + ThinClientErrorKind::Unauthorized => { + if self.message.trim().is_empty() { + "unauthorized" + } else { + &self.message + } + } + ThinClientErrorKind::BadRequest + | ThinClientErrorKind::Server + | ThinClientErrorKind::Decode => &self.message, + } + } +} + +impl fmt::Display for ThinClientError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "{}", self.message) + } +} + +impl std::error::Error for ThinClientError {} + +#[derive(serde::Deserialize)] +struct HealthPayload { + ready: bool, +} + +pub struct ThinClient { + options: ThinClientOptions, + agent: ureq::Agent, +} + +impl ThinClient { + pub fn new(options: ThinClientOptions) -> Self { + let agent = ureq::AgentBuilder::new() + .timeout_connect(Duration::from_millis(CONNECT_TIMEOUT_MS)) + .timeout_read(Duration::from_millis(READ_TIMEOUT_MS)) + .build(); + Self { options, agent } + } + + pub fn probe_health(&self) -> Result<(), ThinClientError> { + let query: Vec<(String, String)> = Vec::new(); + let health: HealthPayload = self.get_json("/api/v1/health", &query)?; + if !health.ready { + return Err(ThinClientError::new( + ThinClientErrorKind::Unavailable, + "server health check reported ready=false", + )); + } + Ok(()) + } + + pub fn get_json( + &self, + path: &str, + query: &[(String, String)], + ) -> Result { + let url = build_url(&self.options.base_url, path, query)?; + let mut request = self.agent.get(url.as_str()); + if let Some(token) = &self.options.token { + request = request.set("Authorization", &format!("Bearer {token}")); + } + + let response = request.call().map_err(classify_ureq_error)?; + response + .into_json::() + .map_err(|err| ThinClientError::new(ThinClientErrorKind::Decode, err.to_string())) + } +} + +fn build_url( + base_url: &str, + path: &str, + query: &[(String, String)], +) -> Result { + let base = Url::parse(base_url) + .map_err(|err| ThinClientError::new(ThinClientErrorKind::BadRequest, err.to_string()))?; + let mut url = base + .join(path.trim_start_matches('/')) + .map_err(|err| ThinClientError::new(ThinClientErrorKind::BadRequest, err.to_string()))?; + { + let mut pairs = url.query_pairs_mut(); + for (key, value) in query { + pairs.append_pair(key, value); + } + } + Ok(url) +} + +fn classify_ureq_error(err: ureq::Error) -> ThinClientError { + match err { + ureq::Error::Status(code, response) => { + let body = response + .into_string() + .unwrap_or_else(|_| String::from("request failed")); + match code { + 401 => ThinClientError::new(ThinClientErrorKind::Unauthorized, body), + 400..=499 => ThinClientError::new(ThinClientErrorKind::BadRequest, body), + _ => ThinClientError::new(ThinClientErrorKind::Server, body), + } + } + ureq::Error::Transport(transport) => { + ThinClientError::new(ThinClientErrorKind::Unavailable, transport.to_string()) + } + } +} + +fn simplify_unavailable_message(message: &str) -> &str { + let detail = if message.starts_with("http://") || message.starts_with("https://") { + message + .split_once(": ") + .map(|(_, rest)| rest) + .unwrap_or(message) + } else { + message + }; + + detail + .strip_prefix("Connection Failed: ") + .unwrap_or(detail) + .strip_prefix("Connect error: ") + .unwrap_or(detail.strip_prefix("Connection Failed: ").unwrap_or(detail)) +} diff --git a/crates/wx-cli/src/cmd/thin_client_tests.rs b/crates/wx-cli/src/cmd/thin_client_tests.rs new file mode 100644 index 0000000..db9aea8 --- /dev/null +++ b/crates/wx-cli/src/cmd/thin_client_tests.rs @@ -0,0 +1,175 @@ +use super::thin_client::{ + ServerMode, ThinClient, ThinClientCliArgs, ThinClientDecision, ThinClientErrorKind, + ThinClientOptions, +}; +use std::io::{Read, Write}; +use std::net::TcpListener; +use std::thread; + +#[test] +fn resolve_prefers_cli_over_env_and_defaults_to_loopback() { + let args = ThinClientCliArgs { + server_url: Some("http://127.0.0.1:9200".into()), + server_token: Some("cli-token".into()), + server_only: true, + no_server: false, + }; + let options = ThinClientOptions::resolve( + args, + Some("http://127.0.0.1:9300".into()), + Some("env-token".into()), + ); + + assert_eq!(options.base_url, "http://127.0.0.1:9200"); + assert_eq!(options.token.as_deref(), Some("cli-token")); + assert_eq!(options.mode, ServerMode::ServerOnly); +} + +#[test] +fn resolve_uses_env_when_cli_missing() { + let options = ThinClientOptions::resolve( + ThinClientCliArgs::default(), + Some("http://127.0.0.1:9400".into()), + Some("env-token".into()), + ); + + assert_eq!(options.base_url, "http://127.0.0.1:9400"); + assert_eq!(options.token.as_deref(), Some("env-token")); + assert_eq!(options.mode, ServerMode::Auto); +} + +#[test] +fn resolve_respects_no_server() { + let options = ThinClientOptions::resolve( + ThinClientCliArgs { + server_url: None, + server_token: None, + server_only: false, + no_server: true, + }, + None, + None, + ); + + assert_eq!(options.mode, ServerMode::Disabled); + assert_eq!(options.base_url, "http://127.0.0.1:9100"); +} + +#[test] +fn fallback_policy_distinguishes_retryable_and_terminal_errors() { + assert_eq!( + ThinClientDecision::from_error(ServerMode::Auto, ThinClientErrorKind::Unavailable), + ThinClientDecision::FallbackToLocal + ); + assert_eq!( + ThinClientDecision::from_error(ServerMode::Auto, ThinClientErrorKind::Unauthorized), + ThinClientDecision::FallbackToLocal + ); + assert_eq!( + ThinClientDecision::from_error(ServerMode::Auto, ThinClientErrorKind::BadRequest), + ThinClientDecision::Fail + ); + assert_eq!( + ThinClientDecision::from_error(ServerMode::ServerOnly, ThinClientErrorKind::Unavailable), + ThinClientDecision::Fail + ); +} + +#[test] +fn probe_health_requires_ready_true() { + let (base_url, handle) = spawn_mock_server(|request| { + assert!(request.starts_with("GET /api/v1/health")); + http_response("200 OK", "{\"ready\":false}") + }); + + let client = ThinClient::new(ThinClientOptions { + base_url, + token: None, + mode: ServerMode::Auto, + }); + let err = client.probe_health().expect_err("health probe should fail"); + assert_eq!(err.kind, ThinClientErrorKind::Unavailable); + handle.join().unwrap(); +} + +#[test] +fn probe_health_does_not_append_empty_query_marker() { + let (base_url, handle) = spawn_mock_server(|request| { + assert!(request.starts_with("GET /api/v1/health HTTP/1.1\r\n")); + http_response("200 OK", "{\"ready\":true}") + }); + + let client = ThinClient::new(ThinClientOptions { + base_url, + token: None, + mode: ServerMode::Auto, + }); + client.probe_health().expect("health probe should succeed"); + handle.join().unwrap(); +} + +#[test] +fn get_json_sends_bearer_token() { + let (base_url, handle) = spawn_mock_server(|request| { + assert!(request.starts_with("GET /api/v1/contacts?limit=10")); + assert!(request.contains("Authorization: Bearer secret-token\r\n")); + http_response("200 OK", "{\"ok\":true}") + }); + + let client = ThinClient::new(ThinClientOptions { + base_url, + token: Some("secret-token".into()), + mode: ServerMode::Auto, + }); + let query = vec![("limit".to_string(), "10".to_string())]; + let value: serde_json::Value = client + .get_json("/api/v1/contacts", &query) + .expect("request should succeed"); + assert_eq!(value["ok"], true); + handle.join().unwrap(); +} + +#[test] +fn get_json_classifies_unauthorized() { + let (base_url, handle) = spawn_mock_server(|request| { + assert!(request.starts_with("GET /api/v1/sessions")); + http_response("401 Unauthorized", "{\"error\":\"unauthorized\"}") + }); + + let client = ThinClient::new(ThinClientOptions { + base_url, + token: None, + mode: ServerMode::Auto, + }); + let err = client + .get_json::("/api/v1/sessions", &Vec::new()) + .expect_err("request should fail"); + assert_eq!(err.kind, ThinClientErrorKind::Unauthorized); + assert!(err.should_fallback(ServerMode::Auto)); + handle.join().unwrap(); +} + +fn spawn_mock_server( + responder: impl FnOnce(String) -> String + Send + 'static, +) -> (String, thread::JoinHandle<()>) { + let listener = TcpListener::bind("127.0.0.1:0").expect("bind mock server"); + let addr = listener.local_addr().expect("mock server addr"); + let handle = thread::spawn(move || { + let (mut stream, _) = listener.accept().expect("accept"); + let mut buf = [0_u8; 8192]; + let n = stream.read(&mut buf).expect("read request"); + let request = String::from_utf8_lossy(&buf[..n]).into_owned(); + let response = responder(request); + stream + .write_all(response.as_bytes()) + .expect("write response"); + }); + (format!("http://{}", addr), handle) +} + +fn http_response(status: &str, body: &str) -> String { + format!( + "HTTP/1.1 {status}\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}", + body.len() + ) +} diff --git a/crates/wx-cli/src/cmd/watch.rs b/crates/wx-cli/src/cmd/watch.rs new file mode 100644 index 0000000..ea9b2dd --- /dev/null +++ b/crates/wx-cli/src/cmd/watch.rs @@ -0,0 +1,384 @@ +use std::path::PathBuf; +use std::time::Duration; + +use serde::Serialize; +use wx_context::{AccountContext, ContactResolver, Direction, ResolveParams, VisibilityIndex}; +use wx_monitor::SessionEventKind; + +use super::contacts::build_visibility; +use crate::schema::{enrich_session_event, project_session_sender}; +use crate::util::{open_db_core, print_cache_stats, print_detection_note}; +use crate::OutputFormat; + +#[derive(Serialize)] +struct WatchEvent { + #[serde(flatten)] + enriched: crate::schema::EnrichedSession, + kind: SessionEventKind, +} + +#[cfg(test)] +fn format_watch_line( + ev: &wx_monitor::SessionEvent, + resolver: &ContactResolver, + self_wxid: &str, +) -> String { + use crate::schema::derive_session_direction; + let time = chrono::Local::now().format("%H:%M:%S"); + let is_group = wx_db::is_group_chat(&ev.username); + + let content = if !ev.summary.is_empty() { + ev.summary.clone() + } else if let Some(mt) = ev.last_msg_type { + format!("[{}]", wx_db::msg_type_label(mt)) + } else { + String::new() + }; + + let direction = derive_session_direction(ev.last_msg_sender.as_deref(), self_wxid); + + if is_group { + let group_name = resolver.display_with_id(&ev.username); + match direction { + Some(Direction::Outgoing) => { + format!("{time} 发送到群「{group_name}」:{content}") + } + Some(Direction::Incoming) => { + let sender = ev.last_msg_sender.as_deref().unwrap_or(""); + if sender.is_empty() { + format!("{time} 来自群「{group_name}」:{content}") + } else { + let sender_name = resolver.display_with_id(sender); + format!("{time} 来自群「{group_name}」的{sender_name}:{content}") + } + } + None => format!("{time} 群消息更新「{group_name}」:{content}"), + } + } else { + let name = resolver.display_with_id(&ev.username); + match direction { + Some(Direction::Outgoing) => format!("{time} 发送给{name}:{content}"), + Some(Direction::Incoming) => format!("{time} 来自{name}:{content}"), + None => format!("{time} 会话更新「{name}」:{content}"), + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn make_event( + username: &str, + summary: &str, + last_msg_sender: Option<&str>, + ) -> wx_monitor::SessionEvent { + wx_monitor::SessionEvent { + username: username.to_string(), + sort_timestamp: 1, + detected_at: 2, + kind: SessionEventKind::Updated, + summary: summary.to_string(), + last_msg_type: Some(1), + last_msg_sender: last_msg_sender.map(str::to_string), + last_sender_display_name: None, + } + } + + #[test] + fn watch_text_uses_neutral_wording_when_private_direction_unknown() { + let resolver = ContactResolver::empty(); + let line = format_watch_line( + &make_event("wxid_friend", "hello", None), + &resolver, + "wxid_me", + ); + + assert!(line.contains("会话更新")); + assert!(!line.contains("来自")); + assert!(!line.contains("发送给")); + } + + #[test] + fn watch_text_keeps_explicit_wording_for_known_private_outgoing() { + let resolver = ContactResolver::empty(); + let line = format_watch_line( + &make_event("wxid_friend", "hello", Some("wxid_me")), + &resolver, + "wxid_me", + ); + + assert!(line.contains("发送给wxid_friend")); + } + + #[test] + fn watch_text_uses_neutral_group_wording_when_direction_unknown() { + let resolver = ContactResolver::empty(); + let line = format_watch_line( + &make_event("team@chatroom", "hello", Some("")), + &resolver, + "wxid_me", + ); + + assert!(line.contains("群消息更新")); + assert!(!line.contains("来自群")); + assert!(!line.contains("发送到群")); + } + + #[test] + fn watch_text_legacy_non_wxid_self_id_outgoing() { + let resolver = ContactResolver::empty(); + let line = format_watch_line( + &make_event("wxid_friend", "hello", Some("testuser001")), + &resolver, + "testuser001", + ); + + assert!( + line.contains("发送给"), + "expected outgoing wording, got: {line}" + ); + } + + #[test] + fn watch_text_legacy_non_wxid_self_id_incoming() { + let resolver = ContactResolver::empty(); + let line = format_watch_line( + &make_event("wxid_friend", "hello", Some("wxid_friend")), + &resolver, + "testuser001", + ); + + assert!( + line.contains("来自"), + "expected incoming wording, got: {line}" + ); + } + + #[test] + fn hidden_talker_is_dropped_by_default() { + let visibility = + VisibilityIndex::build(&["wxid_secret".to_string()], &[], &ContactResolver::empty()); + let event = make_event("wxid_secret", "hello", None); + + assert!(!should_emit_event(&visibility, false, &event)); + } + + #[test] + fn show_hidden_restores_hidden_talker_output() { + let visibility = + VisibilityIndex::build(&["wxid_secret".to_string()], &[], &ContactResolver::empty()); + let event = make_event("wxid_secret", "hello", None); + + assert!(should_emit_event(&visibility, true, &event)); + } + + // --- Phase 2: sender-level tests --- + + fn make_enriched_session( + username: &str, + summary: &str, + last_msg_sender: Option<&str>, + ) -> crate::schema::EnrichedSession { + use crate::schema::enrich_session_event; + let ev = wx_monitor::SessionEvent { + username: username.to_string(), + sort_timestamp: 1, + detected_at: 2, + kind: wx_monitor::SessionEventKind::Updated, + summary: summary.to_string(), + last_msg_type: Some(1), + last_msg_sender: last_msg_sender.map(str::to_string), + last_sender_display_name: last_msg_sender.map(str::to_string), + }; + enrich_session_event(ev, "wxid_me", &ContactResolver::empty()) + } + + #[test] + fn watch_text_hidden_sender_shows_placeholder() { + use crate::schema::project_session_sender; + let visibility = VisibilityIndex::build( + &["wxid_spam".to_string()], &[], &ContactResolver::empty(), + ); + let mut enriched = make_enriched_session("group@chatroom", "spam msg", Some("wxid_spam")); + project_session_sender(&mut enriched, &visibility); + + let line = format_watch_line_from_enriched(&enriched, &ContactResolver::empty(), "wxid_me"); + assert!(line.contains("[消息已隐藏]"), "should show placeholder: {line}"); + assert!(!line.contains("wxid_spam"), "should not leak sender wxid: {line}"); + } + + #[test] + fn watch_text_visible_sender_shows_normal() { + let enriched = make_enriched_session("group@chatroom", "hello", Some("wxid_normal")); + let line = format_watch_line_from_enriched(&enriched, &ContactResolver::empty(), "wxid_me"); + assert!(line.contains("hello"), "should show normal summary: {line}"); + } +} + +/// Format a watch line from an EnrichedSession (Phase 2: includes sender redaction). +fn format_watch_line_from_enriched( + enriched: &crate::schema::EnrichedSession, + resolver: &ContactResolver, + _self_wxid: &str, +) -> String { + let time = chrono::Local::now().format("%H:%M:%S"); + let session = &enriched.session; + let is_group = wx_db::is_group_chat(&session.username); + + let content = if !session.summary.is_empty() { + session.summary.clone() + } else if let Some(mt) = session.last_msg_type { + format!("[{}]", wx_db::msg_type_label(mt)) + } else { + String::new() + }; + + let direction = &enriched.direction; + + if is_group { + let group_name = resolver.display_with_id(&session.username); + match direction { + Some(Direction::Outgoing) => { + format!("{time} 发送到群「{group_name}」:{content}") + } + Some(Direction::Incoming) => { + let sender = session.last_msg_sender.as_deref().unwrap_or(""); + if sender.is_empty() { + format!("{time} 来自群「{group_name}」:{content}") + } else { + let sender_name = resolver.display_with_id(sender); + format!("{time} 来自群「{group_name}」的{sender_name}:{content}") + } + } + None => format!("{time} 群消息更新「{group_name}」:{content}"), + } + } else { + let name = resolver.display_with_id(&session.username); + match direction { + Some(Direction::Outgoing) => format!("{time} 发送给{name}:{content}"), + Some(Direction::Incoming) => format!("{time} 来自{name}:{content}"), + None => format!("{time} 会话更新「{name}」:{content}"), + } + } +} + +fn resolve_watch_mode(poll: bool, fsnotify: bool) -> wx_monitor::WatchMode { + if poll { + wx_monitor::WatchMode::Poll + } else if fsnotify { + wx_monitor::WatchMode::Fsnotify + } else { + wx_monitor::WatchMode::Auto + } +} + +fn should_emit_event( + visibility: &VisibilityIndex, + show_hidden: bool, + event: &wx_monitor::SessionEvent, +) -> bool { + show_hidden || !visibility.is_hidden_talker(&event.username) +} + +pub async fn cmd_watch( + key_hex: Option, + data_dir: Option, + account: Option, + poll: bool, + fsnotify: bool, + poll_ms: u64, + format: OutputFormat, + show_hidden: bool, +) -> Result<(), Box> { + let params = &wx_decrypt::MACOS_4_1_7_31; + + let acct = AccountContext::resolve(&ResolveParams { + account: account.as_deref(), + data_dir: data_dir.as_deref(), + key_hex: key_hex.as_deref(), + })?; + print_detection_note(&acct); + + let (db, stats) = open_db_core(&acct, crate::util::decrypt_progress_callback)?; + if let Some(ref s) = stats { + print_cache_stats(s); + } + let resolver = ContactResolver::build(&db)?; + let visibility = build_visibility(&acct, &resolver); + let self_wxid = acct.base_wxid.clone(); + + let encrypted_session_dir = acct.data_dir.join("db_storage").join("session"); + if !encrypted_session_dir.exists() { + return Err(format!( + "session directory not found: {}", + encrypted_session_dir.display() + ) + .into()); + } + + let watch_mode = resolve_watch_mode(poll, fsnotify); + let config = wx_monitor::MonitorConfig { + encrypted_session_dir, + key_material: acct.key_material.clone(), + params, + watch_mode: watch_mode.clone(), + poll_interval: Duration::from_millis(poll_ms), + channel_capacity: 1000, + raw_key: acct.raw_key, + encrypted_root: if acct.raw_key.is_some() { + Some(acct.data_dir.join("db_storage")) + } else { + None + }, + }; + + let resolved = wx_monitor::resolve_watch_mode(&watch_mode); + eprintln!("Starting monitor (mode={watch_mode:?} -> {resolved:?}, interval={poll_ms}ms)..."); + let mut monitor = wx_monitor::WechatMonitor::start(config)?; + eprintln!("Monitor started. Press Ctrl+C to stop."); + + loop { + tokio::select! { + _ = tokio::signal::ctrl_c() => { + monitor.stop(); + break; + } + event = monitor.recv() => { + match event { + Some(ev) => { + if !should_emit_event(&visibility, show_hidden, &ev) { + continue; + } + let kind = ev.kind.clone(); + let mut enriched = enrich_session_event(ev, &self_wxid, &resolver); + // Phase 2: redact hidden sender in group session + if !show_hidden { + project_session_sender(&mut enriched, &visibility); + } + match format { + OutputFormat::Json => { + let watch_event = WatchEvent { enriched, kind }; + println!("{}", serde_json::to_string(&watch_event).unwrap()); + } + OutputFormat::Text => { + println!( + "{}", + format_watch_line_from_enriched( + &enriched, + &resolver, + &self_wxid + ) + ); + } + } + } + None => break, + } + } + } + } + + eprintln!("Monitor stopped."); + Ok(()) +} diff --git a/crates/wx-cli/src/contact_id.rs b/crates/wx-cli/src/contact_id.rs new file mode 100644 index 0000000..f29e948 --- /dev/null +++ b/crates/wx-cli/src/contact_id.rs @@ -0,0 +1,471 @@ +use wx_context::{ContactResolver, VisibilityIndex}; +use wx_db::is_group_chat; + +/// Result of contact resolution: the canonical wxid and optional display name. +#[derive(Debug)] +pub struct ResolvedContact { + pub wxid: String, + pub display_name: Option, +} + +/// Errors from contact resolution. +#[derive(Debug)] +pub enum ContactResolveError { + NotFound(String), + Ambiguous(String), + Hidden(String), +} + +impl std::fmt::Display for ContactResolveError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::NotFound(s) | Self::Ambiguous(s) | Self::Hidden(s) => f.write_str(s), + } + } +} + +/// Shared contact/session identifier resolver used by `query`, `serve`, and `export`. +/// +/// Precedence: +/// 1. Unambiguous special identifiers: `wxid_*`, `*@chatroom`, `filehelper` +/// 2. Exact username match in ContactResolver +/// 3. Conservative legacy-ID short-circuit for account-like ASCII identifiers +/// (requires digit; prevents fuzzy mis-resolution of bare account IDs) +/// 4. Fuzzy match via ContactResolver (`find_candidates`) +/// 5. DB keyword fallback +/// +/// When `visibility` is provided and `show_hidden` is false, hidden talkers +/// are rejected with `ContactResolveError::Hidden` without leaking existence. +pub fn resolve_contact_id( + contact: &str, + resolver: &ContactResolver, + db: &wx_db::WechatDb, + visibility: Option<&VisibilityIndex>, + show_hidden: bool, +) -> Result { + // 0. Empty / whitespace-only guard + if contact.trim().is_empty() { + return Err(ContactResolveError::NotFound( + "empty contact identifier".to_string(), + )); + } + + // 1. Unambiguous special identifiers + if contact.starts_with("wxid_") || is_group_chat(contact) || contact == "filehelper" { + let resolved = ResolvedContact { + wxid: contact.to_string(), + display_name: None, + }; + return check_visibility(resolved, visibility, show_hidden); + } + + // 2. Exact username match (covers legacy IDs like "testuser001" that exist as contact usernames) + if let Some(display) = resolver.resolve(contact) { + let resolved = ResolvedContact { + wxid: contact.to_string(), + display_name: Some(display.to_string()), + }; + return check_visibility(resolved, visibility, show_hidden); + } + + // 3. Conservative legacy-ID short-circuit: ASCII alphanumeric + underscore with at least + // one digit (e.g. "testuser001"). This prevents fuzzy matching from mis-resolving a bare + // account ID to an unrelated contact whose display name happens to substring-match. + // Pure-alpha names like "Alice" and hyphenated names like "team-alpha" are NOT matched. + if is_account_like(contact) { + let resolved = ResolvedContact { + wxid: contact.to_string(), + display_name: None, + }; + return check_visibility(resolved, visibility, show_hidden); + } + + // 4. Fuzzy match via ContactResolver + let candidates = resolver.find_candidates(contact); + match candidates.len() { + 1 => { + let (name, wxid) = candidates[0]; + let resolved = ResolvedContact { + wxid: wxid.to_string(), + display_name: Some(name.to_string()), + }; + return check_visibility(resolved, visibility, show_hidden); + } + n if n > 1 => { + // Filter out hidden candidates before reporting ambiguity + let visible: Vec<_> = if let Some(vis) = visibility { + candidates + .into_iter() + .filter(|(_, wxid)| show_hidden || !vis.is_hidden_talker(wxid)) + .collect() + } else { + candidates + }; + match visible.len() { + 0 => { + return Err(ContactResolveError::NotFound(format!( + "contact \"{contact}\" not found" + ))); + } + 1 => { + let (name, wxid) = visible[0]; + return Ok(ResolvedContact { + wxid: wxid.to_string(), + display_name: Some(name.to_string()), + }); + } + _ => { + return Err(ContactResolveError::Ambiguous(format_ambiguous( + contact, + &visible + .iter() + .map(|(name, wxid)| (name.to_string(), wxid.to_string())) + .collect::>(), + ))); + } + } + } + _ => {} + } + + // 5. DB keyword fallback + let result = db + .query_contacts(&wx_db::ContactQuery::new().keyword(contact).limit(10)) + .map_err(|e| ContactResolveError::NotFound(format!("db error: {e}")))?; + + // Filter hidden contacts from DB results + let visible_items: Vec<_> = if let Some(vis) = visibility { + result + .items + .into_iter() + .filter(|c| show_hidden || !vis.is_hidden_talker(&c.user_name)) + .collect() + } else { + result.items + }; + + match visible_items.len() { + 1 => { + let c = &visible_items[0]; + let name = if !c.remark.is_empty() { + &c.remark + } else if !c.nick_name.is_empty() { + &c.nick_name + } else { + &c.user_name + }; + Ok(ResolvedContact { + wxid: c.user_name.clone(), + display_name: Some(name.to_string()), + }) + } + 0 => Err(ContactResolveError::NotFound(format!( + "contact \"{contact}\" not found" + ))), + _ => Err(ContactResolveError::Ambiguous(format_ambiguous( + contact, + &visible_items + .iter() + .map(|c| { + let name = if !c.remark.is_empty() { + c.remark.clone() + } else { + c.nick_name.clone() + }; + (name, c.user_name.clone()) + }) + .collect::>(), + ))), + } +} + +/// Check if a resolved contact is hidden. Returns the contact if visible, +/// or a generic "not found" error that does not leak hidden status. +fn check_visibility( + resolved: ResolvedContact, + visibility: Option<&VisibilityIndex>, + show_hidden: bool, +) -> Result { + if let Some(vis) = visibility { + if !show_hidden && vis.is_hidden_talker(&resolved.wxid) { + return Err(ContactResolveError::Hidden(format!( + "contact \"{}\" not found", + resolved.wxid + ))); + } + } + Ok(resolved) +} + +/// Check if a string looks like a WeChat account-like identifier. +/// +/// Conservative: requires pure ASCII alphanumeric + underscore, AND at least one digit. +/// The digit requirement avoids treating ASCII display names like "Alice" or "Bob" +/// as direct identifiers. Real account IDs virtually always contain digits +/// (e.g. `testuser001`, `user123456`). +fn is_account_like(s: &str) -> bool { + !s.is_empty() + && s.is_ascii() + && s.chars().all(|c| c.is_ascii_alphanumeric() || c == '_') + && s.chars().any(|c| c.is_ascii_digit()) +} + +fn format_ambiguous(contact: &str, items: &[(String, String)]) -> String { + let mut msg = format!( + "ambiguous contact \"{contact}\": {} matches found. Use wxid directly:\n", + items.len() + ); + for (name, wxid) in items { + msg.push_str(&format!(" {name}({wxid})\n")); + } + msg +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn is_account_like_accepts_wxid_pattern() { + assert!(is_account_like("wxid_example123abc")); + } + + #[test] + fn is_account_like_accepts_legacy_id() { + assert!(is_account_like("testuser001")); + } + + #[test] + fn is_account_like_accepts_underscore_with_digits() { + assert!(is_account_like("user_name_123")); + } + + #[test] + fn is_account_like_rejects_pure_alpha() { + // "Alice", "Bob", "team" should NOT bypass fuzzy match + assert!(!is_account_like("Alice")); + assert!(!is_account_like("Bob")); + } + + #[test] + fn is_account_like_rejects_hyphen() { + // Hyphens not in real account IDs + assert!(!is_account_like("team-alpha")); + } + + #[test] + fn is_account_like_rejects_chinese() { + assert!(!is_account_like("张三")); + } + + #[test] + fn is_account_like_rejects_spaces() { + assert!(!is_account_like("John Doe")); + } + + #[test] + fn is_account_like_rejects_empty() { + assert!(!is_account_like("")); + } + + #[test] + fn empty_contact_returns_error() { + let resolver = ContactResolver::empty(); + let db = { + let tmp = tempfile::tempdir().unwrap(); + create_minimal_db(tmp.path()) + }; + + let err = resolve_contact_id("", &resolver, &db, None, false).unwrap_err(); + assert!(err.to_string().contains("empty contact identifier")); + + let err = resolve_contact_id(" ", &resolver, &db, None, false).unwrap_err(); + assert!(err.to_string().contains("empty contact identifier")); + } + + #[test] + fn direct_identifiers_pass_through() { + let resolver = ContactResolver::empty(); + let db = { + let tmp = tempfile::tempdir().unwrap(); + create_minimal_db(tmp.path()) + }; + + let r = resolve_contact_id("wxid_abc", &resolver, &db, None, false).unwrap(); + assert_eq!(r.wxid, "wxid_abc"); + + let r = resolve_contact_id("group@chatroom", &resolver, &db, None, false).unwrap(); + assert_eq!(r.wxid, "group@chatroom"); + + let r = resolve_contact_id("filehelper", &resolver, &db, None, false).unwrap(); + assert_eq!(r.wxid, "filehelper"); + } + + #[test] + fn legacy_ascii_id_passes_through() { + let resolver = ContactResolver::empty(); + let db = { + let tmp = tempfile::tempdir().unwrap(); + create_minimal_db(tmp.path()) + }; + + let r = resolve_contact_id("testuser001", &resolver, &db, None, false).unwrap(); + assert_eq!(r.wxid, "testuser001"); + } + + #[test] + fn exact_username_match_keeps_display_name_before_legacy_short_circuit() { + let tmp = tempfile::tempdir().unwrap(); + let db = create_db_with_contacts( + tmp.path(), + &[ + ("testuser001", "", "小明", "xming"), + ("wxid_other", "", "其他人", ""), + ], + ); + let resolver = ContactResolver::build(&db).unwrap(); + + let r = resolve_contact_id("testuser001", &resolver, &db, None, false).unwrap(); + assert_eq!(r.wxid, "testuser001"); + assert_eq!(r.display_name.as_deref(), Some("小明")); + } + + #[test] + fn legacy_ascii_id_bypasses_fuzzy_resolution() { + let tmp = tempfile::tempdir().unwrap(); + let db = create_db_with_contacts( + tmp.path(), + &[ + ("wxid_friend", "testuser001 的同学", "", ""), + ("wxid_other", "其他人", "", ""), + ], + ); + let resolver = ContactResolver::build(&db).unwrap(); + + let r = resolve_contact_id("testuser001", &resolver, &db, None, false).unwrap(); + assert_eq!(r.wxid, "testuser001"); + assert_eq!(r.display_name, None); + } + + #[test] + fn chinese_name_not_treated_as_direct_id() { + let resolver = ContactResolver::empty(); + let db = { + let tmp = tempfile::tempdir().unwrap(); + create_minimal_db(tmp.path()) + }; + + // Chinese name should fall through to fuzzy/db lookup and fail as "not found" + let result = resolve_contact_id("张三", &resolver, &db, None, false); + assert!(result.is_err()); + } + + #[test] + fn hidden_talker_returns_not_found() { + let resolver = ContactResolver::empty(); + let db = { + let tmp = tempfile::tempdir().unwrap(); + create_minimal_db(tmp.path()) + }; + let vis = VisibilityIndex::build(&["wxid_hidden".to_string()], &[], &resolver); + + let result = resolve_contact_id("wxid_hidden", &resolver, &db, Some(&vis), false); + assert!(result.is_err()); + // Error should look like "not found", not leak "hidden" + let err_msg = result.unwrap_err().to_string(); + assert!(err_msg.contains("not found")); + } + + #[test] + fn show_hidden_bypasses_visibility() { + let resolver = ContactResolver::empty(); + let db = { + let tmp = tempfile::tempdir().unwrap(); + create_minimal_db(tmp.path()) + }; + let vis = VisibilityIndex::build(&["wxid_hidden".to_string()], &[], &resolver); + + let r = resolve_contact_id("wxid_hidden", &resolver, &db, Some(&vis), true).unwrap(); + assert_eq!(r.wxid, "wxid_hidden"); + } + + #[test] + fn hidden_talker_hidden_error_does_not_leak() { + let resolver = ContactResolver::empty(); + let db = { + let tmp = tempfile::tempdir().unwrap(); + create_minimal_db(tmp.path()) + }; + let vis = VisibilityIndex::build(&["wxid_secret".to_string()], &[], &resolver); + + let result = resolve_contact_id("wxid_secret", &resolver, &db, Some(&vis), false); + let err = result.unwrap_err(); + // The error message should not contain the word "hidden" + let msg = err.to_string(); + assert!(!msg.contains("hidden"), "error leaked visibility: {msg}"); + assert!(msg.contains("not found")); + } + + fn create_minimal_db(dir: &std::path::Path) -> wx_db::WechatDb { + let msg_dir = dir.join("message"); + std::fs::create_dir_all(&msg_dir).unwrap(); + let contact_dir = dir.join("contact"); + std::fs::create_dir_all(&contact_dir).unwrap(); + let session_dir = dir.join("session"); + std::fs::create_dir_all(&session_dir).unwrap(); + + let msg_conn = rusqlite::Connection::open(msg_dir.join("message_0.db")).unwrap(); + msg_conn + .execute_batch( + "CREATE TABLE msg_0 ( + localId INTEGER, mesSvrID INTEGER, msgCreateTime INTEGER, msgType INTEGER, + msgSubType INTEGER, msgSeq INTEGER, msgContent TEXT, msgSource TEXT, + msgStatus INTEGER, compressContent BLOB, mesMsgSender TEXT + );", + ) + .unwrap(); + + let contact_conn = rusqlite::Connection::open(contact_dir.join("contact.db")).unwrap(); + contact_conn + .execute_batch( + "CREATE TABLE contact ( + username TEXT PRIMARY KEY, alias TEXT DEFAULT '', remark TEXT DEFAULT '', + nick_name TEXT DEFAULT '', description TEXT DEFAULT NULL, + extra_buffer BLOB DEFAULT NULL + ); + CREATE TABLE contact_label ( + label_id_ TEXT, label_name_ TEXT, sort_order_ INTEGER + );", + ) + .unwrap(); + + let session_conn = rusqlite::Connection::open(session_dir.join("session.db")).unwrap(); + session_conn + .execute_batch( + "CREATE TABLE SessionTable ( + userName TEXT, summary TEXT, sortTimestamp INTEGER, + lastMsgType INTEGER, lastMsgSender TEXT, lastSenderDisplayName TEXT + );", + ) + .unwrap(); + + wx_db::WechatDb::open(dir).unwrap() + } + + fn create_db_with_contacts( + dir: &std::path::Path, + contacts: &[(&str, &str, &str, &str)], + ) -> wx_db::WechatDb { + let db = create_minimal_db(dir); + let contact_db = dir.join("contact").join("contact.db"); + let conn = rusqlite::Connection::open(contact_db).unwrap(); + for (username, remark, nick_name, alias) in contacts { + conn.execute( + "INSERT INTO contact (username, remark, nick_name, alias) VALUES (?1, ?2, ?3, ?4)", + rusqlite::params![username, remark, nick_name, alias], + ) + .unwrap(); + } + db + } +} diff --git a/crates/wx-cli/src/main.rs b/crates/wx-cli/src/main.rs new file mode 100644 index 0000000..9246400 --- /dev/null +++ b/crates/wx-cli/src/main.rs @@ -0,0 +1,770 @@ +use std::path::PathBuf; + +use clap::{Args, CommandFactory, FromArgMatches, Parser, Subcommand, ValueEnum}; + +use crate::cmd::server::ServerAction; +use crate::cmd::thin_client::ThinClientCliArgs; + +mod cmd; +pub(crate) mod contact_id; +mod output; +pub(crate) mod schema; +pub(crate) mod settings; +mod util; +mod version; +pub(crate) mod visibility_projection; + +#[derive(Parser)] +#[command( + name = "wx-cli", + about = "WeChat database decryption tool (macOS 4.1.x)" +)] +struct Cli { + #[command(subcommand)] + command: Commands, +} + +#[derive(Args, Clone, Debug, Default)] +struct ServerRoutingArgs { + /// Reuse a running server instance at this base URL (default: http://127.0.0.1:9100) + #[arg(long)] + server_url: Option, + + /// Bearer token for the running server instance + #[arg(long)] + server_token: Option, + + /// Only use the running server instance; do not fall back to local queries + #[arg(long, conflicts_with = "no_server")] + server_only: bool, + + /// Force local direct-query mode even if a server instance is available + #[arg(long, conflicts_with = "server_only")] + no_server: bool, +} + +impl From for ThinClientCliArgs { + fn from(value: ServerRoutingArgs) -> Self { + Self { + server_url: value.server_url, + server_token: value.server_token, + server_only: value.server_only, + no_server: value.no_server, + } + } +} + +#[derive(Subcommand)] +enum Commands { + /// Manage encryption keys + Key { + #[command(subcommand)] + action: KeyAction, + }, + /// Decrypt WeChat databases + Decrypt { + /// 32-byte hex key (overrides KeyStore lookup) + #[arg(short, long)] + key: Option, + + /// WeChat data directory (auto-detect if omitted) + #[arg(short, long)] + data_dir: Option, + + /// Account directory name or base account ID + #[arg(long)] + account: Option, + + /// Output directory (default: cache dir) + #[arg(short, long)] + output: Option, + + /// Only re-decrypt files modified since last run + #[arg(long)] + incremental: bool, + }, + /// Show info about a database file + Info { + /// Path to a WeChat database file + db_file: PathBuf, + }, + /// Media operations (image decrypt, hardlink query, voice extract) + Media { + #[command(subcommand)] + action: MediaAction, + }, + /// Watch for real-time session changes in an encrypted WeChat database + Watch { + /// 32-byte hex key (overrides KeyStore lookup) + #[arg(short, long)] + key: Option, + + /// WeChat data directory (auto-detect if omitted) + #[arg(short, long)] + data_dir: Option, + + /// Account directory name or base account ID + #[arg(long)] + account: Option, + + /// Force mtime polling instead of fsnotify + #[arg(long, conflicts_with = "fsnotify")] + poll: bool, + + /// Force fsnotify backend (opt-in on macOS, where polling is the default) + #[arg(long, conflicts_with = "poll")] + fsnotify: bool, + + /// Polling interval in milliseconds + #[arg(long, default_value = "2000")] + poll_ms: u64, + + /// Output format + #[arg(long, default_value = "text", value_enum)] + format: OutputFormat, + + /// Show hidden contacts (bypass contact hiding rules) + #[arg(long)] + show_hidden: bool, + }, + /// Manage the long-running HTTP API service + Server { + #[command(subcommand)] + action: ServerAction, + }, + /// Query recent sessions (conversations) + Sessions { + /// WeChat data directory (auto-detect if omitted) + #[arg(short, long)] + data_dir: Option, + /// Account directory name or base account ID + #[arg(long)] + account: Option, + /// 32-byte hex key (overrides KeyStore lookup) + #[arg(short, long)] + key: Option, + /// Maximum number of results (overridden by --all) + #[arg(long, default_value = "20")] + limit: usize, + /// Pagination offset + #[arg(long, default_value = "0")] + offset: usize, + /// Sort order + #[arg(long, default_value = "desc", value_enum)] + order: SortOrderArg, + /// Return up to 20,000 results (overrides --limit) + #[arg(long)] + all: bool, + /// Output format + #[arg(long, default_value = "text", value_enum)] + format: OutputFormat, + /// Show hidden contacts (bypass contact hiding rules) + #[arg(long)] + show_hidden: bool, + + #[command(flatten)] + server: ServerRoutingArgs, + }, + /// Search contacts by name, wxid, phone, labels, or other fields + Contacts { + /// WeChat data directory (auto-detect if omitted) + #[arg(short, long)] + data_dir: Option, + /// Account directory name or base account ID + #[arg(long)] + account: Option, + /// 32-byte hex key (overrides KeyStore lookup) + #[arg(short, long)] + key: Option, + /// Search keyword (matches name, wxid, phone, labels, memo, signature, region) + #[arg(long)] + search: Option, + /// Maximum number of results (overridden by --all) + #[arg(long, default_value = "50")] + limit: usize, + /// Pagination offset + #[arg(long, default_value = "0")] + offset: usize, + /// Return up to 20,000 results (overrides --limit) + #[arg(long)] + all: bool, + /// Output format + #[arg(long, default_value = "text", value_enum)] + format: OutputFormat, + /// Show hidden contacts (bypass contact hiding rules) + #[arg(long)] + show_hidden: bool, + + #[command(flatten)] + server: ServerRoutingArgs, + }, + /// Query chat messages for a contact or group + Query { + /// Contact name, wxid, or chatroom ID + contact: String, + /// WeChat data directory (auto-detect if omitted) + #[arg(short, long)] + data_dir: Option, + /// Account directory name or base account ID + #[arg(long)] + account: Option, + /// 32-byte hex key (overrides KeyStore lookup) + #[arg(short, long)] + key: Option, + /// Start time (Unix seconds) + #[arg(long, conflicts_with_all = ["around_sort_seq", "around_server_id", "after_sort_seq"])] + since: Option, + /// End time (Unix seconds) + #[arg(long, conflicts_with_all = ["around_sort_seq", "around_server_id", "after_sort_seq"])] + until: Option, + /// Message type filter (text/image/voice/video/emoji/app/system/revoke or numeric, e.g. 49) + #[arg(long = "type")] + msg_type: Option, + /// Maximum number of results (overridden by --all) + #[arg(long, default_value = "50")] + limit: usize, + /// Pagination offset + #[arg(long, default_value = "0")] + offset: usize, + /// Sort order + #[arg(long, default_value = "desc", value_enum)] + order: SortOrderArg, + /// Return up to 20,000 results (overrides --limit) + #[arg(long, conflicts_with_all = ["around_sort_seq", "around_server_id", "after_sort_seq"])] + all: bool, + /// Output format + #[arg(long, default_value = "text", value_enum)] + format: OutputFormat, + /// Show N messages before and after this sort_seq + #[arg(long, conflicts_with_all = ["since", "until", "all", "around_server_id", "after_sort_seq"])] + around_sort_seq: Option, + /// Show N messages before and after the message with this server_id + #[arg(long, conflicts_with_all = ["since", "until", "all", "around_sort_seq", "after_sort_seq"])] + around_server_id: Option, + /// Context window size (messages before/after --around-*, default 50) + #[arg(long)] + context: Option, + /// Return messages after this sort_seq (incremental pull, ASC order) + #[arg(long, conflicts_with_all = ["since", "until", "all", "around_sort_seq", "around_server_id"])] + after_sort_seq: Option, + /// Show hidden contacts (bypass contact hiding rules) + #[arg(long)] + show_hidden: bool, + + #[command(flatten)] + server: ServerRoutingArgs, + }, + /// Export a conversation to TXT or JSON with media files + Export { + /// Contact name, wxid, or chatroom ID + contact: String, + /// Output directory + #[arg(short, long)] + output: PathBuf, + /// WeChat data directory (auto-detect if omitted) + #[arg(short, long)] + data_dir: Option, + /// Account directory name or base account ID + #[arg(long)] + account: Option, + /// 32-byte hex key (overrides KeyStore lookup) + #[arg(short, long)] + key: Option, + /// Start time filter (Unix seconds) + #[arg(long)] + since: Option, + /// End time filter (Unix seconds) + #[arg(long)] + until: Option, + /// Maximum number of results (overridden by --all) + #[arg(long, default_value = "50")] + limit: usize, + /// Pagination offset + #[arg(long, default_value = "0")] + offset: usize, + /// Sort order (default: asc for chronological export) + #[arg(long, default_value = "asc", value_enum)] + order: SortOrderArg, + /// Export all messages (paged internally; overrides --limit) + #[arg(long)] + all: bool, + /// Output format + #[arg(long, default_value = "txt", value_enum)] + format: ExportFormat, + /// Skip media file export (text-only, faster) + #[arg(long)] + no_media: bool, + /// Max parallel threads for media resolve (default: min(CPU, 4), 1 = serial) + #[arg(long)] + parallel: Option, + /// Show emoji/sticker detail instead of [动画表情] + #[arg(long)] + show_emoji: bool, + /// Show hidden contacts (bypass contact hiding rules) + #[arg(long)] + show_hidden: bool, + }, + /// Full-text search across all conversations + Search { + /// Search keyword + keyword: String, + /// WeChat data directory (auto-detect if omitted) + #[arg(short, long)] + data_dir: Option, + /// Account directory name or base account ID + #[arg(long)] + account: Option, + /// 32-byte hex key (overrides KeyStore lookup) + #[arg(short, long)] + key: Option, + /// Maximum number of results (overridden by --all) + #[arg(long, default_value = "20")] + limit: usize, + /// Pagination offset + #[arg(long, default_value = "0")] + offset: usize, + /// Return up to 20,000 results (overrides --limit) + #[arg(long)] + all: bool, + /// Output format + #[arg(long, default_value = "text", value_enum)] + format: OutputFormat, + + #[command(flatten)] + server: ServerRoutingArgs, + }, + /// Decrypt .dat image file(s) (shortcut for `media decrypt-dat`) + #[command(name = "decode-image")] + DecodeImage { + /// Input .dat file or directory of .dat files + input: PathBuf, + /// Output file or directory (auto-named if omitted) + #[arg(short, long)] + output: Option, + /// Account directory name for KeyStore V2 key lookup + #[arg(long)] + account: Option, + /// WeChat account data directory for automatic V2 key derivation + #[arg(short, long)] + data_dir: Option, + }, + /// Show WeChat process and account status + Status, + /// Show all managed file paths with existence status + Paths { + /// Output as JSON + #[arg(long)] + json: bool, + }, + /// Check prerequisites for key extraction + Doctor { + /// Output fix commands for failing checks + #[arg(long)] + fix: bool, + }, + /// Query decrypted WeChat databases (dev/debug tool) + #[command(name = "db-dev", hide = true)] + DbDev { + /// Path to decrypted db_storage directory (containing contact/, session/, message/) + #[arg(long)] + path: PathBuf, + + #[command(subcommand)] + action: DbDevAction, + }, +} + +#[derive(Subcommand)] +pub enum DbDevAction { + /// Query contacts + Contacts { + /// Search keyword (matches userName, alias, remark, nickName, description, phone, labels, signature, region) + #[arg(long)] + keyword: Option, + /// Maximum number of results (default: 1000, max: 20000) + #[arg(long, default_value = "0")] + limit: usize, + /// Pagination offset + #[arg(long, default_value = "0")] + offset: usize, + }, + /// Query recent sessions (conversations) + Sessions { + /// Maximum number of results + #[arg(long, default_value = "0")] + limit: usize, + /// Pagination offset + #[arg(long, default_value = "0")] + offset: usize, + }, + /// Query messages for a specific conversation + Messages { + /// Talker wxid or chatroom ID + #[arg(long)] + talker: String, + /// Start time filter (Unix seconds, inclusive) + #[arg(long)] + start: Option, + /// End time filter (Unix seconds, inclusive) + #[arg(long)] + end: Option, + /// Content keyword filter (case-insensitive) + #[arg(long)] + keyword: Option, + /// Maximum number of results + #[arg(long, default_value = "0")] + limit: usize, + /// Pagination offset + #[arg(long, default_value = "0")] + offset: usize, + }, + /// Query chatrooms (group chats) + Chatrooms { + /// Filter by specific chatroom username + #[arg(long)] + username: Option, + /// Maximum number of results + #[arg(long, default_value = "0")] + limit: usize, + /// Pagination offset + #[arg(long, default_value = "0")] + offset: usize, + }, +} + +#[derive(Subcommand)] +pub enum MediaAction { + /// Decrypt .dat image file(s) (XOR/V1/V2). Accepts a file or directory. + #[command(name = "decrypt-dat")] + DecryptDat { + /// Input .dat file or directory of .dat files + input: PathBuf, + /// Output file or directory (auto-named if omitted) + #[arg(short, long)] + output: Option, + /// V2 AES key (16-byte ASCII string); auto-read from KeyStore if omitted + #[arg(long)] + v2_key: Option, + /// Account directory name for KeyStore V2 key lookup (used when --v2-key is omitted) + #[arg(long)] + account: Option, + /// WeChat account data directory for automatic V2 key derivation (MD5(UIN+WXID)) + #[arg(long)] + data_dir: Option, + /// XOR key for V2 tail (hex byte, e.g. "37"); auto-detected from thumbnails if omitted + #[arg(long)] + xor_key: Option, + }, + /// Resolve media path from hardlink.db + #[command(name = "resolve-path")] + ResolvePath { + /// Path to hardlink.db + #[arg(long)] + db: PathBuf, + /// Media type: image, video, file + #[arg(long, default_value = "image")] + media_type: String, + /// MD5 key or file name prefix + key: String, + }, + /// Extract voice BLOB from media_N.db + #[command(name = "extract-voice")] + ExtractVoice { + /// Path to media/ directory containing media_*.db files + #[arg(long)] + media_dir: PathBuf, + /// Server ID (svr_id) of the voice message + svr_id: String, + /// Output file path + #[arg(short, long)] + output: Option, + /// Output raw SILK instead of transcoding to MP3 + #[arg(long)] + raw: bool, + }, + /// Decrypt WeChat Channels encrypted video using Isaac64 PRNG + #[command(name = "decrypt-video")] + DecryptVideo { + /// Path to encrypted video file + input: PathBuf, + /// Decryption seed (decimal or hex with 0x prefix) + #[arg(long)] + seed: String, + /// Output file path (default: .mp4) + #[arg(short, long)] + output: Option, + }, +} + +#[derive(Clone, Debug, ValueEnum)] +pub enum OutputFormat { + Text, + Json, +} + +#[derive(Clone, ValueEnum)] +pub enum ExportFormat { + Txt, + Json, +} + +#[derive(Clone, ValueEnum)] +pub enum SortOrderArg { + Asc, + Desc, +} + +impl From for wx_db::SortOrder { + fn from(v: SortOrderArg) -> Self { + match v { + SortOrderArg::Asc => wx_db::SortOrder::Asc, + SortOrderArg::Desc => wx_db::SortOrder::Desc, + } + } +} + +#[derive(Subcommand)] +enum KeyAction { + /// Extract key via LLDB (requires SIP disabled, auto-detects current account) + Extract { + /// Timeout in seconds for LLDB capture + #[arg(long, default_value = "120")] + timeout: u64, + }, + /// List stored keys + List, + /// Manually set a key for an account + Set { + /// Account directory name as stored in KeyStore (e.g. wxid_xxx_ab12, testuser001_1662) + account: String, + /// 32-byte hex key + hex_key: String, + }, + /// Manually set V2 image AES key for an account + #[command(name = "set-image")] + SetImage { + /// Account directory name as stored in KeyStore (e.g. wxid_xxx_ab12, testuser001_1662) + account: String, + /// 16-byte image AES key (ASCII string or hex) + image_key: String, + }, + /// Scan WeChat process memory for pre-derived encryption keys (requires SIP disabled + sudo, no restart needed) + Scan, +} + +#[tokio::main] +async fn main() { + tracing_subscriber::fmt() + .with_env_filter( + tracing_subscriber::EnvFilter::try_from_default_env() + .unwrap_or_else(|_| tracing_subscriber::EnvFilter::new("warn")), + ) + .with_writer(std::io::stderr) + .init(); + + let matches = Cli::command() + .version(version::cli_version_long()) + .long_version(version::cli_version_long()) + .get_matches(); + let cli = Cli::from_arg_matches(&matches).unwrap(); + + let result = match cli.command { + Commands::Key { action } => match action { + KeyAction::Extract { timeout: t } => cmd::key::cmd_key_extract(t).await, + KeyAction::List => cmd::key::cmd_key_list(), + KeyAction::Set { account, hex_key } => cmd::key::cmd_key_set(&account, &hex_key), + KeyAction::SetImage { account, image_key } => { + cmd::key::cmd_key_set_image(&account, &image_key) + } + KeyAction::Scan => cmd::key::cmd_key_scan(), + }, + Commands::Media { action } => cmd::media::cmd_media(action), + Commands::Watch { + key, + data_dir, + account, + poll, + fsnotify, + poll_ms, + format, + show_hidden, + } => { + cmd::watch::cmd_watch( + key, + data_dir, + account, + poll, + fsnotify, + poll_ms, + format, + show_hidden, + ) + .await + } + Commands::Server { action } => cmd::server::cmd_server(action).await, + Commands::Decrypt { + key, + data_dir, + account, + output, + incremental, + } => cmd::decrypt::cmd_decrypt(key, data_dir, account, output, incremental), + Commands::Info { db_file } => cmd::info::cmd_info(&db_file), + Commands::Sessions { + data_dir, + account, + key, + limit, + offset, + order, + all, + format, + show_hidden, + server, + } => cmd::sessions::cmd_sessions( + data_dir, + account, + key, + limit, + offset, + order, + all, + format, + show_hidden, + server.into(), + ), + Commands::Contacts { + data_dir, + account, + key, + search, + limit, + offset, + all, + format, + show_hidden, + server, + } => cmd::contacts::cmd_contacts( + data_dir, + account, + key, + search, + limit, + offset, + all, + format, + show_hidden, + server.into(), + ), + Commands::Query { + contact, + data_dir, + account, + key, + since, + until, + msg_type, + limit, + offset, + order, + all, + format, + around_sort_seq, + around_server_id, + context, + after_sort_seq, + show_hidden, + server, + } => cmd::query::cmd_query( + &contact, + data_dir, + account, + key, + since, + until, + msg_type, + limit, + offset, + order, + all, + format, + around_sort_seq, + around_server_id, + context, + after_sort_seq, + show_hidden, + server.into(), + ), + Commands::Export { + contact, + output, + data_dir, + account, + key, + since, + until, + limit, + offset, + order, + all, + format, + no_media, + show_emoji, + show_hidden, + parallel, + } => cmd::export::cmd_export( + &contact, + output, + data_dir, + account, + key, + since, + until, + limit, + offset, + order, + all, + format, + no_media, + show_emoji, + show_hidden, + parallel, + ), + Commands::Search { + keyword, + data_dir, + account, + key, + limit, + offset, + all, + format, + server, + } => cmd::search::cmd_search( + &keyword, + data_dir, + account, + key, + limit, + offset, + all, + format, + server.into(), + ), + Commands::DecodeImage { + input, + output, + account, + data_dir, + } => cmd::decode_image::cmd_decode_image(input, output, account, data_dir), + Commands::Status => cmd::status::cmd_status(), + Commands::Paths { json } => cmd::paths::cmd_paths(json), + Commands::Doctor { fix } => cmd::doctor::cmd_doctor(fix), + Commands::DbDev { path, action } => cmd::db_dev::cmd_db_dev(&path, action), + }; + + if let Err(e) = result { + eprintln!("error: {e}"); + std::process::exit(1); + } +} diff --git a/crates/wx-cli/src/output.rs b/crates/wx-cli/src/output.rs new file mode 100644 index 0000000..72a607f --- /dev/null +++ b/crates/wx-cli/src/output.rs @@ -0,0 +1,120 @@ +use serde::{Deserialize, Serialize}; + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub struct JsonEnvelope { + pub items: Vec, + pub paging: PagingMeta, + pub stats: StatsMeta, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub struct PagingMeta { + pub limit: usize, + pub offset: usize, + pub returned: usize, + pub has_more: bool, + pub total: usize, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub struct StatsMeta { + pub scanned: usize, + pub skipped: usize, + #[serde(skip_serializing_if = "Option::is_none")] + pub elapsed_ms: Option, + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub shard_warnings: Vec, +} + +impl JsonEnvelope { + pub fn from_query_result( + result: wx_db::QueryResult, + limit: usize, + offset: usize, + map_item: impl FnMut(U) -> T, + ) -> Self { + let returned = result.items.len(); + // total: use filtered_count if available (messages), otherwise total_rows + let total = result + .stats + .filtered_count + .unwrap_or(result.stats.total_rows); + let has_more = offset + returned < total; + Self { + items: result.items.into_iter().map(map_item).collect(), + paging: PagingMeta { + limit, + offset, + returned, + has_more, + total, + }, + stats: StatsMeta { + scanned: result.stats.total_rows, + skipped: result.stats.skipped, + elapsed_ms: None, + shard_warnings: Vec::new(), + }, + } + } + + pub fn from_message_query_result( + result: wx_db::MessageQueryResult, + limit: usize, + offset: usize, + map_item: impl FnMut(wx_db::Message) -> T, + ) -> Self { + let returned = result.items.len(); + let total = result + .stats + .filtered_count + .unwrap_or(result.stats.total_rows); + let has_more = offset + returned < total; + let shard_warnings = result.shard_warnings; + Self { + items: result.items.into_iter().map(map_item).collect(), + paging: PagingMeta { + limit, + offset, + returned, + has_more, + total, + }, + stats: StatsMeta { + scanned: result.stats.total_rows, + skipped: result.stats.skipped, + elapsed_ms: None, + shard_warnings, + }, + } + } + + /// Kept for Task 6: remove self-built FTS index code. + #[allow(dead_code)] + pub fn from_fts_result( + items: Vec, + total: usize, + limit: usize, + offset: usize, + map_item: impl FnMut(U) -> T, + ) -> Self { + let returned = items.len(); + let has_more = offset + returned < total; + Self { + items: items.into_iter().map(map_item).collect(), + paging: PagingMeta { + limit, + offset, + returned, + has_more, + total, + }, + stats: StatsMeta { + scanned: total, + skipped: 0, + elapsed_ms: None, + shard_warnings: Vec::new(), + }, + } + } +} diff --git a/crates/wx-cli/src/schema.rs b/crates/wx-cli/src/schema.rs new file mode 100644 index 0000000..7a06b05 --- /dev/null +++ b/crates/wx-cli/src/schema.rs @@ -0,0 +1,765 @@ +use serde::{Deserialize, Serialize}; +use wx_context::{ContactResolver, Direction, VisibilityIndex}; +use wx_db::{extract_quote_fromusr, is_group_chat, FtsHit, Message, NativeFtsHit, Session}; + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub struct EnrichedMessage { + #[serde(flatten)] + pub message: Message, + pub sender_display_name: String, + pub direction: Direction, + pub snippet: String, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub struct EnrichedSession { + #[serde(flatten)] + pub session: Session, + pub display_name: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub direction: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub detected_at: Option, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub struct SearchHit { + pub server_id: i64, + pub talker: String, + pub talker_display_name: String, + pub sender: String, + pub sender_display_name: String, + pub direction: Direction, + pub create_time: i64, + pub sort_seq: i64, + pub msg_type: u32, + pub sub_type: u32, + pub snippet: String, + pub hit_type: String, +} + +pub fn enrich_message( + msg: Message, + self_wxid: &str, + resolver: &ContactResolver, +) -> EnrichedMessage { + let direction = Direction::detect(&msg.sender, self_wxid); + let sender_display_name = resolver.display_with_id(&msg.sender); + let snippet = format_content(&msg); + EnrichedMessage { + message: msg, + sender_display_name, + direction, + snippet, + } +} + +pub fn enrich_session( + session: Session, + self_wxid: &str, + resolver: &ContactResolver, + detected_at: Option, +) -> EnrichedSession { + let display_name = resolver.display_with_id(&session.username); + let direction = derive_session_direction(session.last_msg_sender.as_deref(), self_wxid); + EnrichedSession { + session, + display_name, + direction, + detected_at, + } +} + +pub fn enrich_session_event( + ev: wx_monitor::SessionEvent, + self_wxid: &str, + resolver: &ContactResolver, +) -> EnrichedSession { + let display_name = resolver.display_with_id(&ev.username); + let direction = derive_session_direction(ev.last_msg_sender.as_deref(), self_wxid); + let session = Session { + username: ev.username, + summary: ev.summary, + sort_timestamp: ev.sort_timestamp, + last_msg_type: ev.last_msg_type, + last_msg_sender: ev.last_msg_sender, + last_sender_display_name: ev.last_sender_display_name, + }; + EnrichedSession { + session, + display_name, + direction, + detected_at: Some(ev.detected_at), + } +} + +pub fn derive_session_direction( + last_msg_sender: Option<&str>, + self_wxid: &str, +) -> Option { + let sender = last_msg_sender?.trim(); + if sender.is_empty() { + None + } else { + Some(Direction::detect(sender, self_wxid)) + } +} + +/// Kept for Task 6: remove self-built FTS index code. +#[allow(dead_code)] +pub fn enrich_fts_hit(hit: FtsHit, self_wxid: &str, resolver: &ContactResolver) -> SearchHit { + let direction = Direction::detect(&hit.sender, self_wxid); + SearchHit { + server_id: hit.server_id, + talker_display_name: resolver.display_with_id(&hit.talker), + talker: hit.talker, + sender_display_name: resolver.display_with_id(&hit.sender), + sender: hit.sender, + direction, + create_time: hit.create_time, + sort_seq: hit.sort_seq, + msg_type: hit.msg_type, + sub_type: hit.sub_type, + snippet: hit.snippet, + hit_type: "message".to_string(), + } +} + +pub fn enrich_native_fts_hit( + hit: NativeFtsHit, + self_wxid: &str, + resolver: &ContactResolver, +) -> SearchHit { + let direction = Direction::detect(&hit.sender, self_wxid); + SearchHit { + server_id: 0, + talker_display_name: resolver.display_with_id(&hit.talker), + talker: hit.talker, + sender_display_name: resolver.display_with_id(&hit.sender), + sender: hit.sender, + direction, + create_time: hit.create_time, + sort_seq: hit.sort_seq, + msg_type: hit.msg_type, + sub_type: hit.sub_type, + snippet: hit.snippet, + hit_type: "message".to_string(), + } +} + +pub fn enrich_message_as_hit( + msg: Message, + talker: String, + self_wxid: &str, + resolver: &ContactResolver, +) -> SearchHit { + let direction = Direction::detect(&msg.sender, self_wxid); + let snippet = format_content(&msg); + SearchHit { + server_id: msg.server_id, + talker_display_name: resolver.display_with_id(&talker), + talker, + sender_display_name: resolver.display_with_id(&msg.sender), + sender: msg.sender, + direction, + create_time: msg.create_time, + sort_seq: msg.sort_seq, + msg_type: msg.msg_type, + sub_type: msg.sub_type, + snippet, + hit_type: "message".to_string(), + } +} + +// --- Phase 2: sender-level projection helpers --- + +/// Filter a single message by sender visibility, and redact quote references if needed. +/// +/// Returns `None` if the message's sender is hidden in this group. +/// For Quote messages referencing a hidden sender, clears refer_sender/refer_content/raw_xml. +pub fn project_message_item( + mut msg: EnrichedMessage, + talker: &str, + visibility: &VisibilityIndex, +) -> Option { + if visibility.is_hidden_sender_in_group(talker, &msg.message.sender) { + return None; + } + + // Redact quote references to hidden senders + if let wx_db::MessageContent::Quote { + ref raw_xml, + ref mut refer_sender, + ref mut refer_content, + .. + } = msg.message.content + { + if let Some(fromusr) = extract_quote_fromusr(raw_xml) { + if visibility.is_hidden_sender_in_group(talker, &fromusr) { + *refer_sender = None; + *refer_content = None; + // Clear raw_xml and recalculate snippet + if let wx_db::MessageContent::Quote { + ref mut raw_xml, .. + } = msg.message.content + { + *raw_xml = String::new(); + } + msg.snippet = format_content(&msg.message); + } + } + } + + Some(msg) +} + +/// Filter a list of messages by sender visibility. +/// +/// When `show_hidden` is true, returns the original list unchanged. +pub fn project_message_items( + items: Vec, + talker: &str, + visibility: &VisibilityIndex, + show_hidden: bool, +) -> Vec { + if show_hidden { + return items; + } + items + .into_iter() + .filter_map(|msg| project_message_item(msg, talker, visibility)) + .collect() +} + +/// Redact session sender info if the last message sender is hidden. +/// +/// For group chats where the last_msg_sender is a hidden sender: +/// - summary → "[消息已隐藏]" +/// - last_msg_sender → None +/// - last_sender_display_name → None +/// - direction → None +pub fn project_session_sender(session: &mut EnrichedSession, visibility: &VisibilityIndex) { + if !is_group_chat(&session.session.username) { + return; + } + if let Some(ref sender) = session.session.last_msg_sender { + if visibility.is_hidden_sender_in_group(&session.session.username, sender) { + session.session.summary = "[消息已隐藏]".to_string(); + session.session.last_msg_sender = None; + session.session.last_sender_display_name = None; + session.direction = None; + } + } +} + +pub fn format_content(msg: &Message) -> String { + match &msg.content { + wx_db::MessageContent::Text(s) => s.clone(), + wx_db::MessageContent::Image { md5 } => { + format!("[image {}]", md5.as_deref().unwrap_or("")) + } + wx_db::MessageContent::Voice => "[voice]".into(), + wx_db::MessageContent::Video { md5 } => { + format!("[video {}]", md5.as_deref().unwrap_or("")) + } + wx_db::MessageContent::Emoji(s) => format!("[emoji {s}]"), + wx_db::MessageContent::Location(s) => format!("[location {s}]"), + wx_db::MessageContent::Link { title, .. } => { + format!("[链接] {}", title.as_deref().unwrap_or("")) + } + wx_db::MessageContent::File { title, .. } => { + format!("[文件] {}", title.as_deref().unwrap_or("")) + } + wx_db::MessageContent::MiniProgram { title, .. } => { + format!("[小程序] {}", title.as_deref().unwrap_or("")) + } + wx_db::MessageContent::MergedMessages { title, .. } => { + format!("[聊天记录] {}", title.as_deref().unwrap_or("")) + } + wx_db::MessageContent::Quote { + reply_text, + refer_sender, + refer_content, + .. + } => { + let reply = reply_text.as_deref().unwrap_or(""); + match (refer_sender.as_deref(), refer_content.as_deref()) { + (Some(sender), Some(content)) => { + format!("[引用 @{sender}: {content}] {reply}") + } + _ => format!("[引用] {reply}"), + } + } + wx_db::MessageContent::Transfer { amount_desc, .. } => { + format!("[转账] {}", amount_desc.as_deref().unwrap_or("")) + } + wx_db::MessageContent::RedEnvelope { title, .. } => { + format!("[红包] {}", title.as_deref().unwrap_or("")) + } + wx_db::MessageContent::ChannelVideo { title, .. } => { + format!("[视频号] {}", title.as_deref().unwrap_or("")) + } + wx_db::MessageContent::Pat { .. } => "[拍一拍]".into(), + wx_db::MessageContent::AppGeneric { + sub_type, title, .. + } => { + let label = app_sub_type_display_label(*sub_type); + let t = title.as_deref().unwrap_or(""); + format!("[{label}] {t}") + } + wx_db::MessageContent::System(s) => format!("[system] {s}"), + wx_db::MessageContent::Revoke(s) => format!("[revoke] {s}"), + wx_db::MessageContent::Unknown { msg_type, .. } => { + format!("[type={msg_type}]") + } + } +} + +/// Chinese display labels for app sub_types that fall through to `AppGeneric`. +/// +/// Only covers sub_types NOT already handled by dedicated `MessageContent` variants +/// (link, file, mini-program, quote, transfer, etc.). The authoritative full mapping +/// is in `wx_db::msg_sub_type_label` (English). +fn app_sub_type_display_label(sub_type: u32) -> String { + match sub_type { + 1 => "文本分享".into(), + 2 => "图片分享".into(), + 3 => "音频分享".into(), + 7 => "网页应用".into(), + 8 => "GIF表情".into(), + 10 => "位置共享".into(), + 13 => "品牌消息".into(), + 14 => "聊天备份".into(), + 15 => "聊天迁移".into(), + 16 => "卡券".into(), + 17 => "实时位置".into(), + 21 => "小程序推广".into(), + 24 => "笔记".into(), + 35 => "消息历史".into(), + 40 => "视频号转发".into(), + 44 => "直播商品".into(), + 53 => "群聊引用".into(), + 74 => "视频号文件".into(), + 87 => "群公告".into(), + 88 => "群笔记".into(), + 100 => "表情包".into(), + 101 => "广告".into(), + 107 => "微信链接".into(), + 113 => "视频号名片".into(), + 116 => "视频号橱窗".into(), + 117 => "视频号商品".into(), + 124 => "微信礼物".into(), + 2003 => "红包封面".into(), + _ => format!("app type={sub_type}"), + } +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::Value; + + fn make_session(username: &str, last_msg_sender: Option<&str>) -> Session { + Session { + username: username.to_string(), + summary: "hello".to_string(), + sort_timestamp: 1, + last_msg_type: Some(1), + last_msg_sender: last_msg_sender.map(str::to_string), + last_sender_display_name: None, + } + } + + #[test] + fn enrich_session_detects_known_outgoing_direction() { + let resolver = ContactResolver::empty(); + let enriched = enrich_session( + make_session("wxid_friend", Some("wxid_me")), + "wxid_me", + &resolver, + None, + ); + + assert_eq!(enriched.direction, Some(Direction::Outgoing)); + } + + #[test] + fn enrich_session_detects_known_incoming_direction() { + let resolver = ContactResolver::empty(); + let enriched = enrich_session( + make_session("wxid_friend", Some("wxid_friend")), + "wxid_me", + &resolver, + None, + ); + + assert_eq!(enriched.direction, Some(Direction::Incoming)); + } + + #[test] + fn enrich_session_omits_direction_when_last_sender_missing() { + let resolver = ContactResolver::empty(); + let enriched = enrich_session( + make_session("wxid_friend", None), + "wxid_me", + &resolver, + None, + ); + + assert_eq!(enriched.direction, None); + + let json = serde_json::to_value(&enriched).unwrap(); + assert!(json.get("direction").is_none()); + } + + #[test] + fn enrich_session_omits_direction_when_last_sender_empty_for_group() { + let resolver = ContactResolver::empty(); + let enriched = enrich_session( + make_session("group@chatroom", Some("")), + "wxid_me", + &resolver, + None, + ); + + assert_eq!(enriched.direction, None); + + let json = serde_json::to_value(&enriched).unwrap(); + assert!(json.get("direction").is_none()); + } + + #[test] + fn message_level_direction_remains_required_in_json() { + let resolver = ContactResolver::empty(); + let enriched = enrich_message( + Message { + sort_seq: 1, + server_id: 2, + msg_type: 1, + sub_type: 0, + sender: "wxid_me".to_string(), + talker: "wxid_friend".to_string(), + create_time: 3, + content: wx_db::MessageContent::Text("hello".to_string()), + status: 0, + }, + "wxid_me", + &resolver, + ); + + let json = serde_json::to_value(&enriched).unwrap(); + assert_eq!( + json.get("direction"), + Some(&Value::String("outgoing".to_string())) + ); + } + + #[test] + fn legacy_non_wxid_self_id_outgoing_direction() { + let resolver = ContactResolver::empty(); + let enriched = enrich_message( + Message { + sort_seq: 1, + server_id: 2, + msg_type: 1, + sub_type: 0, + sender: "testuser001".to_string(), + talker: "wxid_friend".to_string(), + create_time: 3, + content: wx_db::MessageContent::Text("hello".to_string()), + status: 0, + }, + "testuser001", + &resolver, + ); + + assert_eq!(enriched.direction, Direction::Outgoing); + } + + #[test] + fn legacy_non_wxid_self_id_incoming_direction() { + let resolver = ContactResolver::empty(); + let enriched = enrich_message( + Message { + sort_seq: 1, + server_id: 2, + msg_type: 1, + sub_type: 0, + sender: "wxid_friend".to_string(), + talker: "wxid_friend".to_string(), + create_time: 3, + content: wx_db::MessageContent::Text("hello".to_string()), + status: 0, + }, + "testuser001", + &resolver, + ); + + assert_eq!(enriched.direction, Direction::Incoming); + } + + #[test] + fn legacy_non_wxid_session_outgoing_direction() { + let resolver = ContactResolver::empty(); + let enriched = enrich_session( + make_session("wxid_friend", Some("testuser001")), + "testuser001", + &resolver, + None, + ); + + assert_eq!(enriched.direction, Some(Direction::Outgoing)); + } + + #[test] + fn legacy_non_wxid_session_incoming_direction() { + let resolver = ContactResolver::empty(); + let enriched = enrich_session( + make_session("wxid_friend", Some("wxid_friend")), + "testuser001", + &resolver, + None, + ); + + assert_eq!(enriched.direction, Some(Direction::Incoming)); + } + + // --- Phase 2: projection helper tests --- + + fn make_message(sender: &str, talker: &str, content: wx_db::MessageContent) -> Message { + Message { + sort_seq: 1, + server_id: 1, + msg_type: 1, + sub_type: 0, + sender: sender.to_string(), + talker: talker.to_string(), + create_time: 1000, + content, + status: 0, + } + } + + fn make_enriched(sender: &str, talker: &str, content: wx_db::MessageContent) -> EnrichedMessage { + let msg = make_message(sender, talker, content); + let snippet = format_content(&msg); + EnrichedMessage { + message: msg, + sender_display_name: sender.to_string(), + direction: Direction::Incoming, + snippet, + } + } + + fn vis_with_hidden_persons(persons: &[&str]) -> VisibilityIndex { + VisibilityIndex::build( + &persons.iter().map(|s| s.to_string()).collect::>(), + &[], + &ContactResolver::empty(), + ) + } + + #[test] + fn project_message_item_non_group_does_not_filter() { + let vis = vis_with_hidden_persons(&["wxid_spam"]); + let msg = make_enriched("wxid_spam", "wxid_spam", wx_db::MessageContent::Text("hi".into())); + assert!(project_message_item(msg, "wxid_spam", &vis).is_some()); + } + + #[test] + fn project_message_item_group_hidden_sender_filtered() { + let vis = vis_with_hidden_persons(&["wxid_spam"]); + let msg = make_enriched("wxid_spam", "group@chatroom", wx_db::MessageContent::Text("spam".into())); + assert!(project_message_item(msg, "group@chatroom", &vis).is_none()); + } + + #[test] + fn project_message_item_group_visible_sender_kept() { + let vis = vis_with_hidden_persons(&["wxid_spam"]); + let msg = make_enriched("wxid_normal", "group@chatroom", wx_db::MessageContent::Text("hi".into())); + assert!(project_message_item(msg, "group@chatroom", &vis).is_some()); + } + + #[test] + fn project_message_item_quote_redaction_hidden_refer() { + let vis = vis_with_hidden_persons(&["wxid_hidden"]); + let raw_xml = r#"replywxid_hiddensecret"#; + let content = wx_db::MessageContent::Quote { + reply_text: Some("my reply".to_string()), + refer_sender: Some("Hidden User".to_string()), + refer_content: Some("secret".to_string()), + refer_type: Some(1), + raw_xml: raw_xml.to_string(), + }; + let msg = make_enriched("wxid_visible", "group@chatroom", content); + let result = project_message_item(msg, "group@chatroom", &vis).unwrap(); + + match &result.message.content { + wx_db::MessageContent::Quote { + reply_text, + refer_sender, + refer_content, + raw_xml, + .. + } => { + assert_eq!(reply_text.as_deref(), Some("my reply")); + assert!(refer_sender.is_none()); + assert!(refer_content.is_none()); + assert!(raw_xml.is_empty()); + } + _ => panic!("expected Quote"), + } + // Snippet should degrade to "[引用] my reply" + assert!(result.snippet.contains("my reply")); + assert!(!result.snippet.contains("Hidden User")); + } + + #[test] + fn project_message_item_quote_visible_refer_preserved() { + let vis = vis_with_hidden_persons(&["wxid_hidden"]); + let raw_xml = r#"replywxid_normalvisible msg"#; + let content = wx_db::MessageContent::Quote { + reply_text: Some("ok".to_string()), + refer_sender: Some("Normal User".to_string()), + refer_content: Some("visible msg".to_string()), + refer_type: Some(1), + raw_xml: raw_xml.to_string(), + }; + let msg = make_enriched("wxid_visible", "group@chatroom", content); + let result = project_message_item(msg, "group@chatroom", &vis).unwrap(); + + match &result.message.content { + wx_db::MessageContent::Quote { + refer_sender, + refer_content, + .. + } => { + assert_eq!(refer_sender.as_deref(), Some("Normal User")); + assert_eq!(refer_content.as_deref(), Some("visible msg")); + } + _ => panic!("expected Quote"), + } + } + + #[test] + fn project_message_item_quote_redaction_group_chatusr() { + // Real WeChat group chat XML: is chatroom, is actual sender + let vis = vis_with_hidden_persons(&["wxid_hidden"]); + let raw_xml = r#"replygroup@chatroomwxid_hiddenHidden Usersecret"#; + let content = wx_db::MessageContent::Quote { + reply_text: Some("my reply".to_string()), + refer_sender: Some("Hidden User".to_string()), + refer_content: Some("secret".to_string()), + refer_type: Some(1), + raw_xml: raw_xml.to_string(), + }; + let msg = make_enriched("wxid_visible", "group@chatroom", content); + let result = project_message_item(msg, "group@chatroom", &vis).unwrap(); + + match &result.message.content { + wx_db::MessageContent::Quote { + reply_text, + refer_sender, + refer_content, + raw_xml, + .. + } => { + assert_eq!(reply_text.as_deref(), Some("my reply")); + assert!(refer_sender.is_none(), "refer_sender should be redacted"); + assert!(refer_content.is_none(), "refer_content should be redacted"); + assert!(raw_xml.is_empty(), "raw_xml should be cleared"); + } + _ => panic!("expected Quote"), + } + assert!(result.snippet.contains("my reply")); + assert!(!result.snippet.contains("Hidden User")); + } + + #[test] + fn project_message_items_show_hidden_bypasses() { + let vis = vis_with_hidden_persons(&["wxid_spam"]); + let items = vec![ + make_enriched("wxid_spam", "group@chatroom", wx_db::MessageContent::Text("spam".into())), + make_enriched("wxid_normal", "group@chatroom", wx_db::MessageContent::Text("hi".into())), + ]; + let result = project_message_items(items, "group@chatroom", &vis, true); + assert_eq!(result.len(), 2); + } + + #[test] + fn project_message_items_filters_hidden_sender() { + let vis = vis_with_hidden_persons(&["wxid_spam"]); + let items = vec![ + make_enriched("wxid_spam", "group@chatroom", wx_db::MessageContent::Text("spam".into())), + make_enriched("wxid_normal", "group@chatroom", wx_db::MessageContent::Text("hi".into())), + ]; + let result = project_message_items(items, "group@chatroom", &vis, false); + assert_eq!(result.len(), 1); + assert_eq!(result[0].message.sender, "wxid_normal"); + } + + #[test] + fn project_session_sender_non_group_noop() { + let vis = vis_with_hidden_persons(&["wxid_spam"]); + let mut session = EnrichedSession { + session: Session { + username: "wxid_spam".to_string(), + summary: "hello".to_string(), + sort_timestamp: 1, + last_msg_type: Some(1), + last_msg_sender: Some("wxid_spam".to_string()), + last_sender_display_name: None, + }, + display_name: "Spam".to_string(), + direction: Some(Direction::Incoming), + detected_at: None, + }; + project_session_sender(&mut session, &vis); + assert_eq!(session.session.summary, "hello"); + assert!(session.session.last_msg_sender.is_some()); + } + + #[test] + fn project_session_sender_group_hidden_sender_redacted() { + let vis = vis_with_hidden_persons(&["wxid_spam"]); + let mut session = EnrichedSession { + session: Session { + username: "group@chatroom".to_string(), + summary: "spam message".to_string(), + sort_timestamp: 1, + last_msg_type: Some(1), + last_msg_sender: Some("wxid_spam".to_string()), + last_sender_display_name: Some("Spammer".to_string()), + }, + display_name: "Group".to_string(), + direction: Some(Direction::Incoming), + detected_at: None, + }; + project_session_sender(&mut session, &vis); + assert_eq!(session.session.summary, "[消息已隐藏]"); + assert!(session.session.last_msg_sender.is_none()); + assert!(session.session.last_sender_display_name.is_none()); + assert!(session.direction.is_none()); + } + + #[test] + fn project_session_sender_group_visible_sender_kept() { + let vis = vis_with_hidden_persons(&["wxid_spam"]); + let mut session = EnrichedSession { + session: Session { + username: "group@chatroom".to_string(), + summary: "normal message".to_string(), + sort_timestamp: 1, + last_msg_type: Some(1), + last_msg_sender: Some("wxid_normal".to_string()), + last_sender_display_name: Some("Normal".to_string()), + }, + display_name: "Group".to_string(), + direction: Some(Direction::Incoming), + detected_at: None, + }; + project_session_sender(&mut session, &vis); + assert_eq!(session.session.summary, "normal message"); + assert_eq!(session.session.last_msg_sender.as_deref(), Some("wxid_normal")); + } +} diff --git a/crates/wx-cli/src/settings.rs b/crates/wx-cli/src/settings.rs new file mode 100644 index 0000000..8020f10 --- /dev/null +++ b/crates/wx-cli/src/settings.rs @@ -0,0 +1,131 @@ +use std::collections::BTreeMap; +use std::fs; +use std::path::Path; + +use serde::{Deserialize, Serialize}; + +/// Per-account visibility settings loaded from `/settings.toml`. +#[derive(Debug, Default, Serialize, Deserialize)] +pub struct Settings { + #[serde(default)] + pub accounts: BTreeMap, +} + +#[derive(Debug, Clone, Default, Serialize, Deserialize)] +pub struct AccountSettings { + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub ignore_contacts: Vec, + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub ignore_tags: Vec, +} + +impl Settings { + /// Load from the default path, returning empty settings if the file doesn't exist. + pub fn load_default() -> Result> { + let ap = wx_paths::AppPaths::new()?; + ap.migrate_config()?; + let path = ap.settings_file(); + Self::load(&path) + } + + /// Load from a specific path, returning empty settings if the file doesn't exist. + pub fn load(path: &Path) -> Result> { + if !path.exists() { + return Ok(Self::default()); + } + let content = fs::read_to_string(path)?; + let settings: Self = toml::from_str(&content)?; + Ok(settings.sanitize()) + } + + /// Get settings for a specific account, returning empty settings if not configured. + pub fn for_account(&self, account_id: &str) -> AccountSettings { + self.accounts.get(account_id).cloned().unwrap_or_default() + } + + /// Trim whitespace and remove empty strings from all lists. + fn sanitize(mut self) -> Self { + for settings in self.accounts.values_mut() { + settings.ignore_contacts = sanitize_list(&settings.ignore_contacts); + settings.ignore_tags = sanitize_list(&settings.ignore_tags); + } + self + } +} + +fn sanitize_list(items: &[String]) -> Vec { + items + .iter() + .map(|s| s.trim().to_string()) + .filter(|s| !s.is_empty()) + .collect() +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn empty_config_returns_defaults() { + let settings = Settings::default(); + let acct = settings.for_account("wxid_test_ab12"); + assert!(acct.ignore_contacts.is_empty()); + assert!(acct.ignore_tags.is_empty()); + } + + #[test] + fn loads_account_settings() { + let toml = r#" +[accounts."wxid_test_ab12"] +ignore_contacts = ["wxid_hidden", "12345@chatroom"] +ignore_tags = ["同事"] +"#; + let settings: Settings = toml::from_str(toml).unwrap(); + let acct = settings.for_account("wxid_test_ab12"); + assert_eq!(acct.ignore_contacts, vec!["wxid_hidden", "12345@chatroom"]); + assert_eq!(acct.ignore_tags, vec!["同事"]); + } + + #[test] + fn unknown_account_returns_empty() { + let toml = r#" +[accounts."wxid_other"] +ignore_contacts = ["wxid_hidden"] +"#; + let settings: Settings = toml::from_str(toml).unwrap(); + let acct = settings.for_account("wxid_test_ab12"); + assert!(acct.ignore_contacts.is_empty()); + } + + #[test] + fn sanitize_trims_and_removes_empty() { + let toml = r#" +[accounts."wxid_test"] +ignore_contacts = [" wxid_a ", "", " ", "wxid_b"] +ignore_tags = [" 同事 ", ""] +"#; + let settings: Settings = toml::from_str::(toml).unwrap().sanitize(); + let acct = settings.for_account("wxid_test"); + assert_eq!(acct.ignore_contacts, vec!["wxid_a", "wxid_b"]); + assert_eq!(acct.ignore_tags, vec!["同事"]); + } + + #[test] + fn legacy_removed_field_silently_ignored() { + // Verify that a TOML with an unknown field (formerly used, now removed) + // deserializes without error — serde without deny_unknown_fields ignores it. + let legacy_key = format!("ignore_{}", "senders"); + let toml = format!( + "[accounts.\"wxid_test\"]\nignore_contacts = [\"wxid_a\"]\n{legacy_key} = [\"wxid_spam\", \"wxid_noisy\"]\n" + ); + let settings: Settings = toml::from_str(&toml).unwrap(); + let acct = settings.for_account("wxid_test"); + assert_eq!(acct.ignore_contacts, vec!["wxid_a"]); + } + + #[test] + fn missing_file_returns_empty() { + let settings = Settings::load(Path::new("/nonexistent/settings.toml")).unwrap(); + assert!(settings.accounts.is_empty()); + } +} diff --git a/crates/wx-cli/src/util.rs b/crates/wx-cli/src/util.rs new file mode 100644 index 0000000..a05c9f7 --- /dev/null +++ b/crates/wx-cli/src/util.rs @@ -0,0 +1,289 @@ +use std::path::PathBuf; + +use crate::cmd::thin_client::{ThinClient, ThinClientError, ThinClientOptions}; +use wx_context::{AccountContext, DecryptRequest, DecryptStats, PersistentCache}; + +/// Open a WechatDb: direct encrypted open if raw_key available, else decrypt+cache (core only). +pub fn open_db_core( + acct: &AccountContext, + progress: impl Fn(wx_context::DecryptProgress) + Send + Sync, +) -> Result<(wx_db::WechatDb, Option), Box> { + if acct.raw_key.is_some() { + eprintln!("Direct encrypted open (SQLCipher)"); + let db = wx_context::open_encrypted_db(acct)?; + Ok((db, None)) + } else { + let params = &wx_decrypt::MACOS_4_1_7_31; + let cache = PersistentCache::new(acct, params)?; + let stats = DecryptRequest::new() + .core() + .execute_with_progress(&cache, progress)?; + let db = wx_db::WechatDb::open(cache.decrypted_root())?; + Ok((db, Some(stats))) + } +} + +/// Open a WechatDb: direct encrypted open if raw_key available, else decrypt+cache (all DBs). +pub fn open_db_all( + acct: &AccountContext, + progress: impl Fn(wx_context::DecryptProgress) + Send + Sync, +) -> Result< + ( + wx_db::WechatDb, + Option, + Option, + ), + Box, +> { + if acct.raw_key.is_some() { + eprintln!("Direct encrypted open (SQLCipher)"); + let db = wx_context::open_encrypted_db(acct)?; + Ok((db, None, None)) + } else { + let params = &wx_decrypt::MACOS_4_1_7_31; + let cache = PersistentCache::new(acct, params)?; + let stats = DecryptRequest::new() + .all() + .execute_with_progress(&cache, progress)?; + let db = wx_db::WechatDb::open(cache.decrypted_root())?; + Ok((db, Some(cache), Some(stats))) + } +} + +/// Print the account detection note (if any) to stderr. +pub fn print_detection_note(acct: &wx_context::AccountContext) { + if let Some(note) = &acct.detection_note { + eprintln!("{note}"); + } +} + +/// Progress callback for `ensure_decrypted_with_progress` — prints to stderr. +pub fn decrypt_progress_callback(event: wx_context::DecryptProgress) { + match event { + wx_context::DecryptProgress::Starting { total } => { + eprintln!("Decrypting {total} databases..."); + } + wx_context::DecryptProgress::Decrypting { .. } => {} + wx_context::DecryptProgress::Decrypted { .. } => {} + wx_context::DecryptProgress::Skipped { .. } => {} + wx_context::DecryptProgress::Failed { .. } => {} + _ => {} + } +} + +/// Print cache decrypt warnings and summary to stderr. +pub fn print_cache_stats(stats: &wx_context::DecryptStats) { + for w in &stats.warnings { + eprintln!(" {w}"); + } + if stats.decrypted > 0 || stats.errors > 0 { + eprintln!( + "Cache: {} decrypted, {} cached, {} errors, {} WAL patched", + stats.decrypted, stats.skipped, stats.errors, stats.wal_patched + ); + } +} + +/// Print FTS index build stats to stderr (silent if index was already fresh). +/// Kept for Task 6: remove self-built FTS index code. +#[allow(dead_code)] +pub fn print_fts_stats(stats: &wx_db::FtsBuildStats) { + if stats.was_fresh { + return; + } + eprintln!( + "Search index: {} messages indexed in {:.1}s", + stats.indexed, stats.duration_secs + ); +} + +pub fn parse_hex_key_32( + hex_key: &str, + source: &str, +) -> Result<[u8; 32], Box> { + let bytes = hex::decode(hex_key).map_err(|e| format!("invalid {source}: {e}"))?; + if bytes.len() != 32 { + return Err(format!("{source} must be 32 bytes, got {}", bytes.len()).into()); + } + + let mut key = [0u8; 32]; + key.copy_from_slice(&bytes); + Ok(key) +} + +pub fn find_db_files(dir: &std::path::Path) -> Result, std::io::Error> { + wx_context::discover_db_files(dir) + .map(|files| files.into_iter().map(|f| f.path).collect()) + .map_err(|e| std::io::Error::other(e.to_string())) +} + +pub fn sanitize_filename(name: &str) -> String { + name.chars() + .map(|c| { + if c.is_alphanumeric() || c == '-' || c == '_' || c == '.' { + c + } else { + '_' + } + }) + .collect() +} + +/// When `all` is true, return [`wx_db::MAX_QUERY_LIMIT`]; otherwise delegate to +/// [`wx_db::effective_limit`] which clamps `limit` into `[DEFAULT_QUERY_LIMIT, MAX_QUERY_LIMIT]`. +pub fn effective_limit_all(all: bool, limit: usize) -> usize { + if all { + wx_db::MAX_QUERY_LIMIT + } else { + wx_db::effective_limit(limit) + } +} + +/// Attempt a remote API call via ThinClient; on connection/auth failure fall back to the +/// local path. This encapsulates the `probe_health → remote_fn → should_fallback → local_fn` +/// pattern shared by `search`, `contacts`, and `sessions`. +pub fn try_remote_or_local( + options: &ThinClientOptions, + remote_fn: impl FnOnce(&ThinClient) -> Result, + local_fn: impl FnOnce() -> Result>, + label: &str, +) -> Result> { + if options.is_enabled() { + let client = ThinClient::new(options.clone()); + match client.probe_health().and_then(|_| remote_fn(&client)) { + Ok(result) => return Ok(result), + Err(err) if err.should_fallback(options.mode) => { + eprintln!( + "note: remote server unavailable, falling back to local {label} ({})", + err.fallback_detail() + ); + } + Err(err) => return Err(err.into()), + } + } + local_fn() +} + +pub fn format_month(create_time: i64) -> String { + chrono::DateTime::from_timestamp(create_time, 0) + .map(|dt| dt.with_timezone(&chrono::Local).format("%Y-%m").to_string()) + .unwrap_or_else(|| "1970-01".to_string()) +} + +pub fn walkdir_dat_files(dir: &std::path::Path) -> Vec { + let mut results = Vec::new(); + if let Ok(entries) = std::fs::read_dir(dir) { + for entry in entries.flatten() { + if let Ok(ft) = entry.file_type() { + if ft.is_dir() { + results.extend(walkdir_dat_files(&entry.path())); + } else if ft.is_file() { + let name = entry.file_name().to_string_lossy().to_string(); + if name.ends_with(".dat") { + results.push(entry); + } + } + } + } + } + results +} + +pub fn lookup_or_resolve_nickname( + store: &mut wx_keychain::KeyStore, + account: &wx_keychain::AccountDirInfo, +) -> Option { + if let Some(existing) = store + .get(&account.account_id) + .and_then(|k| k.nickname.as_ref()) + .cloned() + { + return Some(existing); + } + + let key_material = store.resolve_key_material(&account.account_id)?; + let nickname = + wx_keychain::resolve_nickname(&account.data_dir, &key_material, &account.base_wxid) + .ok() + .flatten()?; + + if let Some(key) = store.accounts.get_mut(&account.account_id) { + key.nickname = Some(nickname.clone()); + // Opportunistically repair base_wxid with canonical value + key.base_wxid = Some(account.base_wxid.clone()); + } + + Some(nickname) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn sanitize_filename_basic() { + assert_eq!(sanitize_filename("hello world.txt"), "hello_world.txt"); + } + + #[test] + fn sanitize_filename_cjk() { + // CJK chars are alphanumeric in Rust's Unicode definition + assert_eq!(sanitize_filename("你好世界"), "你好世界"); + } + + #[test] + fn sanitize_filename_path_separators() { + assert_eq!(sanitize_filename("a/b\\c:d"), "a_b_c_d"); + } + + #[test] + fn sanitize_filename_empty() { + assert_eq!(sanitize_filename(""), ""); + } + + #[test] + fn sanitize_filename_allowed_chars() { + assert_eq!(sanitize_filename("foo-bar_baz.qux"), "foo-bar_baz.qux"); + } + + #[test] + fn effective_limit_all_true_returns_max() { + assert_eq!( + effective_limit_all(true, 0), + wx_db::MAX_QUERY_LIMIT + ); + assert_eq!( + effective_limit_all(true, 50), + wx_db::MAX_QUERY_LIMIT + ); + } + + #[test] + fn effective_limit_all_false_delegates() { + assert_eq!(effective_limit_all(false, 0), wx_db::DEFAULT_QUERY_LIMIT); + assert_eq!(effective_limit_all(false, 50), 50); + assert_eq!( + effective_limit_all(false, wx_db::MAX_QUERY_LIMIT + 1), + wx_db::MAX_QUERY_LIMIT + ); + } + + #[test] + fn format_month_zero() { + assert_eq!(format_month(0), "1970-01"); + } + + #[test] + fn format_month_normal() { + // 2024-03-15 12:00:00 UTC + let result = format_month(1710504000); + assert!(result.starts_with("2024-03"), "got: {result}"); + } + + #[test] + fn format_month_leap_year() { + // 2024-02-29 00:00:00 UTC + let result = format_month(1709164800); + assert!(result.starts_with("2024-02"), "got: {result}"); + } +} diff --git a/crates/wx-cli/src/version.rs b/crates/wx-cli/src/version.rs new file mode 100644 index 0000000..ef4644b --- /dev/null +++ b/crates/wx-cli/src/version.rs @@ -0,0 +1,18 @@ +use std::sync::OnceLock; + +pub(crate) fn cli_version_long() -> &'static str { + static CLI_VERSION: OnceLock = OnceLock::new(); + CLI_VERSION + .get_or_init(|| { + let version = env!("CARGO_PKG_VERSION"); + let sha = option_env!("VERGEN_GIT_SHA").unwrap_or("unknown"); + let sha_short = if sha.len() >= 7 { &sha[..7] } else { sha }; + let date = option_env!("VERGEN_BUILD_DATE").unwrap_or("unknown"); + format!("{version} ({sha_short} {date})") + }) + .as_str() +} + +pub(crate) fn cli_version_string() -> String { + cli_version_long().to_string() +} diff --git a/crates/wx-cli/src/visibility_projection.rs b/crates/wx-cli/src/visibility_projection.rs new file mode 100644 index 0000000..f36b404 --- /dev/null +++ b/crates/wx-cli/src/visibility_projection.rs @@ -0,0 +1,158 @@ +use wx_context::VisibilityIndex; +use wx_db::Contact; + +use crate::output::{JsonEnvelope, PagingMeta, StatsMeta}; +use crate::schema::{project_session_sender, EnrichedSession}; + +/// Filter a fully collected result set, then rebuild paging metadata for the visible slice. +/// +/// Phase 1 uses this only on relatively small list-shaped outputs (`contacts` / `sessions`), +/// where loading the current matched result set into memory is acceptable. +pub fn project_visible_envelope( + items: Vec, + limit: usize, + offset: usize, + stats: &StatsMeta, + show_hidden: bool, + should_hide: impl Fn(&T) -> bool, +) -> JsonEnvelope { + let visible: Vec = if show_hidden { + items + } else { + items + .into_iter() + .filter(|item| !should_hide(item)) + .collect() + }; + + let total = visible.len(); + let start = offset.min(total); + let paged: Vec = visible.into_iter().skip(start).take(limit).collect(); + let returned = paged.len(); + let has_more = start + returned < total; + + JsonEnvelope { + items: paged, + paging: PagingMeta { + limit, + offset, + returned, + has_more, + total, + }, + stats: StatsMeta { + scanned: stats.scanned, + skipped: stats.skipped, + elapsed_ms: stats.elapsed_ms, + shard_warnings: stats.shard_warnings.clone(), + }, + } +} + +pub fn project_contacts_envelope( + contacts: Vec, + visibility: &VisibilityIndex, + limit: usize, + offset: usize, + stats: &StatsMeta, + show_hidden: bool, +) -> JsonEnvelope { + project_visible_envelope(contacts, limit, offset, stats, show_hidden, |contact| { + visibility.is_hidden_talker(&contact.user_name) + }) +} + +#[allow(dead_code)] +pub fn project_sessions_envelope( + sessions: Vec, + visibility: &VisibilityIndex, + limit: usize, + offset: usize, + stats: &StatsMeta, + show_hidden: bool, + talker_of: impl Fn(&T) -> &str, +) -> JsonEnvelope { + project_visible_envelope(sessions, limit, offset, stats, show_hidden, |session| { + visibility.is_hidden_talker(talker_of(session)) + }) +} + +/// Phase 2: project_sessions_envelope with sender-level redaction for EnrichedSession. +/// +/// After talker-level filtering, applies `project_session_sender` to redact +/// hidden senders in group chat session summaries. +pub fn project_sessions_envelope_enriched( + sessions: Vec, + visibility: &VisibilityIndex, + limit: usize, + offset: usize, + stats: &StatsMeta, + show_hidden: bool, +) -> JsonEnvelope { + let mut envelope = project_visible_envelope( + sessions, + limit, + offset, + stats, + show_hidden, + |session: &EnrichedSession| visibility.is_hidden_talker(&session.session.username), + ); + + if !show_hidden { + for session in &mut envelope.items { + project_session_sender(session, visibility); + } + } + + envelope +} + +#[cfg(test)] +mod tests { + use super::*; + + fn stats_meta(scanned: usize) -> StatsMeta { + StatsMeta { + scanned, + skipped: 0, + elapsed_ms: Some(1), + shard_warnings: Vec::new(), + } + } + + #[test] + fn visible_projection_rebuilds_total_and_offset_on_filtered_set() { + let envelope = project_visible_envelope( + vec!["wxid_a", "wxid_hidden", "wxid_b"], + 1, + 1, + &stats_meta(3), + false, + |talker| *talker == "wxid_hidden", + ); + + assert_eq!(envelope.items, vec!["wxid_b"]); + assert_eq!(envelope.paging.total, 2); + assert_eq!(envelope.paging.returned, 1); + assert_eq!(envelope.paging.offset, 1); + assert!(!envelope.paging.has_more); + assert_eq!(envelope.stats.scanned, 3); + } + + #[test] + fn show_hidden_bypasses_filtering() { + let envelope = project_visible_envelope( + vec!["wxid_a", "wxid_hidden", "wxid_b"], + 3, + 0, + &stats_meta(3), + true, + |talker| *talker == "wxid_hidden", + ); + + assert_eq!(envelope.items, vec!["wxid_a", "wxid_hidden", "wxid_b"]); + assert_eq!(envelope.paging.total, 3); + assert_eq!(envelope.paging.returned, 3); + assert!(!envelope.paging.has_more); + } +} diff --git a/crates/wx-cli/tests/ignore-filters.rs b/crates/wx-cli/tests/ignore-filters.rs new file mode 100644 index 0000000..fe1dfb0 --- /dev/null +++ b/crates/wx-cli/tests/ignore-filters.rs @@ -0,0 +1,935 @@ +use std::fs; +use std::path::Path; +use std::process::Command; + +use rusqlite::{params, Connection}; +use serde_json::Value; +use tempfile::TempDir; +use wx_db::encode_extra_buffer_for_test; + +const TEST_KEY_HEX: &str = "abababababababababababababababababababababababababababababababab"; +const ACCOUNT_SCOPE_A: &str = "wxid_scope_a_ab12"; +const ACCOUNT_SCOPE_B: &str = "wxid_scope_b_cd34"; +const ACCOUNT_TAGS: &str = "wxid_scope_tags_ef56"; +const ACCOUNT_SENDERS: &str = "wxid_scope_senders_gh78"; + +const TALKER_ALICE: &str = "wxid_alice"; +const TALKER_BOB: &str = "wxid_bob"; +const TALKER_TAGGED: &str = "wxid_hidden_tagged"; +const TALKER_GROUP: &str = "team@chatroom"; +const TALKER_SPAM: &str = "wxid_spam"; + +const TABLE_ALICE: &str = "Msg_29a6db07e8bbdb53f5d54cc3c309f3f1"; +const TABLE_BOB: &str = "Msg_8a7b11f2fd24e19a60664a5fe5d56342"; +const TABLE_GROUP: &str = "Msg_adcd19623ae4b1f076f9731d6c37b266"; + +fn bin() -> &'static str { + env!("CARGO_BIN_EXE_wx-cli") +} + +#[test] +fn ignore_rules_are_scoped_by_account_id() { + let fixture = create_fixture(); + let scope_a_dir = fixture.account_dir(ACCOUNT_SCOPE_A); + let scope_b_dir = fixture.account_dir(ACCOUNT_SCOPE_B); + + let scope_a = run_json( + fixture.path(), + &[ + "contacts", + "--data-dir", + scope_a_dir.as_str(), + "--key", + TEST_KEY_HEX, + "--format", + "json", + ], + ); + assert_contact_ids(&scope_a, &[TALKER_BOB, TALKER_TAGGED]); + + let scope_b = run_json( + fixture.path(), + &[ + "contacts", + "--data-dir", + scope_b_dir.as_str(), + "--key", + TEST_KEY_HEX, + "--format", + "json", + ], + ); + assert_contact_ids(&scope_b, &[TALKER_ALICE, TALKER_TAGGED]); +} + +#[test] +fn ignore_tags_hide_matching_contact_at_both_talker_and_sender_level() { + let fixture = create_fixture(); + let tags_dir = fixture.account_dir(ACCOUNT_TAGS); + + let contacts = run_json( + fixture.path(), + &[ + "contacts", + "--data-dir", + tags_dir.as_str(), + "--key", + TEST_KEY_HEX, + "--format", + "json", + ], + ); + assert_contact_ids(&contacts, &[TALKER_ALICE, TALKER_BOB]); + + // Group messages from tagged contact should be sender-level filtered + let group_messages = run_json( + fixture.path(), + &[ + "query", + TALKER_GROUP, + "--data-dir", + tags_dir.as_str(), + "--key", + TEST_KEY_HEX, + "--format", + "json", + ], + ); + let items = group_messages["items"].as_array().expect("query items array"); + // The only group message is from wxid_hidden_tagged; it should be filtered out + assert_eq!(items.len(), 0, "tagged contact's group messages should be sender-level filtered: {group_messages}"); + + // Session should show placeholder for group where last sender is tagged + let sessions = run_json( + fixture.path(), + &[ + "sessions", + "--data-dir", + tags_dir.as_str(), + "--key", + TEST_KEY_HEX, + "--format", + "json", + ], + ); + let items = sessions["items"] + .as_array() + .expect("sessions items array"); + let group_session = items + .iter() + .find(|item| item["username"].as_str() == Some(TALKER_GROUP)); + assert!( + group_session.is_some(), + "group session should stay visible: {sessions}" + ); + let group_session = group_session.unwrap(); + assert_eq!( + group_session["summary"].as_str(), + Some("[消息已隐藏]"), + "summary should be placeholder when last sender is tag-hidden: {group_session}" + ); + assert!( + group_session["last_msg_sender"].is_null(), + "last_msg_sender should be null when tag-hidden: {group_session}" + ); +} + +#[test] +fn hidden_talker_defaults_to_not_found_and_show_hidden_restores_query() { + let fixture = create_fixture(); + let scope_a_dir = fixture.account_dir(ACCOUNT_SCOPE_A); + + let hidden = run_failure( + fixture.path(), + &[ + "query", + TALKER_ALICE, + "--data-dir", + scope_a_dir.as_str(), + "--key", + TEST_KEY_HEX, + "--format", + "json", + ], + ); + assert!(hidden.contains("not found"), "{hidden}"); + assert!(!hidden.contains("hidden"), "{hidden}"); + + let visible = run_json( + fixture.path(), + &[ + "query", + TALKER_ALICE, + "--data-dir", + scope_a_dir.as_str(), + "--key", + TEST_KEY_HEX, + "--show-hidden", + "--format", + "json", + ], + ); + let items = visible["items"].as_array().expect("query items array"); + assert_eq!(items.len(), 1, "{visible}"); + assert_eq!(items[0]["talker"], TALKER_ALICE); + assert_eq!(items[0]["server_id"], 2001); +} + +#[test] +fn contacts_and_sessions_paging_metadata_reflect_visible_result_sets() { + let fixture = create_fixture(); + let scope_a_dir = fixture.account_dir(ACCOUNT_SCOPE_A); + + let contacts = run_json( + fixture.path(), + &[ + "contacts", + "--data-dir", + scope_a_dir.as_str(), + "--key", + TEST_KEY_HEX, + "--limit", + "1", + "--offset", + "1", + "--format", + "json", + ], + ); + assert_eq!(contacts["paging"]["total"], 2); + assert_eq!(contacts["paging"]["returned"], 1); + assert_eq!(contacts["paging"]["has_more"], false); + assert_eq!(contacts["stats"]["scanned"], 3); + assert_contact_ids(&contacts, &[TALKER_TAGGED]); + + let sessions = run_json( + fixture.path(), + &[ + "sessions", + "--data-dir", + scope_a_dir.as_str(), + "--key", + TEST_KEY_HEX, + "--limit", + "1", + "--offset", + "1", + "--format", + "json", + ], + ); + assert_eq!(sessions["paging"]["total"], 2); + assert_eq!(sessions["paging"]["returned"], 1); + assert_eq!(sessions["paging"]["has_more"], false); + assert_eq!(sessions["stats"]["scanned"], 3); + let items = sessions["items"].as_array().expect("sessions items array"); + assert_eq!(items.len(), 1, "{sessions}"); + assert_eq!(items[0]["username"], TALKER_GROUP); +} + +struct Fixture { + root: TempDir, +} + +impl Fixture { + fn path(&self) -> &Path { + self.root.path() + } + + fn account_dir(&self, account_id: &str) -> String { + self.root + .path() + .join(account_id) + .to_str() + .expect("fixture path utf8") + .to_string() + } +} + +fn create_fixture() -> Fixture { + let root = TempDir::new().expect("tempdir"); + for account_id in [ACCOUNT_SCOPE_A, ACCOUNT_SCOPE_B, ACCOUNT_TAGS] { + create_account(root.path(), account_id); + } + create_sender_account(root.path()); + write_settings(root.path()); + Fixture { root } +} + +fn write_settings(home: &Path) { + let config_dir = home.join(".config").join("wechat-utils"); + fs::create_dir_all(&config_dir).expect("create config dir"); + fs::write( + config_dir.join("settings.toml"), + format!( + r#"[accounts."{ACCOUNT_SCOPE_A}"] +ignore_contacts = ["{TALKER_ALICE}"] + +[accounts."{ACCOUNT_SCOPE_B}"] +ignore_contacts = ["{TALKER_BOB}"] + +[accounts."{ACCOUNT_TAGS}"] +ignore_tags = ["Sensitive"] + +[accounts."{ACCOUNT_SENDERS}"] +ignore_contacts = ["{TALKER_SPAM}"] +"# + ), + ) + .expect("write settings"); +} + +fn create_account(root: &Path, account_id: &str) { + let account_dir = root.join(account_id); + let db_root = account_dir.join("db_storage"); + let contact_dir = db_root.join("contact"); + let session_dir = db_root.join("session"); + let message_dir = db_root.join("message"); + + fs::create_dir_all(&contact_dir).expect("create contact dir"); + fs::create_dir_all(&session_dir).expect("create session dir"); + fs::create_dir_all(&message_dir).expect("create message dir"); + + let raw_key = test_raw_key(); + create_encrypted_contact_db(&contact_dir.join("contact.db"), &raw_key); + create_encrypted_session_db(&session_dir.join("session.db"), &raw_key); + create_encrypted_message_db(&message_dir.join("message_0.db"), &raw_key); +} + +fn create_encrypted_contact_db(path: &Path, raw_key: &[u8; 32]) { + create_encrypted_db( + path, + raw_key, + "CREATE TABLE contact ( + username TEXT PRIMARY KEY, + alias TEXT DEFAULT '', + remark TEXT DEFAULT '', + nick_name TEXT DEFAULT '', + description TEXT DEFAULT NULL, + extra_buffer BLOB DEFAULT NULL + ); + CREATE TABLE contact_label ( + label_id_ TEXT, + label_name_ TEXT, + sort_order_ INTEGER + );", + |conn| { + conn.execute( + "INSERT INTO contact_label VALUES (?1, ?2, ?3)", + params!["1", "Sensitive", 0], + ) + .expect("insert label"); + + conn.execute( + "INSERT INTO contact (username, nick_name) VALUES (?1, ?2)", + params![TALKER_ALICE, "Alice"], + ) + .expect("insert alice"); + conn.execute( + "INSERT INTO contact (username, nick_name) VALUES (?1, ?2)", + params![TALKER_BOB, "Bob"], + ) + .expect("insert bob"); + + let tagged_extra = encode_extra_buffer_for_test( + None, + None, + None, + None, + None, + None, + None, + Some("1"), + ); + conn.execute( + "INSERT INTO contact (username, nick_name, extra_buffer) VALUES (?1, ?2, ?3)", + params![TALKER_TAGGED, "Sensitive Person", tagged_extra], + ) + .expect("insert tagged contact"); + }, + ); +} + +fn create_encrypted_session_db(path: &Path, raw_key: &[u8; 32]) { + create_encrypted_db( + path, + raw_key, + "CREATE TABLE SessionTable ( + username TEXT, + sort_timestamp INTEGER, + summary TEXT, + last_msg_type INTEGER DEFAULT NULL, + last_msg_sender TEXT DEFAULT NULL, + last_sender_display_name TEXT DEFAULT NULL + );", + |conn| { + conn.execute( + "INSERT INTO SessionTable VALUES (?1, ?2, ?3, ?4, ?5, ?6)", + params![ + TALKER_ALICE, + 1_700_000_300_i64, + "alice summary", + Some(1_i64), + Some(TALKER_ALICE), + Some("Alice"), + ], + ) + .expect("insert alice session"); + conn.execute( + "INSERT INTO SessionTable VALUES (?1, ?2, ?3, ?4, ?5, ?6)", + params![ + TALKER_BOB, + 1_700_000_200_i64, + "bob summary", + Some(1_i64), + Some(TALKER_BOB), + Some("Bob"), + ], + ) + .expect("insert bob session"); + conn.execute( + "INSERT INTO SessionTable VALUES (?1, ?2, ?3, ?4, ?5, ?6)", + params![ + TALKER_GROUP, + 1_700_000_100_i64, + "group summary", + Some(1_i64), + Some(TALKER_TAGGED), + Some("Sensitive Person"), + ], + ) + .expect("insert group session"); + }, + ); +} + +fn create_encrypted_message_db(path: &Path, raw_key: &[u8; 32]) { + create_encrypted_db( + path, + raw_key, + &format!( + "CREATE TABLE Timestamp (timestamp INTEGER); + CREATE TABLE Name2Id ( + rowid INTEGER PRIMARY KEY, + user_name TEXT + ); + CREATE TABLE [{alice}] ( + sort_seq INTEGER, + server_id INTEGER, + local_type INTEGER, + real_sender_id INTEGER, + create_time INTEGER, + message_content BLOB, + packed_info_data BLOB, + status INTEGER, + WCDB_CT_message_content INTEGER + ); + CREATE TABLE [{bob}] ( + sort_seq INTEGER, + server_id INTEGER, + local_type INTEGER, + real_sender_id INTEGER, + create_time INTEGER, + message_content BLOB, + packed_info_data BLOB, + status INTEGER, + WCDB_CT_message_content INTEGER + ); + CREATE TABLE [{group}] ( + sort_seq INTEGER, + server_id INTEGER, + local_type INTEGER, + real_sender_id INTEGER, + create_time INTEGER, + message_content BLOB, + packed_info_data BLOB, + status INTEGER, + WCDB_CT_message_content INTEGER + );", + alice = TABLE_ALICE, + bob = TABLE_BOB, + group = TABLE_GROUP, + ), + |conn| { + conn.execute("INSERT INTO Timestamp VALUES (?1)", params![1_700_000_000_i64]) + .expect("insert timestamp"); + conn.execute( + "INSERT INTO Name2Id VALUES (?1, ?2)", + params![1_i64, TALKER_ALICE], + ) + .expect("insert alice mapping"); + conn.execute( + "INSERT INTO Name2Id VALUES (?1, ?2)", + params![2_i64, TALKER_BOB], + ) + .expect("insert bob mapping"); + conn.execute( + "INSERT INTO Name2Id VALUES (?1, ?2)", + params![3_i64, TALKER_GROUP], + ) + .expect("insert group mapping"); + conn.execute( + "INSERT INTO Name2Id VALUES (?1, ?2)", + params![4_i64, TALKER_TAGGED], + ) + .expect("insert tagged mapping"); + + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + table = TABLE_ALICE + ), + params![ + 100_i64, + 2001_i64, + 1_i64, + 1_i64, + 1_700_000_301_i64, + b"alice says hi" as &[u8], + None::>, + 0_i32, + None::, + ], + ) + .expect("insert alice message"); + + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + table = TABLE_BOB + ), + params![ + 110_i64, + 2002_i64, + 1_i64, + 2_i64, + 1_700_000_201_i64, + b"bob says hi" as &[u8], + None::>, + 0_i32, + None::, + ], + ) + .expect("insert bob message"); + + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + table = TABLE_GROUP + ), + params![ + 120_i64, + 2003_i64, + 1_i64, + 4_i64, + 1_700_000_101_i64, + b"group message from tagged sender" as &[u8], + None::>, + 0_i32, + None::, + ], + ) + .expect("insert group message"); + }, + ); +} + +fn create_encrypted_db( + path: &Path, + raw_key: &[u8; 32], + schema_sql: &str, + seed: impl FnOnce(&Connection), +) { + let conn = Connection::open(path).expect("open sqlite"); + unsafe { + let rc = rusqlite::ffi::sqlite3_key( + conn.handle(), + raw_key.as_ptr() as *const _, + raw_key.len() as i32, + ); + assert_eq!(rc, 0, "sqlite3_key rc={rc}"); + } + conn.execute_batch(schema_sql).expect("apply schema"); + seed(&conn); +} + +fn test_raw_key() -> [u8; 32] { + let bytes = hex::decode(TEST_KEY_HEX).expect("decode test key"); + let mut raw_key = [0_u8; 32]; + raw_key.copy_from_slice(&bytes); + raw_key +} + +fn run_json(home: &Path, args: &[&str]) -> Value { + let output = Command::new(bin()) + .args(args) + .env("HOME", home) + .output() + .expect("run wx-cli"); + assert!( + output.status.success(), + "command failed: {:?}\nstdout={}\nstderr={}", + args, + String::from_utf8_lossy(&output.stdout), + String::from_utf8_lossy(&output.stderr) + ); + serde_json::from_slice(&output.stdout).expect("parse json output") +} + +fn run_failure(home: &Path, args: &[&str]) -> String { + let output = Command::new(bin()) + .args(args) + .env("HOME", home) + .output() + .expect("run failing wx-cli command"); + assert!( + !output.status.success(), + "command unexpectedly succeeded: {:?}\nstdout={}\nstderr={}", + args, + String::from_utf8_lossy(&output.stdout), + String::from_utf8_lossy(&output.stderr) + ); + String::from_utf8_lossy(&output.stderr).to_string() +} + +fn assert_contact_ids(envelope: &Value, expected: &[&str]) { + let items = envelope["items"].as_array().expect("contacts items array"); + let actual = items + .iter() + .map(|item| item["user_name"].as_str().expect("contact user_name")) + .collect::>(); + assert_eq!(actual, expected, "{envelope}"); +} + +// ── Phase 2: sender-level hiding tests ────────────────────────────── + +const TABLE_GROUP_SENDER: &str = "Msg_adcd19623ae4b1f076f9731d6c37b266"; + +fn create_sender_account(root: &Path) { + let account_dir = root.join(ACCOUNT_SENDERS); + let db_root = account_dir.join("db_storage"); + let contact_dir = db_root.join("contact"); + let session_dir = db_root.join("session"); + let message_dir = db_root.join("message"); + + fs::create_dir_all(&contact_dir).expect("create contact dir"); + fs::create_dir_all(&session_dir).expect("create session dir"); + fs::create_dir_all(&message_dir).expect("create message dir"); + + let raw_key = test_raw_key(); + + // Contact DB: alice, spam, group + create_encrypted_db( + &contact_dir.join("contact.db"), + &raw_key, + "CREATE TABLE contact ( + username TEXT PRIMARY KEY, + alias TEXT DEFAULT '', + remark TEXT DEFAULT '', + nick_name TEXT DEFAULT '', + description TEXT DEFAULT NULL, + extra_buffer BLOB DEFAULT NULL + ); + CREATE TABLE contact_label ( + label_id_ TEXT, + label_name_ TEXT, + sort_order_ INTEGER + );", + |conn| { + conn.execute( + "INSERT INTO contact (username, nick_name) VALUES (?1, ?2)", + params![TALKER_ALICE, "Alice"], + ) + .expect("insert alice"); + conn.execute( + "INSERT INTO contact (username, nick_name) VALUES (?1, ?2)", + params![TALKER_SPAM, "Spammer"], + ) + .expect("insert spam"); + }, + ); + + // Session DB: alice (private), group (last_msg_sender = spam) + create_encrypted_db( + &session_dir.join("session.db"), + &raw_key, + "CREATE TABLE SessionTable ( + username TEXT, + sort_timestamp INTEGER, + summary TEXT, + last_msg_type INTEGER DEFAULT NULL, + last_msg_sender TEXT DEFAULT NULL, + last_sender_display_name TEXT DEFAULT NULL + );", + |conn| { + conn.execute( + "INSERT INTO SessionTable VALUES (?1, ?2, ?3, ?4, ?5, ?6)", + params![ + TALKER_ALICE, + 1_700_000_300_i64, + "private msg", + Some(1_i64), + Some(TALKER_ALICE), + Some("Alice"), + ], + ) + .expect("insert alice session"); + conn.execute( + "INSERT INTO SessionTable VALUES (?1, ?2, ?3, ?4, ?5, ?6)", + params![ + TALKER_GROUP, + 1_700_000_200_i64, + "spam message in group", + Some(1_i64), + Some(TALKER_SPAM), + Some("Spammer"), + ], + ) + .expect("insert group session"); + }, + ); + + // Message DB: group messages from alice and spam, plus a quote referencing spam + let quote_xml = format!( + r#"my reply to spam5715001{TALKER_SPAM}Spammerspam content"# + ); + + create_encrypted_db( + &message_dir.join("message_0.db"), + &raw_key, + &format!( + "CREATE TABLE Timestamp (timestamp INTEGER); + CREATE TABLE Name2Id ( + rowid INTEGER PRIMARY KEY, + user_name TEXT + ); + CREATE TABLE [{alice}] ( + sort_seq INTEGER, + server_id INTEGER, + local_type INTEGER, + real_sender_id INTEGER, + create_time INTEGER, + message_content BLOB, + packed_info_data BLOB, + status INTEGER, + WCDB_CT_message_content INTEGER + ); + CREATE TABLE [{group}] ( + sort_seq INTEGER, + server_id INTEGER, + local_type INTEGER, + real_sender_id INTEGER, + create_time INTEGER, + message_content BLOB, + packed_info_data BLOB, + status INTEGER, + WCDB_CT_message_content INTEGER + );", + alice = TABLE_ALICE, + group = TABLE_GROUP_SENDER, + ), + |conn| { + conn.execute("INSERT INTO Timestamp VALUES (?1)", params![1_700_000_000_i64]) + .expect("insert timestamp"); + conn.execute( + "INSERT INTO Name2Id VALUES (?1, ?2)", + params![1_i64, TALKER_ALICE], + ) + .expect("insert alice mapping"); + conn.execute( + "INSERT INTO Name2Id VALUES (?1, ?2)", + params![3_i64, TALKER_GROUP], + ) + .expect("insert group mapping"); + conn.execute( + "INSERT INTO Name2Id VALUES (?1, ?2)", + params![5_i64, TALKER_SPAM], + ) + .expect("insert spam mapping"); + + // Private chat message (sender hiding should NOT apply) + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + table = TABLE_ALICE + ), + params![ + 100_i64, 3001_i64, 1_i64, 1_i64, 1_700_000_301_i64, + b"private hello" as &[u8], None::>, 0_i32, None::, + ], + ) + .expect("insert alice private message"); + + // Group: message from alice (visible sender) + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + table = TABLE_GROUP_SENDER + ), + params![ + 200_i64, 4001_i64, 1_i64, 1_i64, 1_700_000_101_i64, + b"alice says hello in group" as &[u8], None::>, 0_i32, None::, + ], + ) + .expect("insert alice group message"); + + // Group: message from spam (hidden sender) + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + table = TABLE_GROUP_SENDER + ), + params![ + 210_i64, 4002_i64, 1_i64, 5_i64, 1_700_000_102_i64, + b"spam content" as &[u8], None::>, 0_i32, None::, + ], + ) + .expect("insert spam group message"); + + // Group: quote from alice referencing spam + // local_type = (sub_type << 32) | msg_type = (57 << 32) | 49 + let quote_local_type: i64 = (57_i64 << 32) | 49; + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + table = TABLE_GROUP_SENDER + ), + params![ + 220_i64, 4003_i64, + quote_local_type, + 1_i64, 1_700_000_103_i64, + quote_xml.as_bytes(), None::>, 0_i32, None::, + ], + ) + .expect("insert quote message"); + }, + ); +} + +#[test] +fn sender_hiding_filters_hidden_sender_messages_in_group() { + let fixture = create_fixture(); + let senders_dir = fixture.account_dir(ACCOUNT_SENDERS); + + let result = run_json( + fixture.path(), + &[ + "query", TALKER_GROUP, + "--data-dir", senders_dir.as_str(), + "--key", TEST_KEY_HEX, + "--format", "json", + ], + ); + let items = result["items"].as_array().expect("query items array"); + + // spam message (server_id=4002) should be filtered out + let senders: Vec<&str> = items.iter().map(|i| i["sender"].as_str().unwrap()).collect(); + assert!(!senders.contains(&TALKER_SPAM), "hidden sender message should be filtered: {result}"); + assert!(senders.contains(&TALKER_ALICE), "visible sender should remain: {result}"); + + // paging.total should NOT change (DB-level count) + // paging.returned should reflect filtered items + assert_eq!(result["paging"]["returned"].as_u64().unwrap(), items.len() as u64); +} + +#[test] +fn sender_hiding_does_not_affect_private_chat() { + let fixture = create_fixture(); + let senders_dir = fixture.account_dir(ACCOUNT_SENDERS); + + // TALKER_ALICE is not in ignore_contacts, so her private chat is visible + let result = run_json( + fixture.path(), + &[ + "query", TALKER_ALICE, + "--data-dir", senders_dir.as_str(), + "--key", TEST_KEY_HEX, + "--format", "json", + ], + ); + let items = result["items"].as_array().expect("query items array"); + assert_eq!(items.len(), 1, "private chat should not be filtered: {result}"); +} + +#[test] +fn sender_hiding_redacts_quote_referring_hidden_sender() { + let fixture = create_fixture(); + let senders_dir = fixture.account_dir(ACCOUNT_SENDERS); + + let result = run_json( + fixture.path(), + &[ + "query", TALKER_GROUP, + "--data-dir", senders_dir.as_str(), + "--key", TEST_KEY_HEX, + "--format", "json", + ], + ); + let items = result["items"].as_array().expect("query items array"); + + // Find the quote message (server_id=4003) + let quote = items.iter().find(|i| i["server_id"].as_i64() == Some(4003)); + assert!(quote.is_some(), "quote message should be present (alice's reply): items={:?}", items.iter().map(|i| i["server_id"].as_i64()).collect::>()); + let quote = quote.unwrap(); + + // refer_sender and refer_content should be null (redacted) + let q = "e["content"]["Quote"]; + assert!(q["refer_sender"].is_null(), "refer_sender should be redacted: {quote}"); + assert!(q["refer_content"].is_null(), "refer_content should be redacted: {quote}"); + // reply_text should be preserved + assert!(q["reply_text"].as_str().unwrap().contains("my reply"), "reply_text should be preserved: {quote}"); + // raw_xml should be cleared + assert_eq!(q["raw_xml"].as_str(), Some(""), "raw_xml should be empty: {quote}"); +} + +#[test] +fn sender_hiding_show_hidden_restores_all_messages_and_quotes() { + let fixture = create_fixture(); + let senders_dir = fixture.account_dir(ACCOUNT_SENDERS); + + let result = run_json( + fixture.path(), + &[ + "query", TALKER_GROUP, + "--data-dir", senders_dir.as_str(), + "--key", TEST_KEY_HEX, + "--show-hidden", + "--format", "json", + ], + ); + let items = result["items"].as_array().expect("query items array"); + assert_eq!(items.len(), 3, "show_hidden should restore all 3 messages: {result}"); + + // Quote should have refer_sender intact + let quote = items.iter().find(|i| i["server_id"].as_i64() == Some(4003)).unwrap(); + assert!(!quote["content"]["Quote"]["refer_sender"].is_null(), "refer_sender should be intact with show_hidden: {quote}"); +} + +#[test] +fn sender_hiding_session_placeholder_when_last_sender_hidden() { + let fixture = create_fixture(); + let senders_dir = fixture.account_dir(ACCOUNT_SENDERS); + + let sessions = run_json( + fixture.path(), + &[ + "sessions", + "--data-dir", senders_dir.as_str(), + "--key", TEST_KEY_HEX, + "--format", "json", + ], + ); + let items = sessions["items"].as_array().expect("sessions items array"); + + let group_session = items.iter().find(|i| i["username"].as_str() == Some(TALKER_GROUP)); + assert!(group_session.is_some(), "group session should be visible: {sessions}"); + let group_session = group_session.unwrap(); + + // Summary should be placeholder, sender fields should be null + assert_eq!(group_session["summary"].as_str(), Some("[消息已隐藏]"), + "summary should be placeholder: {group_session}"); + assert!(group_session["last_msg_sender"].is_null(), + "last_msg_sender should be null: {group_session}"); + assert!(group_session["direction"].is_null(), + "direction should be null: {group_session}"); +} diff --git a/crates/wx-cli/tests/native_fts_e2e.rs b/crates/wx-cli/tests/native_fts_e2e.rs new file mode 100644 index 0000000..5bfcce6 --- /dev/null +++ b/crates/wx-cli/tests/native_fts_e2e.rs @@ -0,0 +1,294 @@ +/// End-to-end integration test for the native FTS pipeline. +/// +/// Tests the complete path: +/// tokenizer registration (wx_context) → +/// FTS5 MATCH query (rusqlite) → +/// search_message_fts (wx_db) → +/// result enrichment (wx_cli::schema) +/// +/// This test lives in wx-cli (which depends on all three lower crates) +/// to avoid introducing a circular dependency between wx-db and wx-context. +use rusqlite::Connection; +use wx_context::register_mm_fts_tokenizer; +use wx_db::native_fts::search_message_fts; + +// --------------------------------------------------------------------------- +// Test fixture setup +// --------------------------------------------------------------------------- + +fn build_test_db() -> Connection { + let conn = Connection::open_in_memory().unwrap(); + register_mm_fts_tokenizer(&conn).expect("tokenizer registration should succeed"); + + // Use real column names matching message_fts.db schema + conn.execute_batch( + "CREATE VIRTUAL TABLE message_fts_v4_0 USING fts5( + acontent, message_local_id UNINDEXED, sort_seq UNINDEXED, local_type UNINDEXED, + session_id UNINDEXED, sender_id UNINDEXED, create_time UNINDEXED, + tokenize='MMFtsTokenizer disable_pinyin' + ); + CREATE VIRTUAL TABLE message_fts_v4_1 USING fts5( + acontent, message_local_id UNINDEXED, sort_seq UNINDEXED, local_type UNINDEXED, + session_id UNINDEXED, sender_id UNINDEXED, create_time UNINDEXED, + tokenize='MMFtsTokenizer disable_pinyin' + ); + CREATE VIRTUAL TABLE message_fts_v4_2 USING fts5( + acontent, message_local_id UNINDEXED, sort_seq UNINDEXED, local_type UNINDEXED, + session_id UNINDEXED, sender_id UNINDEXED, create_time UNINDEXED, + tokenize='MMFtsTokenizer disable_pinyin' + ); + CREATE VIRTUAL TABLE message_fts_v4_3 USING fts5( + acontent, message_local_id UNINDEXED, sort_seq UNINDEXED, local_type UNINDEXED, + session_id UNINDEXED, sender_id UNINDEXED, create_time UNINDEXED, + tokenize='MMFtsTokenizer disable_pinyin' + ); + CREATE TABLE name2id (rowid INTEGER PRIMARY KEY, username TEXT NOT NULL);", + ) + .unwrap(); + + conn +} + +#[allow(clippy::too_many_arguments)] +fn insert_msg( + conn: &Connection, + shard: usize, + content: &str, + sort_seq: i64, + local_type: i64, + session_id: i64, + sender_id: i64, + create_time: i64, +) { + let table = format!("message_fts_v4_{shard}"); + conn.execute( + &format!("INSERT INTO {table}(acontent,message_local_id,sort_seq,local_type,session_id,sender_id,create_time) VALUES(?1,?2,?3,?4,?5,?6,?7)"), + rusqlite::params![content, 0i64, sort_seq, local_type, session_id, sender_id, create_time], + ) + .unwrap(); +} + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +/// Chinese single-char search ("你") matches messages containing that character. +#[test] +fn chinese_single_char_search() { + let conn = build_test_db(); + conn.execute_batch( + "INSERT INTO name2id VALUES (1, 'wxid_alice'); + INSERT INTO name2id VALUES (2, 'wxid_self');", + ) + .unwrap(); + insert_msg(&conn, 0, "你好世界", 1000, 1, 1, 2, 1700000001); + insert_msg(&conn, 1, "hello world", 2000, 1, 1, 2, 1700000002); + + let result = search_message_fts(&conn, "你", 10, 0).unwrap(); + assert_eq!(result.total_hits, 1, "Should find Chinese message"); + assert!(result.hits[0].snippet.contains("你")); +} + +/// English word search ("hello") matches. +#[test] +fn english_word_search() { + let conn = build_test_db(); + conn.execute_batch( + "INSERT INTO name2id VALUES (1, 'wxid_alice'); + INSERT INTO name2id VALUES (2, 'wxid_self');", + ) + .unwrap(); + insert_msg(&conn, 0, "你好世界", 1000, 1, 1, 2, 1700000001); + insert_msg(&conn, 1, "hello world", 2000, 1, 1, 2, 1700000002); + + let result = search_message_fts(&conn, "hello", 10, 0).unwrap(); + assert_eq!(result.total_hits, 1, "Should find English message"); + assert!(result.hits[0].snippet.contains("hello")); +} + +/// Porter stemming: "running" inserted, searching "run" should match. +#[test] +fn porter_stemming_match() { + let conn = build_test_db(); + conn.execute_batch( + "INSERT INTO name2id VALUES (1, 'wxid_alice'); + INSERT INTO name2id VALUES (2, 'wxid_self');", + ) + .unwrap(); + insert_msg( + &conn, + 0, + "I am running fast today", + 1000, + 1, + 1, + 2, + 1700000001, + ); + + // Searching for the stem "run" should match "running" + let result = search_message_fts(&conn, "run", 10, 0).unwrap(); + assert_eq!( + result.total_hits, 1, + "Porter stemming: 'run' should match 'running'" + ); +} + +/// Pagination: limit and offset work correctly across multi-shard results. +#[test] +fn pagination_across_shards() { + let conn = build_test_db(); + conn.execute_batch("INSERT INTO name2id VALUES (1, 'wxid_alice');") + .unwrap(); + + // Insert 6 messages across 4 shards + for i in 0..6usize { + insert_msg( + &conn, + i % 4, + "test message content", + (i as i64) * 100, + 1, + 1, + 1, + 1700000000 + i as i64, + ); + } + + // All 6 + let all = search_message_fts(&conn, "test", 10, 0).unwrap(); + assert_eq!(all.total_hits, 6); + assert_eq!(all.hits.len(), 6); + + // First 3 + let page1 = search_message_fts(&conn, "test", 3, 0).unwrap(); + assert_eq!(page1.total_hits, 6); + assert_eq!(page1.hits.len(), 3); + + // Skip 4, get remaining 2 + let page2 = search_message_fts(&conn, "test", 10, 4).unwrap(); + assert_eq!(page2.total_hits, 6); + assert_eq!(page2.hits.len(), 2); +} + +/// name2id resolution produces correct talker/sender strings. +#[test] +fn name2id_resolution() { + let conn = build_test_db(); + conn.execute_batch( + "INSERT INTO name2id VALUES (10, 'wxid_alice'); + INSERT INTO name2id VALUES (20, 'wxid_bob');", + ) + .unwrap(); + // session_id=10 (talker=wxid_alice), sender_id=20 (sender=wxid_bob) + insert_msg(&conn, 0, "hello from bob", 1000, 1, 10, 20, 1700000001); + + let result = search_message_fts(&conn, "hello", 10, 0).unwrap(); + assert_eq!(result.hits.len(), 1); + let hit = &result.hits[0]; + assert_eq!(hit.talker, "wxid_alice"); + assert_eq!(hit.sender, "wxid_bob"); +} + +/// Result ordering: create_time DESC, sort_seq DESC. +#[test] +fn result_ordering() { + let conn = build_test_db(); + conn.execute_batch("INSERT INTO name2id VALUES (1, 'wxid_alice');") + .unwrap(); + + // Insert with known create_times + insert_msg(&conn, 0, "first message hello", 100, 1, 1, 1, 1700000001); + insert_msg(&conn, 1, "second message hello", 200, 1, 1, 1, 1700000003); + insert_msg(&conn, 2, "third message hello", 300, 1, 1, 1, 1700000002); + + let result = search_message_fts(&conn, "hello", 10, 0).unwrap(); + assert_eq!(result.hits.len(), 3); + // Ordered by create_time DESC + let times: Vec = result.hits.iter().map(|h| h.create_time).collect(); + assert!( + times[0] >= times[1] && times[1] >= times[2], + "Results not ordered by create_time DESC: {times:?}" + ); +} + +/// hit_type is Message and server_id is 0 for native FTS results. +#[test] +fn hit_type_and_server_id() { + let conn = build_test_db(); + conn.execute_batch("INSERT INTO name2id VALUES (1, 'wxid_alice');") + .unwrap(); + insert_msg(&conn, 0, "hello native fts", 1000, 1, 1, 1, 1700000001); + + let result = search_message_fts(&conn, "hello", 10, 0).unwrap(); + assert_eq!(result.hits.len(), 1); + let hit = &result.hits[0]; + assert_eq!(hit.server_id, 0, "server_id must be 0 for native FTS"); + assert!( + matches!(hit.hit_type, wx_db::FtsHitType::Message), + "hit_type must be Message" + ); +} + +/// Mixed CJK + English text: both are searchable. +#[test] +fn mixed_cjk_english() { + let conn = build_test_db(); + conn.execute_batch("INSERT INTO name2id VALUES (1, 'wxid_alice');") + .unwrap(); + insert_msg(&conn, 0, "我在用iPhone发消息", 1000, 1, 1, 1, 1700000001); + + // CJK char search + let r = search_message_fts(&conn, "我", 10, 0).unwrap(); + assert_eq!(r.total_hits, 1, "CJK search should find mixed message"); + + // English word search + let r2 = search_message_fts(&conn, "iphone", 10, 0).unwrap(); + assert_eq!(r2.total_hits, 1, "English search should find mixed message"); +} + +/// Special character queries (@, [, *) must not trigger fallback (BUG-3 regression guard). +/// After BUG-3 fix, load_name2id() succeeds and the MATCH query executes normally. +/// Results may be empty (no indexed messages contain these chars), but the call must succeed. +#[test] +fn special_char_queries_do_not_fail() { + let conn = build_test_db(); + conn.execute_batch("INSERT INTO name2id VALUES (1, 'wxid_alice');") + .unwrap(); + insert_msg( + &conn, + 0, + "hello world @user test", + 1000, + 1, + 1, + 1, + 1700000001, + ); + insert_msg( + &conn, + 1, + "this [is] a bracket test", + 2000, + 1, + 1, + 1, + 1700000002, + ); + + // '@' — must not error (BUG-3 fix: load_name2id used to fail before MATCH) + let r = search_message_fts(&conn, "@", 10, 0); + assert!(r.is_ok(), "@ query should not fail: {:?}", r.err()); + + // '[' — must not error + let r = search_message_fts(&conn, "[", 10, 0); + assert!(r.is_ok(), "[ query should not fail: {:?}", r.err()); + + // '*' — must not error + let r = search_message_fts(&conn, "*", 10, 0); + assert!(r.is_ok(), "* query should not fail: {:?}", r.err()); + + // '@user' — can find indexed text containing '@user' + let r = search_message_fts(&conn, "@user", 10, 0).unwrap(); + assert_eq!(r.total_hits, 1, "@user should match the first message"); +} diff --git a/crates/wx-cli/tests/paths_status_cli.rs b/crates/wx-cli/tests/paths_status_cli.rs new file mode 100644 index 0000000..8dfb35e --- /dev/null +++ b/crates/wx-cli/tests/paths_status_cli.rs @@ -0,0 +1,88 @@ +use std::process::Command; + +fn bin() -> &'static str { + env!("CARGO_BIN_EXE_wx-cli") +} + +#[test] +fn paths_json_outputs_valid_json_with_expected_fields() { + let output = Command::new(bin()) + .args(["paths", "--json"]) + .output() + .expect("run paths --json"); + assert!(output.status.success(), "paths --json failed: {output:?}"); + + let json: serde_json::Value = + serde_json::from_slice(&output.stdout).expect("valid JSON from paths --json"); + let obj = json.as_object().expect("JSON is an object"); + + let expected_fields = [ + "platform", + "config_dir", + "keys_file", + "settings_file", + "cache_root", + "state_root", + "logs_dir", + "server_state_dir", + "server_stdout_log", + "server_stderr_log", + "temp_root", + ]; + for field in &expected_fields { + assert!( + obj.contains_key(*field), + "missing field '{field}' in paths --json output" + ); + } +} + +#[test] +fn paths_text_contains_expected_labels() { + let output = Command::new(bin()) + .args(["paths"]) + .output() + .expect("run paths"); + assert!(output.status.success(), "paths failed: {output:?}"); + + let stdout = String::from_utf8_lossy(&output.stdout); + + // Platform header should be present + assert!( + stdout.contains("Platform:"), + "paths text output should contain 'Platform:' header" + ); + + let expected_labels = [ + "config_dir", + "keys_file", + "settings_file", + "cache_root", + "state_root", + "logs_dir", + "server_state_dir", + "server_stdout_log", + "server_stderr_log", + "temp_root", + ]; + for label in &expected_labels { + assert!( + stdout.contains(label), + "missing label '{label}' in paths text output" + ); + } +} + +#[test] +fn status_outputs_paths_line_even_without_accounts() { + let output = Command::new(bin()) + .args(["status"]) + .output() + .expect("run status"); + // status may fail if WeChat version check fails, but stdout should still have Paths: + let stdout = String::from_utf8_lossy(&output.stdout); + assert!( + stdout.contains("Paths:"), + "status output should contain 'Paths:' line, got: {stdout}" + ); +} diff --git a/crates/wx-cli/tests/serve-media.rs b/crates/wx-cli/tests/serve-media.rs new file mode 100644 index 0000000..c6dad3e --- /dev/null +++ b/crates/wx-cli/tests/serve-media.rs @@ -0,0 +1,1129 @@ +use std::collections::HashMap; +use std::fs; +use std::io::{Read, Write}; +use std::net::{TcpListener, TcpStream}; +use std::path::Path; +use std::process::{Child, Command, Stdio}; +use std::thread; +use std::time::{Duration, Instant}; + +use rusqlite::{params, Connection}; +use tempfile::TempDir; +use wx_db::encode_packed_info_for_test; + +const TEST_KEY_HEX: &str = "abababababababababababababababababababababababababababababababab"; +const TEST_ACCOUNT_ID: &str = "wxid_test_account"; +const TALKER: &str = "wxid_alice"; +const MSG_TABLE: &str = "Msg_29a6db07e8bbdb53f5d54cc3c309f3f1"; +const GROUP_TALKER: &str = "test@chatroom"; +const GROUP_MSG_TABLE: &str = "Msg_1d282e28b02b5c9f9522f855de32f9a8"; +const HIDDEN_SENDER: &str = "wxid_spam"; + +fn bin() -> &'static str { + env!("CARGO_BIN_EXE_wx-cli") +} + +struct TestServer { + _fixture: TempDir, + child: Child, + base_url: String, +} + +impl Drop for TestServer { + fn drop(&mut self) { + let _ = self.child.kill(); + let _ = self.child.wait(); + } +} + +#[test] +fn serve_media_missing_server_id_returns_400() { + let server = spawn_test_server(); + let response = http_get(&server.base_url, "/api/v1/media?talker=wxid_alice"); + assert_eq!(response.status_code, 400, "{response:#?}"); +} + +#[test] +fn serve_media_missing_talker_returns_400() { + let server = spawn_test_server(); + let response = http_get(&server.base_url, "/api/v1/media?server_id=2001"); + assert_eq!(response.status_code, 400, "{response:#?}"); +} + +#[test] +fn serve_media_invalid_format_returns_400() { + let server = spawn_test_server(); + let response = http_get( + &server.base_url, + "/api/v1/media?server_id=2001&talker=wxid_alice&format=wav", + ); + assert_eq!(response.status_code, 400, "{response:#?}"); +} + +#[test] +fn serve_media_missing_asset_returns_404() { + let server = spawn_test_server(); + let response = http_get( + &server.base_url, + "/api/v1/media?server_id=2001&talker=wxid_alice", + ); + assert_eq!(response.status_code, 404, "{response:#?}"); +} + +#[test] +fn serve_media_hidden_sender_in_group_returns_404() { + let server = spawn_test_server_with_hidden_contacts(&[HIDDEN_SENDER]); + let response = http_get( + &server.base_url, + &format!("/api/v1/media?server_id=7001&talker={GROUP_TALKER}"), + ); + assert_eq!(response.status_code, 404, "{response:#?}"); + let body = String::from_utf8_lossy(&response.body); + assert!(body.contains("message not found"), "{body}"); + assert!(!body.contains("hidden"), "{body}"); +} + +#[test] +fn serve_media_visible_sender_in_group_with_hidden_persons_not_visibility_blocked() { + let server = spawn_test_server_with_hidden_contacts(&[HIDDEN_SENDER]); + // server_id=7002 is from visible sender (wxid_alice), same group + let response = http_get( + &server.base_url, + &format!("/api/v1/media?server_id=7002&talker={GROUP_TALKER}"), + ); + // May be 404 (image asset not physically present for this talker path) but the + // error should NOT be visibility-related — it should be about the missing asset, + // not "message not found" (which is the visibility error message). + if response.status_code == 404 { + let body = String::from_utf8_lossy(&response.body); + // The media endpoint was reached (not blocked by visibility) but the asset + // isn't on disk. Acceptable. + assert!(!body.contains("message not found"), + "visible sender should NOT get visibility 404: {body}"); + } + // If 200, even better — means asset resolution succeeded +} + +#[test] +fn serve_media_hidden_talker_returns_404_without_leaking_hidden_state() { + let server = spawn_test_server_with_hidden_contacts(&[TALKER]); + let response = http_get( + &server.base_url, + "/api/v1/media?server_id=3001&talker=wxid_alice", + ); + assert_eq!(response.status_code, 404, "{response:#?}"); + let body = String::from_utf8_lossy(&response.body); + assert!(body.contains("message not found"), "{body}"); + assert!(!body.contains("hidden"), "{body}"); +} + +#[test] +fn serve_media_missing_video_mentions_not_downloaded() { + let server = spawn_test_server(); + let response = http_get( + &server.base_url, + "/api/v1/media?server_id=2003&talker=wxid_alice", + ); + assert_eq!(response.status_code, 404, "{response:#?}"); + let body = String::from_utf8_lossy(&response.body); + assert!(body.contains("downloaded locally"), "{body}"); +} + +#[test] +fn serve_media_unsupported_message_returns_415() { + let server = spawn_test_server(); + let response = http_get( + &server.base_url, + "/api/v1/media?server_id=2002&talker=wxid_alice", + ); + assert_eq!(response.status_code, 415, "{response:#?}"); +} + +#[test] +fn serve_media_dispatch_image_returns_png_bytes() { + let server = spawn_test_server(); + let response = http_get( + &server.base_url, + "/api/v1/media?server_id=3001&talker=wxid_alice", + ); + assert_eq!(response.status_code, 200, "{response:#?}"); + assert_eq!(response.header("content-type"), Some("image/png")); + assert_eq!(&response.body[..8], b"\x89PNG\r\n\x1a\n"); +} + +#[test] +fn serve_media_wxgf_embedded_png_returns_png_without_ffmpeg() { + let server = spawn_test_server_with_env(&[("FFMPEG_PATH", "/definitely-missing-ffmpeg")]); + let response = http_get( + &server.base_url, + "/api/v1/media?server_id=3006&talker=wxid_alice", + ); + assert_eq!(response.status_code, 200, "{response:#?}"); + assert_eq!(response.header("content-type"), Some("image/png")); + assert_eq!(&response.body[..8], b"\x89PNG\r\n\x1a\n"); +} + +#[test] +fn serve_media_wxgf_hevc_without_ffmpeg_returns_actionable_415() { + let server = spawn_test_server_with_env(&[("FFMPEG_PATH", "/definitely-missing-ffmpeg")]); + let response = http_get( + &server.base_url, + "/api/v1/media?server_id=3007&talker=wxid_alice", + ); + assert_eq!(response.status_code, 415, "{response:#?}"); + let body = String::from_utf8_lossy(&response.body); + assert!(body.contains("install ffmpeg"), "{body}"); + assert!(body.contains("HEVC"), "{body}"); +} + +#[test] +fn serve_media_wxgf_hevc_returns_png_when_ffmpeg_is_available() { + if !wx_media::ffmpeg_available() { + return; + } + + let server = spawn_test_server(); + let response = http_get( + &server.base_url, + "/api/v1/media?server_id=3008&talker=wxid_alice", + ); + assert_eq!(response.status_code, 200, "{response:#?}"); + assert_eq!(response.header("content-type"), Some("image/png")); + assert_eq!(&response.body[..8], b"\x89PNG\r\n\x1a\n"); +} + +#[test] +fn serve_media_dispatch_voice_returns_ogg_by_default() { + if !wx_media::ffmpeg_available() { + return; + } + + let server = spawn_test_server(); + let response = http_get( + &server.base_url, + "/api/v1/media?server_id=3002&talker=wxid_alice", + ); + assert_eq!(response.status_code, 200, "{response:#?}"); + assert_eq!(response.header("content-type"), Some("audio/ogg")); + assert!(!response.body.is_empty()); +} + +#[test] +fn serve_media_dispatch_voice_mp3_returns_audio_mpeg() { + if !wx_media::ffmpeg_available() { + return; + } + + let server = spawn_test_server(); + let response = http_get( + &server.base_url, + "/api/v1/media?server_id=3002&talker=wxid_alice&format=mp3", + ); + assert_eq!(response.status_code, 200, "{response:#?}"); + assert_eq!(response.header("content-type"), Some("audio/mpeg")); + assert!(!response.body.is_empty()); +} + +#[test] +fn serve_media_voice_ogg_without_ffmpeg_returns_actionable_415() { + let server = spawn_test_server_with_env(&[("FFMPEG_PATH", "/definitely-missing-ffmpeg")]); + let response = http_get( + &server.base_url, + "/api/v1/media?server_id=3002&talker=wxid_alice", + ); + assert_eq!(response.status_code, 415, "{response:#?}"); + let body = String::from_utf8_lossy(&response.body); + assert!(body.contains("install ffmpeg"), "{body}"); + assert!(body.contains("server_id=3002"), "{body}"); +} + +#[test] +fn serve_media_voice_mp3_without_ffmpeg_returns_actionable_415() { + let server = spawn_test_server_with_env(&[("FFMPEG_PATH", "/definitely-missing-ffmpeg")]); + let response = http_get( + &server.base_url, + "/api/v1/media?server_id=3002&talker=wxid_alice&format=mp3", + ); + assert_eq!(response.status_code, 415, "{response:#?}"); + let body = String::from_utf8_lossy(&response.body); + assert!(body.contains("install ffmpeg"), "{body}"); + assert!(body.contains("mp3"), "{body}"); +} + +#[test] +fn serve_media_dispatch_video_returns_inline_file() { + let server = spawn_test_server(); + let response = http_get( + &server.base_url, + "/api/v1/media?server_id=3003&talker=wxid_alice", + ); + assert_eq!(response.status_code, 200, "{response:#?}"); + assert_eq!(response.header("content-type"), Some("video/mp4")); + assert_eq!( + response.header("content-disposition"), + Some("inline; filename=\"vid001.mp4\"") + ); + assert!(response.body.starts_with(b"video payload")); +} + +#[test] +fn serve_media_dispatch_video_from_msg_video_path() { + let server = spawn_test_server(); + let response = http_get( + &server.base_url, + "/api/v1/media?server_id=3005&talker=wxid_alice", + ); + assert_eq!(response.status_code, 200, "{response:#?}"); + assert_eq!(response.header("content-type"), Some("video/mp4")); + assert_eq!( + response.header("content-disposition"), + Some("inline; filename=\"custom-name.mp4\"") + ); + assert_eq!(response.body, b"video via msg/video path".to_vec()); +} + +#[test] +fn serve_media_headers_file_attachment_and_range() { + let server = spawn_test_server(); + let file_response = http_get( + &server.base_url, + "/api/v1/media?server_id=3004&talker=wxid_alice", + ); + assert_eq!(file_response.status_code, 200, "{file_response:#?}"); + assert_eq!( + file_response.header("content-disposition"), + Some("attachment; filename=\"report.txt\"") + ); + + let range_response = http_get_with_headers( + &server.base_url, + "/api/v1/media?server_id=3003&talker=wxid_alice", + &[("Range", "bytes=0-4")], + ); + assert_eq!(range_response.status_code, 206, "{range_response:#?}"); + assert_eq!(range_response.body, b"video".to_vec()); + assert_eq!(range_response.header("content-range"), Some("bytes 0-4/23")); +} + +fn spawn_test_server() -> TestServer { + spawn_test_server_with_env(&[]) +} + +fn spawn_test_server_with_hidden_contacts(hidden_contacts: &[&str]) -> TestServer { + spawn_test_server_with_setup(&[], hidden_contacts) +} + +fn spawn_test_server_with_env(envs: &[(&str, &str)]) -> TestServer { + spawn_test_server_with_setup(envs, &[]) +} + +fn spawn_test_server_with_setup( + envs: &[(&str, &str)], + hidden_contacts: &[&str], +) -> TestServer { + let fixture = create_fixture(); + if !hidden_contacts.is_empty() { + write_settings(fixture.path(), hidden_contacts); + } + let account_dir = fixture.path().join(TEST_ACCOUNT_ID); + let runtime_root = fixture.path().join("runtime"); + let port = find_open_port(); + let mut command = Command::new(bin()); + command + .args([ + "server", + "_worker", + "--data-dir", + account_dir.to_str().expect("fixture path utf8"), + "--key", + TEST_KEY_HEX, + "--host", + "127.0.0.1", + "--port", + &port.to_string(), + "--poll", + "--poll-ms", + "1000", + "--runtime-root", + runtime_root.to_str().expect("runtime root utf8"), + ]) + .env("HOME", fixture.path()) + .stdout(Stdio::null()) + .stderr(Stdio::piped()); + for (key, value) in envs { + command.env(key, value); + } + let mut child = command.spawn().expect("spawn wx-cli server worker"); + + wait_for_server(port, &mut child); + + TestServer { + _fixture: fixture, + child, + base_url: format!("http://127.0.0.1:{port}"), + } +} + +fn wait_for_server(port: u16, child: &mut Child) { + let deadline = Instant::now() + Duration::from_secs(10); + let mut last_error = String::new(); + + while Instant::now() < deadline { + if let Some(status) = child.try_wait().expect("poll child") { + let mut stderr = String::new(); + if let Some(mut pipe) = child.stderr.take() { + let _ = pipe.read_to_string(&mut stderr); + } + panic!("server worker exited early with {status}: {stderr}"); + } + + match TcpStream::connect(("127.0.0.1", port)) { + Ok(mut stream) => { + let _ = stream.write_all( + b"GET /api/v1/health HTTP/1.1\r\nHost: 127.0.0.1\r\nConnection: close\r\n\r\n", + ); + let mut buf = String::new(); + let _ = stream.read_to_string(&mut buf); + if buf.starts_with("HTTP/1.1") || buf.starts_with("HTTP/1.0") { + return; + } + last_error = format!("unexpected health response: {buf}"); + } + Err(err) => { + last_error = err.to_string(); + } + } + + thread::sleep(Duration::from_millis(50)); + } + + panic!("server worker did not start on port {port}: {last_error}"); +} + +#[derive(Debug)] +struct HttpResponse { + status_code: u16, + headers: HashMap, + body: Vec, +} + +fn http_get(base_url: &str, path_and_query: &str) -> HttpResponse { + http_get_with_headers(base_url, path_and_query, &[]) +} + +fn http_get_with_headers( + base_url: &str, + path_and_query: &str, + headers: &[(&str, &str)], +) -> HttpResponse { + let url = format!("{base_url}{path_and_query}"); + let mut request = ureq::get(&url); + for (name, value) in headers { + request = request.set(name, value); + } + + match request.call() { + Ok(response) => build_http_response(response.status(), response), + Err(ureq::Error::Status(status_code, response)) => { + build_http_response(status_code, response) + } + Err(err) => panic!("request failed for {url}: {err}"), + } +} + +fn build_http_response(status_code: u16, response: ureq::Response) -> HttpResponse { + let headers = response + .headers_names() + .into_iter() + .filter_map(|name| { + response + .header(&name) + .map(|value| (name.to_ascii_lowercase(), value.to_string())) + }) + .collect::>(); + let mut body = Vec::new(); + response + .into_reader() + .read_to_end(&mut body) + .expect("read response body"); + HttpResponse { + status_code, + headers, + body, + } +} + +impl HttpResponse { + fn header(&self, name: &str) -> Option<&str> { + self.headers + .get(&name.to_ascii_lowercase()) + .map(String::as_str) + } +} + +fn write_settings(home: &Path, hidden_contacts: &[&str]) { + let config_dir = home.join(".config").join("wechat-utils"); + fs::create_dir_all(&config_dir).expect("create config dir"); + let contacts = hidden_contacts + .iter() + .map(|contact| format!("\"{contact}\"")) + .collect::>() + .join(", "); + let mut toml = format!("[accounts.\"{TEST_ACCOUNT_ID}\"]\n"); + if !hidden_contacts.is_empty() { + toml.push_str(&format!("ignore_contacts = [{contacts}]\n")); + } + fs::write(config_dir.join("settings.toml"), toml).expect("write settings"); +} + +fn create_fixture() -> TempDir { + let dir = TempDir::new().expect("tempdir"); + let account_dir = dir.path().join(TEST_ACCOUNT_ID); + let db_root = account_dir.join("db_storage"); + let contact_dir = db_root.join("contact"); + let session_dir = db_root.join("session"); + let message_dir = db_root.join("message"); + let attach_dir = account_dir.join("msg").join("attach"); + let file_dir = account_dir.join("msg").join("file"); + let video_dir = account_dir.join("msg").join("video"); + let hardlink_dir = db_root.join("hardlink"); + + fs::create_dir_all(&contact_dir).expect("create contact dir"); + fs::create_dir_all(&session_dir).expect("create session dir"); + fs::create_dir_all(&message_dir).expect("create message dir"); + fs::create_dir_all(&hardlink_dir).expect("create hardlink dir"); + fs::create_dir_all(&file_dir).expect("create file dir"); + fs::create_dir_all(&video_dir).expect("create video dir"); + fs::create_dir_all(&attach_dir).expect("create attach dir"); + + let raw_key = test_raw_key(); + create_encrypted_contact_db(&contact_dir.join("contact.db"), &raw_key); + create_encrypted_session_db(&session_dir.join("session.db"), &raw_key); + create_encrypted_message_db(&message_dir.join("message_0.db"), &raw_key); + create_encrypted_voice_db( + &message_dir.join("media_0.db"), + &message_dir.join("media_1.db"), + &raw_key, + ); + create_encrypted_hardlink_db(&hardlink_dir.join("hardlink.db"), &raw_key); + create_image_fixture(&attach_dir); + create_video_fixture(&attach_dir); + create_video_fixture_under_video_dir(&video_dir); + create_file_fixture(&file_dir); + + dir +} + +fn test_raw_key() -> [u8; 32] { + let bytes = hex::decode(TEST_KEY_HEX).expect("decode test key"); + let mut raw_key = [0_u8; 32]; + raw_key.copy_from_slice(&bytes); + raw_key +} + +fn create_encrypted_contact_db(path: &Path, raw_key: &[u8; 32]) { + create_encrypted_db( + path, + raw_key, + "CREATE TABLE contact ( + username TEXT PRIMARY KEY, + alias TEXT DEFAULT '', + remark TEXT DEFAULT '', + nick_name TEXT DEFAULT '', + description TEXT DEFAULT NULL, + extra_buffer BLOB DEFAULT NULL + ); + CREATE TABLE contact_label ( + label_id_ TEXT, + label_name_ TEXT, + sort_order_ INTEGER + );", + |conn| { + conn.execute( + "INSERT INTO contact (username, alias, remark, nick_name) VALUES (?1, ?2, ?3, ?4)", + params![TALKER, "", "", "Alice"], + ) + .expect("insert contact"); + conn.execute( + "INSERT INTO contact (username, alias, remark, nick_name) VALUES (?1, ?2, ?3, ?4)", + params![HIDDEN_SENDER, "", "", "Spammer"], + ) + .expect("insert spam contact"); + }, + ); +} + +fn create_encrypted_session_db(path: &Path, raw_key: &[u8; 32]) { + create_encrypted_db( + path, + raw_key, + "CREATE TABLE SessionTable ( + username TEXT, + sort_timestamp INTEGER, + summary TEXT + );", + |conn| { + conn.execute( + "INSERT INTO SessionTable VALUES (?1, ?2, ?3)", + params![TALKER, 1_700_000_000_i64, "fixture summary"], + ) + .expect("insert session"); + conn.execute( + "INSERT INTO SessionTable VALUES (?1, ?2, ?3)", + params![GROUP_TALKER, 1_700_000_001_i64, "group summary"], + ) + .expect("insert group session"); + }, + ); +} + +fn create_encrypted_message_db(path: &Path, raw_key: &[u8; 32]) { + create_encrypted_db( + path, + raw_key, + &format!( + "CREATE TABLE Timestamp (timestamp INTEGER); + CREATE TABLE Name2Id ( + rowid INTEGER PRIMARY KEY, + user_name TEXT + ); + CREATE TABLE [{table}] ( + sort_seq INTEGER, + server_id INTEGER, + local_type INTEGER, + real_sender_id INTEGER, + create_time INTEGER, + message_content BLOB, + packed_info_data BLOB, + status INTEGER, + WCDB_CT_message_content INTEGER + ); + CREATE TABLE [{group_table}] ( + sort_seq INTEGER, + server_id INTEGER, + local_type INTEGER, + real_sender_id INTEGER, + create_time INTEGER, + message_content BLOB, + packed_info_data BLOB, + status INTEGER, + WCDB_CT_message_content INTEGER + );", + table = MSG_TABLE, + group_table = GROUP_MSG_TABLE, + ), + |conn| { + conn.execute( + "INSERT INTO Timestamp VALUES (?1)", + params![1_700_000_000_i64], + ) + .expect("insert timestamp"); + conn.execute( + "INSERT INTO Name2Id VALUES (?1, ?2)", + params![1_i64, TALKER], + ) + .expect("insert name2id"); + conn.execute( + "INSERT INTO Name2Id VALUES (?1, ?2)", + params![2_i64, GROUP_TALKER], + ) + .expect("insert group name2id"); + conn.execute( + "INSERT INTO Name2Id VALUES (?1, ?2)", + params![3_i64, HIDDEN_SENDER], + ) + .expect("insert spam name2id"); + + let image_info = encode_packed_info_for_test(Some("md5_image_missing_asset"), None); + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + table = MSG_TABLE + ), + params![ + 100_i64, + 2001_i64, + 3_i64, + 1_i64, + 1_700_000_100_i64, + Vec::::new(), + image_info, + 0_i32, + None::, + ], + ) + .expect("insert image message"); + + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + table = MSG_TABLE + ), + params![ + 200_i64, + 2002_i64, + 1_i64, + 1_i64, + 1_700_000_200_i64, + b"plain text" as &[u8], + None::>, + 0_i32, + None::, + ], + ) + .expect("insert text message"); + + let missing_video_info = encode_packed_info_for_test(None, Some("vid_missing")); + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + table = MSG_TABLE + ), + params![ + 250_i64, + 2003_i64, + 43_i64, + 1_i64, + 1_700_000_250_i64, + Vec::::new(), + missing_video_info, + 0_i32, + None::, + ], + ) + .expect("insert missing video message"); + + let image_ok_info = encode_packed_info_for_test(Some("md5_image_ok"), None); + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + table = MSG_TABLE + ), + params![ + 300_i64, + 3001_i64, + 3_i64, + 1_i64, + 1_709_251_200_i64, + Vec::::new(), + image_ok_info, + 0_i32, + None::, + ], + ) + .expect("insert image ok message"); + + let image_wxgf_png_info = encode_packed_info_for_test(Some("md5_image_wxgf_png"), None); + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + table = MSG_TABLE + ), + params![ + 350_i64, + 3006_i64, + 3_i64, + 1_i64, + 1_709_251_205_i64, + Vec::::new(), + image_wxgf_png_info, + 0_i32, + None::, + ], + ) + .expect("insert wxgf png image message"); + + let image_wxgf_hevc_info = + encode_packed_info_for_test(Some("md5_image_wxgf_hevc"), None); + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + table = MSG_TABLE + ), + params![ + 360_i64, + 3007_i64, + 3_i64, + 1_i64, + 1_709_251_206_i64, + Vec::::new(), + image_wxgf_hevc_info, + 0_i32, + None::, + ], + ) + .expect("insert wxgf hevc image message"); + + let image_wxgf_hevc_valid_info = + encode_packed_info_for_test(Some("md5_image_wxgf_hevc_valid"), None); + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + table = MSG_TABLE + ), + params![ + 370_i64, + 3008_i64, + 3_i64, + 1_i64, + 1_709_251_207_i64, + Vec::::new(), + image_wxgf_hevc_valid_info, + 0_i32, + None::, + ], + ) + .expect("insert valid wxgf hevc image message"); + + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + table = MSG_TABLE + ), + params![ + 400_i64, + 3002_i64, + 34_i64, + 1_i64, + 1_709_251_201_i64, + Vec::::new(), + None::>, + 0_i32, + None::, + ], + ) + .expect("insert voice message"); + + let video_info = encode_packed_info_for_test(None, Some("vid001")); + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + table = MSG_TABLE + ), + params![ + 500_i64, + 3003_i64, + 43_i64, + 1_i64, + 1_709_251_202_i64, + Vec::::new(), + video_info, + 0_i32, + None::, + ], + ) + .expect("insert video message"); + + let video_info_msg_video = encode_packed_info_for_test(None, Some("vid002")); + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + table = MSG_TABLE + ), + params![ + 550_i64, + 3005_i64, + 43_i64, + 1_i64, + 1_709_251_204_i64, + Vec::::new(), + video_info_msg_video, + 0_i32, + None::, + ], + ) + .expect("insert msg/video video message"); + + let file_xml = r#"report.txttxt11doc123"#; + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + table = MSG_TABLE + ), + params![ + 600_i64, + 3004_i64, + (6_i64 << 32) | 49_i64, + 1_i64, + 1_709_251_203_i64, + file_xml.as_bytes(), + None::>, + 0_i32, + None::, + ], + ) + .expect("insert file message"); + + // Group chatroom: image message from hidden sender (server_id=7001) + let group_image_info = encode_packed_info_for_test(Some("md5_group_spam_image"), None); + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + table = GROUP_MSG_TABLE + ), + params![ + 900_i64, + 7001_i64, + 3_i64, // msg_type=3 (image) + 3_i64, // real_sender_id=3 → HIDDEN_SENDER + 1_700_000_900_i64, + Vec::::new(), + group_image_info, + 0_i32, + None::, + ], + ) + .expect("insert group spam image"); + + // Group chatroom: image message from visible sender (server_id=7002) + let group_visible_info = encode_packed_info_for_test(Some("md5_image_png_ok"), None); + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + table = GROUP_MSG_TABLE + ), + params![ + 910_i64, + 7002_i64, + 3_i64, // msg_type=3 (image) + 1_i64, // real_sender_id=1 → TALKER (visible) + 1_700_000_910_i64, + Vec::::new(), + group_visible_info, + 0_i32, + None::, + ], + ) + .expect("insert group visible image"); + }, + ); +} + +fn create_encrypted_voice_db(schema_only_path: &Path, voice_path: &Path, raw_key: &[u8; 32]) { + create_encrypted_db( + schema_only_path, + raw_key, + "CREATE TABLE Metadata (value TEXT);", + |conn| { + conn.execute("INSERT INTO Metadata (value) VALUES ('schema-only')", []) + .expect("insert metadata row"); + }, + ); + + create_encrypted_db( + voice_path, + raw_key, + "CREATE TABLE VoiceInfo ( + svr_id TEXT, + voice_data BLOB + );", + |conn| { + conn.execute( + "INSERT INTO VoiceInfo (svr_id, voice_data) VALUES (?1, ?2)", + params!["3002", sample_silk()], + ) + .expect("insert voice blob"); + }, + ); +} + +fn create_encrypted_hardlink_db(path: &Path, raw_key: &[u8; 32]) { + create_encrypted_db( + path, + raw_key, + "CREATE TABLE dir2id (rowid INTEGER PRIMARY KEY, username TEXT); + CREATE TABLE image_hardlink_info_v3 ( + md5 TEXT, file_name TEXT, file_size INTEGER, modify_time INTEGER, + dir1 INTEGER, dir2 INTEGER + ); + CREATE TABLE video_hardlink_info_v3 ( + md5 TEXT, file_name TEXT, file_size INTEGER, modify_time INTEGER, + dir1 INTEGER, dir2 INTEGER + ); + CREATE TABLE file_hardlink_info_v3 ( + md5 TEXT, file_name TEXT, file_size INTEGER, modify_time INTEGER, + dir1 INTEGER, dir2 INTEGER + );", + |conn| { + conn.execute( + "INSERT INTO dir2id (rowid, username) VALUES (?1, ?2)", + params![1_i64, TALKER], + ) + .expect("insert hardlink dir1"); + conn.execute( + "INSERT INTO dir2id (rowid, username) VALUES (?1, ?2)", + params![2_i64, "2026-03"], + ) + .expect("insert hardlink dir2"); + conn.execute( + "INSERT INTO video_hardlink_info_v3 (md5, file_name, file_size, modify_time, dir1, dir2) VALUES (?1, ?2, ?3, ?4, ?5, ?6)", + params!["vid001", "vid001.mp4", 23_i64, 1_709_251_202_i64, 1_i64, 2_i64], + ) + .expect("insert video hardlink"); + conn.execute( + "INSERT INTO video_hardlink_info_v3 (md5, file_name, file_size, modify_time, dir1, dir2) VALUES (?1, ?2, ?3, ?4, ?5, ?6)", + params!["vid002", "custom-name.mp4", 22_i64, 1_709_251_204_i64, 1_i64, 2_i64], + ) + .expect("insert msg/video hardlink"); + conn.execute( + "INSERT INTO file_hardlink_info_v3 (md5, file_name, file_size, modify_time, dir1, dir2) VALUES (?1, ?2, ?3, ?4, ?5, ?6)", + params!["doc123", "report.txt", 11_i64, 1_709_251_203_i64, 1_i64, 2_i64], + ) + .expect("insert file hardlink"); + }, + ); +} + +fn create_image_fixture(attach_dir: &Path) { + let month_dir = attach_dir + .join(format!("{:x}", wx_media::md5_hash(TALKER.as_bytes()))) + .join("2026-03") + .join("Img"); + fs::create_dir_all(&month_dir).expect("create image month dir"); + + let xor_key = 0x5A_u8; + let png = sample_png(); + let encrypted = xor_bytes(&png, xor_key); + fs::write(month_dir.join("md5_image_ok_t.dat"), &encrypted).expect("write thumb dat"); + fs::write(month_dir.join("md5_image_ok.dat"), &encrypted).expect("write image dat"); + + let wxgf_png = sample_wxgf_with_embedded_png(); + let wxgf_png_encrypted = xor_bytes(&wxgf_png, xor_key); + fs::write( + month_dir.join("md5_image_wxgf_png.dat"), + &wxgf_png_encrypted, + ) + .expect("write wxgf embedded png dat"); + + let wxgf_hevc = sample_wxgf_with_hevc(); + let wxgf_hevc_encrypted = xor_bytes(&wxgf_hevc, xor_key); + fs::write( + month_dir.join("md5_image_wxgf_hevc.dat"), + &wxgf_hevc_encrypted, + ) + .expect("write wxgf hevc dat"); + + let wxgf_hevc_valid = sample_wxgf_with_valid_hevc(); + let wxgf_hevc_valid_encrypted = xor_bytes(&wxgf_hevc_valid, xor_key); + fs::write( + month_dir.join("md5_image_wxgf_hevc_valid.dat"), + &wxgf_hevc_valid_encrypted, + ) + .expect("write valid wxgf hevc dat"); +} + +fn create_video_fixture(attach_dir: &Path) { + let video_dir = attach_dir.join(TALKER).join("2026-03").join("Video"); + fs::create_dir_all(&video_dir).expect("create video dir"); + fs::write(video_dir.join("vid001.mp4"), b"video payload bytes 123").expect("write video"); +} + +fn create_file_fixture(file_dir: &Path) { + let month_dir = file_dir.join(TALKER).join("2026-03"); + fs::create_dir_all(&month_dir).expect("create file month dir"); + fs::write(month_dir.join("report.txt"), b"hello file!").expect("write file"); +} + +fn create_video_fixture_under_video_dir(video_dir: &Path) { + let month_dir = video_dir.join(TALKER).join("2026-03"); + fs::create_dir_all(&month_dir).expect("create msg/video dir"); + fs::write( + month_dir.join("custom-name.mp4"), + b"video via msg/video path", + ) + .expect("write msg/video file"); +} + +fn sample_png() -> Vec { + vec![ + 0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A, 0x00, 0x00, 0x00, 0x0D, 0x49, 0x48, 0x44, + 0x52, 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01, 0x08, 0x06, 0x00, 0x00, 0x00, 0x1F, + 0x15, 0xC4, 0x89, 0x00, 0x00, 0x00, 0x0D, 0x49, 0x44, 0x41, 0x54, 0x78, 0x9C, 0x63, 0xF8, + 0xCF, 0xC0, 0xF0, 0x1F, 0x00, 0x05, 0x00, 0x01, 0xFF, 0x89, 0x99, 0x3D, 0x1D, 0x00, 0x00, + 0x00, 0x00, 0x49, 0x45, 0x4E, 0x44, 0xAE, 0x42, 0x60, 0x82, + ] +} + +fn sample_wxgf_with_embedded_png() -> Vec { + let mut data = b"wxgfmetadata".to_vec(); + data.extend_from_slice(&sample_png()); + data +} + +fn sample_wxgf_with_hevc() -> Vec { + let mut data = b"wxgfmetadata".to_vec(); + data.extend_from_slice(&[0x00, 0x00, 0x00, 0x01, 0x26, 0x01, 0x02, 0x03, 0x04]); + data +} + +fn sample_wxgf_with_valid_hevc() -> Vec { + let mut data = b"wxgfmetadata".to_vec(); + data.extend_from_slice(&sample_valid_hevc()); + data +} + +fn sample_valid_hevc() -> Vec { + if !wx_media::ffmpeg_available() { + return vec![0x00, 0x00, 0x00, 0x01, 0x26, 0x01, 0x02, 0x03, 0x04]; + } + + let temp = TempDir::new().expect("tempdir for hevc sample"); + let output = temp.path().join("frame.hevc"); + let ffmpeg = std::env::var("FFMPEG_PATH").unwrap_or_else(|_| "ffmpeg".to_string()); + let status = Command::new(ffmpeg) + .args([ + "-hide_banner", + "-loglevel", + "error", + "-f", + "lavfi", + "-i", + "color=c=red:s=64x64:d=0.04:r=1", + "-frames:v", + "1", + "-c:v", + "libx265", + "-x265-params", + "log-level=error", + "-f", + "hevc", + output.to_str().expect("hevc output path utf8"), + ]) + .status() + .expect("run ffmpeg for hevc sample"); + assert!(status.success(), "failed to create HEVC sample"); + fs::read(output).expect("read hevc sample") +} + +fn xor_bytes(data: &[u8], key: u8) -> Vec { + data.iter().map(|byte| byte ^ key).collect() +} + +fn sample_silk() -> Vec { + let pcm = vec![0_u8; 24_000 / 1_000 * 40 * 2]; + silk_rs::encode_silk(pcm, 24_000, 24_000, true).expect("encode silk") +} + +fn create_encrypted_db( + path: &Path, + raw_key: &[u8; 32], + schema_sql: &str, + seed: impl FnOnce(&Connection), +) { + let conn = Connection::open(path).expect("open sqlite"); + unsafe { + let rc = rusqlite::ffi::sqlite3_key( + conn.handle(), + raw_key.as_ptr() as *const std::ffi::c_void, + 32, + ); + assert_eq!(rc, 0, "sqlite3_key failed for {}", path.display()); + } + conn.execute_batch(schema_sql).expect("apply schema"); + seed(&conn); +} + +fn find_open_port() -> u16 { + let listener = TcpListener::bind("127.0.0.1:0").expect("bind ephemeral port"); + listener.local_addr().expect("listener addr").port() +} diff --git a/crates/wx-cli/tests/server_manager_cli.rs b/crates/wx-cli/tests/server_manager_cli.rs new file mode 100644 index 0000000..00d2173 --- /dev/null +++ b/crates/wx-cli/tests/server_manager_cli.rs @@ -0,0 +1,533 @@ +use std::fs; +use std::path::{Path, PathBuf}; +use std::process::Command; +use std::thread; +use std::time::{Duration, Instant}; + +use rusqlite::{params, Connection}; +use tempfile::TempDir; + +const TEST_KEY_HEX: &str = "abababababababababababababababababababababababababababababababab"; +const TEST_ACCOUNT_ID: &str = "wxid_test_account"; +const TALKER: &str = "wxid_alice"; +const MSG_TABLE: &str = "Msg_29a6db07e8bbdb53f5d54cc3c309f3f1"; + +fn bin() -> &'static str { + env!("CARGO_BIN_EXE_wx-cli") +} + +#[test] +fn help_surface_exposes_server_group_only() { + let root_help = Command::new(bin()) + .arg("--help") + .output() + .expect("run root help"); + assert!(root_help.status.success(), "{root_help:?}"); + let root_help = String::from_utf8_lossy(&root_help.stdout); + assert!(root_help.contains("server")); + assert!(!root_help.lines().any(|line| line.trim() == "serve")); + + let server_help = Command::new(bin()) + .args(["server", "--help"]) + .output() + .expect("run server help"); + assert!(server_help.status.success(), "{server_help:?}"); + let server_help = String::from_utf8_lossy(&server_help.stdout); + assert!(server_help.contains("run")); + assert!(server_help.contains("status")); + assert!(server_help.contains("stop")); + assert!(server_help.contains("restart")); + assert!(!server_help.contains("_worker")); +} + +#[test] +fn server_run_status_restart_and_stop_round_trip() { + let fixture = create_fixture(); + let runtime_root = fixture.path().join("runtime"); + let _guard = ManagedServerGuard::new(runtime_root.clone()); + let account_dir = fixture.path().join(TEST_ACCOUNT_ID); + let port = find_open_port(); + + let run = Command::new(bin()) + .args([ + "server", + "run", + "--data-dir", + account_dir.to_str().expect("fixture path utf8"), + "--key", + TEST_KEY_HEX, + "--host", + "127.0.0.1", + "--port", + &port.to_string(), + "--poll", + "--poll-ms", + "1000", + "--runtime-root", + runtime_root.to_str().expect("runtime root utf8"), + ]) + .output() + .expect("run server"); + assert_runtime_success(&run, &runtime_root); + + let status = server_status_json(&runtime_root); + assert_eq!(status["status"], "running"); + assert_eq!(status["health"], "healthy"); + assert_eq!(status["ready"], true); + assert_eq!(status["base_url"], format!("http://127.0.0.1:{port}")); + assert!(status["pid"].as_u64().unwrap_or_default() > 0); + assert!(status["cli_version"].as_str().is_some()); + + let duplicate = Command::new(bin()) + .args([ + "server", + "run", + "--data-dir", + account_dir.to_str().expect("fixture path utf8"), + "--key", + TEST_KEY_HEX, + "--host", + "127.0.0.1", + "--port", + &port.to_string(), + "--poll", + "--poll-ms", + "1000", + "--runtime-root", + runtime_root.to_str().expect("runtime root utf8"), + ]) + .output() + .expect("run duplicate server"); + assert_runtime_success(&duplicate, &runtime_root); + assert!(String::from_utf8_lossy(&duplicate.stdout).contains("already running")); + + let different_config = Command::new(bin()) + .args([ + "server", + "run", + "--data-dir", + account_dir.to_str().expect("fixture path utf8"), + "--key", + TEST_KEY_HEX, + "--host", + "127.0.0.1", + "--port", + &(port + 1).to_string(), + "--poll", + "--poll-ms", + "1000", + "--runtime-root", + runtime_root.to_str().expect("runtime root utf8"), + ]) + .output() + .expect("run duplicate server with different config"); + assert!(!different_config.status.success(), "{different_config:?}"); + assert!(String::from_utf8_lossy(&different_config.stderr) + .contains("different launch configuration")); + + let unchanged = server_status_json(&runtime_root); + assert_eq!(unchanged["base_url"], format!("http://127.0.0.1:{port}")); + + let restart = Command::new(bin()) + .args([ + "server", + "restart", + "--runtime-root", + runtime_root.to_str().expect("runtime root utf8"), + ]) + .output() + .expect("restart server"); + assert_runtime_success(&restart, &runtime_root); + + let restarted = server_status_json(&runtime_root); + assert_eq!(restarted["status"], "running"); + assert_eq!(restarted["health"], "healthy"); + + let stop = Command::new(bin()) + .args([ + "server", + "stop", + "--runtime-root", + runtime_root.to_str().expect("runtime root utf8"), + ]) + .output() + .expect("stop server"); + assert_runtime_success(&stop, &runtime_root); + assert!(String::from_utf8_lossy(&stop.stdout).contains("server stopped")); + + let stopped = server_status_json(&runtime_root); + assert_eq!(stopped["status"], "not_running"); +} + +#[test] +fn stale_runtime_state_is_reported_and_recovered() { + let fixture = create_fixture(); + let runtime_root = fixture.path().join("runtime"); + let _guard = ManagedServerGuard::new(runtime_root.clone()); + let account_dir = fixture.path().join(TEST_ACCOUNT_ID); + let port = find_open_port(); + + let run_args = [ + "server", + "run", + "--data-dir", + account_dir.to_str().expect("fixture path utf8"), + "--key", + TEST_KEY_HEX, + "--host", + "127.0.0.1", + "--port", + &port.to_string(), + "--poll", + "--poll-ms", + "1000", + "--runtime-root", + runtime_root.to_str().expect("runtime root utf8"), + ]; + + let run = Command::new(bin()) + .args(run_args) + .output() + .expect("run server"); + assert_runtime_success(&run, &runtime_root); + + let status = server_status_json(&runtime_root); + let pid = status["pid"].as_u64().expect("pid in status") as i32; + let kill_result = unsafe { libc::kill(pid, libc::SIGKILL) }; + assert_eq!(kill_result, 0, "kill stale pid"); + wait_for_pid_exit(pid as u32); + + let stale = server_status_json(&runtime_root); + assert_eq!(stale["status"], "stale"); + + let cleanup = Command::new(bin()) + .args([ + "server", + "stop", + "--runtime-root", + runtime_root.to_str().expect("runtime root utf8"), + ]) + .output() + .expect("stop stale server"); + assert!(cleanup.status.success(), "{cleanup:?}"); + assert!(String::from_utf8_lossy(&cleanup.stdout).contains("removed stale server state")); + + let rerun = Command::new(bin()) + .args(run_args) + .output() + .expect("rerun server after stale state"); + assert_runtime_success(&rerun, &runtime_root); + + let recovered = server_status_json(&runtime_root); + assert_eq!(recovered["status"], "running"); + assert_eq!(recovered["health"], "healthy"); +} + +#[test] +fn live_worker_with_bad_health_does_not_spawn_duplicate_and_can_be_stopped() { + let fixture = create_fixture(); + let runtime_root = fixture.path().join("runtime"); + let _guard = ManagedServerGuard::new(runtime_root.clone()); + let account_dir = fixture.path().join(TEST_ACCOUNT_ID); + let port = find_open_port(); + + let run_args = [ + "server", + "run", + "--data-dir", + account_dir.to_str().expect("fixture path utf8"), + "--key", + TEST_KEY_HEX, + "--host", + "127.0.0.1", + "--port", + &port.to_string(), + "--poll", + "--poll-ms", + "1000", + "--runtime-root", + runtime_root.to_str().expect("runtime root utf8"), + ]; + + let run = Command::new(bin()) + .args(run_args) + .output() + .expect("run server"); + assert_runtime_success(&run, &runtime_root); + + let state_path = runtime_root.join("state.json"); + let mut state: serde_json::Value = + serde_json::from_slice(&fs::read(&state_path).expect("read state json")) + .expect("parse state json"); + state["base_url"] = serde_json::Value::String(format!("http://127.0.0.1:{}", port + 10)); + state["port"] = serde_json::Value::from((port + 10) as u64); + fs::write( + &state_path, + serde_json::to_vec_pretty(&state).expect("serialize state"), + ) + .expect("write corrupted state"); + + let rerun = Command::new(bin()) + .args(run_args) + .output() + .expect("rerun unhealthy live server"); + assert!(!rerun.status.success(), "{rerun:?}"); + assert!( + String::from_utf8_lossy(&rerun.stderr).contains("still running but unhealthy"), + "{rerun:?}" + ); + + let stop = Command::new(bin()) + .args([ + "server", + "stop", + "--runtime-root", + runtime_root.to_str().expect("runtime root utf8"), + ]) + .output() + .expect("stop unhealthy live server"); + assert!(stop.status.success(), "{stop:?}"); + assert!(String::from_utf8_lossy(&stop.stdout).contains("server stopped")); +} + +struct ManagedServerGuard { + runtime_root: PathBuf, +} + +impl ManagedServerGuard { + fn new(runtime_root: PathBuf) -> Self { + Self { runtime_root } + } +} + +impl Drop for ManagedServerGuard { + fn drop(&mut self) { + let _ = Command::new(bin()) + .args([ + "server", + "stop", + "--runtime-root", + self.runtime_root.to_str().unwrap_or_default(), + ]) + .output(); + } +} + +fn server_status_json(runtime_root: &Path) -> serde_json::Value { + let deadline = Instant::now() + Duration::from_secs(10); + let mut last = serde_json::Value::Null; + while Instant::now() < deadline { + let output = Command::new(bin()) + .args([ + "server", + "status", + "--format", + "json", + "--runtime-root", + runtime_root.to_str().expect("runtime root utf8"), + ]) + .output() + .expect("run server status"); + assert!(output.status.success(), "{output:?}"); + let parsed: serde_json::Value = + serde_json::from_slice(&output.stdout).expect("parse server status json"); + if matches!( + parsed["status"].as_str(), + Some("running" | "stale" | "not_running") + ) { + return parsed; + } + last = parsed; + thread::sleep(Duration::from_millis(100)); + } + last +} + +fn assert_runtime_success(output: &std::process::Output, runtime_root: &Path) { + if output.status.success() { + return; + } + let stderr_log = runtime_root.join("stderr.log"); + let log = fs::read_to_string(&stderr_log).unwrap_or_else(|_| "".to_string()); + panic!("{output:?}\nworker stderr log:\n{log}"); +} + +fn wait_for_pid_exit(pid: u32) { + let deadline = Instant::now() + Duration::from_secs(5); + while Instant::now() < deadline { + let alive = unsafe { libc::kill(pid as i32, 0) == 0 }; + if !alive { + return; + } + thread::sleep(Duration::from_millis(50)); + } +} + +fn create_fixture() -> TempDir { + let dir = TempDir::new().expect("tempdir"); + let account_dir = dir.path().join(TEST_ACCOUNT_ID); + let db_root = account_dir.join("db_storage"); + let contact_dir = db_root.join("contact"); + let session_dir = db_root.join("session"); + let message_dir = db_root.join("message"); + let attach_dir = account_dir.join("msg").join("attach"); + let file_dir = account_dir.join("msg").join("file"); + let video_dir = account_dir.join("msg").join("video"); + + fs::create_dir_all(&contact_dir).expect("create contact dir"); + fs::create_dir_all(&session_dir).expect("create session dir"); + fs::create_dir_all(&message_dir).expect("create message dir"); + fs::create_dir_all(&attach_dir).expect("create attach dir"); + fs::create_dir_all(&file_dir).expect("create file dir"); + fs::create_dir_all(&video_dir).expect("create video dir"); + + let raw_key = test_raw_key(); + create_encrypted_contact_db(&contact_dir.join("contact.db"), &raw_key); + create_encrypted_session_db(&session_dir.join("session.db"), &raw_key); + create_encrypted_message_db(&message_dir.join("message_0.db"), &raw_key); + + dir +} + +fn test_raw_key() -> [u8; 32] { + let bytes = hex::decode(TEST_KEY_HEX).expect("decode test key"); + let mut raw_key = [0_u8; 32]; + raw_key.copy_from_slice(&bytes); + raw_key +} + +fn create_encrypted_contact_db(path: &Path, raw_key: &[u8; 32]) { + create_encrypted_db( + path, + raw_key, + "CREATE TABLE contact ( + username TEXT PRIMARY KEY, + alias TEXT DEFAULT '', + remark TEXT DEFAULT '', + nick_name TEXT DEFAULT '', + description TEXT DEFAULT NULL, + extra_buffer BLOB DEFAULT NULL + ); + CREATE TABLE contact_label ( + label_id_ TEXT, + label_name_ TEXT, + sort_order_ INTEGER + );", + |conn| { + conn.execute( + "INSERT INTO contact (username, alias, remark, nick_name) VALUES (?1, ?2, ?3, ?4)", + params![TALKER, "", "", "Alice"], + ) + .expect("insert contact"); + }, + ); +} + +fn create_encrypted_session_db(path: &Path, raw_key: &[u8; 32]) { + create_encrypted_db( + path, + raw_key, + "CREATE TABLE SessionTable ( + username TEXT, + sort_timestamp INTEGER, + summary TEXT, + last_msg_type INTEGER DEFAULT NULL, + last_msg_sender TEXT DEFAULT NULL, + last_sender_display_name TEXT DEFAULT NULL + );", + |conn| { + conn.execute( + "INSERT INTO SessionTable VALUES (?1, ?2, ?3, ?4, ?5, ?6)", + params![ + TALKER, + 1_700_000_000_i64, + "fixture summary", + None::, + None::, + None::, + ], + ) + .expect("insert session"); + }, + ); +} + +fn create_encrypted_message_db(path: &Path, raw_key: &[u8; 32]) { + create_encrypted_db( + path, + raw_key, + &format!( + "CREATE TABLE Timestamp (timestamp INTEGER); + CREATE TABLE Name2Id ( + rowid INTEGER PRIMARY KEY, + user_name TEXT + ); + CREATE TABLE [{table}] ( + sort_seq INTEGER, + server_id INTEGER, + local_type INTEGER, + real_sender_id INTEGER, + create_time INTEGER, + message_content BLOB, + packed_info_data BLOB, + status INTEGER, + WCDB_CT_message_content INTEGER + );", + table = MSG_TABLE + ), + |conn| { + conn.execute( + "INSERT INTO Timestamp VALUES (?1)", + params![1_700_000_000_i64], + ) + .expect("insert timestamp"); + conn.execute( + "INSERT INTO Name2Id VALUES (?1, ?2)", + params![1_i64, TALKER], + ) + .expect("insert name2id"); + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + table = MSG_TABLE + ), + params![ + 100_i64, + 2001_i64, + 1_i64, + 1_i64, + 1_700_000_100_i64, + b"plain text" as &[u8], + None::>, + 0_i32, + None::, + ], + ) + .expect("insert text message"); + }, + ); +} + +fn create_encrypted_db( + path: &Path, + raw_key: &[u8; 32], + schema_sql: &str, + seed: impl FnOnce(&Connection), +) { + let conn = Connection::open(path).expect("open sqlite"); + unsafe { + let rc = rusqlite::ffi::sqlite3_key( + conn.handle(), + raw_key.as_ptr() as *const std::ffi::c_void, + 32, + ); + assert_eq!(rc, 0, "sqlite3_key failed for {}", path.display()); + } + conn.execute_batch(schema_sql).expect("apply schema"); + seed(&conn); +} + +fn find_open_port() -> u16 { + let listener = std::net::TcpListener::bind("127.0.0.1:0").expect("bind ephemeral port"); + listener.local_addr().expect("listener addr").port() +} diff --git a/crates/wx-cli/tests/thin_client_cli.rs b/crates/wx-cli/tests/thin_client_cli.rs new file mode 100644 index 0000000..4ae928d --- /dev/null +++ b/crates/wx-cli/tests/thin_client_cli.rs @@ -0,0 +1,369 @@ +use std::io::{Read, Write}; +use std::net::TcpListener; +use std::process::Command; +use std::thread; +use std::time::{Duration, Instant}; + +fn bin() -> &'static str { + env!("CARGO_BIN_EXE_wx-cli") +} + +#[test] +fn sessions_can_use_remote_json_path() { + let (base_url, handle) = spawn_sequence_server(2, |request, index| match index { + 0 => { + assert!(request.starts_with("GET /api/v1/health")); + http_response("200 OK", "{\"ready\":true}") + } + 1 => { + assert!(request.starts_with("GET /api/v1/sessions?")); + assert!(request.contains("limit=20")); + assert!(request.contains("offset=0")); + assert!(request.contains("order=desc")); + assert!(!request.contains("show_hidden=")); + assert!(request.contains("Authorization: Bearer secret-token\r\n")); + empty_envelope_response() + } + _ => unreachable!(), + }); + + let output = Command::new(bin()) + .args([ + "sessions", + "--server-only", + "--server-url", + &base_url, + "--server-token", + "secret-token", + "--format", + "json", + ]) + .output() + .expect("run sessions"); + + assert!(output.status.success(), "{output:?}"); + let stdout = String::from_utf8_lossy(&output.stdout); + assert!(stdout.contains("\"items\": []")); + handle.join().unwrap(); +} + +#[test] +fn contacts_can_use_remote_json_path() { + let (base_url, handle) = spawn_sequence_server(2, |request, index| match index { + 0 => http_response("200 OK", "{\"ready\":true}"), + 1 => { + assert!(request.starts_with("GET /api/v1/contacts?")); + assert!(!request.contains("show_hidden=")); + empty_envelope_response() + } + _ => unreachable!(), + }); + + let output = Command::new(bin()) + .args([ + "contacts", + "--server-only", + "--server-url", + &base_url, + "--format", + "json", + ]) + .output() + .expect("run contacts"); + + assert!(output.status.success(), "{output:?}"); + assert!(String::from_utf8_lossy(&output.stdout).contains("\"items\": []")); + handle.join().unwrap(); +} + +#[test] +fn contacts_show_hidden_is_forwarded_only_when_requested() { + let (base_url, handle) = spawn_sequence_server(2, |request, index| match index { + 0 => http_response("200 OK", "{\"ready\":true}"), + 1 => { + assert!(request.starts_with("GET /api/v1/contacts?")); + assert!(request.contains("show_hidden=1")); + empty_envelope_response() + } + _ => unreachable!(), + }); + + let output = Command::new(bin()) + .args([ + "contacts", + "--server-only", + "--server-url", + &base_url, + "--show-hidden", + "--format", + "json", + ]) + .output() + .expect("run contacts"); + + assert!(output.status.success(), "{output:?}"); + handle.join().unwrap(); +} + +#[test] +fn contacts_all_uses_global_max_limit() { + let (base_url, handle) = spawn_sequence_server(2, |request, index| match index { + 0 => http_response("200 OK", "{\"ready\":true}"), + 1 => { + assert!(request.starts_with("GET /api/v1/contacts?")); + assert!(request.contains("limit=20000")); + assert!(request.contains("show_hidden=1")); + empty_envelope_response() + } + _ => unreachable!(), + }); + + let output = Command::new(bin()) + .args([ + "contacts", + "--server-only", + "--server-url", + &base_url, + "--all", + "--show-hidden", + "--format", + "json", + ]) + .output() + .expect("run contacts"); + + assert!(output.status.success(), "{output:?}"); + handle.join().unwrap(); +} + +#[test] +fn query_can_use_remote_json_path() { + let (base_url, handle) = spawn_sequence_server(2, |request, index| match index { + 0 => http_response("200 OK", "{\"ready\":true}"), + 1 => { + assert!(request.starts_with("GET /api/v1/messages?")); + assert!(request.contains("contact=%E5%BC%A0%E4%B8%89")); + assert!(!request.contains("show_hidden=")); + empty_envelope_response() + } + _ => unreachable!(), + }); + + let output = Command::new(bin()) + .args([ + "query", + "张三", + "--server-only", + "--server-url", + &base_url, + "--format", + "json", + ]) + .output() + .expect("run query"); + + assert!(output.status.success(), "{output:?}"); + assert!(String::from_utf8_lossy(&output.stdout).contains("\"items\": []")); + handle.join().unwrap(); +} + +#[test] +fn query_show_hidden_is_forwarded_only_when_requested() { + let (base_url, handle) = spawn_sequence_server(2, |request, index| match index { + 0 => http_response("200 OK", "{\"ready\":true}"), + 1 => { + assert!(request.starts_with("GET /api/v1/messages?")); + assert!(request.contains("contact=%E5%BC%A0%E4%B8%89")); + assert!(request.contains("show_hidden=1")); + empty_envelope_response() + } + _ => unreachable!(), + }); + + let output = Command::new(bin()) + .args([ + "query", + "张三", + "--server-only", + "--server-url", + &base_url, + "--show-hidden", + "--format", + "json", + ]) + .output() + .expect("run query"); + + assert!(output.status.success(), "{output:?}"); + handle.join().unwrap(); +} + +#[test] +fn search_can_use_remote_json_path() { + let (base_url, handle) = spawn_sequence_server(2, |request, index| match index { + 0 => http_response("200 OK", "{\"ready\":true}"), + 1 => { + assert!(request.starts_with("GET /api/v1/search?")); + assert!(request.contains("q=%E5%91%A8%E6%9C%AB")); + assert!(!request.contains("show_hidden=")); + empty_envelope_response() + } + _ => unreachable!(), + }); + + let output = Command::new(bin()) + .args([ + "search", + "周末", + "--server-only", + "--server-url", + &base_url, + "--format", + "json", + ]) + .output() + .expect("run search"); + + assert!(output.status.success(), "{output:?}"); + assert!(String::from_utf8_lossy(&output.stdout).contains("\"items\": []")); + handle.join().unwrap(); +} + +#[test] +fn sessions_show_hidden_is_forwarded_only_when_requested() { + let (base_url, handle) = spawn_sequence_server(2, |request, index| match index { + 0 => http_response("200 OK", "{\"ready\":true}"), + 1 => { + assert!(request.starts_with("GET /api/v1/sessions?")); + assert!(request.contains("show_hidden=1")); + empty_envelope_response() + } + _ => unreachable!(), + }); + + let output = Command::new(bin()) + .args([ + "sessions", + "--server-only", + "--server-url", + &base_url, + "--show-hidden", + "--format", + "json", + ]) + .output() + .expect("run sessions"); + + assert!(output.status.success(), "{output:?}"); + handle.join().unwrap(); +} + +#[test] +fn server_only_fails_when_remote_unavailable() { + let output = Command::new(bin()) + .args([ + "sessions", + "--server-only", + "--server-url", + "http://127.0.0.1:9", + ]) + .output() + .expect("run sessions server-only"); + + assert!(!output.status.success()); + let stderr = String::from_utf8_lossy(&output.stderr); + assert!(stderr.contains("error:")); + assert!(!stderr.contains("falling back to local")); +} + +#[test] +fn unavailable_remote_falls_back_to_local() { + let output = Command::new(bin()) + .args(["sessions", "--server-url", "http://127.0.0.1:9"]) + .output() + .expect("run sessions fallback"); + + assert!(!output.status.success()); + let stderr = String::from_utf8_lossy(&output.stderr); + assert!(stderr.contains("falling back to local sessions")); + assert!(!stderr.contains("/api/v1/health")); + assert!(!stderr.contains("127.0.0.1:9")); + assert!(stderr.contains("no account found")); +} + +#[test] +fn no_server_bypasses_remote_probe() { + let listener = TcpListener::bind("127.0.0.1:0").expect("bind listener"); + listener + .set_nonblocking(true) + .expect("set listener nonblocking"); + let addr = listener.local_addr().expect("listener addr"); + let handle = thread::spawn(move || { + let deadline = Instant::now() + Duration::from_millis(750); + loop { + match listener.accept() { + Ok((_stream, _)) => return true, + Err(err) if err.kind() == std::io::ErrorKind::WouldBlock => { + if Instant::now() >= deadline { + return false; + } + thread::sleep(Duration::from_millis(10)); + } + Err(err) => panic!("accept error: {err}"), + } + } + }); + + let output = Command::new(bin()) + .args([ + "sessions", + "--no-server", + "--server-url", + &format!("http://{}", addr), + ]) + .output() + .expect("run sessions no-server"); + + assert!(!output.status.success()); + let stderr = String::from_utf8_lossy(&output.stderr); + assert!(!stderr.contains("remote server unavailable")); + assert!( + !handle.join().unwrap(), + "no-server should not contact the mock server" + ); +} + +fn spawn_sequence_server( + expected_requests: usize, + responder: impl Fn(String, usize) -> String + Send + 'static, +) -> (String, thread::JoinHandle<()>) { + let listener = TcpListener::bind("127.0.0.1:0").expect("bind mock server"); + let addr = listener.local_addr().expect("mock server addr"); + let handle = thread::spawn(move || { + for index in 0..expected_requests { + let (mut stream, _) = listener.accept().expect("accept"); + let mut buf = [0_u8; 8192]; + let n = stream.read(&mut buf).expect("read request"); + let request = String::from_utf8_lossy(&buf[..n]).into_owned(); + let response = responder(request, index); + stream + .write_all(response.as_bytes()) + .expect("write response"); + } + }); + (format!("http://{}", addr), handle) +} + +fn empty_envelope_response() -> String { + http_response( + "200 OK", + r#"{"items":[],"paging":{"limit":20,"offset":0,"returned":0,"has_more":false,"total":0},"stats":{"scanned":0,"skipped":0,"elapsed_ms":1,"shard_warnings":[]}}"#, + ) +} + +fn http_response(status: &str, body: &str) -> String { + format!( + "HTTP/1.1 {status}\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}", + body.len() + ) +} diff --git a/crates/wx-context/Cargo.toml b/crates/wx-context/Cargo.toml new file mode 100644 index 0000000..12e920b --- /dev/null +++ b/crates/wx-context/Cargo.toml @@ -0,0 +1,27 @@ +[package] +name = "wx-context" +version.workspace = true +edition.workspace = true + +[dependencies] +wx-decrypt = { path = "../wx-decrypt" } +wx-keychain = { path = "../wx-keychain" } +wx-db = { path = "../wx-db" } +wx-paths = { path = "../wx-paths" } +serde = { version = "1", features = ["derive"] } +thiserror = "2" +hex = "0.4" +rusqlite = { version = "0.32", features = ["bundled-sqlcipher"] } +dashmap = "6" +rayon = "1" +rust-stemmers = "1" + +[dev-dependencies] +tempfile = "3" +filetime = "0.2" +aes = "0.8" +cbc = "0.1" +hmac = "0.12" +sha2 = "0.10" +pbkdf2 = { version = "0.12", features = ["hmac"] } +hex = "0.4" diff --git a/crates/wx-context/src/account.rs b/crates/wx-context/src/account.rs new file mode 100644 index 0000000..b492095 --- /dev/null +++ b/crates/wx-context/src/account.rs @@ -0,0 +1,522 @@ +use std::path::{Path, PathBuf}; + +use wx_decrypt::KeyMaterial; + +use crate::ContextError; + +pub struct AccountContext { + pub account_id: String, + pub base_wxid: String, + pub data_dir: PathBuf, + pub key_material: KeyMaterial, + /// The original 32-byte raw key, always populated when `data_key` exists in the store + /// or when `--key` is provided via CLI. + pub raw_key: Option<[u8; 32]>, + /// Whether KDF cache writeback is allowed. False when `--key` CLI flag is used + /// (ephemeral key source), true when key comes from KeyStore. + pub writeback_enabled: bool, + /// Human-readable note about how the account was resolved (e.g. "Auto-detected account: ..."). + /// The CLI layer should print this; the library itself never writes to stderr. + pub detection_note: Option, +} + +/// 解析参数。三种入口互斥,优先级:account > data_dir > 自动检测。 +pub struct ResolveParams<'a> { + pub account: Option<&'a str>, + pub data_dir: Option<&'a Path>, + pub key_hex: Option<&'a str>, +} + +impl AccountContext { + pub fn resolve(params: &ResolveParams<'_>) -> Result { + let (account_id, base_wxid, data_dir, detection_note) = Self::resolve_account(params)?; + let (key_material, raw_key, writeback_enabled) = + Self::resolve_key(params.key_hex, &account_id)?; + Ok(Self { + account_id, + base_wxid, + data_dir, + key_material, + raw_key, + writeback_enabled, + detection_note, + }) + } + + fn resolve_account( + params: &ResolveParams<'_>, + ) -> Result<(String, String, PathBuf, Option), ContextError> { + // 1. --account 指定 → 在已知账号中查找(exact account_id > exact base alias > ambiguity error) + if let Some(acct) = params.account { + let accounts = wx_keychain::find_account_dirs()?; + + // Exact account_id match wins + if let Some(exact) = accounts.iter().find(|a| a.account_id == acct) { + return Ok(( + exact.account_id.clone(), + exact.base_wxid.clone(), + exact.data_dir.clone(), + None, + )); + } + + // Canonical base alias match (e.g. "testuser001" matches dir "testuser001_1662") + let alias_matches: Vec<_> = accounts + .iter() + .filter(|a| { + let id = wx_keychain::AccountId::parse(&a.account_id); + id.matches(acct) + }) + .collect(); + + return match alias_matches.len() { + 1 => Ok(( + alias_matches[0].account_id.clone(), + alias_matches[0].base_wxid.clone(), + alias_matches[0].data_dir.clone(), + None, + )), + 0 => Err(ContextError::NoAccount(format!("'{acct}' not found"))), + _ => { + let candidates = alias_matches + .iter() + .map(|a| format!(" - {} (base: {})", a.account_id, a.base_wxid)) + .collect::>() + .join("\n"); + Err(ContextError::NoAccount(format!( + "ambiguous account '{acct}': multiple directories match\n{candidates}" + ))) + } + }; + } + + // 2. -d 指定 + if let Some(dir) = params.data_dir { + if wx_keychain::is_xwechat_files_root(dir) { + // xwechat_files 根目录 → 自动检测活跃账号 + let accounts = wx_keychain::find_account_dirs_under(dir)?; + if accounts.is_empty() { + return Err(ContextError::NoAccount("no account dirs under root".into())); + } + let active = wx_keychain::detect_active_account(&accounts)?; + let note = format!( + "Auto-detected account: {} (source: {})", + active.info.account_id, active.source + ); + return Ok(( + active.info.account_id, + active.info.base_wxid, + active.info.data_dir, + Some(note), + )); + } + // 直接账号目录 — confirmed directory, use aggressive normalization + let account_id = dir + .file_name() + .and_then(|n| n.to_str()) + .ok_or_else(|| ContextError::NoAccount("invalid data_dir".into()))? + .to_string(); + let base_wxid = dir + .parent() + .map(|root| { + wx_keychain::process::extract_base_wxid_for_account_dir_under_root( + root, + &account_id, + ) + }) + .unwrap_or_else(|| { + wx_keychain::process::extract_base_wxid_for_account_dir(&account_id) + }); + return Ok((account_id, base_wxid, dir.to_path_buf(), None)); + } + + // 3. 全自动检测 + let accounts = wx_keychain::find_account_dirs()?; + if accounts.is_empty() { + return Err(ContextError::NoAccount( + "no WeChat account directories found; use --data-dir or --account".into(), + )); + } + if accounts.len() == 1 { + let note = format!("Auto-detected account: {}", accounts[0].account_id); + let a = accounts.into_iter().next().unwrap(); + return Ok((a.account_id, a.base_wxid, a.data_dir, Some(note))); + } + let active = wx_keychain::detect_active_account(&accounts)?; + let note = format!( + "Auto-detected account: {} (source: {})", + active.info.account_id, active.source + ); + Ok(( + active.info.account_id, + active.info.base_wxid, + active.info.data_dir, + Some(note), + )) + } + + /// Returns (key_material, raw_key, writeback_enabled). + fn resolve_key( + key_hex: Option<&str>, + account_id: &str, + ) -> Result<(KeyMaterial, Option<[u8; 32]>, bool), ContextError> { + // CLI --key flag always produces a RawKey; writeback disabled (ephemeral source). + if let Some(h) = key_hex { + let km = Self::parse_raw_key(h)?; + let raw = match &km { + KeyMaterial::RawKey(k) => Some(*k), + _ => None, + }; + return Ok((km, raw, false)); + } + + // No CLI key — resolve from KeyStore, preferring raw_key when available. + let store = wx_keychain::KeyStore::load_default()?; + let km = Self::resolve_key_from_store(&store, account_id)?; + + // Extract raw_key from store entry's data_key field. + let raw_key = store.get(account_id).and_then(|entry| { + if entry.data_key.is_empty() { + return None; + } + let bytes = hex::decode(&entry.data_key).ok()?; + if bytes.len() != 32 { + return None; + } + let mut key = [0u8; 32]; + key.copy_from_slice(&bytes); + Some(key) + }); + + Ok((km, raw_key, true)) + } + + fn parse_raw_key(hex_str: &str) -> Result { + let bytes = hex::decode(hex_str) + .map_err(|e| ContextError::Cache(format!("invalid hex key: {e}")))?; + if bytes.len() != 32 { + return Err(ContextError::Cache(format!( + "key must be 32 bytes, got {}", + bytes.len() + ))); + } + let mut key = [0u8; 32]; + key.copy_from_slice(&bytes); + Ok(KeyMaterial::RawKey(key)) + } + + /// Resolve key from store, preferring raw_key when available (covers all DBs). + /// + /// The store layer prefers `EncKeys`/`EncKey` (faster decrypt), but at runtime + /// `RawKey` is more universal — it can decrypt any DB without salt matching. + fn resolve_key_from_store( + store: &wx_keychain::KeyStore, + account_id: &str, + ) -> Result { + let entry = store + .get(account_id) + .ok_or_else(|| ContextError::NoKey(account_id.to_string()))?; + + // Prefer raw_key when available — it covers all DBs. + if !entry.data_key.is_empty() { + if let Ok(key_bytes) = hex::decode(&entry.data_key) { + if key_bytes.len() == 32 { + let mut key = [0u8; 32]; + key.copy_from_slice(&key_bytes); + return Ok(KeyMaterial::RawKey(key)); + } + } + } + + // Fall back to store's default resolution (EncKeys > EncKey). + store + .resolve_key_material(account_id) + .ok_or_else(|| ContextError::NoKey(account_id.to_string())) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use wx_decrypt::EncKeyPair; + + #[test] + fn resolve_key_from_store_with_enc_key_only_returns_enc_key() { + let mut store = wx_keychain::KeyStore::default(); + let enc_key = "ab".repeat(32); + let enc_salt = "cd".repeat(16); + store.set_enc_key( + "wxid_test_1234", + &enc_key, + &enc_salt, + "4.1.7.31", + None, + None, + ); + + let km = AccountContext::resolve_key_from_store(&store, "wxid_test_1234").unwrap(); + match km { + KeyMaterial::EncKey { key, salt } => { + assert_eq!(hex::encode(key), enc_key); + assert_eq!(hex::encode(salt), enc_salt); + } + _ => panic!("expected EncKey variant when only enc_key exists"), + } + } + + #[test] + fn resolve_key_from_store_raw_key_preferred_over_enc_key() { + let mut store = wx_keychain::KeyStore::default(); + let data_key = "ef".repeat(32); + let enc_key = "ab".repeat(32); + let enc_salt = "cd".repeat(16); + store.set("wxid_test_1234", &data_key, "4.1.7.31", None, None); + store.set_enc_key( + "wxid_test_1234", + &enc_key, + &enc_salt, + "4.1.7.31", + None, + None, + ); + + let km = AccountContext::resolve_key_from_store(&store, "wxid_test_1234").unwrap(); + match km { + KeyMaterial::RawKey(key) => { + assert_eq!(hex::encode(key), data_key); + } + _ => panic!("expected RawKey variant (raw_key should be preferred)"), + } + } + + #[test] + fn resolve_key_from_store_raw_key_preferred_over_enc_keys() { + let mut store = wx_keychain::KeyStore::default(); + let data_key = "ef".repeat(32); + store.set("wxid_test_1234", &data_key, "4.1.7.31", None, None); + + let pairs = vec![ + EncKeyPair { + key: [0xAAu8; 32], + salt: [0x01u8; 16], + }, + EncKeyPair { + key: [0xBBu8; 32], + salt: [0x02u8; 16], + }, + ]; + store.set_enc_keys("wxid_test_1234", &pairs, "4.1.8.0", None, None); + + let km = AccountContext::resolve_key_from_store(&store, "wxid_test_1234").unwrap(); + match km { + KeyMaterial::RawKey(key) => { + assert_eq!(hex::encode(key), data_key); + } + _ => panic!("expected RawKey variant (raw_key should be preferred over EncKeys)"), + } + } + + #[test] + fn resolve_key_from_store_with_data_key_only_returns_raw_key() { + let mut store = wx_keychain::KeyStore::default(); + let data_key = "ef".repeat(32); + store.set("wxid_test_1234", &data_key, "4.1.7.31", None, None); + + let km = AccountContext::resolve_key_from_store(&store, "wxid_test_1234").unwrap(); + match km { + KeyMaterial::RawKey(key) => { + assert_eq!(hex::encode(key), data_key); + } + _ => panic!("expected RawKey variant"), + } + } + + #[test] + fn resolve_key_from_store_enc_keys_only_returns_enc_keys() { + let mut store = wx_keychain::KeyStore::default(); + let pairs = vec![EncKeyPair { + key: [0xAAu8; 32], + salt: [0x01u8; 16], + }]; + store.set_enc_keys("wxid_test_1234", &pairs, "4.1.8.0", None, None); + + let km = AccountContext::resolve_key_from_store(&store, "wxid_test_1234").unwrap(); + assert!(matches!(km, KeyMaterial::EncKeys(_))); + } + + #[test] + fn resolve_key_with_key_hex_always_returns_raw_key() { + let hex_key = "ab".repeat(32); + let km = AccountContext::parse_raw_key(&hex_key).unwrap(); + match km { + KeyMaterial::RawKey(key) => { + assert_eq!(hex::encode(key), hex_key); + } + _ => panic!("expected RawKey variant"), + } + } + + #[test] + fn resolve_key_from_store_missing_account_returns_error() { + let store = wx_keychain::KeyStore::default(); + let result = AccountContext::resolve_key_from_store(&store, "wxid_nonexistent"); + assert!(result.is_err()); + } + + #[test] + fn resolve_key_raw_key_populated_when_data_key_exists() { + // Test the KeyStore path: when data_key exists in store, + // resolve_key(None, ...) should return raw_key = Some(...) and writeback_enabled = true. + let tmp = tempfile::tempdir().unwrap(); + let config_dir = tmp.path().join(".config").join("wechat-utils"); + std::fs::create_dir_all(&config_dir).unwrap(); + + let data_key = "ab".repeat(32); + let account_id = "wxid_test_kdf_cache"; + let toml_content = format!( + r#"[accounts.{account_id}] +account_id = "{account_id}" +data_key = "{data_key}" +extracted_at = "2026-01-01T00:00:00Z" +wechat_version = "4.1.7.31" +enc_keys = [] +"# + ); + std::fs::write(config_dir.join("keys.toml"), &toml_content).unwrap(); + + // Temporarily override HOME so KeyStore::load_default() finds our temp store. + let original_home = std::env::var("HOME").ok(); + std::env::set_var("HOME", tmp.path()); + // Unset SUDO_USER to avoid getpwnam path + let original_sudo = std::env::var("SUDO_USER").ok(); + std::env::remove_var("SUDO_USER"); + + let result = AccountContext::resolve_key(None, account_id); + + // Restore env + if let Some(h) = original_home { + std::env::set_var("HOME", h); + } + if let Some(s) = original_sudo { + std::env::set_var("SUDO_USER", s); + } + + let (km, raw_key, writeback_enabled) = result.unwrap(); + assert!( + matches!(km, KeyMaterial::RawKey(_)), + "should resolve as RawKey" + ); + assert!( + raw_key.is_some(), + "raw_key should be populated when data_key exists in store" + ); + assert_eq!(hex::encode(raw_key.unwrap()), data_key); + assert!( + writeback_enabled, + "writeback should be true when key comes from KeyStore" + ); + } + + #[test] + fn resolve_key_writeback_disabled_when_key_hex_provided() { + let hex_key = "cd".repeat(32); + let (_km, raw_key, writeback_enabled) = + AccountContext::resolve_key(Some(&hex_key), "wxid_test").unwrap(); + assert!(raw_key.is_some(), "raw_key should be populated from --key"); + assert!( + !writeback_enabled, + "writeback_enabled must be false for CLI --key path" + ); + } + + // ── resolve_account alias matching tests ────────────────────────── + + fn make_account_dir( + root: &std::path::Path, + name: &str, + confirmed_base: Option<&str>, + ) -> wx_keychain::AccountDirInfo { + let dir = root.join(name); + let db_dir = dir.join("db_storage/message"); + std::fs::create_dir_all(&db_dir).unwrap(); + std::fs::write(db_dir.join("message_0.db"), b"fake").unwrap(); + if let Some(base) = confirmed_base { + std::fs::create_dir_all(root.join("all_users/login").join(base)).unwrap(); + } + wx_keychain::find_account_dirs_under(root) + .unwrap() + .into_iter() + .find(|a| a.account_id == name) + .unwrap() + } + + #[test] + fn resolve_account_alias_matches_legacy_dir() { + let tmp = tempfile::tempdir().unwrap(); + let acct = make_account_dir(tmp.path(), "testuser001_1662", Some("testuser001")); + + // base_wxid should be canonicalized for confirmed dir + assert_eq!(acct.base_wxid, "testuser001"); + + // AccountId::matches should find it via the base alias + let id = wx_keychain::AccountId::parse(&acct.account_id); + assert!(id.matches("testuser001"), "alias should match"); + assert!(id.matches("testuser001_1662"), "raw should match"); + } + + #[test] + fn resolve_account_alias_ambiguity_detected() { + let tmp = tempfile::tempdir().unwrap(); + // Two directories that share the same base alias + make_account_dir(tmp.path(), "user123_ab12", None); + make_account_dir(tmp.path(), "user123_cd34", None); + + let accounts = wx_keychain::find_account_dirs_under(tmp.path()).unwrap(); + assert_eq!(accounts.len(), 2); + + // Both should match "user123" + let matches: Vec<_> = accounts + .iter() + .filter(|a| { + let id = wx_keychain::AccountId::parse(&a.account_id); + id.matches("user123") + }) + .collect(); + + assert_eq!( + matches.len(), + 2, + "ambiguity: two dirs share the same base alias" + ); + } + + #[test] + fn resolve_account_exact_id_takes_priority() { + let tmp = tempfile::tempdir().unwrap(); + make_account_dir(tmp.path(), "wxid_foo123_ab12", None); + + let accounts = wx_keychain::find_account_dirs_under(tmp.path()).unwrap(); + let exact = accounts.iter().find(|a| a.account_id == "wxid_foo123_ab12"); + assert!(exact.is_some(), "exact account_id match should be found"); + } + + #[test] + fn resolve_account_data_dir_keeps_legacy_dir_without_login_hint() { + let tmp = tempfile::tempdir().unwrap(); + let dir = tmp.path().join("testuser001_1662"); + std::fs::create_dir_all(&dir).unwrap(); + + let (account_id, base_wxid, resolved_dir, note) = + AccountContext::resolve_account(&ResolveParams { + account: None, + data_dir: Some(&dir), + key_hex: None, + }) + .unwrap(); + + assert_eq!(account_id, "testuser001_1662"); + assert_eq!(base_wxid, "testuser001_1662"); + assert_eq!(resolved_dir, dir); + assert_eq!(note, None); + } +} diff --git a/crates/wx-context/src/cache.rs b/crates/wx-context/src/cache.rs new file mode 100644 index 0000000..d8366bd --- /dev/null +++ b/crates/wx-context/src/cache.rs @@ -0,0 +1,1286 @@ +use std::path::{Path, PathBuf}; +use std::sync::{Arc, Mutex}; +use std::time::SystemTime; + +use wx_decrypt::KeyMaterial; + +use crate::account::AccountContext; +use crate::file_lock::FileLockMap; +use crate::kdf_cache::KdfCache; +use crate::patch_state::{clear_wal_failed_marker, is_wal_patch_failed, write_wal_failed_marker}; +use crate::progress::{DbOutcome, DecryptProgress, DecryptStats, WalOutcome}; +use crate::wal_patch::{apply_wal_patch, WalPatchResult}; +use crate::ContextError; + +pub struct PersistentCache { + cache_root: PathBuf, + encrypted_root: PathBuf, + raw_key: Option<[u8; 32]>, + account_id: String, + base_wxid: String, + writeback_enabled: bool, + kdf_cache: Mutex, + params: &'static wx_decrypt::CryptoParams, + file_locks: Arc, +} + +impl PersistentCache { + pub fn new( + account: &AccountContext, + params: &'static wx_decrypt::CryptoParams, + ) -> Result { + let cache_base = wx_paths::AppPaths::new() + .map_err(|e| ContextError::Cache(e.to_string()))? + .account_cache_dir(&account.account_id); + + let encrypted_root = account.data_dir.join("db_storage"); + if !encrypted_root.exists() { + return Err(ContextError::Cache(format!( + "db_storage not found in {}", + account.data_dir.display() + ))); + } + + // Initialize KdfCache from stored enc_keys (independent of account.key_material). + let kdf_cache = match wx_keychain::KeyStore::load_default() { + Ok(store) => Self::load_kdf_cache_from_store(&store, &account.account_id), + Err(_) => KdfCache::empty(), // Best-effort: don't block on KeyStore failure + }; + + Ok(Self { + cache_root: cache_base.join("db_storage"), + encrypted_root, + raw_key: account.raw_key, + account_id: account.account_id.clone(), + base_wxid: account.base_wxid.clone(), + writeback_enabled: account.writeback_enabled, + kdf_cache: Mutex::new(kdf_cache), + params, + file_locks: Arc::new(FileLockMap::default()), + }) + } + + /// 返回解密后的 db_storage 根路径。 + pub fn decrypted_root(&self) -> &Path { + &self.cache_root + } + + /// Test-only constructor that bypasses AccountContext resolution. + #[doc(hidden)] + pub fn new_for_test( + cache_root: PathBuf, + encrypted_root: PathBuf, + raw_key: Option<[u8; 32]>, + params: &'static wx_decrypt::CryptoParams, + ) -> Self { + Self { + cache_root, + encrypted_root, + raw_key, + account_id: String::new(), + base_wxid: String::new(), + writeback_enabled: false, + kdf_cache: Mutex::new(KdfCache::empty()), + params, + file_locks: Arc::new(FileLockMap::default()), + } + } + + /// 确保所有 DB 已解密到缓存目录,按 .mtime 标记跳过未变化的文件。 + pub fn ensure_decrypted(&self) -> Result { + self.ensure_decrypted_scoped(&crate::DecryptScope::All, |_| {}) + } + + /// Like `ensure_decrypted`, but fires progress events via the callback. + pub fn ensure_decrypted_with_progress( + &self, + on_progress: impl Fn(DecryptProgress) + Send + Sync, + ) -> Result { + self.ensure_decrypted_scoped(&crate::DecryptScope::All, on_progress) + } + + /// Decrypt only the databases matching `scope`. + pub fn ensure_decrypted_scoped( + &self, + scope: &crate::DecryptScope, + on_progress: impl Fn(DecryptProgress) + Send + Sync, + ) -> Result { + use crate::progress::AtomicStats; + use rayon::prelude::*; + + // Phase 1: Discover + let all_db_files = crate::db_category::discover_db_files(&self.encrypted_root)?; + if all_db_files.is_empty() { + return Err(ContextError::Cache("no .db files in db_storage/".into())); + } + + // Phase 2: Filter + let db_files: Vec<_> = all_db_files + .into_iter() + .filter(|db| scope.matches(db)) + .collect(); + + if db_files.is_empty() { + return Ok(DecryptStats { + decrypted: 0, + skipped: 0, + errors: 0, + wal_patched: 0, + warnings: Vec::new(), + }); + } + + // Phase 3: Parallel execute + let atomic_stats = AtomicStats::new(); + + db_files.par_iter().for_each(|db| { + let rel = db + .path + .strip_prefix(&self.encrypted_root) + .expect("db path discovered by discover_db_files must be under encrypted_root"); + let dst = self.cache_root.join(rel); + let rel_str = rel.display().to_string(); + + match self.decrypt_one_db(&db.path, &dst, rel) { + DbOutcome::Decrypted { + wal_patched, + wal_warnings, + } => { + atomic_stats.inc_decrypted(); + if wal_patched { + atomic_stats.inc_wal_patched(); + } + for w in wal_warnings { + atomic_stats.add_warning(w); + } + on_progress(DecryptProgress::Decrypted { + path: rel_str, + wal_patched, + }); + } + DbOutcome::Skipped { + wal_patched, + wal_warnings, + } => { + atomic_stats.inc_skipped(); + if wal_patched { + atomic_stats.inc_wal_patched(); + } + for w in wal_warnings { + atomic_stats.add_warning(w); + } + on_progress(DecryptProgress::Skipped { + path: rel_str, + wal_patched, + }); + } + DbOutcome::Failed { warning } => { + atomic_stats.inc_errors(); + let error = warning.clone(); + atomic_stats.add_warning(warning); + on_progress(DecryptProgress::Failed { + path: rel_str, + error, + }); + } + } + }); + + // Phase 4: Aggregate + let mut stats = atomic_stats.into_stats(); + + if stats.errors > 0 && stats.decrypted == 0 && stats.skipped == 0 { + return Err(ContextError::Cache(format!( + "all {} database files failed to decrypt", + stats.errors + ))); + } + + // Write back newly-derived enc_keys to keys.toml (best-effort). + if let Err(e) = self.writeback_enc_keys() { + stats + .warnings + .push(format!("enc_keys writeback failed: {e}")); + } + + Ok(stats) + } + + /// Per-DB decrypt logic. Acquires per-file lock, performs authoritative needs_decrypt + /// check inside the lock, decrypts if needed, and applies WAL patch. + fn decrypt_one_db(&self, src: &Path, dst: &Path, rel: &Path) -> DbOutcome { + // Acquire per-file lock before any needs_decrypt check (P7 dedup) + let lock = self.file_locks.lock_for(dst); + let _guard = lock.lock().unwrap(); + + // Authoritative needs_decrypt check inside the lock + if !self.needs_decrypt(src, dst) { + // DB unchanged — but check if WAL was updated independently + let (wal_patched, wal_warnings) = self.try_wal_patch(src, dst, rel); + return DbOutcome::Skipped { + wal_patched, + wal_warnings, + }; + } + + if let Some(parent) = dst.parent() { + if let Err(e) = std::fs::create_dir_all(parent) { + return DbOutcome::Failed { + warning: format!("mkdir err {}: {e}", rel.display()), + }; + } + } + + match self.decrypt_db(src, dst) { + Ok(()) => { + if let Err(e) = self.write_mtime_marker(src, dst) { + return DbOutcome::Failed { + warning: format!("mtime marker err {}: {e}", rel.display()), + }; + } + let (wal_patched, wal_warnings) = self.try_wal_patch(src, dst, rel); + DbOutcome::Decrypted { + wal_patched, + wal_warnings, + } + } + Err(wx_decrypt::DecryptError::AlreadyDecrypted) => { + if let Err(e) = std::fs::copy(src, dst) { + return DbOutcome::Failed { + warning: format!("copy err {}: {e}", rel.display()), + }; + } + if let Err(e) = self.write_mtime_marker(src, dst) { + return DbOutcome::Failed { + warning: format!("mtime marker err {}: {e}", rel.display()), + }; + } + DbOutcome::Decrypted { + wal_patched: false, + wal_warnings: Vec::new(), + } + } + Err(e) => DbOutcome::Failed { + warning: format!("decrypt err {}: {e}", rel.display()), + }, + } + } + + /// Try WAL patch if WAL file exists. Returns (patched, warnings). + /// Caller must hold the per-file lock. + fn try_wal_patch(&self, src: &Path, dst: &Path, rel: &Path) -> (bool, Vec) { + let wal = src.with_extension("db-wal"); + if !wal.exists() { + return (false, Vec::new()); + } + match self.guarded_wal_patch(src, &wal, dst, rel) { + WalOutcome::Patched => (true, Vec::new()), + WalOutcome::NothingToDo => (false, Vec::new()), + WalOutcome::Warning(w) => (false, vec![w]), + } + } + + /// Apply WAL patch with failed-marker awareness. Caller must hold the per-file lock. + /// Returns `WalOutcome` instead of mutating stats directly. + fn guarded_wal_patch(&self, src: &Path, wal: &Path, dst: &Path, rel: &Path) -> WalOutcome { + // Check failed marker — skip if previous attempt failed with same mtimes + if is_wal_patch_failed(src, wal, dst) { + return WalOutcome::Warning(format!( + "WAL skipped (previously failed) {}", + rel.display() + )); + } + + // Check mtime — skip if WAL hasn't changed since last successful patch + if !self.needs_wal_patch(wal, dst) { + return WalOutcome::NothingToDo; + } + + // Resolve enc_key for this WAL's parent DB salt via KdfCache. + let salt = match wx_decrypt::read_main_db_salt_for_path(wal) { + Ok(s) => s, + Err(wx_decrypt::DecryptError::AlreadyDecrypted) => { + return WalOutcome::NothingToDo; + } + Err(e) => { + return WalOutcome::Warning(format!( + "WAL skipped (salt read err: {e}) {}", + rel.display() + )); + } + }; + + let km = { + let mut cache_guard = self.kdf_cache.lock().unwrap(); + match (cache_guard.lookup(&salt), &self.raw_key) { + (Some(cached), _) => KeyMaterial::EncKey { key: cached, salt }, + (None, Some(raw)) => { + let enc_key = cache_guard.get_or_derive(&salt, raw, self.params); + KeyMaterial::EncKey { key: enc_key, salt } + } + (None, None) => { + return WalOutcome::Warning(format!( + "WAL skipped (no raw_key for derivation) {}", + rel.display() + )); + } + } + }; + + self.apply_wal_and_handle_retry(src, wal, dst, rel, &salt, &km) + } + + /// Core WAL patch logic with stale-key retry. Extracted to avoid duplication. + /// Translate a `WalPatchResult` into a `WalOutcome`, handling bookkeeping + /// (clearing/setting WAL markers, writing mtime markers). + fn handle_wal_result( + &self, + result: WalPatchResult, + src: &Path, + wal: &Path, + dst: &Path, + rel: &Path, + ) -> WalOutcome { + match result { + WalPatchResult::Patched(n) => { + clear_wal_failed_marker(dst); + let _ = self.write_wal_mtime_marker(wal, dst); + if n > 0 { + WalOutcome::Patched + } else { + WalOutcome::NothingToDo + } + } + WalPatchResult::NoFrames => { + clear_wal_failed_marker(dst); + let _ = self.write_wal_mtime_marker(wal, dst); + WalOutcome::NothingToDo + } + WalPatchResult::ContentFailed(e) => { + let _ = write_wal_failed_marker(src, wal, dst); + WalOutcome::Warning(format!("WAL content err {}: {e}", rel.display())) + } + WalPatchResult::IoFailed(e) => { + WalOutcome::Warning(format!("WAL IO err {}: {e}", rel.display())) + } + } + } + + fn apply_wal_and_handle_retry( + &self, + src: &Path, + wal: &Path, + dst: &Path, + rel: &Path, + salt: &[u8; 16], + km: &KeyMaterial, + ) -> WalOutcome { + let result = apply_wal_patch(wal, dst, km, self.params); + match result { + WalPatchResult::ContentFailed(ref e) + if e.contains("incorrect key") && self.raw_key.is_some() => + { + // Stale enc_key — re-derive from raw_key and retry once + let raw = self.raw_key.as_ref().unwrap(); + let fresh_enc_key = wx_decrypt::kdf::derive_enc_key(raw, salt, self.params); + { + let mut cache_guard = self.kdf_cache.lock().unwrap(); + cache_guard.insert(salt, &fresh_enc_key); + } + let fresh_km = KeyMaterial::EncKey { + key: fresh_enc_key, + salt: *salt, + }; + self.apply_wal_final(src, wal, dst, rel, &fresh_km) + } + _ => self.handle_wal_result(result, src, wal, dst, rel), + } + } + + /// Final WAL patch attempt (after retry). No further retries. + fn apply_wal_final( + &self, + src: &Path, + wal: &Path, + dst: &Path, + rel: &Path, + km: &KeyMaterial, + ) -> WalOutcome { + let result = apply_wal_patch(wal, dst, km, self.params); + self.handle_wal_result(result, src, wal, dst, rel) + } + + fn decrypt_db(&self, src: &Path, dst: &Path) -> Result<(), wx_decrypt::DecryptError> { + let salt = wx_decrypt::read_db_salt(src)?; + + let enc_key = { + let mut cache_guard = self.kdf_cache.lock().unwrap(); + match (cache_guard.lookup(&salt), &self.raw_key) { + (Some(cached), _) => cached, + (None, Some(raw)) => cache_guard.get_or_derive(&salt, raw, self.params), + (None, None) => return Err(wx_decrypt::DecryptError::NoMatchingEncKey), + } + }; // Mutex guard dropped here + + // Try direct decrypt with cached/derived enc_key + match wx_decrypt::decrypt_db_direct(src, dst, &enc_key, &salt, self.params) { + Ok(()) => Ok(()), + Err(wx_decrypt::DecryptError::IncorrectKey) if self.raw_key.is_some() => { + // Stale enc_key — re-derive from raw_key and retry once + let raw = self.raw_key.as_ref().unwrap(); + let fresh_enc_key = wx_decrypt::kdf::derive_enc_key(raw, &salt, self.params); + { + let mut cache_guard = self.kdf_cache.lock().unwrap(); + cache_guard.insert(&salt, &fresh_enc_key); + } + wx_decrypt::decrypt_db_direct(src, dst, &fresh_enc_key, &salt, self.params) + } + Err(e) => Err(e), + } + } + + fn writeback_enc_keys(&self) -> Result<(), ContextError> { + if !self.writeback_enabled { + return Ok(()); + } + + let cache_guard = self.kdf_cache.lock().unwrap(); + if !cache_guard.has_new_derivations() { + return Ok(()); + } + + let session_pairs = cache_guard.all_pairs(); + drop(cache_guard); // Release before file I/O + + // Union merge with existing enc_keys on disk + let mut store = wx_keychain::KeyStore::load_default()?; + let mut merged_pairs = session_pairs; + if let Some(entry) = store.get(&self.account_id) { + for existing in &entry.enc_keys { + if let (Ok(key_bytes), Ok(salt_bytes)) = + (hex::decode(&existing.enc_key), hex::decode(&existing.salt)) + { + if key_bytes.len() == 32 && salt_bytes.len() == 16 { + let mut key = [0u8; 32]; + let mut salt = [0u8; 16]; + key.copy_from_slice(&key_bytes); + salt.copy_from_slice(&salt_bytes); + merged_pairs.push(wx_decrypt::EncKeyPair { key, salt }); + } + } + } + } + let version = store + .get(&self.account_id) + .map(|e| e.wechat_version.clone()) + .unwrap_or_default(); + // Opportunistically repair base_wxid with canonical value + store.set_enc_keys( + &self.account_id, + &merged_pairs, + &version, + None, + Some(self.base_wxid.clone()), + ); + store.save_default()?; + Ok(()) + } + + fn load_kdf_cache_from_store(store: &wx_keychain::KeyStore, account_id: &str) -> KdfCache { + let entry = match store.get(account_id) { + Some(e) => e, + None => return KdfCache::empty(), + }; + + // Try new per-DB enc_keys format first. + if !entry.enc_keys.is_empty() { + let pairs: Vec = entry + .enc_keys + .iter() + .filter_map(|e| { + let key_bytes = hex::decode(&e.enc_key).ok()?; + let salt_bytes = hex::decode(&e.salt).ok()?; + if key_bytes.len() == 32 && salt_bytes.len() == 16 { + let mut key = [0u8; 32]; + let mut salt = [0u8; 16]; + key.copy_from_slice(&key_bytes); + salt.copy_from_slice(&salt_bytes); + Some(wx_decrypt::EncKeyPair { key, salt }) + } else { + None + } + }) + .collect(); + if !pairs.is_empty() { + return KdfCache::from_pairs(&pairs); + } + } + + // Legacy single enc_key path. + if let (Some(ek), Some(es)) = (&entry.enc_key, &entry.enc_key_salt) { + if !ek.is_empty() && !es.is_empty() { + if let (Ok(key_bytes), Ok(salt_bytes)) = (hex::decode(ek), hex::decode(es)) { + if key_bytes.len() == 32 && salt_bytes.len() == 16 { + let mut key = [0u8; 32]; + let mut salt = [0u8; 16]; + key.copy_from_slice(&key_bytes); + salt.copy_from_slice(&salt_bytes); + return KdfCache::from_pairs(&[wx_decrypt::EncKeyPair { key, salt }]); + } + } + } + } + + KdfCache::empty() + } + + /// 检查源文件是否需要重新解密。 + fn needs_decrypt(&self, src: &Path, dst: &Path) -> bool { + if !dst.exists() { + return true; + } + let marker = mtime_marker_path(dst); + let recorded = match std::fs::read_to_string(&marker) { + Ok(s) => s, + Err(_) => return true, + }; + let src_mtime = match src.metadata().and_then(|m| m.modified()) { + Ok(t) => format_system_time(t), + Err(_) => return false, + }; + recorded.trim() != src_mtime + } + + fn write_mtime_marker(&self, src: &Path, dst: &Path) -> Result<(), ContextError> { + let src_mtime = src.metadata()?.modified()?; + let marker = mtime_marker_path(dst); + std::fs::write(&marker, format_system_time(src_mtime))?; + Ok(()) + } + + /// Check if the WAL file has changed since the last patch. + fn needs_wal_patch(&self, wal: &Path, dst: &Path) -> bool { + let marker = wal_mtime_marker_path(dst); + let recorded = match std::fs::read_to_string(&marker) { + Ok(s) => s, + Err(_) => return true, + }; + let wal_mtime = match wal.metadata().and_then(|m| m.modified()) { + Ok(t) => format_system_time(t), + Err(_) => return false, + }; + recorded.trim() != wal_mtime + } + + fn write_wal_mtime_marker(&self, wal: &Path, dst: &Path) -> Result<(), ContextError> { + let wal_mtime = wal.metadata()?.modified()?; + let marker = wal_mtime_marker_path(dst); + std::fs::write(&marker, format_system_time(wal_mtime))?; + Ok(()) + } +} + +fn mtime_marker_path(dst: &Path) -> PathBuf { + dst.with_extension(format!( + "{}.mtime", + dst.extension().unwrap_or_default().to_string_lossy() + )) +} + +fn wal_mtime_marker_path(dst: &Path) -> PathBuf { + dst.with_extension("db.wal_mtime") +} + +pub(crate) fn format_system_time(t: SystemTime) -> String { + t.duration_since(SystemTime::UNIX_EPOCH) + .map(|d| d.as_nanos().to_string()) + .unwrap_or_default() +} + +#[cfg(test)] +mod tests { + use super::*; + use tempfile::TempDir; + + /// Test helper: construct a PersistentCache with optional pre-loaded KdfCache entries. + fn test_cache( + cache_root: PathBuf, + encrypted_root: PathBuf, + raw_key: Option<[u8; 32]>, + kdf_pairs: &[wx_decrypt::EncKeyPair], + params: &'static wx_decrypt::CryptoParams, + ) -> PersistentCache { + let kdf_cache = if kdf_pairs.is_empty() { + KdfCache::empty() + } else { + KdfCache::from_pairs(kdf_pairs) + }; + PersistentCache { + cache_root, + encrypted_root, + raw_key, + account_id: String::new(), + base_wxid: String::new(), + writeback_enabled: false, + kdf_cache: Mutex::new(kdf_cache), + params, + file_locks: Arc::new(FileLockMap::default()), + } + } + + #[test] + fn needs_decrypt_no_marker() { + let tmp = TempDir::new().unwrap(); + let src = tmp.path().join("test.db"); + let dst = tmp.path().join("out/test.db"); + std::fs::create_dir_all(tmp.path().join("out")).unwrap(); + std::fs::write(&src, b"data").unwrap(); + std::fs::write(&dst, b"cached").unwrap(); + let cache = test_cache( + tmp.path().join("out"), + tmp.path().to_path_buf(), + None, + &[], + &wx_decrypt::MACOS_4_1_7_31, + ); + assert!(cache.needs_decrypt(&src, &dst)); + } + + #[test] + fn needs_decrypt_matching_marker() { + let tmp = TempDir::new().unwrap(); + let src = tmp.path().join("test.db"); + let dst = tmp.path().join("out/test.db"); + std::fs::create_dir_all(tmp.path().join("out")).unwrap(); + std::fs::write(&src, b"data").unwrap(); + std::fs::write(&dst, b"cached").unwrap(); + let mtime = src.metadata().unwrap().modified().unwrap(); + let marker = mtime_marker_path(&dst); + std::fs::write(&marker, format_system_time(mtime)).unwrap(); + + let cache = test_cache( + tmp.path().join("out"), + tmp.path().to_path_buf(), + None, + &[], + &wx_decrypt::MACOS_4_1_7_31, + ); + assert!(!cache.needs_decrypt(&src, &dst)); + } + + #[test] + fn mtime_marker_path_format() { + let p = PathBuf::from("/cache/message_0.db"); + assert_eq!( + mtime_marker_path(&p), + PathBuf::from("/cache/message_0.db.mtime") + ); + } + + #[test] + fn enc_keys_decrypt_db_selects_matching_pair() { + use wx_decrypt::EncKeyPair; + + // Build two encrypted DBs with different salts, same raw key + let raw_key = [0xABu8; 32]; + let salt1 = [0x01u8; 16]; + let salt2 = [0x02u8; 16]; + let params = &wx_decrypt::MACOS_4_1_7_31; + + let enc_key1 = derive_enc_key(&raw_key, &salt1, params); + let enc_key2 = derive_enc_key(&raw_key, &salt2, params); + + let tmp = TempDir::new().unwrap(); + let enc_root = tmp.path().join("encrypted"); + let cache_root = tmp.path().join("cache"); + let sub = enc_root.join("sub"); + std::fs::create_dir_all(&sub).unwrap(); + std::fs::create_dir_all(cache_root.join("sub")).unwrap(); + + build_encrypted_db(&sub.join("a.db"), &raw_key, &salt1, params); + build_encrypted_db(&sub.join("b.db"), &raw_key, &salt2, params); + + let cache = test_cache( + cache_root.clone(), + enc_root.clone(), + None, + &[ + EncKeyPair { + key: enc_key1, + salt: salt1, + }, + EncKeyPair { + key: enc_key2, + salt: salt2, + }, + ], + params, + ); + + // decrypt a.db (salt1 → enc_key1) + let src_a = sub.join("a.db"); + let dst_a = cache_root.join("sub").join("a.db"); + cache.decrypt_db(&src_a, &dst_a).unwrap(); + let data = std::fs::read(&dst_a).unwrap(); + assert_eq!( + &data[..16], + b"SQLite format 3\0", + "a.db should be valid SQLite" + ); + + // decrypt b.db (salt2 → enc_key2) + let src_b = sub.join("b.db"); + let dst_b = cache_root.join("sub").join("b.db"); + cache.decrypt_db(&src_b, &dst_b).unwrap(); + let data = std::fs::read(&dst_b).unwrap(); + assert_eq!( + &data[..16], + b"SQLite format 3\0", + "b.db should be valid SQLite" + ); + } + + #[test] + fn enc_keys_decrypt_db_no_match_returns_error() { + use wx_decrypt::EncKeyPair; + + let raw_key = [0xABu8; 32]; + let salt_db = [0x01u8; 16]; + let salt_wrong = [0x99u8; 16]; + let params = &wx_decrypt::MACOS_4_1_7_31; + + let enc_key_wrong = derive_enc_key(&raw_key, &salt_wrong, params); + + let tmp = TempDir::new().unwrap(); + let enc_root = tmp.path().join("encrypted"); + std::fs::create_dir_all(&enc_root).unwrap(); + build_encrypted_db(&enc_root.join("test.db"), &raw_key, &salt_db, params); + + let cache = test_cache( + tmp.path().join("cache"), + enc_root.clone(), + None, + &[EncKeyPair { + key: enc_key_wrong, + salt: salt_wrong, + }], + params, + ); + + let err = cache + .decrypt_db( + &enc_root.join("test.db"), + &tmp.path().join("cache").join("test.db"), + ) + .unwrap_err(); + assert!(matches!( + err, + wx_decrypt::DecryptError::NoMatchingEncKey + )); + } + + #[test] + fn guarded_wal_patch_skips_when_failed_marker_exists() { + use crate::patch_state::{ + is_wal_patch_failed, wal_failed_marker_path, write_wal_failed_marker, + }; + + let tmp = TempDir::new().unwrap(); + let enc_root = tmp.path().join("encrypted"); + let cache_root = tmp.path().join("cache"); + std::fs::create_dir_all(&enc_root).unwrap(); + std::fs::create_dir_all(&cache_root).unwrap(); + + let src = enc_root.join("test.db"); + let wal = enc_root.join("test.db-wal"); + let dst = cache_root.join("test.db"); + let rel = Path::new("test.db"); + + std::fs::write(&src, b"db-data").unwrap(); + std::fs::write(&wal, b"wal-data").unwrap(); + std::fs::write(&dst, b"cached-db").unwrap(); + + // Write a failed marker for the current mtimes + write_wal_failed_marker(&src, &wal, &dst).unwrap(); + assert!(is_wal_patch_failed(&src, &wal, &dst)); + + let cache = test_cache( + cache_root.clone(), + enc_root.clone(), + None, + &[], + &wx_decrypt::MACOS_4_1_7_31, + ); + + let outcome = cache.guarded_wal_patch(&src, &wal, &dst, rel); + + // Should have been skipped with a warning + match outcome { + WalOutcome::Warning(w) => assert!(w.contains("previously failed")), + other => panic!("expected Warning, got {other:?}"), + } + + // The failed marker should still exist + assert!(wal_failed_marker_path(&dst).exists()); + } + + #[test] + fn guarded_wal_patch_retries_after_mtime_change() { + use crate::patch_state::{is_wal_patch_failed, write_wal_failed_marker}; + use std::time::Duration; + + let tmp = TempDir::new().unwrap(); + let enc_root = tmp.path().join("encrypted"); + let cache_root = tmp.path().join("cache"); + std::fs::create_dir_all(&enc_root).unwrap(); + std::fs::create_dir_all(&cache_root).unwrap(); + + let src = enc_root.join("test.db"); + let wal = enc_root.join("test.db-wal"); + let dst = cache_root.join("test.db"); + + std::fs::write(&src, b"db-data").unwrap(); + std::fs::write(&wal, b"wal-data").unwrap(); + std::fs::write(&dst, b"cached-db").unwrap(); + + // Write a failed marker for the current mtimes + write_wal_failed_marker(&src, &wal, &dst).unwrap(); + assert!(is_wal_patch_failed(&src, &wal, &dst)); + + // Advance the WAL mtime — the failed marker should no longer match + let new_mtime = filetime::FileTime::from_system_time( + std::time::SystemTime::now() + Duration::from_secs(2), + ); + filetime::set_file_mtime(&wal, new_mtime).unwrap(); + + assert!(!is_wal_patch_failed(&src, &wal, &dst)); + // This proves that guarded_wal_patch would proceed past the failed-marker check + // and attempt the actual patch (which would fail on our dummy data, but that's + // testing apply_wal_patch, not the guard logic). + } + + #[test] + fn guarded_wal_patch_noframes_clears_failed_marker() { + use crate::patch_state::{ + is_wal_patch_failed, wal_failed_marker_path, write_wal_failed_marker, + }; + use std::time::Duration; + + let tmp = TempDir::new().unwrap(); + let enc_root = tmp.path().join("encrypted"); + let cache_root = tmp.path().join("cache"); + std::fs::create_dir_all(&enc_root).unwrap(); + std::fs::create_dir_all(&cache_root).unwrap(); + + let src = enc_root.join("test.db"); + let wal = enc_root.join("test.db-wal"); + let dst = cache_root.join("test.db"); + let rel = Path::new("test.db"); + + // Use pre-loaded KdfCache with enc_key for salt [0u8; 16] (the src file's salt). + // src must be at least page_size (4096) for read_db_salt, but for the WAL patch + // we use read_main_db_salt_for_path which reads from the .db companion. + // The WAL's parent DB is "test.db" → salt comes from src's first 16 bytes. + std::fs::write(&src, [0u8; 4096]).unwrap(); + std::fs::write(&dst, b"cached-db").unwrap(); + + // Minimal valid WAL: 32-byte header, zero frames → dispatch_decrypt_wal returns Ok(0) + let mut wal_header = vec![0u8; 32]; + wal_header[..4].copy_from_slice(&0x377f0682u32.to_be_bytes()); // WAL magic + wal_header[4..8].copy_from_slice(&3007000u32.to_be_bytes()); // version + wal_header[8..12].copy_from_slice(&4096u32.to_be_bytes()); // page size + std::fs::write(&wal, &wal_header).unwrap(); + + // Write a wal_failed marker for current mtimes + write_wal_failed_marker(&src, &wal, &dst).unwrap(); + assert!(is_wal_patch_failed(&src, &wal, &dst)); + + // Advance WAL mtime so the marker becomes stale (otherwise guarded_wal_patch skips) + let new_mtime = filetime::FileTime::from_system_time( + std::time::SystemTime::now() + Duration::from_secs(2), + ); + filetime::set_file_mtime(&wal, new_mtime).unwrap(); + assert!(!is_wal_patch_failed(&src, &wal, &dst)); + + let cache = test_cache( + cache_root.clone(), + enc_root.clone(), + None, + &[wx_decrypt::EncKeyPair { + key: [0u8; 32], + salt: [0u8; 16], + }], + &wx_decrypt::MACOS_4_1_7_31, + ); + + let outcome = cache.guarded_wal_patch(&src, &wal, &dst, rel); + + // NoFrames should have cleared the failed marker + assert!( + !wal_failed_marker_path(&dst).exists(), + "wal_failed marker should have been cleared by NoFrames path" + ); + // Should be NothingToDo (no warnings) + assert!( + matches!(outcome, WalOutcome::NothingToDo), + "expected NothingToDo, got {outcome:?}", + ); + } + + // --- test helper --- + fn derive_enc_key( + raw_key: &[u8; 32], + salt: &[u8; 16], + params: &wx_decrypt::CryptoParams, + ) -> [u8; 32] { + let mut key = [0u8; 32]; + pbkdf2::pbkdf2_hmac::(raw_key, salt, params.kdf_iter, &mut key); + key + } + + fn build_encrypted_db( + path: &std::path::Path, + raw_key: &[u8; 32], + salt: &[u8; 16], + params: &wx_decrypt::CryptoParams, + ) { + use aes::cipher::{BlockEncryptMut, KeyIvInit}; + use hmac::{Hmac, Mac}; + use sha2::Sha512; + + let enc_key = derive_enc_key(raw_key, salt, params); + let mut mac_salt = [0u8; 16]; + for (i, b) in salt.iter().enumerate() { + mac_salt[i] = b ^ 0x3a; + } + let mut mac_key = [0u8; 32]; + pbkdf2::pbkdf2_hmac::(&enc_key, &mac_salt, 2, &mut mac_key); + + let iv = [0x42u8; 16]; + let data_size = params.page_size - params.reserve - params.salt_size; + let plaintext = vec![0u8; data_size]; + + let mut ciphertext = plaintext; + type Aes256CbcEnc = cbc::Encryptor; + Aes256CbcEnc::new((&enc_key).into(), (&iv).into()) + .encrypt_padded_mut::(&mut ciphertext, data_size) + .unwrap(); + + let mut page = Vec::with_capacity(params.page_size); + page.extend_from_slice(salt); + page.extend_from_slice(&ciphertext); + page.extend_from_slice(&iv); + page.resize(params.page_size, 0); + + let hmac_data_end = params.page_size - params.reserve + params.iv_size; + let mut mac = as Mac>::new_from_slice(&mac_key).unwrap(); + mac.update(&page[params.salt_size..hmac_data_end]); + mac.update(&1u32.to_le_bytes()); + let hmac_result = mac.finalize().into_bytes(); + let hmac_start = params.page_size - params.reserve + params.iv_size; + page[hmac_start..hmac_start + params.hmac_size] + .copy_from_slice(&hmac_result[..params.hmac_size]); + + if let Some(parent) = path.parent() { + std::fs::create_dir_all(parent).unwrap(); + } + std::fs::write(path, &page).unwrap(); + } + + #[test] + fn ensure_decrypted_with_progress_emits_decrypted_when_files_need_decrypt() { + let tmp = TempDir::new().unwrap(); + let enc_root = tmp.path().join("encrypted").join("db_storage"); + let cache_root = tmp.path().join("cache").join("db_storage"); + std::fs::create_dir_all(&enc_root).unwrap(); + std::fs::create_dir_all(&cache_root).unwrap(); + + // Build two encrypted DBs + let raw_key = [0xABu8; 32]; + let salt = [0x01u8; 16]; + let params = &wx_decrypt::MACOS_4_1_7_31; + build_encrypted_db(&enc_root.join("a.db"), &raw_key, &salt, params); + build_encrypted_db(&enc_root.join("b.db"), &raw_key, &salt, params); + + let enc_key = derive_enc_key(&raw_key, &salt, params); + let cache = test_cache( + cache_root.clone(), + enc_root.clone(), + None, + &[wx_decrypt::EncKeyPair { key: enc_key, salt }], + params, + ); + + let events = Mutex::new(Vec::::new()); + let stats = cache + .ensure_decrypted_with_progress(|e| { + events.lock().unwrap().push(e); + }) + .unwrap(); + let events = events.into_inner().unwrap(); + + assert_eq!(stats.decrypted, 2); + let decrypted_count = events + .iter() + .filter(|e| matches!(e, DecryptProgress::Decrypted { .. })) + .count(); + assert_eq!(decrypted_count, 2, "should emit 2 Decrypted events"); + assert_eq!( + decrypted_count + + events + .iter() + .filter(|e| matches!(e, DecryptProgress::Skipped { .. })) + .count(), + 2, + "Decrypted + Skipped should equal total DB count" + ); + } + + #[test] + fn ensure_decrypted_with_progress_only_skipped_when_all_cached() { + let tmp = TempDir::new().unwrap(); + let enc_root = tmp.path().join("encrypted").join("db_storage"); + let cache_root = tmp.path().join("cache").join("db_storage"); + std::fs::create_dir_all(&enc_root).unwrap(); + std::fs::create_dir_all(&cache_root).unwrap(); + + let raw_key = [0xABu8; 32]; + let salt = [0x01u8; 16]; + let params = &wx_decrypt::MACOS_4_1_7_31; + build_encrypted_db(&enc_root.join("a.db"), &raw_key, &salt, params); + + let enc_key = derive_enc_key(&raw_key, &salt, params); + let cache = test_cache( + cache_root.clone(), + enc_root.clone(), + None, + &[wx_decrypt::EncKeyPair { key: enc_key, salt }], + params, + ); + + // First run: decrypt everything + cache.ensure_decrypted().unwrap(); + + // Second run with progress: only Skipped events, no Decrypted or Failed + let events = Mutex::new(Vec::::new()); + let stats = cache + .ensure_decrypted_with_progress(|e| { + events.lock().unwrap().push(e); + }) + .unwrap(); + let events = events.into_inner().unwrap(); + + assert_eq!(stats.skipped, 1); + assert_eq!(stats.decrypted, 0); + assert!( + events + .iter() + .all(|e| matches!(e, DecryptProgress::Skipped { .. })), + "all events should be Skipped when all cached, got: {events:?}" + ); + assert_eq!(events.len(), 1, "should emit exactly 1 Skipped event"); + } + + // --- Task 5g integration tests --- + + #[test] + fn rawkey_path_caches_across_files() { + // Decrypt two DBs with same raw_key but different salts via RawKey. + // KdfCache should have 2 entries after, both newly derived. + let raw_key = [0xABu8; 32]; + let salt1 = [0x01u8; 16]; + let salt2 = [0x02u8; 16]; + let params = &wx_decrypt::MACOS_4_1_7_31; + + let tmp = TempDir::new().unwrap(); + let enc_root = tmp.path().join("encrypted"); + let cache_root = tmp.path().join("cache"); + std::fs::create_dir_all(&enc_root).unwrap(); + std::fs::create_dir_all(&cache_root).unwrap(); + + build_encrypted_db(&enc_root.join("a.db"), &raw_key, &salt1, params); + build_encrypted_db(&enc_root.join("b.db"), &raw_key, &salt2, params); + + // No pre-loaded pairs, only raw_key → all derivations go through get_or_derive + let cache = test_cache( + cache_root.clone(), + enc_root.clone(), + Some(raw_key), + &[], + params, + ); + + cache + .decrypt_db(&enc_root.join("a.db"), &cache_root.join("a.db")) + .unwrap(); + cache + .decrypt_db(&enc_root.join("b.db"), &cache_root.join("b.db")) + .unwrap(); + + let guard = cache.kdf_cache.lock().unwrap(); + assert!(guard.has_new_derivations()); + assert_eq!( + guard.new_pairs().len(), + 2, + "two unique salts should be derived" + ); + assert_eq!(guard.all_pairs().len(), 2); + } + + #[test] + fn enc_keys_miss_falls_back_to_raw_key() { + // Pre-load KdfCache with enc_key for salt1 only. + // Decrypt a DB with salt2 → should derive via raw_key fallback. + let raw_key = [0xABu8; 32]; + let salt1 = [0x01u8; 16]; + let salt2 = [0x02u8; 16]; + let params = &wx_decrypt::MACOS_4_1_7_31; + + let enc_key1 = derive_enc_key(&raw_key, &salt1, params); + + let tmp = TempDir::new().unwrap(); + let enc_root = tmp.path().join("encrypted"); + let cache_root = tmp.path().join("cache"); + std::fs::create_dir_all(&enc_root).unwrap(); + std::fs::create_dir_all(&cache_root).unwrap(); + + build_encrypted_db(&enc_root.join("b.db"), &raw_key, &salt2, params); + + let cache = test_cache( + cache_root.clone(), + enc_root.clone(), + Some(raw_key), + &[wx_decrypt::EncKeyPair { + key: enc_key1, + salt: salt1, + }], + params, + ); + + // Decrypt DB with salt2 (not in pre-loaded pairs) → should succeed via raw_key + cache + .decrypt_db(&enc_root.join("b.db"), &cache_root.join("b.db")) + .unwrap(); + + let guard = cache.kdf_cache.lock().unwrap(); + assert!(guard.has_new_derivations(), "salt2 should be newly derived"); + assert_eq!(guard.new_pairs().len(), 1); + assert_eq!(guard.all_pairs().len(), 2, "pre-loaded + newly derived"); + } + + #[test] + fn stale_enc_key_triggers_raw_key_fallback() { + // Pre-load KdfCache with an INCORRECT enc_key for salt1. + // With raw_key present, decrypt_db should retry with fresh derivation. + let raw_key = [0xABu8; 32]; + let salt1 = [0x01u8; 16]; + let params = &wx_decrypt::MACOS_4_1_7_31; + + let tmp = TempDir::new().unwrap(); + let enc_root = tmp.path().join("encrypted"); + let cache_root = tmp.path().join("cache"); + std::fs::create_dir_all(&enc_root).unwrap(); + std::fs::create_dir_all(&cache_root).unwrap(); + + build_encrypted_db(&enc_root.join("a.db"), &raw_key, &salt1, params); + + // Pre-load with wrong enc_key for this salt + let wrong_enc_key = [0xFFu8; 32]; + let cache = test_cache( + cache_root.clone(), + enc_root.clone(), + Some(raw_key), + &[wx_decrypt::EncKeyPair { + key: wrong_enc_key, + salt: salt1, + }], + params, + ); + + // Should succeed — stale enc_key triggers re-derive from raw_key + cache + .decrypt_db(&enc_root.join("a.db"), &cache_root.join("a.db")) + .unwrap(); + + let guard = cache.kdf_cache.lock().unwrap(); + // The stale entry should have been refreshed + let cached = guard.lookup(&salt1).unwrap(); + let correct = derive_enc_key(&raw_key, &salt1, params); + assert_eq!(cached, correct, "cache should now have the correct enc_key"); + } + + #[test] + fn newly_derived_pairs_tracked_in_kdf_cache() { + // After RawKey derivation, has_new_derivations should be true + // and new_pairs should contain the derived entry. + let raw_key = [0xABu8; 32]; + let salt = [0x01u8; 16]; + let params = &wx_decrypt::MACOS_4_1_7_31; + + let tmp = TempDir::new().unwrap(); + let enc_root = tmp.path().join("encrypted"); + let cache_root = tmp.path().join("cache"); + std::fs::create_dir_all(&enc_root).unwrap(); + std::fs::create_dir_all(&cache_root).unwrap(); + + build_encrypted_db(&enc_root.join("a.db"), &raw_key, &salt, params); + + let cache = test_cache( + cache_root.clone(), + enc_root.clone(), + Some(raw_key), + &[], + params, + ); + + cache + .decrypt_db(&enc_root.join("a.db"), &cache_root.join("a.db")) + .unwrap(); + + let guard = cache.kdf_cache.lock().unwrap(); + assert!(guard.has_new_derivations()); + let pairs = guard.new_pairs(); + assert_eq!(pairs.len(), 1); + assert_eq!(pairs[0].salt, salt); + let expected_key = derive_enc_key(&raw_key, &salt, params); + assert_eq!(pairs[0].key, expected_key); + } + + #[test] + fn cli_key_flag_does_not_trigger_writeback() { + // When writeback_enabled is false, writeback_enc_keys should be a no-op. + let raw_key = [0xABu8; 32]; + let salt = [0x01u8; 16]; + let params = &wx_decrypt::MACOS_4_1_7_31; + + let tmp = TempDir::new().unwrap(); + let enc_root = tmp.path().join("encrypted"); + let cache_root = tmp.path().join("cache"); + std::fs::create_dir_all(&enc_root).unwrap(); + std::fs::create_dir_all(&cache_root).unwrap(); + + build_encrypted_db(&enc_root.join("a.db"), &raw_key, &salt, params); + + let cache = PersistentCache { + cache_root: cache_root.clone(), + encrypted_root: enc_root.clone(), + raw_key: Some(raw_key), + account_id: "wxid_test_cli".into(), + base_wxid: "wxid_test_cli".into(), + writeback_enabled: false, // CLI --key flag + kdf_cache: Mutex::new(KdfCache::empty()), + params, + file_locks: Arc::new(FileLockMap::default()), + }; + + cache + .decrypt_db(&enc_root.join("a.db"), &cache_root.join("a.db")) + .unwrap(); + + // writeback should be a no-op (returns Ok immediately) + assert!(cache.writeback_enc_keys().is_ok()); + // The KdfCache has new derivations, but writeback_enabled is false + let guard = cache.kdf_cache.lock().unwrap(); + assert!(guard.has_new_derivations(), "derivation happened"); + } +} diff --git a/crates/wx-context/src/contact.rs b/crates/wx-context/src/contact.rs new file mode 100644 index 0000000..0931416 --- /dev/null +++ b/crates/wx-context/src/contact.rs @@ -0,0 +1,346 @@ +use std::collections::HashMap; + +use serde::{Deserialize, Serialize}; + +use crate::ContextError; + +struct ResolvedContact { + display_name: String, + remark: String, + nick_name: String, + alias: String, + user_name: String, + phone: Option, + memo: Option, + signature: Option, + region: Option, + labels: Vec, +} + +pub struct ContactResolver { + contacts: HashMap, +} + +impl ContactResolver { + pub fn empty() -> Self { + Self { + contacts: HashMap::new(), + } + } + + /// 从 WechatDb 构建映射表。 + /// 显示名优先级:remark > nick_name > alias > username。 + pub fn build(db: &wx_db::WechatDb) -> Result { + const BATCH: usize = 10_000; + let mut contacts = HashMap::new(); + let mut offset = 0; + + loop { + let query = wx_db::ContactQuery::new().limit(BATCH).offset(offset); + let result = db.query_contacts(&query)?; + let count = result.items.len(); + + for c in result.items { + let display_name = effective_display_name(&c); + let wxid = c.user_name.clone(); + contacts.insert( + wxid, + ResolvedContact { + display_name, + remark: c.remark, + nick_name: c.nick_name, + alias: c.alias, + user_name: c.user_name, + phone: c.phone, + memo: c.memo, + signature: c.signature, + region: c.region, + labels: c.labels, + }, + ); + } + + if count < BATCH { + break; + } + offset += count; + } + + Ok(Self { contacts }) + } + + /// 解析 wxid → 显示名,未找到返回 None。 + pub fn resolve(&self, wxid: &str) -> Option<&str> { + self.contacts.get(wxid).map(|r| r.display_name.as_str()) + } + + /// 解析 wxid → 显示名,未找到时返回原始 wxid。 + pub fn display_name<'a>(&'a self, wxid: &'a str) -> &'a str { + self.resolve(wxid).unwrap_or(wxid) + } + + /// Check if a wxid exists in the contact index. + pub fn contains(&self, wxid: &str) -> bool { + self.contacts.contains_key(wxid) + } + + /// Get labels for a contact. Returns empty slice if not found. + pub fn labels(&self, wxid: &str) -> &[String] { + self.contacts + .get(wxid) + .map(|r| r.labels.as_slice()) + .unwrap_or(&[]) + } + + /// Iterate all contacts with their wxid and labels. + /// Used by VisibilityIndex to expand ignore_tags. + pub fn all_labels(&self) -> impl Iterator { + self.contacts + .iter() + .map(|(wxid, r)| (wxid, r.labels.as_slice())) + } + + /// 格式化为 "显示名(wxid)",显示名与 wxid 相同时只返回 wxid。 + pub fn display_with_id(&self, wxid: &str) -> String { + match self.resolve(wxid) { + Some(name) if name != wxid => format!("{name}({wxid})"), + _ => wxid.to_string(), + } + } + + /// 反向模糊匹配:从所有字段查找候选 wxid。 + /// 返回 (display_name, wxid) 列表。 + pub fn find_candidates(&self, keyword: &str) -> Vec<(&str, &str)> { + let kw_lower = keyword.to_lowercase(); + self.contacts + .iter() + .filter(|(_, r)| { + r.display_name.to_lowercase().contains(&kw_lower) + || r.remark.to_lowercase().contains(&kw_lower) + || r.nick_name.to_lowercase().contains(&kw_lower) + || r.alias.to_lowercase().contains(&kw_lower) + || r.user_name.to_lowercase().contains(&kw_lower) + || r.phone + .as_deref() + .is_some_and(|s| s.to_lowercase().contains(&kw_lower)) + || r.memo + .as_deref() + .is_some_and(|s| s.to_lowercase().contains(&kw_lower)) + || r.signature + .as_deref() + .is_some_and(|s| s.to_lowercase().contains(&kw_lower)) + || r.region + .as_deref() + .is_some_and(|s| s.to_lowercase().contains(&kw_lower)) + || r.labels + .iter() + .any(|l| l.to_lowercase().contains(&kw_lower)) + }) + .map(|(wxid, r)| (r.display_name.as_str(), wxid.as_str())) + .collect() + } +} + +fn effective_display_name(c: &wx_db::Contact) -> String { + if !c.remark.is_empty() { + c.remark.clone() + } else if !c.nick_name.is_empty() { + c.nick_name.clone() + } else if !c.alias.is_empty() { + c.alias.clone() + } else { + c.user_name.clone() + } +} + +/// 消息方向。 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum Direction { + Incoming, + Outgoing, +} + +impl Direction { + /// 从 sender 与 self_wxid 的比较判断方向。 + pub fn detect(sender: &str, self_wxid: &str) -> Self { + if sender == self_wxid { + Direction::Outgoing + } else { + Direction::Incoming + } + } + + pub fn as_str(&self) -> &'static str { + match self { + Direction::Incoming => "incoming", + Direction::Outgoing => "outgoing", + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn make_resolver(entries: &[(&str, &str)]) -> ContactResolver { + let contacts = entries + .iter() + .map(|(k, v)| { + ( + k.to_string(), + ResolvedContact { + display_name: v.to_string(), + remark: v.to_string(), + nick_name: String::new(), + alias: String::new(), + user_name: k.to_string(), + phone: None, + memo: None, + signature: None, + region: None, + labels: Vec::new(), + }, + ) + }) + .collect(); + ContactResolver { contacts } + } + + #[allow(clippy::type_complexity)] + fn make_resolver_full( + entries: &[( + &str, // wxid + &str, // display_name / remark + &str, // nick_name + &str, // alias + Option<&str>, // phone + Option<&str>, // memo + &[&str], // labels + )], + ) -> ContactResolver { + let contacts = entries + .iter() + .map(|(wxid, remark, nick, alias, phone, memo, labels)| { + ( + wxid.to_string(), + ResolvedContact { + display_name: remark.to_string(), + remark: remark.to_string(), + nick_name: nick.to_string(), + alias: alias.to_string(), + user_name: wxid.to_string(), + phone: phone.map(|s| s.to_string()), + memo: memo.map(|s| s.to_string()), + signature: None, + region: None, + labels: labels.iter().map(|s| s.to_string()).collect(), + }, + ) + }) + .collect(); + ContactResolver { contacts } + } + + #[test] + fn display_name_found() { + let r = make_resolver(&[("wxid_abc", "张三")]); + assert_eq!(r.display_name("wxid_abc"), "张三"); + } + + #[test] + fn display_name_not_found() { + let r = make_resolver(&[]); + assert_eq!(r.display_name("wxid_abc"), "wxid_abc"); + } + + #[test] + fn display_with_id_format() { + let r = make_resolver(&[("wxid_abc", "张三")]); + assert_eq!(r.display_with_id("wxid_abc"), "张三(wxid_abc)"); + } + + #[test] + fn display_with_id_same_as_wxid() { + let r = make_resolver(&[("wxid_abc", "wxid_abc")]); + assert_eq!(r.display_with_id("wxid_abc"), "wxid_abc"); + } + + #[test] + fn find_candidates_case_insensitive() { + let r = make_resolver(&[("wxid_a", "Alice"), ("wxid_b", "Bob")]); + let c = r.find_candidates("ali"); + assert_eq!(c.len(), 1); + assert_eq!(c[0], ("Alice", "wxid_a")); + } + + #[test] + fn find_candidates_matches_alias() { + let r = make_resolver_full(&[( + "wxid_a", + "Alice Remark", + "Alice Nick", + "alice_alias", + None, + None, + &[], + )]); + let c = r.find_candidates("alice_alias"); + assert_eq!(c.len(), 1); + assert_eq!(c[0].1, "wxid_a"); + } + + #[test] + fn find_candidates_matches_nick_name() { + let r = + make_resolver_full(&[("wxid_a", "Alice Remark", "Alice Nick", "", None, None, &[])]); + let c = r.find_candidates("Nick"); + assert_eq!(c.len(), 1); + assert_eq!(c[0].1, "wxid_a"); + } + + #[test] + fn find_candidates_matches_phone() { + let r = make_resolver_full(&[("wxid_a", "Alice", "", "", Some("13800138000"), None, &[])]); + let c = r.find_candidates("138001"); + assert_eq!(c.len(), 1); + assert_eq!(c[0].1, "wxid_a"); + } + + #[test] + fn find_candidates_matches_label() { + let r = make_resolver_full(&[("wxid_a", "Alice", "", "", None, None, &["体育生", "同事"])]); + let c = r.find_candidates("体育"); + assert_eq!(c.len(), 1); + assert_eq!(c[0].1, "wxid_a"); + } + + #[test] + fn find_candidates_matches_memo() { + let r = make_resolver_full(&[("wxid_a", "Alice", "", "", None, Some("QQ 442007516"), &[])]); + let c = r.find_candidates("442007"); + assert_eq!(c.len(), 1); + assert_eq!(c[0].1, "wxid_a"); + } + + #[test] + fn direction_detect() { + assert_eq!(Direction::detect("wxid_me", "wxid_me"), Direction::Outgoing); + assert_eq!( + Direction::detect("wxid_other", "wxid_me"), + Direction::Incoming + ); + } + + #[test] + fn direction_detect_legacy_non_wxid_self_id() { + assert_eq!( + Direction::detect("testuser001", "testuser001"), + Direction::Outgoing + ); + assert_eq!( + Direction::detect("wxid_friend", "testuser001"), + Direction::Incoming + ); + } +} diff --git a/crates/wx-context/src/db_category.rs b/crates/wx-context/src/db_category.rs new file mode 100644 index 0000000..dc93536 --- /dev/null +++ b/crates/wx-context/src/db_category.rs @@ -0,0 +1,123 @@ +use std::path::{Path, PathBuf}; + +/// Categorization of WeChat database files. +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +pub enum DbCategory { + Contact, + Session, + MessageShard { shard_id: u32 }, + Fts, + Other, +} + +/// A database file with its path and category. +#[derive(Debug, Clone)] +pub struct DbFile { + pub path: PathBuf, + pub category: DbCategory, +} + +impl DbFile { + /// Categorize a database file by its stem. + pub fn categorize(path: PathBuf) -> Self { + let category = match path.file_stem().and_then(|s| s.to_str()) { + Some("contact") => DbCategory::Contact, + Some("session") => DbCategory::Session, + Some("message_fts") => DbCategory::Fts, + Some(stem) => { + if let Some(suffix) = stem.strip_prefix("message_") { + if let Ok(id) = suffix.parse::() { + DbCategory::MessageShard { shard_id: id } + } else { + DbCategory::Other + } + } else { + DbCategory::Other + } + } + None => DbCategory::Other, + }; + Self { path, category } + } +} + +/// Recursively discover and categorize all `.db` files under `dir`. +pub fn discover_db_files(dir: &Path) -> Result, std::io::Error> { + let mut result = Vec::new(); + collect_db_files(dir, &mut result)?; + result.sort_by(|a, b| a.path.cmp(&b.path)); + Ok(result) +} + +fn collect_db_files(dir: &Path, out: &mut Vec) -> Result<(), std::io::Error> { + for entry in std::fs::read_dir(dir)? { + let entry = entry?; + let ft = entry.file_type()?; + if ft.is_dir() { + collect_db_files(&entry.path(), out)?; + } else if ft.is_file() && entry.path().extension().is_some_and(|e| e == "db") { + out.push(DbFile::categorize(entry.path())); + } + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + use tempfile::TempDir; + + #[test] + fn categorize_contact() { + let db = DbFile::categorize(PathBuf::from("/tmp/contact.db")); + assert_eq!(db.category, DbCategory::Contact); + } + + #[test] + fn categorize_session() { + let db = DbFile::categorize(PathBuf::from("/tmp/session.db")); + assert_eq!(db.category, DbCategory::Session); + } + + #[test] + fn categorize_message_shard() { + let db = DbFile::categorize(PathBuf::from("/tmp/message_0.db")); + assert_eq!(db.category, DbCategory::MessageShard { shard_id: 0 }); + + let db = DbFile::categorize(PathBuf::from("/tmp/message_42.db")); + assert_eq!(db.category, DbCategory::MessageShard { shard_id: 42 }); + } + + #[test] + fn categorize_fts() { + let db = DbFile::categorize(PathBuf::from("/tmp/message_fts.db")); + assert_eq!(db.category, DbCategory::Fts); + } + + #[test] + fn categorize_unknown() { + let db = DbFile::categorize(PathBuf::from("/tmp/something_else.db")); + assert_eq!(db.category, DbCategory::Other); + } + + #[test] + fn discover_db_files_with_temp_dir() { + let tmp = TempDir::new().unwrap(); + let sub = tmp.path().join("contact"); + std::fs::create_dir_all(&sub).unwrap(); + std::fs::write(sub.join("contact.db"), b"").unwrap(); + std::fs::write(tmp.path().join("session.db"), b"").unwrap(); + std::fs::write(tmp.path().join("message_0.db"), b"").unwrap(); + std::fs::write(tmp.path().join("message_fts.db"), b"").unwrap(); + std::fs::write(tmp.path().join("not_a_db.txt"), b"").unwrap(); + + let files = discover_db_files(tmp.path()).unwrap(); + assert_eq!(files.len(), 4); + + let categories: Vec<_> = files.iter().map(|f| &f.category).collect(); + assert!(categories.contains(&&DbCategory::Contact)); + assert!(categories.contains(&&DbCategory::Session)); + assert!(categories.contains(&&DbCategory::MessageShard { shard_id: 0 })); + assert!(categories.contains(&&DbCategory::Fts)); + } +} diff --git a/crates/wx-context/src/decrypt_request.rs b/crates/wx-context/src/decrypt_request.rs new file mode 100644 index 0000000..4a8217d --- /dev/null +++ b/crates/wx-context/src/decrypt_request.rs @@ -0,0 +1,56 @@ +use crate::cache::PersistentCache; +use crate::decrypt_scope::DecryptScope; +use crate::progress::{DecryptProgress, DecryptStats}; +use crate::ContextError; + +/// Fluent builder for scoped decryption requests. +pub struct DecryptRequest { + scope: DecryptScope, +} + +impl Default for DecryptRequest { + fn default() -> Self { + Self::new() + } +} + +impl DecryptRequest { + /// Start with `All` scope (decrypt everything). + pub fn new() -> Self { + Self { + scope: DecryptScope::All, + } + } + + /// Restrict to core DBs only (contact + session). + pub fn core(mut self) -> Self { + self.scope = DecryptScope::Core; + self + } + + /// Explicitly request all DBs. + pub fn all(mut self) -> Self { + self.scope = DecryptScope::All; + self + } + + /// Add specific message shards (implies core). + pub fn shards(mut self, shard_ids: &[u32]) -> Self { + self.scope = self.scope.with_shards(shard_ids.to_vec()); + self + } + + /// Execute the decrypt request. + pub fn execute(self, cache: &PersistentCache) -> Result { + cache.ensure_decrypted_scoped(&self.scope, |_| {}) + } + + /// Execute with progress callback. + pub fn execute_with_progress( + self, + cache: &PersistentCache, + on_progress: impl Fn(DecryptProgress) + Send + Sync, + ) -> Result { + cache.ensure_decrypted_scoped(&self.scope, on_progress) + } +} diff --git a/crates/wx-context/src/decrypt_scope.rs b/crates/wx-context/src/decrypt_scope.rs new file mode 100644 index 0000000..e88c2a3 --- /dev/null +++ b/crates/wx-context/src/decrypt_scope.rs @@ -0,0 +1,95 @@ +use std::collections::HashSet; + +use crate::db_category::{DbCategory, DbFile}; + +/// Which databases to decrypt. +#[derive(Debug, Clone)] +pub enum DecryptScope { + /// Decrypt everything. + All, + /// Only contact.db + session.db. + Core, + /// Arbitrary set of categories. + Categories(HashSet), + /// Core DBs + specific message shards. + MessageShards { shard_ids: Vec }, +} + +impl DecryptScope { + /// Convenience: core-only scope. + pub fn core() -> Self { + Self::Core + } + + /// Upgrade to include specific message shards (implies Core). + pub fn with_shards(self, shard_ids: Vec) -> Self { + Self::MessageShards { shard_ids } + } + + /// Check whether a given DB file is included in this scope. + pub fn matches(&self, db: &DbFile) -> bool { + match self { + Self::All => true, + Self::Core => matches!(db.category, DbCategory::Contact | DbCategory::Session), + Self::Categories(cats) => cats.contains(&db.category), + Self::MessageShards { shard_ids } => match &db.category { + DbCategory::Contact | DbCategory::Session => true, + DbCategory::MessageShard { shard_id } => shard_ids.contains(shard_id), + _ => false, + }, + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::path::PathBuf; + + fn db(name: &str) -> DbFile { + DbFile::categorize(PathBuf::from(format!("/tmp/{name}.db"))) + } + + #[test] + fn all_matches_everything() { + let scope = DecryptScope::All; + assert!(scope.matches(&db("contact"))); + assert!(scope.matches(&db("session"))); + assert!(scope.matches(&db("message_0"))); + assert!(scope.matches(&db("message_fts"))); + assert!(scope.matches(&db("other"))); + } + + #[test] + fn core_matches_only_contact_and_session() { + let scope = DecryptScope::core(); + assert!(scope.matches(&db("contact"))); + assert!(scope.matches(&db("session"))); + assert!(!scope.matches(&db("message_0"))); + assert!(!scope.matches(&db("message_fts"))); + assert!(!scope.matches(&db("other"))); + } + + #[test] + fn message_shards_includes_core_and_selected() { + let scope = DecryptScope::core().with_shards(vec![0, 2]); + assert!(scope.matches(&db("contact"))); + assert!(scope.matches(&db("session"))); + assert!(scope.matches(&db("message_0"))); + assert!(!scope.matches(&db("message_1"))); + assert!(scope.matches(&db("message_2"))); + assert!(!scope.matches(&db("message_fts"))); + } + + #[test] + fn categories_matches_specified() { + let mut cats = HashSet::new(); + cats.insert(DbCategory::Fts); + cats.insert(DbCategory::Contact); + let scope = DecryptScope::Categories(cats); + assert!(scope.matches(&db("contact"))); + assert!(scope.matches(&db("message_fts"))); + assert!(!scope.matches(&db("session"))); + assert!(!scope.matches(&db("message_0"))); + } +} diff --git a/crates/wx-context/src/error.rs b/crates/wx-context/src/error.rs new file mode 100644 index 0000000..79fb74f --- /dev/null +++ b/crates/wx-context/src/error.rs @@ -0,0 +1,28 @@ +use thiserror::Error; + +#[derive(Debug, Error)] +pub enum ContextError { + #[error("keychain: {0}")] + Keychain(#[from] wx_keychain::KeychainError), + + #[error("database: {0}")] + Db(#[from] wx_db::DbError), + + #[error("decrypt: {0}")] + Decrypt(#[from] wx_decrypt::DecryptError), + + #[error("io: {0}")] + Io(#[from] std::io::Error), + + #[error("no account found: {0}")] + NoAccount(String), + + #[error("no key for account {0}: use `key extract` or `-k`")] + NoKey(String), + + #[error("cache: {0}")] + Cache(String), + + #[error("sqlite: {0}")] + Sqlite(String), +} diff --git a/crates/wx-context/src/file_lock.rs b/crates/wx-context/src/file_lock.rs new file mode 100644 index 0000000..a5ae5ac --- /dev/null +++ b/crates/wx-context/src/file_lock.rs @@ -0,0 +1,58 @@ +use dashmap::DashMap; +use std::path::{Path, PathBuf}; +use std::sync::{Arc, Mutex}; + +/// Per-file mutex map to prevent concurrent `ensure_decrypted` calls +/// from operating on the same DB file simultaneously. +#[derive(Default)] +pub(crate) struct FileLockMap { + inner: DashMap>>, +} + +impl FileLockMap { + /// Returns an `Arc>` for the given path. + /// + /// Callers should `.lock().unwrap()` on the returned Arc in their own + /// scope — never while holding a DashMap entry reference. + pub fn lock_for(&self, path: &Path) -> Arc> { + self.inner + .entry(path.to_path_buf()) + .or_insert_with(|| Arc::new(Mutex::new(()))) + .clone() + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::Arc as StdArc; + + #[test] + fn same_path_returns_same_mutex() { + let map = FileLockMap::default(); + let path = Path::new("/tmp/test.db"); + + let a = map.lock_for(path); + let b = map.lock_for(path); + + assert!(StdArc::ptr_eq(&a, &b)); + } + + #[test] + fn different_paths_do_not_block_each_other() { + let map = StdArc::new(FileLockMap::default()); + let path_a = PathBuf::from("/tmp/a.db"); + let path_b = PathBuf::from("/tmp/b.db"); + + let lock_a = map.lock_for(&path_a); + let _guard_a = lock_a.lock().unwrap(); + + // Locking a different path must succeed immediately (not blocked by path_a). + let lock_b = map.lock_for(&path_b); + let result = lock_b.try_lock(); + assert!( + result.is_ok(), + "lock on different path should not be blocked" + ); + } +} diff --git a/crates/wx-context/src/fts_tokenizer.rs b/crates/wx-context/src/fts_tokenizer.rs new file mode 100644 index 0000000..694509c --- /dev/null +++ b/crates/wx-context/src/fts_tokenizer.rs @@ -0,0 +1,348 @@ +use std::os::raw::{c_char, c_int, c_void}; +use std::path::Path; +use std::ptr; + +use rusqlite::ffi; + +use crate::error::ContextError; +use crate::tokenizer::{tokenize, Token, TokenizerConfig}; + +// --------------------------------------------------------------------------- +// FTS5 token colocated flag +// --------------------------------------------------------------------------- + +const FTS5_TOKEN_COLOCATED: c_int = 0x0001; + +// --------------------------------------------------------------------------- +// MmTokenizer struct (C-ABI compatible) +// --------------------------------------------------------------------------- + +#[repr(C)] +struct MmTokenizer { + config: TokenizerConfig, +} + +// --------------------------------------------------------------------------- +// C-ABI callbacks +// --------------------------------------------------------------------------- + +unsafe extern "C" fn x_create( + user_data: *mut c_void, + az_arg: *mut *const c_char, + n_arg: c_int, + pp_out: *mut *mut ffi::Fts5Tokenizer, +) -> c_int { + let _ = user_data; + let mut config = TokenizerConfig::default(); + + // Parse azArg parameters + if !az_arg.is_null() && n_arg > 0 { + for i in 0..n_arg as usize { + let arg_ptr = *az_arg.add(i); + if arg_ptr.is_null() { + continue; + } + let Ok(arg) = std::ffi::CStr::from_ptr(arg_ptr).to_str() else { + continue; + }; + match arg { + "enable_special_char" => config.enable_special_char = true, + "enable_num_token" => config.enable_num_token = true, + // disable_pinyin and disable_origin are deferred to Task 8 (contact pinyin support) + // Unknown params: silently ignored for forward-compatibility + _ => {} + } + } + } + + let tok = Box::new(MmTokenizer { config }); + *pp_out = Box::into_raw(tok) as *mut ffi::Fts5Tokenizer; + ffi::SQLITE_OK as c_int +} + +unsafe extern "C" fn x_delete(tokenizer: *mut ffi::Fts5Tokenizer) { + if !tokenizer.is_null() { + drop(Box::from_raw(tokenizer as *mut MmTokenizer)); + } +} + +unsafe extern "C" fn x_tokenize( + tokenizer: *mut ffi::Fts5Tokenizer, + ctx: *mut c_void, + _flags: c_int, + text: *const c_char, + n_text: c_int, + x_token: Option< + unsafe extern "C" fn(*mut c_void, c_int, *const c_char, c_int, c_int, c_int) -> c_int, + >, +) -> c_int { + if tokenizer.is_null() || text.is_null() { + return ffi::SQLITE_OK as c_int; + } + + let mm = &*(tokenizer as *mut MmTokenizer); + + let text_slice = if n_text < 0 { + // null-terminated + let cstr = std::ffi::CStr::from_ptr(text); + cstr.to_bytes() + } else { + std::slice::from_raw_parts(text as *const u8, n_text as usize) + }; + + let Some(x_tok_fn) = x_token else { + return ffi::SQLITE_OK as c_int; + }; + + let tokens: Vec = tokenize(text_slice, &mm.config); + + for tok in &tokens { + let tflags = if tok.colocated { + FTS5_TOKEN_COLOCATED + } else { + 0 + }; + let rc = x_tok_fn( + ctx, + tflags, + tok.text.as_ptr() as *const c_char, + tok.text.len() as c_int, + tok.start as c_int, + tok.end as c_int, + ); + if rc != ffi::SQLITE_OK as c_int { + return rc; + } + } + + ffi::SQLITE_OK as c_int +} + +// --------------------------------------------------------------------------- +// Internal helper: get_fts5_api + register +// --------------------------------------------------------------------------- + +unsafe fn errmsg(db: *mut ffi::sqlite3) -> String { + let ptr = ffi::sqlite3_errmsg(db); + if ptr.is_null() { + return "(no message)".into(); + } + std::ffi::CStr::from_ptr(ptr).to_string_lossy().into_owned() +} + +unsafe fn get_fts5_api(db: *mut ffi::sqlite3) -> Result<*mut ffi::fts5_api, String> { + let sql = c"SELECT fts5(?1)"; + let mut stmt: *mut ffi::sqlite3_stmt = ptr::null_mut(); + let rc = ffi::sqlite3_prepare_v2(db, sql.as_ptr(), -1, &mut stmt, ptr::null_mut()); + if rc != ffi::SQLITE_OK as c_int { + return Err(format!("prepare fts5 query: rc={rc} msg={}", errmsg(db))); + } + + let mut api: *mut ffi::fts5_api = ptr::null_mut(); + let rc = ffi::sqlite3_bind_pointer( + stmt, + 1, + &mut api as *mut _ as *mut c_void, + c"fts5_api_ptr".as_ptr(), + None, + ); + if rc != ffi::SQLITE_OK as c_int { + ffi::sqlite3_finalize(stmt); + return Err(format!("bind_pointer: rc={rc} msg={}", errmsg(db))); + } + + let rc = ffi::sqlite3_step(stmt); + ffi::sqlite3_finalize(stmt); + if rc != ffi::SQLITE_ROW as c_int { + return Err(format!("step fts5: rc={rc} msg={}", errmsg(db))); + } + + if api.is_null() { + return Err("fts5_api unavailable".into()); + } + + Ok(api) +} + +// --------------------------------------------------------------------------- +// Public API +// --------------------------------------------------------------------------- + +/// Register the real `MMFtsTokenizer` on `conn`. +/// +/// This replaces the no-op tokenizer in `wal_patch.rs` and enables actual +/// FTS5 MATCH queries against WeChat's native `message_fts.db`. +pub fn register_mm_fts_tokenizer(conn: &rusqlite::Connection) -> Result<(), String> { + unsafe { + let db = conn.handle(); + let api = get_fts5_api(db)?; + + let mut tok = ffi::fts5_tokenizer { + xCreate: Some(x_create), + xDelete: Some(x_delete), + xTokenize: Some(x_tokenize), + }; + let name = c"MMFtsTokenizer"; + let x_create_tok = (*api) + .xCreateTokenizer + .ok_or("fts5_api.xCreateTokenizer is NULL")?; + let rc = x_create_tok(api, name.as_ptr(), ptr::null_mut(), &mut tok, None); + if rc != ffi::SQLITE_OK as c_int { + return Err(format!( + "xCreateTokenizer failed: rc={rc} msg={}", + errmsg(db) + )); + } + + Ok(()) + } +} + +/// Open a read-only connection to an FTS database with `MMFtsTokenizer` registered. +pub fn open_fts_connection(path: &Path) -> Result { + open_fts_connection_with_key(path, None) +} + +/// Open a read-only connection to an FTS database (optionally encrypted) +/// with `MMFtsTokenizer` registered. +pub fn open_fts_connection_with_key( + path: &Path, + raw_key: Option<&[u8; 32]>, +) -> Result { + if !path.exists() { + return Err(ContextError::Sqlite(format!( + "FTS database not found: {}", + path.display() + ))); + } + let conn = + rusqlite::Connection::open_with_flags(path, rusqlite::OpenFlags::SQLITE_OPEN_READ_ONLY) + .map_err(|e| ContextError::Sqlite(e.to_string()))?; + + if let Some(key) = raw_key { + unsafe { + let rc = ffi::sqlite3_key(conn.handle(), key.as_ptr() as *const c_void, 32); + if rc != 0 { + return Err(ContextError::Sqlite(format!("sqlite3_key failed: rc={rc}"))); + } + } + conn.query_row("SELECT count(*) FROM sqlite_master", [], |r| { + r.get::<_, i64>(0) + }) + .map_err(|_| { + ContextError::Sqlite("incorrect key or not an encrypted FTS database".into()) + })?; + conn.execute_batch("PRAGMA query_only = ON") + .map_err(|e| ContextError::Sqlite(e.to_string()))?; + } + + register_mm_fts_tokenizer(&conn).map_err(ContextError::Sqlite)?; + + Ok(conn) +} + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +#[cfg(test)] +mod tests { + use super::*; + use rusqlite::Connection; + + fn open_test_conn() -> Connection { + Connection::open_in_memory().unwrap() + } + + // Test 1: register_mm_fts_tokenizer succeeds on a fresh connection + #[test] + fn register_succeeds() { + let conn = open_test_conn(); + let result = register_mm_fts_tokenizer(&conn); + assert!(result.is_ok(), "registration should succeed: {result:?}"); + } + + // Test 2: After registration, CREATE VIRTUAL TABLE with MMFtsTokenizer succeeds + #[test] + fn create_fts5_table_with_tokenizer() { + let conn = open_test_conn(); + register_mm_fts_tokenizer(&conn).unwrap(); + conn.execute_batch( + "CREATE VIRTUAL TABLE t USING fts5(content, tokenize='MMFtsTokenizer disable_pinyin');", + ) + .expect("should create FTS5 table with MMFtsTokenizer"); + } + + // Test 3: INSERT + MATCH for CJK works + #[test] + fn fts_match_cjk() { + let conn = open_test_conn(); + register_mm_fts_tokenizer(&conn).unwrap(); + conn.execute_batch( + "CREATE VIRTUAL TABLE t USING fts5(content, tokenize='MMFtsTokenizer'); + INSERT INTO t(content) VALUES ('你好world');", + ) + .unwrap(); + let count: i64 = conn + .query_row("SELECT count(*) FROM t WHERE t MATCH '\"你\"'", [], |r| { + r.get(0) + }) + .unwrap(); + assert_eq!(count, 1, "CJK MATCH should find the row"); + } + + // Test 4: INSERT + MATCH with English stemming: insert "running fast", MATCH "run" returns the row + #[test] + fn fts_match_english_stemming() { + let conn = open_test_conn(); + register_mm_fts_tokenizer(&conn).unwrap(); + conn.execute_batch( + "CREATE VIRTUAL TABLE t USING fts5(content, tokenize='MMFtsTokenizer'); + INSERT INTO t(content) VALUES ('running fast');", + ) + .unwrap(); + let count: i64 = conn + .query_row("SELECT count(*) FROM t WHERE t MATCH 'run'", [], |r| { + r.get(0) + }) + .unwrap(); + assert_eq!(count, 1, "Stemming: 'run' should match 'running'"); + } + + // Test 5: open_fts_connection on non-existent path returns error + #[test] + fn open_nonexistent_path_errors() { + let result = open_fts_connection(Path::new("/nonexistent/path/fts.db")); + assert!(result.is_err(), "should error on non-existent path"); + } + + #[test] + fn real_tokenizer_fts_round_trip() { + let conn = open_test_conn(); + register_mm_fts_tokenizer(&conn).unwrap(); + conn.execute_batch( + "CREATE VIRTUAL TABLE msg USING fts5(body, tokenize='MMFtsTokenizer disable_pinyin'); + INSERT INTO msg VALUES ('你好世界'); + INSERT INTO msg VALUES ('hello world running');", + ) + .unwrap(); + + // CJK unigram search + let n: i64 = conn + .query_row( + "SELECT count(*) FROM msg WHERE msg MATCH '\"你\"'", + [], + |r| r.get(0), + ) + .unwrap(); + assert_eq!(n, 1); + + // English stemmed search + let n: i64 = conn + .query_row("SELECT count(*) FROM msg WHERE msg MATCH 'run'", [], |r| { + r.get(0) + }) + .unwrap(); + assert_eq!(n, 1); + } +} diff --git a/crates/wx-context/src/kdf_cache.rs b/crates/wx-context/src/kdf_cache.rs new file mode 100644 index 0000000..f50e956 --- /dev/null +++ b/crates/wx-context/src/kdf_cache.rs @@ -0,0 +1,170 @@ +use std::collections::HashMap; + +use wx_decrypt::{CryptoParams, EncKeyPair}; + +/// In-memory cache for PBKDF2-derived encryption keys, keyed by DB salt. +/// +/// Tracks which entries were derived during the current session (vs pre-loaded) +/// so callers can persist only the new derivations. +pub(crate) struct KdfCache { + /// salt → enc_key + entries: HashMap<[u8; 16], [u8; 32]>, + /// Salts derived this session (not pre-loaded). + newly_derived: Vec<[u8; 16]>, +} + +impl KdfCache { + /// Pre-load from existing `EncKeyPair`s; `newly_derived` starts empty. + pub fn from_pairs(pairs: &[EncKeyPair]) -> Self { + let entries = pairs + .iter() + .map(|p| (p.salt, p.key)) + .collect::>(); + Self { + entries, + newly_derived: Vec::new(), + } + } + + /// Empty cache with no pre-loaded entries. + pub fn empty() -> Self { + Self { + entries: HashMap::new(), + newly_derived: Vec::new(), + } + } + + /// Cache hit check (returns owned copy). + pub fn lookup(&self, salt: &[u8; 16]) -> Option<[u8; 32]> { + self.entries.get(salt).copied() + } + + /// Lookup; on miss, derive enc_key via PBKDF2, insert to entries, + /// push salt to `newly_derived`. + pub fn get_or_derive( + &mut self, + salt: &[u8; 16], + raw_key: &[u8; 32], + params: &CryptoParams, + ) -> [u8; 32] { + if let Some(enc_key) = self.entries.get(salt) { + return *enc_key; + } + let enc_key = wx_decrypt::kdf::derive_enc_key(raw_key, salt, params); + self.entries.insert(*salt, enc_key); + self.newly_derived.push(*salt); + enc_key + } + + /// Overwrite an entry (for stale enc_key refresh). + /// Also marks as newly derived if not already tracked. + pub fn insert(&mut self, salt: &[u8; 16], enc_key: &[u8; 32]) { + self.entries.insert(*salt, *enc_key); + if !self.newly_derived.contains(salt) { + self.newly_derived.push(*salt); + } + } + + /// Any new pairs to write back? + pub fn has_new_derivations(&self) -> bool { + !self.newly_derived.is_empty() + } + + /// Only the session-derived pairs (filtered by `newly_derived`). + #[allow(dead_code)] // Used in tests; will be used by future consumers + pub fn new_pairs(&self) -> Vec { + self.newly_derived + .iter() + .filter_map(|salt| { + self.entries.get(salt).map(|key| EncKeyPair { + key: *key, + salt: *salt, + }) + }) + .collect() + } + + /// All entries (for merged writeback). + pub fn all_pairs(&self) -> Vec { + self.entries + .iter() + .map(|(salt, key)| EncKeyPair { + key: *key, + salt: *salt, + }) + .collect() + } +} + +#[cfg(test)] +mod tests { + use super::*; + use wx_decrypt::MACOS_4_1_7_31; + + fn make_pair(salt_byte: u8, key_byte: u8) -> EncKeyPair { + EncKeyPair { + salt: [salt_byte; 16], + key: [key_byte; 32], + } + } + + #[test] + fn pre_load_from_pairs_lookup_hits() { + let pair = make_pair(0x01, 0xAA); + let cache = KdfCache::from_pairs(std::slice::from_ref(&pair)); + + assert_eq!(cache.lookup(&[0x01; 16]), Some([0xAA; 32])); + assert_eq!(cache.lookup(&[0x02; 16]), None); + assert!(!cache.has_new_derivations()); + } + + #[test] + fn empty_cache_miss_triggers_derive() { + let mut cache = KdfCache::empty(); + let raw_key = [0xBB; 32]; + let salt = [0x01; 16]; + + // Miss → derive + assert_eq!(cache.lookup(&salt), None); + let enc_key = cache.get_or_derive(&salt, &raw_key, &MACOS_4_1_7_31); + + // Second lookup hits + assert_eq!(cache.lookup(&salt), Some(enc_key)); + assert!(cache.has_new_derivations()); + + // Verify derivation matches direct call + let expected = wx_decrypt::kdf::derive_enc_key(&raw_key, &salt, &MACOS_4_1_7_31); + assert_eq!(enc_key, expected); + } + + #[test] + fn new_pairs_only_includes_derived() { + let preloaded = make_pair(0x01, 0xAA); + let mut cache = KdfCache::from_pairs(&[preloaded]); + + // Derive a new entry + let raw_key = [0xCC; 32]; + let new_salt = [0x02; 16]; + cache.get_or_derive(&new_salt, &raw_key, &MACOS_4_1_7_31); + + let new_pairs = cache.new_pairs(); + assert_eq!(new_pairs.len(), 1); + assert_eq!(new_pairs[0].salt, new_salt); + + // all_pairs includes both + assert_eq!(cache.all_pairs().len(), 2); + } + + #[test] + fn get_or_derive_same_salt_twice_derives_once() { + let mut cache = KdfCache::empty(); + let raw_key = [0xDD; 32]; + let salt = [0x03; 16]; + + let first = cache.get_or_derive(&salt, &raw_key, &MACOS_4_1_7_31); + let second = cache.get_or_derive(&salt, &raw_key, &MACOS_4_1_7_31); + + assert_eq!(first, second); + assert_eq!(cache.new_pairs().len(), 1); + } +} diff --git a/crates/wx-context/src/lib.rs b/crates/wx-context/src/lib.rs new file mode 100644 index 0000000..6e50381 --- /dev/null +++ b/crates/wx-context/src/lib.rs @@ -0,0 +1,65 @@ +//! WeChat database context layer — account resolution, decrypt orchestration, +//! and encrypted direct-open support. +//! +//! When `raw_key` is available, use [`open_encrypted_db`] or +//! [`open_encrypted_db_with_pool`] to open encrypted databases directly +//! via SQLCipher's `sqlite3_key()` C API, bypassing the decrypt-cache pipeline. + +mod account; +mod cache; +mod contact; +mod db_category; +mod decrypt_request; +mod decrypt_scope; +mod error; +mod file_lock; +mod fts_tokenizer; +mod kdf_cache; +mod patch_state; +mod progress; +pub mod shard_routing; +pub mod tokenizer; +mod visibility; +mod wal_patch; + +pub use account::{AccountContext, ResolveParams}; +pub use cache::PersistentCache; +pub use contact::{ContactResolver, Direction}; +pub use db_category::{discover_db_files, DbCategory, DbFile}; +pub use decrypt_request::DecryptRequest; +pub use decrypt_scope::DecryptScope; +pub use error::ContextError; +pub use fts_tokenizer::{ + open_fts_connection, open_fts_connection_with_key, register_mm_fts_tokenizer, +}; +pub use progress::{DecryptProgress, DecryptStats}; +pub use shard_routing::{route_shards_for_query, write_shard_metadata_sidecar}; +pub use visibility::VisibilityIndex; + +/// Open encrypted WeChat DB directory directly (no pool, no FTS). +/// For one-shot commands: contacts, sessions, query, search, export. +pub fn open_encrypted_db(account: &AccountContext) -> Result { + let raw_key = account + .raw_key + .ok_or_else(|| ContextError::Cache("raw_key required for encrypted direct open".into()))?; + let encrypted_root = account.data_dir.join("db_storage"); + let db = wx_db::WechatDb::open_encrypted(&encrypted_root, raw_key)?; + Ok(db) +} + +/// Open encrypted WeChat DB directory with connection pool and FTS tokenizer. +/// For long-running serve mode only. +pub fn open_encrypted_db_with_pool( + account: &AccountContext, +) -> Result { + let raw_key = account + .raw_key + .ok_or_else(|| ContextError::Cache("raw_key required for encrypted direct open".into()))?; + let encrypted_root = account.data_dir.join("db_storage"); + let db = wx_db::WechatDb::open_encrypted_with_pool( + &encrypted_root, + raw_key, + register_mm_fts_tokenizer, + )?; + Ok(db) +} diff --git a/crates/wx-context/src/patch_state.rs b/crates/wx-context/src/patch_state.rs new file mode 100644 index 0000000..80849fa --- /dev/null +++ b/crates/wx-context/src/patch_state.rs @@ -0,0 +1,168 @@ +use std::path::{Path, PathBuf}; + +use crate::cache::format_system_time; + +/// Returns the path for the WAL failed marker: `dst.with_extension("db.wal_failed")`. +pub(crate) fn wal_failed_marker_path(dst: &Path) -> PathBuf { + dst.with_extension("db.wal_failed") +} + +/// Checks whether a previous WAL patch attempt failed for the current `(db_mtime, wal_mtime)`. +/// +/// Returns `true` if the marker exists and its composite key matches the current mtimes, +/// meaning the same failing patch should not be retried. +pub(crate) fn is_wal_patch_failed(src: &Path, wal: &Path, dst: &Path) -> bool { + let marker = wal_failed_marker_path(dst); + let recorded = match std::fs::read_to_string(&marker) { + Ok(s) => s, + Err(_) => return false, + }; + let current = match composite_key(src, wal) { + Some(k) => k, + None => return false, + }; + recorded.trim() == current +} + +/// Writes a WAL failed marker with the composite key `"db_mtime_nanos:wal_mtime_nanos"`. +pub(crate) fn write_wal_failed_marker( + src: &Path, + wal: &Path, + dst: &Path, +) -> Result<(), std::io::Error> { + let key = composite_key(src, wal) + .ok_or_else(|| std::io::Error::other("cannot read source mtimes"))?; + let marker = wal_failed_marker_path(dst); + std::fs::write(&marker, key) +} + +/// Removes the WAL failed marker. Ignores "not found" errors. +pub(crate) fn clear_wal_failed_marker(dst: &Path) { + let marker = wal_failed_marker_path(dst); + let _ = std::fs::remove_file(&marker); +} + +/// Builds the composite key `"db_mtime_nanos:wal_mtime_nanos"`. +fn composite_key(src: &Path, wal: &Path) -> Option { + let db_mtime = src.metadata().ok()?.modified().ok()?; + let wal_mtime = wal.metadata().ok()?.modified().ok()?; + Some(format!( + "{}:{}", + format_system_time(db_mtime), + format_system_time(wal_mtime), + )) +} + +#[cfg(test)] +mod tests { + use super::*; + use std::time::Duration; + use tempfile::TempDir; + + #[test] + fn marker_path_format() { + let p = PathBuf::from("/cache/message_0.db"); + assert_eq!( + wal_failed_marker_path(&p), + PathBuf::from("/cache/message_0.db.wal_failed") + ); + } + + #[test] + fn write_then_detect_failed() { + let tmp = TempDir::new().unwrap(); + let src = tmp.path().join("test.db"); + let wal = tmp.path().join("test.db-wal"); + let dst = tmp.path().join("out/test.db"); + std::fs::create_dir_all(tmp.path().join("out")).unwrap(); + + std::fs::write(&src, b"db-data").unwrap(); + std::fs::write(&wal, b"wal-data").unwrap(); + std::fs::write(&dst, b"cached").unwrap(); + + // Before writing marker, should return false + assert!(!is_wal_patch_failed(&src, &wal, &dst)); + + // Write marker + write_wal_failed_marker(&src, &wal, &dst).unwrap(); + + // Now should return true + assert!(is_wal_patch_failed(&src, &wal, &dst)); + } + + #[test] + fn mtime_change_allows_retry() { + let tmp = TempDir::new().unwrap(); + let src = tmp.path().join("test.db"); + let wal = tmp.path().join("test.db-wal"); + let dst = tmp.path().join("out/test.db"); + std::fs::create_dir_all(tmp.path().join("out")).unwrap(); + + std::fs::write(&src, b"db-data").unwrap(); + std::fs::write(&wal, b"wal-data").unwrap(); + std::fs::write(&dst, b"cached").unwrap(); + + write_wal_failed_marker(&src, &wal, &dst).unwrap(); + assert!(is_wal_patch_failed(&src, &wal, &dst)); + + // Modify WAL file to change its mtime + let new_mtime = filetime::FileTime::from_system_time( + std::time::SystemTime::now() + Duration::from_secs(2), + ); + filetime::set_file_mtime(&wal, new_mtime).unwrap(); + + // Composite key no longer matches → retry allowed + assert!(!is_wal_patch_failed(&src, &wal, &dst)); + } + + #[test] + fn db_mtime_change_allows_retry() { + let tmp = TempDir::new().unwrap(); + let src = tmp.path().join("test.db"); + let wal = tmp.path().join("test.db-wal"); + let dst = tmp.path().join("out/test.db"); + std::fs::create_dir_all(tmp.path().join("out")).unwrap(); + + std::fs::write(&src, b"db-data").unwrap(); + std::fs::write(&wal, b"wal-data").unwrap(); + std::fs::write(&dst, b"cached").unwrap(); + + write_wal_failed_marker(&src, &wal, &dst).unwrap(); + assert!(is_wal_patch_failed(&src, &wal, &dst)); + + // Modify DB file mtime + let new_mtime = filetime::FileTime::from_system_time( + std::time::SystemTime::now() + Duration::from_secs(2), + ); + filetime::set_file_mtime(&src, new_mtime).unwrap(); + + assert!(!is_wal_patch_failed(&src, &wal, &dst)); + } + + #[test] + fn clear_marker_allows_retry() { + let tmp = TempDir::new().unwrap(); + let src = tmp.path().join("test.db"); + let wal = tmp.path().join("test.db-wal"); + let dst = tmp.path().join("out/test.db"); + std::fs::create_dir_all(tmp.path().join("out")).unwrap(); + + std::fs::write(&src, b"db-data").unwrap(); + std::fs::write(&wal, b"wal-data").unwrap(); + std::fs::write(&dst, b"cached").unwrap(); + + write_wal_failed_marker(&src, &wal, &dst).unwrap(); + assert!(is_wal_patch_failed(&src, &wal, &dst)); + + clear_wal_failed_marker(&dst); + assert!(!is_wal_patch_failed(&src, &wal, &dst)); + } + + #[test] + fn clear_nonexistent_marker_is_noop() { + let tmp = TempDir::new().unwrap(); + let dst = tmp.path().join("out/test.db"); + // Should not panic + clear_wal_failed_marker(&dst); + } +} diff --git a/crates/wx-context/src/progress.rs b/crates/wx-context/src/progress.rs new file mode 100644 index 0000000..40b9c88 --- /dev/null +++ b/crates/wx-context/src/progress.rs @@ -0,0 +1,96 @@ +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::Mutex; + +// ── Public types ── + +#[non_exhaustive] +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum DecryptProgress { + Starting { total: usize }, + Decrypting { path: String }, + Decrypted { path: String, wal_patched: bool }, + Skipped { path: String, wal_patched: bool }, + Failed { path: String, error: String }, +} + +pub struct DecryptStats { + pub decrypted: usize, + pub skipped: usize, + pub errors: usize, + pub wal_patched: usize, + /// Non-fatal warnings (e.g. WAL decrypt failures) for the CLI to display. + pub warnings: Vec, +} + +// ── Internal types ── + +pub(crate) enum DbOutcome { + Decrypted { + wal_patched: bool, + wal_warnings: Vec, + }, + Skipped { + wal_patched: bool, + wal_warnings: Vec, + }, + Failed { + warning: String, + }, +} + +#[derive(Debug)] +pub(crate) enum WalOutcome { + Patched, + NothingToDo, + Warning(String), +} + +pub(crate) struct AtomicStats { + decrypted: AtomicUsize, + skipped: AtomicUsize, + errors: AtomicUsize, + wal_patched: AtomicUsize, + warnings: Mutex>, +} + +impl AtomicStats { + pub(crate) fn new() -> Self { + Self { + decrypted: AtomicUsize::new(0), + skipped: AtomicUsize::new(0), + errors: AtomicUsize::new(0), + wal_patched: AtomicUsize::new(0), + warnings: Mutex::new(Vec::new()), + } + } + + pub(crate) fn inc_decrypted(&self) { + self.decrypted.fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn inc_skipped(&self) { + self.skipped.fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn inc_errors(&self) { + self.errors.fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn inc_wal_patched(&self) { + self.wal_patched.fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn add_warning(&self, msg: String) { + self.warnings.lock().unwrap().push(msg); + } + + pub(crate) fn into_stats(self) -> DecryptStats { + DecryptStats { + decrypted: self.decrypted.into_inner(), + skipped: self.skipped.into_inner(), + errors: self.errors.into_inner(), + wal_patched: self.wal_patched.into_inner(), + warnings: self.warnings.into_inner().unwrap(), + } + } +} diff --git a/crates/wx-context/src/shard_routing.rs b/crates/wx-context/src/shard_routing.rs new file mode 100644 index 0000000..d4f4e2c --- /dev/null +++ b/crates/wx-context/src/shard_routing.rs @@ -0,0 +1,245 @@ +//! Shard routing: determine which message shard IDs to decrypt based on +//! cached shard metadata and query time bounds. + +use std::path::Path; + +use wx_db::shard_metadata::{read_shard_metadata, write_shard_metadata, ShardMetadataFile}; +use wx_db::WechatDb; + +/// Determine which shard IDs to decrypt for a given query. +/// +/// Returns `Some(shard_ids)` when routing can reduce the set of shards needed. +/// Returns `None` when routing is not possible or not beneficial (caller should +/// fall back to decrypting all shards). +/// +/// Routing only activates when at least one of `since`/`until` is explicitly +/// provided. Without explicit time bounds, we cannot safely determine a subset. +pub fn route_shards_for_query( + decrypted_root: &Path, + talker: &str, + since: Option, + until: Option, +) -> Option> { + // Only route when explicit time bounds are provided + if since.is_none() && until.is_none() { + return None; + } + + // Read cached shard metadata + let msg_dir = decrypted_root.join("message"); + let meta = read_shard_metadata(&msg_dir)?; + + // Staleness check: verify no shard DB has been modified after metadata was written + if is_metadata_stale(&meta, &msg_dir) { + return None; + } + + // Look up talker's sort_timestamp from session.db for upper bound + let session_path = decrypted_root.join("session").join("session.db"); + let sort_timestamp = read_sort_timestamp(&session_path, talker)?; + + // Compute effective time range + let start = since.unwrap_or(0); + let end = until.unwrap_or(sort_timestamp); + + // Filter shards overlapping [start, end] + let filtered: Vec = meta + .shards + .iter() + .filter(|s| s.start_unix <= end && s.end_unix >= start) + .map(|s| s.shard_id) + .collect(); + + // No meaningful reduction → fall back + if filtered.is_empty() || filtered.len() >= meta.shards.len() { + return None; + } + + Some(filtered) +} + +/// Write the shard metadata sidecar from a `WechatDb` instance. +/// This keeps the write responsibility in wx-context, not in wx-db. +pub fn write_shard_metadata_sidecar( + db: &WechatDb, + decrypted_root: &Path, +) -> Result<(), std::io::Error> { + let meta = db.shard_metadata(); + let msg_dir = decrypted_root.join("message"); + if msg_dir.is_dir() { + write_shard_metadata(&msg_dir, &meta)?; + } + Ok(()) +} + +/// Check whether the metadata is stale by comparing shard DB mtimes. +fn is_metadata_stale(meta: &ShardMetadataFile, msg_dir: &Path) -> bool { + for shard in &meta.shards { + let db_path = msg_dir.join(format!("message_{}.db", shard.shard_id)); + if let Ok(file_meta) = std::fs::metadata(&db_path) { + if let Ok(mtime) = file_meta.modified() { + let mtime_nanos = mtime + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_nanos(); + if mtime_nanos > meta.written_at_ns { + return true; + } + } + } + // If we can't read mtime, be conservative — don't mark stale + } + false +} + +/// Read a talker's `sort_timestamp` from session.db. +fn read_sort_timestamp(session_path: &Path, talker: &str) -> Option { + let conn = rusqlite::Connection::open_with_flags( + session_path, + rusqlite::OpenFlags::SQLITE_OPEN_READ_ONLY, + ) + .ok()?; + conn.query_row( + "SELECT sort_timestamp FROM SessionTable WHERE username = ?1", + [talker], + |row| row.get(0), + ) + .ok() +} + +#[cfg(test)] +mod tests { + use super::*; + use tempfile::TempDir; + + /// Create a minimal test environment with session.db and shard-metadata.json. + fn setup_test_env(shards: &[(u32, i64, i64)], talker: &str, sort_ts: i64) -> TempDir { + let dir = TempDir::new().unwrap(); + let root = dir.path(); + + // Create session dir + session.db + let session_dir = root.join("session"); + std::fs::create_dir_all(&session_dir).unwrap(); + let session_path = session_dir.join("session.db"); + let conn = rusqlite::Connection::open(&session_path).unwrap(); + conn.execute_batch( + "CREATE TABLE SessionTable (username TEXT, sort_timestamp INTEGER, summary TEXT)", + ) + .unwrap(); + conn.execute( + "INSERT INTO SessionTable (username, sort_timestamp) VALUES (?1, ?2)", + rusqlite::params![talker, sort_ts], + ) + .unwrap(); + + // Create message dir + shard files + metadata + let msg_dir = root.join("message"); + std::fs::create_dir_all(&msg_dir).unwrap(); + + let meta_shards: Vec = shards + .iter() + .map(|&(id, start, end)| { + // Create a dummy shard file + std::fs::write(msg_dir.join(format!("message_{id}.db")), b"dummy").unwrap(); + wx_db::shard_metadata::ShardMeta { + shard_id: id, + start_unix: start, + end_unix: end, + } + }) + .collect(); + + let meta = wx_db::shard_metadata::ShardMetadataFile { + shards: meta_shards, + written_at_ns: wx_db::shard_metadata::now_nanos(), + }; + wx_db::shard_metadata::write_shard_metadata(&msg_dir, &meta).unwrap(); + + dir + } + + #[test] + fn routes_to_subset_with_since() { + // 3 shards: [1000,2000], [2001,3000], [3001,i64::MAX] + let dir = setup_test_env( + &[(0, 1000, 2000), (1, 2001, 3000), (2, 3001, i64::MAX)], + "wxid_alice", + 5000, + ); + let result = route_shards_for_query(dir.path(), "wxid_alice", Some(2500), None); + // since=2500 → overlaps shard 1 ([2001,3000]) and shard 2 ([3001,MAX]) + // end = sort_timestamp = 5000 → overlaps shards 1, 2 + assert_eq!(result, Some(vec![1, 2])); + } + + #[test] + fn routes_to_subset_with_until() { + let dir = setup_test_env( + &[(0, 1000, 2000), (1, 2001, 3000), (2, 3001, i64::MAX)], + "wxid_alice", + 5000, + ); + let result = route_shards_for_query(dir.path(), "wxid_alice", None, Some(1500)); + // start=0, end=1500 → only overlaps shard 0 [1000,2000] + assert_eq!(result, Some(vec![0])); + } + + #[test] + fn routes_to_subset_with_since_and_until() { + let dir = setup_test_env( + &[(0, 1000, 2000), (1, 2001, 3000), (2, 3001, i64::MAX)], + "wxid_alice", + 5000, + ); + let result = route_shards_for_query(dir.path(), "wxid_alice", Some(1500), Some(2500)); + // [1500,2500] → overlaps shard 0 and shard 1 + assert_eq!(result, Some(vec![0, 1])); + } + + #[test] + fn no_time_bounds_returns_none() { + let dir = setup_test_env(&[(0, 1000, 2000), (1, 2001, 3000)], "wxid_alice", 5000); + let result = route_shards_for_query(dir.path(), "wxid_alice", None, None); + assert!(result.is_none()); + } + + #[test] + fn unknown_talker_returns_none() { + let dir = setup_test_env(&[(0, 1000, 2000), (1, 2001, 3000)], "wxid_alice", 5000); + let result = route_shards_for_query(dir.path(), "wxid_unknown", Some(1500), None); + assert!(result.is_none()); + } + + #[test] + fn missing_metadata_returns_none() { + let dir = TempDir::new().unwrap(); + let result = route_shards_for_query(dir.path(), "wxid_alice", Some(1000), None); + assert!(result.is_none()); + } + + #[test] + fn all_shards_match_returns_none() { + // If routing includes all shards, no benefit + let dir = setup_test_env(&[(0, 1000, 2000), (1, 2001, i64::MAX)], "wxid_alice", 5000); + let result = route_shards_for_query(dir.path(), "wxid_alice", Some(500), None); + // [500, 5000] overlaps both shards → returns None (no reduction) + assert!(result.is_none()); + } + + #[test] + fn stale_metadata_returns_none() { + let dir = setup_test_env( + &[(0, 1000, 2000), (1, 2001, 3000), (2, 3001, i64::MAX)], + "wxid_alice", + 5000, + ); + + // Touch a shard file to make it newer than the metadata + std::thread::sleep(std::time::Duration::from_millis(50)); + let shard_path = dir.path().join("message").join("message_0.db"); + std::fs::write(&shard_path, b"updated").unwrap(); + + let result = route_shards_for_query(dir.path(), "wxid_alice", Some(2500), None); + assert!(result.is_none(), "stale metadata should trigger fallback"); + } +} diff --git a/crates/wx-context/src/tokenizer.rs b/crates/wx-context/src/tokenizer.rs new file mode 100644 index 0000000..c3afd80 --- /dev/null +++ b/crates/wx-context/src/tokenizer.rs @@ -0,0 +1,384 @@ +use rust_stemmers::{Algorithm, Stemmer}; + +/// Configuration flags for the MMFtsTokenizer. +/// +/// Mirrors WCDB's `OneOrBinaryTokenizer` config parameters. +#[derive(Debug, Clone, Default)] +pub struct TokenizerConfig { + /// Emit each special/symbol character as a token. + pub enable_special_char: bool, + /// Emit ASCII digit sequences as tokens. + pub enable_num_token: bool, + /// Skip Porter stemming (emit lowercased form only). + pub skip_stemming: bool, +} + +/// A single token emitted by the tokenizer. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct Token { + pub text: String, + /// Byte offset of the first byte in the input. + pub start: usize, + /// Byte offset one past the last byte in the input. + pub end: usize, + /// True for colocated (synonym/variant) tokens — they share position with + /// the preceding primary token (FTS5_TOKEN_COLOCATED). + pub colocated: bool, +} + +impl Token { + fn primary(text: String, start: usize, end: usize) -> Self { + Token { + text, + start, + end, + colocated: false, + } + } + + fn colocated(text: String, start: usize, end: usize) -> Self { + Token { + text, + start, + end, + colocated: true, + } + } +} + +/// Tokenize `input` according to WCDB's `OneOrBinaryTokenizer` (m_needBinary=false). +/// +/// - ASCII letters → one token (lowercased + Porter-stemmed; original lowercased +/// form emitted as colocated token if it differs from stem) +/// - ASCII digits → one token if `enable_num_token`, else skipped +/// - CJK / BMP multi-byte / Auxiliary Plane → each codepoint is one token +/// - Symbols → emitted as individual token if `enable_special_char`, else skipped +pub fn tokenize(input: &[u8], config: &TokenizerConfig) -> Vec { + let Ok(text) = std::str::from_utf8(input) else { + return Vec::new(); + }; + if text.is_empty() { + return Vec::new(); + } + + let stemmer = if config.skip_stemming { + None + } else { + Some(Stemmer::create(Algorithm::English)) + }; + + let mut tokens = Vec::new(); + let bytes = text.as_bytes(); + let mut pos = 0; + + while pos < bytes.len() { + let byte = bytes[pos]; + + if byte < 0xC0 { + // Single-byte UTF-8 (ASCII) + let ch = byte as char; + + if ch.is_ascii_alphabetic() { + // Consume consecutive ASCII letters + let start = pos; + while pos < bytes.len() && (bytes[pos] as char).is_ascii_alphabetic() { + pos += 1; + } + let word = &text[start..pos]; + let lowercased = word.to_lowercase(); + + if let Some(ref s) = stemmer { + let stemmed = s.stem(&lowercased).into_owned(); + if stemmed != lowercased { + // Emit stemmed form as primary + tokens.push(Token::primary(stemmed, start, pos)); + // Emit original lowercased form as colocated (for exact match) + tokens.push(Token::colocated(lowercased, start, pos)); + } else { + tokens.push(Token::primary(lowercased, start, pos)); + } + } else { + tokens.push(Token::primary(lowercased, start, pos)); + } + } else if ch.is_ascii_digit() { + // Consume consecutive ASCII digits + let start = pos; + while pos < bytes.len() && (bytes[pos] as char).is_ascii_digit() { + pos += 1; + } + if config.enable_num_token { + let num_str = &text[start..pos]; + tokens.push(Token::primary(num_str.to_string(), start, pos)); + } + } else { + // Symbol or whitespace (0x00-0x2F range minus letters/digits, etc.) + if config.enable_special_char && !ch.is_ascii_whitespace() && !ch.is_ascii_control() + { + let sym = ch.to_string(); + tokens.push(Token::primary(sym, pos, pos + 1)); + } + pos += 1; + } + } else if byte < 0xF0 { + // 2-3 byte UTF-8 → BMP multi-byte character + let char_len = utf8_char_len(byte); + let start = pos; + pos += char_len; + let ch_str = &text[start..pos]; + + // Check if this is a symbol that should be emitted or skipped + let is_sym = is_symbol_char(ch_str); + if is_sym { + if config.enable_special_char { + tokens.push(Token::primary(ch_str.to_string(), start, pos)); + } + } else { + // BMP-Other (including CJK): emit as individual unigram token + tokens.push(Token::primary(ch_str.to_string(), start, pos)); + } + } else { + // 4+ byte UTF-8 → Auxiliary Plane character + let char_len = utf8_char_len(byte); + let start = pos; + pos += char_len; + let ch_str = &text[start..pos]; + // Each codepoint emitted individually + tokens.push(Token::primary(ch_str.to_string(), start, pos)); + } + } + + tokens +} + +/// Return the byte-length of a UTF-8 character given its leading byte. +fn utf8_char_len(leading: u8) -> usize { + if leading < 0x80 { + 1 + } else if leading < 0xE0 { + 2 + } else if leading < 0xF0 { + 3 + } else if leading < 0xF8 { + 4 + } else if leading < 0xFC { + 5 + } else { + 6 + } +} + +/// Heuristic: is this 2-3 byte BMP character a "symbol"? +/// +/// In practice WCDB's symbol detector is rarely configured for message_fts.db, +/// so we conservatively only mark obvious punctuation Unicode blocks as symbols. +/// This preserves CJK characters, Hangul, etc. as normal tokens. +fn is_symbol_char(s: &str) -> bool { + s.chars().all(|c| { + // General punctuation and symbol blocks + matches!(c, + '\u{2000}'..='\u{206F}' // General Punctuation + | '\u{2E00}'..='\u{2E7F}' // Supplemental Punctuation + | '\u{3000}'..='\u{303F}' // CJK Symbols and Punctuation + ) + }) +} + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +#[cfg(test)] +mod tests { + use super::*; + + fn cfg_default() -> TokenizerConfig { + TokenizerConfig::default() + } + + fn primary_texts(tokens: &[Token]) -> Vec { + tokens + .iter() + .filter(|t| !t.colocated) + .map(|t| t.text.clone()) + .collect() + } + + fn all_texts(tokens: &[Token]) -> Vec { + tokens.iter().map(|t| t.text.clone()).collect() + } + + // Test 1: Empty input → empty tokens + #[test] + fn empty_input() { + let tokens = tokenize(b"", &cfg_default()); + assert!(tokens.is_empty()); + } + + // Test 2: Pure ASCII letters: "hello" → stemmed form (no diff for "hello") + #[test] + fn pure_ascii_letters() { + let tokens = tokenize(b"hello", &cfg_default()); + let primary: Vec<_> = tokens.iter().filter(|t| !t.colocated).collect(); + assert_eq!(primary.len(), 1); + assert_eq!(primary[0].start, 0); + assert_eq!(primary[0].end, 5); + // "hello" → Porter stem → "hello" (no stem difference) + assert_eq!(primary[0].text, "hello"); + } + + // Test 3: Pure CJK: "你好世界" → 4 individual unigram tokens with correct byte offsets + #[test] + fn pure_cjk() { + let tokens = tokenize("你好世界".as_bytes(), &cfg_default()); + let primary = primary_texts(&tokens); + assert_eq!(primary, vec!["你", "好", "世", "界"]); + // Each CJK char is 3 bytes in UTF-8 + assert_eq!(tokens[0].start, 0); + assert_eq!(tokens[0].end, 3); + assert_eq!(tokens[1].start, 3); + assert_eq!(tokens[1].end, 6); + assert_eq!(tokens[2].start, 6); + assert_eq!(tokens[2].end, 9); + assert_eq!(tokens[3].start, 9); + assert_eq!(tokens[3].end, 12); + } + + // Test 4: Mixed CJK + ASCII: "你好world" → ["你", "好", "world"] + #[test] + fn mixed_cjk_ascii() { + let tokens = tokenize("你好world".as_bytes(), &cfg_default()); + let primary = primary_texts(&tokens); + assert_eq!(primary, vec!["你", "好", "world"]); + } + + // Test 5: Digits with enable_num_token + #[test] + fn digits_disabled() { + let tokens = tokenize(b"abc123", &cfg_default()); + let primary = primary_texts(&tokens); + // digits skipped by default + assert!(primary.iter().all(|t| t != "123")); + } + + #[test] + fn digits_enabled() { + let config = TokenizerConfig { + enable_num_token: true, + ..Default::default() + }; + let tokens = tokenize(b"abc123", &config); + let primary = primary_texts(&tokens); + assert!(primary.contains(&"123".to_string())); + } + + // Test 6: Symbols skip/emit + #[test] + fn symbols_default_skipped() { + let tokens = tokenize(b"hello!world", &cfg_default()); + let primary = primary_texts(&tokens); + // "!" should not appear in tokens + assert!(!primary.contains(&"!".to_string())); + } + + #[test] + fn symbols_enabled() { + let config = TokenizerConfig { + enable_special_char: true, + ..Default::default() + }; + let tokens = tokenize(b"hello!world", &config); + let texts = all_texts(&tokens); + assert!(texts.contains(&"!".to_string())); + } + + // Test 7: Porter stemming: "running" → "run" + #[test] + fn porter_stemming_running() { + let tokens = tokenize(b"running", &cfg_default()); + // Primary token should be the stemmed form + let primary: Vec<_> = tokens.iter().filter(|t| !t.colocated).collect(); + assert_eq!(primary.len(), 1); + assert_eq!(primary[0].text, "run"); + // Colocated token should be the original lowercased form + let colocated: Vec<_> = tokens.iter().filter(|t| t.colocated).collect(); + assert_eq!(colocated.len(), 1); + assert_eq!(colocated[0].text, "running"); + } + + // Test 8: Consecutive types: "abc123你好" (digits disabled by default) + #[test] + fn consecutive_types() { + let tokens = tokenize("abc123你好".as_bytes(), &cfg_default()); + let primary = primary_texts(&tokens); + // abc → stemmed (no stem diff for "abc"), then CJK chars + assert!(primary.contains(&"你".to_string())); + assert!(primary.contains(&"好".to_string())); + // digits skipped + assert!(!primary.contains(&"123".to_string())); + } + + #[test] + fn consecutive_types_num_enabled() { + let config = TokenizerConfig { + enable_num_token: true, + ..Default::default() + }; + let tokens = tokenize("abc123你好".as_bytes(), &config); + let primary = primary_texts(&tokens); + assert!(primary.contains(&"123".to_string())); + assert!(primary.contains(&"abc".to_string())); + assert!(primary.contains(&"你".to_string())); + assert!(primary.contains(&"好".to_string())); + } + + // Test 9: 4-byte emoji "😀" → single token + #[test] + fn emoji_single_token() { + let emoji = "😀"; + let tokens = tokenize(emoji.as_bytes(), &cfg_default()); + assert_eq!(tokens.len(), 1); + assert_eq!(tokens[0].text, "😀"); + assert_eq!(tokens[0].start, 0); + assert_eq!(tokens[0].end, 4); // 😀 is 4 bytes in UTF-8 + } + + // Test 10: Fullwidth "!" (U+FF01) — in Fullwidth Forms, treated as BMP-Other + // With enable_special_char=false, if it's classified as symbol it should be skipped. + // Since U+FF01 is NOT in our symbol ranges, it's emitted as a normal token. + #[test] + fn fullwidth_forms() { + let s = "!"; // U+FF01 fullwidth exclamation mark (3 bytes) + let tokens = tokenize(s.as_bytes(), &cfg_default()); + // U+FF01 is NOT in our symbol ranges → emitted as BMP-Other token + assert_eq!(tokens.len(), 1); + assert_eq!(tokens[0].text, "!"); + } + + // Additional: skip_stemming mode + #[test] + fn skip_stemming_mode() { + let config = TokenizerConfig { + skip_stemming: true, + ..Default::default() + }; + let tokens = tokenize(b"running", &config); + let primary = primary_texts(&tokens); + assert_eq!(primary, vec!["running"]); + // No colocated tokens + assert!(!tokens.iter().any(|t| t.colocated)); + } + + // Additional: byte offsets for mixed content + #[test] + fn byte_offsets_mixed() { + // "hi你" — "hi" is 2 bytes, "你" starts at byte 2 + let tokens = tokenize("hi你".as_bytes(), &cfg_default()); + let primary: Vec<_> = tokens.iter().filter(|t| !t.colocated).collect(); + assert_eq!(primary[0].text, "hi"); + assert_eq!(primary[0].start, 0); + assert_eq!(primary[0].end, 2); + assert_eq!(primary[1].text, "你"); + assert_eq!(primary[1].start, 2); + assert_eq!(primary[1].end, 5); // 你 is 3 bytes + } +} diff --git a/crates/wx-context/src/visibility.rs b/crates/wx-context/src/visibility.rs new file mode 100644 index 0000000..0d60af6 --- /dev/null +++ b/crates/wx-context/src/visibility.rs @@ -0,0 +1,269 @@ +use std::collections::HashSet; + +use crate::ContactResolver; + +/// Compiled visibility index for talker-level and sender-level hiding. +/// +/// Built once from account settings + ContactResolver. +/// `ignore_contacts` and `ignore_tags` cascade to both talker-level and sender-level hiding +/// through a single `hidden_persons` set. +pub struct VisibilityIndex { + hidden_persons: HashSet, +} + +impl VisibilityIndex { + /// Empty index that hides nothing. + pub fn empty() -> Self { + Self { + hidden_persons: HashSet::new(), + } + } + + /// Build from raw rule lists and a ContactResolver. + /// + /// - `ignore_contacts`: direct wxid/chatroom IDs to hide + /// - `ignore_tags`: tag names; any contact matching these tags is hidden + /// + /// Both `ignore_contacts` and tag-expanded contacts cascade to sender-level hiding + /// within visible group chats. + pub fn build( + ignore_contacts: &[String], + ignore_tags: &[String], + resolver: &ContactResolver, + ) -> Self { + let mut hidden_persons: HashSet = ignore_contacts.iter().cloned().collect(); + + if !ignore_tags.is_empty() { + let tag_set: HashSet<&str> = ignore_tags.iter().map(|s| s.as_str()).collect(); + for (wxid, labels) in resolver.all_labels() { + if labels.iter().any(|l| tag_set.contains(l.as_str())) { + hidden_persons.insert(wxid.clone()); + } + } + } + + Self { hidden_persons } + } + + /// Whether this talker is hidden. + pub fn is_hidden_talker(&self, talker: &str) -> bool { + self.hidden_persons.contains(talker) + } + + /// Whether a target talker can be resolved (i.e. is not hidden). + /// Used by the shared contact resolution entry point. + pub fn can_resolve_target(&self, talker: &str) -> bool { + !self.hidden_persons.contains(talker) + } + + /// Whether media access is allowed for this talker. + pub fn allows_media(&self, talker: &str) -> bool { + !self.hidden_persons.contains(talker) + } + + /// Whether this sender should be hidden within a visible group chat. + /// + /// Returns true only when: + /// - talker ends with `@chatroom` (is a group) + /// - talker is NOT in hidden_persons (the group itself is visible) + /// - sender IS in hidden_persons + pub fn is_hidden_sender_in_group(&self, talker: &str, sender: &str) -> bool { + wx_db::is_group_chat(talker) + && !self.hidden_persons.contains(talker) + && self.hidden_persons.contains(sender) + } + + /// Whether media access is allowed for this talker + sender combination. + /// + /// Hidden talkers OR hidden senders in visible groups cannot access media. + pub fn allows_media_for_sender(&self, talker: &str, sender: &str) -> bool { + !self.hidden_persons.contains(talker) + && !self.is_hidden_sender_in_group(talker, sender) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + use std::fs; + use std::path::Path; + + use rusqlite::{params, Connection}; + use tempfile::TempDir; + use wx_db::{encode_extra_buffer_for_test, WechatDb}; + + #[test] + fn empty_index_hides_nothing() { + let idx = VisibilityIndex::empty(); + assert!(!idx.is_hidden_talker("wxid_test")); + assert!(idx.can_resolve_target("wxid_test")); + assert!(idx.allows_media("wxid_test")); + } + + #[test] + fn direct_ignore_contacts() { + let idx = VisibilityIndex { + hidden_persons: vec!["wxid_hidden".to_string(), "group@chatroom".to_string()] + .into_iter() + .collect(), + }; + assert!(idx.is_hidden_talker("wxid_hidden")); + assert!(idx.is_hidden_talker("group@chatroom")); + assert!(!idx.can_resolve_target("wxid_hidden")); + assert!(!idx.allows_media("wxid_hidden")); + assert!(!idx.is_hidden_talker("wxid_visible")); + assert!(idx.can_resolve_target("wxid_visible")); + } + + #[test] + fn build_expands_ignore_tags_to_matching_contact_talkers_only() { + let fixture = create_contact_fixture(&[ + ("wxid_tagged", "Tagged", Some("1")), + ("team@chatroom", "Group", None), + ("wxid_visible", "Visible", None), + ]); + let db = WechatDb::open(fixture.path()).unwrap(); + let resolver = ContactResolver::build(&db).unwrap(); + + let idx = VisibilityIndex::build(&[], &["Sensitive".to_string()], &resolver); + + assert!(idx.is_hidden_talker("wxid_tagged")); + assert!(!idx.is_hidden_talker("team@chatroom")); + assert!(!idx.is_hidden_talker("wxid_visible")); + } + + fn create_contact_fixture(entries: &[(&str, &str, Option<&str>)]) -> TempDir { + let dir = TempDir::new().unwrap(); + let contact_dir = dir.path().join("contact"); + let session_dir = dir.path().join("session"); + let message_dir = dir.path().join("message"); + fs::create_dir_all(&contact_dir).unwrap(); + fs::create_dir_all(&session_dir).unwrap(); + fs::create_dir_all(&message_dir).unwrap(); + + let contact_db = contact_dir.join("contact.db"); + let conn = Connection::open(&contact_db).unwrap(); + conn.execute_batch( + "CREATE TABLE contact ( + username TEXT PRIMARY KEY, + alias TEXT DEFAULT '', + remark TEXT DEFAULT '', + nick_name TEXT DEFAULT '', + description TEXT DEFAULT NULL, + extra_buffer BLOB DEFAULT NULL + ); + CREATE TABLE contact_label ( + label_id_ TEXT, + label_name_ TEXT, + sort_order_ INTEGER + );", + ) + .unwrap(); + conn.execute( + "INSERT INTO contact_label VALUES (?1, ?2, ?3)", + params!["1", "Sensitive", 0], + ) + .unwrap(); + + for (username, nickname, label_ids_csv) in entries { + let extra_buffer = label_ids_csv.map(|ids| { + encode_extra_buffer_for_test(None, None, None, None, None, None, None, Some(ids)) + }); + conn.execute( + "INSERT INTO contact (username, nick_name, extra_buffer) VALUES (?1, ?2, ?3)", + params![username, nickname, extra_buffer], + ) + .unwrap(); + } + + create_empty_session_db(&session_dir.join("session.db")); + dir + } + + fn create_empty_session_db(path: &Path) { + let conn = Connection::open(path).unwrap(); + conn.execute_batch( + "CREATE TABLE SessionTable ( + username TEXT, + sort_timestamp INTEGER, + summary TEXT, + last_msg_type INTEGER DEFAULT NULL, + last_msg_sender TEXT DEFAULT NULL, + last_sender_display_name TEXT DEFAULT NULL + );", + ) + .unwrap(); + } + + // --- Sender-level hiding tests --- + + #[test] + fn is_hidden_sender_in_group_non_chatroom_returns_false() { + let idx = VisibilityIndex { + hidden_persons: vec!["wxid_spam".to_string()].into_iter().collect(), + }; + // Private chat — sender hiding does not apply + assert!(!idx.is_hidden_sender_in_group("wxid_spam", "wxid_spam")); + } + + #[test] + fn is_hidden_sender_in_group_hidden_talker_returns_false() { + let idx = VisibilityIndex { + hidden_persons: vec!["group@chatroom".to_string(), "wxid_spam".to_string()] + .into_iter() + .collect(), + }; + // Talker-level hiding already handles this group — sender check not needed + assert!(!idx.is_hidden_sender_in_group("group@chatroom", "wxid_spam")); + } + + #[test] + fn is_hidden_sender_in_group_visible_group_hidden_sender() { + let idx = VisibilityIndex { + hidden_persons: vec!["wxid_spam".to_string()].into_iter().collect(), + }; + assert!(idx.is_hidden_sender_in_group("group@chatroom", "wxid_spam")); + assert!(!idx.is_hidden_sender_in_group("group@chatroom", "wxid_visible")); + } + + #[test] + fn allows_media_for_sender_covers_both_levels() { + let idx = VisibilityIndex { + hidden_persons: vec![ + "hidden_group@chatroom".to_string(), + "wxid_spam".to_string(), + ] + .into_iter() + .collect(), + }; + // Hidden talker → no media + assert!(!idx.allows_media_for_sender("hidden_group@chatroom", "wxid_anyone")); + // Visible group + hidden sender → no media + assert!(!idx.allows_media_for_sender("visible@chatroom", "wxid_spam")); + // Visible group + visible sender → media allowed + assert!(idx.allows_media_for_sender("visible@chatroom", "wxid_normal")); + // Hidden person as talker in private chat → no media (unified model) + assert!(!idx.allows_media_for_sender("wxid_spam", "wxid_spam")); + // Non-hidden person in private chat → media allowed + assert!(idx.allows_media_for_sender("wxid_normal", "wxid_normal")); + } + + #[test] + fn build_expands_ignore_tags_to_sender_level() { + let fixture = create_contact_fixture(&[ + ("wxid_tagged", "Tagged", Some("1")), + ("team@chatroom", "Group", None), + ("wxid_visible", "Visible", None), + ]); + let db = WechatDb::open(fixture.path()).unwrap(); + let resolver = ContactResolver::build(&db).unwrap(); + + let idx = VisibilityIndex::build(&[], &["Sensitive".to_string()], &resolver); + + // Tag-expanded contact should trigger sender-level hiding in visible groups + assert!(idx.is_hidden_sender_in_group("team@chatroom", "wxid_tagged")); + // Non-tagged contacts should not be hidden + assert!(!idx.is_hidden_sender_in_group("team@chatroom", "wxid_visible")); + } +} diff --git a/crates/wx-context/src/wal_patch.rs b/crates/wx-context/src/wal_patch.rs new file mode 100644 index 0000000..6b5c51e --- /dev/null +++ b/crates/wx-context/src/wal_patch.rs @@ -0,0 +1,299 @@ +use std::path::Path; + +use wx_decrypt::{dispatch_decrypt_wal, CryptoParams, DecryptError, KeyMaterial}; + +/// Result of an atomic WAL patch attempt. +/// +/// Distinguishes *content* failures (bad data that won't self-heal) from +/// *environment/IO* failures (transient issues that may resolve on retry). +#[derive(Debug)] +pub(crate) enum WalPatchResult { + /// Successfully patched N frames into the destination DB. + Patched(usize), + /// WAL contained no valid frames — destination unchanged. + NoFrames, + /// HMAC / decryption / WAL-header failure — the WAL content itself is bad. + /// Caller should write a failed marker; retrying with the same mtime is pointless. + ContentFailed(String), + /// Copy / rename / disk-full or other transient I/O error. + /// Caller should log a warning but NOT write a failed marker — next run may succeed. + IoFailed(String), +} + +/// Atomically apply encrypted WAL frames to a cached (plaintext) DB. +/// +/// Sequence: copy `dst` to tmp -> decrypt WAL frames into tmp -> rename tmp over `dst`. +/// On any failure the tmp file is cleaned up and `dst` is left untouched. +pub(crate) fn apply_wal_patch( + wal: &Path, + dst: &Path, + km: &KeyMaterial, + params: &CryptoParams, +) -> WalPatchResult { + let tmp_path = dst.with_extension("db.tmp"); + + // Clean up any residual tmp from a previous crash. + std::fs::remove_file(&tmp_path).ok(); + + let result = (|| -> WalPatchResult { + // Step 1: copy dst -> tmp + if let Err(e) = std::fs::copy(dst, &tmp_path) { + return WalPatchResult::IoFailed(format!("copy dst to tmp: {e}")); + } + + // Step 2: decrypt WAL frames into tmp + let frame_count = match dispatch_decrypt_wal(wal, &tmp_path, km, params) { + Ok(n) => n, + Err(e) => return classify_decrypt_error(e), + }; + + if frame_count == 0 { + return WalPatchResult::NoFrames; + } + + // Step 3: atomic rename tmp -> dst + if let Err(e) = std::fs::rename(&tmp_path, dst) { + return WalPatchResult::IoFailed(format!("rename tmp to dst: {e}")); + } + + WalPatchResult::Patched(frame_count) + })(); + + // Ensure tmp is cleaned up on any non-success path. + // For Patched, the rename already consumed the tmp file. + // For all other variants (including NoFrames), remove the tmp if it exists. + if !matches!(result, WalPatchResult::Patched(_)) { + std::fs::remove_file(&tmp_path).ok(); + } + + result +} + +/// Classify a `DecryptError` into either `ContentFailed` or `IoFailed`. +fn classify_decrypt_error(e: DecryptError) -> WalPatchResult { + match &e { + // Content failures — the WAL data itself is bad. + DecryptError::HmacVerificationFailed { .. } + | DecryptError::AesDecryptFailed { .. } + | DecryptError::InvalidWalHeader { .. } + | DecryptError::IncorrectKey => WalPatchResult::ContentFailed(e.to_string()), + + // Environment / transient failures — may self-heal on retry. + DecryptError::SaltMismatch + | DecryptError::NoMatchingEncKey + | DecryptError::Io(_) + | DecryptError::FileTooSmall { .. } + | DecryptError::AlreadyDecrypted => WalPatchResult::IoFailed(e.to_string()), + } +} + +#[cfg(test)] +mod tests { + use super::*; + use tempfile::TempDir; + + /// Helper: create a minimal valid SQLite database at `path`. + fn create_sqlite_db(path: &Path) { + let conn = rusqlite::Connection::open(path).unwrap(); + conn.execute_batch("CREATE TABLE t(x INTEGER);").unwrap(); + } + + // --------------------------------------------------------------- + // Acceptance: successful patch — tmp renamed to dst, content updated + // --------------------------------------------------------------- + // NOTE: Testing the full happy path requires a real encrypted WAL + matching + // KeyMaterial, which is impractical in a unit test. Instead we test the + // observable contract via the sub-components and a decrypt-returns-0 scenario. + + #[test] + fn no_frames_leaves_dst_intact_and_cleans_tmp() { + // When dispatch_decrypt_wal returns Ok(0), apply_wal_patch should + // return NoFrames, leave dst unchanged, and remove the tmp file. + // + // We simulate this by providing a WAL file that is a valid but empty + // WAL (just the 32-byte header with zero frames). + let tmp = TempDir::new().unwrap(); + let dst = tmp.path().join("test.db"); + create_sqlite_db(&dst); + let original = std::fs::read(&dst).unwrap(); + + // Create a minimal WAL header (32 bytes). The magic and version are + // enough for dispatch_decrypt_wal to parse but find 0 frames. + let wal = tmp.path().join("test.db-wal"); + let mut wal_header = vec![0u8; 32]; + // WAL magic: 0x377f0682 (little-endian WAL) or 0x377f0683 (big-endian) + wal_header[..4].copy_from_slice(&0x377f0682u32.to_be_bytes()); + // File format version: 3007000 + wal_header[4..8].copy_from_slice(&3007000u32.to_be_bytes()); + // Page size: 4096 + wal_header[8..12].copy_from_slice(&4096u32.to_be_bytes()); + std::fs::write(&wal, &wal_header).unwrap(); + + let km = KeyMaterial::EncKey { + key: [0u8; 32], + salt: [0u8; 16], + }; + let params = wx_decrypt::MACOS_4_1_7_31; + + let result = apply_wal_patch(&wal, &dst, &km, ¶ms); + assert!( + matches!(result, WalPatchResult::NoFrames), + "expected NoFrames, got {result:?}" + ); + + // dst should be unchanged + assert_eq!(std::fs::read(&dst).unwrap(), original); + + // tmp should be cleaned up + let tmp_path = dst.with_extension("db.tmp"); + assert!(!tmp_path.exists(), "tmp file should be removed"); + } + + // --------------------------------------------------------------- + // Acceptance: decrypt failure — dst preserved, tmp cleaned up + // --------------------------------------------------------------- + #[test] + fn decrypt_failure_preserves_dst_and_cleans_tmp() { + let tmp = TempDir::new().unwrap(); + let dst = tmp.path().join("test.db"); + create_sqlite_db(&dst); + let original = std::fs::read(&dst).unwrap(); + + // Create a WAL with a valid-looking header but garbage frame data. + // This should cause an HMAC or decrypt error. + let wal = tmp.path().join("test.db-wal"); + let mut wal_data = vec![0u8; 32 + 24 + 4096]; // header + frame header + one page + wal_data[..4].copy_from_slice(&0x377f0682u32.to_be_bytes()); + wal_data[4..8].copy_from_slice(&3007000u32.to_be_bytes()); + wal_data[8..12].copy_from_slice(&4096u32.to_be_bytes()); + // Frame header: page number = 1, commit_size = 1 + wal_data[32..36].copy_from_slice(&1u32.to_be_bytes()); + wal_data[36..40].copy_from_slice(&1u32.to_be_bytes()); + // Fill frame body with garbage + for b in &mut wal_data[56..] { + *b = 0xAB; + } + std::fs::write(&wal, &wal_data).unwrap(); + + let km = KeyMaterial::EncKey { + key: [0u8; 32], + salt: [0u8; 16], + }; + let params = wx_decrypt::MACOS_4_1_7_31; + + let result = apply_wal_patch(&wal, &dst, &km, ¶ms); + assert!( + matches!( + result, + WalPatchResult::ContentFailed(_) | WalPatchResult::IoFailed(_) + ), + "expected a failure variant, got {result:?}" + ); + + // dst must be untouched + assert_eq!(std::fs::read(&dst).unwrap(), original); + + // tmp must be cleaned up + let tmp_path = dst.with_extension("db.tmp"); + assert!(!tmp_path.exists(), "tmp file should be removed on failure"); + } + + // --------------------------------------------------------------- + // classify_decrypt_error coverage + // --------------------------------------------------------------- + #[test] + fn classify_content_errors() { + let cases = vec![ + DecryptError::HmacVerificationFailed { page_num: 1 }, + DecryptError::AesDecryptFailed { + page_num: 1, + reason: "test".into(), + }, + DecryptError::InvalidWalHeader { + reason: "bad".into(), + }, + DecryptError::IncorrectKey, + ]; + for e in cases { + let r = classify_decrypt_error(e); + assert!( + matches!(r, WalPatchResult::ContentFailed(_)), + "expected ContentFailed, got {r:?}" + ); + } + } + + #[test] + fn classify_io_errors() { + let cases = vec![ + DecryptError::SaltMismatch, + DecryptError::NoMatchingEncKey, + DecryptError::FileTooSmall { + expected: 100, + actual: 10, + }, + DecryptError::AlreadyDecrypted, + ]; + for e in cases { + let r = classify_decrypt_error(e); + assert!( + matches!(r, WalPatchResult::IoFailed(_)), + "expected IoFailed, got {r:?}" + ); + } + } + + // --------------------------------------------------------------- + // IO failure: copy fails when dst does not exist + // --------------------------------------------------------------- + #[test] + fn copy_failure_returns_io_failed() { + let tmp = TempDir::new().unwrap(); + let dst = tmp.path().join("nonexistent.db"); + let wal = tmp.path().join("test.db-wal"); + std::fs::write(&wal, b"dummy").unwrap(); + + let km = KeyMaterial::EncKey { + key: [0u8; 32], + salt: [0u8; 16], + }; + let params = wx_decrypt::MACOS_4_1_7_31; + + let result = apply_wal_patch(&wal, &dst, &km, ¶ms); + assert!( + matches!(result, WalPatchResult::IoFailed(_)), + "expected IoFailed when dst missing, got {result:?}" + ); + } + + // --------------------------------------------------------------- + // Residual tmp cleanup on entry + // --------------------------------------------------------------- + #[test] + fn residual_tmp_is_cleaned_on_entry() { + let tmp = TempDir::new().unwrap(); + let dst = tmp.path().join("test.db"); + create_sqlite_db(&dst); + + let tmp_path = dst.with_extension("db.tmp"); + std::fs::write(&tmp_path, b"leftover from crash").unwrap(); + + // Create empty WAL (will produce NoFrames) + let wal = tmp.path().join("test.db-wal"); + let mut wal_header = vec![0u8; 32]; + wal_header[..4].copy_from_slice(&0x377f0682u32.to_be_bytes()); + wal_header[4..8].copy_from_slice(&3007000u32.to_be_bytes()); + wal_header[8..12].copy_from_slice(&4096u32.to_be_bytes()); + std::fs::write(&wal, &wal_header).unwrap(); + + let km = KeyMaterial::EncKey { + key: [0u8; 32], + salt: [0u8; 16], + }; + let params = wx_decrypt::MACOS_4_1_7_31; + + let result = apply_wal_patch(&wal, &dst, &km, ¶ms); + assert!(matches!(result, WalPatchResult::NoFrames)); + assert!(!tmp_path.exists(), "tmp should be cleaned up"); + } +} diff --git a/crates/wx-context/tests/contact_resolver_pagination.rs b/crates/wx-context/tests/contact_resolver_pagination.rs new file mode 100644 index 0000000..3fa91eb --- /dev/null +++ b/crates/wx-context/tests/contact_resolver_pagination.rs @@ -0,0 +1,92 @@ +use std::fs; + +use rusqlite::{params, Connection}; +use tempfile::TempDir; +use wx_context::ContactResolver; +use wx_db::WechatDb; + +/// Create a fixture directory with 10,005 contacts to test pagination beyond the 10K boundary. +fn create_fixture_with_many_contacts(count: usize) -> TempDir { + let dir = TempDir::new().unwrap(); + let base = dir.path(); + + // contact/contact.db + let contact_dir = base.join("contact"); + fs::create_dir_all(&contact_dir).unwrap(); + { + let conn = Connection::open(contact_dir.join("contact.db")).unwrap(); + conn.execute_batch( + "CREATE TABLE contact ( + username TEXT PRIMARY KEY, + alias TEXT DEFAULT '', + remark TEXT DEFAULT '', + nick_name TEXT DEFAULT '', + description TEXT DEFAULT NULL, + extra_buffer BLOB DEFAULT NULL + ); + CREATE TABLE contact_label ( + label_id_ TEXT, + label_name_ TEXT, + sort_order_ INTEGER + );", + ) + .unwrap(); + + // Insert contacts in a transaction for speed + conn.execute_batch("BEGIN").unwrap(); + for i in 0..count { + conn.execute( + "INSERT INTO contact (username, nick_name) VALUES (?1, ?2)", + params![format!("wxid_{i:06}"), format!("User {i}")], + ) + .unwrap(); + } + conn.execute_batch("COMMIT").unwrap(); + } + + // session/session.db (required by WechatDb::open) + let session_dir = base.join("session"); + fs::create_dir_all(&session_dir).unwrap(); + { + let conn = Connection::open(session_dir.join("session.db")).unwrap(); + conn.execute_batch( + "CREATE TABLE SessionTable ( + username TEXT, + sort_timestamp INTEGER, + summary TEXT, + last_msg_type INTEGER, + last_msg_sender TEXT, + last_sender_display_name TEXT + );", + ) + .unwrap(); + } + + // message/ (empty dir — no shards needed for contact resolution) + let msg_dir = base.join("message"); + fs::create_dir_all(&msg_dir).unwrap(); + + dir +} + +#[test] +fn contact_resolver_loads_all_contacts_beyond_10k() { + let dir = create_fixture_with_many_contacts(10_005); + let db = WechatDb::open(dir.path()).unwrap(); + let resolver = ContactResolver::build(&db).unwrap(); + + // Contacts sorted by username: wxid_000000 .. wxid_010004 + // The last 5 (wxid_010000 .. wxid_010004) would be beyond the 10K cutoff without pagination + for i in 10_000..10_005 { + let wxid = format!("wxid_{i:06}"); + let expected = format!("User {i}"); + assert_eq!( + resolver.display_name(&wxid), + expected, + "contact {wxid} should be resolvable (beyond 10K boundary)" + ); + } + + // Also verify first contact is present + assert_eq!(resolver.display_name("wxid_000000"), "User 0"); +} diff --git a/crates/wx-context/tests/selective_decrypt.rs b/crates/wx-context/tests/selective_decrypt.rs new file mode 100644 index 0000000..2f742ff --- /dev/null +++ b/crates/wx-context/tests/selective_decrypt.rs @@ -0,0 +1,163 @@ +use std::path::{Path, PathBuf}; +use tempfile::TempDir; +use wx_context::{DecryptRequest, PersistentCache}; + +fn build_encrypted_db( + path: &Path, + raw_key: &[u8; 32], + salt: &[u8; 16], + params: &wx_decrypt::CryptoParams, +) { + use aes::cipher::{BlockEncryptMut, KeyIvInit}; + use hmac::{Hmac, Mac}; + use sha2::Sha512; + + let enc_key = wx_decrypt::kdf::derive_enc_key(raw_key, salt, params); + let mac_key = wx_decrypt::kdf::derive_mac_key(&enc_key, salt, params); + + let iv = [0x42u8; 16]; + let data_size = params.page_size - params.reserve - params.salt_size; + let plaintext = vec![0u8; data_size]; + + let mut ciphertext = plaintext; + type Aes256CbcEnc = cbc::Encryptor; + Aes256CbcEnc::new((&enc_key).into(), (&iv).into()) + .encrypt_padded_mut::(&mut ciphertext, data_size) + .unwrap(); + + let mut page = Vec::with_capacity(params.page_size); + page.extend_from_slice(salt); + page.extend_from_slice(&ciphertext); + page.extend_from_slice(&iv); + page.resize(params.page_size, 0); + + let hmac_data_end = params.page_size - params.reserve + params.iv_size; + let mut mac = as Mac>::new_from_slice(&mac_key).unwrap(); + mac.update(&page[params.salt_size..hmac_data_end]); + mac.update(&1u32.to_le_bytes()); + let hmac_result = mac.finalize().into_bytes(); + let hmac_start = params.page_size - params.reserve + params.iv_size; + page[hmac_start..hmac_start + params.hmac_size] + .copy_from_slice(&hmac_result[..params.hmac_size]); + + if let Some(parent) = path.parent() { + std::fs::create_dir_all(parent).unwrap(); + } + std::fs::write(path, &page).unwrap(); +} + +fn setup_encrypted_dbs(enc_root: &Path, raw_key: &[u8; 32], salt: &[u8; 16]) { + let params = &wx_decrypt::MACOS_4_1_7_31; + let contact_dir = enc_root.join("contact"); + let session_dir = enc_root.join("session"); + let msg_dir = enc_root.join("message"); + std::fs::create_dir_all(&contact_dir).unwrap(); + std::fs::create_dir_all(&session_dir).unwrap(); + std::fs::create_dir_all(&msg_dir).unwrap(); + + build_encrypted_db(&contact_dir.join("contact.db"), raw_key, salt, params); + build_encrypted_db(&session_dir.join("session.db"), raw_key, salt, params); + build_encrypted_db(&msg_dir.join("message_0.db"), raw_key, salt, params); + build_encrypted_db(&msg_dir.join("message_1.db"), raw_key, salt, params); + build_encrypted_db(&msg_dir.join("message_2.db"), raw_key, salt, params); +} + +fn make_cache(enc_root: PathBuf, cache_root: PathBuf, raw_key: [u8; 32]) -> PersistentCache { + PersistentCache::new_for_test( + cache_root, + enc_root, + Some(raw_key), + &wx_decrypt::MACOS_4_1_7_31, + ) +} + +// Slow by design: builds encrypted fixtures and runs the real PBKDF2-backed +// decrypt path. Skip by default for routine cargo test runs. +#[test] +#[ignore = "slow decrypt test; runs real PBKDF2/decrypt path"] +fn contacts_only_decrypts_core_dbs() { + let tmp = TempDir::new().unwrap(); + let enc_root = tmp.path().join("encrypted"); + let cache_root = tmp.path().join("cache"); + let raw_key = [0xABu8; 32]; + let salt = [0x01u8; 16]; + + setup_encrypted_dbs(&enc_root, &raw_key, &salt); + let cache = make_cache(enc_root, cache_root.clone(), raw_key); + + let stats = DecryptRequest::new().core().execute(&cache).unwrap(); + + assert_eq!(stats.decrypted, 2, "only contact.db and session.db"); + assert_eq!(stats.errors, 0); + + assert!(cache_root.join("contact/contact.db").exists()); + assert!(cache_root.join("session/session.db").exists()); + assert!(!cache_root.join("message/message_0.db").exists()); + assert!(!cache_root.join("message/message_1.db").exists()); + assert!(!cache_root.join("message/message_2.db").exists()); +} + +// Slow by design: builds encrypted fixtures and runs the real PBKDF2-backed +// decrypt path. Skip by default for routine cargo test runs. +#[test] +#[ignore = "slow decrypt test; runs real PBKDF2/decrypt path"] +fn core_plus_selected_shards_decrypts_only_requested_dbs() { + let tmp = TempDir::new().unwrap(); + let enc_root = tmp.path().join("encrypted"); + let cache_root = tmp.path().join("cache"); + let raw_key = [0xABu8; 32]; + let salt = [0x01u8; 16]; + + setup_encrypted_dbs(&enc_root, &raw_key, &salt); + let cache = make_cache(enc_root, cache_root.clone(), raw_key); + + let stats = DecryptRequest::new() + .core() + .shards(&[0, 2]) + .execute(&cache) + .unwrap(); + + assert_eq!( + stats.decrypted, 4, + "contact + session + message_0 + message_2" + ); + assert_eq!(stats.errors, 0); + + assert!(cache_root.join("contact/contact.db").exists()); + assert!(cache_root.join("session/session.db").exists()); + assert!(cache_root.join("message/message_0.db").exists()); + assert!( + !cache_root.join("message/message_1.db").exists(), + "shard 1 NOT decrypted" + ); + assert!(cache_root.join("message/message_2.db").exists()); +} + +// Slow by design: performs multiple decrypt passes to verify incremental +// caching behavior. Skip by default for routine cargo test runs. +#[test] +#[ignore = "slow decrypt test; runs real PBKDF2/decrypt path"] +fn incremental_decrypt_skips_cached() { + let tmp = TempDir::new().unwrap(); + let enc_root = tmp.path().join("encrypted"); + let cache_root = tmp.path().join("cache"); + let raw_key = [0xABu8; 32]; + let salt = [0x01u8; 16]; + + setup_encrypted_dbs(&enc_root, &raw_key, &salt); + let cache = make_cache(enc_root, cache_root.clone(), raw_key); + + // First: decrypt core only + let stats1 = DecryptRequest::new().core().execute(&cache).unwrap(); + assert_eq!(stats1.decrypted, 2); + + // Second: decrypt core + shard 0 → core should be skipped + let stats2 = DecryptRequest::new() + .core() + .shards(&[0]) + .execute(&cache) + .unwrap(); + assert_eq!(stats2.skipped, 2, "contact + session already cached"); + assert_eq!(stats2.decrypted, 1, "only message_0 newly decrypted"); + assert_eq!(stats2.errors, 0); +} diff --git a/crates/wx-db/Cargo.toml b/crates/wx-db/Cargo.toml new file mode 100644 index 0000000..deed595 --- /dev/null +++ b/crates/wx-db/Cargo.toml @@ -0,0 +1,18 @@ +[package] +name = "wx-db" +version.workspace = true +edition.workspace = true + +[dependencies] +rusqlite = { version = "0.32", features = ["bundled-sqlcipher"] } +prost = "0.13" +zstd = "0.13" +md5 = "0.7" +thiserror = "2" +serde = { version = "1", features = ["derive"] } +serde_json = "1" +hex = "0.4" + +[dev-dependencies] +insta = { version = "1", features = ["yaml"] } +tempfile = "3" diff --git a/crates/wx-db/src/chatrooms.rs b/crates/wx-db/src/chatrooms.rs new file mode 100644 index 0000000..412f059 --- /dev/null +++ b/crates/wx-db/src/chatrooms.rs @@ -0,0 +1,92 @@ +use crate::decode::decode_room_data; +use crate::error::DbError; +use crate::model::{effective_limit, ChatRoom, ChatRoomQuery, QueryResult, QueryStats}; +use crate::open::WechatDb; + +impl WechatDb { + /// Query chatrooms, optionally filtered by username. + pub fn query_chatrooms(&self, query: &ChatRoomQuery) -> Result, DbError> { + let limit = effective_limit(query.limit); + + let (sql, params_vec) = if let Some(ref username) = query.username { + ( + "SELECT username, owner, ext_buffer \ + FROM chat_room \ + WHERE username = ?1 \ + ORDER BY username ASC \ + LIMIT ?2 OFFSET ?3" + .to_string(), + vec![ + rusqlite::types::Value::Text(username.clone()), + rusqlite::types::Value::Integer(limit as i64), + rusqlite::types::Value::Integer(query.offset as i64), + ], + ) + } else { + ( + "SELECT username, owner, ext_buffer \ + FROM chat_room \ + ORDER BY username ASC \ + LIMIT ?1 OFFSET ?2" + .to_string(), + vec![ + rusqlite::types::Value::Integer(limit as i64), + rusqlite::types::Value::Integer(query.offset as i64), + ], + ) + }; + + // Count total matching rows before LIMIT/OFFSET + let count_sql = if query.username.is_some() { + "SELECT COUNT(*) FROM chat_room WHERE username = ?1" + } else { + "SELECT COUNT(*) FROM chat_room" + }; + let total_rows: usize = if let Some(ref username) = query.username { + self.contact_conn + .query_row(count_sql, [username.as_str()], |row| row.get::<_, i64>(0))? + as usize + } else { + self.contact_conn + .query_row(count_sql, [], |row| row.get::<_, i64>(0))? as usize + }; + + let mut stmt = self.contact_conn.prepare(&sql)?; + let rows = stmt.query_map(rusqlite::params_from_iter(params_vec.iter()), |row| { + let username: String = row.get(0)?; + let owner: String = row.get::<_, String>(1).unwrap_or_default(); + let ext_buffer: Vec = row.get::<_, Vec>(2).unwrap_or_default(); + Ok((username, owner, ext_buffer)) + })?; + + let mut items = Vec::new(); + + for row_result in rows { + let (username, owner, ext_buffer) = match row_result { + Ok(r) => r, + Err(_) => continue, + }; + + let members = if ext_buffer.is_empty() { + Vec::new() + } else { + decode_room_data(&ext_buffer) + }; + + items.push(ChatRoom { + username, + owner, + members, + }); + } + + Ok(QueryResult { + items, + stats: QueryStats { + total_rows, + filtered_count: None, + skipped: 0, + }, + }) + } +} diff --git a/crates/wx-db/src/contact_proto.rs b/crates/wx-db/src/contact_proto.rs new file mode 100644 index 0000000..667487e --- /dev/null +++ b/crates/wx-db/src/contact_proto.rs @@ -0,0 +1,189 @@ +use prost::Message; + +// --- Protobuf structs for extra_buffer --- + +#[derive(prost::Message)] +pub(crate) struct ContactExtraBufferProto { + #[prost(uint32, optional, tag = "2")] + pub gender: Option, + #[prost(string, optional, tag = "4")] + pub signature: Option, + #[prost(string, optional, tag = "5")] + pub country: Option, + #[prost(string, optional, tag = "6")] + pub province: Option, + #[prost(string, optional, tag = "7")] + pub city: Option, + #[prost(uint32, optional, tag = "8")] + pub source_scene: Option, + #[prost(message, optional, tag = "14")] + pub phone_entry: Option, + #[prost(string, optional, tag = "30")] + pub label_ids_csv: Option, +} + +#[derive(prost::Message)] +pub(crate) struct PhoneEntryProto { + #[prost(uint32, optional, tag = "1")] + pub has_phone: Option, + #[prost(message, optional, tag = "2")] + pub phone_detail: Option, +} + +#[derive(prost::Message)] +pub(crate) struct PhoneDetailProto { + #[prost(string, optional, tag = "1")] + pub number: Option, +} + +// --- Decoded intermediate type --- + +#[derive(Default)] +pub(crate) struct ContactExtra { + pub gender: Option, + pub signature: Option, + pub region: Option, + pub source_scene: Option, + pub phone: Option, + pub label_ids: Vec, +} + +// --- Decode function --- + +pub(crate) fn decode_extra_buffer(blob: &[u8]) -> ContactExtra { + let proto = match ContactExtraBufferProto::decode(blob) { + Ok(p) => p, + Err(_) => return ContactExtra::default(), + }; + + let gender = proto.gender.filter(|&v| v > 0); + + let signature = proto.signature.filter(|s| !s.is_empty()); + + let region = { + let parts: Vec<&str> = [ + proto.country.as_deref(), + proto.province.as_deref(), + proto.city.as_deref(), + ] + .iter() + .filter_map(|p| p.filter(|s| !s.is_empty())) + .collect(); + if parts.is_empty() { + None + } else { + Some(parts.join(" · ")) + } + }; + + let source_scene = proto.source_scene.filter(|&v| v > 0); + + let phone = proto + .phone_entry + .filter(|pe| pe.has_phone == Some(1)) + .and_then(|pe| pe.phone_detail) + .and_then(|pd| pd.number) + .filter(|n| !n.is_empty()); + + let label_ids = proto + .label_ids_csv + .as_deref() + .unwrap_or("") + .split(',') + .map(|s| s.trim().to_string()) + .filter(|s| !s.is_empty()) + .collect(); + + ContactExtra { + gender, + signature, + region, + source_scene, + phone, + label_ids, + } +} + +// --- Test encode helper --- + +#[doc(hidden)] +#[allow(clippy::too_many_arguments)] +pub fn encode_extra_buffer_for_test( + gender: Option, + signature: Option<&str>, + country: Option<&str>, + province: Option<&str>, + city: Option<&str>, + source_scene: Option, + phone: Option<&str>, + label_ids_csv: Option<&str>, +) -> Vec { + let phone_entry = phone.map(|number| PhoneEntryProto { + has_phone: Some(1), + phone_detail: Some(PhoneDetailProto { + number: Some(number.to_string()), + }), + }); + + let proto = ContactExtraBufferProto { + gender, + signature: signature.map(|s| s.to_string()), + country: country.map(|s| s.to_string()), + province: province.map(|s| s.to_string()), + city: city.map(|s| s.to_string()), + source_scene, + phone_entry, + label_ids_csv: label_ids_csv.map(|s| s.to_string()), + }; + proto.encode_to_vec() +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn decode_all_fields() { + let blob = encode_extra_buffer_for_test( + Some(1), + Some("hello world"), + Some("CN"), + Some("Beijing"), + Some("Haidian"), + Some(30), + Some("13800138000"), + Some("5,6"), + ); + let extra = decode_extra_buffer(&blob); + assert_eq!(extra.gender, Some(1)); + assert_eq!(extra.signature.as_deref(), Some("hello world")); + assert_eq!(extra.region.as_deref(), Some("CN · Beijing · Haidian")); + assert_eq!(extra.source_scene, Some(30)); + assert_eq!(extra.phone.as_deref(), Some("13800138000")); + assert_eq!(extra.label_ids, vec!["5", "6"]); + } + + #[test] + fn decode_empty_blob() { + let extra = decode_extra_buffer(&[]); + assert_eq!(extra.gender, None); + assert_eq!(extra.signature, None); + assert_eq!(extra.region, None); + assert_eq!(extra.source_scene, None); + assert_eq!(extra.phone, None); + assert!(extra.label_ids.is_empty()); + } + + #[test] + fn decode_country_only() { + let blob = + encode_extra_buffer_for_test(None, None, Some("CN"), None, None, None, None, None); + let extra = decode_extra_buffer(&blob); + assert_eq!(extra.region.as_deref(), Some("CN")); + assert_eq!(extra.gender, None); + assert_eq!(extra.signature, None); + assert_eq!(extra.source_scene, None); + assert_eq!(extra.phone, None); + assert!(extra.label_ids.is_empty()); + } +} diff --git a/crates/wx-db/src/contacts.rs b/crates/wx-db/src/contacts.rs new file mode 100644 index 0000000..38dcfc4 --- /dev/null +++ b/crates/wx-db/src/contacts.rs @@ -0,0 +1,224 @@ +use std::collections::HashMap; + +use rusqlite::types::Value; + +use crate::contact_proto; +use crate::decode; +use crate::error::DbError; +use crate::model::{effective_limit, Contact, ContactQuery, QueryResult, QueryStats}; +use crate::open::WechatDb; + +impl WechatDb { + /// Query contacts, optionally filtered by keyword. + /// + /// When a keyword is present, filtering is done at the application level + /// (after decoding extra_buffer) so that phone, labels, signature, and + /// region can also be searched. SQL-level LIMIT/OFFSET is only used for + /// the fast path (no keyword). + pub fn query_contacts(&self, query: &ContactQuery) -> Result, DbError> { + let limit = effective_limit(query.limit); + // Ensure label_map is loaded (lazy init with error propagation). + { + let guard = self.label_cache.read().unwrap(); + if guard.is_none() { + drop(guard); + let map = self.load_label_map()?; + self.label_cache.write().unwrap().replace(map); + } + } + let guard = self.label_cache.read().unwrap(); + let label_map = guard.as_ref().unwrap(); + + if query.keyword.is_some() { + self.query_contacts_with_keyword(query, limit, label_map) + } else { + self.query_contacts_fast(query, limit, label_map) + } + } + + /// Fast path: no keyword, use SQL-level LIMIT/OFFSET. + fn query_contacts_fast( + &self, + query: &ContactQuery, + limit: usize, + label_map: &HashMap, + ) -> Result, DbError> { + let total_rows: usize = + self.contact_conn + .query_row("SELECT COUNT(*) FROM contact", [], |row| { + row.get::<_, i64>(0) + })? as usize; + + let mut stmt = self.contact_conn.prepare( + "SELECT username, alias, remark, nick_name, description, extra_buffer \ + FROM contact ORDER BY username ASC LIMIT ?1 OFFSET ?2", + )?; + let rows = stmt.query_map( + [ + Value::Integer(limit as i64), + Value::Integer(query.offset as i64), + ], + |row| self.map_contact_row(row, label_map), + )?; + + let items: Vec = rows.filter_map(|r| r.ok()).collect(); + + Ok(QueryResult { + items, + stats: QueryStats { + total_rows, + filtered_count: None, + skipped: 0, + }, + }) + } + + /// Keyword path: fetch all contacts, decode extra_buffer, filter in Rust. + fn query_contacts_with_keyword( + &self, + query: &ContactQuery, + limit: usize, + label_map: &HashMap, + ) -> Result, DbError> { + let kw_lower = query.keyword.as_ref().unwrap().to_lowercase(); + + let total_rows: usize = + self.contact_conn + .query_row("SELECT COUNT(*) FROM contact", [], |row| { + row.get::<_, i64>(0) + })? as usize; + + let mut stmt = self.contact_conn.prepare( + "SELECT username, alias, remark, nick_name, description, extra_buffer \ + FROM contact ORDER BY username ASC", + )?; + let rows = stmt.query_map([], |row| self.map_contact_row(row, label_map))?; + + let all_contacts: Vec = rows.filter_map(|r| r.ok()).collect(); + + let matched: Vec = all_contacts + .into_iter() + .filter(|c| contact_matches_keyword(c, &kw_lower)) + .collect(); + + let filtered_count = matched.len(); + let items: Vec = matched.into_iter().skip(query.offset).take(limit).collect(); + + Ok(QueryResult { + items, + stats: QueryStats { + total_rows, + filtered_count: Some(filtered_count), + skipped: 0, + }, + }) + } + + /// Map a single row to a Contact, decoding extra_buffer and resolving labels. + fn map_contact_row( + &self, + row: &rusqlite::Row<'_>, + label_map: &HashMap, + ) -> rusqlite::Result { + let user_name: String = row.get(0)?; + let alias: String = row.get::<_, String>(1).unwrap_or_default(); + let remark: String = row.get::<_, String>(2).unwrap_or_default(); + let nick_name: String = row.get::<_, String>(3).unwrap_or_default(); + let memo: Option = row.get::<_, Option>(4).unwrap_or(None); + let extra_buffer: Vec = row.get::<_, Vec>(5).unwrap_or_default(); + + let extra = contact_proto::decode_extra_buffer(&extra_buffer); + + let labels: Vec = extra + .label_ids + .iter() + .filter_map(|id| label_map.get(id).cloned()) + .collect(); + + Ok(Contact { + user_name, + alias, + remark, + nick_name, + memo: memo.filter(|s| !s.is_empty()), + gender: extra.gender, + signature: extra.signature, + region: extra.region, + source_scene: extra.source_scene, + phone: extra.phone, + labels, + }) + } + + /// Load `contact_label` table into a HashMap. + /// Returns empty map if the table does not exist. + fn load_label_map(&self) -> Result, DbError> { + if !decode::table_exists(&self.contact_conn, "contact_label")? { + return Ok(HashMap::new()); + } + + let mut stmt = self + .contact_conn + .prepare("SELECT label_id_, label_name_ FROM contact_label")?; + let rows = stmt.query_map([], |row| { + let id_val: Value = row.get(0)?; + let id = match id_val { + Value::Integer(n) => n.to_string(), + Value::Text(s) => s, + _ => String::new(), + }; + let name: String = row.get::<_, String>(1).unwrap_or_default(); + Ok((id, name)) + })?; + + let mut map = HashMap::new(); + for (id, name) in rows.flatten() { + if !id.is_empty() { + map.insert(id, name); + } + } + Ok(map) + } +} + +/// Check if a contact matches a keyword (case-insensitive) across all fields. +fn contact_matches_keyword(c: &Contact, kw_lower: &str) -> bool { + if c.user_name.to_lowercase().contains(kw_lower) { + return true; + } + if c.alias.to_lowercase().contains(kw_lower) { + return true; + } + if c.remark.to_lowercase().contains(kw_lower) { + return true; + } + if c.nick_name.to_lowercase().contains(kw_lower) { + return true; + } + if let Some(ref memo) = c.memo { + if memo.to_lowercase().contains(kw_lower) { + return true; + } + } + if let Some(ref phone) = c.phone { + if phone.to_lowercase().contains(kw_lower) { + return true; + } + } + if let Some(ref sig) = c.signature { + if sig.to_lowercase().contains(kw_lower) { + return true; + } + } + if let Some(ref region) = c.region { + if region.to_lowercase().contains(kw_lower) { + return true; + } + } + for label in &c.labels { + if label.to_lowercase().contains(kw_lower) { + return true; + } + } + false +} diff --git a/crates/wx-db/src/decode.rs b/crates/wx-db/src/decode.rs new file mode 100644 index 0000000..75ed4d7 --- /dev/null +++ b/crates/wx-db/src/decode.rs @@ -0,0 +1,381 @@ +use rusqlite::Connection; + +use crate::error::DbError; +use crate::model::{split_local_type, ChatRoomMember, Message, MessageContent, PackedInfo}; + +// --- Protobuf structs (hand-written prost, no .proto files) --- + +#[derive(prost::Message)] +pub(crate) struct PackedInfoProto { + #[prost(uint32, tag = "1")] + pub r#type: u32, + #[prost(uint32, tag = "2")] + pub version: u32, + #[prost(message, optional, tag = "3")] + pub image: Option, + #[prost(message, optional, tag = "4")] + pub video: Option, +} + +#[derive(prost::Message)] +pub(crate) struct ImageHashProto { + #[prost(string, tag = "4")] + pub md5: String, +} + +#[derive(prost::Message)] +pub(crate) struct VideoHashProto { + #[prost(string, tag = "8")] + pub md5: String, +} + +#[derive(prost::Message)] +pub(crate) struct RoomDataProto { + #[prost(message, repeated, tag = "1")] + pub users: Vec, +} + +#[derive(prost::Message)] +pub(crate) struct RoomDataUserProto { + #[prost(string, tag = "1")] + pub user_name: String, + #[prost(string, optional, tag = "2")] + pub display_name: Option, +} + +// Zstd magic bytes: 0x28 0xB5 0x2F 0xFD +const ZSTD_MAGIC: [u8; 4] = [0x28, 0xB5, 0x2F, 0xFD]; + +// --- Decode functions --- + +/// Compute the message table name for a given talker (wxid / chatroom id). +/// The table name is `Msg_` followed by the full 32-char MD5 hex digest. +pub(crate) fn msg_table_name(talker: &str) -> String { + let hash = md5::compute(talker.as_bytes()); + format!("Msg_{:x}", hash) +} + +/// Check whether a table exists in the given SQLite connection. +pub(crate) fn table_exists(conn: &Connection, name: &str) -> Result { + let count: i64 = conn.query_row( + "SELECT COUNT(*) FROM sqlite_master WHERE type='table' AND name=?1", + [name], + |row| row.get(0), + )?; + Ok(count > 0) +} + +/// Check whether a specific column exists in a table. +pub(crate) fn check_column_exists( + conn: &Connection, + table: &str, + col: &str, +) -> Result { + let sql = format!("PRAGMA table_info([{}])", table); + let mut stmt = conn.prepare(&sql)?; + let exists = stmt + .query_map([], |row| { + let name: String = row.get(1)?; + Ok(name) + })? + .any(|r| r.is_ok_and(|name| name == col)); + Ok(exists) +} + +/// Decode raw content bytes, optionally using wcdb compression type. +/// +/// - If `wcdb_ct == Some(4)`, treat as zstd compressed data. +/// - Else if raw starts with zstd magic bytes, decompress as zstd. +/// - Otherwise, interpret as UTF-8 (lossy). +pub(crate) fn decode_content(raw: &[u8], wcdb_ct: Option) -> Result { + let is_zstd = wcdb_ct == Some(4) || (raw.len() >= 4 && raw[..4] == ZSTD_MAGIC); + + if is_zstd { + let decompressed = zstd::decode_all(raw).map_err(|e| DbError::Zstd(e.to_string()))?; + Ok(String::from_utf8_lossy(&decompressed).into_owned()) + } else { + Ok(String::from_utf8_lossy(raw).into_owned()) + } +} + +/// Decode a PackedInfo protobuf blob into our model type. +/// Returns `None` on any decode error (swallows errors). +pub(crate) fn decode_packed_info(blob: &[u8]) -> Option { + use prost::Message; + let proto = PackedInfoProto::decode(blob).ok()?; + Some(PackedInfo { + image_md5: proto.image.map(|img| img.md5).filter(|s| !s.is_empty()), + video_md5: proto.video.map(|vid| vid.md5).filter(|s| !s.is_empty()), + }) +} + +/// Decode room data protobuf blob into a list of ChatRoomMembers. +pub(crate) fn decode_room_data(blob: &[u8]) -> Vec { + use prost::Message; + let proto = match RoomDataProto::decode(blob) { + Ok(p) => p, + Err(_) => return Vec::new(), + }; + proto + .users + .into_iter() + .map(|u| ChatRoomMember { + user_name: u.user_name, + display_name: u.display_name.filter(|s| !s.is_empty()), + }) + .collect() +} + +/// Encode a `PackedInfo` protobuf blob for use in test fixtures. +/// +/// This is a test-only helper; it is **not** part of the public API. +#[doc(hidden)] +pub fn encode_packed_info_for_test(image_md5: Option<&str>, video_md5: Option<&str>) -> Vec { + use prost::Message; + let proto = PackedInfoProto { + r#type: 106, + version: 14, + image: image_md5.map(|md5| ImageHashProto { + md5: md5.to_string(), + }), + video: video_md5.map(|md5| VideoHashProto { + md5: md5.to_string(), + }), + }; + proto.encode_to_vec() +} + +/// Encode a room-data protobuf blob for use in test fixtures. +/// +/// This is a test-only helper; it is **not** part of the public API. +#[doc(hidden)] +pub fn encode_room_data_for_test(members: &[(&str, Option<&str>)]) -> Vec { + use prost::Message; + let proto = RoomDataProto { + users: members + .iter() + .map(|(name, display)| RoomDataUserProto { + user_name: name.to_string(), + display_name: display.map(|s| s.to_string()), + }) + .collect(), + }; + proto.encode_to_vec() +} + +/// Decode raw DB columns into a Message for test use. +/// +/// This mirrors the logic in `decode_message_row()` (same decode steps in the +/// same order) but takes explicit arguments instead of a `rusqlite::Row`. +#[doc(hidden)] +#[allow(clippy::too_many_arguments)] +pub fn decode_message_for_test( + sort_seq: i64, + server_id: i64, + local_type: i64, + sender: &str, + talker: &str, + create_time: i64, + raw_content: &[u8], + packed_info_data: Option<&[u8]>, + status: i32, + wcdb_ct: Option, + compress_content: Option<&[u8]>, + is_group: bool, +) -> Result { + // Decode content (zstd decompression if needed) + let decoded_text = decode_content(raw_content, wcdb_ct)?; + + // Group sender parsing: extract sender from content prefix + let (sender, content_text) = parse_group_sender(is_group, decoded_text, sender.to_string()); + + // Decode packed info + let packed_info = packed_info_data.and_then(|b| { + if b.is_empty() { + None + } else { + decode_packed_info(b) + } + }); + + // Split local_type into msg_type and sub_type + let (msg_type, sub_type) = split_local_type(local_type); + + // Parse content into typed enum + let content = parse_content( + msg_type, + sub_type, + &content_text, + server_id, + packed_info.as_ref(), + compress_content, + ); + + Ok(Message { + sort_seq, + server_id, + msg_type, + sub_type, + sender, + talker: talker.to_string(), + create_time, + content, + status, + }) +} + +/// Parse group-sender info from decoded text. +/// +/// In group chats, the decoded text often starts with a `"sender:\n"` prefix. +/// This function splits on `":\n"` to extract the sender and the remaining content. +/// +/// Returns `(sender, content)` tuple. If `is_group` is false or no `":\n"` separator +/// is found, returns `(fallback_sender, decoded_text)` unchanged. +pub(crate) fn parse_group_sender( + is_group: bool, + decoded_text: String, + fallback_sender: String, +) -> (String, String) { + if is_group { + if let Some((sender_prefix, rest)) = decoded_text.split_once(":\n") { + (sender_prefix.to_string(), rest.to_string()) + } else { + (fallback_sender, decoded_text) + } + } else { + (fallback_sender, decoded_text) + } +} + +/// Parse message content into a typed MessageContent enum variant. +/// +/// `compress_content` is an optional zstd-compressed blob from the DB's +/// `compress_content` column, used by app messages (especially sub_type=57 +/// quotes) as an alternative content source. +pub(crate) fn parse_content( + msg_type: u32, + sub_type: u32, + content: &str, + _server_id: i64, + packed: Option<&PackedInfo>, + compress_content: Option<&[u8]>, +) -> MessageContent { + use crate::model::*; + match msg_type { + MSG_TYPE_TEXT => MessageContent::Text(content.to_string()), + MSG_TYPE_IMAGE => MessageContent::Image { + md5: packed.and_then(|p| p.image_md5.clone()), + }, + MSG_TYPE_VOICE => MessageContent::Voice, + MSG_TYPE_VIDEO => MessageContent::Video { + md5: packed.and_then(|p| p.video_md5.clone()), + }, + MSG_TYPE_EMOJI => MessageContent::Emoji(content.to_string()), + MSG_TYPE_LOCATION => MessageContent::Location(content.to_string()), + MSG_TYPE_APP => { + // Try compress_content first (zstd-compressed XML), fall back to content + let xml = if let Some(blob) = compress_content { + decode_content(blob, None).unwrap_or_else(|_| content.to_string()) + } else { + content.to_string() + }; + crate::xml_extract::dispatch_app_message(sub_type, &xml) + } + MSG_TYPE_SYSTEM => { + // Try to extract readable text from sysmsg XML (e.g. revokemsg); + // fall back to raw content for plain-text system messages. + let text = crate::xml_extract::extract_system_message_text(content) + .unwrap_or_else(|| content.to_string()); + MessageContent::System(text) + } + MSG_TYPE_REVOKE => MessageContent::Revoke(content.to_string()), + _ => MessageContent::Unknown { + msg_type, + raw: content.to_string(), + }, + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn non_group_returns_fallback_unchanged() { + let (sender, content) = parse_group_sender( + false, + "hello world".to_string(), + "fallback".to_string(), + ); + assert_eq!(sender, "fallback"); + assert_eq!(content, "hello world"); + } + + #[test] + fn group_with_sender_prefix() { + let (sender, content) = parse_group_sender( + true, + "wxid_abc:\nHello group".to_string(), + "fallback".to_string(), + ); + assert_eq!(sender, "wxid_abc"); + assert_eq!(content, "Hello group"); + } + + #[test] + fn group_without_colon_newline() { + let (sender, content) = parse_group_sender( + true, + "no colon newline here".to_string(), + "fallback".to_string(), + ); + assert_eq!(sender, "fallback"); + assert_eq!(content, "no colon newline here"); + } + + #[test] + fn group_empty_content_after_separator() { + let (sender, content) = parse_group_sender( + true, + "wxid_abc:\n".to_string(), + "fallback".to_string(), + ); + assert_eq!(sender, "wxid_abc"); + assert_eq!(content, ""); + } + + #[test] + fn group_only_colon_before_newline() { + let (sender, content) = parse_group_sender( + true, + ":\nsome content".to_string(), + "fallback".to_string(), + ); + assert_eq!(sender, ""); + assert_eq!(content, "some content"); + } + + #[test] + fn group_colon_but_no_newline() { + // ":\n" is the separator; colon without newline should not split + let (sender, content) = parse_group_sender( + true, + "wxid_abc: no newline".to_string(), + "fallback".to_string(), + ); + assert_eq!(sender, "fallback"); + assert_eq!(content, "wxid_abc: no newline"); + } + + #[test] + fn non_group_ignores_sender_prefix() { + // Even if text has ":\n", non-group should not split + let (sender, content) = parse_group_sender( + false, + "wxid_abc:\nHello".to_string(), + "fallback".to_string(), + ); + assert_eq!(sender, "fallback"); + assert_eq!(content, "wxid_abc:\nHello"); + } +} diff --git a/crates/wx-db/src/error.rs b/crates/wx-db/src/error.rs new file mode 100644 index 0000000..52552cd --- /dev/null +++ b/crates/wx-db/src/error.rs @@ -0,0 +1,35 @@ +use serde::{Deserialize, Serialize}; +use thiserror::Error; + +/// A non-fatal warning about a shard that was skipped during query. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ShardWarning { + pub path: String, + pub reason: String, +} + +/// Errors that can occur when opening or querying a WeChat database. +#[derive(Debug, Error)] +pub enum DbError { + /// A required file or directory was not found at the given path. + #[error("path not found: {0}")] + NotFound(String), + /// No message shard database files (`message_N.db`) were found. + #[error("no message shards found")] + NoShards, + /// An error from the underlying SQLite layer. + #[error("sqlite error: {0}")] + Sqlite(#[from] rusqlite::Error), + /// An error during zstd decompression of message content. + #[error("zstd error: {0}")] + Zstd(String), + /// An error applying the encryption key (sqlite3_key failed or wrong key). + #[error("encryption key error: {0}")] + EncryptionKey(String), + /// An error during FTS tokenizer initialization. + #[error("fts init error: {0}")] + FtsInit(String), + /// An I/O error (e.g. reading the message shard directory). + #[error("io error: {0}")] + Io(#[from] std::io::Error), +} diff --git a/crates/wx-db/src/fts.rs b/crates/wx-db/src/fts.rs new file mode 100644 index 0000000..124b982 --- /dev/null +++ b/crates/wx-db/src/fts.rs @@ -0,0 +1,1244 @@ +use std::collections::HashMap; +use std::path::{Path, PathBuf}; +use std::time::Instant; + +use rusqlite::types::ValueRef; +use rusqlite::Connection; +use serde::Serialize; + +use crate::decode::{check_column_exists, decode_content, msg_table_name, parse_content, parse_group_sender}; +use crate::error::DbError; +use crate::model::{split_local_type, MessageContent}; +use crate::open::WechatDb; + +// --------------------------------------------------------------------------- +// CJK tokenization +// --------------------------------------------------------------------------- + +/// Returns true if the character is in a CJK range that should be space-split. +fn is_cjk(c: char) -> bool { + matches!(c, + '\u{3000}'..='\u{303F}' // CJK punctuation + | '\u{3400}'..='\u{9FFF}' // CJK Unified + Ext A + | '\u{F900}'..='\u{FAFF}' // CJK Compatibility Ideographs + | '\u{FF00}'..='\u{FFEF}' // Fullwidth Forms + | '\u{20000}'..='\u{2FA1F}' // CJK Ext B-F + Compat Supplement + ) +} + +/// Space-split CJK characters while leaving non-CJK text intact. +/// +/// Used both at index-build time (INSERT) and query time (MATCH) to ensure +/// symmetric tokenization. +/// +/// ```text +/// "你好world 测试" → "你 好 world 测 试" +/// ``` +pub(crate) fn cjk_tokenize(input: &str) -> String { + let mut out = String::with_capacity(input.len() * 2); + for c in input.chars() { + if is_cjk(c) { + if !out.is_empty() && !out.ends_with(' ') { + out.push(' '); + } + out.push(c); + out.push(' '); + } else { + out.push(c); + } + } + // Trim trailing space added by CJK characters + let trimmed = out.trim_end(); + trimmed.to_string() +} + +/// Convert a user keyword into an FTS5 MATCH expression. +/// +/// - Split by whitespace into tokens +/// - Each token is `cjk_tokenize`-d, then wrapped in double quotes (phrase query) +/// - Internal double quotes are escaped as `""` +/// - Multiple tokens joined with AND +/// +/// ```text +/// "你好 world" → "你 好" AND "world" +/// "测试" → "测 试" +/// 'he said "hi"' → "he" AND "said" AND """hi""" +/// ``` +pub(crate) fn build_fts_query(keyword: &str) -> String { + let tokens: Vec<&str> = keyword.split_whitespace().collect(); + if tokens.is_empty() { + return String::new(); + } + + let phrases: Vec = tokens + .into_iter() + .map(|tok| { + let tokenized = cjk_tokenize(tok); + let escaped = tokenized.replace('"', "\"\""); + format!("\"{}\"", escaped) + }) + .collect(); + + phrases.join(" AND ") +} + +// --------------------------------------------------------------------------- +// Text extraction +// --------------------------------------------------------------------------- + +/// Extract indexable text from a `MessageContent` variant. +/// +/// Returns `None` for media types (Image, Voice, Video) that have no text. +pub(crate) fn extract_fts_text(content: &MessageContent) -> Option { + match content { + MessageContent::Text(s) => Some(s.clone()), + MessageContent::Emoji(s) => Some(s.clone()), + MessageContent::Location(s) => Some(s.clone()), + MessageContent::System(_) => None, + MessageContent::Revoke(s) => Some(s.clone()), + MessageContent::Link { + title, + des, + raw_xml, + .. + } => { + let text = extract_app_fields(raw_xml); + if text.is_empty() { + let parts: Vec<&str> = [title.as_deref(), des.as_deref()] + .iter() + .filter_map(|x| *x) + .collect(); + let text = parts.join(" "); + if text.is_empty() { + None + } else { + Some(text) + } + } else { + Some(text) + } + } + MessageContent::File { title, raw_xml, .. } => { + let text = title.as_deref().unwrap_or(""); + if text.is_empty() { + Some(extract_app_fields(raw_xml)).filter(|s| !s.is_empty()) + } else { + Some(text.to_string()) + } + } + MessageContent::MiniProgram { title, raw_xml, .. } => { + let text = title.as_deref().unwrap_or(""); + if text.is_empty() { + Some(extract_app_fields(raw_xml)).filter(|s| !s.is_empty()) + } else { + Some(text.to_string()) + } + } + MessageContent::MergedMessages { title, raw_xml, .. } => { + let text = title.as_deref().unwrap_or(""); + if text.is_empty() { + Some(extract_app_fields(raw_xml)).filter(|s| !s.is_empty()) + } else { + Some(text.to_string()) + } + } + MessageContent::Quote { + reply_text, + refer_content, + raw_xml, + .. + } => { + let parts: Vec<&str> = [reply_text.as_deref(), refer_content.as_deref()] + .iter() + .filter_map(|x| *x) + .collect(); + let text = parts.join(" "); + if text.is_empty() { + Some(extract_app_fields(raw_xml)).filter(|s| !s.is_empty()) + } else { + Some(text) + } + } + MessageContent::Transfer { + amount_desc, + pay_memo, + .. + } => { + let parts: Vec<&str> = [amount_desc.as_deref(), pay_memo.as_deref()] + .iter() + .filter_map(|x| *x) + .collect(); + let text = parts.join(" "); + if text.is_empty() { + None + } else { + Some(text) + } + } + MessageContent::RedEnvelope { title, .. } => title.clone(), + MessageContent::ChannelVideo { title, .. } => title.clone(), + MessageContent::Pat { .. } => None, + MessageContent::AppGeneric { + title, + des, + raw_xml, + .. + } => { + let text = extract_app_fields(raw_xml); + if text.is_empty() { + let parts: Vec<&str> = [title.as_deref(), des.as_deref()] + .iter() + .filter_map(|x| *x) + .collect(); + let text = parts.join(" "); + if text.is_empty() { + None + } else { + Some(text) + } + } else { + Some(text) + } + } + MessageContent::Image { .. } => None, + MessageContent::Voice => None, + MessageContent::Video { .. } => None, + MessageContent::Unknown { raw, .. } => Some(raw.clone()), + } +} + +/// Extract `` and `<des>` fields from App message XML. +/// +/// Handles `<![CDATA[...]]>` wrapping. Returns empty string on failure. +pub(crate) fn extract_app_fields(raw_xml: &str) -> String { + let title = extract_xml_field(raw_xml, "title"); + let des = extract_xml_field(raw_xml, "des"); + + match (title, des) { + (Some(t), Some(d)) => format!("{}\n{}", t, d), + (Some(t), None) => t, + (None, Some(d)) => d, + (None, None) => String::new(), + } +} + +/// Extract content between `<tag>` and `</tag>`, stripping CDATA if present. +fn extract_xml_field(xml: &str, tag: &str) -> Option<String> { + let open = format!("<{}>", tag); + let close = format!("</{}>", tag); + + let start = xml.find(&open)?; + let after_open = start + open.len(); + let end = xml[after_open..].find(&close)?; + let inner = &xml[after_open..after_open + end]; + + let content = inner.trim(); + if content.is_empty() { + return None; + } + + // Strip CDATA wrapper if present + let stripped = if content.starts_with("<![CDATA[") && content.ends_with("]]>") { + &content[9..content.len() - 3] + } else { + content + }; + + let stripped = stripped.trim(); + if stripped.is_empty() { + None + } else { + Some(stripped.to_string()) + } +} + +// --------------------------------------------------------------------------- +// Public types +// --------------------------------------------------------------------------- + +/// A single hit from the FTS search index. +#[derive(Debug, Clone, Serialize)] +pub struct FtsHit { + pub server_id: i64, + pub talker: String, + pub sender: String, + pub create_time: i64, + pub sort_seq: i64, + pub msg_type: u32, + pub sub_type: u32, + /// Original (un-tokenized) text, for display. + pub snippet: String, +} + +/// Statistics from an FTS index build operation. +#[derive(Debug, Clone, Serialize)] +pub struct FtsBuildStats { + pub indexed: usize, + pub skipped: usize, + pub duration_secs: f64, + pub was_fresh: bool, +} + +/// FTS search result (does not reuse `QueryResult` to avoid semantic clash +/// with `total_rows` / `scanned`). +#[derive(Debug, Clone)] +pub struct FtsSearchResult { + pub hits: Vec<FtsHit>, + /// Total number of FTS MATCH hits (for pagination). + pub total_hits: usize, +} + +// --------------------------------------------------------------------------- +// Paths + schema +// --------------------------------------------------------------------------- + +const SCHEMA_VERSION: &str = "1"; + +fn index_db_path(decrypted_root: &Path) -> PathBuf { + decrypted_root.join("search_index.db") +} + +fn index_tmp_path(decrypted_root: &Path) -> PathBuf { + decrypted_root.join("search_index.tmp.db") +} + +fn ensure_schema(conn: &Connection) -> Result<(), DbError> { + conn.execute_batch( + "CREATE VIRTUAL TABLE IF NOT EXISTS message_fts USING fts5( + body, + raw_text UNINDEXED, + talker UNINDEXED, + sender UNINDEXED, + create_time UNINDEXED, + sort_seq UNINDEXED, + server_id UNINDEXED, + msg_type UNINDEXED, + sub_type UNINDEXED, + tokenize = 'unicode61' + ); + + CREATE TABLE IF NOT EXISTS meta (key TEXT PRIMARY KEY, value TEXT NOT NULL); + + CREATE TABLE IF NOT EXISTS shard_state ( + shard_path TEXT PRIMARY KEY, + last_mtime TEXT NOT NULL + );", + )?; + Ok(()) +} + +/// Check if the existing index covers all current shards with matching mtimes. +fn is_index_fresh(conn: &Connection, shards: &[crate::open::MessageShard]) -> bool { + for shard in shards { + let current_mtime = match std::fs::metadata(&shard.path).and_then(|m| m.modified()) { + Ok(t) => format!("{:?}", t), + Err(_) => return false, + }; + + let stored: Result<String, _> = conn.query_row( + "SELECT last_mtime FROM shard_state WHERE shard_path = ?1", + [shard.path.to_string_lossy().as_ref()], + |row| row.get(0), + ); + + match stored { + Ok(s) if s == current_mtime => {} + _ => return false, + } + } + true +} + +// --------------------------------------------------------------------------- +// Index build +// --------------------------------------------------------------------------- + +/// Rows between COMMIT + BEGIN cycles. +const COMMIT_INTERVAL: usize = 20_000; + +impl WechatDb { + /// Build (or skip if fresh) the FTS5 search index for all message shards. + /// + /// The index is written atomically: `.tmp.db` → `rename` → final path. + pub fn build_fts_index(&self, decrypted_root: &Path) -> Result<FtsBuildStats, DbError> { + let start = Instant::now(); + let final_path = index_db_path(decrypted_root); + + // Check freshness of existing index + if final_path.exists() { + let conn = Connection::open_with_flags( + &final_path, + rusqlite::OpenFlags::SQLITE_OPEN_READ_ONLY, + )?; + if is_index_fresh(&conn, &self.shards) { + return Ok(FtsBuildStats { + indexed: 0, + skipped: 0, + duration_secs: start.elapsed().as_secs_f64(), + was_fresh: true, + }); + } + } + + // Build session → table_name lookup + let talker_map = self.build_talker_map()?; + + // Create tmp db + let tmp_path = index_tmp_path(decrypted_root); + if tmp_path.exists() { + std::fs::remove_file(&tmp_path)?; + } + + let conn = Connection::open(&tmp_path)?; + conn.execute_batch( + "PRAGMA journal_mode=DELETE; + PRAGMA synchronous=OFF; + PRAGMA temp_store=MEMORY;", + )?; + ensure_schema(&conn)?; + + let mut indexed: usize = 0; + let mut skipped: usize = 0; + let mut rows_in_tx: usize = 0; + + conn.execute_batch("BEGIN")?; + + // Prepare INSERT statements once (reused across all shards + tables) + let mut insert_stmt = conn.prepare( + "INSERT INTO message_fts (body, raw_text, talker, sender, create_time, \ + sort_seq, server_id, msg_type, sub_type) \ + VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + )?; + let mut shard_state_stmt = conn.prepare( + "INSERT OR REPLACE INTO shard_state (shard_path, last_mtime) VALUES (?1, ?2)", + )?; + + for shard in &self.shards { + let shard_conn = WechatDb::open_shard_with_key(shard, self.raw_key.as_ref())?; + + // List Msg_* tables in this shard + let mut table_stmt = shard_conn.prepare( + "SELECT name FROM sqlite_master WHERE type='table' AND name LIKE 'Msg_%'", + )?; + let table_names: Vec<String> = table_stmt + .query_map([], |row| row.get(0))? + .filter_map(|r| r.ok()) + .collect(); + + for table_name in &table_names { + let talker = match talker_map.get(table_name.as_str()) { + Some(t) => t.as_str(), + None => continue, // unknown table, skip + }; + let is_group = crate::model::is_group_chat(talker); + + let has_ct_col = + check_column_exists(&shard_conn, table_name, "WCDB_CT_message_content")?; + + let sql = if has_ct_col { + format!( + "SELECT m.sort_seq, m.server_id, m.local_type, \ + COALESCE(n.user_name, ''), m.create_time, \ + m.message_content, m.packed_info_data, \ + m.WCDB_CT_message_content \ + FROM [{table}] m \ + LEFT JOIN Name2Id n ON m.real_sender_id = n.rowid", + table = table_name, + ) + } else { + format!( + "SELECT m.sort_seq, m.server_id, m.local_type, \ + COALESCE(n.user_name, ''), m.create_time, \ + m.message_content, m.packed_info_data \ + FROM [{table}] m \ + LEFT JOIN Name2Id n ON m.real_sender_id = n.rowid", + table = table_name, + ) + }; + + let mut stmt = shard_conn.prepare(&sql)?; + let mut rows = stmt.query([])?; + + while let Some(row) = rows.next()? { + match Self::index_one_row(&mut insert_stmt, row, has_ct_col, is_group, talker) { + Ok(true) => { + indexed += 1; + rows_in_tx += 1; + } + Ok(false) => { /* no text to index (image, voice, etc.) */ } + Err(_) => { + skipped += 1; + } + } + + if rows_in_tx >= COMMIT_INTERVAL { + conn.execute_batch("COMMIT; BEGIN")?; + rows_in_tx = 0; + } + } + } + + // Record shard mtime + let mtime = std::fs::metadata(&shard.path) + .and_then(|m| m.modified()) + .map(|t| format!("{:?}", t)) + .unwrap_or_default(); + shard_state_stmt.execute(rusqlite::params![ + shard.path.to_string_lossy().as_ref(), + mtime, + ])?; + } + + // Drop prepared statements to release borrow on conn before meta writes + drop(insert_stmt); + drop(shard_state_stmt); + + // Write meta + let now = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .map(|d| d.as_secs().to_string()) + .unwrap_or_default(); + conn.execute( + "INSERT OR REPLACE INTO meta (key, value) VALUES ('built_at', ?1)", + [&now], + )?; + conn.execute( + "INSERT OR REPLACE INTO meta (key, value) VALUES ('message_count', ?1)", + [indexed.to_string()], + )?; + conn.execute( + "INSERT OR REPLACE INTO meta (key, value) VALUES ('schema_version', ?1)", + [SCHEMA_VERSION], + )?; + + conn.execute_batch("COMMIT")?; + drop(conn); + + // Atomic rename + std::fs::rename(&tmp_path, &final_path)?; + + Ok(FtsBuildStats { + indexed, + skipped, + duration_secs: start.elapsed().as_secs_f64(), + was_fresh: false, + }) + } + + /// Build a HashMap<table_name, talker_wxid> from session.db. + fn build_talker_map(&self) -> Result<HashMap<String, String>, DbError> { + let mut stmt = self + .session_conn + .prepare("SELECT username FROM SessionTable")?; + let mut rows = stmt.query([])?; + let mut map = HashMap::new(); + while let Some(row) = rows.next()? { + let username: String = row.get(0)?; + let table = msg_table_name(&username); + map.insert(table, username); + } + Ok(map) + } + + /// Decode one message row and insert into the FTS index via a prepared statement. + /// Returns Ok(true) if a row was inserted, Ok(false) if skipped (no text). + fn index_one_row( + insert_stmt: &mut rusqlite::Statement<'_>, + row: &rusqlite::Row<'_>, + has_ct_col: bool, + is_group: bool, + talker: &str, + ) -> Result<bool, DbError> { + let sort_seq: i64 = row.get(0)?; + let server_id: i64 = row.get(1)?; + let local_type: u32 = row.get(2)?; + let sender_from_name2id: String = row.get(3)?; + let create_time: i64 = row.get(4)?; + + let raw_content: Vec<u8> = match row.get_ref(5)? { + ValueRef::Blob(b) => b.to_vec(), + ValueRef::Text(b) => b.to_vec(), + ValueRef::Null => Vec::new(), + _ => Vec::new(), + }; + + // packed_info_data — not needed for FTS text extraction, skip decoding + // WCDB_CT column + let wcdb_ct: Option<i32> = if has_ct_col { + row.get::<_, Option<i32>>(7)? + } else { + None + }; + + let decoded_text = decode_content(&raw_content, wcdb_ct)?; + + // Group sender parsing + let (sender, content_text) = parse_group_sender(is_group, decoded_text, sender_from_name2id); + + let (msg_type, sub_type) = split_local_type(local_type as i64); + + // Parse into MessageContent (packed_info not needed for text extraction) + let content = parse_content(msg_type, sub_type, &content_text, server_id, None, None); + + // Extract indexable text + let raw_text = match extract_fts_text(&content) { + Some(t) => t, + None => return Ok(false), + }; + + let tokenized_body = cjk_tokenize(&raw_text); + + insert_stmt.execute(rusqlite::params![ + tokenized_body, + raw_text, + talker, + sender, + create_time, + sort_seq, + server_id, + msg_type, + sub_type, + ])?; + + Ok(true) + } +} + +// --------------------------------------------------------------------------- +// FTS search +// --------------------------------------------------------------------------- + +impl WechatDb { + /// Search the FTS index. Returns `Ok(None)` if the index file does not exist. + pub fn search_fts( + decrypted_root: &Path, + keyword: &str, + limit: usize, + offset: usize, + ) -> Result<Option<FtsSearchResult>, DbError> { + let path = index_db_path(decrypted_root); + if !path.exists() { + return Ok(None); + } + + let fts_query = build_fts_query(keyword); + if fts_query.is_empty() { + return Ok(Some(FtsSearchResult { + hits: Vec::new(), + total_hits: 0, + })); + } + + let conn = Connection::open_with_flags(&path, rusqlite::OpenFlags::SQLITE_OPEN_READ_ONLY)?; + + // Count total matches + let total_hits: usize = conn.query_row( + "SELECT count(*) FROM message_fts WHERE message_fts MATCH ?1", + [&fts_query], + |row| row.get::<_, i64>(0), + )? as usize; + + // Fetch paginated results + let mut stmt = conn.prepare( + "SELECT server_id, talker, sender, create_time, sort_seq, \ + msg_type, sub_type, raw_text \ + FROM message_fts WHERE message_fts MATCH ?1 \ + ORDER BY create_time DESC, sort_seq DESC \ + LIMIT ?2 OFFSET ?3", + )?; + + let hits: Vec<FtsHit> = stmt + .query_map( + rusqlite::params![fts_query, limit as i64, offset as i64], + |row| { + Ok(FtsHit { + server_id: row.get(0)?, + talker: row.get(1)?, + sender: row.get(2)?, + create_time: row.get(3)?, + sort_seq: row.get(4)?, + msg_type: row.get::<_, i64>(5)? as u32, + sub_type: row.get::<_, i64>(6)? as u32, + snippet: row.get(7)?, + }) + }, + )? + .collect::<Result<Vec<_>, _>>()?; + + Ok(Some(FtsSearchResult { hits, total_hits })) + } +} + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +#[cfg(test)] +mod tests { + use super::*; + + // -- cjk_tokenize -- + + #[test] + fn tokenize_pure_chinese() { + assert_eq!(cjk_tokenize("你好世界"), "你 好 世 界"); + } + + #[test] + fn tokenize_pure_english() { + assert_eq!(cjk_tokenize("hello world"), "hello world"); + } + + #[test] + fn tokenize_mixed() { + assert_eq!(cjk_tokenize("你好world 测试"), "你 好 world 测 试"); + } + + #[test] + fn tokenize_empty() { + assert_eq!(cjk_tokenize(""), ""); + } + + #[test] + fn tokenize_cjk_punctuation() { + // U+3001 (、) is in CJK punctuation range + assert_eq!(cjk_tokenize("你、好"), "你 、 好"); + } + + #[test] + fn tokenize_ext_b_character() { + // U+20000 (𠀀) is in CJK Ext B range + assert_eq!(cjk_tokenize("𠀀test"), "𠀀 test"); + } + + #[test] + fn tokenize_fullwidth() { + // U+FF01 (!) is in Fullwidth Forms + assert_eq!(cjk_tokenize("hello!world"), "hello ! world"); + } + + // -- build_fts_query -- + + #[test] + fn query_single_chinese() { + assert_eq!(build_fts_query("测试"), "\"测 试\""); + } + + #[test] + fn query_multi_word() { + assert_eq!(build_fts_query("你好 world"), "\"你 好\" AND \"world\""); + } + + #[test] + fn query_pure_english() { + assert_eq!(build_fts_query("hello"), "\"hello\""); + } + + #[test] + fn query_mixed() { + assert_eq!(build_fts_query("你好world"), "\"你 好 world\""); + } + + #[test] + fn query_with_double_quotes() { + // Input: he said "hi" (3 whitespace-separated tokens) + assert_eq!( + build_fts_query("he said \"hi\""), + "\"he\" AND \"said\" AND \"\"\"hi\"\"\"" + ); + } + + #[test] + fn query_empty() { + assert_eq!(build_fts_query(""), ""); + assert_eq!(build_fts_query(" "), ""); + } + + // -- extract_fts_text -- + + #[test] + fn extract_text_message() { + let content = MessageContent::Text("hello".into()); + assert_eq!(extract_fts_text(&content), Some("hello".into())); + } + + #[test] + fn extract_image_none() { + let content = MessageContent::Image { md5: None }; + assert_eq!(extract_fts_text(&content), None); + } + + #[test] + fn extract_voice_none() { + assert_eq!(extract_fts_text(&MessageContent::Voice), None); + } + + #[test] + fn extract_video_none() { + let content = MessageContent::Video { md5: None }; + assert_eq!(extract_fts_text(&content), None); + } + + #[test] + fn extract_emoji() { + let content = MessageContent::Emoji("<emoji>".into()); + assert_eq!(extract_fts_text(&content), Some("<emoji>".into())); + } + + #[test] + fn extract_location() { + let content = MessageContent::Location("loc".into()); + assert_eq!(extract_fts_text(&content), Some("loc".into())); + } + + #[test] + fn extract_system() { + let content = MessageContent::System("sys".into()); + assert_eq!(extract_fts_text(&content), None); + } + + #[test] + fn extract_revoke() { + let content = MessageContent::Revoke("revoke".into()); + assert_eq!(extract_fts_text(&content), Some("revoke".into())); + } + + #[test] + fn extract_unknown() { + let content = MessageContent::Unknown { + msg_type: 999, + raw: "raw content".into(), + }; + assert_eq!(extract_fts_text(&content), Some("raw content".into())); + } + + #[test] + fn extract_app_with_title_and_des() { + let content = MessageContent::Link { + sub_type: 5, + title: Some("Link Title".into()), + des: Some("Description text".into()), + url: None, + raw_xml: "<msg><appmsg><title><![CDATA[Link Title]]>".into(), + }; + assert_eq!( + extract_fts_text(&content), + Some("Link Title\nDescription text".into()) + ); + } + + #[test] + fn extract_app_title_only() { + let content = MessageContent::Link { + sub_type: 5, + title: Some("Just Title".into()), + des: None, + url: None, + raw_xml: "Just Title".into(), + }; + assert_eq!(extract_fts_text(&content), Some("Just Title".into())); + } + + #[test] + fn extract_app_empty_xml() { + let content = MessageContent::Link { + sub_type: 5, + title: None, + des: None, + url: None, + raw_xml: "".into(), + }; + assert_eq!(extract_fts_text(&content), None); + } + + // -- extract_app_fields -- + + #[test] + fn app_fields_cdata() { + let xml = "<![CDATA[Hello]]>"; + assert_eq!(extract_app_fields(xml), "Hello\nWorld"); + } + + #[test] + fn app_fields_no_cdata() { + let xml = "HelloWorld"; + assert_eq!(extract_app_fields(xml), "Hello\nWorld"); + } + + #[test] + fn app_fields_no_title() { + let xml = "Only Des"; + assert_eq!(extract_app_fields(xml), "Only Des"); + } + + #[test] + fn app_fields_no_des() { + let xml = "Only Title"; + assert_eq!(extract_app_fields(xml), "Only Title"); + } + + #[test] + fn app_fields_empty() { + assert_eq!(extract_app_fields(""), ""); + assert_eq!(extract_app_fields(""), ""); + } + + // ==================================================================== + // Round-trip tests: build_fts_index + search_fts + // ==================================================================== + + use rusqlite::{params, Connection as SqlConn}; + use std::fs; + use tempfile::TempDir; + + /// Create a minimal fixture with session.db, contact.db, and one message shard. + fn create_fts_fixture() -> TempDir { + let dir = TempDir::new().unwrap(); + let base = dir.path(); + + // contact/contact.db + let contact_dir = base.join("contact"); + fs::create_dir_all(&contact_dir).unwrap(); + let conn = SqlConn::open(contact_dir.join("contact.db")).unwrap(); + crate::test_ddl::create_test_contact_table_minimal(&conn); + conn.execute( + "INSERT INTO contact VALUES (?1, ?2, ?3, ?4)", + params!["wxid_alice", "", "", "Alice"], + ) + .unwrap(); + drop(conn); + + // session/session.db — MUST have rows for session lookup + let session_dir = base.join("session"); + fs::create_dir_all(&session_dir).unwrap(); + let conn = SqlConn::open(session_dir.join("session.db")).unwrap(); + crate::test_ddl::create_test_session_table(&conn); + // Insert session for "wxid_alice" + conn.execute( + "INSERT INTO SessionTable VALUES (?1, ?2, ?3)", + params!["wxid_alice", 1700000000_i64, "last msg"], + ) + .unwrap(); + drop(conn); + + // message/message_0.db + let msg_dir = base.join("message"); + fs::create_dir_all(&msg_dir).unwrap(); + + // md5("wxid_alice") = 29a6db07e8bbdb53f5d54cc3c309f3f1 + let alice_table = "Msg_29a6db07e8bbdb53f5d54cc3c309f3f1"; + + let conn = SqlConn::open(msg_dir.join("message_0.db")).unwrap(); + conn.execute_batch("CREATE TABLE Timestamp (timestamp INTEGER);") + .unwrap(); + conn.execute( + "INSERT INTO Timestamp VALUES (?1)", + params![1_700_000_000_i64], + ) + .unwrap(); + + conn.execute_batch("CREATE TABLE Name2Id (rowid INTEGER PRIMARY KEY, user_name TEXT);") + .unwrap(); + conn.execute( + "INSERT INTO Name2Id VALUES (?1, ?2)", + params![1, "wxid_alice"], + ) + .unwrap(); + + conn.execute_batch(&format!( + "CREATE TABLE [{table}] ( + sort_seq INTEGER, + server_id INTEGER, + local_type INTEGER, + real_sender_id INTEGER, + create_time INTEGER, + message_content BLOB, + packed_info_data BLOB, + status INTEGER + );", + table = alice_table, + )) + .unwrap(); + + // Row 1: Chinese text "你好世界" + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1,?2,?3,?4,?5,?6,?7,?8)", + table = alice_table + ), + params![ + 100_i64, + 1001_i64, + 1_u32, + 1_i64, + 1700000001_i64, + "你好世界".as_bytes(), + rusqlite::types::Null, + 0_i32 + ], + ) + .unwrap(); + + // Row 2: English text "hello world" + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1,?2,?3,?4,?5,?6,?7,?8)", + table = alice_table + ), + params![ + 101_i64, + 1002_i64, + 1_u32, + 1_i64, + 1700000002_i64, + "hello world".as_bytes(), + rusqlite::types::Null, + 0_i32 + ], + ) + .unwrap(); + + // Row 3: Image (msg_type=3) — should NOT be indexed + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1,?2,?3,?4,?5,?6,?7,?8)", + table = alice_table + ), + params![ + 102_i64, + 1003_i64, + 3_u32, + 1_i64, + 1700000003_i64, + "img_content".as_bytes(), + rusqlite::types::Null, + 0_i32 + ], + ) + .unwrap(); + + // Row 4: App message with XML + let app_xml = "<![CDATA[Link Title]]>"; + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1,?2,?3,?4,?5,?6,?7,?8)", + table = alice_table + ), + params![ + 103_i64, + 1004_i64, + 49_u32, + 1_i64, + 1700000004_i64, + app_xml.as_bytes(), + rusqlite::types::Null, + 0_i32 + ], + ) + .unwrap(); + + // Row 5: Mixed CJK+English "我在用 iPhone 15" + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1,?2,?3,?4,?5,?6,?7,?8)", + table = alice_table + ), + params![ + 104_i64, + 1005_i64, + 1_u32, + 1_i64, + 1700000005_i64, + "我在用 iPhone 15".as_bytes(), + rusqlite::types::Null, + 0_i32 + ], + ) + .unwrap(); + + // Row 6: Another text message "hello again" for pagination testing + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1,?2,?3,?4,?5,?6,?7,?8)", + table = alice_table + ), + params![ + 105_i64, + 1006_i64, + 1_u32, + 1_i64, + 1700000006_i64, + "hello again".as_bytes(), + rusqlite::types::Null, + 0_i32 + ], + ) + .unwrap(); + + // Row 7: System message (msg_type=10000) with sysmsg XML — should NOT be indexed + let sysmsg_xml = "wxid_alicewxid_bob"; + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1,?2,?3,?4,?5,?6,?7,?8)", + table = alice_table + ), + params![ + 106_i64, + 1007_i64, + 10000_u32, + 1_i64, + 1700000007_i64, + sysmsg_xml.as_bytes(), + rusqlite::types::Null, + 0_i32 + ], + ) + .unwrap(); + + drop(conn); + dir + } + + #[test] + fn build_and_search_chinese() { + let dir = create_fts_fixture(); + let db = crate::open::WechatDb::open(dir.path()).unwrap(); + + let stats = db.build_fts_index(dir.path()).unwrap(); + assert!(!stats.was_fresh); + assert!(stats.indexed >= 5); // text*3 + app + mixed (not image) + assert_eq!(stats.skipped, 0); + + let result = crate::open::WechatDb::search_fts(dir.path(), "你好", 10, 0) + .unwrap() + .expect("index should exist"); + assert!(result.total_hits >= 1); + assert!(result.hits.iter().any(|h| h.snippet.contains("你好世界"))); + } + + #[test] + fn build_and_search_english() { + let dir = create_fts_fixture(); + let db = crate::open::WechatDb::open(dir.path()).unwrap(); + db.build_fts_index(dir.path()).unwrap(); + + let result = crate::open::WechatDb::search_fts(dir.path(), "hello", 10, 0) + .unwrap() + .expect("index should exist"); + assert!(result.total_hits >= 1); + assert!(result + .hits + .iter() + .any(|h| h.snippet.contains("hello world"))); + } + + #[test] + fn search_image_not_indexed() { + let dir = create_fts_fixture(); + let db = crate::open::WechatDb::open(dir.path()).unwrap(); + db.build_fts_index(dir.path()).unwrap(); + + let result = crate::open::WechatDb::search_fts(dir.path(), "img_content", 10, 0) + .unwrap() + .expect("index should exist"); + assert_eq!(result.total_hits, 0); + } + + #[test] + fn search_nonexistent_index() { + let dir = TempDir::new().unwrap(); + let result = crate::open::WechatDb::search_fts(dir.path(), "test", 10, 0).unwrap(); + assert!(result.is_none()); + } + + #[test] + fn search_pagination() { + let dir = create_fts_fixture(); + let db = crate::open::WechatDb::open(dir.path()).unwrap(); + db.build_fts_index(dir.path()).unwrap(); + + // "hello" appears in "hello world" and "hello again" — guaranteed 2 hits + let all = crate::open::WechatDb::search_fts(dir.path(), "hello", 10, 0) + .unwrap() + .unwrap(); + assert_eq!(all.total_hits, 2); + assert_eq!(all.hits.len(), 2); + + // offset=1 should return 1 hit + let page2 = crate::open::WechatDb::search_fts(dir.path(), "hello", 10, 1) + .unwrap() + .unwrap(); + assert_eq!(page2.total_hits, 2); // total unchanged + assert_eq!(page2.hits.len(), 1); + + // limit=1 should return 1 hit + let limited = crate::open::WechatDb::search_fts(dir.path(), "hello", 1, 0) + .unwrap() + .unwrap(); + assert_eq!(limited.total_hits, 2); // total unchanged + assert_eq!(limited.hits.len(), 1); + } + + #[test] + fn search_iphone_in_mixed_text() { + let dir = create_fts_fixture(); + let db = crate::open::WechatDb::open(dir.path()).unwrap(); + db.build_fts_index(dir.path()).unwrap(); + + let result = crate::open::WechatDb::search_fts(dir.path(), "iPhone", 10, 0) + .unwrap() + .expect("index should exist"); + assert!( + result.total_hits >= 1, + "iPhone should match in mixed CJK+English text" + ); + } + + #[test] + fn index_freshness_skip_rebuild() { + let dir = create_fts_fixture(); + let db = crate::open::WechatDb::open(dir.path()).unwrap(); + + let stats1 = db.build_fts_index(dir.path()).unwrap(); + assert!(!stats1.was_fresh); + + let stats2 = db.build_fts_index(dir.path()).unwrap(); + assert!(stats2.was_fresh); + } + + #[test] + fn search_app_message() { + let dir = create_fts_fixture(); + let db = crate::open::WechatDb::open(dir.path()).unwrap(); + db.build_fts_index(dir.path()).unwrap(); + + let result = crate::open::WechatDb::search_fts(dir.path(), "Description", 10, 0) + .unwrap() + .expect("index should exist"); + assert!(result.total_hits >= 1, "Should match App message des field"); + } + + #[test] + fn search_system_not_indexed() { + let dir = create_fts_fixture(); + let db = crate::open::WechatDb::open(dir.path()).unwrap(); + db.build_fts_index(dir.path()).unwrap(); + + // sysmsg XML keywords should not appear in FTS index + let result = crate::open::WechatDb::search_fts(dir.path(), "sysmsg", 10, 0) + .unwrap() + .expect("index should exist"); + assert_eq!( + result.total_hits, 0, + "System messages should not be indexed" + ); + + let result2 = crate::open::WechatDb::search_fts(dir.path(), "pattedusername", 10, 0) + .unwrap() + .expect("index should exist"); + assert_eq!( + result2.total_hits, 0, + "System message XML content should not be indexed" + ); + } +} diff --git a/crates/wx-db/src/lib.rs b/crates/wx-db/src/lib.rs new file mode 100644 index 0000000..068afa5 --- /dev/null +++ b/crates/wx-db/src/lib.rs @@ -0,0 +1,148 @@ +//! Read-only query layer for WeChat macOS databases (decrypted or encrypted). +//! +//! This crate provides [`WechatDb`] — a handle to an opened WeChat database +//! directory — along with typed query structs for messages, contacts, +//! chatrooms, and sessions. +//! +//! Databases can be opened in two modes: +//! - **Decrypted**: via [`WechatDb::open`] on a pre-decrypted directory +//! - **Encrypted (direct)**: via [`WechatDb::open_encrypted`] using a raw key +//! and SQLCipher's `sqlite3_key()` C API +//! +//! # Usage +//! +//! ```ignore +//! use wx_db::{WechatDb, MessageQuery, ContactQuery}; +//! +//! // Open decrypted DB +//! let db = WechatDb::open("/path/to/decrypted/db")?; +//! +//! // Or open encrypted DB directly with raw key +//! let db = WechatDb::open_encrypted("/path/to/db_storage", raw_key)?; +//! +//! let messages = db.query_messages(&MessageQuery::for_talker("wxid_alice"))?; +//! let contacts = db.query_contacts(&ContactQuery::new())?; +//! ``` + +mod chatrooms; +mod contact_proto; +mod contacts; +mod decode; +mod error; +mod fts; +mod messages; +mod model; +pub mod native_fts; +mod open; +mod pool; +mod sessions; +pub mod shard_metadata; +mod xml_extract; + +pub use error::{DbError, ShardWarning}; +pub use fts::{FtsBuildStats, FtsHit, FtsSearchResult}; +pub use model::*; +pub use native_fts::{load_name2id, FtsHitType, NativeFtsHit, NativeFtsResult}; +pub use open::{open_readonly_connection, WechatDb}; +pub use pool::ShardPool; +pub use xml_extract::extract_quote_fromusr; + +// Test-only helpers for building protobuf fixtures. +#[doc(hidden)] +pub use contact_proto::encode_extra_buffer_for_test; +#[doc(hidden)] +pub use decode::decode_message_for_test; +#[doc(hidden)] +pub use decode::encode_packed_info_for_test; +#[doc(hidden)] +pub use decode::encode_room_data_for_test; + +/// Shared test-fixture DDL helpers. +/// +/// These functions create the standard table schemas used across integration +/// tests so that the DDL strings live in one place. All helpers accept an +/// `&rusqlite::Connection` and execute the `CREATE TABLE` statement +/// directly. +#[doc(hidden)] +pub mod test_ddl { + use rusqlite::Connection; + + /// Create the full `contact` table (7 columns) used by most integration tests. + /// + /// Schema: `username, alias, remark, nick_name, description, extra_buffer`. + pub fn create_test_contact_table(conn: &Connection) { + conn.execute_batch( + "CREATE TABLE contact ( + username TEXT PRIMARY KEY, + alias TEXT DEFAULT '', + remark TEXT DEFAULT '', + nick_name TEXT DEFAULT '', + description TEXT DEFAULT NULL, + extra_buffer BLOB DEFAULT NULL + );", + ) + .unwrap(); + } + + /// Create the minimal `contact` table (4 columns) used by simpler fixtures. + /// + /// Schema: `username, alias, remark, nick_name`. + pub fn create_test_contact_table_minimal(conn: &Connection) { + conn.execute_batch( + "CREATE TABLE contact ( + username TEXT PRIMARY KEY, + alias TEXT DEFAULT '', + remark TEXT DEFAULT '', + nick_name TEXT DEFAULT '' + );", + ) + .unwrap(); + } + + /// Create the `SessionTable` used by session and message query tests. + /// + /// Schema: `username, sort_timestamp, summary`. + pub fn create_test_session_table(conn: &Connection) { + conn.execute_batch( + "CREATE TABLE SessionTable ( + username TEXT, + sort_timestamp INTEGER, + summary TEXT + );", + ) + .unwrap(); + } + + /// Create the `SessionTable` with extended columns used by session tests + /// that need `last_msg_type`, `last_msg_sender`, and `last_sender_display_name`. + /// + /// Schema: `username, sort_timestamp, summary, last_msg_type, last_msg_sender, + /// last_sender_display_name`. + pub fn create_test_session_table_extended(conn: &Connection) { + conn.execute_batch( + "CREATE TABLE SessionTable ( + username TEXT, + sort_timestamp INTEGER, + summary TEXT, + last_msg_type INTEGER, + last_msg_sender TEXT, + last_sender_display_name TEXT + );", + ) + .unwrap(); + } + + /// Create the `contact_label` table used by label-aware tests. + /// + /// Schema: `label_id_, label_name_, sort_order_`. + pub fn create_test_contact_label_table(conn: &Connection) { + conn.execute_batch( + "CREATE TABLE contact_label ( + label_id_ TEXT, + label_name_ TEXT, + sort_order_ INTEGER + );", + ) + .unwrap(); + } +} diff --git a/crates/wx-db/src/messages.rs b/crates/wx-db/src/messages.rs new file mode 100644 index 0000000..f652481 --- /dev/null +++ b/crates/wx-db/src/messages.rs @@ -0,0 +1,1181 @@ +use std::collections::HashMap; + +use rusqlite::types::ValueRef; +use rusqlite::Connection; + +use crate::decode::{ + check_column_exists, decode_content, decode_packed_info, msg_table_name, parse_content, + parse_group_sender, table_exists, +}; +use crate::error::{DbError, ShardWarning}; +use crate::model::{ + effective_limit, split_local_type, AnchorMode, Message, MessageQuery, MessageQueryResult, + QueryStats, SortOrder, +}; +use crate::open::{MessageShard, WechatDb}; + +/// Dispatch mode for regular (non-anchor) queries. +enum RegularQueryMode { + /// Full table scan — used when keyword, filtered_count, or anchor is active. + FullScan, + /// SQL LIMIT pushdown — per-shard `ORDER BY sort_seq {order} LIMIT ?`. + LimitPushdown { + /// `offset + effective_limit` — each shard fetches at most this many rows. + sql_limit: usize, + /// Optional msg_type pushed to SQL WHERE clause. + msg_type_filter: Option, + }, +} + +/// Result of preparing a shard for querying: connection + column metadata. +enum ShardConnection<'a> { + Borrowed(&'a Connection), + Owned(Connection), +} + +impl ShardConnection<'_> { + fn as_conn(&self) -> &Connection { + match self { + ShardConnection::Borrowed(conn) => conn, + ShardConnection::Owned(conn) => conn, + } + } +} + +struct PreparedShard<'a> { + conn: ShardConnection<'a>, + select_cols: String, + has_ct_col: bool, + has_compress_col: bool, +} + +/// Try to open a shard and check table/column availability. +/// Returns `None` (with a warning pushed) if the shard cannot be used. +fn prepare_shard_query<'a>( + shard: &MessageShard, + table_name: &str, + warnings: &mut Vec, + pooled_conn: Option<&'a Connection>, + raw_key: Option<&[u8; 32]>, +) -> Option> { + let shard_path = shard.path.display().to_string(); + + let conn = match pooled_conn { + Some(conn) => ShardConnection::Borrowed(conn), + None => match WechatDb::open_shard_with_key(shard, raw_key) { + Ok(c) => ShardConnection::Owned(c), + Err(e) => { + warnings.push(ShardWarning { + path: shard_path, + reason: format!("open failed: {e}"), + }); + return None; + } + }, + }; + let conn_ref = conn.as_conn(); + + match table_exists(conn_ref, table_name) { + Ok(true) => {} + Ok(false) => return None, + Err(e) => { + warnings.push(ShardWarning { + path: shard_path, + reason: format!("table_exists check failed: {e}"), + }); + return None; + } + } + + let has_ct_col = match check_column_exists(conn_ref, table_name, "WCDB_CT_message_content") { + Ok(v) => v, + Err(e) => { + warnings.push(ShardWarning { + path: shard_path, + reason: format!("check WCDB_CT column failed: {e}"), + }); + return None; + } + }; + + let has_compress_col = match check_column_exists(conn_ref, table_name, "compress_content") { + Ok(v) => v, + Err(e) => { + warnings.push(ShardWarning { + path: shard_path, + reason: format!("check compress_content column failed: {e}"), + }); + return None; + } + }; + + let mut select_cols = String::from( + "m.sort_seq, m.server_id, m.local_type, \ + COALESCE(n.user_name, ''), m.create_time, \ + m.message_content, m.packed_info_data, m.status", + ); + if has_ct_col { + select_cols.push_str(", m.WCDB_CT_message_content"); + } + if has_compress_col { + select_cols.push_str(", m.compress_content"); + } + + Some(PreparedShard { + conn, + select_cols, + has_ct_col, + has_compress_col, + }) +} + +impl WechatDb { + /// Query messages for a given talker (contact or chatroom). + /// + /// Pipeline: + /// 1. Compute the Msg table name from talker via MD5 + /// 2. Find shards overlapping the time range + /// 3. Decide query mode: LIMIT pushdown (index-backed) or full scan + /// 4. For each shard: open, check table exists, build SQL, decode rows + /// (individual shard failures are recorded as warnings, not errors) + /// 5. Merge results across shards, sort by (sort_seq, create_time, server_id) + /// 6. Apply post-filters (keyword, msg_type) in full-scan mode + /// 7. Apply offset + limit (Rust is the authoritative paginator) + pub fn query_messages(&self, query: &MessageQuery) -> Result { + let limit = effective_limit(query.limit); + let table_name = msg_table_name(&query.talker); + let is_group = crate::model::is_group_chat(&query.talker); + + let shards = self.shards_for_range(query.start_time, query.end_time); + if self.shards.is_empty() { + return Err(DbError::NoShards); + } + + // Decide query mode + let mode = if query.limit_pushdown_eligible() { + // Safely compute per-shard SQL LIMIT = offset + limit. + // Use saturating_add to avoid usize overflow, then cap at i64::MAX + // so the subsequent `as i64` cast is always non-negative. + let sql_limit = query.offset.saturating_add(limit).min(i64::MAX as usize); + RegularQueryMode::LimitPushdown { + sql_limit, + msg_type_filter: query.msg_type_filter, + } + } else { + RegularQueryMode::FullScan + }; + + let mut all_messages: Vec = Vec::new(); + let mut total_rows: usize = 0; + let mut skipped: usize = 0; + let mut shard_warnings: Vec = Vec::new(); + + for shard in &shards { + let prepared = match prepare_shard_query( + shard, + &table_name, + &mut shard_warnings, + self.pool().and_then(|pool| pool.get(&shard.path)), + self.raw_key.as_ref(), + ) { + Some(p) => p, + None => continue, + }; + let shard_path = shard.path.display().to_string(); + + let (sql, params) = build_regular_shard_sql( + &mode, + &prepared.select_cols, + &table_name, + query.order, + query.start_time, + query.end_time, + ); + + query_shard_sql( + &prepared, + &sql, + ¶ms, + is_group, + &query.talker, + &shard_path, + &mut all_messages, + &mut total_rows, + &mut skipped, + &mut shard_warnings, + ); + } + + // Sort all messages by (sort_seq, create_time, server_id) in requested direction + match query.order { + SortOrder::Asc => { + all_messages.sort_unstable_by_key(|m| (m.sort_seq, m.create_time, m.server_id)) + } + SortOrder::Desc => all_messages.sort_unstable_by(|a, b| { + (b.sort_seq, b.create_time, b.server_id).cmp(&( + a.sort_seq, + a.create_time, + a.server_id, + )) + }), + } + + // Apply post-filters based on mode + match &mode { + RegularQueryMode::FullScan => { + // Keyword filter post-SQL (content is compressed in DB) + if let Some(ref kw) = query.keyword { + let kw_lower = kw.to_lowercase(); + all_messages.retain(|m| content_contains_keyword(&m.content, &kw_lower)); + } + // msg_type filter + if let Some(mt) = query.msg_type_filter { + all_messages.retain(|m| m.msg_type == mt); + } + } + RegularQueryMode::LimitPushdown { .. } => { + // keyword is None (precondition), msg_type already in SQL WHERE + // No post-filters needed + } + } + + // Compute filtered_count before pagination (opt-in, only in FullScan mode) + let filtered_count = if query.with_filtered_count { + Some(all_messages.len()) + } else { + None + }; + + // Apply offset + limit — Rust is the authoritative paginator + let after_offset: Vec = all_messages + .into_iter() + .skip(query.offset) + .take(limit) + .collect(); + + Ok(MessageQueryResult { + items: after_offset, + stats: QueryStats { + total_rows, + filtered_count, + skipped, + }, + shard_warnings, + }) + } + + /// Count total messages for a talker across all shards (lightweight, no content decoding). + /// + /// Uses `SELECT COUNT(*)` per shard, which is fast (index-only scan, no row decoding). + /// Applies the same time-range and msg_type filters as `query_messages` but skips + /// keyword post-filters since those require content decoding. + pub fn count_messages( + &self, + talker: &str, + start_time: i64, + end_time: i64, + msg_type_filter: Option, + ) -> usize { + let table_name = msg_table_name(talker); + let shards = self.shards_for_range(start_time, end_time); + let mut total: usize = 0; + + let sql = if msg_type_filter.is_some() { + format!( + "SELECT COUNT(*) FROM [{table_name}] \ + WHERE create_time >= ?1 AND create_time <= ?2 \ + AND (local_type & 4294967295) = ?3" + ) + } else { + format!( + "SELECT COUNT(*) FROM [{table_name}] \ + WHERE create_time >= ?1 AND create_time <= ?2" + ) + }; + + for shard in &shards { + let count = if let Some(pool) = self.pool() { + if let Some(conn) = pool.get(&shard.path) { + Self::count_shard(&conn, &sql, start_time, end_time, msg_type_filter) + } else { + continue; + } + } else { + match crate::open::open_connection(&shard.path, self.raw_key.as_ref()) { + Ok(conn) => { + Self::count_shard(&conn, &sql, start_time, end_time, msg_type_filter) + } + Err(_) => continue, + } + }; + + total += count; + } + + total + } + + fn count_shard( + conn: &Connection, + sql: &str, + start_time: i64, + end_time: i64, + msg_type_filter: Option, + ) -> usize { + let result = if let Some(mt) = msg_type_filter { + conn.query_row(sql, [start_time, end_time, mt as i64], |row: &rusqlite::Row<'_>| { + row.get::<_, i64>(0) + }) + } else { + conn.query_row(sql, [start_time, end_time], |row: &rusqlite::Row<'_>| { + row.get::<_, i64>(0) + }) + }; + result.unwrap_or(0).max(0) as usize + } + + /// Query messages using an anchor mode (around/after a specific position). + /// + /// Requires `query.anchor` to be `Some`. Uses full-shard scan for correctness. + pub fn query_messages_anchor( + &self, + query: &MessageQuery, + ) -> Result { + let anchor = query + .anchor + .as_ref() + .expect("query_messages_anchor called without anchor mode"); + + if self.shards.is_empty() { + return Err(DbError::NoShards); + } + + let table_name = msg_table_name(&query.talker); + let is_group = crate::model::is_group_chat(&query.talker); + + match anchor { + AnchorMode::AfterSortSeq(seq) => { + self.query_after_sort_seq(query, &table_name, is_group, *seq) + } + AnchorMode::AroundSortSeq(seq) => { + self.query_around_sort_seq(query, &table_name, is_group, *seq) + } + AnchorMode::AroundServerId(id) => { + self.query_around_server_id(query, &table_name, is_group, *id) + } + } + } + + fn query_after_sort_seq( + &self, + query: &MessageQuery, + table_name: &str, + is_group: bool, + seq: i64, + ) -> Result { + let limit = effective_limit(query.limit); + let mut all_messages: Vec = Vec::new(); + let mut total_rows: usize = 0; + let mut skipped: usize = 0; + let mut shard_warnings: Vec = Vec::new(); + + for shard in self.all_shards() { + let prepared = match prepare_shard_query( + shard, + table_name, + &mut shard_warnings, + self.pool().and_then(|pool| pool.get(&shard.path)), + self.raw_key.as_ref(), + ) { + Some(p) => p, + None => continue, + }; + let shard_path = shard.path.display().to_string(); + + let sql = format!( + "SELECT {select_cols} \ + FROM [{table}] m \ + LEFT JOIN Name2Id n ON m.real_sender_id = n.rowid \ + WHERE m.sort_seq > ?1 \ + ORDER BY m.sort_seq ASC, m.create_time ASC, m.server_id ASC", + select_cols = prepared.select_cols, + table = table_name, + ); + + let mut stmt = match prepared.conn.as_conn().prepare(&sql) { + Ok(s) => s, + Err(e) => { + shard_warnings.push(ShardWarning { + path: shard_path, + reason: format!("prepare failed: {e}"), + }); + continue; + } + }; + + let mut rows = match stmt.query([seq]) { + Ok(r) => r, + Err(e) => { + shard_warnings.push(ShardWarning { + path: shard_path, + reason: format!("query failed: {e}"), + }); + continue; + } + }; + + collect_rows( + &mut rows, + &prepared, + is_group, + &query.talker, + &shard_path, + &mut all_messages, + &mut total_rows, + &mut skipped, + &mut shard_warnings, + ); + } + + // Sort ASC by compound key + all_messages.sort_unstable_by_key(|m| (m.sort_seq, m.create_time, m.server_id)); + + // Apply filters then take(limit) + apply_post_filters(&mut all_messages, &query.keyword, query.msg_type_filter); + all_messages.truncate(limit); + + Ok(MessageQueryResult { + items: all_messages, + stats: QueryStats { + total_rows, + filtered_count: None, + skipped, + }, + shard_warnings, + }) + } + + fn query_around_sort_seq( + &self, + query: &MessageQuery, + table_name: &str, + is_group: bool, + seq: i64, + ) -> Result { + let context = query.context; + let mut before: Vec = Vec::new(); + let mut pivot: Vec = Vec::new(); + let mut after: Vec = Vec::new(); + let mut total_rows: usize = 0; + let mut skipped: usize = 0; + let mut shard_warnings: Vec = Vec::new(); + + for shard in self.all_shards() { + let prepared = match prepare_shard_query( + shard, + table_name, + &mut shard_warnings, + self.pool().and_then(|pool| pool.get(&shard.path)), + self.raw_key.as_ref(), + ) { + Some(p) => p, + None => continue, + }; + let shard_path = shard.path.display().to_string(); + + // Before segment + let sql_before = format!( + "SELECT {select_cols} \ + FROM [{table}] m \ + LEFT JOIN Name2Id n ON m.real_sender_id = n.rowid \ + WHERE m.sort_seq < ?1 \ + ORDER BY m.sort_seq DESC, m.create_time DESC, m.server_id DESC \ + LIMIT ?2", + select_cols = prepared.select_cols, + table = table_name, + ); + query_shard_sql( + &prepared, + &sql_before, + &[seq, context as i64], + is_group, + &query.talker, + &shard_path, + &mut before, + &mut total_rows, + &mut skipped, + &mut shard_warnings, + ); + + // Pivot bucket (no LIMIT — return all messages at this sort_seq) + let sql_pivot = format!( + "SELECT {select_cols} \ + FROM [{table}] m \ + LEFT JOIN Name2Id n ON m.real_sender_id = n.rowid \ + WHERE m.sort_seq = ?1 \ + ORDER BY m.create_time ASC, m.server_id ASC", + select_cols = prepared.select_cols, + table = table_name, + ); + query_shard_sql( + &prepared, + &sql_pivot, + &[seq], + is_group, + &query.talker, + &shard_path, + &mut pivot, + &mut total_rows, + &mut skipped, + &mut shard_warnings, + ); + + // After segment + let sql_after = format!( + "SELECT {select_cols} \ + FROM [{table}] m \ + LEFT JOIN Name2Id n ON m.real_sender_id = n.rowid \ + WHERE m.sort_seq > ?1 \ + ORDER BY m.sort_seq ASC, m.create_time ASC, m.server_id ASC \ + LIMIT ?2", + select_cols = prepared.select_cols, + table = table_name, + ); + query_shard_sql( + &prepared, + &sql_after, + &[seq, context as i64], + is_group, + &query.talker, + &shard_path, + &mut after, + &mut total_rows, + &mut skipped, + &mut shard_warnings, + ); + } + + // Sort and truncate each segment across shards + // Before: sort DESC then take(context), then reverse to ASC + before.sort_unstable_by(|a, b| { + (b.sort_seq, b.create_time, b.server_id).cmp(&(a.sort_seq, a.create_time, a.server_id)) + }); + before.truncate(context); + before.reverse(); + + // Pivot: sort ASC by (create_time, server_id) + pivot.sort_unstable_by_key(|m| (m.create_time, m.server_id)); + + // After: sort ASC then take(context) + after.sort_unstable_by_key(|m| (m.sort_seq, m.create_time, m.server_id)); + after.truncate(context); + + // Merge: before + pivot + after + let mut all_messages = before; + all_messages.append(&mut pivot); + all_messages.append(&mut after); + + // Apply post-filters (may reduce count below 2*context + pivot) + apply_post_filters(&mut all_messages, &query.keyword, query.msg_type_filter); + + Ok(MessageQueryResult { + items: all_messages, + stats: QueryStats { + total_rows, + filtered_count: None, + skipped, + }, + shard_warnings, + }) + } + + fn query_around_server_id( + &self, + query: &MessageQuery, + table_name: &str, + is_group: bool, + target_server_id: i64, + ) -> Result { + let context = query.context; + let mut shard_warnings: Vec = Vec::new(); + + // Phase 1: Locate the target message by server_id across all shards + let mut pivot_msg: Option<(i64, i64, i64)> = None; // (sort_seq, create_time, server_id) + for shard in self.all_shards() { + let prepared = match prepare_shard_query( + shard, + table_name, + &mut shard_warnings, + self.pool().and_then(|pool| pool.get(&shard.path)), + self.raw_key.as_ref(), + ) { + Some(p) => p, + None => continue, + }; + + let sql = format!( + "SELECT m.sort_seq, m.create_time, m.server_id \ + FROM [{table}] m \ + WHERE m.server_id = ?1 \ + LIMIT 1", + table = table_name, + ); + let result: Result, _> = prepared + .conn + .as_conn() + .query_row(&sql, [target_server_id], |row| { + Ok(( + row.get::<_, i64>(0)?, + row.get::<_, i64>(1)?, + row.get::<_, i64>(2)?, + )) + }) + .map(Some) + .or_else(|e| { + if matches!(e, rusqlite::Error::QueryReturnedNoRows) { + Ok(None) + } else { + Err(e) + } + }); + match result { + Ok(Some(row)) => { + pivot_msg = Some(row); + break; + } + Ok(None) => continue, + Err(e) => { + shard_warnings.push(ShardWarning { + path: shard.path.display().to_string(), + reason: format!("locate server_id failed: {e}"), + }); + continue; + } + } + } + + let (pivot_seq, pivot_ct, pivot_sid) = match pivot_msg { + Some(t) => t, + None => { + // server_id not found — return empty result + return Ok(MessageQueryResult { + items: vec![], + stats: QueryStats { + total_rows: 0, + filtered_count: None, + skipped: 0, + }, + shard_warnings, + }); + } + }; + + // Phase 2: Query context around the located message + let mut pivot_messages: Vec = Vec::new(); + let mut before: Vec = Vec::new(); + let mut after: Vec = Vec::new(); + let mut total_rows: usize = 0; + let mut skipped: usize = 0; + + for shard in self.all_shards() { + let prepared = match prepare_shard_query( + shard, + table_name, + &mut shard_warnings, + self.pool().and_then(|pool| pool.get(&shard.path)), + self.raw_key.as_ref(), + ) { + Some(p) => p, + None => continue, + }; + let shard_path = shard.path.display().to_string(); + + // Pivot: exact server_id match + let sql_pivot = format!( + "SELECT {select_cols} \ + FROM [{table}] m \ + LEFT JOIN Name2Id n ON m.real_sender_id = n.rowid \ + WHERE m.server_id = ?1", + select_cols = prepared.select_cols, + table = table_name, + ); + query_shard_sql( + &prepared, + &sql_pivot, + &[target_server_id], + is_group, + &query.talker, + &shard_path, + &mut pivot_messages, + &mut total_rows, + &mut skipped, + &mut shard_warnings, + ); + + // Before: messages strictly before the pivot's compound key + let sql_before = format!( + "SELECT {select_cols} \ + FROM [{table}] m \ + LEFT JOIN Name2Id n ON m.real_sender_id = n.rowid \ + WHERE (m.sort_seq < ?1) \ + OR (m.sort_seq = ?1 AND m.create_time < ?2) \ + OR (m.sort_seq = ?1 AND m.create_time = ?2 AND m.server_id < ?3) \ + ORDER BY m.sort_seq DESC, m.create_time DESC, m.server_id DESC \ + LIMIT ?4", + select_cols = prepared.select_cols, + table = table_name, + ); + query_shard_sql( + &prepared, + &sql_before, + &[pivot_seq, pivot_ct, pivot_sid, context as i64], + is_group, + &query.talker, + &shard_path, + &mut before, + &mut total_rows, + &mut skipped, + &mut shard_warnings, + ); + + // After: messages strictly after the pivot's compound key + let sql_after = format!( + "SELECT {select_cols} \ + FROM [{table}] m \ + LEFT JOIN Name2Id n ON m.real_sender_id = n.rowid \ + WHERE (m.sort_seq > ?1) \ + OR (m.sort_seq = ?1 AND m.create_time > ?2) \ + OR (m.sort_seq = ?1 AND m.create_time = ?2 AND m.server_id > ?3) \ + ORDER BY m.sort_seq ASC, m.create_time ASC, m.server_id ASC \ + LIMIT ?4", + select_cols = prepared.select_cols, + table = table_name, + ); + query_shard_sql( + &prepared, + &sql_after, + &[pivot_seq, pivot_ct, pivot_sid, context as i64], + is_group, + &query.talker, + &shard_path, + &mut after, + &mut total_rows, + &mut skipped, + &mut shard_warnings, + ); + } + + // Sort and truncate + before.sort_unstable_by(|a, b| { + (b.sort_seq, b.create_time, b.server_id).cmp(&(a.sort_seq, a.create_time, a.server_id)) + }); + before.truncate(context); + before.reverse(); + + after.sort_unstable_by_key(|m| (m.sort_seq, m.create_time, m.server_id)); + after.truncate(context); + + // Merge: before + pivot + after + let mut all_messages = before; + all_messages.append(&mut pivot_messages); + all_messages.append(&mut after); + + apply_post_filters(&mut all_messages, &query.keyword, query.msg_type_filter); + + Ok(MessageQueryResult { + items: all_messages, + stats: QueryStats { + total_rows, + filtered_count: None, + skipped, + }, + shard_warnings, + }) + } + + /// Bulk-query `MAX(sort_seq)` per talker across all shards. + /// + /// For each shard, opens one connection, discovers which `Msg_*` tables exist, + /// and queries `MAX(sort_seq)` for each. Results are merged across shards + /// (per-username max). Shard/table failures are logged and skipped. + pub fn bulk_max_sort_seq(&self, known_usernames: &[String]) -> HashMap { + // Build reverse map: table_name -> username + let mut table_to_username: HashMap = HashMap::new(); + for u in known_usernames { + let tbl = msg_table_name(u); + table_to_username.insert(tbl, u.as_str()); + } + + // Pre-fill all known usernames with 0 so sessions without Msg_* tables + // still get a per-talker baseline (rather than falling back to startup_watermark). + let mut result: HashMap = + known_usernames.iter().map(|u| (u.clone(), 0)).collect(); + + for shard in self.all_shards() { + let conn = match WechatDb::open_shard_with_key(shard, self.raw_key.as_ref()) { + Ok(c) => c, + Err(e) => { + eprintln!( + "warn: bulk_max_sort_seq: open shard {} failed: {e}", + shard.path.display() + ); + continue; + } + }; + + // Discover Msg_* tables in this shard + let mut stmt = match conn + .prepare("SELECT name FROM sqlite_master WHERE type='table' AND name LIKE 'Msg_%'") + { + Ok(s) => s, + Err(e) => { + eprintln!( + "warn: bulk_max_sort_seq: list tables in {} failed: {e}", + shard.path.display() + ); + continue; + } + }; + + let table_names: Vec = match stmt.query_map([], |row| row.get::<_, String>(0)) { + Ok(rows) => rows.filter_map(|r| r.ok()).collect(), + Err(e) => { + eprintln!( + "warn: bulk_max_sort_seq: query tables in {} failed: {e}", + shard.path.display() + ); + continue; + } + }; + + for tbl in &table_names { + let username = match table_to_username.get(tbl.as_str()) { + Some(u) => *u, + None => continue, + }; + + let sql = format!("SELECT MAX(sort_seq) FROM [{}]", tbl); + match conn.query_row(&sql, [], |row| row.get::<_, Option>(0)) { + Ok(max_seq) => { + let seq = max_seq.unwrap_or(0); + let entry = result.entry(username.to_string()).or_insert(0); + if seq > *entry { + *entry = seq; + } + } + Err(e) => { + eprintln!( + "warn: bulk_max_sort_seq: MAX(sort_seq) for {tbl} in {} failed: {e}", + shard.path.display() + ); + } + } + } + } + + result + } +} + +/// Build per-shard SQL and parameters for a regular (non-anchor) query. +fn build_regular_shard_sql( + mode: &RegularQueryMode, + select_cols: &str, + table_name: &str, + order: SortOrder, + start_time: i64, + end_time: i64, +) -> (String, Vec) { + match mode { + RegularQueryMode::FullScan => { + let sql = format!( + "SELECT {select_cols} \ + FROM [{table}] m \ + LEFT JOIN Name2Id n ON m.real_sender_id = n.rowid \ + WHERE m.create_time >= ?1 AND m.create_time <= ?2 \ + ORDER BY m.sort_seq {order}, m.create_time {order}, m.server_id {order}", + table = table_name, + order = order.sql_keyword(), + ); + (sql, vec![start_time, end_time]) + } + RegularQueryMode::LimitPushdown { + sql_limit, + msg_type_filter, + } => { + if let Some(mt) = msg_type_filter { + let sql = format!( + "SELECT {select_cols} \ + FROM [{table}] m \ + LEFT JOIN Name2Id n ON m.real_sender_id = n.rowid \ + WHERE m.create_time >= ?1 AND m.create_time <= ?2 \ + AND (m.local_type & 4294967295) = ?3 \ + ORDER BY m.sort_seq {order} \ + LIMIT ?4", + table = table_name, + order = order.sql_keyword(), + ); + ( + sql, + vec![start_time, end_time, *mt as i64, *sql_limit as i64], + ) + } else { + let sql = format!( + "SELECT {select_cols} \ + FROM [{table}] m \ + LEFT JOIN Name2Id n ON m.real_sender_id = n.rowid \ + WHERE m.create_time >= ?1 AND m.create_time <= ?2 \ + ORDER BY m.sort_seq {order} \ + LIMIT ?3", + table = table_name, + order = order.sql_keyword(), + ); + (sql, vec![start_time, end_time, *sql_limit as i64]) + } + } + } +} + +/// Execute a SQL query on a prepared shard and collect decoded message rows. +#[allow(clippy::too_many_arguments)] +fn query_shard_sql( + prepared: &PreparedShard<'_>, + sql: &str, + params: &[i64], + is_group: bool, + talker: &str, + shard_path: &str, + messages: &mut Vec, + total_rows: &mut usize, + skipped: &mut usize, + shard_warnings: &mut Vec, +) { + let mut stmt = match prepared.conn.as_conn().prepare(sql) { + Ok(s) => s, + Err(e) => { + shard_warnings.push(ShardWarning { + path: shard_path.to_string(), + reason: format!("prepare failed: {e}"), + }); + return; + } + }; + + let param_refs: Vec<&dyn rusqlite::types::ToSql> = params + .iter() + .map(|p| p as &dyn rusqlite::types::ToSql) + .collect(); + + let mut rows = match stmt.query(param_refs.as_slice()) { + Ok(r) => r, + Err(e) => { + shard_warnings.push(ShardWarning { + path: shard_path.to_string(), + reason: format!("query failed: {e}"), + }); + return; + } + }; + + collect_rows( + &mut rows, + prepared, + is_group, + talker, + shard_path, + messages, + total_rows, + skipped, + shard_warnings, + ); +} + +/// Apply keyword and msg_type post-filters to a message vector. +fn apply_post_filters( + messages: &mut Vec, + keyword: &Option, + msg_type_filter: Option, +) { + if let Some(ref kw) = keyword { + let kw_lower = kw.to_lowercase(); + messages.retain(|m| content_contains_keyword(&m.content, &kw_lower)); + } + if let Some(mt) = msg_type_filter { + messages.retain(|m| m.msg_type == mt); + } +} + +/// Collect decoded message rows from a query result set into the accumulator vectors. +#[allow(clippy::too_many_arguments)] +fn collect_rows( + rows: &mut rusqlite::Rows<'_>, + prepared: &PreparedShard, + is_group: bool, + talker: &str, + shard_path: &str, + all_messages: &mut Vec, + total_rows: &mut usize, + skipped: &mut usize, + shard_warnings: &mut Vec, +) { + loop { + match rows.next() { + Ok(Some(row)) => { + *total_rows += 1; + match decode_message_row( + row, + prepared.has_ct_col, + prepared.has_compress_col, + is_group, + talker, + ) { + Ok(msg) => all_messages.push(msg), + Err(_) => { + *skipped += 1; + } + } + } + Ok(None) => break, + Err(e) => { + shard_warnings.push(ShardWarning { + path: shard_path.to_string(), + reason: format!("row iteration failed: {e}"), + }); + break; + } + } + } +} + +/// Decode a single message row from a rusqlite Row reference. +fn decode_message_row( + row: &rusqlite::Row<'_>, + has_ct_col: bool, + has_compress_col: bool, + is_group: bool, + talker: &str, +) -> Result { + let sort_seq: i64 = row.get(0)?; + let server_id: i64 = row.get(1)?; + let local_type: i64 = row.get(2)?; + let sender_from_name2id: String = row.get(3)?; + let create_time: i64 = row.get(4)?; + + // message_content can be Text or Blob + let raw_content: Vec = match row.get_ref(5)? { + ValueRef::Blob(b) => b.to_vec(), + ValueRef::Text(b) => b.to_vec(), + ValueRef::Null => Vec::new(), + _ => Vec::new(), + }; + + // packed_info_data (BLOB, nullable) + let packed_blob: Option> = match row.get_ref(6)? { + ValueRef::Blob(b) => Some(b.to_vec()), + ValueRef::Null => None, + _ => None, + }; + + let status: i32 = row.get(7)?; + + // WCDB_CT column (optional, index=8 if present) + let wcdb_ct: Option = if has_ct_col { + row.get::<_, Option>(8)? + } else { + None + }; + + // compress_content column (optional BLOB, index depends on has_ct_col) + let compress_content: Option> = if has_compress_col { + let col_idx = 8 + (has_ct_col as usize); + match row.get_ref(col_idx)? { + ValueRef::Blob(b) if !b.is_empty() => Some(b.to_vec()), + _ => None, + } + } else { + None + }; + + // Decode content (zstd decompression if needed) + let decoded_text = decode_content(&raw_content, wcdb_ct)?; + + // Group sender parsing: extract sender from content prefix + let (sender, content_text) = parse_group_sender(is_group, decoded_text, sender_from_name2id); + + // Decode packed info + let packed_info = packed_blob.as_deref().and_then(|b| { + if b.is_empty() { + None + } else { + decode_packed_info(b) + } + }); + + // Split local_type into msg_type and sub_type + let (msg_type, sub_type) = split_local_type(local_type); + + // Parse content into typed enum + let content = parse_content( + msg_type, + sub_type, + &content_text, + server_id, + packed_info.as_ref(), + compress_content.as_deref(), + ); + + Ok(Message { + sort_seq, + server_id, + msg_type, + sub_type, + sender, + talker: talker.to_string(), + create_time, + content, + status, + }) +} + +/// Check if a MessageContent matches a keyword (case-insensitive). +fn content_contains_keyword(content: &crate::model::MessageContent, kw_lower: &str) -> bool { + use crate::model::MessageContent; + match content { + MessageContent::Text(s) => s.to_lowercase().contains(kw_lower), + MessageContent::Image { .. } => false, + MessageContent::Voice => false, + MessageContent::Video { .. } => false, + MessageContent::Emoji(s) => s.to_lowercase().contains(kw_lower), + MessageContent::Location(s) => s.to_lowercase().contains(kw_lower), + MessageContent::Link { title, des, .. } => matches_any_opt(&[title, des], kw_lower), + MessageContent::File { title, .. } => matches_opt(title, kw_lower), + MessageContent::MiniProgram { title, .. } => matches_opt(title, kw_lower), + MessageContent::MergedMessages { title, .. } => matches_opt(title, kw_lower), + MessageContent::Quote { + reply_text, + refer_content, + .. + } => matches_any_opt(&[reply_text, refer_content], kw_lower), + MessageContent::Transfer { + amount_desc, + pay_memo, + .. + } => matches_any_opt(&[amount_desc, pay_memo], kw_lower), + MessageContent::RedEnvelope { title, .. } => matches_opt(title, kw_lower), + MessageContent::ChannelVideo { title, .. } => matches_opt(title, kw_lower), + MessageContent::Pat { .. } => false, + MessageContent::AppGeneric { title, des, .. } => matches_any_opt(&[title, des], kw_lower), + MessageContent::System(s) => s.to_lowercase().contains(kw_lower), + MessageContent::Revoke(s) => s.to_lowercase().contains(kw_lower), + MessageContent::Unknown { raw, .. } => raw.to_lowercase().contains(kw_lower), + } +} + +fn matches_opt(opt: &Option, kw_lower: &str) -> bool { + opt.as_deref() + .is_some_and(|s| s.to_lowercase().contains(kw_lower)) +} + +fn matches_any_opt(opts: &[&Option], kw_lower: &str) -> bool { + opts.iter().any(|o| matches_opt(o, kw_lower)) +} diff --git a/crates/wx-db/src/model.rs b/crates/wx-db/src/model.rs new file mode 100644 index 0000000..813c37c --- /dev/null +++ b/crates/wx-db/src/model.rs @@ -0,0 +1,868 @@ +use serde::{Deserialize, Serialize}; + +use crate::error::ShardWarning; + +/// Maximum number of rows that a single query can return. +pub const MAX_QUERY_LIMIT: usize = 20_000; + +/// Default number of rows returned when no explicit limit is specified (i.e. `limit == 0`). +pub const DEFAULT_QUERY_LIMIT: usize = 1_000; + +/// Check whether a username refers to a group chat. +pub fn is_group_chat(username: &str) -> bool { + username.ends_with("@chatroom") +} + +/// Sort direction for query results. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)] +pub enum SortOrder { + Asc, + #[default] + Desc, +} + +impl SortOrder { + /// Return the SQL keyword for this sort direction. + pub fn sql_keyword(self) -> &'static str { + match self { + SortOrder::Asc => "ASC", + SortOrder::Desc => "DESC", + } + } +} + +// Message type constants + +/// Message type constant: plain text message. +pub const MSG_TYPE_TEXT: u32 = 1; +/// Message type constant: image message. +pub const MSG_TYPE_IMAGE: u32 = 3; +/// Message type constant: voice / audio message. +pub const MSG_TYPE_VOICE: u32 = 34; +/// Message type constant: video message. +pub const MSG_TYPE_VIDEO: u32 = 43; +/// Message type constant: custom emoji / sticker. +pub const MSG_TYPE_EMOJI: u32 = 47; +/// Message type constant: location sharing. +pub const MSG_TYPE_LOCATION: u32 = 48; +/// Message type constant: app / rich-media message (links, mini-programs, etc.). +pub const MSG_TYPE_APP: u32 = 49; +/// Message type constant: system notification. +pub const MSG_TYPE_SYSTEM: u32 = 10000; +/// Message type constant: message recall / revoke notification. +pub const MSG_TYPE_REVOKE: u32 = 10002; + +/// Return a human-readable label for the given message type. +pub fn msg_type_label(msg_type: u32) -> &'static str { + match msg_type { + MSG_TYPE_TEXT => "text", + MSG_TYPE_IMAGE => "image", + MSG_TYPE_VOICE => "voice", + MSG_TYPE_VIDEO => "video", + MSG_TYPE_EMOJI => "emoji", + MSG_TYPE_LOCATION => "location", + MSG_TYPE_APP => "app", + MSG_TYPE_SYSTEM => "system", + MSG_TYPE_REVOKE => "revoke", + _ => "unknown", + } +} + +/// Parse a message type string (name or numeric) into a `msg_type` value. +/// +/// Accepts the same labels returned by [`msg_type_label`] (case-insensitive) +/// plus raw numeric values (e.g. `"49"`). +pub fn parse_msg_type(s: &str) -> Option { + match s.to_lowercase().as_str() { + "text" => Some(MSG_TYPE_TEXT), + "image" => Some(MSG_TYPE_IMAGE), + "voice" => Some(MSG_TYPE_VOICE), + "video" => Some(MSG_TYPE_VIDEO), + "emoji" => Some(MSG_TYPE_EMOJI), + "location" => Some(MSG_TYPE_LOCATION), + "app" => Some(MSG_TYPE_APP), + "system" => Some(MSG_TYPE_SYSTEM), + "revoke" => Some(MSG_TYPE_REVOKE), + _ => s.parse().ok(), + } +} + +// App message sub_type constants (for first-class structured variants) + +pub const APP_SUB_TYPE_LINK: u32 = 5; +pub const APP_SUB_TYPE_FILE: u32 = 6; +pub const APP_SUB_TYPE_MINI_PROGRAM: u32 = 33; +pub const APP_SUB_TYPE_MINI_PROGRAM_2: u32 = 36; +pub const APP_SUB_TYPE_MERGED: u32 = 19; +pub const APP_SUB_TYPE_CHANNEL: u32 = 51; +pub const APP_SUB_TYPE_QUOTE: u32 = 57; +pub const APP_SUB_TYPE_PAT: u32 = 62; +pub const APP_SUB_TYPE_CHANNEL_LIVE: u32 = 63; +pub const APP_SUB_TYPE_MUSIC: u32 = 92; +pub const APP_SUB_TYPE_TRANSFER: u32 = 2000; +pub const APP_SUB_TYPE_RED_ENVELOPE: u32 = 2001; + +/// Return a human-readable label for the given message type and sub_type. +/// +/// For `MSG_TYPE_APP` messages, dispatches on `sub_type` to return a specific +/// label (e.g. `"link"`, `"quote"`, `"transfer"`). For other message types, +/// falls back to [`msg_type_label`]. Unknown app sub_types return `"app"`. +/// +/// This is the **authoritative source** for sub_type labels (English). +/// The `schema.rs` `app_sub_type_display_label` covers only the `AppGeneric` +/// fallback display path (Chinese) — sub_types that already have dedicated +/// `MessageContent` variants (link, file, mini-program, quote, transfer, etc.) +/// are handled before reaching `AppGeneric`. +pub fn msg_sub_type_label(msg_type: u32, sub_type: u32) -> &'static str { + if msg_type != MSG_TYPE_APP { + return msg_type_label(msg_type); + } + match sub_type { + 1 => "text_share", + 2 => "image_share", + 3 => "audio_share", + 4 => "video_share", + 5 => "link", + 6 => "file", + 7 => "webview", + 8 => "gif", + 10 => "location_sharing", + 13 => "brand", + 14 => "chat_log_backup", + 15 => "chat_log_migrate", + 16 => "card_ticket", + 17 => "realtime_location", + 19 => "merged_messages", + 21 => "mini_program_promo", + 24 => "note", + 33 => "mini_program", + 35 => "message_history", + 36 => "mini_program", + 40 => "channel_forward", + 44 => "channel_live_product", + 51 => "channel_video", + 53 => "group_chat_reference", + 57 => "quote", + 62 => "pat", + 63 => "channel_live", + 74 => "channel_files", + 87 => "group_announcement", + 88 => "group_note", + 92 => "music", + 100 => "sticker_set", + 101 => "ad", + 107 => "open_link", + 113 => "video_account_intro", + 116 => "channel_show_card", + 117 => "channel_product", + 124 => "wechat_gift", + 2000 => "transfer", + 2001 => "red_envelope", + 2003 => "red_envelope_cover", + _ => "app", + } +} + +/// Extract `msg_type` and `sub_type` from a `local_type` value stored in the database. +/// +/// WeChat 4.x encodes `local_type` as `(sub_type << 32) | msg_type`. +/// `msg_type` occupies the lower 32 bits, `sub_type` the upper 32 bits. +pub fn split_local_type(local_type: i64) -> (u32, u32) { + let msg_type = (local_type & 0xFFFFFFFF) as u32; + let sub_type = ((local_type >> 32) & 0xFFFFFFFF) as u32; + (msg_type, sub_type) +} + +/// Compute the effective query limit from a caller-supplied value. +/// +/// - If `limit == 0`, returns [`DEFAULT_QUERY_LIMIT`] (1000). +/// - If `limit > MAX_QUERY_LIMIT`, clamps to [`MAX_QUERY_LIMIT`] (20 000). +/// - Otherwise returns `limit` unchanged. +pub fn effective_limit(limit: usize) -> usize { + let l = if limit == 0 { + DEFAULT_QUERY_LIMIT + } else { + limit + }; + l.min(MAX_QUERY_LIMIT) +} + +/// Statistics about a query execution. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct QueryStats { + /// Total number of database rows scanned (before filtering and pagination). + pub total_rows: usize, + /// Number of rows matching keyword + type filters, before pagination. + /// `Some(n)` when application-level filtering is used (messages with `with_filtered_count`, + /// contacts with keyword search); `None` otherwise. + pub filtered_count: Option, + /// Number of rows that were skipped due to decode errors (e.g. corrupted zstd data). + pub skipped: usize, +} + +/// A paginated query result containing items and execution statistics. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct QueryResult { + /// The result items after filtering, sorting, and pagination. + pub items: Vec, + /// Statistics about the query execution. + pub stats: QueryStats, +} + +/// A message query result with shard-level fault tolerance. +/// +/// Unlike [`QueryResult`] which propagates errors, this type collects +/// warnings about individual shards that could not be read, allowing the +/// remaining shards to still return results. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct MessageQueryResult { + /// The result items after filtering, sorting, and pagination. + pub items: Vec, + /// Statistics about the query execution. + pub stats: QueryStats, + /// Warnings about shards that were skipped due to errors. + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub shard_warnings: Vec, +} + +/// A decoded WeChat message. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct Message { + /// Sort sequence number used for ordering within a shard. + pub sort_seq: i64, + /// Server-assigned unique message identifier. + pub server_id: i64, + /// The primary message type (lower 32 bits of `local_type`). + pub msg_type: u32, + /// The message sub-type (upper 32 bits of `local_type`), used by app messages. + pub sub_type: u32, + /// The wxid of the message sender. + pub sender: String, + /// The wxid of the conversation partner or chatroom. + pub talker: String, + /// Unix timestamp (seconds) when the message was created. + pub create_time: i64, + /// Typed message content parsed from raw data. + pub content: MessageContent, + /// Message status code from the database. + pub status: i32, +} + +/// Typed message content, parsed from the raw database blob based on `msg_type`. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub enum MessageContent { + /// Plain text message. + Text(String), + /// Image message with an optional MD5 hash from packed info. + Image { + /// MD5 hash of the image, extracted from protobuf `packed_info_data`. + md5: Option, + }, + /// Voice / audio message (content not decoded). + Voice, + /// Video message with an optional MD5 hash from packed info. + Video { + /// MD5 hash of the video, extracted from protobuf `packed_info_data`. + md5: Option, + }, + /// Custom emoji / sticker (raw XML content). + Emoji(String), + /// Location sharing (raw XML content). + Location(String), + /// Link share (sub_type=5, 4, 7, 92 etc.). + Link { + sub_type: u32, + title: Option, + des: Option, + url: Option, + raw_xml: String, + }, + /// File transfer (sub_type=6). + File { + title: Option, + file_ext: Option, + file_size: Option, + md5: Option, + raw_xml: String, + }, + /// Mini program (sub_type=33, 36). + MiniProgram { + sub_type: u32, + title: Option, + url: Option, + raw_xml: String, + }, + /// Merged forwarded messages (sub_type=19). + MergedMessages { + title: Option, + raw_xml: String, + }, + /// Quote / reply (sub_type=57). + Quote { + reply_text: Option, + refer_sender: Option, + refer_content: Option, + refer_type: Option, + raw_xml: String, + }, + /// Transfer / payment (sub_type=2000). + Transfer { + amount_desc: Option, + pay_memo: Option, + pay_sub_type: Option, + raw_xml: String, + }, + /// Red envelope (sub_type=2001, 2003). + RedEnvelope { + title: Option, + raw_xml: String, + }, + /// Channel video / live (sub_type=51, 63). + ChannelVideo { + sub_type: u32, + title: Option, + raw_xml: String, + }, + /// Pat message (sub_type=62). + Pat { raw_xml: String }, + /// Generic app message (unknown sub_type fallback). + AppGeneric { + sub_type: u32, + title: Option, + des: Option, + url: Option, + raw_xml: String, + }, + /// System notification message. + System(String), + /// Message recall / revoke notification. + Revoke(String), + /// Unknown or unsupported message type, preserved as raw text. + Unknown { + /// The unrecognized `msg_type` value. + msg_type: u32, + /// The raw content text. + raw: String, + }, +} + +/// Decoded protobuf packed info attached to image/video messages. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct PackedInfo { + /// MD5 hash of the image, if present in the protobuf. + pub image_md5: Option, + /// MD5 hash of the video, if present in the protobuf. + pub video_md5: Option, +} + +/// A WeChat contact entry from `contact.db`. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct Contact { + /// The unique WeChat user name (wxid). + pub user_name: String, + /// The user-set alias (WeChat ID). + pub alias: String, + /// The remark name set by the account owner. + pub remark: String, + /// The user's display nickname. + pub nick_name: String, + /// Memo / description set by the account owner. + pub memo: Option, + /// Gender from extra_buffer (1=male, 2=female). + pub gender: Option, + /// Personal signature from extra_buffer. + pub signature: Option, + /// Region from extra_buffer (country · province · city). + pub region: Option, + /// Source scene code from extra_buffer. + pub source_scene: Option, + /// Phone number from extra_buffer. + pub phone: Option, + /// Resolved label names from extra_buffer + contact_label table. + pub labels: Vec, +} + +/// A WeChat chatroom (group chat) entry with its member list. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ChatRoom { + /// The chatroom identifier (e.g. `"12345@chatroom"`). + pub username: String, + /// The wxid of the chatroom owner. + pub owner: String, + /// List of chatroom members decoded from the protobuf `ext_buffer`. + pub members: Vec, +} + +/// A single member within a chatroom. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ChatRoomMember { + /// The member's wxid. + pub user_name: String, + /// The member's in-group display name, if set. + pub display_name: Option, +} + +/// A recent conversation session from `session.db`. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct Session { + /// The wxid of the conversation partner or chatroom. + pub username: String, + /// Summary text of the last message in this session. + pub summary: String, + /// Unix timestamp (seconds) of the last activity, used for sort order. + pub sort_timestamp: i64, + /// The message type of the last message (e.g. 1=text, 3=image). + pub last_msg_type: Option, + /// The wxid of the last message sender. + pub last_msg_sender: Option, + /// The display name of the last message sender. + pub last_sender_display_name: Option, +} + +// --- Query structs with builder methods --- + +/// Anchor mode for context-based message queries. +#[derive(Debug, Clone)] +pub enum AnchorMode { + /// Query messages around a sort_seq (before + pivot bucket + after). + /// When multiple messages share this sort_seq, all are included as the pivot bucket. + AroundSortSeq(i64), + /// Query messages around a server_id (two-phase: locate → context). + /// Locates the exact message by server_id, retrieves its (sort_seq, create_time, server_id) + /// as compound pivot key, then queries context around that precise position. + AroundServerId(i64), + /// Query messages strictly after a sort_seq (incremental pull). + AfterSortSeq(i64), +} + +/// Default context window size for around queries (messages before/after pivot). +pub const DEFAULT_CONTEXT: usize = 50; + +/// Parameters for querying messages from a specific conversation. +/// +/// Use [`MessageQuery::for_talker`] to create a query, then chain builder +/// methods to refine it: +/// +/// ```ignore +/// let q = MessageQuery::for_talker("wxid_alice") +/// .time_range(1700000000, 1710000000) +/// .keyword("hello") +/// .limit(50) +/// .offset(0); +/// ``` +#[derive(Debug, Clone)] +pub struct MessageQuery { + /// The wxid or chatroom ID to query messages for. + pub talker: String, + /// Start of the time range filter (inclusive, Unix seconds). Default: `0`. + pub start_time: i64, + /// End of the time range filter (inclusive, Unix seconds). Default: `i64::MAX`. + pub end_time: i64, + /// Optional keyword for post-SQL content filtering (case-insensitive). + pub keyword: Option, + /// Maximum number of results to return. `0` means use [`DEFAULT_QUERY_LIMIT`]. + pub limit: usize, + /// Number of results to skip for pagination. + pub offset: usize, + /// Sort direction for results. Default: [`SortOrder::Desc`]. + pub order: SortOrder, + /// Whether to compute `filtered_count` in [`QueryStats`]. + pub with_filtered_count: bool, + /// Optional message type filter (applied after keyword filter, before pagination). + pub msg_type_filter: Option, + /// Anchor mode for context-based queries (mutually exclusive with time range / offset). + pub anchor: Option, + /// Context window size for around queries (messages before/after pivot). Default: 50. + pub context: usize, +} + +impl MessageQuery { + /// Create a new query targeting the given talker (wxid or chatroom ID). + pub fn for_talker(talker: impl Into) -> Self { + Self { + talker: talker.into(), + start_time: 0, + end_time: i64::MAX, + keyword: None, + limit: 0, + offset: 0, + order: SortOrder::default(), + with_filtered_count: false, + msg_type_filter: None, + anchor: None, + context: DEFAULT_CONTEXT, + } + } + + /// Set the time range filter `[start, end]` (inclusive, Unix seconds). + pub fn time_range(mut self, start: i64, end: i64) -> Self { + self.start_time = start; + self.end_time = end; + self + } + + /// Set only the start of the time range (inclusive, Unix seconds). + pub fn since(mut self, start: i64) -> Self { + self.start_time = start; + self + } + + /// Set only the end of the time range (inclusive, Unix seconds). + pub fn until(mut self, end: i64) -> Self { + self.end_time = end; + self + } + + /// Set a keyword filter. Only messages whose text content contains + /// this keyword (case-insensitive) will be returned. + pub fn keyword(mut self, kw: impl Into) -> Self { + self.keyword = Some(kw.into()); + self + } + + /// Set the maximum number of results. Values of `0` use + /// [`DEFAULT_QUERY_LIMIT`]; values above [`MAX_QUERY_LIMIT`] are clamped. + pub fn limit(mut self, limit: usize) -> Self { + self.limit = limit; + self + } + + /// Set the pagination offset (number of results to skip). + pub fn offset(mut self, offset: usize) -> Self { + self.offset = offset; + self + } + + /// Set the sort direction for results. + pub fn order(mut self, order: SortOrder) -> Self { + self.order = order; + self + } + + /// Enable or disable `filtered_count` computation in [`QueryStats`]. + pub fn with_filtered_count(mut self, yes: bool) -> Self { + self.with_filtered_count = yes; + self + } + + /// Set a message type filter. Only messages with this `msg_type` will be returned. + pub fn msg_type(mut self, mt: u32) -> Self { + self.msg_type_filter = Some(mt); + self + } + + /// Set anchor mode to query messages around a specific sort_seq. + pub fn around_sort_seq(mut self, seq: i64) -> Self { + self.anchor = Some(AnchorMode::AroundSortSeq(seq)); + self + } + + /// Set anchor mode to query messages around the message with a specific server_id. + pub fn around_server_id(mut self, id: i64) -> Self { + self.anchor = Some(AnchorMode::AroundServerId(id)); + self + } + + /// Set anchor mode to query messages strictly after a specific sort_seq. + pub fn after_sort_seq(mut self, seq: i64) -> Self { + self.anchor = Some(AnchorMode::AfterSortSeq(seq)); + self + } + + /// Set the context window size for around queries. + pub fn context(mut self, n: usize) -> Self { + self.context = n; + self + } + + /// Whether this query is eligible for SQL LIMIT pushdown. + /// + /// Returns `true` when there is no keyword search, no `filtered_count` + /// request, and no anchor mode — i.e. a straightforward paginated browse. + pub fn limit_pushdown_eligible(&self) -> bool { + self.keyword.is_none() && !self.with_filtered_count && self.anchor.is_none() + } +} + +/// Parameters for querying contacts. +/// +/// ```ignore +/// let q = ContactQuery::new().keyword("alice").limit(10); +/// ``` +#[derive(Debug, Clone)] +pub struct ContactQuery { + /// Optional keyword to search across userName, alias, remark, nickName, description, + /// phone, labels, signature, and region. + pub keyword: Option, + /// Maximum number of results. `0` means use [`DEFAULT_QUERY_LIMIT`]. + pub limit: usize, + /// Number of results to skip for pagination. + pub offset: usize, +} + +impl ContactQuery { + /// Create a new contact query with default parameters. + pub fn new() -> Self { + Self { + keyword: None, + limit: 0, + offset: 0, + } + } + + /// Set a keyword filter. Matches against userName, alias, remark, nickName, + /// description, phone, labels, signature, and region. + pub fn keyword(mut self, kw: impl Into) -> Self { + self.keyword = Some(kw.into()); + self + } + + /// Set the maximum number of results. + pub fn limit(mut self, limit: usize) -> Self { + self.limit = limit; + self + } + + /// Set the pagination offset. + pub fn offset(mut self, offset: usize) -> Self { + self.offset = offset; + self + } +} + +impl Default for ContactQuery { + fn default() -> Self { + Self::new() + } +} + +/// Parameters for querying chatrooms. +/// +/// ```ignore +/// let q = ChatRoomQuery::new().username("12345@chatroom"); +/// ``` +#[derive(Debug, Clone)] +pub struct ChatRoomQuery { + /// Optional chatroom username to filter by. If `None`, returns all chatrooms. + pub username: Option, + /// Maximum number of results. `0` means use [`DEFAULT_QUERY_LIMIT`]. + pub limit: usize, + /// Number of results to skip for pagination. + pub offset: usize, +} + +impl ChatRoomQuery { + /// Create a new chatroom query with default parameters. + pub fn new() -> Self { + Self { + username: None, + limit: 0, + offset: 0, + } + } + + /// Filter to a specific chatroom by its username. + pub fn username(mut self, name: impl Into) -> Self { + self.username = Some(name.into()); + self + } + + /// Set the maximum number of results. + pub fn limit(mut self, limit: usize) -> Self { + self.limit = limit; + self + } + + /// Set the pagination offset. + pub fn offset(mut self, offset: usize) -> Self { + self.offset = offset; + self + } +} + +impl Default for ChatRoomQuery { + fn default() -> Self { + Self::new() + } +} + +/// Parameters for querying recent sessions (conversations). +/// +/// ```ignore +/// let q = SessionQuery::new().limit(20); +/// ``` +#[derive(Debug, Clone)] +pub struct SessionQuery { + /// Maximum number of results. `0` means use [`DEFAULT_QUERY_LIMIT`]. + pub limit: usize, + /// Number of results to skip for pagination. + pub offset: usize, + /// Sort direction for results. Default: [`SortOrder::Desc`]. + pub order: SortOrder, +} + +impl SessionQuery { + /// Create a new session query with default parameters. + pub fn new() -> Self { + Self { + limit: 0, + offset: 0, + order: SortOrder::default(), + } + } + + /// Set the maximum number of results. + pub fn limit(mut self, limit: usize) -> Self { + self.limit = limit; + self + } + + /// Set the pagination offset. + pub fn offset(mut self, offset: usize) -> Self { + self.offset = offset; + self + } + + /// Set the sort direction for results. + pub fn order(mut self, order: SortOrder) -> Self { + self.order = order; + self + } +} + +impl Default for SessionQuery { + fn default() -> Self { + Self::new() + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn parse_msg_type_names() { + assert_eq!(parse_msg_type("text"), Some(MSG_TYPE_TEXT)); + assert_eq!(parse_msg_type("IMAGE"), Some(MSG_TYPE_IMAGE)); + assert_eq!(parse_msg_type("Voice"), Some(MSG_TYPE_VOICE)); + assert_eq!(parse_msg_type("video"), Some(MSG_TYPE_VIDEO)); + assert_eq!(parse_msg_type("emoji"), Some(MSG_TYPE_EMOJI)); + assert_eq!(parse_msg_type("location"), Some(MSG_TYPE_LOCATION)); + assert_eq!(parse_msg_type("app"), Some(MSG_TYPE_APP)); + assert_eq!(parse_msg_type("system"), Some(MSG_TYPE_SYSTEM)); + assert_eq!(parse_msg_type("revoke"), Some(MSG_TYPE_REVOKE)); + } + + #[test] + fn parse_msg_type_numeric_fallback() { + assert_eq!(parse_msg_type("49"), Some(49)); + assert_eq!(parse_msg_type("10002"), Some(10002)); + } + + #[test] + fn parse_msg_type_invalid() { + assert_eq!(parse_msg_type("unknown"), None); + assert_eq!(parse_msg_type(""), None); + } + + #[test] + fn parse_msg_type_roundtrip() { + for &mt in &[ + MSG_TYPE_TEXT, + MSG_TYPE_IMAGE, + MSG_TYPE_VOICE, + MSG_TYPE_VIDEO, + MSG_TYPE_EMOJI, + MSG_TYPE_LOCATION, + MSG_TYPE_APP, + MSG_TYPE_SYSTEM, + MSG_TYPE_REVOKE, + ] { + let label = msg_type_label(mt); + assert_eq!( + parse_msg_type(label), + Some(mt), + "roundtrip failed for {label}" + ); + } + } + + #[test] + fn message_query_since_until() { + let q = MessageQuery::for_talker("test").since(100).until(200); + assert_eq!(q.start_time, 100); + assert_eq!(q.end_time, 200); + } + + #[test] + fn message_query_since_only() { + let q = MessageQuery::for_talker("test").since(100); + assert_eq!(q.start_time, 100); + assert_eq!(q.end_time, i64::MAX); + } + + #[test] + fn message_query_until_only() { + let q = MessageQuery::for_talker("test").until(200); + assert_eq!(q.start_time, 0); + assert_eq!(q.end_time, 200); + } + + #[test] + fn limit_pushdown_eligible_basic() { + let q = MessageQuery::for_talker("test"); + assert!(q.limit_pushdown_eligible()); + } + + #[test] + fn limit_pushdown_blocked_by_keyword() { + let q = MessageQuery::for_talker("test").keyword("hello"); + assert!(!q.limit_pushdown_eligible()); + } + + #[test] + fn limit_pushdown_blocked_by_filtered_count() { + let q = MessageQuery::for_talker("test").with_filtered_count(true); + assert!(!q.limit_pushdown_eligible()); + } + + #[test] + fn limit_pushdown_blocked_by_anchor() { + let q = MessageQuery::for_talker("test").around_sort_seq(100); + assert!(!q.limit_pushdown_eligible()); + } + + #[test] + fn limit_pushdown_eligible_with_msg_type_and_limit() { + let q = MessageQuery::for_talker("test") + .msg_type(1) + .limit(50) + .offset(10); + assert!(q.limit_pushdown_eligible()); + } + + #[test] + fn message_query_anchor_builders() { + let q = MessageQuery::for_talker("test") + .around_sort_seq(12345) + .context(20); + assert!(matches!(q.anchor, Some(AnchorMode::AroundSortSeq(12345)))); + assert_eq!(q.context, 20); + + let q = MessageQuery::for_talker("test").around_server_id(999); + assert!(matches!(q.anchor, Some(AnchorMode::AroundServerId(999)))); + assert_eq!(q.context, DEFAULT_CONTEXT); + + let q = MessageQuery::for_talker("test").after_sort_seq(5000); + assert!(matches!(q.anchor, Some(AnchorMode::AfterSortSeq(5000)))); + } + + #[test] + fn test_is_group_chat() { + assert!(is_group_chat("12345678@chatroom")); + assert!(is_group_chat("@chatroom")); + assert!(!is_group_chat("wxid_abc123")); + assert!(!is_group_chat("filehelper")); + assert!(!is_group_chat("")); + } +} diff --git a/crates/wx-db/src/native_fts.rs b/crates/wx-db/src/native_fts.rs new file mode 100644 index 0000000..8610f81 --- /dev/null +++ b/crates/wx-db/src/native_fts.rs @@ -0,0 +1,478 @@ +use std::collections::HashMap; + +use rusqlite::Connection; +use serde::Serialize; + +use crate::error::DbError; +use crate::fts::build_fts_query; +use crate::model::split_local_type; + +// --------------------------------------------------------------------------- +// Public types +// --------------------------------------------------------------------------- + +/// The type of content hit (message, contact, or image OCR). +#[derive(Debug, Clone, Serialize)] +#[serde(rename_all = "lowercase")] +pub enum FtsHitType { + Message, + Contact, + Image, +} + +/// A single hit from native FTS search. +#[derive(Debug, Clone, Serialize)] +pub struct NativeFtsHit { + /// Hit type — always `Message` for Task 3 results. + pub hit_type: FtsHitType, + /// Resolved wxid or chatroom identifier, via name2id table. + pub talker: String, + /// Resolved wxid of the sender, via name2id table. + pub sender: String, + /// Raw message content text (acontent column verbatim). + pub snippet: String, + /// Unix seconds creation time. + pub create_time: i64, + /// Millisecond sort sequence. + pub sort_seq: i64, + /// Message local ID (c1), used as unique tie-breaker for stable pagination. + pub message_local_id: i64, + /// Primary message type (lower 32 bits of local_type). + pub msg_type: u32, + /// Message sub-type (upper 32 bits of local_type). + pub sub_type: u32, + /// Always `0` for native FTS results (breaking change — see plan). + pub server_id: i64, +} + +/// Result of a native FTS search operation. +#[derive(Debug, Clone)] +pub struct NativeFtsResult { + pub hits: Vec, + /// Total matched rows (before limit/offset). + pub total_hits: usize, +} + +// --------------------------------------------------------------------------- +// Helpers +// --------------------------------------------------------------------------- + +/// Load the `name2id` lookup table from the given connection. +/// +/// Returns a `HashMap` for resolving session_id and sender_id. +pub fn load_name2id(conn: &Connection) -> Result, DbError> { + let mut stmt = conn.prepare("SELECT rowid, username FROM name2id")?; + let mut map = HashMap::new(); + let mut rows = stmt.query([])?; + while let Some(row) = rows.next()? { + let rowid: i64 = row.get(0)?; + let username: String = row.get(1)?; + map.insert(rowid, username); + } + Ok(map) +} + +// --------------------------------------------------------------------------- +// Message FTS search +// --------------------------------------------------------------------------- + +/// Search the native WeChat message FTS database across all 4 shards. +/// +/// `conn` must already have `MMFtsTokenizer` registered (via `open_fts_connection`). +/// +/// Column layout (named columns in message_fts.db): +/// - acontent (the text that was indexed) +/// - message_local_id +/// - sort_seq (milliseconds) +/// - local_type (encodes msg_type + sub_type) +/// - session_id (rowid into name2id → talker wxid) +/// - sender_id (rowid into name2id → sender wxid) +/// - create_time (Unix seconds) +pub fn search_message_fts( + conn: &Connection, + keyword: &str, + limit: usize, + offset: usize, +) -> Result { + search_message_fts_with_cache(conn, keyword, limit, offset, None) +} + +/// Search the native WeChat message FTS database with an optional pre-loaded name2id cache. +/// +/// When `name2id_cache` is `Some`, it is used directly instead of querying the database. +/// This avoids re-loading the name2id table on every search when the caller already has it. +pub fn search_message_fts_with_cache( + conn: &Connection, + keyword: &str, + limit: usize, + offset: usize, + name2id_cache: Option<&HashMap>, +) -> Result { + let fts_query = build_fts_query(keyword); + if fts_query.is_empty() { + return Ok(NativeFtsResult { + hits: Vec::new(), + total_hits: 0, + }); + } + + // Use the provided cache or load from the database. + let owned_name2id; + let name2id: &HashMap = match name2id_cache { + Some(cache) => cache, + None => { + owned_name2id = load_name2id(conn)?; + &owned_name2id + } + }; + + // Build UNION ALL across all 4 shards (using real column names from message_fts.db) + let union_sql = " + SELECT acontent, message_local_id, sort_seq, local_type, session_id, sender_id, create_time + FROM message_fts_v4_0 WHERE message_fts_v4_0 MATCH ?1 + UNION ALL + SELECT acontent, message_local_id, sort_seq, local_type, session_id, sender_id, create_time + FROM message_fts_v4_1 WHERE message_fts_v4_1 MATCH ?1 + UNION ALL + SELECT acontent, message_local_id, sort_seq, local_type, session_id, sender_id, create_time + FROM message_fts_v4_2 WHERE message_fts_v4_2 MATCH ?1 + UNION ALL + SELECT acontent, message_local_id, sort_seq, local_type, session_id, sender_id, create_time + FROM message_fts_v4_3 WHERE message_fts_v4_3 MATCH ?1 + ORDER BY create_time DESC, sort_seq DESC, message_local_id DESC + LIMIT ?2 OFFSET ?3 + "; + + let count_sql = " + SELECT count(*) FROM ( + SELECT 1 FROM message_fts_v4_0 WHERE message_fts_v4_0 MATCH ?1 + UNION ALL + SELECT 1 FROM message_fts_v4_1 WHERE message_fts_v4_1 MATCH ?1 + UNION ALL + SELECT 1 FROM message_fts_v4_2 WHERE message_fts_v4_2 MATCH ?1 + UNION ALL + SELECT 1 FROM message_fts_v4_3 WHERE message_fts_v4_3 MATCH ?1 + ) + "; + + // Count total matches + let total_hits: usize = conn.query_row(count_sql, rusqlite::params![fts_query], |row| { + row.get::<_, i64>(0) + })? as usize; + + // Fetch paginated results + let mut stmt = conn.prepare(union_sql)?; + + let rows: Vec<(String, i64, i64, i64, i64, i64, i64)> = stmt + .query_map( + rusqlite::params![fts_query, limit as i64, offset as i64], + |row| { + Ok(( + row.get::<_, String>(0)?, // acontent + row.get::<_, i64>(1)?, // message_local_id + row.get::<_, i64>(2)?, // sort_seq + row.get::<_, i64>(3)?, // local_type + row.get::<_, i64>(4)?, // session_id + row.get::<_, i64>(5)?, // sender_id + row.get::<_, i64>(6)?, // create_time + )) + }, + )? + .collect::, _>>()?; + + let hits: Vec = rows + .into_iter() + .map( + |( + snippet, + message_local_id, + sort_seq, + local_type, + session_id, + sender_id, + create_time, + )| { + let (msg_type, sub_type) = split_local_type(local_type); + let talker = name2id + .get(&session_id) + .cloned() + .unwrap_or_else(|| format!("unknown:{session_id}")); + let sender = name2id + .get(&sender_id) + .cloned() + .unwrap_or_else(|| format!("unknown:{sender_id}")); + NativeFtsHit { + hit_type: FtsHitType::Message, + talker, + sender, + snippet, + create_time, + sort_seq, + message_local_id, + msg_type, + sub_type, + server_id: 0, + } + }, + ) + .collect(); + + Ok(NativeFtsResult { hits, total_hits }) +} + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +#[cfg(test)] +mod tests { + use super::*; + use rusqlite::Connection; + + fn open_fts_test_db() -> Connection { + let conn = Connection::open_in_memory().unwrap(); + // Register the MMFtsTokenizer + // We need to call register_mm_fts_tokenizer from wx_context, but since + // wx-db doesn't depend on wx-context, we use the unicode61 tokenizer + // for unit tests here. The integration test in wx-cli uses the real tokenizer. + // Use real column names matching message_fts.db schema + conn.execute_batch( + "CREATE VIRTUAL TABLE message_fts_v4_0 USING fts5(acontent, message_local_id UNINDEXED, sort_seq UNINDEXED, local_type UNINDEXED, session_id UNINDEXED, sender_id UNINDEXED, create_time UNINDEXED, tokenize='unicode61'); + CREATE VIRTUAL TABLE message_fts_v4_1 USING fts5(acontent, message_local_id UNINDEXED, sort_seq UNINDEXED, local_type UNINDEXED, session_id UNINDEXED, sender_id UNINDEXED, create_time UNINDEXED, tokenize='unicode61'); + CREATE VIRTUAL TABLE message_fts_v4_2 USING fts5(acontent, message_local_id UNINDEXED, sort_seq UNINDEXED, local_type UNINDEXED, session_id UNINDEXED, sender_id UNINDEXED, create_time UNINDEXED, tokenize='unicode61'); + CREATE VIRTUAL TABLE message_fts_v4_3 USING fts5(acontent, message_local_id UNINDEXED, sort_seq UNINDEXED, local_type UNINDEXED, session_id UNINDEXED, sender_id UNINDEXED, create_time UNINDEXED, tokenize='unicode61'); + CREATE TABLE name2id (rowid INTEGER PRIMARY KEY, username TEXT NOT NULL);", + ) + .unwrap(); + conn + } + + #[allow(clippy::too_many_arguments)] + fn insert_row( + conn: &Connection, + shard: usize, + content: &str, + message_local_id: i64, + sort_seq: i64, + local_type: i64, + session_id: i64, + sender_id: i64, + create_time: i64, + ) { + let table = format!("message_fts_v4_{shard}"); + conn.execute( + &format!( + "INSERT INTO {table}(acontent,message_local_id,sort_seq,local_type,session_id,sender_id,create_time) VALUES(?1,?2,?3,?4,?5,?6,?7)" + ), + rusqlite::params![content, message_local_id, sort_seq, local_type, session_id, sender_id, create_time], + ) + .unwrap(); + } + + // Test 1: load_name2id correctly builds the HashMap using real schema (username column) + #[test] + fn load_name2id_basic() { + let conn = Connection::open_in_memory().unwrap(); + conn.execute_batch( + "CREATE TABLE name2id (rowid INTEGER PRIMARY KEY, username TEXT); + INSERT INTO name2id VALUES (1, 'wxid_alice'); + INSERT INTO name2id VALUES (2, 'wxid_bob');", + ) + .unwrap(); + let map = load_name2id(&conn).unwrap(); + assert_eq!(map.get(&1), Some(&"wxid_alice".to_string())); + assert_eq!(map.get(&2), Some(&"wxid_bob".to_string())); + assert_eq!(map.len(), 2); + } + + // Test: load_name2id fails with wrong column name (regression guard for BUG-3) + #[test] + fn load_name2id_wrong_column_fails() { + let conn = Connection::open_in_memory().unwrap(); + conn.execute_batch( + "CREATE TABLE name2id (rowid INTEGER PRIMARY KEY, user_name TEXT); + INSERT INTO name2id VALUES (1, 'wxid_alice');", + ) + .unwrap(); + assert!( + load_name2id(&conn).is_err(), + "load_name2id should fail when name2id has user_name column instead of username" + ); + } + + // Test 2: Integration: messages in multiple shards, verify merge + ordering + #[test] + fn search_across_shards() { + let conn = open_fts_test_db(); + // Add name2id entries + conn.execute_batch( + "INSERT INTO name2id VALUES (1, 'wxid_alice'); + INSERT INTO name2id VALUES (2, 'wxid_self');", + ) + .unwrap(); + + // Insert rows in different shards + insert_row( + &conn, + 0, + "hello world from shard 0", + 1, + 1000, + 1, + 1, + 2, + 1700000001, + ); + insert_row( + &conn, + 1, + "hello again from shard 1", + 2, + 2000, + 1, + 1, + 2, + 1700000002, + ); + insert_row(&conn, 2, "goodbye shard 2", 3, 3000, 1, 1, 2, 1700000003); + insert_row( + &conn, + 3, + "hello final shard 3", + 4, + 4000, + 1, + 1, + 2, + 1700000004, + ); + + let result = search_message_fts(&conn, "hello", 10, 0).unwrap(); + assert_eq!(result.total_hits, 3, "should find 3 hello messages"); + assert_eq!(result.hits.len(), 3); + // Ordered by create_time DESC + assert!(result.hits[0].create_time >= result.hits[1].create_time); + assert!(result.hits[1].create_time >= result.hits[2].create_time); + } + + // Test 3: Empty keyword returns empty results + #[test] + fn empty_keyword() { + let conn = open_fts_test_db(); + let result = search_message_fts(&conn, "", 10, 0).unwrap(); + assert_eq!(result.total_hits, 0); + assert!(result.hits.is_empty()); + } + + // Test 4: No matches returns 0 hits + #[test] + fn no_matches() { + let conn = open_fts_test_db(); + conn.execute_batch("INSERT INTO name2id VALUES (1, 'wxid_alice');") + .unwrap(); + insert_row(&conn, 0, "hello world", 1, 1000, 1, 1, 1, 1700000001); + + let result = search_message_fts(&conn, "nonexistent", 10, 0).unwrap(); + assert_eq!(result.total_hits, 0); + assert!(result.hits.is_empty()); + } + + // Test 5: Pagination: limit/offset work correctly + #[test] + fn pagination() { + let conn = open_fts_test_db(); + conn.execute_batch("INSERT INTO name2id VALUES (1, 'wxid_alice');") + .unwrap(); + + for i in 0..5usize { + insert_row( + &conn, + i % 4, + "hello message", + i as i64 + 1, // message_local_id (distinct) + (i as i64) * 100, + 1, + 1, + 1, + 1700000000 + i as i64, + ); + } + + // All 5 + let all = search_message_fts(&conn, "hello", 10, 0).unwrap(); + assert_eq!(all.total_hits, 5); + assert_eq!(all.hits.len(), 5); + + // First 2 + let page1 = search_message_fts(&conn, "hello", 2, 0).unwrap(); + assert_eq!(page1.total_hits, 5); + assert_eq!(page1.hits.len(), 2); + + // Offset 3, limit 10 → 2 remaining + let page2 = search_message_fts(&conn, "hello", 10, 3).unwrap(); + assert_eq!(page2.total_hits, 5); + assert_eq!(page2.hits.len(), 2); + } + + // name2id resolution + #[test] + fn name2id_resolution() { + let conn = open_fts_test_db(); + conn.execute_batch( + "INSERT INTO name2id VALUES (10, 'wxid_alice'); + INSERT INTO name2id VALUES (20, 'wxid_bob');", + ) + .unwrap(); + insert_row(&conn, 0, "hello", 1, 1000, 1, 10, 20, 1700000001); + + let result = search_message_fts(&conn, "hello", 10, 0).unwrap(); + assert_eq!(result.hits.len(), 1); + assert_eq!(result.hits[0].talker, "wxid_alice"); + assert_eq!(result.hits[0].sender, "wxid_bob"); + assert_eq!(result.hits[0].server_id, 0); + assert_eq!(result.hits[0].message_local_id, 1); + } + + // Test: pagination is stable even when (create_time, sort_seq) ties exist + // Rows with same (c6, c2) but distinct c1 must not appear in both pages. + #[test] + fn pagination_stable_with_same_sort_key() { + let conn = open_fts_test_db(); + conn.execute_batch("INSERT INTO name2id VALUES (1, 'wxid_alice');") + .unwrap(); + + // Insert 6 rows all with the same (create_time=1700000000, sort_seq=1000) + // but distinct message_local_id (c1) to verify tie-breaker works. + for mid in 1i64..=6 { + insert_row( + &conn, + ((mid - 1) % 4) as usize, + "same key message", + mid, // message_local_id (distinct tie-breaker) + 1000, // sort_seq (same for all) + 1, + 1, + 1, + 1700000000, // create_time (same for all) + ); + } + + let page1 = search_message_fts(&conn, "same", 3, 0).unwrap(); + let page2 = search_message_fts(&conn, "same", 3, 3).unwrap(); + + assert_eq!(page1.total_hits, 6); + assert_eq!(page1.hits.len(), 3); + assert_eq!(page2.total_hits, 6); + assert_eq!(page2.hits.len(), 3); + + // No overlap: message_local_ids in page1 and page2 must be disjoint + let ids1: std::collections::HashSet = + page1.hits.iter().map(|h| h.message_local_id).collect(); + let ids2: std::collections::HashSet = + page2.hits.iter().map(|h| h.message_local_id).collect(); + assert!( + ids1.is_disjoint(&ids2), + "Pagination overlap detected! page1={ids1:?}, page2={ids2:?}" + ); + } +} diff --git a/crates/wx-db/src/open.rs b/crates/wx-db/src/open.rs new file mode 100644 index 0000000..94dbc42 --- /dev/null +++ b/crates/wx-db/src/open.rs @@ -0,0 +1,466 @@ +use std::collections::HashMap; +use std::fmt; +use std::os::raw::c_void; +use std::path::{Path, PathBuf}; +use std::sync::{Arc, RwLock}; + +use rusqlite::Connection; + +use crate::error::DbError; +use crate::pool::ShardPool; +use crate::shard_metadata::{now_nanos, ShardMeta, ShardMetadataFile}; + +/// Metadata for a single message shard database file. +#[derive(Debug)] +pub(crate) struct MessageShard { + pub path: PathBuf, + pub start_unix: i64, + pub end_unix: i64, +} + +/// Handle to an opened (decrypted) WeChat database directory. +/// +/// Holds connections to contact/session databases and metadata about +/// message shard files. Created via [`WechatDb::open`]. +pub struct WechatDb { + pub(crate) contact_conn: Connection, + pub(crate) contact_path: PathBuf, + pub(crate) session_conn: Connection, + pub(crate) session_path: PathBuf, + pub(crate) shards: Vec, + /// Path to `message/message_fts.db` if it exists. + pub message_fts_path: Option, + /// Path to `message/contact_fts.db` if it exists. + pub contact_fts_path: Option, + /// Optional pre-opened connection pool for serve mode. + pub(crate) pool: Option, + /// Raw key for encrypted direct open. Stored for reopen operations. + pub(crate) raw_key: Option<[u8; 32]>, + /// Lazily initialized cache of label_id -> label_name from contact_label table. + /// Cleared on `reopen_contacts()` so label changes are visible. + pub(crate) label_cache: RwLock>>, +} + +impl fmt::Debug for WechatDb { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("WechatDb") + .field("shards", &self.shards) + .finish_non_exhaustive() + } +} + +/// Open a read-only connection, optionally applying sqlite3_key for encrypted DBs. +pub fn open_readonly_connection( + path: &Path, + raw_key: Option<&[u8; 32]>, +) -> Result { + let conn = Connection::open_with_flags(path, rusqlite::OpenFlags::SQLITE_OPEN_READ_ONLY)?; + if let Some(key) = raw_key { + unsafe { + let rc = rusqlite::ffi::sqlite3_key(conn.handle(), key.as_ptr() as *const c_void, 32); + if rc != 0 { + return Err(DbError::EncryptionKey(format!( + "sqlite3_key failed: rc={rc}" + ))); + } + } + conn.query_row("SELECT count(*) FROM sqlite_master", [], |r| { + r.get::<_, i64>(0) + }) + .map_err(|_| DbError::EncryptionKey("incorrect key or not an encrypted database".into()))?; + conn.execute_batch("PRAGMA query_only = ON")?; + } + Ok(conn) +} + +pub(crate) fn open_connection( + path: &Path, + raw_key: Option<&[u8; 32]>, +) -> Result { + open_readonly_connection(path, raw_key) +} + +impl WechatDb { + /// Open a decrypted WeChat database directory. + /// + /// Returns `DbError::NotFound` if the path, contact.db, or session.db + /// does not exist. Message shards are optional here; message queries will + /// return `DbError::NoShards` if no numbered shard is available. + pub fn open(path: impl AsRef) -> Result { + Self::open_internal(path.as_ref(), None) + } + + /// Open a decrypted WeChat database directory with a pre-opened + /// connection pool for all message shards and FTS. + /// + /// `fts_init` is called on the FTS connection to register custom + /// tokenizers (e.g. `register_mm_fts_tokenizer`). + pub fn open_with_pool( + path: impl AsRef, + fts_init: impl Fn(&Connection) -> Result<(), String> + Send + Sync + 'static, + ) -> Result { + Self::open_with_pool_internal(path, None, fts_init) + } + + /// Open an encrypted WeChat database directory directly using `sqlite3_key()`. + pub fn open_encrypted(path: impl AsRef, raw_key: [u8; 32]) -> Result { + Self::open_internal(path.as_ref(), Some(raw_key)) + } + + /// Open an encrypted WeChat database directory with a pre-opened + /// connection pool for all message shards and FTS. + pub fn open_encrypted_with_pool( + path: impl AsRef, + raw_key: [u8; 32], + fts_init: impl Fn(&Connection) -> Result<(), String> + Send + Sync + 'static, + ) -> Result { + Self::open_with_pool_internal(path, Some(raw_key), fts_init) + } + + fn open_internal(path: &Path, raw_key: Option<[u8; 32]>) -> Result { + if !path.exists() { + return Err(DbError::NotFound(path.display().to_string())); + } + + let key_ref = raw_key.as_ref(); + + // Open contact.db + let contact_path = path.join("contact").join("contact.db"); + if !contact_path.exists() { + return Err(DbError::NotFound(contact_path.display().to_string())); + } + let contact_conn = open_readonly_connection(&contact_path, key_ref)?; + + // Open session.db + let session_path = path.join("session").join("session.db"); + if !session_path.exists() { + return Err(DbError::NotFound(session_path.display().to_string())); + } + let session_conn = open_readonly_connection(&session_path, key_ref)?; + + // Scan message shards + let msg_dir = path.join("message"); + let mut shards = Vec::new(); + + if msg_dir.is_dir() { + let mut entries: Vec = std::fs::read_dir(&msg_dir)? + .filter_map(|e| e.ok()) + .map(|e| e.path()) + .filter(|p| is_numbered_message_shard(p)) + .collect(); + entries.sort(); + + for shard_path in entries { + let start_unix = read_shard_timestamp(&shard_path, key_ref); + shards.push(MessageShard { + path: shard_path, + start_unix, + end_unix: 0, // assigned below + }); + } + } + + // Sort shards by start_unix ASC + shards.sort_by_key(|s| s.start_unix); + + // Assign end_unix: each shard ends at next shard's start - 1; last = i64::MAX + let n = shards.len(); + for i in 0..n { + if i + 1 < n { + shards[i].end_unix = shards[i + 1].start_unix - 1; + } else { + shards[i].end_unix = i64::MAX; + } + } + + Ok(WechatDb { + contact_conn, + contact_path, + session_conn, + session_path, + shards, + message_fts_path: { + let p = msg_dir.join("message_fts.db"); + if p.exists() { + Some(p) + } else { + None + } + }, + contact_fts_path: { + let p = msg_dir.join("contact_fts.db"); + if p.exists() { + Some(p) + } else { + None + } + }, + pool: None, + raw_key, + label_cache: RwLock::new(None), + }) + } + + fn open_with_pool_internal( + path: impl AsRef, + raw_key: Option<[u8; 32]>, + fts_init: impl Fn(&Connection) -> Result<(), String> + Send + Sync + 'static, + ) -> Result { + let mut db = Self::open_internal(path.as_ref(), raw_key)?; + let fts_init_arc: Arc = Arc::new(fts_init); + let pool = ShardPool::open( + &db.shards, + db.message_fts_path.as_deref(), + Some(fts_init_arc), + raw_key, + )?; + db.pool = Some(pool); + Ok(db) + } + + /// Re-open the session.db connection to pick up external changes. + pub fn reopen_sessions(&mut self) -> Result<(), DbError> { + self.session_conn = open_connection(&self.session_path, self.raw_key.as_ref())?; + Ok(()) + } + + /// Re-open the contact.db connection to pick up external changes. + /// Also invalidates the label cache so it is reloaded on next query. + pub fn reopen_contacts(&mut self) -> Result<(), DbError> { + self.contact_conn = open_connection(&self.contact_path, self.raw_key.as_ref())?; + *self.label_cache.write().unwrap() = None; + Ok(()) + } + + /// Reopen a specific pooled shard connection. + /// Returns `Ok(true)` if the path was in the pool (and reopened), + /// `Ok(false)` if the path was not in the pool (unknown shard — possible topology change). + /// No-op if pool is not initialized (returns `Ok(false)`). + pub fn reopen_pooled_shard(&mut self, path: &Path) -> Result { + if let Some(pool) = &mut self.pool { + if pool.get(path).is_some() { + pool.reopen_shard(path)?; + return Ok(true); + } + return Ok(false); + } + Ok(false) + } + + /// Reopen all pooled connections (shards + FTS). + /// No-op if pool is not initialized. + pub fn reopen_all_pooled(&mut self) -> Result<(), DbError> { + if let Some(pool) = &mut self.pool { + pool.reopen_all()?; + } + Ok(()) + } + + /// Reopen only the FTS connection in the pool. + /// No-op if pool is not initialized. + pub fn reopen_fts(&mut self) -> Result<(), DbError> { + if let Some(pool) = &mut self.pool { + pool.reopen_fts()?; + } + Ok(()) + } + + /// Borrow the connection pool, if initialized. + pub fn pool(&self) -> Option<&ShardPool> { + self.pool.as_ref() + } + + /// Return shards whose time range overlaps `[start, end]`. + pub(crate) fn shards_for_range(&self, start: i64, end: i64) -> Vec<&MessageShard> { + self.shards + .iter() + .filter(|s| s.start_unix <= end && s.end_unix >= start) + .collect() + } + + /// Return all message shards (for full-shard scan in anchor queries). + pub(crate) fn all_shards(&self) -> &[MessageShard] { + &self.shards + } + + /// Build a `ShardMetadataFile` from the current shard metadata. + /// Callers can persist this to a sidecar file for future routing. + pub fn shard_metadata(&self) -> ShardMetadataFile { + let shards = self + .shards + .iter() + .filter_map(|s| { + let shard_id = extract_shard_id(&s.path)?; + Some(ShardMeta { + shard_id, + start_unix: s.start_unix, + end_unix: s.end_unix, + }) + }) + .collect(); + ShardMetadataFile { + shards, + written_at_ns: now_nanos(), + } + } + + /// Open a SQLite connection to a specific shard, optionally encrypted. + pub(crate) fn open_shard_with_key( + shard: &MessageShard, + raw_key: Option<&[u8; 32]>, + ) -> Result { + open_readonly_connection(&shard.path, raw_key) + } +} + +/// Extract the numeric shard ID from a path like `message_N.db`. +fn extract_shard_id(path: &Path) -> Option { + let stem = path.file_stem()?.to_str()?; + let suffix = stem.strip_prefix("message_")?; + suffix.parse().ok() +} + +fn is_numbered_message_shard(path: &Path) -> bool { + let Some(ext) = path.extension().and_then(|ext| ext.to_str()) else { + return false; + }; + if ext != "db" { + return false; + } + + let Some(stem) = path.file_stem().and_then(|stem| stem.to_str()) else { + return false; + }; + + let Some(suffix) = stem.strip_prefix("message_") else { + return false; + }; + + !suffix.is_empty() && suffix.bytes().all(|byte| byte.is_ascii_digit()) +} + +/// Try to read the timestamp from a message shard's Timestamp table. +/// Returns 0 if the table does not exist or is empty. +fn read_shard_timestamp(path: &Path, raw_key: Option<&[u8; 32]>) -> i64 { + let conn = match open_connection(path, raw_key) { + Ok(c) => c, + Err(_) => return 0, + }; + + conn.query_row("SELECT timestamp FROM Timestamp LIMIT 1", [], |row| { + row.get(0) + }) + .unwrap_or(0) +} + +#[cfg(test)] +mod tests { + use super::*; + use std::os::raw::c_void; + use tempfile::TempDir; + + /// Create an encrypted SQLite DB at `path` using SQLCipher's `sqlite3_key()`. + fn create_encrypted_db(path: &Path, raw_key: &[u8; 32], setup_sql: &str) { + let conn = Connection::open(path).unwrap(); + unsafe { + let rc = + rusqlite::ffi::sqlite3_key(conn.handle(), raw_key.as_ptr() as *const c_void, 32); + assert_eq!(rc, 0, "sqlite3_key failed during test DB creation"); + } + conn.execute_batch(setup_sql).unwrap(); + } + + /// Build a minimal encrypted db_storage directory for `open_encrypted` tests. + fn build_encrypted_db_storage(root: &Path, raw_key: &[u8; 32]) { + std::fs::create_dir_all(root.join("contact")).unwrap(); + std::fs::create_dir_all(root.join("session")).unwrap(); + std::fs::create_dir_all(root.join("message")).unwrap(); + + create_encrypted_db( + &root.join("contact").join("contact.db"), + raw_key, + "CREATE TABLE contact (username TEXT PRIMARY KEY, alias TEXT, remark TEXT, nick_name TEXT, description TEXT, extra_buffer BLOB);", + ); + create_encrypted_db( + &root.join("session").join("session.db"), + raw_key, + "CREATE TABLE SessionTable (username TEXT, sort_timestamp INTEGER, summary TEXT);", + ); + create_encrypted_db( + &root.join("message").join("message_0.db"), + raw_key, + "CREATE TABLE Timestamp (timestamp INTEGER); INSERT INTO Timestamp VALUES (1700000000);", + ); + } + + #[test] + fn open_encrypted_succeeds_with_correct_key() { + let tmp = TempDir::new().unwrap(); + let root = tmp.path().join("db_storage"); + let raw_key = [0xAB_u8; 32]; + build_encrypted_db_storage(&root, &raw_key); + + let db = WechatDb::open_encrypted(&root, raw_key).unwrap(); + assert_eq!(db.shards.len(), 1); + assert_eq!(db.shards[0].start_unix, 1700000000); + } + + #[test] + fn open_encrypted_fails_with_wrong_key() { + let tmp = TempDir::new().unwrap(); + let root = tmp.path().join("db_storage"); + let raw_key = [0xAB_u8; 32]; + build_encrypted_db_storage(&root, &raw_key); + + let wrong_key = [0xCD_u8; 32]; + let result = WechatDb::open_encrypted(&root, wrong_key); + assert!(result.is_err()); + } + + #[test] + fn open_encrypted_reopen_sessions_works() { + let tmp = TempDir::new().unwrap(); + let root = tmp.path().join("db_storage"); + let raw_key = [0xAB_u8; 32]; + build_encrypted_db_storage(&root, &raw_key); + + let mut db = WechatDb::open_encrypted(&root, raw_key).unwrap(); + // Reopen should succeed (re-applies sqlite3_key) + db.reopen_sessions().unwrap(); + db.reopen_contacts().unwrap(); + } + + #[test] + fn open_connection_plaintext_works() { + let tmp = TempDir::new().unwrap(); + let path = tmp.path().join("test.db"); + Connection::open(&path) + .unwrap() + .execute_batch("CREATE TABLE t (id INTEGER)") + .unwrap(); + + let conn = open_connection(&path, None).unwrap(); + let count: i64 = conn + .query_row("SELECT count(*) FROM t", [], |r| r.get(0)) + .unwrap(); + assert_eq!(count, 0); + } + + #[test] + fn open_connection_encrypted_works() { + let tmp = TempDir::new().unwrap(); + let path = tmp.path().join("enc.db"); + let raw_key = [0xAB_u8; 32]; + create_encrypted_db( + &path, + &raw_key, + "CREATE TABLE t (id INTEGER); INSERT INTO t VALUES (42);", + ); + + let conn = open_connection(&path, Some(&raw_key)).unwrap(); + let val: i64 = conn + .query_row("SELECT id FROM t", [], |r| r.get(0)) + .unwrap(); + assert_eq!(val, 42); + } +} diff --git a/crates/wx-db/src/pool.rs b/crates/wx-db/src/pool.rs new file mode 100644 index 0000000..99294a0 --- /dev/null +++ b/crates/wx-db/src/pool.rs @@ -0,0 +1,293 @@ +use std::collections::HashMap; +use std::path::{Path, PathBuf}; +use std::sync::Arc; + +use rusqlite::Connection; + +use crate::error::DbError; +use crate::open::MessageShard; + +pub(crate) type FtsInitFn = dyn Fn(&Connection) -> Result<(), String> + Send + Sync; + +/// Pre-opened connection pool for message shards and FTS. +/// +/// All connections are opened at construction time and held persistently. +/// Use [`ShardPool::reopen_all`] to close and reopen everything after +/// a background decrypt cycle. +pub struct ShardPool { + conns: HashMap, + fts_conn: Option, + fts_path: Option, + fts_init: Option>, + raw_key: Option<[u8; 32]>, +} + +impl std::fmt::Debug for ShardPool { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("ShardPool") + .field("shard_count", &self.conns.len()) + .field("has_fts", &self.fts_conn.is_some()) + .finish() + } +} + +impl ShardPool { + /// Open all shard connections and optionally an FTS connection. + /// + /// `fts_init` is called on the FTS connection after opening to register + /// custom tokenizers. The `Arc` is stored internally so that `reopen_fts()` + /// and `reopen_all()` can re-invoke it without the caller passing it again. + pub(crate) fn open( + shards: &[MessageShard], + fts_path: Option<&Path>, + fts_init: Option>, + raw_key: Option<[u8; 32]>, + ) -> Result { + let mut conns = HashMap::with_capacity(shards.len()); + for shard in shards { + let conn = crate::open::open_connection(&shard.path, raw_key.as_ref())?; + conns.insert(shard.path.clone(), conn); + } + + let fts_conn = match (fts_path, &fts_init) { + (Some(path), Some(init)) => { + let conn = crate::open::open_connection(path, raw_key.as_ref())?; + init(&conn).map_err(DbError::FtsInit)?; + Some(conn) + } + (Some(path), None) => { + let conn = crate::open::open_connection(path, raw_key.as_ref())?; + Some(conn) + } + _ => None, + }; + + Ok(ShardPool { + conns, + fts_conn, + fts_path: fts_path.map(|p| p.to_path_buf()), + fts_init, + raw_key, + }) + } + + /// Borrow a shard connection by path. + pub fn get(&self, path: &Path) -> Option<&Connection> { + self.conns.get(path) + } + + /// Close and reopen one shard connection. + pub fn reopen_shard(&mut self, path: &Path) -> Result<(), DbError> { + if self.conns.contains_key(path) { + let conn = crate::open::open_connection(path, self.raw_key.as_ref())?; + self.conns.insert(path.to_path_buf(), conn); + } + Ok(()) + } + + /// Borrow the FTS connection. + pub fn fts_conn(&self) -> Option<&Connection> { + self.fts_conn.as_ref() + } + + /// Close and reopen the FTS connection, re-registering the tokenizer. + pub fn reopen_fts(&mut self) -> Result<(), DbError> { + if let Some(path) = &self.fts_path { + let conn = crate::open::open_connection(path, self.raw_key.as_ref())?; + if let Some(init) = &self.fts_init { + init(&conn).map_err(DbError::FtsInit)?; + } + self.fts_conn = Some(conn); + } + Ok(()) + } + + /// Close and reopen all shard connections and FTS. + pub fn reopen_all(&mut self) -> Result<(), DbError> { + let paths: Vec = self.conns.keys().cloned().collect(); + for path in paths { + let conn = crate::open::open_connection(&path, self.raw_key.as_ref())?; + self.conns.insert(path, conn); + } + self.reopen_fts()?; + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use tempfile::TempDir; + + fn create_test_shard(dir: &Path, name: &str) -> MessageShard { + let path = dir.join(name); + Connection::open(&path) + .unwrap() + .execute_batch("CREATE TABLE test (id INTEGER PRIMARY KEY)") + .unwrap(); + MessageShard { + path, + start_unix: 0, + end_unix: i64::MAX, + } + } + + #[test] + fn open_and_get_connections() { + let tmp = TempDir::new().unwrap(); + let s1 = create_test_shard(tmp.path(), "message_0.db"); + let s2 = create_test_shard(tmp.path(), "message_1.db"); + let shards = vec![s1, s2]; + + let pool = ShardPool::open(&shards, None, None, None).unwrap(); + + assert!(pool.get(&shards[0].path).is_some()); + assert!(pool.get(&shards[1].path).is_some()); + assert!(pool.get(Path::new("/nonexistent")).is_none()); + assert!(pool.fts_conn().is_none()); + } + + #[test] + fn reopen_shard_picks_up_changes() { + let tmp = TempDir::new().unwrap(); + let shard = create_test_shard(tmp.path(), "message_0.db"); + let shards = vec![shard]; + + let mut pool = ShardPool::open(&shards, None, None, None).unwrap(); + + // Write new data outside the pool + { + let ext_conn = Connection::open(&shards[0].path).unwrap(); + ext_conn + .execute("INSERT INTO test (id) VALUES (42)", []) + .unwrap(); + } + + // Before reopen, read-only conn may or may not see it (WAL mode). + // After reopen, it must see it. + pool.reopen_shard(&shards[0].path).unwrap(); + let conn = pool.get(&shards[0].path).unwrap(); + let count: i64 = conn + .query_row("SELECT COUNT(*) FROM test", [], |r| r.get(0)) + .unwrap(); + assert_eq!(count, 1); + } + + #[test] + fn reopen_all_reopens_everything() { + let tmp = TempDir::new().unwrap(); + let s1 = create_test_shard(tmp.path(), "message_0.db"); + let s2 = create_test_shard(tmp.path(), "message_1.db"); + let shards = vec![s1, s2]; + + let mut pool = ShardPool::open(&shards, None, None, None).unwrap(); + pool.reopen_all().unwrap(); + + assert!(pool.get(&shards[0].path).is_some()); + assert!(pool.get(&shards[1].path).is_some()); + } + + #[test] + fn fts_connection_with_init() { + let tmp = TempDir::new().unwrap(); + let fts_path = tmp.path().join("message_fts.db"); + Connection::open(&fts_path) + .unwrap() + .execute_batch("CREATE TABLE fts_test (id INTEGER)") + .unwrap(); + + let init_called = Arc::new(std::sync::atomic::AtomicBool::new(false)); + let init_called_clone = Arc::clone(&init_called); + let fts_init: Arc = Arc::new(move |_conn| { + init_called_clone.store(true, std::sync::atomic::Ordering::SeqCst); + Ok(()) + }); + + let pool = ShardPool::open(&[], Some(&fts_path), Some(fts_init), None).unwrap(); + assert!(pool.fts_conn().is_some()); + assert!(init_called.load(std::sync::atomic::Ordering::SeqCst)); + } + + #[test] + fn reopen_fts_reinvokes_init() { + let tmp = TempDir::new().unwrap(); + let fts_path = tmp.path().join("message_fts.db"); + Connection::open(&fts_path).unwrap(); + + let call_count = Arc::new(std::sync::atomic::AtomicU32::new(0)); + let call_count_clone = Arc::clone(&call_count); + let fts_init: Arc = Arc::new(move |_conn| { + call_count_clone.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + Ok(()) + }); + + let mut pool = ShardPool::open(&[], Some(&fts_path), Some(fts_init), None).unwrap(); + assert_eq!(call_count.load(std::sync::atomic::Ordering::SeqCst), 1); + + pool.reopen_fts().unwrap(); + assert_eq!(call_count.load(std::sync::atomic::Ordering::SeqCst), 2); + } + + #[test] + fn reopen_fts_after_file_deleted() { + let tmp = TempDir::new().unwrap(); + let fts_path = tmp.path().join("message_fts.db"); + Connection::open(&fts_path).unwrap(); + + let fts_init: Arc = Arc::new(|_conn| Ok(())); + let mut pool = ShardPool::open(&[], Some(&fts_path), Some(fts_init), None).unwrap(); + assert!(pool.fts_conn().is_some()); + + // Delete the FTS file + std::fs::remove_file(&fts_path).unwrap(); + + // reopen_fts should fail because the file no longer exists + let result = pool.reopen_fts(); + assert!( + result.is_err(), + "reopen_fts should fail when file is deleted" + ); + } + + #[test] + fn reopen_fts_after_file_replaced() { + let tmp = TempDir::new().unwrap(); + let fts_path = tmp.path().join("message_fts.db"); + + // Create original FTS file with a marker table + { + let conn = Connection::open(&fts_path).unwrap(); + conn.execute_batch("CREATE TABLE marker (id INTEGER)") + .unwrap(); + } + + let fts_init: Arc = Arc::new(|_conn| Ok(())); + let mut pool = ShardPool::open(&[], Some(&fts_path), Some(fts_init), None).unwrap(); + assert!(pool.fts_conn().is_some()); + + // Replace the file with a new one containing a different table + std::fs::remove_file(&fts_path).unwrap(); + { + let conn = Connection::open(&fts_path).unwrap(); + conn.execute_batch("CREATE TABLE replaced_marker (id INTEGER)") + .unwrap(); + } + + // reopen_fts should succeed and connect to the new file + pool.reopen_fts().unwrap(); + let conn = pool.fts_conn().unwrap(); + + // Verify we see the replaced_marker table (proving we connected to the new file) + let has_replaced: bool = conn + .query_row( + "SELECT COUNT(*) > 0 FROM sqlite_master WHERE type='table' AND name='replaced_marker'", + [], + |row| row.get(0), + ) + .unwrap(); + assert!( + has_replaced, + "reopened connection should see replaced_marker table" + ); + } +} diff --git a/crates/wx-db/src/sessions.rs b/crates/wx-db/src/sessions.rs new file mode 100644 index 0000000..0becaf6 --- /dev/null +++ b/crates/wx-db/src/sessions.rs @@ -0,0 +1,96 @@ +use rusqlite::types::ValueRef; + +use crate::decode::decode_content; +use crate::error::DbError; +use crate::model::{effective_limit, QueryResult, QueryStats, Session, SessionQuery}; +use crate::open::WechatDb; + +/// Strip the `"wxid_xxx:\n"` sender prefix from group chat summaries. +/// +/// Only strips when `username` ends with `@chatroom`. For non-group sessions, +/// the summary is returned unchanged. +fn strip_group_summary_prefix(username: &str, summary: String) -> String { + if !crate::model::is_group_chat(username) { + return summary; + } + if let Some(newline_pos) = summary.find('\n') { + let prefix = &summary[..newline_pos]; + // Prefix should end with ':' and contain no spaces (looks like "wxid_xxx:") + if prefix.ends_with(':') && !prefix.contains(' ') { + return summary[newline_pos + 1..].to_string(); + } + } + summary +} + +impl WechatDb { + /// Query recent sessions (conversations). + pub fn query_sessions(&self, query: &SessionQuery) -> Result, DbError> { + let limit = effective_limit(query.limit); + + // Count total rows before LIMIT/OFFSET + let total_rows: usize = + self.session_conn + .query_row("SELECT COUNT(*) FROM SessionTable", [], |row| { + row.get::<_, i64>(0) + })? as usize; + + let sql = format!( + "SELECT username, sort_timestamp, summary, \ + last_msg_type, last_msg_sender, last_sender_display_name \ + FROM SessionTable \ + ORDER BY sort_timestamp {order}, username ASC \ + LIMIT ?1 OFFSET ?2", + order = query.order.sql_keyword(), + ); + + let mut stmt = self.session_conn.prepare(&sql)?; + let mut rows = stmt.query([limit as i64, query.offset as i64])?; + + let mut items = Vec::new(); + let mut skipped: usize = 0; + + while let Some(row) = rows.next()? { + let username: String = row.get(0)?; + let sort_timestamp: i64 = row.get(1)?; + + // summary can be Text or Blob (zstd-compressed) + let summary = match row.get_ref(2)? { + ValueRef::Text(bytes) => String::from_utf8_lossy(bytes).into_owned(), + ValueRef::Blob(bytes) => match decode_content(bytes, None) { + Ok(s) => s, + Err(_) => { + skipped += 1; + continue; + } + }, + ValueRef::Null => String::new(), + _ => String::new(), + }; + + let last_msg_type = row.get::<_, Option>(3)?.map(|v| v as u32); + let last_msg_sender: Option = row.get(4)?; + let last_sender_display_name: Option = row.get(5)?; + + let summary = strip_group_summary_prefix(&username, summary); + + items.push(Session { + username, + summary, + sort_timestamp, + last_msg_type, + last_msg_sender, + last_sender_display_name, + }); + } + + Ok(QueryResult { + items, + stats: QueryStats { + total_rows, + filtered_count: None, + skipped, + }, + }) + } +} diff --git a/crates/wx-db/src/shard_metadata.rs b/crates/wx-db/src/shard_metadata.rs new file mode 100644 index 0000000..b3a3f03 --- /dev/null +++ b/crates/wx-db/src/shard_metadata.rs @@ -0,0 +1,104 @@ +//! Shard metadata types and read/write helpers for persistent shard time-range cache. +//! +//! The sidecar file `shard-metadata.json` is written alongside decrypted message +//! shard DBs. It records each shard's time range so that callers can route queries +//! to a minimal subset of shards without opening every DB. + +use std::path::Path; +use std::time::{SystemTime, UNIX_EPOCH}; + +use serde::{Deserialize, Serialize}; + +/// Time-range metadata for a single message shard. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ShardMeta { + pub shard_id: u32, + pub start_unix: i64, + pub end_unix: i64, +} + +/// Persistent shard metadata file contents. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ShardMetadataFile { + pub shards: Vec, + /// Unix nanoseconds when this metadata was written. + pub written_at_ns: u128, +} + +const SIDECAR_FILENAME: &str = "shard-metadata.json"; + +/// Read shard metadata from `dir/shard-metadata.json`. +/// Returns `None` on any error (missing file, corrupt JSON, IO error). +pub fn read_shard_metadata(dir: &Path) -> Option { + let path = dir.join(SIDECAR_FILENAME); + let data = std::fs::read_to_string(&path).ok()?; + serde_json::from_str(&data).ok() +} + +/// Atomically write shard metadata to `dir/shard-metadata.json`. +/// Writes to a `.tmp` file first, then renames. +pub fn write_shard_metadata(dir: &Path, meta: &ShardMetadataFile) -> Result<(), std::io::Error> { + let path = dir.join(SIDECAR_FILENAME); + let tmp_path = dir.join(format!("{SIDECAR_FILENAME}.tmp")); + let data = serde_json::to_string_pretty(meta).map_err(std::io::Error::other)?; + std::fs::write(&tmp_path, data)?; + std::fs::rename(&tmp_path, &path)?; + Ok(()) +} + +/// Build a `ShardMetadataFile` timestamp (current time in nanoseconds since epoch). +pub fn now_nanos() -> u128 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_nanos() +} + +#[cfg(test)] +mod tests { + use super::*; + use tempfile::TempDir; + + #[test] + fn roundtrip_write_read() { + let dir = TempDir::new().unwrap(); + let meta = ShardMetadataFile { + shards: vec![ + ShardMeta { + shard_id: 0, + start_unix: 1000, + end_unix: 2000, + }, + ShardMeta { + shard_id: 1, + start_unix: 2001, + end_unix: 3000, + }, + ], + written_at_ns: 123456789, + }; + + write_shard_metadata(dir.path(), &meta).unwrap(); + let loaded = read_shard_metadata(dir.path()).unwrap(); + + assert_eq!(loaded.shards.len(), 2); + assert_eq!(loaded.shards[0].shard_id, 0); + assert_eq!(loaded.shards[0].start_unix, 1000); + assert_eq!(loaded.shards[1].shard_id, 1); + assert_eq!(loaded.shards[1].end_unix, 3000); + assert_eq!(loaded.written_at_ns, 123456789); + } + + #[test] + fn read_nonexistent_returns_none() { + let dir = TempDir::new().unwrap(); + assert!(read_shard_metadata(dir.path()).is_none()); + } + + #[test] + fn read_corrupt_json_returns_none() { + let dir = TempDir::new().unwrap(); + std::fs::write(dir.path().join("shard-metadata.json"), "not valid json").unwrap(); + assert!(read_shard_metadata(dir.path()).is_none()); + } +} diff --git a/crates/wx-db/src/xml_extract.rs b/crates/wx-db/src/xml_extract.rs new file mode 100644 index 0000000..65a4e23 --- /dev/null +++ b/crates/wx-db/src/xml_extract.rs @@ -0,0 +1,570 @@ +//! XML field extraction for app messages (type=49) and system messages (type=10000). +//! +//! Pure string-based extraction using `str::find` — no regex or XML parser dependencies. +//! All functions are `pub(crate)` and never panic on malformed input. + +use crate::model::*; + +// --- Helper functions --- + +/// Extract the text content of a single XML tag from the given string. +/// +/// Handles CDATA sections and HTML entity unescaping. Returns `None` if: +/// - The tag is not found +/// - The tag content is empty +pub(crate) fn extract_tag_text(xml: &str, tag: &str) -> Option { + let open = format!("<{}>", tag); + let close = format!("", tag); + + let start_idx = xml.find(&open)?; + let content_start = start_idx + open.len(); + let end_idx = xml[content_start..].find(&close)?; + let raw = &xml[content_start..content_start + end_idx]; + + // Strip CDATA wrapper if present + let text = if let Some(inner) = raw + .strip_prefix("")) + { + inner + } else { + raw + }; + + if text.is_empty() { + return None; + } + + Some(unescape_xml_entities(text)) +} + +/// Unescape the five standard XML entities. +fn unescape_xml_entities(s: &str) -> String { + if !s.contains('&') { + return s.to_string(); + } + s.replace("&", "&") + .replace("<", "<") + .replace(">", ">") + .replace(""", "\"") + .replace("'", "'") +} + +/// Filter out the literal string "null" used by WeChat as a sentinel for empty values. +/// Apply only to display fields (title, des) where "null" is never a valid value. +fn filter_null_sentinel(value: Option) -> Option { + value.filter(|s| s != "null") +} + +// --- Extraction structs --- + +pub(crate) struct QuoteFields { + pub reply_text: Option, + pub refer_sender: Option, + pub refer_display_name: Option, + pub refer_content: Option, + pub refer_type: Option, +} + +pub(crate) struct TransferFields { + pub amount_desc: Option, + pub pay_memo: Option, + pub pay_sub_type: Option, +} + +pub(crate) struct FileFields { + pub title: Option, + pub file_ext: Option, + pub file_size: Option, + pub md5: Option, +} + +// --- Extraction functions --- + +/// Extract the wxid of the quoted message's original sender from raw_xml. +/// +/// In group chats, `` contains the chatroom ID (not the sender wxid), +/// while `` contains the actual sender's wxid. This function prefers +/// `` when present, falling back to `` for private chats. +/// +/// This is needed because `MessageContent::Quote.refer_sender` stores the display name +/// (merged from `refer_display_name.or(refer_sender)`), not the wxid. +pub fn extract_quote_fromusr(raw_xml: &str) -> Option { + let refermsg_block = extract_inner_block(raw_xml, "refermsg")?; + extract_tag_text(refermsg_block, "chatusr") + .or_else(|| extract_tag_text(refermsg_block, "fromusr")) +} + +/// Extract common app fields: title, des, url. +pub(crate) fn extract_app_fields(xml: &str) -> (Option, Option, Option) { + ( + filter_null_sentinel(extract_tag_text(xml, "title")), + filter_null_sentinel(extract_tag_text(xml, "des")), + extract_tag_text(xml, "url"), + ) +} + +/// Extract fields for quote/reply messages (sub_type=57). +pub(crate) fn extract_quote_fields(xml: &str) -> QuoteFields { + let reply_text = extract_tag_text(xml, "title"); + + // Extract refermsg block, then extract fields within it + let refermsg_block = extract_inner_block(xml, "refermsg"); + let (refer_sender, refer_display_name, refer_content, refer_type) = + if let Some(block) = refermsg_block { + let sender = extract_tag_text(block, "fromusr"); + let display = extract_tag_text(block, "displayname"); + let content = extract_tag_text(block, "content"); + let rtype = extract_tag_text(block, "type").and_then(|s| s.parse::().ok()); + + // If the referred message contains raw XML, convert to readable placeholder + let content = match (&content, rtype) { + (Some(xml), Some(49)) if xml.contains(" { + extract_tag_text(xml, "title").or(content) + } + (Some(xml), Some(3)) if xml.contains(" Some("[图片]".to_string()), + (Some(xml), Some(43)) if xml.contains(" Some("[视频]".to_string()), + (Some(xml), Some(47)) if xml.contains(" Some("[表情]".to_string()), + _ => content, + }; + + (sender, display, content, rtype) + } else { + (None, None, None, None) + }; + + QuoteFields { + reply_text, + refer_sender, + refer_display_name, + refer_content, + refer_type, + } +} + +/// Extract fields for transfer messages (sub_type=2000). +pub(crate) fn extract_transfer_fields(xml: &str) -> TransferFields { + let pay_block = extract_inner_block(xml, "wcpayinfo"); + if let Some(block) = pay_block { + TransferFields { + amount_desc: extract_tag_text(block, "feedesc"), + pay_memo: extract_tag_text(block, "pay_memo"), + pay_sub_type: extract_tag_text(block, "paysubtype").and_then(|s| s.parse().ok()), + } + } else { + TransferFields { + amount_desc: None, + pay_memo: None, + pay_sub_type: None, + } + } +} + +/// Extract fields for file messages (sub_type=6). +pub(crate) fn extract_file_fields(xml: &str) -> FileFields { + FileFields { + title: extract_tag_text(xml, "title"), + file_ext: extract_tag_text(xml, "fileext"), + file_size: extract_tag_text(xml, "totallen").and_then(|s| s.parse().ok()), + md5: extract_tag_text(xml, "md5"), + } +} + +/// Extract the inner content (including nested tags) of a block element. +fn extract_inner_block<'a>(xml: &'a str, tag: &str) -> Option<&'a str> { + let open = format!("<{}>", tag); + let close = format!("", tag); + let start = xml.find(&open)?; + let inner_start = start + open.len(); + let end = xml[inner_start..].find(&close)?; + Some(&xml[inner_start..inner_start + end]) +} + +// --- Dispatch --- + +/// Parse an app message (type=49) XML into a typed MessageContent variant. +pub(crate) fn dispatch_app_message(sub_type: u32, xml: &str) -> MessageContent { + match sub_type { + // Link (5) and link-like types (4, 7, 92 for music) + APP_SUB_TYPE_LINK | 4 | 7 => { + let (title, des, url) = extract_app_fields(xml); + MessageContent::Link { + sub_type, + title, + des, + url, + raw_xml: xml.to_string(), + } + } + APP_SUB_TYPE_MUSIC => { + let (title, des, url) = extract_app_fields(xml); + MessageContent::Link { + sub_type, + title, + des, + url, + raw_xml: xml.to_string(), + } + } + APP_SUB_TYPE_FILE => { + let f = extract_file_fields(xml); + MessageContent::File { + title: f.title, + file_ext: f.file_ext, + file_size: f.file_size, + md5: f.md5, + raw_xml: xml.to_string(), + } + } + APP_SUB_TYPE_MINI_PROGRAM | APP_SUB_TYPE_MINI_PROGRAM_2 => { + let (title, _, url) = extract_app_fields(xml); + MessageContent::MiniProgram { + sub_type, + title, + url, + raw_xml: xml.to_string(), + } + } + APP_SUB_TYPE_MERGED => { + let title = filter_null_sentinel(extract_tag_text(xml, "title")); + MessageContent::MergedMessages { + title, + raw_xml: xml.to_string(), + } + } + APP_SUB_TYPE_QUOTE => { + let q = extract_quote_fields(xml); + MessageContent::Quote { + reply_text: q.reply_text, + refer_sender: q.refer_display_name.or(q.refer_sender), + refer_content: q.refer_content, + refer_type: q.refer_type, + raw_xml: xml.to_string(), + } + } + APP_SUB_TYPE_TRANSFER => { + let t = extract_transfer_fields(xml); + MessageContent::Transfer { + amount_desc: t.amount_desc, + pay_memo: t.pay_memo, + pay_sub_type: t.pay_sub_type, + raw_xml: xml.to_string(), + } + } + APP_SUB_TYPE_RED_ENVELOPE | 2003 => { + let title = filter_null_sentinel(extract_tag_text(xml, "title")); + MessageContent::RedEnvelope { + title, + raw_xml: xml.to_string(), + } + } + APP_SUB_TYPE_CHANNEL | APP_SUB_TYPE_CHANNEL_LIVE => { + let title = filter_null_sentinel(extract_tag_text(xml, "title")) + .or_else(|| filter_null_sentinel(extract_tag_text(xml, "des"))); + MessageContent::ChannelVideo { + sub_type, + title, + raw_xml: xml.to_string(), + } + } + APP_SUB_TYPE_PAT => MessageContent::Pat { + raw_xml: xml.to_string(), + }, + _ => { + let (title, des, url) = extract_app_fields(xml); + MessageContent::AppGeneric { + sub_type, + title, + des, + url, + raw_xml: xml.to_string(), + } + } + } +} + +// --- System message extraction --- + +/// Try to extract readable text from a system message (type=10000) that may +/// contain `` XML. +/// +/// Returns `Some(readable_text)` if the content is a revokemsg sysmsg XML +/// and the `` tag inside `` contains text. +/// Returns `None` if the content is not sysmsg XML or extraction fails, +/// in which case the caller should use the original content as-is. +pub(crate) fn extract_system_message_text(content: &str) -> Option { + // Quick check: only process XML-like system messages + if !content.contains(" from within the block + if content.contains("type=\"revokemsg\"") || content.contains("type='revokemsg'") { + if let Some(block) = extract_inner_block(content, "revokemsg") { + return extract_tag_text(block, "content"); + } + } + + None +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn extract_tag_text_plain() { + assert_eq!( + extract_tag_text("Hello World", "title"), + Some("Hello World".to_string()) + ); + } + + #[test] + fn extract_tag_text_cdata() { + assert_eq!( + extract_tag_text("<![CDATA[Hello]]>", "title"), + Some("Hello".to_string()) + ); + } + + #[test] + fn extract_tag_text_entities() { + assert_eq!( + extract_tag_text("A&B<C>D", "title"), + Some("A&BD".to_string()) + ); + } + + #[test] + fn extract_tag_text_empty() { + assert_eq!(extract_tag_text("", "title"), None); + } + + #[test] + fn extract_tag_text_missing() { + assert_eq!(extract_tag_text("no title here", "title"), None); + } + + #[test] + fn extract_tag_text_cdata_with_entities() { + // CDATA content should NOT be entity-unescaped (it's literal), + // but we still run unescape for simplicity — in practice CDATA + // won't contain XML entities. + assert_eq!( + extract_tag_text("", "des"), + Some("A&B".to_string()) + ); + } + + #[test] + fn extract_app_fields_full() { + let xml = r#"Test TitleDescriptionhttps://example.com"#; + let (title, des, url) = extract_app_fields(xml); + assert_eq!(title, Some("Test Title".to_string())); + assert_eq!(des, Some("Description".to_string())); + assert_eq!(url, Some("https://example.com".to_string())); + } + + #[test] + fn extract_app_fields_filters_null_title() { + let xml = r#"nullnullhttps://example.com"#; + let (title, des, url) = extract_app_fields(xml); + assert_eq!(title, None); + assert_eq!(des, None); + assert_eq!(url, Some("https://example.com".to_string())); + } + + #[test] + fn extract_app_fields_preserves_non_null() { + let xml = r#"nullable"#; + let (title, _, _) = extract_app_fields(xml); + assert_eq!(title, Some("nullable".to_string())); + } + + #[test] + fn extract_tag_text_returns_null_literal() { + // extract_tag_text itself should NOT filter "null" — that's filter_null_sentinel's job + assert_eq!( + extract_tag_text("null", "title"), + Some("null".to_string()) + ); + } + + #[test] + fn extract_quote_fields_full() { + let xml = r#"reply textwxid_aliceAliceoriginal message1"#; + let q = extract_quote_fields(xml); + assert_eq!(q.reply_text, Some("reply text".to_string())); + assert_eq!(q.refer_sender, Some("wxid_alice".to_string())); + assert_eq!(q.refer_display_name, Some("Alice".to_string())); + assert_eq!(q.refer_content, Some("original message".to_string())); + assert_eq!(q.refer_type, Some(1)); + } + + #[test] + fn extract_quote_fields_nested_app() { + let xml = r#"reply textwxid_bobBob49被引原文desc"#; + let q = extract_quote_fields(xml); + assert_eq!(q.refer_content, Some("被引原文".to_string())); + assert_eq!(q.refer_type, Some(49)); + } + + #[test] + fn extract_quote_fields_nested_text() { + let xml = r#"replywxid_alice1普通文本"#; + let q = extract_quote_fields(xml); + assert_eq!(q.refer_content, Some("普通文本".to_string())); + assert_eq!(q.refer_type, Some(1)); + } + + #[test] + fn extract_transfer_fields_full() { + let xml = r#"¥66.00lunch1"#; + let t = extract_transfer_fields(xml); + assert_eq!(t.amount_desc, Some("¥66.00".to_string())); + assert_eq!(t.pay_memo, Some("lunch".to_string())); + assert_eq!(t.pay_sub_type, Some(1)); + } + + #[test] + fn extract_file_fields_full() { + let xml = r#"report.pdfpdf12345abc123"#; + let f = extract_file_fields(xml); + assert_eq!(f.title, Some("report.pdf".to_string())); + assert_eq!(f.file_ext, Some("pdf".to_string())); + assert_eq!(f.file_size, Some(12345)); + assert_eq!(f.md5, Some("abc123".to_string())); + } + + #[test] + fn dispatch_link_message() { + let xml = r#"ArticleDeschttps://mp.weixin.qq.com"#; + match dispatch_app_message(5, xml) { + MessageContent::Link { + title, des, url, .. + } => { + assert_eq!(title, Some("Article".to_string())); + assert_eq!(des, Some("Desc".to_string())); + assert_eq!(url, Some("https://mp.weixin.qq.com".to_string())); + } + other => panic!("expected Link, got: {:?}", other), + } + } + + #[test] + fn dispatch_channel_video_fallback_des() { + let xml = r#"今日新闻"#; + match dispatch_app_message(APP_SUB_TYPE_CHANNEL, xml) { + MessageContent::ChannelVideo { title, .. } => { + assert_eq!(title, Some("今日新闻".to_string())); + } + other => panic!("expected ChannelVideo, got: {:?}", other), + } + } + + #[test] + fn dispatch_channel_video_title_preferred_over_des() { + let xml = r#"视频标题描述"#; + match dispatch_app_message(APP_SUB_TYPE_CHANNEL, xml) { + MessageContent::ChannelVideo { title, .. } => { + assert_eq!(title, Some("视频标题".to_string())); + } + other => panic!("expected ChannelVideo, got: {:?}", other), + } + } + + #[test] + fn dispatch_unknown_sub_type() { + let xml = "Unknown"; + match dispatch_app_message(9999, xml) { + MessageContent::AppGeneric { sub_type, .. } => { + assert_eq!(sub_type, 9999); + } + other => panic!("expected AppGeneric, got: {:?}", other), + } + } + + #[test] + fn extract_quote_refer_image() { + let xml = r#"replywxid_alice3<?xml version="1.0"?><msg><img aeskey="abc" /></msg>"#; + let q = extract_quote_fields(xml); + assert_eq!(q.refer_content, Some("[图片]".to_string())); + assert_eq!(q.refer_type, Some(3)); + } + + #[test] + fn extract_quote_refer_video() { + let xml = r#"replywxid_bob43<msg><videomsg length="30" /></msg>"#; + let q = extract_quote_fields(xml); + assert_eq!(q.refer_content, Some("[视频]".to_string())); + assert_eq!(q.refer_type, Some(43)); + } + + #[test] + fn extract_quote_refer_emoji() { + let xml = r#"replywxid_carol47<msg><emoji md5="abc123" /></msg>"#; + let q = extract_quote_fields(xml); + assert_eq!(q.refer_content, Some("[表情]".to_string())); + assert_eq!(q.refer_type, Some(47)); + } + + #[test] + fn extract_quote_refer_image_plain_text_preserved() { + let xml = r#"replywxid_alice3just plain text"#; + let q = extract_quote_fields(xml); + assert_eq!(q.refer_content, Some("just plain text".to_string())); + } + + #[test] + fn extract_quote_refer_image_with_wxid_prefix() { + // In group chats, content may be prefixed with "wxid_xxx:\n" + let xml = r#"replywxid_alice3wxid_someone: +<?xml version="1.0"?><msg><img aeskey="def" /></msg>"#; + let q = extract_quote_fields(xml); + assert_eq!(q.refer_content, Some("[图片]".to_string())); + } + + // --- extract_quote_fromusr tests --- + + #[test] + fn extract_quote_fromusr_normal() { + let xml = r#"replywxid_alicehi"#; + assert_eq!( + extract_quote_fromusr(xml), + Some("wxid_alice".to_string()) + ); + } + + #[test] + fn extract_quote_fromusr_missing_refermsg() { + let xml = r#"reply"#; + assert_eq!(extract_quote_fromusr(xml), None); + } + + #[test] + fn extract_quote_fromusr_missing_fromusr() { + let xml = r#"replyhi"#; + assert_eq!(extract_quote_fromusr(xml), None); + } + + #[test] + fn extract_quote_fromusr_prefers_chatusr_in_group() { + // In group chats, is the chatroom ID, is the actual sender + let xml = r#"replygroup@chatroomwxid_senderSenderhi"#; + assert_eq!( + extract_quote_fromusr(xml), + Some("wxid_sender".to_string()) + ); + } + + #[test] + fn extract_quote_fromusr_falls_back_to_fromusr_without_chatusr() { + // In private chats, only exists (no ) + let xml = r#"replywxid_bobBobhi"#; + assert_eq!( + extract_quote_fromusr(xml), + Some("wxid_bob".to_string()) + ); + } +} diff --git a/crates/wx-db/tests/anchor-queries.rs b/crates/wx-db/tests/anchor-queries.rs new file mode 100644 index 0000000..db5045d --- /dev/null +++ b/crates/wx-db/tests/anchor-queries.rs @@ -0,0 +1,418 @@ +use std::fs; +use std::path::Path; + +use rusqlite::{params, Connection}; +use tempfile::TempDir; +use wx_db::test_ddl; +use wx_db::{MessageQuery, WechatDb}; + +// Msg_29a6db07e8bbdb53f5d54cc3c309f3f1 = md5("wxid_alice") +const ALICE_TABLE: &str = "Msg_29a6db07e8bbdb53f5d54cc3c309f3f1"; + +/// Create a fixture with 2 shards containing messages for wxid_alice. +/// Shard 0 (t=1700000000): sort_seq 100..500, server_id 1001..1005 +/// Shard 1 (t=1710000000): sort_seq 600..800, server_id 1006..1008 +/// +/// Also includes a sort_seq collision: two messages at sort_seq=300 +/// with different create_time and server_id. +fn create_anchor_fixture() -> TempDir { + let dir = TempDir::new().unwrap(); + let base = dir.path(); + + // contact/contact.db + let contact_dir = base.join("contact"); + fs::create_dir_all(&contact_dir).unwrap(); + create_contact_db(&contact_dir.join("contact.db")); + + // session/session.db + let session_dir = base.join("session"); + fs::create_dir_all(&session_dir).unwrap(); + create_session_db(&session_dir.join("session.db")); + + // message/message_0.db (shard 0) + let msg_dir = base.join("message"); + fs::create_dir_all(&msg_dir).unwrap(); + create_shard_0(&msg_dir.join("message_0.db")); + + // message/message_1.db (shard 1) + create_shard_1(&msg_dir.join("message_1.db")); + + dir +} + +fn create_contact_db(path: &Path) { + let conn = Connection::open(path).unwrap(); + test_ddl::create_test_contact_table(&conn); +} + +fn create_session_db(path: &Path) { + let conn = Connection::open(path).unwrap(); + test_ddl::create_test_session_table(&conn); +} + +fn insert_msg( + conn: &Connection, + table: &str, + sort_seq: i64, + server_id: i64, + create_time: i64, + text: &str, +) { + conn.execute( + &format!("INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8)"), + params![ + sort_seq, + server_id, + 1_u32, // local_type = MSG_TYPE_TEXT + 1_i32, // real_sender_id + create_time, + text.as_bytes(), + None::>, + 0_i32, + ], + ) + .unwrap(); +} + +/// Shard 0: timestamp=1700000000 +/// Messages (sort_seq, server_id, create_time, text): +/// (100, 1001, 1700000100, "msg-1") +/// (200, 1002, 1700000200, "msg-2") +/// (300, 1003, 1700000300, "msg-3a") ← sort_seq collision +/// (300, 1004, 1700000301, "msg-3b") ← sort_seq collision (different create_time) +/// (500, 1005, 1700000500, "msg-5") +fn create_shard_0(path: &Path) { + let conn = Connection::open(path).unwrap(); + conn.execute_batch("CREATE TABLE Timestamp (timestamp INTEGER);") + .unwrap(); + conn.execute( + "INSERT INTO Timestamp VALUES (?1)", + params![1_700_000_000_i64], + ) + .unwrap(); + conn.execute_batch("CREATE TABLE Name2Id (rowid INTEGER PRIMARY KEY, user_name TEXT);") + .unwrap(); + conn.execute( + "INSERT INTO Name2Id VALUES (?1, ?2)", + params![1, "wxid_alice"], + ) + .unwrap(); + + conn.execute_batch(&format!( + "CREATE TABLE [{table}] ( + sort_seq INTEGER, + server_id INTEGER, + local_type INTEGER, + real_sender_id INTEGER, + create_time INTEGER, + message_content BLOB, + packed_info_data BLOB, + status INTEGER + );", + table = ALICE_TABLE + )) + .unwrap(); + + insert_msg(&conn, ALICE_TABLE, 100, 1001, 1_700_000_100, "msg-1"); + insert_msg(&conn, ALICE_TABLE, 200, 1002, 1_700_000_200, "msg-2"); + insert_msg(&conn, ALICE_TABLE, 300, 1003, 1_700_000_300, "msg-3a"); + insert_msg(&conn, ALICE_TABLE, 300, 1004, 1_700_000_301, "msg-3b"); + insert_msg(&conn, ALICE_TABLE, 500, 1005, 1_700_000_500, "msg-5"); +} + +/// Shard 1: timestamp=1710000000 +/// Messages (sort_seq, server_id, create_time, text): +/// (600, 1006, 1710000100, "msg-6") +/// (700, 1007, 1710000200, "msg-7") +/// (800, 1008, 1710000300, "msg-8") +fn create_shard_1(path: &Path) { + let conn = Connection::open(path).unwrap(); + conn.execute_batch("CREATE TABLE Timestamp (timestamp INTEGER);") + .unwrap(); + conn.execute( + "INSERT INTO Timestamp VALUES (?1)", + params![1_710_000_000_i64], + ) + .unwrap(); + conn.execute_batch("CREATE TABLE Name2Id (rowid INTEGER PRIMARY KEY, user_name TEXT);") + .unwrap(); + conn.execute( + "INSERT INTO Name2Id VALUES (?1, ?2)", + params![1, "wxid_alice"], + ) + .unwrap(); + + conn.execute_batch(&format!( + "CREATE TABLE [{table}] ( + sort_seq INTEGER, + server_id INTEGER, + local_type INTEGER, + real_sender_id INTEGER, + create_time INTEGER, + message_content BLOB, + packed_info_data BLOB, + status INTEGER + );", + table = ALICE_TABLE + )) + .unwrap(); + + insert_msg(&conn, ALICE_TABLE, 600, 1006, 1_710_000_100, "msg-6"); + insert_msg(&conn, ALICE_TABLE, 700, 1007, 1_710_000_200, "msg-7"); + insert_msg(&conn, ALICE_TABLE, 800, 1008, 1_710_000_300, "msg-8"); +} + +// --------------------------------------------------------------------------- +// AfterSortSeq tests +// --------------------------------------------------------------------------- + +#[test] +fn after_sort_seq_returns_messages_after_pivot() { + let dir = create_anchor_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let query = MessageQuery::for_talker("wxid_alice") + .after_sort_seq(300) + .limit(100); + let result = db.query_messages_anchor(&query).unwrap(); + + let seqs: Vec = result.items.iter().map(|m| m.sort_seq).collect(); + // Should return sort_seq > 300: 500, 600, 700, 800 + assert_eq!(seqs, vec![500, 600, 700, 800]); +} + +#[test] +fn after_sort_seq_respects_limit() { + let dir = create_anchor_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let query = MessageQuery::for_talker("wxid_alice") + .after_sort_seq(100) + .limit(2); + let result = db.query_messages_anchor(&query).unwrap(); + + assert_eq!(result.items.len(), 2); + assert_eq!(result.items[0].sort_seq, 200); + assert_eq!(result.items[1].sort_seq, 300); +} + +#[test] +fn after_sort_seq_applies_filter_before_limit() { + let dir = create_anchor_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // Filter for keyword "msg-5" — only sort_seq=500 matches + let query = MessageQuery::for_talker("wxid_alice") + .after_sort_seq(100) + .keyword("msg-5") + .limit(100); + let result = db.query_messages_anchor(&query).unwrap(); + + assert_eq!(result.items.len(), 1); + assert_eq!(result.items[0].sort_seq, 500); +} + +#[test] +fn pooled_after_sort_seq_works_after_shard_file_removed() { + let dir = create_anchor_fixture(); + let db = WechatDb::open_with_pool(dir.path(), |_conn| Ok(())).unwrap(); + + let shard_path = dir.path().join("message").join("message_1.db"); + fs::remove_file(&shard_path).unwrap(); + + let query = MessageQuery::for_talker("wxid_alice") + .after_sort_seq(300) + .limit(100); + let result = db.query_messages_anchor(&query).unwrap(); + + let seqs: Vec = result.items.iter().map(|m| m.sort_seq).collect(); + assert_eq!(seqs, vec![500, 600, 700, 800]); +} + +// --------------------------------------------------------------------------- +// AroundSortSeq tests +// --------------------------------------------------------------------------- + +#[test] +fn around_sort_seq_basic_context() { + let dir = create_anchor_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let query = MessageQuery::for_talker("wxid_alice") + .around_sort_seq(500) + .context(2); + let result = db.query_messages_anchor(&query).unwrap(); + + let seqs: Vec = result.items.iter().map(|m| m.sort_seq).collect(); + // before(2): [300, 300], pivot: [500], after(2): [600, 700] + assert_eq!(seqs, vec![300, 300, 500, 600, 700]); +} + +#[test] +fn around_sort_seq_collision_pivot_bucket() { + let dir = create_anchor_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // Pivot at sort_seq=300 which has 2 messages (collision) + let query = MessageQuery::for_talker("wxid_alice") + .around_sort_seq(300) + .context(1); + let result = db.query_messages_anchor(&query).unwrap(); + + let seqs: Vec = result.items.iter().map(|m| m.sort_seq).collect(); + // before(1): [200], pivot bucket: [300, 300], after(1): [500] + assert_eq!(seqs, vec![200, 300, 300, 500]); + + // Verify pivot bucket is ordered by (create_time, server_id) + let pivot_msgs: Vec<_> = result.items.iter().filter(|m| m.sort_seq == 300).collect(); + assert_eq!(pivot_msgs.len(), 2); + assert_eq!(pivot_msgs[0].server_id, 1003); + assert_eq!(pivot_msgs[1].server_id, 1004); +} + +#[test] +fn around_sort_seq_cross_shard_boundary() { + let dir = create_anchor_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // Pivot at sort_seq=500 (shard 0), after should reach into shard 1 + let query = MessageQuery::for_talker("wxid_alice") + .around_sort_seq(500) + .context(3); + let result = db.query_messages_anchor(&query).unwrap(); + + let seqs: Vec = result.items.iter().map(|m| m.sort_seq).collect(); + // before(3): [200, 300, 300], pivot: [500], after(3): [600, 700, 800] + assert_eq!(seqs, vec![200, 300, 300, 500, 600, 700, 800]); +} + +#[test] +fn around_sort_seq_at_start() { + let dir = create_anchor_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // Pivot at the very first sort_seq — no before messages + let query = MessageQuery::for_talker("wxid_alice") + .around_sort_seq(100) + .context(2); + let result = db.query_messages_anchor(&query).unwrap(); + + let seqs: Vec = result.items.iter().map(|m| m.sort_seq).collect(); + // before: [], pivot: [100], after(2): [200, 300] + assert_eq!(seqs, vec![100, 200, 300]); +} + +#[test] +fn pooled_around_sort_seq_works_after_shard_file_removed() { + let dir = create_anchor_fixture(); + let db = WechatDb::open_with_pool(dir.path(), |_conn| Ok(())).unwrap(); + + let shard_path = dir.path().join("message").join("message_0.db"); + fs::remove_file(&shard_path).unwrap(); + + let query = MessageQuery::for_talker("wxid_alice") + .around_sort_seq(500) + .context(2); + let result = db.query_messages_anchor(&query).unwrap(); + + let seqs: Vec = result.items.iter().map(|m| m.sort_seq).collect(); + assert_eq!(seqs, vec![300, 300, 500, 600, 700]); +} + +// --------------------------------------------------------------------------- +// AroundServerId tests +// --------------------------------------------------------------------------- + +#[test] +fn around_server_id_basic() { + let dir = create_anchor_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let query = MessageQuery::for_talker("wxid_alice") + .around_server_id(1005) + .context(2); + let result = db.query_messages_anchor(&query).unwrap(); + + let seqs: Vec = result.items.iter().map(|m| m.sort_seq).collect(); + // pivot=1005 is at sort_seq=500 + // before(2): [300, 300], pivot: [500], after(2): [600, 700] + assert_eq!(seqs, vec![300, 300, 500, 600, 700]); +} + +#[test] +fn around_server_id_with_collision() { + let dir = create_anchor_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // Target server_id=1003 which shares sort_seq=300 with server_id=1004 + let query = MessageQuery::for_talker("wxid_alice") + .around_server_id(1003) + .context(1); + let result = db.query_messages_anchor(&query).unwrap(); + + let ids: Vec = result.items.iter().map(|m| m.server_id).collect(); + // before(1): [1002], pivot: [1003], after(1): [1004] + // server_id=1004 has same sort_seq but higher create_time, so it's "after" 1003 + assert_eq!(ids, vec![1002, 1003, 1004]); +} + +#[test] +fn around_server_id_not_found_returns_empty() { + let dir = create_anchor_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let query = MessageQuery::for_talker("wxid_alice") + .around_server_id(9999) + .context(5); + let result = db.query_messages_anchor(&query).unwrap(); + + assert!(result.items.is_empty()); +} + +#[test] +fn around_server_id_cross_shard() { + let dir = create_anchor_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // server_id=1006 is in shard 1, context should reach back to shard 0 + let query = MessageQuery::for_talker("wxid_alice") + .around_server_id(1006) + .context(2); + let result = db.query_messages_anchor(&query).unwrap(); + + let seqs: Vec = result.items.iter().map(|m| m.sort_seq).collect(); + // before(2): [300, 500] from shard 0, pivot: [600], after(2): [700, 800] + // Note: 300 appears twice (collision) but LIMIT 2 on the before SQL per shard + // Shard 0 returns the 2 most recent before pivot: sort_seq=500, then 300 + // Shard 1 has nothing before sort_seq=600 (only 600 itself which is the pivot) + // Cross-shard merge DESC → [500, 300/300] → take(2) → [500, 300] → reverse → [300, 500] + // But which 300? The DESC order picks 300(ct=1700000301, sid=1004) first + assert_eq!(seqs.len(), 5); + assert_eq!(seqs[0], 300); // or could be 300 (1004) + assert_eq!(seqs[1], 500); + assert_eq!(seqs[2], 600); // pivot + assert_eq!(seqs[3], 700); + assert_eq!(seqs[4], 800); +} + +#[test] +fn pooled_around_server_id_works_after_shard_file_removed() { + let dir = create_anchor_fixture(); + let db = WechatDb::open_with_pool(dir.path(), |_conn| Ok(())).unwrap(); + + let shard_path = dir.path().join("message").join("message_1.db"); + fs::remove_file(&shard_path).unwrap(); + + let query = MessageQuery::for_talker("wxid_alice") + .around_server_id(1006) + .context(2); + let result = db.query_messages_anchor(&query).unwrap(); + + let seqs: Vec = result.items.iter().map(|m| m.sort_seq).collect(); + assert_eq!(seqs.len(), 5); + assert_eq!(seqs[0], 300); + assert_eq!(seqs[1], 500); + assert_eq!(seqs[2], 600); + assert_eq!(seqs[3], 700); + assert_eq!(seqs[4], 800); +} diff --git a/crates/wx-db/tests/chatrooms-protobuf.rs b/crates/wx-db/tests/chatrooms-protobuf.rs new file mode 100644 index 0000000..8f7ecc4 --- /dev/null +++ b/crates/wx-db/tests/chatrooms-protobuf.rs @@ -0,0 +1,210 @@ +use std::fs; +use std::path::Path; + +use rusqlite::{params, Connection}; +use tempfile::TempDir; +use wx_db::test_ddl; +use wx_db::{encode_room_data_for_test, ChatRoomQuery, WechatDb}; + +// ---- helpers ---- + +/// Create a minimal fixture directory with contact.db (including chat_room table), +/// session.db, and message_0.db. +fn create_fixture() -> TempDir { + let dir = TempDir::new().unwrap(); + let base = dir.path(); + + // contact/contact.db + let contact_dir = base.join("contact"); + fs::create_dir_all(&contact_dir).unwrap(); + create_contact_db(&contact_dir.join("contact.db")); + + // session/session.db + let session_dir = base.join("session"); + fs::create_dir_all(&session_dir).unwrap(); + create_session_db(&session_dir.join("session.db")); + + // message/message_0.db (open needs at least 1 shard) + let msg_dir = base.join("message"); + fs::create_dir_all(&msg_dir).unwrap(); + create_message_shard(&msg_dir.join("message_0.db"), 1000); + + dir +} + +fn create_contact_db(path: &Path) { + let conn = Connection::open(path).unwrap(); + + // contact table (required by WechatDb) + test_ddl::create_test_contact_table(&conn); + + // chat_room table + conn.execute_batch( + "CREATE TABLE chat_room ( + username TEXT, + owner TEXT DEFAULT '', + ext_buffer BLOB + );", + ) + .unwrap(); + + // Row 1: valid protobuf with 2 members (one with display_name, one without) + let ext_buffer = + encode_room_data_for_test(&[("wxid_member1", Some("Member One")), ("wxid_member2", None)]); + conn.execute( + "INSERT INTO chat_room (username, owner, ext_buffer) VALUES (?1, ?2, ?3)", + params!["group@chatroom", "wxid_member1", ext_buffer], + ) + .unwrap(); + + // Row 2: empty ext_buffer (should return empty members, not panic) + conn.execute( + "INSERT INTO chat_room (username, owner, ext_buffer) VALUES (?1, ?2, ?3)", + params!["empty_group@chatroom", "wxid_owner", Vec::::new()], + ) + .unwrap(); + + // Row 3: another group for testing list-all + let ext_buffer2 = encode_room_data_for_test(&[("wxid_solo", Some("Solo Name"))]); + conn.execute( + "INSERT INTO chat_room (username, owner, ext_buffer) VALUES (?1, ?2, ?3)", + params!["another_group@chatroom", "wxid_solo", ext_buffer2], + ) + .unwrap(); +} + +fn create_session_db(path: &Path) { + let conn = Connection::open(path).unwrap(); + test_ddl::create_test_session_table(&conn); +} + +fn create_message_shard(path: &Path, timestamp: i64) { + let conn = Connection::open(path).unwrap(); + conn.execute_batch("CREATE TABLE Timestamp (timestamp INTEGER);") + .unwrap(); + conn.execute("INSERT INTO Timestamp VALUES (?1)", params![timestamp]) + .unwrap(); +} + +// ---- tests ---- + +#[test] +fn chatrooms_protobuf_query_by_username() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db + .query_chatrooms(&ChatRoomQuery::new().username("group@chatroom")) + .unwrap(); + + assert_eq!(result.items.len(), 1); + assert_eq!(result.stats.skipped, 0); + + let room = &result.items[0]; + assert_eq!(room.username, "group@chatroom"); + assert_eq!(room.owner, "wxid_member1"); + assert_eq!(room.members.len(), 2); + + // First member: has display_name + assert_eq!(room.members[0].user_name, "wxid_member1"); + assert_eq!(room.members[0].display_name.as_deref(), Some("Member One")); + + // Second member: no display_name + assert_eq!(room.members[1].user_name, "wxid_member2"); + assert_eq!(room.members[1].display_name, None); +} + +#[test] +fn chatrooms_protobuf_empty_ext_buffer() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db + .query_chatrooms(&ChatRoomQuery::new().username("empty_group@chatroom")) + .unwrap(); + + assert_eq!(result.items.len(), 1); + assert_eq!(result.stats.skipped, 0); // empty ext_buffer is normal, not skipped + + let room = &result.items[0]; + assert_eq!(room.username, "empty_group@chatroom"); + assert_eq!(room.owner, "wxid_owner"); + assert!(room.members.is_empty()); +} + +#[test] +fn chatrooms_protobuf_query_all() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db.query_chatrooms(&ChatRoomQuery::new()).unwrap(); + + assert_eq!(result.stats.total_rows, 3); + assert_eq!(result.items.len(), 3); + assert_eq!(result.stats.skipped, 0); +} + +#[test] +fn chatrooms_protobuf_query_nonexistent_username() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db + .query_chatrooms(&ChatRoomQuery::new().username("nonexistent@chatroom")) + .unwrap(); + + assert_eq!(result.items.len(), 0); + assert_eq!(result.stats.total_rows, 0); +} + +#[test] +fn chatrooms_protobuf_query_with_limit() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db.query_chatrooms(&ChatRoomQuery::new().limit(1)).unwrap(); + + assert_eq!(result.items.len(), 1); + assert_eq!(result.stats.total_rows, 3); // total_rows = pre-limit count +} + +#[test] +fn chatrooms_protobuf_query_with_offset() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db + .query_chatrooms(&ChatRoomQuery::new().limit(100).offset(2)) + .unwrap(); + + assert_eq!(result.items.len(), 1); // 3 total, offset 2 = 1 remaining + assert_eq!(result.stats.total_rows, 3); // total_rows = pre-limit count +} + +#[test] +fn chatrooms_protobuf_pagination_stability() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // Page 1: offset=0, limit=2 + let page1 = db + .query_chatrooms(&ChatRoomQuery::new().limit(2).offset(0)) + .unwrap(); + // Page 2: offset=2, limit=2 + let page2 = db + .query_chatrooms(&ChatRoomQuery::new().limit(2).offset(2)) + .unwrap(); + + assert_eq!(page1.items.len(), 2); + assert_eq!(page2.items.len(), 1); + + // Ensure no overlap + let page1_names: Vec<&str> = page1.items.iter().map(|r| r.username.as_str()).collect(); + let page2_names: Vec<&str> = page2.items.iter().map(|r| r.username.as_str()).collect(); + for name in &page2_names { + assert!( + !page1_names.contains(name), + "chatroom {name} appears in both pages" + ); + } +} diff --git a/crates/wx-db/tests/contacts-extra-buffer.rs b/crates/wx-db/tests/contacts-extra-buffer.rs new file mode 100644 index 0000000..00cb24e --- /dev/null +++ b/crates/wx-db/tests/contacts-extra-buffer.rs @@ -0,0 +1,328 @@ +use std::fs; +use std::path::Path; + +use rusqlite::{params, Connection}; +use tempfile::TempDir; +use wx_db::test_ddl; +use wx_db::{encode_extra_buffer_for_test, ContactQuery, WechatDb}; + +// ---- helpers ---- + +/// Create a fixture with contacts that have extra_buffer data, description, and labels. +fn create_fixture_extended() -> TempDir { + let dir = TempDir::new().unwrap(); + let base = dir.path(); + + // contact/contact.db + let contact_dir = base.join("contact"); + fs::create_dir_all(&contact_dir).unwrap(); + create_contact_db_extended(&contact_dir.join("contact.db")); + + // session/session.db (minimal) + let session_dir = base.join("session"); + fs::create_dir_all(&session_dir).unwrap(); + create_session_db(&session_dir.join("session.db")); + + // message/message_0.db (minimal shard) + let msg_dir = base.join("message"); + fs::create_dir_all(&msg_dir).unwrap(); + create_message_shard(&msg_dir.join("message_0.db"), 1000); + + dir +} + +/// Create a fixture WITHOUT a contact_label table. +fn create_fixture_no_label_table() -> TempDir { + let dir = TempDir::new().unwrap(); + let base = dir.path(); + + let contact_dir = base.join("contact"); + fs::create_dir_all(&contact_dir).unwrap(); + { + let conn = Connection::open(contact_dir.join("contact.db")).unwrap(); + test_ddl::create_test_contact_table(&conn); + // No contact_label table at all + let blob = + encode_extra_buffer_for_test(Some(1), None, None, None, None, None, None, Some("5,6")); + conn.execute( + "INSERT INTO contact (username, alias, remark, nick_name, extra_buffer) \ + VALUES (?1, ?2, ?3, ?4, ?5)", + params!["wxid_nolabel", "", "No Label", "", blob], + ) + .unwrap(); + } + + let session_dir = base.join("session"); + fs::create_dir_all(&session_dir).unwrap(); + create_session_db(&session_dir.join("session.db")); + + let msg_dir = base.join("message"); + fs::create_dir_all(&msg_dir).unwrap(); + create_message_shard(&msg_dir.join("message_0.db"), 1000); + + dir +} + +fn create_contact_db_extended(path: &Path) { + let conn = Connection::open(path).unwrap(); + test_ddl::create_test_contact_table(&conn); + test_ddl::create_test_contact_label_table(&conn); + + // Label rows + conn.execute( + "INSERT INTO contact_label VALUES (?1, ?2, ?3)", + params!["5", "体育生", 0], + ) + .unwrap(); + conn.execute( + "INSERT INTO contact_label VALUES (?1, ?2, ?3)", + params!["6", "直男", 1], + ) + .unwrap(); + + // Contact 1: all fields populated + let blob_full = encode_extra_buffer_for_test( + Some(1), + Some("成长本就是一个孤立无援的过程"), + Some("CN"), + Some("Beijing"), + Some("Haidian"), + Some(30), + Some("15891926830"), + Some("5,6"), + ); + conn.execute( + "INSERT INTO contact (username, alias, remark, nick_name, description, extra_buffer) \ + VALUES (?1, ?2, ?3, ?4, ?5, ?6)", + params![ + "wxid_full", + "Dalong-SS", + "张三", + "大龙", + "QQ 442007516", + blob_full + ], + ) + .unwrap(); + + // Contact 2: empty extra_buffer + conn.execute( + "INSERT INTO contact (username, alias, remark, nick_name) \ + VALUES (?1, ?2, ?3, ?4)", + params!["wxid_empty", "empty_alias", "Empty Remark", "Empty Nick"], + ) + .unwrap(); + + // Contact 3: partial region (only province) + let blob_partial = + encode_extra_buffer_for_test(None, None, None, Some("Guangdong"), None, None, None, None); + conn.execute( + "INSERT INTO contact (username, alias, remark, nick_name, extra_buffer) \ + VALUES (?1, ?2, ?3, ?4, ?5)", + params![ + "wxid_partial", + "", + "Partial Remark", + "Partial", + blob_partial + ], + ) + .unwrap(); + + // Contact 4: phone only (for keyword search test) + let blob_phone = encode_extra_buffer_for_test( + None, + None, + None, + None, + None, + None, + Some("13912345678"), + None, + ); + conn.execute( + "INSERT INTO contact (username, alias, remark, nick_name, extra_buffer) \ + VALUES (?1, ?2, ?3, ?4, ?5)", + params!["wxid_phone", "", "Phone Guy", "PG", blob_phone], + ) + .unwrap(); +} + +fn create_session_db(path: &Path) { + let conn = Connection::open(path).unwrap(); + test_ddl::create_test_session_table_extended(&conn); + conn.execute( + "INSERT INTO SessionTable VALUES (?1, ?2, ?3, NULL, NULL, NULL)", + params!["wxid_full", 1000, "hello"], + ) + .unwrap(); +} + +fn create_message_shard(path: &Path, timestamp: i64) { + let conn = Connection::open(path).unwrap(); + conn.execute_batch("CREATE TABLE Timestamp (timestamp INTEGER);") + .unwrap(); + conn.execute("INSERT INTO Timestamp VALUES (?1)", params![timestamp]) + .unwrap(); +} + +// ---- tests ---- + +#[test] +fn contact_extra_buffer_all_fields() { + let dir = create_fixture_extended(); + let db = WechatDb::open(dir.path()).unwrap(); + let result = db.query_contacts(&ContactQuery::new()).unwrap(); + + let full = result + .items + .iter() + .find(|c| c.user_name == "wxid_full") + .unwrap(); + + assert_eq!(full.gender, Some(1)); + assert_eq!( + full.signature.as_deref(), + Some("成长本就是一个孤立无援的过程") + ); + assert_eq!(full.region.as_deref(), Some("CN · Beijing · Haidian")); + assert_eq!(full.source_scene, Some(30)); + assert_eq!(full.phone.as_deref(), Some("15891926830")); + assert_eq!(full.memo.as_deref(), Some("QQ 442007516")); +} + +#[test] +fn contact_extra_buffer_phone_extraction() { + let dir = create_fixture_extended(); + let db = WechatDb::open(dir.path()).unwrap(); + let result = db.query_contacts(&ContactQuery::new()).unwrap(); + + let phone_guy = result + .items + .iter() + .find(|c| c.user_name == "wxid_phone") + .unwrap(); + assert_eq!(phone_guy.phone.as_deref(), Some("13912345678")); +} + +#[test] +fn contact_extra_buffer_labels_resolved() { + let dir = create_fixture_extended(); + let db = WechatDb::open(dir.path()).unwrap(); + let result = db.query_contacts(&ContactQuery::new()).unwrap(); + + let full = result + .items + .iter() + .find(|c| c.user_name == "wxid_full") + .unwrap(); + assert_eq!(full.labels.len(), 2); + assert!(full.labels.contains(&"体育生".to_string())); + assert!(full.labels.contains(&"直男".to_string())); +} + +#[test] +fn contact_extra_buffer_empty_blob() { + let dir = create_fixture_extended(); + let db = WechatDb::open(dir.path()).unwrap(); + let result = db.query_contacts(&ContactQuery::new()).unwrap(); + + let empty = result + .items + .iter() + .find(|c| c.user_name == "wxid_empty") + .unwrap(); + assert_eq!(empty.gender, None); + assert_eq!(empty.signature, None); + assert_eq!(empty.region, None); + assert_eq!(empty.source_scene, None); + assert_eq!(empty.phone, None); + assert!(empty.labels.is_empty()); +} + +#[test] +fn contact_extra_buffer_no_label_table() { + let dir = create_fixture_no_label_table(); + let db = WechatDb::open(dir.path()).unwrap(); + let result = db.query_contacts(&ContactQuery::new()).unwrap(); + + let c = &result.items[0]; + assert_eq!(c.user_name, "wxid_nolabel"); + // label_ids are "5,6" but no contact_label table → labels is empty + assert!(c.labels.is_empty()); + // gender still decoded + assert_eq!(c.gender, Some(1)); +} + +#[test] +fn contact_extra_buffer_partial_region() { + let dir = create_fixture_extended(); + let db = WechatDb::open(dir.path()).unwrap(); + let result = db.query_contacts(&ContactQuery::new()).unwrap(); + + let partial = result + .items + .iter() + .find(|c| c.user_name == "wxid_partial") + .unwrap(); + assert_eq!(partial.region.as_deref(), Some("Guangdong")); +} + +#[test] +fn contact_memo_from_description() { + let dir = create_fixture_extended(); + let db = WechatDb::open(dir.path()).unwrap(); + let result = db.query_contacts(&ContactQuery::new()).unwrap(); + + let full = result + .items + .iter() + .find(|c| c.user_name == "wxid_full") + .unwrap(); + assert_eq!(full.memo.as_deref(), Some("QQ 442007516")); + + // Contact with no description → memo is None + let empty = result + .items + .iter() + .find(|c| c.user_name == "wxid_empty") + .unwrap(); + assert_eq!(empty.memo, None); +} + +#[test] +fn contact_keyword_matches_description() { + let dir = create_fixture_extended(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db + .query_contacts(&ContactQuery::new().keyword("442007")) + .unwrap(); + assert_eq!(result.items.len(), 1); + assert_eq!(result.items[0].user_name, "wxid_full"); + assert_eq!(result.stats.filtered_count, Some(1)); +} + +#[test] +fn contact_keyword_matches_phone() { + let dir = create_fixture_extended(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db + .query_contacts(&ContactQuery::new().keyword("13912345")) + .unwrap(); + assert_eq!(result.items.len(), 1); + assert_eq!(result.items[0].user_name, "wxid_phone"); +} + +#[test] +fn contact_keyword_matches_label_name() { + let dir = create_fixture_extended(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db + .query_contacts(&ContactQuery::new().keyword("体育生")) + .unwrap(); + assert_eq!(result.items.len(), 1); + assert_eq!(result.items[0].user_name, "wxid_full"); +} diff --git a/crates/wx-db/tests/contacts-sessions.rs b/crates/wx-db/tests/contacts-sessions.rs new file mode 100644 index 0000000..e26a877 --- /dev/null +++ b/crates/wx-db/tests/contacts-sessions.rs @@ -0,0 +1,573 @@ +use std::fs; +use std::path::Path; + +use rusqlite::{params, Connection}; +use tempfile::TempDir; +use wx_db::test_ddl; +use wx_db::{ContactQuery, DbError, MessageQuery, SessionQuery, SortOrder, WechatDb}; + +// ---- helpers ---- + +/// Create a minimal fixture directory with contact.db, session.db, and message_0.db. +fn create_fixture() -> TempDir { + let dir = TempDir::new().unwrap(); + let base = dir.path(); + + // contact/contact.db + let contact_dir = base.join("contact"); + fs::create_dir_all(&contact_dir).unwrap(); + create_contact_db(&contact_dir.join("contact.db")); + + // session/session.db + let session_dir = base.join("session"); + fs::create_dir_all(&session_dir).unwrap(); + create_session_db(&session_dir.join("session.db")); + + // message/message_0.db + let msg_dir = base.join("message"); + fs::create_dir_all(&msg_dir).unwrap(); + create_message_shard(&msg_dir.join("message_0.db"), 1000); + + dir +} + +fn create_contact_db(path: &Path) { + let conn = Connection::open(path).unwrap(); + test_ddl::create_test_contact_table(&conn); + test_ddl::create_test_contact_label_table(&conn); + conn.execute( + "INSERT INTO contact (username, alias, remark, nick_name) VALUES (?1, ?2, ?3, ?4)", + params!["wxid_alice", "alice_a", "Alice Remark", "Alice"], + ) + .unwrap(); + conn.execute( + "INSERT INTO contact (username, alias, remark, nick_name) VALUES (?1, ?2, ?3, ?4)", + params!["wxid_bob", "bob_b", "Bob Test Remark", "Bob"], + ) + .unwrap(); + conn.execute( + "INSERT INTO contact (username, alias, remark, nick_name) VALUES (?1, ?2, ?3, ?4)", + params!["wxid_test_user", "", "Charlie", "Charlie Nick"], + ) + .unwrap(); +} + +fn create_session_db(path: &Path) { + let conn = Connection::open(path).unwrap(); + test_ddl::create_test_session_table_extended(&conn); + + // Session 1: plain text summary + conn.execute( + "INSERT INTO SessionTable VALUES (?1, ?2, ?3, NULL, NULL, NULL)", + params!["wxid_alice", 2000, "hello from alice"], + ) + .unwrap(); + + // Session 2: zstd-compressed blob summary + let compressed = zstd::encode_all(&b"compressed summary"[..], 0).unwrap(); + conn.execute( + "INSERT INTO SessionTable VALUES (?1, ?2, ?3, NULL, NULL, NULL)", + params!["wxid_bob", 3000, compressed], + ) + .unwrap(); +} + +fn create_message_shard(path: &Path, timestamp: i64) { + let conn = Connection::open(path).unwrap(); + conn.execute_batch("CREATE TABLE Timestamp (timestamp INTEGER);") + .unwrap(); + conn.execute("INSERT INTO Timestamp VALUES (?1)", params![timestamp]) + .unwrap(); +} + +fn create_auxiliary_message_db(path: &Path) { + let _conn = Connection::open(path).unwrap(); +} + +// ---- existing smoke test ---- + +#[test] +fn open_nonexistent_dir_returns_not_found() { + let result = WechatDb::open("/tmp/nonexistent-wx-db-dir-that-does-not-exist"); + assert!(result.is_err()); + let err = result.unwrap_err(); + assert!( + matches!(err, DbError::NotFound(_)), + "expected NotFound, got: {err:?}" + ); +} + +// ---- open tests ---- + +#[test] +fn contacts_sessions_open_valid_dir() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + // Should have exactly 1 shard + assert!(format!("{db:?}").contains("shards")); +} + +#[test] +fn contacts_sessions_open_ignores_auxiliary_message_dbs() { + let dir = create_fixture(); + let msg_dir = dir.path().join("message"); + create_auxiliary_message_db(&msg_dir.join("message_fts.db")); + create_auxiliary_message_db(&msg_dir.join("message_resource.db")); + + let db = WechatDb::open(dir.path()).unwrap(); + let debug = format!("{db:?}"); + assert!(debug.contains("message_0.db")); + assert!(!debug.contains("message_fts.db")); + assert!(!debug.contains("message_resource.db")); +} + +#[test] +fn contacts_sessions_open_only_auxiliary_message_dbs_has_empty_shards() { + let dir = TempDir::new().unwrap(); + let base = dir.path(); + let contact_dir = base.join("contact"); + fs::create_dir_all(&contact_dir).unwrap(); + create_contact_db(&contact_dir.join("contact.db")); + let session_dir = base.join("session"); + fs::create_dir_all(&session_dir).unwrap(); + create_session_db(&session_dir.join("session.db")); + let msg_dir = base.join("message"); + fs::create_dir_all(&msg_dir).unwrap(); + create_auxiliary_message_db(&msg_dir.join("message_fts.db")); + create_auxiliary_message_db(&msg_dir.join("message_resource.db")); + + let db = WechatDb::open(base).unwrap(); + assert!(format!("{db:?}").contains("shards: []")); +} + +#[test] +fn contacts_sessions_query_messages_with_only_auxiliary_dbs_returns_no_shards() { + let dir = TempDir::new().unwrap(); + let base = dir.path(); + let contact_dir = base.join("contact"); + fs::create_dir_all(&contact_dir).unwrap(); + create_contact_db(&contact_dir.join("contact.db")); + let session_dir = base.join("session"); + fs::create_dir_all(&session_dir).unwrap(); + create_session_db(&session_dir.join("session.db")); + let msg_dir = base.join("message"); + fs::create_dir_all(&msg_dir).unwrap(); + create_auxiliary_message_db(&msg_dir.join("message_fts.db")); + create_auxiliary_message_db(&msg_dir.join("message_resource.db")); + + let db = WechatDb::open(base).unwrap(); + let result = db.query_messages(&MessageQuery::for_talker("wxid_alice")); + assert!(matches!(result, Err(DbError::NoShards))); +} + +#[test] +fn contacts_sessions_open_missing_contact_db() { + let dir = TempDir::new().unwrap(); + let base = dir.path(); + // Create session and message but NOT contact + let session_dir = base.join("session"); + fs::create_dir_all(&session_dir).unwrap(); + create_session_db(&session_dir.join("session.db")); + let msg_dir = base.join("message"); + fs::create_dir_all(&msg_dir).unwrap(); + create_message_shard(&msg_dir.join("message_0.db"), 1000); + + let result = WechatDb::open(base); + assert!(matches!(result, Err(DbError::NotFound(_)))); +} + +#[test] +fn contacts_sessions_open_missing_session_db() { + let dir = TempDir::new().unwrap(); + let base = dir.path(); + // Create contact and message but NOT session + let contact_dir = base.join("contact"); + fs::create_dir_all(&contact_dir).unwrap(); + create_contact_db(&contact_dir.join("contact.db")); + let msg_dir = base.join("message"); + fs::create_dir_all(&msg_dir).unwrap(); + create_message_shard(&msg_dir.join("message_0.db"), 1000); + + let result = WechatDb::open(base); + assert!(matches!(result, Err(DbError::NotFound(_)))); +} + +#[test] +fn contacts_sessions_open_no_shards() { + let dir = TempDir::new().unwrap(); + let base = dir.path(); + let contact_dir = base.join("contact"); + fs::create_dir_all(&contact_dir).unwrap(); + create_contact_db(&contact_dir.join("contact.db")); + let session_dir = base.join("session"); + fs::create_dir_all(&session_dir).unwrap(); + create_session_db(&session_dir.join("session.db")); + // message dir exists but is empty + let msg_dir = base.join("message"); + fs::create_dir_all(&msg_dir).unwrap(); + + let db = WechatDb::open(base).unwrap(); + assert!(format!("{db:?}").contains("shards: []")); +} + +// ---- contacts tests ---- + +#[test] +fn contacts_sessions_query_contacts_all() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db.query_contacts(&ContactQuery::new()).unwrap(); + assert_eq!(result.items.len(), 3); + assert_eq!(result.stats.total_rows, 3); + assert_eq!(result.stats.skipped, 0); +} + +#[test] +fn contacts_sessions_query_contacts_keyword_username() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // "test" should match wxid_test_user (userName) and Bob Test Remark (remark) + let result = db + .query_contacts(&ContactQuery::new().keyword("test")) + .unwrap(); + assert_eq!(result.items.len(), 2); // wxid_test_user + Bob Test Remark +} + +#[test] +fn contacts_sessions_query_contacts_keyword_alias() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // "alice_a" should match alias of first contact + let result = db + .query_contacts(&ContactQuery::new().keyword("alice_a")) + .unwrap(); + assert_eq!(result.items.len(), 1); + assert_eq!(result.items[0].user_name, "wxid_alice"); +} + +#[test] +fn contacts_sessions_query_contacts_keyword_remark() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // "Charlie" should match remark of wxid_test_user + let result = db + .query_contacts(&ContactQuery::new().keyword("Charlie")) + .unwrap(); + assert!(result.items.iter().any(|c| c.user_name == "wxid_test_user")); +} + +#[test] +fn contacts_sessions_query_contacts_keyword_nick_name() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // "Nick" should match nick_name of wxid_test_user ("Charlie Nick") + let result = db + .query_contacts(&ContactQuery::new().keyword("Nick")) + .unwrap(); + assert_eq!(result.items.len(), 1); + assert_eq!(result.items[0].user_name, "wxid_test_user"); +} + +#[test] +fn contacts_sessions_query_contacts_limit() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db.query_contacts(&ContactQuery::new().limit(1)).unwrap(); + assert_eq!(result.items.len(), 1); +} + +// ---- sessions tests ---- + +#[test] +fn contacts_sessions_query_sessions_all() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db.query_sessions(&SessionQuery::new().limit(10)).unwrap(); + assert_eq!(result.items.len(), 2); + assert_eq!(result.stats.skipped, 0); + + // Should be sorted by sort_timestamp DESC (bob=3000 first, alice=2000 second) + assert_eq!(result.items[0].username, "wxid_bob"); + assert_eq!(result.items[0].sort_timestamp, 3000); + assert_eq!(result.items[1].username, "wxid_alice"); + assert_eq!(result.items[1].sort_timestamp, 2000); +} + +#[test] +fn contacts_sessions_query_sessions_text_summary() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db.query_sessions(&SessionQuery::new().limit(10)).unwrap(); + // Alice has plain text summary + let alice = result + .items + .iter() + .find(|s| s.username == "wxid_alice") + .unwrap(); + assert_eq!(alice.summary, "hello from alice"); +} + +#[test] +fn contacts_sessions_query_sessions_zstd_summary() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db.query_sessions(&SessionQuery::new().limit(10)).unwrap(); + // Bob has zstd-compressed summary + let bob = result + .items + .iter() + .find(|s| s.username == "wxid_bob") + .unwrap(); + assert_eq!(bob.summary, "compressed summary"); +} + +#[test] +fn contacts_sessions_query_sessions_limit() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db.query_sessions(&SessionQuery::new().limit(1)).unwrap(); + assert_eq!(result.items.len(), 1); + // Should be the most recent (bob, 3000) + assert_eq!(result.items[0].username, "wxid_bob"); +} + +#[test] +fn reopen_sessions_refreshes_connection() { + let dir = create_fixture(); + let mut db = WechatDb::open(dir.path()).unwrap(); + + // Initial query + let result = db.query_sessions(&SessionQuery::new().limit(10)).unwrap(); + assert_eq!(result.items.len(), 2); + + // Modify session.db externally (add a new session) + let session_path = dir.path().join("session").join("session.db"); + let conn = Connection::open(&session_path).unwrap(); + conn.execute( + "INSERT INTO SessionTable VALUES (?1, ?2, ?3, NULL, NULL, NULL)", + params!["wxid_charlie", 5000, "hello from charlie"], + ) + .unwrap(); + drop(conn); + + // Without reopen, the old connection won't see the new row (read-only + WAL) + // After reopen, it should see the new data + db.reopen_sessions().unwrap(); + + let result = db.query_sessions(&SessionQuery::new().limit(10)).unwrap(); + assert_eq!(result.items.len(), 3); + assert_eq!(result.items[0].username, "wxid_charlie"); + assert_eq!(result.items[0].sort_timestamp, 5000); +} + +#[test] +fn query_sessions_with_content_fields() { + let dir = TempDir::new().unwrap(); + let base = dir.path(); + + // contact.db (required by WechatDb::open) + let contact_dir = base.join("contact"); + fs::create_dir_all(&contact_dir).unwrap(); + create_contact_db(&contact_dir.join("contact.db")); + + // session.db with content fields populated + let session_dir = base.join("session"); + fs::create_dir_all(&session_dir).unwrap(); + { + let conn = Connection::open(session_dir.join("session.db")).unwrap(); + test_ddl::create_test_session_table_extended(&conn); + conn.execute( + "INSERT INTO SessionTable VALUES (?1, ?2, ?3, ?4, ?5, ?6)", + params![ + "wxid_alice", + 1000, + "hi there", + 1, + "wxid_sender", + "Sender Name" + ], + ) + .unwrap(); + } + + // message shard (required by WechatDb::open) + let msg_dir = base.join("message"); + fs::create_dir_all(&msg_dir).unwrap(); + create_message_shard(&msg_dir.join("message_0.db"), 1000); + + let db = WechatDb::open(base).unwrap(); + let result = db.query_sessions(&SessionQuery::new().limit(10)).unwrap(); + assert_eq!(result.items.len(), 1); + + let session = &result.items[0]; + assert_eq!(session.last_msg_type, Some(1)); + assert_eq!(session.last_msg_sender.as_deref(), Some("wxid_sender")); + assert_eq!( + session.last_sender_display_name.as_deref(), + Some("Sender Name") + ); +} + +#[test] +fn query_sessions_strips_group_summary_prefix() { + let dir = TempDir::new().unwrap(); + let base = dir.path(); + + let contact_dir = base.join("contact"); + fs::create_dir_all(&contact_dir).unwrap(); + create_contact_db(&contact_dir.join("contact.db")); + + let session_dir = base.join("session"); + fs::create_dir_all(&session_dir).unwrap(); + { + let conn = Connection::open(session_dir.join("session.db")).unwrap(); + test_ddl::create_test_session_table_extended(&conn); + // Group chat: prefix should be stripped + conn.execute( + "INSERT INTO SessionTable VALUES (?1, ?2, ?3, NULL, NULL, NULL)", + params!["group@chatroom", 2000, "wxid_abc:\nhello"], + ) + .unwrap(); + // 1-on-1 chat: prefix should NOT be stripped + conn.execute( + "INSERT INTO SessionTable VALUES (?1, ?2, ?3, NULL, NULL, NULL)", + params!["wxid_alice", 1000, "some_id:\nwhatever"], + ) + .unwrap(); + } + + let msg_dir = base.join("message"); + fs::create_dir_all(&msg_dir).unwrap(); + create_message_shard(&msg_dir.join("message_0.db"), 1000); + + let db = WechatDb::open(base).unwrap(); + let result = db.query_sessions(&SessionQuery::new().limit(10)).unwrap(); + assert_eq!(result.items.len(), 2); + + // Group chat (sort_timestamp=2000 is first, DESC order) + let group = result + .items + .iter() + .find(|s| s.username == "group@chatroom") + .unwrap(); + assert_eq!( + group.summary, "hello", + "group summary prefix should be stripped" + ); + + // 1-on-1 chat + let dm = result + .items + .iter() + .find(|s| s.username == "wxid_alice") + .unwrap(); + assert_eq!( + dm.summary, "some_id:\nwhatever", + "non-group summary should NOT be stripped" + ); +} + +#[test] +fn query_sessions_asc_order() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db + .query_sessions(&SessionQuery::new().limit(10).order(SortOrder::Asc)) + .unwrap(); + assert_eq!(result.items.len(), 2); + // ASC: alice (2000) first, bob (3000) second + assert_eq!(result.items[0].username, "wxid_alice"); + assert_eq!(result.items[0].sort_timestamp, 2000); + assert_eq!(result.items[1].username, "wxid_bob"); + assert_eq!(result.items[1].sort_timestamp, 3000); +} + +#[test] +fn contacts_pagination_stability() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // Page 1: offset=0, limit=2 + let page1 = db + .query_contacts(&ContactQuery::new().limit(2).offset(0)) + .unwrap(); + // Page 2: offset=2, limit=2 + let page2 = db + .query_contacts(&ContactQuery::new().limit(2).offset(2)) + .unwrap(); + + assert_eq!(page1.items.len(), 2); + assert_eq!(page2.items.len(), 1); // 3 total, offset 2 → 1 remaining + + // Ensure no overlap + let page1_names: Vec<&str> = page1.items.iter().map(|c| c.user_name.as_str()).collect(); + let page2_names: Vec<&str> = page2.items.iter().map(|c| c.user_name.as_str()).collect(); + for name in &page2_names { + assert!( + !page1_names.contains(name), + "contact {name} appears in both pages" + ); + } +} + +#[test] +fn sessions_pagination_stability_same_timestamp() { + // Sessions with identical sort_timestamp must still paginate deterministically + let dir = TempDir::new().unwrap(); + let base = dir.path(); + + let contact_dir = base.join("contact"); + fs::create_dir_all(&contact_dir).unwrap(); + create_contact_db(&contact_dir.join("contact.db")); + + let session_dir = base.join("session"); + fs::create_dir_all(&session_dir).unwrap(); + { + let conn = Connection::open(session_dir.join("session.db")).unwrap(); + test_ddl::create_test_session_table_extended(&conn); + // 3 sessions all with the same timestamp + for name in &["wxid_aaa", "wxid_bbb", "wxid_ccc"] { + conn.execute( + "INSERT INTO SessionTable VALUES (?1, ?2, ?3, NULL, NULL, NULL)", + params![name, 5000_i64, format!("summary for {name}")], + ) + .unwrap(); + } + } + + let msg_dir = base.join("message"); + fs::create_dir_all(&msg_dir).unwrap(); + create_message_shard(&msg_dir.join("message_0.db"), 1000); + + let db = WechatDb::open(base).unwrap(); + + // Page 1: offset=0, limit=2 + let page1 = db + .query_sessions(&SessionQuery::new().limit(2).offset(0)) + .unwrap(); + // Page 2: offset=2, limit=2 + let page2 = db + .query_sessions(&SessionQuery::new().limit(2).offset(2)) + .unwrap(); + + assert_eq!(page1.items.len(), 2); + assert_eq!(page2.items.len(), 1); + + // No overlap + let p1: Vec<&str> = page1.items.iter().map(|s| s.username.as_str()).collect(); + let p2: Vec<&str> = page2.items.iter().map(|s| s.username.as_str()).collect(); + for name in &p2 { + assert!(!p1.contains(name), "session {name} appears in both pages"); + } +} diff --git a/crates/wx-db/tests/edge-cases.rs b/crates/wx-db/tests/edge-cases.rs new file mode 100644 index 0000000..1b776c9 --- /dev/null +++ b/crates/wx-db/tests/edge-cases.rs @@ -0,0 +1,106 @@ +use wx_db::{ + effective_limit, msg_sub_type_label, split_local_type, DEFAULT_QUERY_LIMIT, MAX_QUERY_LIMIT, + MSG_TYPE_APP, MSG_TYPE_TEXT, +}; + +// ---- effective_limit boundary tests ---- + +#[test] +fn effective_limit_zero_returns_default() { + assert_eq!(effective_limit(0), DEFAULT_QUERY_LIMIT); + assert_eq!(effective_limit(0), 1_000); +} + +#[test] +fn effective_limit_over_max_is_clamped() { + assert_eq!(effective_limit(20_000), MAX_QUERY_LIMIT); + assert_eq!(effective_limit(20_000), 20_000); +} + +#[test] +fn effective_limit_exactly_max() { + assert_eq!(effective_limit(MAX_QUERY_LIMIT), MAX_QUERY_LIMIT); +} + +#[test] +fn effective_limit_one_below_max() { + assert_eq!(effective_limit(MAX_QUERY_LIMIT - 1), MAX_QUERY_LIMIT - 1); +} + +#[test] +fn effective_limit_one_above_max() { + assert_eq!(effective_limit(MAX_QUERY_LIMIT + 1), MAX_QUERY_LIMIT); +} + +#[test] +fn effective_limit_normal_value_unchanged() { + assert_eq!(effective_limit(50), 50); + assert_eq!(effective_limit(1), 1); + assert_eq!(effective_limit(DEFAULT_QUERY_LIMIT), DEFAULT_QUERY_LIMIT); +} + +// ---- split_local_type boundary tests ---- + +#[test] +fn split_local_type_text() { + let (msg_type, sub_type) = split_local_type(1_i64); + assert_eq!(msg_type, 1); + assert_eq!(sub_type, 0); +} + +#[test] +fn split_local_type_app_with_sub_type() { + // local_type = (5 << 32) | 49 — canonical low32/high32 encoding + let local_type: i64 = (5_i64 << 32) | 49; + let (msg_type, sub_type) = split_local_type(local_type); + assert_eq!(msg_type, 49); + assert_eq!(sub_type, 5); +} + +#[test] +fn split_local_type_zero() { + let (msg_type, sub_type) = split_local_type(0_i64); + assert_eq!(msg_type, 0); + assert_eq!(sub_type, 0); +} + +#[test] +fn split_local_type_wechat_4x_quote() { + // WeChat 4.x: quote reply, local_type = (57 << 32) | 49 + let local_type: i64 = (57_i64 << 32) | 49; + let (msg_type, sub_type) = split_local_type(local_type); + assert_eq!(msg_type, 49); + assert_eq!(sub_type, 57); +} + +// ---- msg_sub_type_label tests ---- + +#[test] +fn msg_sub_type_label_app_link() { + assert_eq!(msg_sub_type_label(MSG_TYPE_APP, 5), "link"); +} + +#[test] +fn msg_sub_type_label_app_quote() { + assert_eq!(msg_sub_type_label(MSG_TYPE_APP, 57), "quote"); +} + +#[test] +fn msg_sub_type_label_app_transfer() { + assert_eq!(msg_sub_type_label(MSG_TYPE_APP, 2000), "transfer"); +} + +#[test] +fn msg_sub_type_label_app_group_announcement() { + assert_eq!(msg_sub_type_label(MSG_TYPE_APP, 87), "group_announcement"); +} + +#[test] +fn msg_sub_type_label_app_unknown_fallback() { + assert_eq!(msg_sub_type_label(MSG_TYPE_APP, 9999), "app"); +} + +#[test] +fn msg_sub_type_label_non_app_type() { + assert_eq!(msg_sub_type_label(MSG_TYPE_TEXT, 0), "text"); +} diff --git a/crates/wx-db/tests/fixtures/01-u64-local-type.sql b/crates/wx-db/tests/fixtures/01-u64-local-type.sql new file mode 100644 index 0000000..0700333 --- /dev/null +++ b/crates/wx-db/tests/fixtures/01-u64-local-type.sql @@ -0,0 +1,14 @@ +-- Fixture 01: u64 local_type (sub_type=57, msg_type=49) +-- Tests that local_type values exceeding 32 bits are correctly split. +-- local_type = (57 << 32) | 49 = 244813135921 + +INSERT INTO fixture_messages (sort_seq, server_id, local_type, sender, talker, create_time, message_content) +VALUES ( + 100, + 900001, + 244813135921, + 'wxid_test_alice', + 'wxid_test_bob', + 1700000100, + 'This is my replywxid_test_bobBobHello, how are you?1' +); diff --git a/crates/wx-db/tests/fixtures/02-xml-null-sentinel.sql b/crates/wx-db/tests/fixtures/02-xml-null-sentinel.sql new file mode 100644 index 0000000..190cdd5 --- /dev/null +++ b/crates/wx-db/tests/fixtures/02-xml-null-sentinel.sql @@ -0,0 +1,14 @@ +-- Fixture 02: XML null sentinel filtering +-- Tests that null and null are filtered to None. +-- local_type = (5 << 32) | 49 = 21474836529 + +INSERT INTO fixture_messages (sort_seq, server_id, local_type, sender, talker, create_time, message_content) +VALUES ( + 200, + 900002, + 21474836529, + 'wxid_test_alice', + 'wxid_test_bob', + 1700000200, + 'nullnullhttps://example.com/test-link5' +); diff --git a/crates/wx-db/tests/fixtures/03-nested-quote-xml.sql b/crates/wx-db/tests/fixtures/03-nested-quote-xml.sql new file mode 100644 index 0000000..9c17a61 --- /dev/null +++ b/crates/wx-db/tests/fixtures/03-nested-quote-xml.sql @@ -0,0 +1,15 @@ +-- Fixture 03: Nested quote XML +-- Tests that when inside contains nested , +-- the title is extracted from the inner XML rather than showing raw XML. +-- local_type = (57 << 32) | 49 = 244813135921 + +INSERT INTO fixture_messages (sort_seq, server_id, local_type, sender, talker, create_time, message_content) +VALUES ( + 300, + 900003, + 244813135921, + 'wxid_test_alice', + 'wxid_test_bob', + 1700000300, + 'I agree with thiswxid_test_bobBob49Shared article about testingA comprehensive guide' +); diff --git a/crates/wx-db/tests/fixtures/04-empty-title-channel.sql b/crates/wx-db/tests/fixtures/04-empty-title-channel.sql new file mode 100644 index 0000000..6c0df1b --- /dev/null +++ b/crates/wx-db/tests/fixtures/04-empty-title-channel.sql @@ -0,0 +1,14 @@ +-- Fixture 04: Empty title with channel video +-- Tests title→des fallback for channel video messages. +-- local_type = (51 << 32) | 49 = 219043332145 + +INSERT INTO fixture_messages (sort_seq, server_id, local_type, sender, talker, create_time, message_content) +VALUES ( + 400, + 900004, + 219043332145, + 'wxid_test_alice', + 'wxid_test_bob', + 1700000400, + 'This is a channel video description51' +); diff --git a/crates/wx-db/tests/fixtures/05-zstd-compressed.sql b/crates/wx-db/tests/fixtures/05-zstd-compressed.sql new file mode 100644 index 0000000..203f76c --- /dev/null +++ b/crates/wx-db/tests/fixtures/05-zstd-compressed.sql @@ -0,0 +1,17 @@ +-- Fixture 05: zstd-compressed message content +-- Tests the zstd decode path when wcdb_ct=4 and message_content is a zstd blob. +-- local_type = (5 << 32) | 49 = 21474836529 +-- The blob is zstd-compressed XML: +-- Test Article LinkThis is a test article description for snapshot testinghttps://example.com/test-article5 + +INSERT INTO fixture_messages (sort_seq, server_id, local_type, sender, talker, create_time, message_content, wcdb_ct) +VALUES ( + 500, + 900005, + 21474836529, + 'wxid_test_alice', + 'wxid_test_bob', + 1700000500, + X'28b52ffd045825040092c81a1b80356d034158e9946ec229cdfe6a05625208b283f83689948b0d18813c8eeef09090ca018d7956722838a4cc400fa302a0778573379267259857e3f048db0a3997745b0a61496df32e73358af5b5add8d7be489157030f13690fa6c76c2f83f5b5483dbdcd5c308f230c08002743d8ead627a410b047526d5ccadd78225e5be419f2b045', + 4 +); diff --git a/crates/wx-db/tests/fixtures/06-group-sender-parsing.sql b/crates/wx-db/tests/fixtures/06-group-sender-parsing.sql new file mode 100644 index 0000000..c42d21e --- /dev/null +++ b/crates/wx-db/tests/fixtures/06-group-sender-parsing.sql @@ -0,0 +1,17 @@ +-- Fixture 06: Group sender parsing +-- Tests group sender prefix extraction from message content. +-- is_group=1, talker is a chatroom, message_content has "wxid:\n" prefix. +-- sender column is set to wxid_test_alice to verify it gets overridden +-- by the content prefix (wxid_test_bob). + +INSERT INTO fixture_messages (sort_seq, server_id, local_type, sender, talker, create_time, message_content, is_group) +VALUES ( + 600, + 900006, + 1, + 'wxid_test_alice', + 'group_test@chatroom', + 1700000600, + 'wxid_test_bob:' || char(10) || 'Hello from the group chat', + 1 +); diff --git a/crates/wx-db/tests/fixtures/07-group-quote-chatusr.sql b/crates/wx-db/tests/fixtures/07-group-quote-chatusr.sql new file mode 100644 index 0000000..2aebd7a --- /dev/null +++ b/crates/wx-db/tests/fixtures/07-group-quote-chatusr.sql @@ -0,0 +1,59 @@ +-- Fixture 07: Group chat quote with tag +-- In group chats, contains the chatroom ID (not the sender), +-- while contains the actual quoted sender's wxid. +-- This fixture reproduces the real XML structure for BUG-2 verification. +-- local_type = (57 << 32) | 49 = 244813135921 +-- is_group = 1 (group chat, sender prefix in content) + +INSERT INTO fixture_messages (sort_seq, server_id, local_type, sender, talker, create_time, message_content, is_group) +VALUES ( + 700, + 900007, + 244813135921, + '', + 'group_test@chatroom', + 1700000700, + 'wxid_test_quoter: + + + + 同意这个观点 + 57 + + + + + + 1 + 2041776388084106207 + group_test@chatroom + wxid_test_hidden + 隐藏用户 + 这是被引用的原始消息 + <msgsource><sequence_id>854104363</sequence_id></msgsource> + 1700000600 + + + wxid_test_quoter + 0 + + 1 + + + +', + 1 +); + +-- Also include a private chat quote for comparison (fromusr is the actual sender) +INSERT INTO fixture_messages (sort_seq, server_id, local_type, sender, talker, create_time, message_content, is_group) +VALUES ( + 701, + 900008, + 244813135921, + 'wxid_test_alice', + 'wxid_test_bob', + 1700000701, + '好的收到57wxid_test_bobBob1明天见面吧', + 0 +); diff --git a/crates/wx-db/tests/fixtures/08-system-revokemsg.sql b/crates/wx-db/tests/fixtures/08-system-revokemsg.sql new file mode 100644 index 0000000..602d723 --- /dev/null +++ b/crates/wx-db/tests/fixtures/08-system-revokemsg.sql @@ -0,0 +1,26 @@ +-- System messages: revokemsg XML should be extracted to readable text; +-- plain-text system messages should pass through unchanged. + +-- Row 1: revokemsg XML (should extract content tag) +INSERT INTO fixture_messages (sort_seq, server_id, local_type, sender, talker, create_time, message_content, is_group) +VALUES (1, 100001, 10000, '', 'group_test@chatroom', 1700000001, + '"测试用户A" 撤回了一条消息0', + 1); + +-- Row 2: plain-text system message (should pass through as-is) +INSERT INTO fixture_messages (sort_seq, server_id, local_type, sender, talker, create_time, message_content, is_group) +VALUES (2, 100002, 10000, '', 'group_test@chatroom', 1700000002, + '你邀请"测试用户B"加入了群聊', + 1); + +-- Row 3: revokemsg XML with CDATA content +INSERT INTO fixture_messages (sort_seq, server_id, local_type, sender, talker, create_time, message_content, is_group) +VALUES (3, 100003, 10000, '', 'wxid_test_private', 1700000003, + '0', + 0); + +-- Row 4: group sender prefix + revokemsg XML (group context) +INSERT INTO fixture_messages (sort_seq, server_id, local_type, sender, talker, create_time, message_content, is_group) +VALUES (4, 100004, 10000, '', 'group_test@chatroom', 1700000004, + '"测试用户D" 撤回了一条消息1700000000', + 1); diff --git a/crates/wx-db/tests/fixtures/empty.sql b/crates/wx-db/tests/fixtures/empty.sql new file mode 100644 index 0000000..e69de29 diff --git a/crates/wx-db/tests/messages-limit-pushdown.rs b/crates/wx-db/tests/messages-limit-pushdown.rs new file mode 100644 index 0000000..9878eb0 --- /dev/null +++ b/crates/wx-db/tests/messages-limit-pushdown.rs @@ -0,0 +1,425 @@ +use std::fs; +use std::path::Path; + +use rusqlite::{params, Connection}; +use tempfile::TempDir; +use wx_db::test_ddl; +use wx_db::{MessageQuery, SortOrder, WechatDb, MSG_TYPE_TEXT}; + +// Msg_29a6db07e8bbdb53f5d54cc3c309f3f1 = md5("wxid_alice") +const ALICE_TABLE: &str = "Msg_29a6db07e8bbdb53f5d54cc3c309f3f1"; + +/// Create a fixture with 10 text messages for wxid_alice spread across 2 shards, +/// plus 1 damaged row in shard 0. This tests cross-shard LIMIT pushdown behavior. +/// +/// Shard 0 (ts=1700000000): sort_seq 10,30,50,70,90 (5 text) + sort_seq 45 (damaged) +/// Shard 1 (ts=1710000000): sort_seq 20,40,60,80,100 (5 text) +/// +/// Global ASC order: 10,20,30,40,50,60,70,80,90,100 +fn create_limit_pushdown_fixture() -> TempDir { + let dir = TempDir::new().unwrap(); + let base = dir.path(); + + // contact/contact.db + let contact_dir = base.join("contact"); + fs::create_dir_all(&contact_dir).unwrap(); + create_minimal_contact_db(&contact_dir.join("contact.db")); + + // session/session.db + let session_dir = base.join("session"); + fs::create_dir_all(&session_dir).unwrap(); + create_minimal_session_db(&session_dir.join("session.db")); + + let msg_dir = base.join("message"); + fs::create_dir_all(&msg_dir).unwrap(); + + create_shard_0(&msg_dir.join("message_0.db")); + create_shard_1(&msg_dir.join("message_1.db")); + + dir +} + +fn create_minimal_contact_db(path: &Path) { + let conn = Connection::open(path).unwrap(); + test_ddl::create_test_contact_table_minimal(&conn); +} + +fn create_minimal_session_db(path: &Path) { + let conn = Connection::open(path).unwrap(); + test_ddl::create_test_session_table(&conn); +} + +fn create_msg_table(conn: &Connection, table: &str) { + conn.execute_batch("CREATE TABLE Timestamp (timestamp INTEGER);") + .unwrap(); + conn.execute_batch("CREATE TABLE Name2Id (rowid INTEGER PRIMARY KEY, user_name TEXT);") + .unwrap(); + conn.execute( + "INSERT INTO Name2Id VALUES (?1, ?2)", + params![1, "wxid_alice"], + ) + .unwrap(); + conn.execute_batch(&format!( + "CREATE TABLE [{table}] ( + sort_seq INTEGER, + server_id INTEGER, + local_type INTEGER, + real_sender_id INTEGER, + create_time INTEGER, + message_content BLOB, + packed_info_data BLOB, + status INTEGER + );" + )) + .unwrap(); +} + +fn insert_text_msg( + conn: &Connection, + table: &str, + sort_seq: i64, + server_id: i64, + create_time: i64, + text: &str, +) { + conn.execute( + &format!("INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8)"), + params![ + sort_seq, + server_id, + 1_u32, // MSG_TYPE_TEXT + 1_i32, + create_time, + text.as_bytes(), + None::>, + 0_i32, + ], + ) + .unwrap(); +} + +fn insert_image_msg( + conn: &Connection, + table: &str, + sort_seq: i64, + server_id: i64, + create_time: i64, +) { + conn.execute( + &format!("INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8)"), + params![ + sort_seq, + server_id, + 3_u32, // MSG_TYPE_IMAGE + 1_i32, + create_time, + b"" as &[u8], + None::>, + 0_i32, + ], + ) + .unwrap(); +} + +/// Shard 0: sort_seq 10,30,50,70,90 (text) + 45 (damaged zstd) +fn create_shard_0(path: &Path) { + let conn = Connection::open(path).unwrap(); + create_msg_table(&conn, ALICE_TABLE); + conn.execute( + "INSERT INTO Timestamp VALUES (?1)", + params![1_700_000_000_i64], + ) + .unwrap(); + + insert_text_msg(&conn, ALICE_TABLE, 10, 1001, 1_700_000_010, "msg-10"); + insert_text_msg(&conn, ALICE_TABLE, 30, 1003, 1_700_000_030, "msg-30"); + insert_text_msg(&conn, ALICE_TABLE, 50, 1005, 1_700_000_050, "msg-50"); + insert_text_msg(&conn, ALICE_TABLE, 70, 1007, 1_700_000_070, "msg-70"); + insert_text_msg(&conn, ALICE_TABLE, 90, 1009, 1_700_000_090, "msg-90"); + + // Damaged zstd row at sort_seq=45 + let damaged_zstd: Vec = vec![0x28, 0xB5, 0x2F, 0xFD, 0xFF, 0xFF, 0xFF]; + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8)", + table = ALICE_TABLE + ), + params![ + 45_i64, + 1045_i64, + 1_u32, + 1_i32, + 1_700_000_045, + damaged_zstd, + None::>, + 0_i32, + ], + ) + .unwrap(); + + // An image message at sort_seq=55 for msg_type filter testing + insert_image_msg(&conn, ALICE_TABLE, 55, 1055, 1_700_000_055); +} + +/// Shard 1: sort_seq 20,40,60,80,100 (text) +fn create_shard_1(path: &Path) { + let conn = Connection::open(path).unwrap(); + create_msg_table(&conn, ALICE_TABLE); + conn.execute( + "INSERT INTO Timestamp VALUES (?1)", + params![1_710_000_000_i64], + ) + .unwrap(); + + insert_text_msg(&conn, ALICE_TABLE, 20, 1002, 1_710_000_020, "msg-20"); + insert_text_msg(&conn, ALICE_TABLE, 40, 1004, 1_710_000_040, "msg-40"); + insert_text_msg(&conn, ALICE_TABLE, 60, 1006, 1_710_000_060, "msg-60"); + insert_text_msg(&conn, ALICE_TABLE, 80, 1008, 1_710_000_080, "msg-80"); + insert_text_msg(&conn, ALICE_TABLE, 100, 1010, 1_710_000_100, "msg-100"); +} + +// ---- Tests ---- + +#[test] +fn limit_pushdown_desc_limit_3() { + let dir = create_limit_pushdown_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db + .query_messages( + &MessageQuery::for_talker("wxid_alice") + .order(SortOrder::Desc) + .limit(3), + ) + .unwrap(); + + assert_eq!(result.items.len(), 3); + let seqs: Vec = result.items.iter().map(|m| m.sort_seq).collect(); + assert_eq!(seqs, vec![100, 90, 80], "top-3 DESC by sort_seq"); +} + +#[test] +fn limit_pushdown_asc_limit_offset() { + let dir = create_limit_pushdown_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // ASC, offset=1, limit=2: skip first, take next 2 + // Global ASC order (text only, but image at 55 too): 10,20,30,40,50,55(img),60,70,80,90,100 + // offset=1 skips 10, take 2 → 20, 30 + let result = db + .query_messages( + &MessageQuery::for_talker("wxid_alice") + .order(SortOrder::Asc) + .limit(2) + .offset(1), + ) + .unwrap(); + + assert_eq!(result.items.len(), 2); + let seqs: Vec = result.items.iter().map(|m| m.sort_seq).collect(); + assert_eq!(seqs, vec![20, 30]); +} + +#[test] +fn limit_pushdown_msg_type_filter() { + let dir = create_limit_pushdown_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // msg_type=1 (text), limit=2, ASC → should get sort_seq 10, 20 (text only, skipping image at 55) + let result = db + .query_messages( + &MessageQuery::for_talker("wxid_alice") + .msg_type(MSG_TYPE_TEXT) + .order(SortOrder::Asc) + .limit(2), + ) + .unwrap(); + + assert_eq!(result.items.len(), 2); + assert!(result.items.iter().all(|m| m.msg_type == MSG_TYPE_TEXT)); + let seqs: Vec = result.items.iter().map(|m| m.sort_seq).collect(); + assert_eq!(seqs, vec![10, 20]); +} + +#[test] +fn limit_pushdown_keyword_disables_pushdown() { + let dir = create_limit_pushdown_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // keyword search still works (falls back to full scan path) + let result = db + .query_messages( + &MessageQuery::for_talker("wxid_alice") + .keyword("msg-30") + .limit(2), + ) + .unwrap(); + + assert_eq!(result.items.len(), 1); + assert_eq!(result.items[0].sort_seq, 30); +} + +#[test] +fn limit_pushdown_with_filtered_count_disables_pushdown() { + let dir = create_limit_pushdown_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db + .query_messages( + &MessageQuery::for_talker("wxid_alice") + .msg_type(MSG_TYPE_TEXT) + .with_filtered_count(true) + .limit(2), + ) + .unwrap(); + + // filtered_count should reflect all text messages, not just the page + assert_eq!(result.items.len(), 2); + assert_eq!( + result.stats.filtered_count, + Some(10), + "10 text messages total" + ); +} + +#[test] +fn limit_pushdown_cross_shard_merge_correctness() { + let dir = create_limit_pushdown_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // Get all messages ASC to verify interleaving is correct + let result = db + .query_messages( + &MessageQuery::for_talker("wxid_alice") + .order(SortOrder::Asc) + .limit(100), + ) + .unwrap(); + + // 12 total rows (6 shard0 + 5 shard1 + 1 damaged), 11 valid, 1 skipped + let seqs: Vec = result.items.iter().map(|m| m.sort_seq).collect(); + // Verify global sort order is correct across shards (including image at 55) + let mut sorted = seqs.clone(); + sorted.sort(); + assert_eq!(seqs, sorted, "global sort order must be maintained"); + + // Verify both shards contributed + assert!(seqs.contains(&10), "shard 0 message present"); + assert!(seqs.contains(&20), "shard 1 message present"); +} + +#[test] +fn limit_pushdown_damaged_row_skipped() { + let dir = create_limit_pushdown_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db + .query_messages(&MessageQuery::for_talker("wxid_alice").limit(100)) + .unwrap(); + + assert_eq!( + result.stats.skipped, 1, + "damaged zstd row should be skipped" + ); + assert!( + result.items.iter().all(|m| m.sort_seq != 45), + "damaged row at sort_seq=45 must not appear" + ); +} + +#[test] +fn limit_pushdown_invariant_total_rows() { + let dir = create_limit_pushdown_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // With pushdown (no keyword, no filtered_count, no anchor) + let result = db + .query_messages( + &MessageQuery::for_talker("wxid_alice") + .order(SortOrder::Desc) + .limit(3), + ) + .unwrap(); + + // With LIMIT pushdown, total_rows reflects scanned rows (may be > items + skipped + // because each shard fetches offset+limit candidates) + assert!( + result.stats.total_rows >= result.items.len() + result.stats.skipped, + "invariant: total_rows ({}) >= items ({}) + skipped ({})", + result.stats.total_rows, + result.items.len(), + result.stats.skipped + ); + + // Huge offset must not panic or produce negative SQL LIMIT + let result_huge = db + .query_messages( + &MessageQuery::for_talker("wxid_alice") + .order(SortOrder::Desc) + .offset(usize::MAX - 1) + .limit(100), + ) + .unwrap(); + assert_eq!(result_huge.items.len(), 0, "huge offset returns empty page"); + + // Without pushdown (full scan via with_filtered_count) + let result_full = db + .query_messages( + &MessageQuery::for_talker("wxid_alice") + .with_filtered_count(true) + .limit(100), + ) + .unwrap(); + + // Full scan without pagination: parsed + skipped == total_rows + assert_eq!( + result_full.items.len() + result_full.stats.skipped, + result_full.stats.total_rows, + "full scan: parsed + skipped == total_rows" + ); +} + +/// Simulates the CLI `query` default path: no keyword, no with_filtered_count, no anchor. +/// Verifies that this path is pushdown-eligible and scans fewer rows than full scan. +#[test] +fn cli_query_default_uses_pushdown() { + let dir = create_limit_pushdown_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // CLI query default path: for_talker + limit + order (no keyword, no with_filtered_count) + let query_pushdown = MessageQuery::for_talker("wxid_alice") + .limit(3) + .order(SortOrder::Desc); + assert!( + query_pushdown.limit_pushdown_eligible(), + "CLI query default path must be pushdown-eligible" + ); + + let result_pushdown = db.query_messages(&query_pushdown).unwrap(); + + // Full scan via with_filtered_count(true) + let result_full = db + .query_messages( + &MessageQuery::for_talker("wxid_alice") + .with_filtered_count(true) + .limit(3) + .order(SortOrder::Desc), + ) + .unwrap(); + + // Pushdown scans strictly fewer rows than full scan + assert!( + result_pushdown.stats.total_rows < result_full.stats.total_rows, + "pushdown total_rows ({}) must be < full scan total_rows ({})", + result_pushdown.stats.total_rows, + result_full.stats.total_rows + ); + + // Both return the same items (same sort_seq sequence) + let seqs_pushdown: Vec = result_pushdown.items.iter().map(|m| m.sort_seq).collect(); + let seqs_full: Vec = result_full.items.iter().map(|m| m.sort_seq).collect(); + assert_eq!( + seqs_pushdown, seqs_full, + "pushdown and full scan must return identical items" + ); +} diff --git a/crates/wx-db/tests/messages-routing.rs b/crates/wx-db/tests/messages-routing.rs new file mode 100644 index 0000000..6778a6a --- /dev/null +++ b/crates/wx-db/tests/messages-routing.rs @@ -0,0 +1,1125 @@ +use std::fs; +use std::path::Path; + +use rusqlite::{params, Connection}; +use tempfile::TempDir; +use wx_db::test_ddl; +use wx_db::{ + encode_packed_info_for_test, MessageContent, MessageQuery, SortOrder, WechatDb, MSG_TYPE_TEXT, +}; + +// ---- Constants ---- + +// Msg_29a6db07e8bbdb53f5d54cc3c309f3f1 = md5("wxid_alice") +const ALICE_TABLE: &str = "Msg_29a6db07e8bbdb53f5d54cc3c309f3f1"; +// Msg_141611b52b72df07b2e0733d9a36d3c9 = md5("group@chatroom") +const GROUP_TABLE: &str = "Msg_141611b52b72df07b2e0733d9a36d3c9"; + +// ---- Fixture helpers ---- + +/// Create the full fixture directory with 2 message shards, contact.db, and session.db. +fn create_fixture() -> TempDir { + let dir = TempDir::new().unwrap(); + let base = dir.path(); + + // contact/contact.db (minimal, required by WechatDb::open) + let contact_dir = base.join("contact"); + fs::create_dir_all(&contact_dir).unwrap(); + create_contact_db(&contact_dir.join("contact.db")); + + // session/session.db (minimal, required by WechatDb::open) + let session_dir = base.join("session"); + fs::create_dir_all(&session_dir).unwrap(); + create_session_db(&session_dir.join("session.db")); + + // message/message_0.db (shard 0: timestamp=1700000000, HAS WCDB_CT column) + let msg_dir = base.join("message"); + fs::create_dir_all(&msg_dir).unwrap(); + create_shard_0(&msg_dir.join("message_0.db")); + + // message/message_1.db (shard 1: timestamp=1710000000, NO WCDB_CT column) + create_shard_1(&msg_dir.join("message_1.db")); + + dir +} + +fn create_contact_db(path: &Path) { + let conn = Connection::open(path).unwrap(); + test_ddl::create_test_contact_table(&conn); + conn.execute( + "INSERT INTO contact (username, alias, remark, nick_name) VALUES (?1, ?2, ?3, ?4)", + params!["wxid_alice", "", "", "Alice"], + ) + .unwrap(); +} + +fn create_session_db(path: &Path) { + let conn = Connection::open(path).unwrap(); + test_ddl::create_test_session_table(&conn); +} + +/// Shard 0: timestamp=1700000000, HAS WCDB_CT_message_content column. +/// Contains: +/// - Name2Id: alice (rowid=1), bob (rowid=2) +/// - Msg table for "wxid_alice" with 3 messages + 1 damaged zstd message +fn create_shard_0(path: &Path) { + let conn = Connection::open(path).unwrap(); + + // Timestamp table + conn.execute_batch("CREATE TABLE Timestamp (timestamp INTEGER);") + .unwrap(); + conn.execute( + "INSERT INTO Timestamp VALUES (?1)", + params![1_700_000_000_i64], + ) + .unwrap(); + + // Name2Id table + conn.execute_batch( + "CREATE TABLE Name2Id ( + rowid INTEGER PRIMARY KEY, + user_name TEXT + );", + ) + .unwrap(); + conn.execute( + "INSERT INTO Name2Id VALUES (?1, ?2)", + params![1, "wxid_alice"], + ) + .unwrap(); + conn.execute( + "INSERT INTO Name2Id VALUES (?1, ?2)", + params![2, "wxid_bob"], + ) + .unwrap(); + + // Msg table for wxid_alice — WITH WCDB_CT column + conn.execute_batch(&format!( + "CREATE TABLE [{table}] ( + sort_seq INTEGER, + server_id INTEGER, + local_type INTEGER, + real_sender_id INTEGER, + create_time INTEGER, + message_content BLOB, + packed_info_data BLOB, + status INTEGER, + WCDB_CT_message_content INTEGER + );", + table = ALICE_TABLE + )) + .unwrap(); + + // Message 1: plain text, type=1 (text), sender=alice (rowid=1) + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + table = ALICE_TABLE + ), + params![ + 100_i64, // sort_seq + 1001_i64, // server_id + 1_u32, // local_type = MSG_TYPE_TEXT + 1_i32, // real_sender_id → Name2Id rowid=1 → "wxid_alice" + 1_700_000_100_i64, // create_time + b"hello world" as &[u8], // message_content (plain text) + None::>, // packed_info_data + 0_i32, // status + None::, // WCDB_CT_message_content (NULL = not compressed) + ], + ) + .unwrap(); + + // Message 2: zstd compressed text (CT=4), type=1 (text), sender=bob (rowid=2) + let compressed = zstd::encode_all(&b"compressed message content"[..], 0).unwrap(); + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + table = ALICE_TABLE + ), + params![ + 200_i64, // sort_seq + 1002_i64, // server_id + 1_u32, // local_type = MSG_TYPE_TEXT + 2_i32, // real_sender_id → Name2Id rowid=2 → "wxid_bob" + 1_700_000_200_i64, // create_time + compressed, // message_content (zstd compressed) + None::>, // packed_info_data + 0_i32, // status + 4_i32, // WCDB_CT_message_content = 4 → zstd + ], + ) + .unwrap(); + + // Message 3: image with packed_info, type=3 (image), sender=alice (rowid=1) + let packed_bytes = encode_packed_info_for_test(Some("abc123def456"), None); + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + table = ALICE_TABLE + ), + params![ + 300_i64, // sort_seq + 1003_i64, // server_id + 3_u32, // local_type = MSG_TYPE_IMAGE + 1_i32, // real_sender_id + 1_700_000_300_i64, // create_time + b"" as &[u8], // message_content (empty for image) + packed_bytes, // packed_info_data + 0_i32, // status + None::, // WCDB_CT_message_content + ], + ) + .unwrap(); + + // Message 4: DAMAGED zstd message — should be skipped + // Starts with zstd magic but is corrupted + let damaged_zstd: Vec = vec![0x28, 0xB5, 0x2F, 0xFD, 0xFF, 0xFF, 0xFF]; + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + table = ALICE_TABLE + ), + params![ + 400_i64, // sort_seq + 1004_i64, // server_id + 1_u32, // local_type = MSG_TYPE_TEXT + 1_i32, // real_sender_id + 1_700_000_400_i64, // create_time + damaged_zstd, // message_content (damaged zstd) + None::>, // packed_info_data + 0_i32, // status + 4_i32, // WCDB_CT = 4 → forces zstd decode attempt + ], + ) + .unwrap(); + + // Message 5: app link message, type=49 sub_type=5, local_type = (5 << 32) | 49 + let link_xml = r#"<![CDATA[Test Article]]>Article descriptionhttps://mp.weixin.qq.com/test"#; + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + table = ALICE_TABLE + ), + params![ + 500_i64, // sort_seq + 1005_i64, // server_id + (5_i64 << 32) | 49, // local_type = (5 << 32) | 49 → msg_type=49, sub_type=5 + 1_i32, // real_sender_id → alice + 1_700_000_500_i64, // create_time + link_xml.as_bytes(), // message_content (XML) + None::>, // packed_info_data + 0_i32, // status + None::, // WCDB_CT_message_content + ], + ) + .unwrap(); +} + +/// Shard 1: timestamp=1710000000, NO WCDB_CT_message_content column. +/// Contains: +/// - Name2Id: charlie (rowid=1) +/// - Msg table for "group@chatroom" with 2 group messages +fn create_shard_1(path: &Path) { + let conn = Connection::open(path).unwrap(); + + // Timestamp table + conn.execute_batch("CREATE TABLE Timestamp (timestamp INTEGER);") + .unwrap(); + conn.execute( + "INSERT INTO Timestamp VALUES (?1)", + params![1_710_000_000_i64], + ) + .unwrap(); + + // Name2Id table + conn.execute_batch( + "CREATE TABLE Name2Id ( + rowid INTEGER PRIMARY KEY, + user_name TEXT + );", + ) + .unwrap(); + conn.execute( + "INSERT INTO Name2Id VALUES (?1, ?2)", + params![1, "wxid_charlie"], + ) + .unwrap(); + + // Msg table for group@chatroom — WITHOUT WCDB_CT column + conn.execute_batch(&format!( + "CREATE TABLE [{table}] ( + sort_seq INTEGER, + server_id INTEGER, + local_type INTEGER, + real_sender_id INTEGER, + create_time INTEGER, + message_content BLOB, + packed_info_data BLOB, + status INTEGER + );", + table = GROUP_TABLE + )) + .unwrap(); + + // Message 1: group message, plain text with sender prefix + let group_msg_1 = b"wxid_sender_a:\nhello from group"; + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8)", + table = GROUP_TABLE + ), + params![ + 500_i64, // sort_seq + 2001_i64, // server_id + 1_u32, // local_type = MSG_TYPE_TEXT + 1_i32, // real_sender_id → charlie + 1_710_000_100_i64, // create_time + &group_msg_1[..], // message_content + None::>, // packed_info_data + 0_i32, // status + ], + ) + .unwrap(); + + // Message 2: group message, zstd compressed (detected by magic bytes, no CT column) + let group_content = b"wxid_sender_b:\ncompressed group message"; + let compressed = zstd::encode_all(&group_content[..], 0).unwrap(); + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8)", + table = GROUP_TABLE + ), + params![ + 600_i64, // sort_seq + 2002_i64, // server_id + 1_u32, // local_type = MSG_TYPE_TEXT + 1_i32, // real_sender_id → charlie + 1_710_000_200_i64, // create_time + compressed, // message_content (zstd compressed, magic bytes detection) + None::>, // packed_info_data + 0_i32, // status + ], + ) + .unwrap(); +} + +// ---- Tests ---- + +#[test] +fn messages_routing_alice_query_hits_shard_0_only() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // Query for wxid_alice with time range covering only shard 0 + let result = db + .query_messages( + &MessageQuery::for_talker("wxid_alice").time_range(1_700_000_000, 1_709_999_999), + ) + .unwrap(); + + // 5 rows total (4 valid + 1 damaged), 4 decoded, 1 skipped + assert_eq!(result.items.len(), 4, "expected 4 valid messages"); + assert_eq!( + result.stats.total_rows, 5, + "expected 5 total rows (incl damaged)" + ); + assert_eq!(result.stats.skipped, 1, "expected 1 skipped (damaged zstd)"); + assert_eq!( + result.items.len() + result.stats.skipped, + result.stats.total_rows, + "invariant: parsed + skipped == total_rows" + ); +} + +#[test] +fn messages_routing_zstd_decompression_with_ct_column() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db + .query_messages( + &MessageQuery::for_talker("wxid_alice").time_range(1_700_000_000, 1_709_999_999), + ) + .unwrap(); + + // Message at sort_seq=200 should be decompressed from zstd (CT=4) + let msg = result.items.iter().find(|m| m.sort_seq == 200).unwrap(); + match &msg.content { + MessageContent::Text(s) => assert_eq!(s, "compressed message content"), + other => panic!("expected Text, got: {:?}", other), + } + assert_eq!( + result.items.len() + result.stats.skipped, + result.stats.total_rows, + "invariant: parsed + skipped == total_rows" + ); +} + +#[test] +fn messages_routing_plain_text_message() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db + .query_messages( + &MessageQuery::for_talker("wxid_alice").time_range(1_700_000_000, 1_709_999_999), + ) + .unwrap(); + + // Message at sort_seq=100 should be plain text + let msg = result.items.iter().find(|m| m.sort_seq == 100).unwrap(); + match &msg.content { + MessageContent::Text(s) => assert_eq!(s, "hello world"), + other => panic!("expected Text, got: {:?}", other), + } + assert_eq!( + result.items.len() + result.stats.skipped, + result.stats.total_rows, + "invariant: parsed + skipped == total_rows" + ); +} + +#[test] +fn messages_routing_name2id_sender_resolved() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db + .query_messages( + &MessageQuery::for_talker("wxid_alice").time_range(1_700_000_000, 1_709_999_999), + ) + .unwrap(); + + // sort_seq=100: real_sender_id=1 → Name2Id "wxid_alice" + let msg100 = result.items.iter().find(|m| m.sort_seq == 100).unwrap(); + assert_eq!(msg100.sender, "wxid_alice"); + + // sort_seq=200: real_sender_id=2 → Name2Id "wxid_bob" + let msg200 = result.items.iter().find(|m| m.sort_seq == 200).unwrap(); + assert_eq!(msg200.sender, "wxid_bob"); + assert_eq!( + result.items.len() + result.stats.skipped, + result.stats.total_rows, + "invariant: parsed + skipped == total_rows" + ); +} + +#[test] +fn messages_routing_image_packed_info_md5() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db + .query_messages( + &MessageQuery::for_talker("wxid_alice").time_range(1_700_000_000, 1_709_999_999), + ) + .unwrap(); + + // sort_seq=300: image message with packed_info containing image_md5 + let msg = result.items.iter().find(|m| m.sort_seq == 300).unwrap(); + assert_eq!(msg.msg_type, 3); // MSG_TYPE_IMAGE + match &msg.content { + MessageContent::Image { md5 } => { + assert_eq!(md5.as_deref(), Some("abc123def456")); + } + other => panic!("expected Image, got: {:?}", other), + } + assert_eq!( + result.items.len() + result.stats.skipped, + result.stats.total_rows, + "invariant: parsed + skipped == total_rows" + ); +} + +#[test] +fn messages_routing_group_sender_from_content() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db + .query_messages( + &MessageQuery::for_talker("group@chatroom").time_range(1_710_000_000, 1_719_999_999), + ) + .unwrap(); + + assert_eq!(result.items.len(), 2); + + // sort_seq=500: sender extracted from "wxid_sender_a:\nhello from group" + let msg500 = result.items.iter().find(|m| m.sort_seq == 500).unwrap(); + assert_eq!(msg500.sender, "wxid_sender_a"); + match &msg500.content { + MessageContent::Text(s) => assert_eq!(s, "hello from group"), + other => panic!("expected Text, got: {:?}", other), + } + + // sort_seq=600: zstd compressed, sender from "wxid_sender_b:\ncompressed group message" + let msg600 = result.items.iter().find(|m| m.sort_seq == 600).unwrap(); + assert_eq!(msg600.sender, "wxid_sender_b"); + match &msg600.content { + MessageContent::Text(s) => assert_eq!(s, "compressed group message"), + other => panic!("expected Text, got: {:?}", other), + } + assert_eq!( + result.items.len() + result.stats.skipped, + result.stats.total_rows, + "invariant: parsed + skipped == total_rows" + ); +} + +#[test] +fn messages_routing_no_wcdb_ct_column_works() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // Shard 1 has no WCDB_CT column — should still work fine + let result = db + .query_messages( + &MessageQuery::for_talker("group@chatroom").time_range(1_710_000_000, 1_719_999_999), + ) + .unwrap(); + + assert_eq!(result.items.len(), 2); + assert_eq!(result.stats.skipped, 0); + assert_eq!( + result.items.len() + result.stats.skipped, + result.stats.total_rows, + "invariant: parsed + skipped == total_rows" + ); +} + +#[test] +fn messages_routing_keyword_filter() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // "hello" should match message at sort_seq=100 ("hello world") + let result = db + .query_messages( + &MessageQuery::for_talker("wxid_alice") + .time_range(1_700_000_000, 1_709_999_999) + .keyword("hello"), + ) + .unwrap(); + + assert_eq!(result.items.len(), 1); + assert_eq!(result.items[0].sort_seq, 100); + assert!( + result.stats.total_rows >= result.items.len() + result.stats.skipped, + "total_rows must be >= parsed + skipped" + ); +} + +#[test] +fn messages_routing_keyword_filter_case_insensitive() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // "HELLO" should match case-insensitively + let result = db + .query_messages( + &MessageQuery::for_talker("wxid_alice") + .time_range(1_700_000_000, 1_709_999_999) + .keyword("HELLO"), + ) + .unwrap(); + + assert_eq!(result.items.len(), 1); + assert_eq!(result.items[0].sort_seq, 100); + assert!( + result.stats.total_rows >= result.items.len() + result.stats.skipped, + "total_rows must be >= parsed + skipped" + ); +} + +#[test] +fn messages_routing_limit_offset_pagination() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // Limit=1: should get only the first message (ASC order) + let result = db + .query_messages( + &MessageQuery::for_talker("wxid_alice") + .time_range(1_700_000_000, 1_709_999_999) + .order(SortOrder::Asc) + .limit(1), + ) + .unwrap(); + assert_eq!(result.items.len(), 1); + assert_eq!(result.items[0].sort_seq, 100); + + // Offset=1, Limit=1: should get the second message + let result = db + .query_messages( + &MessageQuery::for_talker("wxid_alice") + .time_range(1_700_000_000, 1_709_999_999) + .order(SortOrder::Asc) + .limit(1) + .offset(1), + ) + .unwrap(); + assert_eq!(result.items.len(), 1); + assert_eq!(result.items[0].sort_seq, 200); + + // Offset=2, Limit=10: should get messages 3 and 5 (sort_seq 300, 500) + let result = db + .query_messages( + &MessageQuery::for_talker("wxid_alice") + .time_range(1_700_000_000, 1_709_999_999) + .order(SortOrder::Asc) + .limit(10) + .offset(2), + ) + .unwrap(); + assert_eq!(result.items.len(), 2); + assert_eq!(result.items[0].sort_seq, 300); + assert_eq!(result.items[1].sort_seq, 500); + assert!( + result.stats.total_rows >= result.items.len() + result.stats.skipped, + "total_rows must be >= parsed + skipped" + ); +} + +#[test] +fn messages_routing_sorted_by_sort_seq_asc() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db + .query_messages( + &MessageQuery::for_talker("wxid_alice") + .time_range(1_700_000_000, 1_709_999_999) + .order(SortOrder::Asc), + ) + .unwrap(); + + // Verify ascending sort_seq order + let seqs: Vec = result.items.iter().map(|m| m.sort_seq).collect(); + assert_eq!(seqs, vec![100, 200, 300, 500]); + assert_eq!( + result.items.len() + result.stats.skipped, + result.stats.total_rows, + "invariant: parsed + skipped == total_rows" + ); +} + +#[test] +fn messages_routing_damaged_zstd_skipped() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db + .query_messages( + &MessageQuery::for_talker("wxid_alice").time_range(1_700_000_000, 1_709_999_999), + ) + .unwrap(); + + // Damaged zstd message (sort_seq=400) should be skipped + assert_eq!(result.stats.skipped, 1); + assert!( + result.items.iter().all(|m| m.sort_seq != 400), + "damaged message should not appear in results" + ); + assert_eq!( + result.items.len() + result.stats.skipped, + result.stats.total_rows, + "invariant: parsed + skipped == total_rows" + ); +} + +#[test] +fn messages_routing_table_not_in_shard_skips_gracefully() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // Query for a talker that has no Msg table in any shard + let result = db + .query_messages(&MessageQuery::for_talker("nonexistent_wxid")) + .unwrap(); + + assert_eq!(result.items.len(), 0); + assert_eq!(result.stats.total_rows, 0); + assert_eq!(result.stats.skipped, 0); + assert_eq!( + result.items.len() + result.stats.skipped, + result.stats.total_rows, + "invariant: parsed + skipped == total_rows" + ); +} + +#[test] +fn messages_routing_zstd_magic_bytes_detection_no_ct() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // Shard 1 has no CT column but message 2 is zstd-compressed; + // should be detected via magic bytes + let result = db + .query_messages( + &MessageQuery::for_talker("group@chatroom").time_range(1_710_000_000, 1_719_999_999), + ) + .unwrap(); + + let msg600 = result.items.iter().find(|m| m.sort_seq == 600).unwrap(); + match &msg600.content { + MessageContent::Text(s) => assert_eq!(s, "compressed group message"), + other => panic!("expected Text, got: {:?}", other), + } + assert_eq!( + result.items.len() + result.stats.skipped, + result.stats.total_rows, + "invariant: parsed + skipped == total_rows" + ); +} + +#[test] +fn messages_routing_cross_shard_merge_sorted() { + // Create a fixture where the same talker has messages in both shards + let dir = TempDir::new().unwrap(); + let base = dir.path(); + + // contact/contact.db + let contact_dir = base.join("contact"); + fs::create_dir_all(&contact_dir).unwrap(); + { + let conn = Connection::open(contact_dir.join("contact.db")).unwrap(); + test_ddl::create_test_contact_table_minimal(&conn); + } + + // session/session.db + let session_dir = base.join("session"); + fs::create_dir_all(&session_dir).unwrap(); + { + let conn = Connection::open(session_dir.join("session.db")).unwrap(); + test_ddl::create_test_session_table(&conn); + } + + let msg_dir = base.join("message"); + fs::create_dir_all(&msg_dir).unwrap(); + + // Shard 0: timestamp=1700000000 + { + let conn = Connection::open(msg_dir.join("message_0.db")).unwrap(); + conn.execute_batch("CREATE TABLE Timestamp (timestamp INTEGER);") + .unwrap(); + conn.execute( + "INSERT INTO Timestamp VALUES (?1)", + params![1_700_000_000_i64], + ) + .unwrap(); + conn.execute_batch("CREATE TABLE Name2Id (rowid INTEGER PRIMARY KEY, user_name TEXT);") + .unwrap(); + conn.execute( + "INSERT INTO Name2Id VALUES (?1, ?2)", + params![1, "wxid_cross"], + ) + .unwrap(); + conn.execute_batch(&format!( + "CREATE TABLE [{table}] ( + sort_seq INTEGER, + server_id INTEGER, + local_type INTEGER, + real_sender_id INTEGER, + create_time INTEGER, + message_content BLOB, + packed_info_data BLOB, + status INTEGER + );", + table = ALICE_TABLE + )) + .unwrap(); + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8)", + table = ALICE_TABLE + ), + params![ + 10_i64, + 1_i64, + 1_u32, + 1_i32, + 1_700_000_100_i64, + b"shard0 msg" as &[u8], + None::>, + 0_i32, + ], + ) + .unwrap(); + } + + // Shard 1: timestamp=1710000000 + { + let conn = Connection::open(msg_dir.join("message_1.db")).unwrap(); + conn.execute_batch("CREATE TABLE Timestamp (timestamp INTEGER);") + .unwrap(); + conn.execute( + "INSERT INTO Timestamp VALUES (?1)", + params![1_710_000_000_i64], + ) + .unwrap(); + conn.execute_batch("CREATE TABLE Name2Id (rowid INTEGER PRIMARY KEY, user_name TEXT);") + .unwrap(); + conn.execute( + "INSERT INTO Name2Id VALUES (?1, ?2)", + params![1, "wxid_cross"], + ) + .unwrap(); + conn.execute_batch(&format!( + "CREATE TABLE [{table}] ( + sort_seq INTEGER, + server_id INTEGER, + local_type INTEGER, + real_sender_id INTEGER, + create_time INTEGER, + message_content BLOB, + packed_info_data BLOB, + status INTEGER + );", + table = ALICE_TABLE + )) + .unwrap(); + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8)", + table = ALICE_TABLE + ), + params![ + 5_i64, + 2_i64, + 1_u32, + 1_i32, + 1_710_000_100_i64, + b"shard1 msg" as &[u8], + None::>, + 0_i32, + ], + ) + .unwrap(); + } + + let db = WechatDb::open(base).unwrap(); + let result = db + .query_messages(&MessageQuery::for_talker("wxid_alice").order(SortOrder::Asc)) + .unwrap(); + + // Both messages found, sorted by sort_seq ASC (5 before 10) + assert_eq!(result.items.len(), 2); + assert_eq!(result.items[0].sort_seq, 5); + assert_eq!(result.items[1].sort_seq, 10); + assert_eq!( + result.items.len() + result.stats.skipped, + result.stats.total_rows, + "invariant: parsed + skipped == total_rows" + ); +} + +#[test] +fn messages_routing_sorted_by_sort_seq_desc() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // Default order is Desc + let result = db + .query_messages( + &MessageQuery::for_talker("wxid_alice").time_range(1_700_000_000, 1_709_999_999), + ) + .unwrap(); + + let seqs: Vec = result.items.iter().map(|m| m.sort_seq).collect(); + assert_eq!(seqs, vec![500, 300, 200, 100]); + assert_eq!( + result.items.len() + result.stats.skipped, + result.stats.total_rows, + "invariant: parsed + skipped == total_rows" + ); +} + +#[test] +fn messages_routing_msg_type_filter() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // Filter text messages only (sort_seq 100, 200 are text; 300 is image) + let result = db + .query_messages( + &MessageQuery::for_talker("wxid_alice") + .time_range(1_700_000_000, 1_709_999_999) + .msg_type(MSG_TYPE_TEXT) + .with_filtered_count(true), + ) + .unwrap(); + + assert_eq!(result.items.len(), 2); + assert!(result.items.iter().all(|m| m.msg_type == MSG_TYPE_TEXT)); + assert_eq!(result.stats.filtered_count, Some(2)); + assert!( + result.stats.total_rows >= result.items.len() + result.stats.skipped, + "total_rows must be >= parsed + skipped" + ); +} + +#[test] +fn messages_routing_filtered_count_without_opt_in() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // without with_filtered_count, filtered_count is None + let result = db + .query_messages( + &MessageQuery::for_talker("wxid_alice").time_range(1_700_000_000, 1_709_999_999), + ) + .unwrap(); + + assert_eq!(result.stats.filtered_count, None); + assert_eq!( + result.items.len() + result.stats.skipped, + result.stats.total_rows, + "invariant: parsed + skipped == total_rows" + ); +} + +#[test] +fn messages_routing_keyword_with_filtered_count() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db + .query_messages( + &MessageQuery::for_talker("wxid_alice") + .time_range(1_700_000_000, 1_709_999_999) + .keyword("hello") + .with_filtered_count(true), + ) + .unwrap(); + + assert_eq!(result.items.len(), 1); + assert_eq!(result.stats.filtered_count, Some(1)); + assert!( + result.stats.total_rows >= result.items.len() + result.stats.skipped, + "total_rows must be >= parsed + skipped" + ); +} + +#[test] +fn messages_routing_link_message_decoded() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db + .query_messages( + &MessageQuery::for_talker("wxid_alice").time_range(1_700_000_000, 1_709_999_999), + ) + .unwrap(); + + // sort_seq=500: app link message (type=49, sub_type=5) + let msg = result.items.iter().find(|m| m.sort_seq == 500).unwrap(); + assert_eq!(msg.msg_type, 49); + assert_eq!(msg.sub_type, 5); + match &msg.content { + MessageContent::Link { + title, url, des, .. + } => { + assert_eq!(title.as_deref(), Some("Test Article")); + assert_eq!(url.as_deref(), Some("https://mp.weixin.qq.com/test")); + assert_eq!(des.as_deref(), Some("Article description")); + } + other => panic!("expected Link, got: {:?}", other), + } + assert_eq!( + result.items.len() + result.stats.skipped, + result.stats.total_rows, + "invariant: parsed + skipped == total_rows" + ); +} + +#[test] +fn messages_routing_keyword_matches_link_title() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // "Article" should match the link message title + let result = db + .query_messages( + &MessageQuery::for_talker("wxid_alice") + .time_range(1_700_000_000, 1_709_999_999) + .keyword("Article"), + ) + .unwrap(); + + assert_eq!(result.items.len(), 1); + assert_eq!(result.items[0].sort_seq, 500); + assert!( + result.stats.total_rows >= result.items.len() + result.stats.skipped, + "total_rows must be >= parsed + skipped" + ); +} + +#[test] +fn messages_routing_compress_content_quote_decoded() { + // Build a custom fixture with compress_content column + let dir = TempDir::new().unwrap(); + let base = dir.path(); + + // Minimal contact.db + let contact_dir = base.join("contact"); + fs::create_dir_all(&contact_dir).unwrap(); + { + let conn = Connection::open(contact_dir.join("contact.db")).unwrap(); + test_ddl::create_test_contact_table_minimal(&conn); + } + + // Minimal session.db + let session_dir = base.join("session"); + fs::create_dir_all(&session_dir).unwrap(); + { + let conn = Connection::open(session_dir.join("session.db")).unwrap(); + test_ddl::create_test_session_table(&conn); + } + + let msg_dir = base.join("message"); + fs::create_dir_all(&msg_dir).unwrap(); + + // Shard with compress_content column + { + let conn = Connection::open(msg_dir.join("message_0.db")).unwrap(); + conn.execute_batch("CREATE TABLE Timestamp (timestamp INTEGER);") + .unwrap(); + conn.execute( + "INSERT INTO Timestamp VALUES (?1)", + params![1_700_000_000_i64], + ) + .unwrap(); + conn.execute_batch("CREATE TABLE Name2Id (rowid INTEGER PRIMARY KEY, user_name TEXT);") + .unwrap(); + conn.execute( + "INSERT INTO Name2Id VALUES (?1, ?2)", + params![1, "wxid_alice"], + ) + .unwrap(); + + // Create table WITH compress_content column + conn.execute_batch(&format!( + "CREATE TABLE [{table}] ( + sort_seq INTEGER, + server_id INTEGER, + local_type INTEGER, + real_sender_id INTEGER, + create_time INTEGER, + message_content BLOB, + packed_info_data BLOB, + status INTEGER, + WCDB_CT_message_content INTEGER, + compress_content BLOB + );", + table = ALICE_TABLE + )) + .unwrap(); + + // Quote message: type=49/sub_type=57, local_type = (57 << 32) | 49 + // message_content is empty placeholder; real XML is in compress_content (zstd) + let quote_xml = r#"reply text herewxid_bobBoboriginal quoted text1"#; + let compressed = zstd::encode_all(quote_xml.as_bytes(), 0).unwrap(); + + conn.execute( + &format!( + "INSERT INTO [{table}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10)", + table = ALICE_TABLE + ), + params![ + 100_i64, // sort_seq + 5001_i64, // server_id + (57_i64 << 32) | 49, // local_type = (57 << 32) | 49 + 1_i32, // real_sender_id + 1_700_000_100_i64, // create_time + b"" as &[u8], // message_content (empty) + None::>, // packed_info_data + 0_i32, // status + None::, // WCDB_CT + compressed, // compress_content (zstd) + ], + ) + .unwrap(); + } + + let db = WechatDb::open(base).unwrap(); + let result = db + .query_messages( + &MessageQuery::for_talker("wxid_alice").time_range(1_700_000_000, 1_709_999_999), + ) + .unwrap(); + + assert_eq!(result.items.len(), 1); + let msg = &result.items[0]; + assert_eq!(msg.msg_type, 49); + assert_eq!(msg.sub_type, 57); + match &msg.content { + MessageContent::Quote { + reply_text, + refer_sender, + refer_content, + refer_type, + .. + } => { + assert_eq!(reply_text.as_deref(), Some("reply text here")); + assert_eq!(refer_sender.as_deref(), Some("Bob")); + assert_eq!(refer_content.as_deref(), Some("original quoted text")); + assert_eq!(*refer_type, Some(1)); + } + other => panic!("expected Quote, got: {:?}", other), + } + assert_eq!( + result.items.len() + result.stats.skipped, + result.stats.total_rows, + "invariant: parsed + skipped == total_rows" + ); +} + +// --------------------------------------------------------------------------- +// bulk_max_sort_seq +// --------------------------------------------------------------------------- + +#[test] +fn bulk_max_sort_seq_known_session_without_msg_table_returns_zero() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + // "wxid_nobody" has no Msg_* table in any shard, but is a "known" session + let usernames = vec!["wxid_nobody".to_string()]; + let result = db.bulk_max_sort_seq(&usernames); + + assert_eq!( + result.get("wxid_nobody"), + Some(&0), + "known session without Msg_* table should return baseline 0" + ); +} + +#[test] +fn bulk_max_sort_seq_returns_max_across_shards() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let usernames = vec!["wxid_alice".to_string(), "group@chatroom".to_string()]; + let result = db.bulk_max_sort_seq(&usernames); + + // wxid_alice: shard 0 has messages with sort_seq 100, 200, 300, 400 (damaged), 500 (link) + let alice_max = *result.get("wxid_alice").unwrap(); + assert_eq!( + alice_max, 500, + "should get max sort_seq across all rows in shard" + ); + + // group@chatroom: shard 1 has messages with sort_seq 500, 600 + let group_max = *result.get("group@chatroom").unwrap(); + assert_eq!(group_max, 600, "should get max sort_seq from shard 1"); +} + +#[test] +fn bulk_max_sort_seq_empty_input_returns_empty() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + + let result = db.bulk_max_sort_seq(&[]); + assert!(result.is_empty()); +} diff --git a/crates/wx-db/tests/pool-integration.rs b/crates/wx-db/tests/pool-integration.rs new file mode 100644 index 0000000..57c22d5 --- /dev/null +++ b/crates/wx-db/tests/pool-integration.rs @@ -0,0 +1,200 @@ +use std::fs; +use std::path::Path; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::sync::Arc; + +use rusqlite::{params, Connection}; +use tempfile::TempDir; +use wx_db::test_ddl; +use wx_db::{MessageQuery, WechatDb}; + +// Msg_29a6db07e8bbdb53f5d54cc3c309f3f1 = md5("wxid_alice") +const ALICE_TABLE: &str = "Msg_29a6db07e8bbdb53f5d54cc3c309f3f1"; + +fn create_fixture() -> TempDir { + let dir = TempDir::new().unwrap(); + let base = dir.path(); + + let contact_dir = base.join("contact"); + fs::create_dir_all(&contact_dir).unwrap(); + let conn = Connection::open(contact_dir.join("contact.db")).unwrap(); + test_ddl::create_test_contact_table(&conn); + + let session_dir = base.join("session"); + fs::create_dir_all(&session_dir).unwrap(); + let conn = Connection::open(session_dir.join("session.db")).unwrap(); + test_ddl::create_test_session_table(&conn); + + let msg_dir = base.join("message"); + fs::create_dir_all(&msg_dir).unwrap(); + create_shard(&msg_dir.join("message_0.db"), 1_700_000_000); + + dir +} + +fn create_shard(path: &Path, timestamp: i64) { + let conn = Connection::open(path).unwrap(); + conn.execute_batch("CREATE TABLE Timestamp (timestamp INTEGER);") + .unwrap(); + conn.execute("INSERT INTO Timestamp VALUES (?1)", params![timestamp]) + .unwrap(); + conn.execute_batch("CREATE TABLE Name2Id (rowid INTEGER PRIMARY KEY, user_name TEXT);") + .unwrap(); + conn.execute_batch(&format!( + "CREATE TABLE [{ALICE_TABLE}] ( + sort_seq INTEGER, + server_id INTEGER, + local_type INTEGER, + real_sender_id INTEGER, + create_time INTEGER, + message_content BLOB, + packed_info_data BLOB, + status INTEGER + );" + )) + .unwrap(); + conn.execute( + &format!("INSERT INTO [{ALICE_TABLE}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8)"), + params![ + 100_i64, + 1001_i64, + 1_u32, + 0_i32, + 1_700_000_100_i64, + b"hello", + None::>, + 0 + ], + ) + .unwrap(); +} + +// ---- Tests ---- + +#[test] +fn open_without_pool_has_no_pool() { + let dir = create_fixture(); + let db = WechatDb::open(dir.path()).unwrap(); + assert!(db.pool().is_none()); +} + +#[test] +fn open_with_pool_has_pool() { + let dir = create_fixture(); + let db = WechatDb::open_with_pool(dir.path(), |_conn| Ok(())).unwrap(); + assert!(db.pool().is_some()); +} + +#[test] +fn fts_init_callback_is_called() { + let dir = create_fixture(); + let called = Arc::new(AtomicBool::new(false)); + let called_clone = Arc::clone(&called); + + // Create a fake FTS db so the pool tries to open it + let fts_path = dir.path().join("message").join("message_fts.db"); + Connection::open(&fts_path).unwrap(); + + let db = WechatDb::open_with_pool(dir.path(), move |_conn| { + called_clone.store(true, Ordering::SeqCst); + Ok(()) + }) + .unwrap(); + + assert!(db.pool().is_some()); + assert!(called.load(Ordering::SeqCst), "fts_init should be called"); +} + +#[test] +fn pooled_query_returns_same_results_as_non_pooled() { + let dir = create_fixture(); + + let db_no_pool = WechatDb::open(dir.path()).unwrap(); + let db_with_pool = WechatDb::open_with_pool(dir.path(), |_conn| Ok(())).unwrap(); + + let query = MessageQuery::for_talker("wxid_alice"); + let result_no_pool = db_no_pool.query_messages(&query).unwrap(); + let result_with_pool = db_with_pool.query_messages(&query).unwrap(); + + assert_eq!(result_no_pool.items.len(), result_with_pool.items.len()); + assert_eq!(result_no_pool.items.len(), 1); + assert_eq!( + result_no_pool.items[0].server_id, + result_with_pool.items[0].server_id + ); +} + +#[test] +fn reopen_all_pooled_picks_up_changes() { + let dir = create_fixture(); + let mut db = WechatDb::open_with_pool(dir.path(), |_conn| Ok(())).unwrap(); + + // Insert a new message outside the pool + let shard_path = dir.path().join("message").join("message_0.db"); + { + let conn = Connection::open(&shard_path).unwrap(); + conn.execute( + &format!("INSERT INTO [{ALICE_TABLE}] VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8)"), + params![ + 200_i64, + 1002_i64, + 1_u32, + 0_i32, + 1_700_000_200_i64, + b"world", + None::>, + 0 + ], + ) + .unwrap(); + } + + // Before reopen: pool connections are stale (may or may not see new data depending on WAL) + // After reopen: must see the new data + db.reopen_all_pooled().unwrap(); + + let query = MessageQuery::for_talker("wxid_alice"); + let result = db.query_messages(&query).unwrap(); + assert_eq!(result.items.len(), 2); +} + +#[test] +fn pooled_query_works_after_shard_file_removed() { + let dir = create_fixture(); + let db = WechatDb::open_with_pool(dir.path(), |_conn| Ok(())).unwrap(); + + let shard_path = dir.path().join("message").join("message_0.db"); + fs::remove_file(&shard_path).unwrap(); + + let query = MessageQuery::for_talker("wxid_alice"); + let result = db.query_messages(&query).unwrap(); + assert_eq!(result.items.len(), 1); + assert_eq!(result.items[0].server_id, 1001); +} + +#[test] +fn reopen_fts_with_pool() { + let dir = create_fixture(); + + // Create a FTS db file so the pool opens it + let fts_path = dir.path().join("message").join("message_fts.db"); + Connection::open(&fts_path).unwrap(); + + let mut db = WechatDb::open_with_pool(dir.path(), |_conn| Ok(())).unwrap(); + assert!(db.pool().unwrap().fts_conn().is_some()); + + // reopen_fts should succeed + db.reopen_fts().unwrap(); + assert!(db.pool().unwrap().fts_conn().is_some()); +} + +#[test] +fn reopen_fts_without_pool_is_noop() { + let dir = create_fixture(); + let mut db = WechatDb::open(dir.path()).unwrap(); + assert!(db.pool().is_none()); + + // reopen_fts on a db without pool should be a no-op + db.reopen_fts().unwrap(); + assert!(db.pool().is_none()); +} diff --git a/crates/wx-db/tests/snapshot-parsing.rs b/crates/wx-db/tests/snapshot-parsing.rs new file mode 100644 index 0000000..d9936a0 --- /dev/null +++ b/crates/wx-db/tests/snapshot-parsing.rs @@ -0,0 +1,188 @@ +use std::path::Path; + +use rusqlite::types::ValueRef; +use rusqlite::Connection; +use wx_db::{decode_message_for_test, Message}; + +struct FixtureResult { + messages: Vec, + skipped: usize, + total: usize, +} + +fn load_fixture(fixture_name: &str) -> FixtureResult { + let fixture_path = Path::new(env!("CARGO_MANIFEST_DIR")) + .join("tests/fixtures") + .join(fixture_name); + let sql = std::fs::read_to_string(&fixture_path) + .unwrap_or_else(|e| panic!("failed to read fixture {}: {}", fixture_path.display(), e)); + + let conn = Connection::open_in_memory().expect("failed to open in-memory db"); + + conn.execute_batch( + "CREATE TABLE fixture_messages ( + sort_seq INTEGER NOT NULL, + server_id INTEGER NOT NULL, + local_type INTEGER NOT NULL, + sender TEXT NOT NULL DEFAULT '', + talker TEXT NOT NULL, + create_time INTEGER NOT NULL DEFAULT 1700000000, + message_content BLOB NOT NULL, + packed_info_data BLOB, + status INTEGER NOT NULL DEFAULT 0, + wcdb_ct INTEGER, + compress_content BLOB, + is_group INTEGER NOT NULL DEFAULT 0 + );", + ) + .expect("failed to create fixture_messages table"); + + if !sql.trim().is_empty() { + conn.execute_batch(&sql) + .unwrap_or_else(|e| panic!("failed to execute fixture SQL {}: {}", fixture_name, e)); + } + + let mut stmt = conn + .prepare( + "SELECT sort_seq, server_id, local_type, sender, talker, + create_time, message_content, packed_info_data, + status, wcdb_ct, compress_content, is_group + FROM fixture_messages ORDER BY sort_seq", + ) + .expect("failed to prepare SELECT"); + + let mut messages = Vec::new(); + let mut skipped = 0usize; + let mut total = 0usize; + + let mut rows = stmt.query([]).expect("failed to query fixture_messages"); + while let Some(row) = rows.next().expect("failed to advance row") { + total += 1; + + let sort_seq: i64 = row.get(0).unwrap(); + let server_id: i64 = row.get(1).unwrap(); + let local_type: i64 = row.get(2).unwrap(); + let sender: String = row.get(3).unwrap(); + let talker: String = row.get(4).unwrap(); + let create_time: i64 = row.get(5).unwrap(); + + // message_content can be Text or Blob (mirrors production decode_message_row) + let raw_content: Vec = match row.get_ref(6).unwrap() { + ValueRef::Blob(b) => b.to_vec(), + ValueRef::Text(b) => b.to_vec(), + other => panic!("unexpected message_content type: {:?}", other), + }; + + let packed_info_data: Option> = row.get(7).unwrap(); + let status: i32 = row.get(8).unwrap(); + let wcdb_ct: Option = row.get(9).unwrap(); + let compress_content: Option> = row.get(10).unwrap(); + let is_group: bool = row.get(11).unwrap(); + + match decode_message_for_test( + sort_seq, + server_id, + local_type, + &sender, + &talker, + create_time, + &raw_content, + packed_info_data.as_deref(), + status, + wcdb_ct, + compress_content.as_deref(), + is_group, + ) { + Ok(msg) => messages.push(msg), + Err(e) => { + eprintln!("decode error in fixture at sort_seq={}: {}", sort_seq, e); + skipped += 1; + } + } + } + + FixtureResult { + messages, + skipped, + total, + } +} + +fn assert_no_message_loss(result: &FixtureResult) { + assert_eq!( + result.messages.len() + result.skipped, + result.total, + "invariant: parsed + skipped == total rows" + ); +} + +#[test] +fn loader_handles_empty_fixture() { + let result = load_fixture("empty.sql"); + assert_eq!(result.total, 0); + assert_eq!(result.messages.len(), 0); +} + +#[test] +fn snapshot_u64_local_type() { + let result = load_fixture("01-u64-local-type.sql"); + assert_no_message_loss(&result); + assert_eq!(result.skipped, 0, "no rows should be skipped"); + insta::assert_yaml_snapshot!(result.messages); +} + +#[test] +fn snapshot_xml_null_sentinel() { + let result = load_fixture("02-xml-null-sentinel.sql"); + assert_no_message_loss(&result); + assert_eq!(result.skipped, 0, "no rows should be skipped"); + insta::assert_yaml_snapshot!(result.messages); +} + +#[test] +fn snapshot_nested_quote_xml() { + let result = load_fixture("03-nested-quote-xml.sql"); + assert_no_message_loss(&result); + assert_eq!(result.skipped, 0, "no rows should be skipped"); + insta::assert_yaml_snapshot!(result.messages); +} + +#[test] +fn snapshot_empty_title_channel() { + let result = load_fixture("04-empty-title-channel.sql"); + assert_no_message_loss(&result); + assert_eq!(result.skipped, 0, "no rows should be skipped"); + insta::assert_yaml_snapshot!(result.messages); +} + +#[test] +fn snapshot_zstd_compressed() { + let result = load_fixture("05-zstd-compressed.sql"); + assert_no_message_loss(&result); + assert_eq!(result.skipped, 0, "no rows should be skipped"); + insta::assert_yaml_snapshot!(result.messages); +} + +#[test] +fn snapshot_group_sender_parsing() { + let result = load_fixture("06-group-sender-parsing.sql"); + assert_no_message_loss(&result); + assert_eq!(result.skipped, 0, "no rows should be skipped"); + insta::assert_yaml_snapshot!(result.messages); +} + +#[test] +fn snapshot_group_quote_chatusr() { + let result = load_fixture("07-group-quote-chatusr.sql"); + assert_no_message_loss(&result); + assert_eq!(result.skipped, 0, "no rows should be skipped"); + insta::assert_yaml_snapshot!(result.messages); +} + +#[test] +fn snapshot_system_revokemsg() { + let result = load_fixture("08-system-revokemsg.sql"); + assert_no_message_loss(&result); + assert_eq!(result.skipped, 0, "no rows should be skipped"); + insta::assert_yaml_snapshot!(result.messages); +} diff --git a/crates/wx-db/tests/snapshots/snapshot_parsing__snapshot_empty_title_channel.snap b/crates/wx-db/tests/snapshots/snapshot_parsing__snapshot_empty_title_channel.snap new file mode 100644 index 0000000..1bd6f72 --- /dev/null +++ b/crates/wx-db/tests/snapshots/snapshot_parsing__snapshot_empty_title_channel.snap @@ -0,0 +1,18 @@ +--- +source: crates/wx-db/tests/snapshot-parsing.rs +assertion_line: 153 +expression: result.messages +--- +- sort_seq: 400 + server_id: 900004 + msg_type: 49 + sub_type: 51 + sender: wxid_test_alice + talker: wxid_test_bob + create_time: 1700000400 + content: + ChannelVideo: + sub_type: 51 + title: This is a channel video description + raw_xml: "This is a channel video description51" + status: 0 diff --git a/crates/wx-db/tests/snapshots/snapshot_parsing__snapshot_group_quote_chatusr.snap b/crates/wx-db/tests/snapshots/snapshot_parsing__snapshot_group_quote_chatusr.snap new file mode 100644 index 0000000..3120517 --- /dev/null +++ b/crates/wx-db/tests/snapshots/snapshot_parsing__snapshot_group_quote_chatusr.snap @@ -0,0 +1,34 @@ +--- +source: crates/wx-db/tests/snapshot-parsing.rs +expression: result.messages +--- +- sort_seq: 700 + server_id: 900007 + msg_type: 49 + sub_type: 57 + sender: wxid_test_quoter + talker: group_test@chatroom + create_time: 1700000700 + content: + Quote: + reply_text: 同意这个观点 + refer_sender: 隐藏用户 + refer_content: 这是被引用的原始消息 + refer_type: 1 + raw_xml: "\n\n\t\n\t\t同意这个观点\n\t\t57\n\t\t\n\t\t\t\n\t\t\t\n\t\t\n\t\t\n\t\t\t1\n\t\t\t2041776388084106207\n\t\t\tgroup_test@chatroom\n\t\t\twxid_test_hidden\n\t\t\t隐藏用户\n\t\t\t这是被引用的原始消息\n\t\t\t<msgsource><sequence_id>854104363</sequence_id></msgsource>\n\t\t\t1700000600\n\t\t\n\t\n\twxid_test_quoter\n\t0\n\t\n\t\t1\n\t\t\n\t\n\t\n" + status: 0 +- sort_seq: 701 + server_id: 900008 + msg_type: 49 + sub_type: 57 + sender: wxid_test_alice + talker: wxid_test_bob + create_time: 1700000701 + content: + Quote: + reply_text: 好的收到 + refer_sender: Bob + refer_content: 明天见面吧 + refer_type: 1 + raw_xml: "好的收到57wxid_test_bobBob1明天见面吧" + status: 0 diff --git a/crates/wx-db/tests/snapshots/snapshot_parsing__snapshot_group_sender_parsing.snap b/crates/wx-db/tests/snapshots/snapshot_parsing__snapshot_group_sender_parsing.snap new file mode 100644 index 0000000..633142e --- /dev/null +++ b/crates/wx-db/tests/snapshots/snapshot_parsing__snapshot_group_sender_parsing.snap @@ -0,0 +1,15 @@ +--- +source: crates/wx-db/tests/snapshot-parsing.rs +assertion_line: 169 +expression: result.messages +--- +- sort_seq: 600 + server_id: 900006 + msg_type: 1 + sub_type: 0 + sender: wxid_test_bob + talker: group_test@chatroom + create_time: 1700000600 + content: + Text: Hello from the group chat + status: 0 diff --git a/crates/wx-db/tests/snapshots/snapshot_parsing__snapshot_nested_quote_xml.snap b/crates/wx-db/tests/snapshots/snapshot_parsing__snapshot_nested_quote_xml.snap new file mode 100644 index 0000000..5576d32 --- /dev/null +++ b/crates/wx-db/tests/snapshots/snapshot_parsing__snapshot_nested_quote_xml.snap @@ -0,0 +1,20 @@ +--- +source: crates/wx-db/tests/snapshot-parsing.rs +assertion_line: 145 +expression: result.messages +--- +- sort_seq: 300 + server_id: 900003 + msg_type: 49 + sub_type: 57 + sender: wxid_test_alice + talker: wxid_test_bob + create_time: 1700000300 + content: + Quote: + reply_text: I agree with this + refer_sender: Bob + refer_content: Shared article about testing + refer_type: 49 + raw_xml: "I agree with thiswxid_test_bobBob49Shared article about testingA comprehensive guide" + status: 0 diff --git a/crates/wx-db/tests/snapshots/snapshot_parsing__snapshot_system_revokemsg.snap b/crates/wx-db/tests/snapshots/snapshot_parsing__snapshot_system_revokemsg.snap new file mode 100644 index 0000000..996f47c --- /dev/null +++ b/crates/wx-db/tests/snapshots/snapshot_parsing__snapshot_system_revokemsg.snap @@ -0,0 +1,45 @@ +--- +source: crates/wx-db/tests/snapshot-parsing.rs +assertion_line: 187 +expression: result.messages +--- +- sort_seq: 1 + server_id: 100001 + msg_type: 10000 + sub_type: 0 + sender: "" + talker: group_test@chatroom + create_time: 1700000001 + content: + System: "\"测试用户A\" 撤回了一条消息" + status: 0 +- sort_seq: 2 + server_id: 100002 + msg_type: 10000 + sub_type: 0 + sender: "" + talker: group_test@chatroom + create_time: 1700000002 + content: + System: "你邀请\"测试用户B\"加入了群聊" + status: 0 +- sort_seq: 3 + server_id: 100003 + msg_type: 10000 + sub_type: 0 + sender: "" + talker: wxid_test_private + create_time: 1700000003 + content: + System: "\"测试用户C\" 撤回了一条消息" + status: 0 +- sort_seq: 4 + server_id: 100004 + msg_type: 10000 + sub_type: 0 + sender: "" + talker: group_test@chatroom + create_time: 1700000004 + content: + System: "\"测试用户D\" 撤回了一条消息" + status: 0 diff --git a/crates/wx-db/tests/snapshots/snapshot_parsing__snapshot_u64_local_type.snap b/crates/wx-db/tests/snapshots/snapshot_parsing__snapshot_u64_local_type.snap new file mode 100644 index 0000000..bcc77bb --- /dev/null +++ b/crates/wx-db/tests/snapshots/snapshot_parsing__snapshot_u64_local_type.snap @@ -0,0 +1,20 @@ +--- +source: crates/wx-db/tests/snapshot-parsing.rs +assertion_line: 129 +expression: result.messages +--- +- sort_seq: 100 + server_id: 900001 + msg_type: 49 + sub_type: 57 + sender: wxid_test_alice + talker: wxid_test_bob + create_time: 1700000100 + content: + Quote: + reply_text: This is my reply + refer_sender: Bob + refer_content: "Hello, how are you?" + refer_type: 1 + raw_xml: "This is my replywxid_test_bobBobHello, how are you?1" + status: 0 diff --git a/crates/wx-db/tests/snapshots/snapshot_parsing__snapshot_xml_null_sentinel.snap b/crates/wx-db/tests/snapshots/snapshot_parsing__snapshot_xml_null_sentinel.snap new file mode 100644 index 0000000..ba8022e --- /dev/null +++ b/crates/wx-db/tests/snapshots/snapshot_parsing__snapshot_xml_null_sentinel.snap @@ -0,0 +1,20 @@ +--- +source: crates/wx-db/tests/snapshot-parsing.rs +assertion_line: 137 +expression: result.messages +--- +- sort_seq: 200 + server_id: 900002 + msg_type: 49 + sub_type: 5 + sender: wxid_test_alice + talker: wxid_test_bob + create_time: 1700000200 + content: + Link: + sub_type: 5 + title: ~ + des: ~ + url: "https://example.com/test-link" + raw_xml: "nullnullhttps://example.com/test-link5" + status: 0 diff --git a/crates/wx-db/tests/snapshots/snapshot_parsing__snapshot_zstd_compressed.snap b/crates/wx-db/tests/snapshots/snapshot_parsing__snapshot_zstd_compressed.snap new file mode 100644 index 0000000..f992fd2 --- /dev/null +++ b/crates/wx-db/tests/snapshots/snapshot_parsing__snapshot_zstd_compressed.snap @@ -0,0 +1,20 @@ +--- +source: crates/wx-db/tests/snapshot-parsing.rs +assertion_line: 161 +expression: result.messages +--- +- sort_seq: 500 + server_id: 900005 + msg_type: 49 + sub_type: 5 + sender: wxid_test_alice + talker: wxid_test_bob + create_time: 1700000500 + content: + Link: + sub_type: 5 + title: Test Article Link + des: This is a test article description for snapshot testing + url: "https://example.com/test-article" + raw_xml: "Test Article LinkThis is a test article description for snapshot testinghttps://example.com/test-article5\n" + status: 0 diff --git a/crates/wx-decrypt/Cargo.toml b/crates/wx-decrypt/Cargo.toml new file mode 100644 index 0000000..37c3ef2 --- /dev/null +++ b/crates/wx-decrypt/Cargo.toml @@ -0,0 +1,15 @@ +[package] +name = "wx-decrypt" +version.workspace = true +edition.workspace = true + +[dependencies] +aes = "0.8" +cbc = "0.1" +pbkdf2 = { version = "0.12", features = ["sha2"] } +sha2 = "0.10" +hmac = "0.12" +thiserror = "2" + +[dev-dependencies] +tempfile = "3" diff --git a/crates/wx-decrypt/src/db.rs b/crates/wx-decrypt/src/db.rs new file mode 100644 index 0000000..7a6e82f --- /dev/null +++ b/crates/wx-decrypt/src/db.rs @@ -0,0 +1,486 @@ +use std::fs::{self, File}; +use std::io::{BufReader, BufWriter, Read, Write}; +use std::path::Path; + +use crate::error::DecryptError; +use crate::kdf::{derive_enc_key, derive_mac_key}; +use crate::page::decrypt_page; +use crate::params::CryptoParams; + +const SQLITE_HEADER: &[u8; 16] = b"SQLite format 3\0"; + +/// Decrypt an entire WeChat database file. +/// +/// Reads from `input`, writes the decrypted SQLite database to `output`. +/// The `raw_key` is the 32-byte key extracted from WeChat's memory. +pub fn decrypt_db( + input: &Path, + output: &Path, + raw_key: &[u8; 32], + params: &CryptoParams, +) -> Result<(), DecryptError> { + let (mut reader, first_page, file_len) = open_and_read_first_page(input, params)?; + let salt = extract_salt(&first_page, params); + + let enc_key = validate_key(&first_page, raw_key, params).ok_or(DecryptError::IncorrectKey)?; + let mac_key = derive_mac_key(&enc_key, &salt, params); + + let mut writer = create_output(output)?; + decrypt_all_pages( + &mut reader, + &mut writer, + &enc_key, + &mac_key, + &first_page, + file_len, + params, + ) +} + +/// Decrypt a database using a pre-derived encryption key and its associated salt. +/// +/// Skips the 256K-iteration PBKDF2 key derivation — only the 2-iteration MAC +/// key derivation is performed. Returns `SaltMismatch` if the DB's header salt +/// does not match the provided `salt`. +pub fn decrypt_db_direct( + input: &Path, + output: &Path, + enc_key: &[u8; 32], + salt: &[u8; 16], + params: &CryptoParams, +) -> Result<(), DecryptError> { + let (mut reader, first_page, file_len) = open_and_read_first_page(input, params)?; + let db_salt = extract_salt(&first_page, params); + + if db_salt != *salt { + return Err(DecryptError::SaltMismatch); + } + + let mac_key = derive_mac_key(enc_key, salt, params); + + // Validate HMAC on first page. + if !verify_first_page_hmac(&first_page, enc_key, &mac_key, params) { + return Err(DecryptError::IncorrectKey); + } + + let mut writer = create_output(output)?; + decrypt_all_pages( + &mut reader, + &mut writer, + enc_key, + &mac_key, + &first_page, + file_len, + params, + ) +} + +/// Validate whether `raw_key` can decrypt the given first page. +/// +/// On success (HMAC passes), returns `Some(enc_key)` — the derived encryption key. +/// On failure, returns `None`. +pub fn validate_key( + first_page: &[u8], + raw_key: &[u8; 32], + params: &CryptoParams, +) -> Option<[u8; 32]> { + if first_page.len() < params.page_size { + return None; + } + + let salt = extract_salt(first_page, params); + let enc_key = derive_enc_key(raw_key, &salt, params); + let mac_key = derive_mac_key(&enc_key, &salt, params); + + if verify_first_page_hmac(first_page, &enc_key, &mac_key, params) { + Some(enc_key) + } else { + None + } +} + +/// Validate whether a pre-derived `enc_key` matches a DB's first page. +/// +/// Extracts the salt from `first_page`, verifies it matches `salt`, +/// derives only the MAC key (2 iterations), and checks the HMAC. +pub fn validate_enc_key( + first_page: &[u8], + enc_key: &[u8; 32], + salt: &[u8; 16], + params: &CryptoParams, +) -> bool { + if first_page.len() < params.page_size { + return false; + } + + let db_salt = extract_salt(first_page, params); + if db_salt != *salt { + return false; + } + + let mac_key = derive_mac_key(enc_key, &db_salt, params); + verify_first_page_hmac(first_page, enc_key, &mac_key, params) +} + +/// Read the 16-byte salt from the first page of an encrypted database. +pub fn read_db_salt(db_path: &Path) -> Result<[u8; 16], DecryptError> { + let mut f = File::open(db_path)?; + let mut salt = [0u8; 16]; + f.read_exact(&mut salt)?; + + if &salt[..] == b"SQLite format 3\0" { + return Err(DecryptError::AlreadyDecrypted); + } + + Ok(salt) +} + +/// Read the 16-byte salt from the main DB file for a given path. +/// +/// If `path` ends with `-wal`, the corresponding main DB path is derived +/// using `wal_path_to_db_path`. Otherwise, `path` is used directly. +pub fn read_main_db_salt_for_path(path: &Path) -> Result<[u8; 16], DecryptError> { + let name = path.file_name().and_then(|n| n.to_str()).unwrap_or(""); + if name.ends_with("-wal") { + let db_path = crate::wal::wal_path_to_db_path(path)?; + read_db_salt(&db_path) + } else { + read_db_salt(path) + } +} + +// --- Internal helpers --- + +/// Open an encrypted DB, read the first page, and return (reader, first_page, file_len). +fn open_and_read_first_page( + input: &Path, + params: &CryptoParams, +) -> Result<(BufReader, Vec, usize), DecryptError> { + let file_len = fs::metadata(input)?.len() as usize; + if file_len < params.page_size { + return Err(DecryptError::FileTooSmall { + expected: params.page_size, + actual: file_len, + }); + } + + let mut reader = BufReader::new(File::open(input)?); + let mut first_page = vec![0u8; params.page_size]; + reader.read_exact(&mut first_page)?; + + if first_page[..SQLITE_HEADER.len()] == SQLITE_HEADER[..] { + return Err(DecryptError::AlreadyDecrypted); + } + + Ok((reader, first_page, file_len)) +} + +/// Extract the 16-byte salt from the first page. +fn extract_salt(first_page: &[u8], params: &CryptoParams) -> [u8; 16] { + let mut salt = [0u8; 16]; + salt.copy_from_slice(&first_page[..params.salt_size]); + salt +} + +/// Create the output file with buffered writer, ensuring parent dirs exist. +fn create_output(output: &Path) -> Result, DecryptError> { + if let Some(parent) = output.parent() { + fs::create_dir_all(parent)?; + } + Ok(BufWriter::new(File::create(output)?)) +} + +/// Verify the HMAC on the first page using pre-derived keys. +fn verify_first_page_hmac( + first_page: &[u8], + _enc_key: &[u8; 32], + mac_key: &[u8; 32], + params: &CryptoParams, +) -> bool { + use hmac::{Hmac, Mac}; + use sha2::Sha512; + + type HmacSha512 = Hmac; + + let offset = params.salt_size; // page 0 + let hmac_data_end = params.page_size - params.reserve + params.iv_size; + + let mut mac = HmacSha512::new_from_slice(mac_key).expect("HMAC key length is always valid"); + mac.update(&first_page[offset..hmac_data_end]); + mac.update(&1u32.to_le_bytes()); // page 1 (1-indexed) + let calculated = mac.finalize().into_bytes(); + + let stored_hmac = &first_page[hmac_data_end..hmac_data_end + params.hmac_size]; + calculated[..params.hmac_size] == *stored_hmac +} + +/// Decrypt all pages (starting from the already-read first page) and write to output. +fn decrypt_all_pages( + reader: &mut impl Read, + writer: &mut impl Write, + enc_key: &[u8; 32], + mac_key: &[u8; 32], + first_page: &[u8], + file_len: usize, + params: &CryptoParams, +) -> Result<(), DecryptError> { + // Write SQLite header. + writer.write_all(SQLITE_HEADER)?; + + // Decrypt page 0 (salt area replaced by SQLite header above). + let page0_decrypted = decrypt_page(first_page, enc_key, mac_key, 0, params)?; + writer.write_all(&page0_decrypted)?; + + // Process remaining pages. + let total_pages = file_len.div_ceil(params.page_size); + let mut page_buf = vec![0u8; params.page_size]; + + for page_num in 1..total_pages as u32 { + let n = read_full_or_eof(reader, &mut page_buf)?; + if n == 0 { + break; + } + if n < params.page_size { + // Partial trailing page — write as-is. + writer.write_all(&page_buf[..n])?; + break; + } + + // Skip all-zero pages (write as-is). + if page_buf.iter().all(|&b| b == 0) { + writer.write_all(&page_buf)?; + continue; + } + + let decrypted = decrypt_page(&page_buf, enc_key, mac_key, page_num, params)?; + writer.write_all(&decrypted)?; + } + + writer.flush()?; + Ok(()) +} + +/// Read exactly `buf.len()` bytes, returning the count read (may be less at EOF). +fn read_full_or_eof(reader: &mut impl Read, buf: &mut [u8]) -> Result { + let mut total = 0; + while total < buf.len() { + match reader.read(&mut buf[total..]) { + Ok(0) => break, + Ok(n) => total += n, + Err(e) if e.kind() == std::io::ErrorKind::Interrupted => continue, + Err(e) => return Err(e), + } + } + Ok(total) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::kdf::{derive_enc_key, derive_mac_key}; + use crate::params::MACOS_4_1_7_31; + + /// Build a fake encrypted page (reverse of decrypt_page). + fn build_encrypted_page( + content: &[u8], + enc_key: &[u8; 32], + mac_key: &[u8; 32], + page_num: u32, + params: &CryptoParams, + salt: Option<&[u8; 16]>, + ) -> Vec { + use aes::cipher::{block_padding::NoPadding, BlockEncryptMut, KeyIvInit}; + use hmac::{Hmac, Mac}; + use sha2::Sha512; + + type HmacSha512 = Hmac; + type Aes256CbcEnc = cbc::Encryptor; + + let iv = [0x42u8; 16]; + let offset = if page_num == 0 { params.salt_size } else { 0 }; + + let data_len = params.page_size - params.reserve - offset; + let mut plaintext = vec![0u8; data_len]; + let copy_len = content.len().min(data_len); + plaintext[..copy_len].copy_from_slice(&content[..copy_len]); + + let mut ciphertext = plaintext.clone(); + Aes256CbcEnc::new(enc_key.into(), (&iv).into()) + .encrypt_padded_mut::(&mut ciphertext, data_len) + .expect("encryption should not fail"); + + let mut page = vec![0u8; params.page_size]; + if page_num == 0 { + page[..params.salt_size].copy_from_slice(salt.unwrap()); + page[offset..offset + data_len].copy_from_slice(&ciphertext); + } else { + page[..data_len].copy_from_slice(&ciphertext); + } + + let iv_start = params.page_size - params.reserve; + page[iv_start..iv_start + params.iv_size].copy_from_slice(&iv); + + let hmac_data_end = params.page_size - params.reserve + params.iv_size; + let mut mac = HmacSha512::new_from_slice(mac_key).unwrap(); + mac.update(&page[offset..hmac_data_end]); + mac.update(&(page_num + 1).to_le_bytes()); + let hmac_result = mac.finalize().into_bytes(); + page[hmac_data_end..hmac_data_end + params.hmac_size] + .copy_from_slice(&hmac_result[..params.hmac_size]); + + page + } + + fn setup_encrypted_db(dir: &std::path::Path) -> ([u8; 32], [u8; 32], [u8; 32], [u8; 16]) { + let params = &MACOS_4_1_7_31; + let raw_key = [0xABu8; 32]; + let salt = [0x01u8; 16]; + let enc_key = derive_enc_key(&raw_key, &salt, params); + let mac_key = derive_mac_key(&enc_key, &salt, params); + + let page0 = + build_encrypted_page(b"hello-page0", &enc_key, &mac_key, 0, params, Some(&salt)); + let page1 = build_encrypted_page(b"hello-page1", &enc_key, &mac_key, 1, params, None); + std::fs::write(dir.join("test.db"), [page0, page1].concat()).unwrap(); + + (raw_key, enc_key, mac_key, salt) + } + + #[test] + fn test_decrypt_db_direct_produces_identical_output() { + let params = &MACOS_4_1_7_31; + let dir = tempfile::tempdir().unwrap(); + let (raw_key, enc_key, _mac_key, salt) = setup_encrypted_db(dir.path()); + + let out_raw = dir.path().join("out_raw.db"); + let out_direct = dir.path().join("out_direct.db"); + let input = dir.path().join("test.db"); + + decrypt_db(&input, &out_raw, &raw_key, params).unwrap(); + decrypt_db_direct(&input, &out_direct, &enc_key, &salt, params).unwrap(); + + assert_eq!( + std::fs::read(&out_raw).unwrap(), + std::fs::read(&out_direct).unwrap(), + "raw and direct decrypt should produce identical output" + ); + } + + #[test] + fn test_decrypt_db_direct_wrong_salt_returns_salt_mismatch() { + let params = &MACOS_4_1_7_31; + let dir = tempfile::tempdir().unwrap(); + let (_raw_key, enc_key, _mac_key, _salt) = setup_encrypted_db(dir.path()); + + let wrong_salt = [0x99u8; 16]; + let input = dir.path().join("test.db"); + let output = dir.path().join("out.db"); + + let err = decrypt_db_direct(&input, &output, &enc_key, &wrong_salt, params).unwrap_err(); + assert!(matches!(err, DecryptError::SaltMismatch)); + } + + #[test] + fn test_validate_enc_key_correct() { + let params = &MACOS_4_1_7_31; + let raw_key = [0xABu8; 32]; + let salt = [0x01u8; 16]; + let enc_key = derive_enc_key(&raw_key, &salt, params); + let mac_key = derive_mac_key(&enc_key, &salt, params); + + let first_page = build_encrypted_page(b"test", &enc_key, &mac_key, 0, params, Some(&salt)); + assert!(validate_enc_key(&first_page, &enc_key, &salt, params)); + } + + #[test] + fn test_validate_enc_key_wrong_salt() { + let params = &MACOS_4_1_7_31; + let raw_key = [0xABu8; 32]; + let salt = [0x01u8; 16]; + let enc_key = derive_enc_key(&raw_key, &salt, params); + let mac_key = derive_mac_key(&enc_key, &salt, params); + + let first_page = build_encrypted_page(b"test", &enc_key, &mac_key, 0, params, Some(&salt)); + let wrong_salt = [0x99u8; 16]; + assert!(!validate_enc_key( + &first_page, + &enc_key, + &wrong_salt, + params + )); + } + + #[test] + fn test_validate_enc_key_wrong_key() { + let params = &MACOS_4_1_7_31; + let raw_key = [0xABu8; 32]; + let salt = [0x01u8; 16]; + let enc_key = derive_enc_key(&raw_key, &salt, params); + let mac_key = derive_mac_key(&enc_key, &salt, params); + + let first_page = build_encrypted_page(b"test", &enc_key, &mac_key, 0, params, Some(&salt)); + let wrong_key = [0xCDu8; 32]; + assert!(!validate_enc_key(&first_page, &wrong_key, &salt, params)); + } + + #[test] + fn test_read_db_salt() { + let dir = tempfile::tempdir().unwrap(); + let db_path = dir.path().join("test.db"); + let mut data = vec![0u8; 4096]; + data[..16].copy_from_slice(&[0x01u8; 16]); + std::fs::write(&db_path, &data).unwrap(); + + let salt = read_db_salt(&db_path).unwrap(); + assert_eq!(salt, [0x01u8; 16]); + } + + #[test] + fn test_read_db_salt_already_decrypted() { + let dir = tempfile::tempdir().unwrap(); + let db_path = dir.path().join("test.db"); + let mut data = vec![0u8; 4096]; + data[..16].copy_from_slice(b"SQLite format 3\0"); + std::fs::write(&db_path, &data).unwrap(); + + let err = read_db_salt(&db_path).unwrap_err(); + assert!(matches!(err, DecryptError::AlreadyDecrypted)); + } + + #[test] + fn test_read_main_db_salt_for_db_path() { + let dir = tempfile::tempdir().unwrap(); + let db_path = dir.path().join("session.db"); + let mut data = vec![0u8; 4096]; + data[..16].copy_from_slice(&[0x42u8; 16]); + std::fs::write(&db_path, &data).unwrap(); + + let salt = read_main_db_salt_for_path(&db_path).unwrap(); + assert_eq!(salt, [0x42u8; 16]); + } + + #[test] + fn test_read_main_db_salt_for_wal_path() { + let dir = tempfile::tempdir().unwrap(); + let db_path = dir.path().join("session.db"); + let mut data = vec![0u8; 4096]; + data[..16].copy_from_slice(&[0x55u8; 16]); + std::fs::write(&db_path, data).unwrap(); + + let wal_path = dir.path().join("session.db-wal"); + std::fs::write(&wal_path, [0u8; 32]).unwrap(); // WAL file exists + + let salt = read_main_db_salt_for_path(&wal_path).unwrap(); + assert_eq!(salt, [0x55u8; 16]); + } + + #[test] + fn test_read_main_db_salt_for_missing_db_returns_error() { + let dir = tempfile::tempdir().unwrap(); + let wal_path = dir.path().join("missing.db-wal"); + std::fs::write(&wal_path, [0u8; 32]).unwrap(); + + let err = read_main_db_salt_for_path(&wal_path).unwrap_err(); + assert!(matches!(err, DecryptError::Io(_))); + } +} diff --git a/crates/wx-decrypt/src/dispatch.rs b/crates/wx-decrypt/src/dispatch.rs new file mode 100644 index 0000000..6b433a4 --- /dev/null +++ b/crates/wx-decrypt/src/dispatch.rs @@ -0,0 +1,282 @@ +//! KeyMaterial dispatch helpers. +//! +//! Centralises the three-arm `KeyMaterial` match that selects between +//! `RawKey` → full KDF path, `EncKey` → direct path, and `EncKeys` → +//! salt-lookup + direct path. + +use std::path::Path; + +use crate::error::DecryptError; +use crate::key_material::KeyMaterial; +use crate::params::CryptoParams; + +/// Decrypt a database file, dispatching on the [`KeyMaterial`] variant. +pub fn dispatch_decrypt_db( + src: &Path, + dst: &Path, + km: &KeyMaterial, + params: &CryptoParams, +) -> Result<(), DecryptError> { + match km { + KeyMaterial::RawKey(key) => crate::db::decrypt_db(src, dst, key, params), + KeyMaterial::EncKey { key, salt } => { + crate::db::decrypt_db_direct(src, dst, key, salt, params) + } + KeyMaterial::EncKeys(pairs) => { + let db_salt = crate::db::read_main_db_salt_for_path(src)?; + let pair = pairs + .iter() + .find(|p| p.salt == db_salt) + .ok_or(DecryptError::NoMatchingEncKey)?; + crate::db::decrypt_db_direct(src, dst, &pair.key, &pair.salt, params) + } + } +} + +/// Decrypt a WAL file and patch it into the decrypted database, +/// dispatching on the [`KeyMaterial`] variant. +pub fn dispatch_decrypt_wal( + wal: &Path, + dst: &Path, + km: &KeyMaterial, + params: &CryptoParams, +) -> Result { + match km { + KeyMaterial::RawKey(key) => crate::wal::decrypt_wal(wal, dst, key, params), + KeyMaterial::EncKey { key, salt } => { + crate::wal::decrypt_wal_direct(wal, dst, key, salt, params) + } + KeyMaterial::EncKeys(pairs) => { + let db_salt = crate::db::read_main_db_salt_for_path(wal)?; + let pair = pairs + .iter() + .find(|p| p.salt == db_salt) + .ok_or(DecryptError::NoMatchingEncKey)?; + crate::wal::decrypt_wal_direct(wal, dst, &pair.key, &pair.salt, params) + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::key_material::EncKeyPair; + use crate::params::MACOS_4_1_7_31; + use tempfile::TempDir; + + // ---- helpers ---- + + fn derive_enc_key(raw_key: &[u8; 32], salt: &[u8; 16], params: &CryptoParams) -> [u8; 32] { + let mut key = [0u8; 32]; + pbkdf2::pbkdf2_hmac::(raw_key, salt, params.kdf_iter, &mut key); + key + } + + /// Build a single-page encrypted DB file that `decrypt_db` can process. + fn build_encrypted_db(path: &Path, raw_key: &[u8; 32], salt: &[u8; 16], params: &CryptoParams) { + use aes::cipher::{BlockEncryptMut, KeyIvInit}; + use hmac::{Hmac, Mac}; + use sha2::Sha512; + + let enc_key = derive_enc_key(raw_key, salt, params); + let mut mac_salt = [0u8; 16]; + for (i, b) in salt.iter().enumerate() { + mac_salt[i] = b ^ 0x3a; + } + let mut mac_key = [0u8; 32]; + pbkdf2::pbkdf2_hmac::(&enc_key, &mac_salt, 2, &mut mac_key); + + let iv = [0x42u8; 16]; + let data_size = params.page_size - params.reserve - params.salt_size; + let plaintext = vec![0u8; data_size]; + + let mut ciphertext = plaintext; + type Aes256CbcEnc = cbc::Encryptor; + Aes256CbcEnc::new((&enc_key).into(), (&iv).into()) + .encrypt_padded_mut::(&mut ciphertext, data_size) + .unwrap(); + + let mut page = Vec::with_capacity(params.page_size); + page.extend_from_slice(salt); + page.extend_from_slice(&ciphertext); + page.extend_from_slice(&iv); + page.resize(params.page_size, 0); + + let hmac_data_end = params.page_size - params.reserve + params.iv_size; + let mut mac = as Mac>::new_from_slice(&mac_key).unwrap(); + mac.update(&page[params.salt_size..hmac_data_end]); + mac.update(&1u32.to_le_bytes()); + let hmac_result = mac.finalize().into_bytes(); + let hmac_start = params.page_size - params.reserve + params.iv_size; + page[hmac_start..hmac_start + params.hmac_size] + .copy_from_slice(&hmac_result[..params.hmac_size]); + + if let Some(parent) = path.parent() { + std::fs::create_dir_all(parent).unwrap(); + } + std::fs::write(path, &page).unwrap(); + } + + // ---- dispatch_decrypt_db tests ---- + + #[test] + fn dispatch_db_raw_key() { + let raw_key = [0xABu8; 32]; + let salt = [0x01u8; 16]; + let params = &MACOS_4_1_7_31; + + let tmp = TempDir::new().unwrap(); + let src = tmp.path().join("test.db"); + let dst = tmp.path().join("out.db"); + build_encrypted_db(&src, &raw_key, &salt, params); + + let km = KeyMaterial::RawKey(raw_key); + dispatch_decrypt_db(&src, &dst, &km, params).unwrap(); + + let data = std::fs::read(&dst).unwrap(); + assert_eq!(&data[..16], b"SQLite format 3\0"); + } + + #[test] + fn dispatch_db_enc_key() { + let raw_key = [0xABu8; 32]; + let salt = [0x01u8; 16]; + let params = &MACOS_4_1_7_31; + let enc_key = derive_enc_key(&raw_key, &salt, params); + + let tmp = TempDir::new().unwrap(); + let src = tmp.path().join("test.db"); + let dst = tmp.path().join("out.db"); + build_encrypted_db(&src, &raw_key, &salt, params); + + let km = KeyMaterial::EncKey { key: enc_key, salt }; + dispatch_decrypt_db(&src, &dst, &km, params).unwrap(); + + let data = std::fs::read(&dst).unwrap(); + assert_eq!(&data[..16], b"SQLite format 3\0"); + } + + #[test] + fn dispatch_db_enc_keys() { + let raw_key = [0xABu8; 32]; + let salt1 = [0x01u8; 16]; + let salt2 = [0x02u8; 16]; + let params = &MACOS_4_1_7_31; + let enc_key1 = derive_enc_key(&raw_key, &salt1, params); + let enc_key2 = derive_enc_key(&raw_key, &salt2, params); + + let tmp = TempDir::new().unwrap(); + // DB with salt1 + let src = tmp.path().join("test.db"); + let dst = tmp.path().join("out.db"); + build_encrypted_db(&src, &raw_key, &salt1, params); + + let km = KeyMaterial::EncKeys(vec![ + EncKeyPair { + key: enc_key2, + salt: salt2, + }, + EncKeyPair { + key: enc_key1, + salt: salt1, + }, + ]); + dispatch_decrypt_db(&src, &dst, &km, params).unwrap(); + + let data = std::fs::read(&dst).unwrap(); + assert_eq!(&data[..16], b"SQLite format 3\0"); + } + + #[test] + fn dispatch_db_enc_keys_no_match() { + let raw_key = [0xABu8; 32]; + let salt_db = [0x01u8; 16]; + let salt_wrong = [0x99u8; 16]; + let params = &MACOS_4_1_7_31; + let enc_key_wrong = derive_enc_key(&raw_key, &salt_wrong, params); + + let tmp = TempDir::new().unwrap(); + let src = tmp.path().join("test.db"); + let dst = tmp.path().join("out.db"); + build_encrypted_db(&src, &raw_key, &salt_db, params); + + let km = KeyMaterial::EncKeys(vec![EncKeyPair { + key: enc_key_wrong, + salt: salt_wrong, + }]); + let err = dispatch_decrypt_db(&src, &dst, &km, params).unwrap_err(); + assert!(matches!(err, DecryptError::NoMatchingEncKey)); + } + + // ---- dispatch_decrypt_wal tests ---- + + #[test] + fn dispatch_wal_raw_key_no_frames() { + let raw_key = [0xABu8; 32]; + let salt = [0x01u8; 16]; + let params = &MACOS_4_1_7_31; + + let tmp = TempDir::new().unwrap(); + // Build a decrypted DB first + let enc_db = tmp.path().join("enc.db"); + let dec_db = tmp.path().join("dec.db"); + build_encrypted_db(&enc_db, &raw_key, &salt, params); + crate::db::decrypt_db(&enc_db, &dec_db, &raw_key, params).unwrap(); + + // Create a minimal WAL with header only (no frames) + let wal = tmp.path().join("enc.db-wal"); + let mut wal_header = [0u8; 32]; + wal_header[0..4].copy_from_slice(&0x377f_0682u32.to_be_bytes()); + std::fs::write(&wal, wal_header).unwrap(); + + let km = KeyMaterial::RawKey(raw_key); + let n = dispatch_decrypt_wal(&wal, &dec_db, &km, params).unwrap(); + assert_eq!(n, 0); + } + + #[test] + fn dispatch_wal_enc_key_no_frames() { + let raw_key = [0xABu8; 32]; + let salt = [0x01u8; 16]; + let params = &MACOS_4_1_7_31; + let enc_key = derive_enc_key(&raw_key, &salt, params); + + let tmp = TempDir::new().unwrap(); + let enc_db = tmp.path().join("enc.db"); + let dec_db = tmp.path().join("dec.db"); + build_encrypted_db(&enc_db, &raw_key, &salt, params); + crate::db::decrypt_db(&enc_db, &dec_db, &raw_key, params).unwrap(); + + let wal = tmp.path().join("enc.db-wal"); + let mut wal_header = [0u8; 32]; + wal_header[0..4].copy_from_slice(&0x377f_0682u32.to_be_bytes()); + std::fs::write(&wal, wal_header).unwrap(); + + let km = KeyMaterial::EncKey { key: enc_key, salt }; + let n = dispatch_decrypt_wal(&wal, &dec_db, &km, params).unwrap(); + assert_eq!(n, 0); + } + + #[test] + fn dispatch_wal_enc_keys_no_frames() { + let raw_key = [0xABu8; 32]; + let salt = [0x01u8; 16]; + let params = &MACOS_4_1_7_31; + let enc_key = derive_enc_key(&raw_key, &salt, params); + + let tmp = TempDir::new().unwrap(); + let enc_db = tmp.path().join("enc.db"); + let dec_db = tmp.path().join("dec.db"); + build_encrypted_db(&enc_db, &raw_key, &salt, params); + crate::db::decrypt_db(&enc_db, &dec_db, &raw_key, params).unwrap(); + + let wal = tmp.path().join("enc.db-wal"); + let mut wal_header = [0u8; 32]; + wal_header[0..4].copy_from_slice(&0x377f_0682u32.to_be_bytes()); + std::fs::write(&wal, wal_header).unwrap(); + + let km = KeyMaterial::EncKeys(vec![EncKeyPair { key: enc_key, salt }]); + let n = dispatch_decrypt_wal(&wal, &dec_db, &km, params).unwrap(); + assert_eq!(n, 0); + } +} diff --git a/crates/wx-decrypt/src/error.rs b/crates/wx-decrypt/src/error.rs new file mode 100644 index 0000000..f42bed7 --- /dev/null +++ b/crates/wx-decrypt/src/error.rs @@ -0,0 +1,31 @@ +use thiserror::Error; + +#[derive(Debug, Error)] +pub enum DecryptError { + #[error("file too small: expected at least {expected} bytes, got {actual}")] + FileTooSmall { expected: usize, actual: usize }, + + #[error("database is already decrypted (SQLite header detected)")] + AlreadyDecrypted, + + #[error("incorrect key: HMAC verification failed on first page")] + IncorrectKey, + + #[error("salt mismatch: DB salt does not match provided salt")] + SaltMismatch, + + #[error("no matching enc_key found for this DB's salt")] + NoMatchingEncKey, + + #[error("HMAC verification failed on page {page_num}")] + HmacVerificationFailed { page_num: u32 }, + + #[error("AES decryption failed on page {page_num}: {reason}")] + AesDecryptFailed { page_num: u32, reason: String }, + + #[error("invalid WAL header: {reason}")] + InvalidWalHeader { reason: String }, + + #[error("I/O error: {0}")] + Io(#[from] std::io::Error), +} diff --git a/crates/wx-decrypt/src/kdf.rs b/crates/wx-decrypt/src/kdf.rs new file mode 100644 index 0000000..1d6ab5d --- /dev/null +++ b/crates/wx-decrypt/src/kdf.rs @@ -0,0 +1,61 @@ +use pbkdf2::pbkdf2_hmac; +use sha2::Sha512; + +use crate::params::CryptoParams; + +/// Derive the AES-256 encryption key from the raw key and salt. +/// +/// Uses PBKDF2-HMAC-SHA512 with `params.kdf_iter` iterations. +pub fn derive_enc_key(raw_key: &[u8; 32], salt: &[u8; 16], params: &CryptoParams) -> [u8; 32] { + let mut enc_key = [0u8; 32]; + pbkdf2_hmac::(raw_key, salt, params.kdf_iter, &mut enc_key); + enc_key +} + +/// Derive the HMAC key from the encryption key and salt. +/// +/// Uses PBKDF2-HMAC-SHA512 with 2 iterations. The MAC salt is `salt XOR 0x3a`. +pub fn derive_mac_key(enc_key: &[u8; 32], salt: &[u8; 16], _params: &CryptoParams) -> [u8; 32] { + let mac_salt: [u8; 16] = std::array::from_fn(|i| salt[i] ^ 0x3a); + let mut mac_key = [0u8; 32]; + pbkdf2_hmac::(enc_key, &mac_salt, 2, &mut mac_key); + mac_key +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::params::MACOS_4_1_7_31; + + #[test] + fn test_derive_keys_deterministic() { + let raw_key = [0xABu8; 32]; + let salt = [0x01u8; 16]; + + let enc1 = derive_enc_key(&raw_key, &salt, &MACOS_4_1_7_31); + let enc2 = derive_enc_key(&raw_key, &salt, &MACOS_4_1_7_31); + assert_eq!(enc1, enc2); + + let mac1 = derive_mac_key(&enc1, &salt, &MACOS_4_1_7_31); + let mac2 = derive_mac_key(&enc2, &salt, &MACOS_4_1_7_31); + assert_eq!(mac1, mac2); + } + + #[test] + fn test_derive_keys_different_salt() { + let raw_key = [0xABu8; 32]; + let salt_a = [0x01u8; 16]; + let salt_b = [0x02u8; 16]; + + let enc_a = derive_enc_key(&raw_key, &salt_a, &MACOS_4_1_7_31); + let enc_b = derive_enc_key(&raw_key, &salt_b, &MACOS_4_1_7_31); + assert_ne!(enc_a, enc_b); + } + + #[test] + fn test_mac_salt_xor() { + let salt = [0x00u8; 16]; + let mac_salt: [u8; 16] = std::array::from_fn(|i| salt[i] ^ 0x3a); + assert!(mac_salt.iter().all(|&b| b == 0x3a)); + } +} diff --git a/crates/wx-decrypt/src/key_material.rs b/crates/wx-decrypt/src/key_material.rs new file mode 100644 index 0000000..b41f3f2 --- /dev/null +++ b/crates/wx-decrypt/src/key_material.rs @@ -0,0 +1,19 @@ +/// A pre-derived encryption key paired with its DB salt. +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +pub struct EncKeyPair { + pub key: [u8; 32], + pub salt: [u8; 16], +} + +/// Represents either a raw LLDB-extracted key or a pre-derived encryption key +/// found via Mach VM memory scanning. +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum KeyMaterial { + /// 32-byte raw key from LLDB capture; requires full PBKDF2 derivation. + RawKey([u8; 32]), + /// Pre-derived encryption key + DB salt; skips the 256K-iteration PBKDF2. + EncKey { key: [u8; 32], salt: [u8; 16] }, + /// Multiple pre-derived enc_keys, each paired with a different DB salt. + /// Used when `key scan` finds keys for multiple DBs within the same account. + EncKeys(Vec), +} diff --git a/crates/wx-decrypt/src/lib.rs b/crates/wx-decrypt/src/lib.rs new file mode 100644 index 0000000..8fab9dd --- /dev/null +++ b/crates/wx-decrypt/src/lib.rs @@ -0,0 +1,18 @@ +pub mod db; +pub mod dispatch; +pub mod error; +pub mod kdf; +pub mod key_material; +pub mod page; +pub mod params; +pub mod wal; + +pub use db::{ + decrypt_db, decrypt_db_direct, read_db_salt, read_main_db_salt_for_path, validate_enc_key, + validate_key, +}; +pub use dispatch::{dispatch_decrypt_db, dispatch_decrypt_wal}; +pub use error::DecryptError; +pub use key_material::{EncKeyPair, KeyMaterial}; +pub use params::{CryptoParams, MACOS_4_1_7_31}; +pub use wal::{decrypt_wal, decrypt_wal_direct}; diff --git a/crates/wx-decrypt/src/page.rs b/crates/wx-decrypt/src/page.rs new file mode 100644 index 0000000..6acd757 --- /dev/null +++ b/crates/wx-decrypt/src/page.rs @@ -0,0 +1,83 @@ +use aes::cipher::{block_padding::NoPadding, BlockDecryptMut, KeyIvInit}; +use hmac::{Hmac, Mac}; +use sha2::Sha512; + +use crate::error::DecryptError; +use crate::params::CryptoParams; + +type HmacSha512 = Hmac; +type Aes256CbcDec = cbc::Decryptor; + +/// Verify HMAC for a page and decrypt it. +/// +/// `page_num` is 0-indexed. The HMAC includes the 1-indexed page number as LE u32. +/// +/// For page 0, the first `salt_size` bytes (salt) are skipped in both HMAC and +/// decryption — the caller is responsible for replacing them with the SQLite header. +pub fn decrypt_page( + page_buf: &[u8], + enc_key: &[u8; 32], + mac_key: &[u8; 32], + page_num: u32, + params: &CryptoParams, +) -> Result, DecryptError> { + let offset = if page_num == 0 { params.salt_size } else { 0 }; + + // --- HMAC verification --- + let hmac_data_end = params.page_size - params.reserve + params.iv_size; + + let mut mac = HmacSha512::new_from_slice(mac_key).expect("HMAC key length is always valid"); + mac.update(&page_buf[offset..hmac_data_end]); + // Page number in HMAC is 1-indexed, little-endian u32. + mac.update(&(page_num + 1).to_le_bytes()); + let calculated = mac.finalize().into_bytes(); + + let stored_hmac = &page_buf[hmac_data_end..hmac_data_end + params.hmac_size]; + if calculated[..params.hmac_size] != *stored_hmac { + return Err(DecryptError::HmacVerificationFailed { page_num }); + } + + // --- Extract IV from reserve area --- + let iv_start = params.page_size - params.reserve; + let iv = &page_buf[iv_start..iv_start + params.iv_size]; + + // --- AES-256-CBC decrypt --- + let encrypted = &page_buf[offset..params.page_size - params.reserve]; + let mut buf = encrypted.to_vec(); + + Aes256CbcDec::new(enc_key.into(), iv.into()) + .decrypt_padded_mut::(&mut buf) + .map_err(|e| DecryptError::AesDecryptFailed { + page_num, + reason: e.to_string(), + })?; + + // Append the reserve area (IV + HMAC) unchanged. + buf.extend_from_slice(&page_buf[params.page_size - params.reserve..params.page_size]); + + Ok(buf) +} + +#[cfg(test)] +mod tests { + use crate::params::MACOS_4_1_7_31; + + #[test] + fn test_page_size_arithmetic() { + let p = &MACOS_4_1_7_31; + // For a non-zero page: decrypted data + reserve = page_size + let data_len = p.page_size - p.reserve; // 4016 + let total = data_len + p.reserve; // 4096 + assert_eq!(total, p.page_size); + } + + #[test] + fn test_page0_size_arithmetic() { + let p = &MACOS_4_1_7_31; + // For page 0: (page_size - reserve - salt_size) + reserve = page_size - salt_size + let data_len = p.page_size - p.reserve - p.salt_size; // 4000 + let total = data_len + p.reserve; // 4080 + // Caller prepends SQLite header (16 bytes) to reach 4096 + assert_eq!(total + p.salt_size, p.page_size); + } +} diff --git a/crates/wx-decrypt/src/params.rs b/crates/wx-decrypt/src/params.rs new file mode 100644 index 0000000..84f895d --- /dev/null +++ b/crates/wx-decrypt/src/params.rs @@ -0,0 +1,21 @@ +/// Encryption parameters for a specific WeChat version. +pub struct CryptoParams { + pub page_size: usize, + pub kdf_iter: u32, + pub hmac_size: usize, + pub reserve: usize, + pub key_size: usize, + pub salt_size: usize, + pub iv_size: usize, +} + +/// macOS WeChat 4.1.7.31: Apple SEE with PBKDF2-HMAC-SHA512. +pub const MACOS_4_1_7_31: CryptoParams = CryptoParams { + page_size: 4096, + kdf_iter: 256_000, + hmac_size: 64, // SHA-512 output + reserve: 80, // IV(16) + HMAC(64) + key_size: 32, + salt_size: 16, + iv_size: 16, +}; diff --git a/crates/wx-decrypt/src/wal.rs b/crates/wx-decrypt/src/wal.rs new file mode 100644 index 0000000..48c009e --- /dev/null +++ b/crates/wx-decrypt/src/wal.rs @@ -0,0 +1,833 @@ +//! WAL (Write-Ahead Log) decryption for encrypted SQLite databases. +//! +//! SQLite WAL files consist of a 32-byte header followed by N frames. +//! Each frame has a 24-byte frame header and a full page of data. +//! In encrypted WeChat databases, the page data within each frame is +//! encrypted using the same parameters as the main database. + +use std::fs::File; +use std::io::{Read, Seek, SeekFrom, Write}; +use std::path::Path; + +use crate::db::read_db_salt; +use crate::error::DecryptError; +use crate::kdf::{derive_enc_key, derive_mac_key}; +use crate::page::decrypt_page; +use crate::params::CryptoParams; + +/// WAL file header size in bytes. +const WAL_HEADER_SIZE: usize = 32; + +/// WAL frame header size in bytes. +const WAL_FRAME_HEADER_SIZE: usize = 24; + +/// SQLite WAL magic numbers (big-endian and little-endian checksum variants). +const WAL_MAGIC_BE: u32 = 0x377f_0682; +const WAL_MAGIC_LE: u32 = 0x377f_0683; + +/// Maximum valid page number (sanity check). +const MAX_PAGE_NUMBER: u32 = 1_000_000; + +/// Parsed WAL file header. +#[derive(Debug, Clone)] +struct WalHeader { + salt1: u32, + salt2: u32, +} + +/// Parsed WAL frame header. +#[derive(Debug, Clone)] +struct WalFrameHeader { + page_number: u32, + commit_size: u32, + salt1: u32, + salt2: u32, +} + +impl WalHeader { + fn parse(buf: &[u8; WAL_HEADER_SIZE]) -> Result { + let magic = u32::from_be_bytes([buf[0], buf[1], buf[2], buf[3]]); + if magic != WAL_MAGIC_BE && magic != WAL_MAGIC_LE { + return Err(DecryptError::InvalidWalHeader { + reason: format!("bad magic: 0x{magic:08x}"), + }); + } + + let salt1 = u32::from_be_bytes([buf[16], buf[17], buf[18], buf[19]]); + let salt2 = u32::from_be_bytes([buf[20], buf[21], buf[22], buf[23]]); + + Ok(WalHeader { salt1, salt2 }) + } +} + +impl WalFrameHeader { + fn parse(buf: &[u8; WAL_FRAME_HEADER_SIZE]) -> Self { + WalFrameHeader { + page_number: u32::from_be_bytes([buf[0], buf[1], buf[2], buf[3]]), + commit_size: u32::from_be_bytes([buf[4], buf[5], buf[6], buf[7]]), + salt1: u32::from_be_bytes([buf[8], buf[9], buf[10], buf[11]]), + salt2: u32::from_be_bytes([buf[12], buf[13], buf[14], buf[15]]), + } + } + + fn is_valid(&self, wal_header: &WalHeader) -> bool { + self.page_number > 0 + && self.page_number <= MAX_PAGE_NUMBER + && self.salt1 == wal_header.salt1 + && self.salt2 == wal_header.salt2 + } +} + +/// Decrypt valid WAL frames and patch them into a decrypted database file. +/// +/// Reads the encrypted WAL file at `wal_path`, decrypts each valid frame's page +/// data, and writes it to the corresponding offset in `decrypted_db_path`. +/// +/// Returns the number of frames successfully patched. +pub fn decrypt_wal( + wal_path: &Path, + decrypted_db_path: &Path, + raw_key: &[u8; 32], + params: &CryptoParams, +) -> Result { + let (mut wal_file, wal_header, wal_len) = open_wal(wal_path)?; + if wal_len < WAL_HEADER_SIZE { + return Ok(0); + } + + let encrypted_db_path = wal_path_to_db_path(wal_path)?; + let salt = read_db_salt(&encrypted_db_path)?; + + let enc_key = derive_enc_key(raw_key, &salt, params); + let mac_key = derive_mac_key(&enc_key, &salt, params); + + let mut db_file = std::fs::OpenOptions::new() + .write(true) + .read(true) + .open(decrypted_db_path)?; + + decrypt_wal_frames( + &mut wal_file, + &mut db_file, + &enc_key, + &mac_key, + &wal_header, + wal_len, + params, + ) +} + +/// Decrypt WAL frames using a pre-derived encryption key and salt. +/// +/// Skips the 256K-iteration PBKDF2 key derivation. +pub fn decrypt_wal_direct( + wal_path: &Path, + decrypted_db_path: &Path, + enc_key: &[u8; 32], + salt: &[u8; 16], + params: &CryptoParams, +) -> Result { + let (mut wal_file, wal_header, wal_len) = open_wal(wal_path)?; + if wal_len < WAL_HEADER_SIZE { + return Ok(0); + } + + let mac_key = derive_mac_key(enc_key, salt, params); + + let mut db_file = std::fs::OpenOptions::new() + .write(true) + .read(true) + .open(decrypted_db_path)?; + + decrypt_wal_frames( + &mut wal_file, + &mut db_file, + enc_key, + &mac_key, + &wal_header, + wal_len, + params, + ) +} + +// --- Internal helpers --- + +/// Open the WAL file, parse its header, and return (file, header, len). +/// Returns Ok with wal_len=0 for empty/too-small WAL files. +fn open_wal(wal_path: &Path) -> Result<(File, WalHeader, usize), DecryptError> { + let wal_len = std::fs::metadata(wal_path)?.len() as usize; + if wal_len < WAL_HEADER_SIZE { + // Return a dummy header — caller checks wal_len and returns Ok(0). + return Ok(( + File::open(wal_path)?, + WalHeader { salt1: 0, salt2: 0 }, + wal_len, + )); + } + + let mut wal_file = File::open(wal_path)?; + let mut wal_hdr_buf = [0u8; WAL_HEADER_SIZE]; + wal_file.read_exact(&mut wal_hdr_buf)?; + let wal_header = WalHeader::parse(&wal_hdr_buf)?; + + Ok((wal_file, wal_header, wal_len)) +} + +/// Process WAL frames: decrypt each valid frame up to the last commit boundary +/// and patch it into the DB file. +/// +/// Uses a two-phase approach to match SQLite WAL reader semantics: +/// - Phase 1: Scan all frame headers to find the last frame with `commit_size > 0`. +/// - Phase 2: Apply frames 0..=last_commit_idx, skipping the rest. +/// +/// Returns `Ok(0)` if there are no committed transactions in the WAL. +fn decrypt_wal_frames( + wal_file: &mut File, + db_file: &mut File, + enc_key: &[u8; 32], + mac_key: &[u8; 32], + wal_header: &WalHeader, + wal_len: usize, + params: &CryptoParams, +) -> Result { + let frame_size = WAL_FRAME_HEADER_SIZE + params.page_size; + let mut frame_hdr_buf = [0u8; WAL_FRAME_HEADER_SIZE]; + + // Phase 1: Find the last committed frame (commit_size > 0). + let total_frames = wal_len.saturating_sub(WAL_HEADER_SIZE) / frame_size; + let mut last_commit_idx: Option = None; + for i in 0..total_frames { + let hdr_offset = (WAL_HEADER_SIZE + i * frame_size) as u64; + wal_file.seek(SeekFrom::Start(hdr_offset))?; + wal_file.read_exact(&mut frame_hdr_buf)?; + let fh = WalFrameHeader::parse(&frame_hdr_buf); + if fh.is_valid(wal_header) && fh.commit_size > 0 { + last_commit_idx = Some(i); + } + } + + let Some(last_commit) = last_commit_idx else { + // No committed transaction found; nothing to apply. + return Ok(0); + }; + + // Phase 2: Apply frames 0..=last_commit sequentially. + let mut page_buf = vec![0u8; params.page_size]; + let mut patched: usize = 0; + + wal_file.seek(SeekFrom::Start(WAL_HEADER_SIZE as u64))?; + for _ in 0..=last_commit { + wal_file.read_exact(&mut frame_hdr_buf)?; + let frame_hdr = WalFrameHeader::parse(&frame_hdr_buf); + wal_file.read_exact(&mut page_buf)?; + + if !frame_hdr.is_valid(wal_header) { + continue; + } + + if page_buf.iter().all(|&b| b == 0) { + continue; + } + + // WAL page_number is 1-indexed; decrypt_page expects 0-indexed. + let page_num_0 = frame_hdr.page_number - 1; + + let decrypted = decrypt_page(&page_buf, enc_key, mac_key, page_num_0, params)?; + + // For page 0, decrypt_page returns 4080 bytes (skips salt area). + // Prepend the SQLite header to restore full page size. + let write_data = if page_num_0 == 0 { + let mut full = Vec::with_capacity(params.page_size); + full.extend_from_slice(b"SQLite format 3\0"); + full.extend_from_slice(&decrypted); + full + } else { + decrypted + }; + + let offset = (page_num_0 as u64) * (params.page_size as u64); + db_file.seek(SeekFrom::Start(offset))?; + db_file.write_all(&write_data)?; + patched += 1; + } + + db_file.flush()?; + Ok(patched) +} + +/// Derive the encrypted DB path from a WAL path by stripping the "-wal" suffix. +pub(crate) fn wal_path_to_db_path(wal_path: &Path) -> Result { + let wal_name = wal_path + .file_name() + .and_then(|n| n.to_str()) + .ok_or_else(|| DecryptError::InvalidWalHeader { + reason: "cannot determine WAL filename".to_string(), + })?; + + if !wal_name.ends_with("-wal") { + return Err(DecryptError::InvalidWalHeader { + reason: format!("WAL path does not end with -wal: {wal_name}"), + }); + } + + let db_name = &wal_name[..wal_name.len() - 4]; + Ok(wal_path.with_file_name(db_name)) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn make_wal_header(salt1: u32, salt2: u32) -> [u8; WAL_HEADER_SIZE] { + let mut buf = [0u8; WAL_HEADER_SIZE]; + buf[0..4].copy_from_slice(&WAL_MAGIC_BE.to_be_bytes()); + buf[4..8].copy_from_slice(&3007000u32.to_be_bytes()); + buf[8..12].copy_from_slice(&4096u32.to_be_bytes()); + buf[12..16].copy_from_slice(&1u32.to_be_bytes()); + buf[16..20].copy_from_slice(&salt1.to_be_bytes()); + buf[20..24].copy_from_slice(&salt2.to_be_bytes()); + buf + } + + fn make_frame_header( + pgno: u32, + commit_size: u32, + salt1: u32, + salt2: u32, + ) -> [u8; WAL_FRAME_HEADER_SIZE] { + let mut buf = [0u8; WAL_FRAME_HEADER_SIZE]; + buf[0..4].copy_from_slice(&pgno.to_be_bytes()); + buf[4..8].copy_from_slice(&commit_size.to_be_bytes()); + buf[8..12].copy_from_slice(&salt1.to_be_bytes()); + buf[12..16].copy_from_slice(&salt2.to_be_bytes()); + buf + } + + #[test] + fn test_wal_header_parse_valid() { + let buf = make_wal_header(0xAABBCCDD, 0x11223344); + let hdr = WalHeader::parse(&buf).unwrap(); + assert_eq!(hdr.salt1, 0xAABBCCDD); + assert_eq!(hdr.salt2, 0x11223344); + } + + #[test] + fn test_wal_header_parse_le_magic() { + let mut buf = make_wal_header(1, 2); + buf[0..4].copy_from_slice(&WAL_MAGIC_LE.to_be_bytes()); + let hdr = WalHeader::parse(&buf).unwrap(); + assert_eq!(hdr.salt1, 1); + } + + #[test] + fn test_wal_header_parse_bad_magic() { + let mut buf = [0u8; WAL_HEADER_SIZE]; + buf[0..4].copy_from_slice(&0xDEADBEEFu32.to_be_bytes()); + let err = WalHeader::parse(&buf).unwrap_err(); + assert!(err.to_string().contains("bad magic")); + } + + #[test] + fn test_frame_header_valid() { + let wal_hdr = WalHeader { + salt1: 100, + salt2: 200, + }; + let fh = WalFrameHeader::parse(&make_frame_header(5, 0, 100, 200)); + assert!(fh.is_valid(&wal_hdr)); + assert_eq!(fh.page_number, 5); + } + + #[test] + fn test_frame_header_stale_salt() { + let wal_hdr = WalHeader { + salt1: 100, + salt2: 200, + }; + let fh = WalFrameHeader::parse(&make_frame_header(5, 0, 99, 200)); + assert!(!fh.is_valid(&wal_hdr)); + } + + #[test] + fn test_frame_header_zero_pgno() { + let wal_hdr = WalHeader { + salt1: 100, + salt2: 200, + }; + let fh = WalFrameHeader::parse(&make_frame_header(0, 0, 100, 200)); + assert!(!fh.is_valid(&wal_hdr)); + } + + #[test] + fn test_frame_header_pgno_too_large() { + let wal_hdr = WalHeader { + salt1: 100, + salt2: 200, + }; + let fh = WalFrameHeader::parse(&make_frame_header(MAX_PAGE_NUMBER + 1, 0, 100, 200)); + assert!(!fh.is_valid(&wal_hdr)); + } + + use crate::kdf::{derive_enc_key, derive_mac_key}; + use crate::params::MACOS_4_1_7_31; + + /// Build a fake encrypted page (reverse of decrypt_page). + fn build_encrypted_page( + content: &[u8], + enc_key: &[u8; 32], + mac_key: &[u8; 32], + page_num: u32, + params: &CryptoParams, + salt: Option<&[u8; 16]>, + ) -> Vec { + use aes::cipher::{block_padding::NoPadding, BlockEncryptMut, KeyIvInit}; + use hmac::{Hmac, Mac}; + use sha2::Sha512; + + type HmacSha512 = Hmac; + type Aes256CbcEnc = cbc::Encryptor; + + let iv = [0x42u8; 16]; // deterministic IV for testing + let offset = if page_num == 0 { params.salt_size } else { 0 }; + + let data_len = params.page_size - params.reserve - offset; + let mut plaintext = vec![0u8; data_len]; + let copy_len = content.len().min(data_len); + plaintext[..copy_len].copy_from_slice(&content[..copy_len]); + + let mut ciphertext = plaintext.clone(); + Aes256CbcEnc::new(enc_key.into(), (&iv).into()) + .encrypt_padded_mut::(&mut ciphertext, data_len) + .expect("encryption should not fail"); + + let mut page = vec![0u8; params.page_size]; + if page_num == 0 { + page[..params.salt_size].copy_from_slice(salt.unwrap()); + page[offset..offset + data_len].copy_from_slice(&ciphertext); + } else { + page[..data_len].copy_from_slice(&ciphertext); + } + + // Place IV in reserve area + let iv_start = params.page_size - params.reserve; + page[iv_start..iv_start + params.iv_size].copy_from_slice(&iv); + + // Compute HMAC + let hmac_data_end = params.page_size - params.reserve + params.iv_size; + let mut mac = HmacSha512::new_from_slice(mac_key).unwrap(); + mac.update(&page[offset..hmac_data_end]); + mac.update(&(page_num + 1).to_le_bytes()); + let hmac_result = mac.finalize().into_bytes(); + page[hmac_data_end..hmac_data_end + params.hmac_size] + .copy_from_slice(&hmac_result[..params.hmac_size]); + + page + } + + /// Build a WAL file in memory: header + frames. + fn build_wal_file(salt1: u32, salt2: u32, frames: &[(u32, u32, u32, u32, &[u8])]) -> Vec { + let mut wal = make_wal_header(salt1, salt2).to_vec(); + for &(pgno, commit_size, fsalt1, fsalt2, page_data) in frames { + let fh = make_frame_header(pgno, commit_size, fsalt1, fsalt2); + wal.extend_from_slice(&fh); + wal.extend_from_slice(page_data); + } + wal + } + + #[test] + fn test_decrypt_wal_patches_valid_frames() { + let params = &MACOS_4_1_7_31; + let raw_key = [0xABu8; 32]; + let salt = [0x01u8; 16]; + let enc_key = derive_enc_key(&raw_key, &salt, params); + let mac_key = derive_mac_key(&enc_key, &salt, params); + + let dir = tempfile::tempdir().unwrap(); + + // Create fake encrypted DB (2 pages) + let enc_db_path = dir.path().join("test.db"); + let page0 = build_encrypted_page( + b"page0-original", + &enc_key, + &mac_key, + 0, + params, + Some(&salt), + ); + let page1 = build_encrypted_page(b"page1-original", &enc_key, &mac_key, 1, params, None); + std::fs::write(&enc_db_path, [page0, page1].concat()).unwrap(); + + // Create decrypted DB placeholder (2 pages of zeros) + let dec_db_path = dir.path().join("test_decrypted.db"); + std::fs::write(&dec_db_path, vec![0u8; params.page_size * 2]).unwrap(); + + // WAL with 1 valid frame (page 2, 1-indexed) + 1 stale frame + let new_page1 = + build_encrypted_page(b"page1-from-wal", &enc_key, &mac_key, 1, params, None); + let stale_page = build_encrypted_page(b"stale-data", &enc_key, &mac_key, 1, params, None); + let wal_salt1: u32 = 0xAAAA; + let wal_salt2: u32 = 0xBBBB; + let wal_data = build_wal_file( + wal_salt1, + wal_salt2, + &[ + (2, 1, wal_salt1, wal_salt2, &new_page1), + (2, 0, 0xDEAD, 0xBEEF, &stale_page), + ], + ); + let wal_path = dir.path().join("test.db-wal"); + std::fs::write(&wal_path, &wal_data).unwrap(); + + let patched = decrypt_wal(&wal_path, &dec_db_path, &raw_key, params).unwrap(); + assert_eq!(patched, 1, "should patch exactly 1 valid frame"); + + // Verify patched page contains the expected decrypted content. + let result = std::fs::read(&dec_db_path).unwrap(); + let patched_page = &result[params.page_size..params.page_size * 2]; + assert_eq!( + &patched_page[..14], + b"page1-from-wal", + "patched page content should match WAL frame plaintext" + ); + assert!( + patched_page[14..params.page_size - params.reserve] + .iter() + .all(|&b| b == 0), + "rest of data area should be zero-padded" + ); + } + + #[test] + fn test_decrypt_wal_patches_page0_with_sqlite_header() { + let params = &MACOS_4_1_7_31; + let raw_key = [0xABu8; 32]; + let salt = [0x01u8; 16]; + let enc_key = derive_enc_key(&raw_key, &salt, params); + let mac_key = derive_mac_key(&enc_key, &salt, params); + + let dir = tempfile::tempdir().unwrap(); + + // Create fake encrypted DB (1 page) + let enc_db_path = dir.path().join("test.db"); + let page0 = build_encrypted_page( + b"page0-original", + &enc_key, + &mac_key, + 0, + params, + Some(&salt), + ); + std::fs::write(&enc_db_path, &page0).unwrap(); + + // Create decrypted DB placeholder (1 page of zeros) + let dec_db_path = dir.path().join("test_decrypted.db"); + std::fs::write(&dec_db_path, vec![0u8; params.page_size]).unwrap(); + + // WAL with 1 valid frame for page 1 (1-indexed = page 0, 0-indexed) + let new_page0 = build_encrypted_page( + b"page0-from-wal", + &enc_key, + &mac_key, + 0, + params, + Some(&salt), + ); + let wal_salt1: u32 = 0x1111; + let wal_salt2: u32 = 0x2222; + let wal_data = build_wal_file( + wal_salt1, + wal_salt2, + &[(1, 1, wal_salt1, wal_salt2, &new_page0)], + ); + let wal_path = dir.path().join("test.db-wal"); + std::fs::write(&wal_path, &wal_data).unwrap(); + + let patched = decrypt_wal(&wal_path, &dec_db_path, &raw_key, params).unwrap(); + assert_eq!(patched, 1); + + let result = std::fs::read(&dec_db_path).unwrap(); + assert_eq!( + &result[..16], + b"SQLite format 3\0", + "page 0 should have SQLite header prepended" + ); + assert_eq!( + &result[16..16 + 14], + b"page0-from-wal", + "page 0 decrypted content should match WAL frame plaintext" + ); + } + + #[test] + fn test_decrypt_wal_header_only_returns_zero() { + let params = &MACOS_4_1_7_31; + let dir = tempfile::tempdir().unwrap(); + + let enc_db_path = dir.path().join("test.db"); + std::fs::write(&enc_db_path, vec![0xAAu8; params.page_size]).unwrap(); + let dec_db_path = dir.path().join("test_dec.db"); + std::fs::write(&dec_db_path, vec![0u8; params.page_size]).unwrap(); + + let wal_data = make_wal_header(1, 2).to_vec(); + let wal_path = dir.path().join("test.db-wal"); + std::fs::write(&wal_path, &wal_data).unwrap(); + + let result = decrypt_wal(&wal_path, &dec_db_path, &[0u8; 32], params).unwrap(); + assert_eq!(result, 0, "header-only WAL should return Ok(0)"); + } + + #[test] + fn test_decrypt_wal_too_small_returns_zero() { + let params = &MACOS_4_1_7_31; + let dir = tempfile::tempdir().unwrap(); + + let wal_path = dir.path().join("tiny.db-wal"); + std::fs::write(&wal_path, [0u8; 16]).unwrap(); + + let result = decrypt_wal(&wal_path, Path::new("/nonexistent"), &[0u8; 32], params); + assert_eq!( + result.unwrap(), + 0, + "WAL smaller than header should return Ok(0)" + ); + } + + #[test] + fn test_decrypt_wal_zero_bytes_returns_zero() { + let params = &MACOS_4_1_7_31; + let dir = tempfile::tempdir().unwrap(); + + let wal_path = dir.path().join("empty.db-wal"); + std::fs::write(&wal_path, []).unwrap(); + + let result = decrypt_wal(&wal_path, Path::new("/nonexistent"), &[0u8; 32], params); + assert_eq!(result.unwrap(), 0, "0-byte WAL should return Ok(0)"); + } + + #[test] + fn test_wal_path_to_db_path() { + let wal = Path::new("/data/session/session.db-wal"); + let db = wal_path_to_db_path(wal).unwrap(); + assert_eq!(db, Path::new("/data/session/session.db")); + } + + #[test] + fn test_wal_path_to_db_path_invalid() { + let bad = Path::new("/data/session/session.db"); + assert!(wal_path_to_db_path(bad).is_err()); + } + + // --- Commit boundary tests --- + + /// Helper: build a minimal WAL + decrypted DB pair for commit boundary tests. + /// + /// Returns `(wal_path, dec_db_path, tempdir)`. + fn setup_commit_boundary_test( + frames: &[(u32, u32, &[u8])], // (pgno, commit_size, page_content) + ) -> (std::path::PathBuf, std::path::PathBuf, tempfile::TempDir) { + let params = &MACOS_4_1_7_31; + let raw_key = [0xCDu8; 32]; + let salt = [0x02u8; 16]; + let enc_key = derive_enc_key(&raw_key, &salt, params); + let mac_key = derive_mac_key(&enc_key, &salt, params); + let wal_salt1: u32 = 0x1234; + let wal_salt2: u32 = 0x5678; + + let dir = tempfile::tempdir().unwrap(); + + // Build encrypted pages for all unique page numbers. + let max_pgno = frames.iter().map(|f| f.0).max().unwrap_or(1); + + // Create a fake encrypted DB large enough for all pages. + let enc_db_path = dir.path().join("cbt.db"); + let mut enc_db = vec![0xFFu8; params.page_size * max_pgno as usize]; + // Page 0 (1-indexed page 1) needs a salt prefix. + enc_db[..16].copy_from_slice(&salt); + std::fs::write(&enc_db_path, &enc_db).unwrap(); + + // Create blank decrypted DB. + let dec_db_path = dir.path().join("cbt_dec.db"); + std::fs::write( + &dec_db_path, + vec![0u8; params.page_size * max_pgno as usize], + ) + .unwrap(); + + // Build WAL frames. + let frame_data: Vec<(u32, u32, u32, u32, Vec)> = frames + .iter() + .map(|&(pgno, commit_size, content)| { + let page_num_0 = pgno - 1; + let page = build_encrypted_page( + content, + &enc_key, + &mac_key, + page_num_0, + params, + if page_num_0 == 0 { Some(&salt) } else { None }, + ); + (pgno, commit_size, wal_salt1, wal_salt2, page) + }) + .collect(); + + let frame_refs: Vec<(u32, u32, u32, u32, &[u8])> = frame_data + .iter() + .map(|(pg, cs, s1, s2, data)| (*pg, *cs, *s1, *s2, data.as_slice())) + .collect(); + + let wal_bytes = build_wal_file(wal_salt1, wal_salt2, &frame_refs); + let wal_path = dir.path().join("cbt.db-wal"); + std::fs::write(&wal_path, &wal_bytes).unwrap(); + + (wal_path, dec_db_path, dir) + } + + /// Content-unique page data for testing (avoids all-zero page skip). + fn page_content(tag: u8) -> Vec { + vec![tag; 32] + } + + #[test] + fn test_commit_boundary_partial_transaction_skipped() { + // Frame A: pgno=1, commit_size=0 (non-commit frame of a transaction) + // Frame B: pgno=2, commit_size=5 (commit frame — this is the boundary) + // Frame C: pgno=3, commit_size=0 (new uncommitted transaction — must be skipped) + let pa = page_content(0xA1); + let pb = page_content(0xA2); + let pc = page_content(0xA3); + let frames: &[(u32, u32, &[u8])] = &[(1, 0, &pa), (2, 5, &pb), (3, 0, &pc)]; + let (wal_path, dec_db_path, _dir) = setup_commit_boundary_test(frames); + + let params = &MACOS_4_1_7_31; + let raw_key = [0xCDu8; 32]; + let patched = decrypt_wal(&wal_path, &dec_db_path, &raw_key, params).unwrap(); + + // Frames A and B applied; frame C skipped. + assert_eq!(patched, 2, "frames A and B should be patched; C skipped"); + } + + #[test] + fn test_no_commit_frame_returns_zero() { + // All frames have commit_size=0 — no committed transaction. + let pa = page_content(0xB1); + let pb = page_content(0xB2); + let frames: &[(u32, u32, &[u8])] = &[(1, 0, &pa), (2, 0, &pb)]; + let (wal_path, dec_db_path, _dir) = setup_commit_boundary_test(frames); + + let params = &MACOS_4_1_7_31; + let raw_key = [0xCDu8; 32]; + let patched = decrypt_wal(&wal_path, &dec_db_path, &raw_key, params).unwrap(); + + assert_eq!( + patched, 0, + "no commit frame means nothing should be applied" + ); + } + + #[test] + fn test_all_committed_frames_applied() { + // Both frames are commit frames — all should be applied. + let pa = page_content(0xC1); + let pb = page_content(0xC2); + let frames: &[(u32, u32, &[u8])] = &[(1, 3, &pa), (2, 7, &pb)]; + let (wal_path, dec_db_path, _dir) = setup_commit_boundary_test(frames); + + let params = &MACOS_4_1_7_31; + let raw_key = [0xCDu8; 32]; + let patched = decrypt_wal(&wal_path, &dec_db_path, &raw_key, params).unwrap(); + + assert_eq!(patched, 2, "both committed frames should be applied"); + } + + #[test] + fn test_commit_boundary_multi_transaction_last_wins() { + // 5 frames, two complete transactions + one incomplete: + // Frame 0 (pgno=2, commit_size=0): non-commit, part of tx1 + // Frame 1 (pgno=3, commit_size=3): commit, ends tx1 + // Frame 2 (pgno=4, commit_size=0): non-commit, part of tx2 + // Frame 3 (pgno=5, commit_size=6): commit, ends tx2 + // Frame 4 (pgno=6, commit_size=0): incomplete tx3, must be SKIPPED + let p0 = page_content(0xD0); + let p1 = page_content(0xD1); + let p2 = page_content(0xD2); + let p3 = page_content(0xD3); + let p4 = page_content(0xD4); + let frames: &[(u32, u32, &[u8])] = &[ + (2, 0, &p0), + (3, 3, &p1), + (4, 0, &p2), + (5, 6, &p3), + (6, 0, &p4), + ]; + let (wal_path, dec_db_path, _dir) = setup_commit_boundary_test(frames); + + let params = &MACOS_4_1_7_31; + let raw_key = [0xCDu8; 32]; + let patched = decrypt_wal(&wal_path, &dec_db_path, &raw_key, params).unwrap(); + + // Frames 0-3 applied (4 frames); frame 4 skipped. + assert_eq!(patched, 4, "frames 0-3 should be patched; frame 4 skipped"); + + // Verify frame 4 (pgno=6, 0-indexed=5) was NOT written to the DB. + let db_contents = std::fs::read(&dec_db_path).unwrap(); + let page5_offset = 5 * params.page_size; + let page5 = &db_contents[page5_offset..page5_offset + params.page_size]; + assert!( + page5.iter().all(|&b| b == 0), + "page 5 (frame 4, pgno=6) should remain zero — not patched" + ); + } + + #[test] + fn test_decrypt_wal_direct_produces_same_result() { + let params = &MACOS_4_1_7_31; + let raw_key = [0xABu8; 32]; + let salt = [0x01u8; 16]; + let enc_key = derive_enc_key(&raw_key, &salt, params); + let mac_key = derive_mac_key(&enc_key, &salt, params); + + let dir = tempfile::tempdir().unwrap(); + + // Create fake encrypted DB (2 pages) + let enc_db_path = dir.path().join("test.db"); + let page0 = build_encrypted_page( + b"wal-direct-test", + &enc_key, + &mac_key, + 0, + params, + Some(&salt), + ); + let page1 = build_encrypted_page(b"page1-data", &enc_key, &mac_key, 1, params, None); + std::fs::write(&enc_db_path, [page0, page1].concat()).unwrap(); + + let wal_salt1: u32 = 0x5555; + let wal_salt2: u32 = 0x6666; + let new_page1 = build_encrypted_page(b"page1-updated", &enc_key, &mac_key, 1, params, None); + let wal_data = build_wal_file( + wal_salt1, + wal_salt2, + &[(2, 1, wal_salt1, wal_salt2, &new_page1)], + ); + let wal_path = dir.path().join("test.db-wal"); + std::fs::write(&wal_path, &wal_data).unwrap(); + + // Decrypt with raw key + let dec_raw = dir.path().join("dec_raw.db"); + std::fs::write(&dec_raw, vec![0u8; params.page_size * 2]).unwrap(); + let patched_raw = decrypt_wal(&wal_path, &dec_raw, &raw_key, params).unwrap(); + + // Decrypt with direct enc_key + let dec_direct = dir.path().join("dec_direct.db"); + std::fs::write(&dec_direct, vec![0u8; params.page_size * 2]).unwrap(); + let patched_direct = + decrypt_wal_direct(&wal_path, &dec_direct, &enc_key, &salt, params).unwrap(); + + assert_eq!(patched_raw, patched_direct); + assert_eq!( + std::fs::read(&dec_raw).unwrap(), + std::fs::read(&dec_direct).unwrap(), + ); + } +} diff --git a/crates/wx-keychain/Cargo.toml b/crates/wx-keychain/Cargo.toml new file mode 100644 index 0000000..c474788 --- /dev/null +++ b/crates/wx-keychain/Cargo.toml @@ -0,0 +1,28 @@ +[package] +name = "wx-keychain" +version.workspace = true +edition.workspace = true + +[dependencies] +wx-decrypt = { path = "../wx-decrypt" } +wx-paths = { path = "../wx-paths" } +tokio = { version = "1", features = ["process", "time", "io-util", "rt-multi-thread", "macros"] } +regex = "1" +serde = { version = "1", features = ["derive"] } +toml = "0.8" +chrono = { version = "0.4", features = ["serde"] } +thiserror = "2" +hex = "0.4" + +[target.'cfg(target_os = "macos")'.dependencies] +mach2 = "0.6" + +[target.'cfg(unix)'.dependencies] +libc = "0.2" + +[dev-dependencies] +aes = "0.8" +cbc = "0.1" +hmac = "0.12" +sha2 = "0.10" +tempfile = "3" diff --git a/crates/wx-keychain/src/account_id.rs b/crates/wx-keychain/src/account_id.rs new file mode 100644 index 0000000..04d2b27 --- /dev/null +++ b/crates/wx-keychain/src/account_id.rs @@ -0,0 +1,316 @@ +/// Normalized account ID, separating raw directory name from canonical base ID +/// and the optional 4-character directory hash suffix used by WeChat macOS. +/// +/// Two levels of canonicalization: +/// - **Conservative** (`base()`): only strips suffix for `wxid_*` prefix accounts +/// where the structure is unambiguous. Safe for arbitrary strings. +/// - **Confirmed** (`confirmed_base()`): strips suffix for non-`wxid_` patterns +/// only when an external signal confirms the base ID candidate. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct AccountId { + raw: String, + /// Conservative base: only wxid_* suffix stripped + conservative_base: String, + /// Alias candidate derived from a trailing `_XXXX` segment. + alias_candidate: Option, + suffix: Option, +} + +impl AccountId { + /// Parse an account directory name into its components. + /// + /// Detects the trailing `_XXXX` directory hash suffix (exactly 4 alphanumeric + /// characters). Conservative `base()` only strips for `wxid_*` accounts; + /// non-`wxid_` inputs expose the stripped form as an alias candidate and need + /// an external confirmation signal before canonicalization. + /// + /// Rules for `base()` (conservative): + /// - `wxid_example123abc_ab12` → `wxid_example123abc` + /// - `testuser001_1662` → `testuser001_1662` (not stripped without confirmation) + /// - `wxid_test` → `wxid_test` + /// - `not_a_wxid` → `not_a_wxid` + /// + /// Rules for `base_for_account_dir()` (confirmed directory without extra signal): + /// - `wxid_example123abc_ab12` → `wxid_example123abc` + /// - `testuser001_1662` → `testuser001_1662` + /// - `wxid_test` → `wxid_test` + /// - `not_a_wxid` → `not_a_wxid` + pub fn parse(dir_name: &str) -> Self { + if let Some(pos) = dir_name.rfind('_') { + let tail = &dir_name[pos + 1..]; + let prefix = &dir_name[..pos]; + + if tail.len() == 4 + && tail.chars().all(|c| c.is_ascii_alphanumeric()) + && !prefix.is_empty() + { + // For wxid_ prefix: only strip if base retains wxid_ + substance + let wxid_safe = if dir_name.starts_with("wxid_") { + prefix.starts_with("wxid_") && prefix.len() > 5 + } else { + true + }; + + let conservative_base = if wxid_safe && dir_name.starts_with("wxid_") { + prefix.to_string() + } else { + dir_name.to_string() + }; + + let alias_candidate = if wxid_safe { + Some(prefix.to_string()) + } else { + // Even alias matching should not collapse wxid_test → wxid. + None + }; + + return Self { + raw: dir_name.to_string(), + conservative_base, + alias_candidate, + suffix: Some(tail.to_string()), + }; + } + } + + Self { + raw: dir_name.to_string(), + conservative_base: dir_name.to_string(), + alias_candidate: None, + suffix: None, + } + } + + /// The raw account directory name as-is. + pub fn raw(&self) -> &str { + &self.raw + } + + /// Conservative canonical base: only strips suffix for `wxid_*` prefix accounts. + /// Safe for arbitrary strings — does not assume the input is a real account directory. + pub fn base(&self) -> &str { + &self.conservative_base + } + + /// Confirmed-dir base without extra signal. + /// + /// This remains conservative for non-`wxid_` inputs. Call `confirmed_base()` + /// with an externally confirmed base-ID hint to canonicalize legacy account IDs. + pub fn base_for_account_dir(&self) -> &str { + &self.conservative_base + } + + /// Alias candidate derived from a trailing `_XXXX` segment. + /// + /// For legacy non-`wxid_` IDs this is used for alias matching and can be + /// promoted to the canonical base only when an external signal confirms it. + pub fn alias_candidate(&self) -> Option<&str> { + self.alias_candidate.as_deref() + } + + /// Canonical base for a confirmed directory, using an independently confirmed + /// base-ID hint when available. + pub fn confirmed_base(&self, confirmed_base: Option<&str>) -> &str { + match (self.alias_candidate(), confirmed_base) { + (Some(candidate), Some(confirmed)) if candidate == confirmed => candidate, + _ => self.base_for_account_dir(), + } + } + + /// The 4-character directory hash suffix candidate, if detected. + /// + /// Present even when the canonical base was NOT stripped (e.g. `wxid_test` + /// has suffix `Some("test")` but base remains `wxid_test`). Useful for + /// ilink directory prioritization in media key derivation. + pub fn suffix(&self) -> Option<&str> { + self.suffix.as_deref() + } + + /// Whether the conservative base differs from the raw name. + pub fn has_stripped_suffix(&self) -> bool { + self.raw != self.conservative_base + } + + /// Check if a user-provided token matches this account. + /// + /// Matches against raw name, conservative base, and the alias candidate. + /// This allows `--account testuser001` to match directory `testuser001_1662`. + pub fn matches(&self, token: &str) -> bool { + self.raw == token + || self.conservative_base == token + || self.alias_candidate.as_deref() == Some(token) + } +} + +/// Convenience: compute conservative canonical base from a directory name. +/// +/// Only strips suffix for `wxid_*` prefix accounts. For non-`wxid_` inputs, +/// returns the input unchanged. Use `canonical_base_for_account_dir()` when +/// the input is a confirmed account directory. +pub fn canonical_base(dir_name: &str) -> String { + AccountId::parse(dir_name).base().to_string() +} + +/// Compute canonical base for a confirmed account directory without extra signal. +pub fn canonical_base_for_account_dir(dir_name: &str) -> String { + AccountId::parse(dir_name) + .base_for_account_dir() + .to_string() +} + +/// Compute canonical base for a confirmed account directory with an independently +/// confirmed base-ID hint (for example `all_users/login/`). +pub fn canonical_base_for_account_dir_with_confirmed_base( + dir_name: &str, + confirmed_base: Option<&str>, +) -> String { + AccountId::parse(dir_name) + .confirmed_base(confirmed_base) + .to_string() +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn wxid_with_suffix() { + let id = AccountId::parse("wxid_example123abc_ab12"); + assert_eq!(id.base(), "wxid_example123abc"); + assert_eq!(id.base_for_account_dir(), "wxid_example123abc"); + assert_eq!(id.alias_candidate(), Some("wxid_example123abc")); + assert_eq!(id.suffix(), Some("ab12")); + assert!(id.has_stripped_suffix()); + assert!(id.matches("wxid_example123abc_ab12")); + assert!(id.matches("wxid_example123abc")); + assert!(!id.matches("wxid_example")); + } + + #[test] + fn legacy_non_wxid_conservative_not_stripped() { + let id = AccountId::parse("testuser001_1662"); + // Conservative: NOT stripped (no independent signal) + assert_eq!(id.base(), "testuser001_1662"); + // Confirmed dir without extra signal: still conservative + assert_eq!(id.base_for_account_dir(), "testuser001_1662"); + assert_eq!(id.alias_candidate(), Some("testuser001")); + assert_eq!(id.suffix(), Some("1662")); + assert!(!id.has_stripped_suffix()); + assert!(id.matches("testuser001_1662")); + assert!(id.matches("testuser001")); + } + + #[test] + fn not_a_wxid_conservative_not_stripped() { + let id = AccountId::parse("not_a_wxid"); + // Conservative: NOT stripped + assert_eq!(id.base(), "not_a_wxid"); + // Confirmed dir without extra signal: still conservative + assert_eq!(id.base_for_account_dir(), "not_a_wxid"); + assert_eq!(id.alias_candidate(), Some("not_a")); + assert_eq!(id.suffix(), Some("wxid")); + } + + #[test] + fn wxid_test_not_stripped() { + let id = AccountId::parse("wxid_test"); + assert_eq!(id.base(), "wxid_test"); + // Even for confirmed dirs, wxid_test → wxid is not safe + assert_eq!(id.base_for_account_dir(), "wxid_test"); + assert_eq!(id.alias_candidate(), None); + assert_eq!(id.suffix(), Some("test")); + assert!(!id.has_stripped_suffix()); + } + + #[test] + fn wxid_bare_short() { + let id = AccountId::parse("wxid_x"); + assert_eq!(id.base(), "wxid_x"); + assert_eq!(id.suffix(), None); + } + + #[test] + fn long_suffix_not_stripped() { + let id = AccountId::parse("wxid_test_abcde"); + assert_eq!(id.base(), "wxid_test_abcde"); + assert_eq!(id.suffix(), None); + } + + #[test] + fn non_alnum_suffix_not_stripped() { + let id = AccountId::parse("wxid_test_ab-c"); + assert_eq!(id.base(), "wxid_test_ab-c"); + assert_eq!(id.suffix(), None); + } + + #[test] + fn multiple_underscores_wxid() { + let id = AccountId::parse("wxid_foobar456def_c3e7"); + assert_eq!(id.base(), "wxid_foobar456def"); + assert_eq!(id.suffix(), Some("c3e7")); + assert!(id.has_stripped_suffix()); + } + + #[test] + fn no_underscore() { + let id = AccountId::parse("nounderscore"); + assert_eq!(id.base(), "nounderscore"); + assert_eq!(id.suffix(), None); + } + + #[test] + fn bare_wxid_prefix() { + let id = AccountId::parse("wxid_"); + assert_eq!(id.base(), "wxid_"); + assert_eq!(id.suffix(), None); + } + + #[test] + fn canonical_base_conservative() { + assert_eq!( + canonical_base("wxid_example123abc_ab12"), + "wxid_example123abc" + ); + // Non-wxid: conservative does NOT strip + assert_eq!(canonical_base("testuser001_1662"), "testuser001_1662"); + assert_eq!(canonical_base("wxid_test"), "wxid_test"); + assert_eq!(canonical_base("not_a_wxid"), "not_a_wxid"); + } + + #[test] + fn canonical_base_for_account_dir_stays_conservative_for_non_wxid() { + assert_eq!( + canonical_base_for_account_dir("wxid_example123abc_ab12"), + "wxid_example123abc" + ); + assert_eq!( + canonical_base_for_account_dir("testuser001_1662"), + "testuser001_1662" + ); + // wxid_test is protected: stripping would leave bare "wxid" + assert_eq!(canonical_base_for_account_dir("wxid_test"), "wxid_test"); + assert_eq!(canonical_base_for_account_dir("not_a_wxid"), "not_a_wxid"); + } + + #[test] + fn canonical_base_for_account_dir_uses_confirmed_base_hint() { + assert_eq!( + canonical_base_for_account_dir_with_confirmed_base( + "testuser001_1662", + Some("testuser001") + ), + "testuser001" + ); + assert_eq!( + canonical_base_for_account_dir_with_confirmed_base("not_a_wxid", Some("not_a")), + "not_a" + ); + assert_eq!( + canonical_base_for_account_dir_with_confirmed_base( + "testuser001_1662", + Some("someone_else") + ), + "testuser001_1662" + ); + } +} diff --git a/crates/wx-keychain/src/error.rs b/crates/wx-keychain/src/error.rs new file mode 100644 index 0000000..97b2ad9 --- /dev/null +++ b/crates/wx-keychain/src/error.rs @@ -0,0 +1,55 @@ +use thiserror::Error; + +#[derive(Debug, Error)] +pub enum KeychainError { + #[error("SIP (System Integrity Protection) is not disabled — run `csrutil disable` in Recovery Mode")] + SipEnabled, + + #[error("DevToolsSecurity is not enabled — run `sudo DevToolsSecurity -enable`")] + DevToolsSecurityDisabled, + + #[error("{0} is not in _developer group — run `sudo dscl . -append /Groups/_developer GroupMembership {0}`")] + NotInDeveloperGroup(String), + + #[error("LLDB not found — install Xcode Command Line Tools: `xcode-select --install`")] + LldbNotFound, + + #[error("python3 not found — install Xcode Command Line Tools: `xcode-select --install`")] + Python3NotFound, + + #[error("WeChat is not running")] + WeChatNotRunning, + + #[error("WeChat version {version} is not supported for key extraction (requires 4.1.7.x or 4.1.8.x)")] + UnsupportedVersion { version: String }, + + #[error("could not detect WeChat account directory")] + AccountNotDetected, + + #[error("cannot detect active account: {reason}\nCandidates:\n{candidates}")] + AccountDetectionFailed { reason: String, candidates: String }, + + #[error("LLDB capture timed out after {seconds}s — did you log in to WeChat?")] + CaptureTimeout { seconds: u64 }, + + #[error("no PBKDF2 calls with rounds=256000 found in LLDB output")] + NoPbkdfCalls, + + #[error("captured key does not match target account salt")] + KeySaltMismatch, + + #[error("task_for_pid failed for PID {pid} (kern_return={kr}) — ensure SIP is disabled (csrutil disable in Recovery Mode) and run with sudo")] + TaskForPidFailed { pid: u32, kr: i32 }, + + #[error("no valid enc_key found in WeChat process memory")] + NoKeysFound, + + #[error("key store error: {0}")] + Store(String), + + #[error("I/O error: {0}")] + Io(#[from] std::io::Error), + + #[error("{0}")] + Other(String), +} diff --git a/crates/wx-keychain/src/lib.rs b/crates/wx-keychain/src/lib.rs new file mode 100644 index 0000000..8090647 --- /dev/null +++ b/crates/wx-keychain/src/lib.rs @@ -0,0 +1,197 @@ +pub mod account_id; +pub mod error; +pub mod lldb; +pub mod mach_vm; +pub mod nickname; +pub mod process; +pub mod script; +pub mod store; + +pub use account_id::AccountId; +pub use error::KeychainError; +pub use lldb::{capture_key, CaptureResult}; +#[cfg(target_os = "macos")] +pub use mach_vm::{capture_key_mach, MachCaptureResult}; +pub use nickname::resolve_nickname; +pub use process::config_dir; +pub use process::detect_active_account; +pub use process::{ + ensure_supported_wechat_version, extract_base_wxid, find_account_dirs, find_account_dirs_under, + find_wechat_pid, is_xwechat_files_root, AccountDirInfo, ActiveAccount, + DetectionSource, SUPPORTED_VERSION, +}; +pub use wx_decrypt::read_db_salt; +pub use store::{AccountKey, EncKeyEntry, KeyStore}; + +use std::process::Command; + +/// 单项前置条件检查的结果。 +pub struct PreflightCheck { + pub name: &'static str, + pub passed: bool, + pub detail: String, + pub fix_cmd: Option, +} + +pub fn check_sip() -> PreflightCheck { + let (passed, detail) = match Command::new("csrutil").arg("status").output() { + Ok(output) => { + let s = String::from_utf8_lossy(&output.stdout).to_lowercase(); + if s.contains("disabled") { + (true, "SIP is disabled".into()) + } else { + (false, "SIP is enabled".into()) + } + } + Err(e) => (false, format!("cannot run csrutil: {e}")), + }; + PreflightCheck { + name: "SIP disabled", + passed, + detail, + fix_cmd: if passed { + None + } else { + Some("csrutil disable # run in Recovery Mode".into()) + }, + } +} + +pub fn check_dev_tools_security() -> PreflightCheck { + let (passed, detail) = match Command::new("DevToolsSecurity").arg("-status").output() { + Ok(output) => { + let s = String::from_utf8_lossy(&output.stdout).to_lowercase(); + if s.contains("enabled") { + (true, "DevToolsSecurity is enabled".into()) + } else { + (false, "DevToolsSecurity is not enabled".into()) + } + } + Err(e) => (false, format!("cannot run DevToolsSecurity: {e}")), + }; + PreflightCheck { + name: "DevToolsSecurity", + passed, + detail, + fix_cmd: if passed { + None + } else { + Some("sudo DevToolsSecurity -enable".into()) + }, + } +} + +pub fn check_developer_group() -> PreflightCheck { + let sudo_user = std::env::var("SUDO_USER").ok(); + + let (passed, detail) = match &sudo_user { + Some(user) => { + // Running under sudo — check the invoking user's groups, not root's. + match Command::new("id").args(["-Gn", user]).output() { + Ok(output) => { + let s = String::from_utf8_lossy(&output.stdout); + if s.split_whitespace().any(|g| g == "_developer") { + (true, format!("{user} is in _developer group")) + } else { + (false, format!("{user} is NOT in _developer group")) + } + } + Err(e) => (false, format!("cannot check groups for {user}: {e}")), + } + } + None => match Command::new("groups").output() { + Ok(output) => { + let s = String::from_utf8_lossy(&output.stdout); + if s.split_whitespace().any(|g| g == "_developer") { + (true, "user is in _developer group".into()) + } else { + (false, "user is NOT in _developer group".into()) + } + } + Err(e) => (false, format!("cannot run groups: {e}")), + }, + }; + + let fix_user = sudo_user.as_deref().unwrap_or("$USER"); + PreflightCheck { + name: "_developer group", + passed, + detail, + fix_cmd: if passed { + None + } else { + Some(format!( + "sudo dscl . -append /Groups/_developer GroupMembership {fix_user}" + )) + }, + } +} + +pub fn check_binary(name: &'static str, version_flag: &str) -> PreflightCheck { + let (passed, detail) = match Command::new(name).arg(version_flag).output() { + Ok(output) if output.status.success() => { + let ver = String::from_utf8_lossy(&output.stdout) + .lines() + .next() + .unwrap_or("") + .trim() + .to_string(); + ( + true, + if ver.is_empty() { + format!("{name} OK") + } else { + ver + }, + ) + } + Ok(_) => (false, format!("{name} found but returned error")), + Err(_) => (false, format!("{name} not found")), + }; + PreflightCheck { + name, + passed, + detail, + fix_cmd: if passed { + None + } else { + Some("xcode-select --install".into()) + }, + } +} + +/// 运行所有前置条件检查,返回结果列表。 +pub fn all_preflight_checks() -> Vec { + vec![ + check_sip(), + check_dev_tools_security(), + check_developer_group(), + check_binary("lldb", "--version"), + check_binary("python3", "-V"), + ] +} + +/// Pre-flight checks before key extraction. +/// +/// Verifies: SIP disabled, DevToolsSecurity enabled, _developer group membership, +/// LLDB and python3 available. +pub fn preflight_checks() -> Result<(), KeychainError> { + for check in all_preflight_checks() { + if !check.passed { + return Err(match check.name { + "SIP disabled" => KeychainError::SipEnabled, + "DevToolsSecurity" => KeychainError::DevToolsSecurityDisabled, + "_developer group" => { + let user = std::env::var("SUDO_USER") + .or_else(|_| std::env::var("USER")) + .unwrap_or_else(|_| "$USER".into()); + KeychainError::NotInDeveloperGroup(user) + } + "lldb" => KeychainError::LldbNotFound, + "python3" => KeychainError::Python3NotFound, + _ => KeychainError::Other(check.detail), + }); + } + } + Ok(()) +} diff --git a/crates/wx-keychain/src/lldb.rs b/crates/wx-keychain/src/lldb.rs new file mode 100644 index 0000000..39c6e12 --- /dev/null +++ b/crates/wx-keychain/src/lldb.rs @@ -0,0 +1,209 @@ +use std::process::Command; +use std::time::Duration; + +use regex::Regex; +use tokio::io::AsyncBufReadExt; +use tokio::process::Command as AsyncCommand; +use tokio::time::timeout; + +use crate::error::KeychainError; +use crate::process::AccountDirInfo; +use crate::script::CAPTURE_KEY_SCRIPT; +use wx_decrypt::params::MACOS_4_1_7_31; +use wx_decrypt::validate_key; + +/// Result of a successful key capture. +#[derive(Debug)] +pub struct CaptureResult { + pub raw_key: [u8; 32], + pub call_count: u32, + /// Which account directory the captured key belongs to. + pub matched_account: AccountDirInfo, +} + +/// Run the full LLDB key capture flow against all known account directories. +/// +/// 1. Read salts from ALL account `message_0.db` files. +/// 2. Kill WeChat. +/// 3. Launch LLDB with `-w -n WeChat` (waits for WeChat to start). +/// 4. Open WeChat; user logs in. +/// 5. Stream LLDB output, parsing PBKDF2 calls. +/// 6. For each call with rounds=256000, check its salt against ALL known salts. +/// 7. On match, validate the full key via HMAC. Return key + matched account. +/// +/// This approach never pre-picks a target account, so it works regardless of +/// which account WeChat decides to auto-login as. +pub async fn capture_key( + accounts: &[AccountDirInfo], + capture_timeout: Duration, +) -> Result { + if accounts.is_empty() { + return Err(KeychainError::Other( + "no account directories provided".into(), + )); + } + + // Pre-read salts from all accounts. Skip unreadable DBs. + let account_salts: Vec<([u8; 16], &AccountDirInfo)> = accounts + .iter() + .filter_map(|a| wx_decrypt::read_db_salt(&a.message_db_path).ok().map(|salt| (salt, a))) + .collect(); + + if account_salts.is_empty() { + return Err(KeychainError::Other( + "could not read salt from any account database".into(), + )); + } + + // Kill WeChat. + let _ = Command::new("killall").arg("WeChat").output(); + tokio::time::sleep(Duration::from_secs(1)).await; + + // Write capture script to temp file. + let script_path = wx_paths::AppPaths::lldb_script_file(); + if let Some(parent) = script_path.parent() { + wx_paths::AppPaths::ensure_dir(parent)?; + } + std::fs::write(&script_path, CAPTURE_KEY_SCRIPT)?; + + // Prepare LLDB output file. + let output_path = wx_paths::AppPaths::lldb_output_file(); + + // Launch LLDB in wait mode. + let mut lldb = AsyncCommand::new("lldb") + .args([ + "-w", + "-n", + "WeChat", + "-o", + &format!("command script import {}", script_path.display()), + "-o", + "capture_keys", + ]) + .stdout(std::process::Stdio::piped()) + .stderr(std::process::Stdio::piped()) + .spawn() + .map_err(|e| KeychainError::Other(format!("failed to start lldb: {e}")))?; + + // Brief pause then open WeChat. + tokio::time::sleep(Duration::from_secs(1)).await; + let _ = Command::new("open").arg("-a").arg("WeChat").output(); + + eprintln!("Waiting for WeChat to start and trigger PBKDF2 calls..."); + eprintln!("Please log in to WeChat when prompted."); + + // Read LLDB stdout line by line, looking for PBKDF2 calls. + let stdout = lldb + .stdout + .take() + .ok_or_else(|| KeychainError::Other("no lldb stdout".into()))?; + let mut reader = tokio::io::BufReader::new(stdout).lines(); + + let re_header = Regex::new(r"^\[PBKDF2 #(\d+)\].*rounds=(\d+)").unwrap(); + let re_password = Regex::new(r"^\s*Password:\s*([0-9a-f]+)").unwrap(); + let re_salt = Regex::new(r"^\s*Salt:\s*([0-9a-f]+)").unwrap(); + + let mut current_call: Option<(u32, u32)> = None; // (call_count, rounds) + let mut current_password: Option = None; + let mut call_count = 0u32; + let mut output_lines = Vec::new(); + + let result = timeout(capture_timeout, async { + loop { + let line = match reader.next_line().await { + Ok(Some(line)) => line, + Ok(None) => break Err(KeychainError::NoPbkdfCalls), + Err(e) => break Err(KeychainError::Other(format!("read error: {e}"))), + }; + + output_lines.push(line.clone()); + + if let Some(caps) = re_header.captures(&line) { + let count: u32 = caps[1].parse().unwrap_or(0); + let rounds: u32 = caps[2].parse().unwrap_or(0); + current_call = Some((count, rounds)); + current_password = None; + call_count = count; + continue; + } + + if let Some(caps) = re_password.captures(&line) { + current_password = Some(caps[1].to_string()); + continue; + } + + if let Some(caps) = re_salt.captures(&line) { + let salt_hex = caps[1].to_string(); + + if let Some((_, rounds)) = current_call { + if rounds == 256000 { + if let Some(ref pwd_hex) = current_password { + if let Ok(salt_bytes) = hex::decode(&salt_hex) { + if salt_bytes.len() == 16 { + let mut pbkdf_salt = [0u8; 16]; + pbkdf_salt.copy_from_slice(&salt_bytes); + + // Match against ALL known account salts. + let mut matched: Option = None; + 'salt_match: for (known_salt, account) in &account_salts { + if pbkdf_salt != *known_salt { + continue; + } + // Salt matched — validate the key. + if let Ok(key_bytes) = hex::decode(pwd_hex) { + if key_bytes.len() == 32 { + let mut raw_key = [0u8; 32]; + raw_key.copy_from_slice(&key_bytes); + + use std::io::Read; + let mut first_page = vec![0u8; 4096]; + if let Ok(mut f) = + std::fs::File::open(&account.message_db_path) + { + if f.read_exact(&mut first_page).is_ok() + && validate_key( + &first_page, + &raw_key, + &MACOS_4_1_7_31, + ) + .is_some() + { + matched = Some(CaptureResult { + raw_key, + call_count, + matched_account: (*account).clone(), + }); + break 'salt_match; + } + } + } + } + } + if let Some(result) = matched { + break Ok(result); + } + } + } + } + } + } + current_call = None; + current_password = None; + } + } + }) + .await; + + // Save output for debugging. + let _ = std::fs::write(&output_path, output_lines.join("\n")); + + // Kill LLDB. + let _ = lldb.kill().await; + + match result { + Ok(inner) => inner, + Err(_) => Err(KeychainError::CaptureTimeout { + seconds: capture_timeout.as_secs(), + }), + } +} diff --git a/crates/wx-keychain/src/mach_vm/mach_reader.rs b/crates/wx-keychain/src/mach_vm/mach_reader.rs new file mode 100644 index 0000000..7765cdd --- /dev/null +++ b/crates/wx-keychain/src/mach_vm/mach_reader.rs @@ -0,0 +1,104 @@ +//! Mach VM reader: attaches to a process and reads its memory regions. +//! +//! This module is macOS-only (`#[cfg(target_os = "macos")]`). + +use mach2::kern_return::KERN_SUCCESS; +use mach2::traps::{mach_task_self, task_for_pid}; +use mach2::vm::{mach_vm_deallocate, mach_vm_read, mach_vm_region}; +use mach2::vm_prot::{VM_PROT_READ, VM_PROT_WRITE}; +use mach2::vm_region::{VM_REGION_BASIC_INFO_64, VM_REGION_BASIC_INFO_COUNT_64}; + +use crate::error::KeychainError; +use crate::mach_vm::reader::{MemRegion, MemoryReader}; + +/// Reads memory from a running process using Mach VM APIs. +pub struct MachVmReader { + task: u32, // mach_port_t +} + +impl MachVmReader { + /// Attach to a process by PID. Requires appropriate privileges + /// (root, or the target process must be ad-hoc signed). + pub fn attach(pid: u32) -> Result { + let mut task: u32 = 0; + let kr = unsafe { task_for_pid(mach_task_self(), pid as i32, &mut task) }; + if kr != KERN_SUCCESS { + return Err(KeychainError::TaskForPidFailed { pid, kr }); + } + Ok(Self { task }) + } +} + +impl MemoryReader for MachVmReader { + fn rw_regions(&self) -> Result, KeychainError> { + let mut regions = Vec::new(); + let mut address: u64 = 0; + + loop { + let mut size: u64 = 0; + let mut info = [0i32; VM_REGION_BASIC_INFO_COUNT_64 as usize]; + let mut info_cnt = VM_REGION_BASIC_INFO_COUNT_64; + let mut object_name: u32 = 0; + + let kr = unsafe { + mach_vm_region( + self.task, + &mut address, + &mut size, + VM_REGION_BASIC_INFO_64, + info.as_mut_ptr(), + &mut info_cnt, + &mut object_name, + ) + }; + + if kr != KERN_SUCCESS { + break; // No more regions. + } + + let protection = info[0]; + if (protection & VM_PROT_READ != 0) && (protection & VM_PROT_WRITE != 0) { + regions.push(MemRegion { + start: address, + end: address + size, + }); + } + + address += size; + } + + Ok(regions) + } + + fn read_bytes(&self, addr: u64, len: usize) -> Result, KeychainError> { + let mut data_ptr: usize = 0; // vm_offset_t + let mut data_cnt: u32 = 0; + + let kr = unsafe { + mach_vm_read( + self.task, + addr, + len as u64, + &mut data_ptr as *mut usize, + &mut data_cnt, + ) + }; + + if kr != KERN_SUCCESS { + return Err(KeychainError::Other(format!( + "mach_vm_read failed at 0x{addr:x} len={len}: kr={kr}" + ))); + } + + let result = unsafe { + std::slice::from_raw_parts(data_ptr as *const u8, data_cnt as usize).to_vec() + }; + + // Deallocate the kernel-allocated buffer. + unsafe { + mach_vm_deallocate(mach_task_self(), data_ptr as u64, data_cnt as u64); + } + + Ok(result) + } +} diff --git a/crates/wx-keychain/src/mach_vm/mod.rs b/crates/wx-keychain/src/mach_vm/mod.rs new file mode 100644 index 0000000..18dd729 --- /dev/null +++ b/crates/wx-keychain/src/mach_vm/mod.rs @@ -0,0 +1,139 @@ +pub mod pattern; +pub mod reader; +pub mod scanner; + +#[cfg(target_os = "macos")] +pub mod mach_reader; + +pub use pattern::{scan_chunk, FoundKey}; +pub use reader::{MemRegion, MemoryReader}; +pub use scanner::{MemoryScanner, ScanResult}; + +#[cfg(target_os = "macos")] +pub use mach_reader::MachVmReader; + +use crate::error::KeychainError; +use crate::process::AccountDirInfo; +use std::collections::HashMap; +use std::path::PathBuf; +use wx_decrypt::{EncKeyPair, KeyMaterial}; + +/// Result of a successful Mach VM key capture for one account. +#[derive(Debug, Clone)] +pub struct MachCaptureResult { + pub key_material: KeyMaterial, + pub matched_account: AccountDirInfo, +} + +/// Recursively find all `.db` files under a directory. +fn find_db_files(dir: &std::path::Path) -> Vec { + let mut result = Vec::new(); + let entries = match std::fs::read_dir(dir) { + Ok(e) => e, + Err(_) => return result, + }; + for entry in entries.flatten() { + let path = entry.path(); + if path.is_dir() { + result.extend(find_db_files(&path)); + } else if path.extension().is_some_and(|e| e == "db") { + result.push(path); + } + } + result +} + +/// Scan WeChat process memory for pre-derived encryption keys. +/// +/// Attaches to the process via `task_for_pid`, enumerates RW regions, scans for +/// `x''` patterns, matches each candidate's salt against the +/// provided account DBs, and HMAC-validates before returning. +#[cfg(target_os = "macos")] +pub fn capture_key_mach( + pid: u32, + accounts: &[AccountDirInfo], + params: &wx_decrypt::CryptoParams, +) -> Result, KeychainError> { + let reader = MachVmReader::attach(pid)?; + + // Collect (salt, db_path) pairs for ALL candidate DBs across all accounts. + let mut db_salts: Vec<([u8; 16], PathBuf)> = Vec::new(); + for account in accounts { + let db_storage = account.data_dir.join("db_storage"); + if !db_storage.exists() { + continue; + } + + for db_path in find_db_files(&db_storage) { + if let Ok(salt) = wx_decrypt::read_db_salt(&db_path) { + db_salts.push((salt, db_path)); + } + } + } + + if db_salts.is_empty() { + return Err(KeychainError::NoKeysFound); + } + + let scanner = MemoryScanner::new(reader); + let db_salt_refs: Vec<([u8; 16], &std::path::Path)> = db_salts + .iter() + .map(|(salt, path)| (*salt, path.as_path())) + .collect(); + let scan_results = scanner.scan(&db_salt_refs, params)?; + + if scan_results.is_empty() { + return Err(KeychainError::NoKeysFound); + } + + // Aggregate scan results by account: collect all (enc_key, salt) pairs per account. + let mut account_pairs: HashMap)> = HashMap::new(); + for sr in scan_results { + if let Some(account) = accounts.iter().find(|a| { + let db_storage = a.data_dir.join("db_storage"); + sr.db_path.starts_with(&db_storage) + }) { + let entry = account_pairs + .entry(account.account_id.clone()) + .or_insert_with(|| (account.clone(), Vec::new())); + entry.1.push(EncKeyPair { + key: sr.enc_key, + salt: sr.salt, + }); + } + } + + if account_pairs.is_empty() { + return Err(KeychainError::NoKeysFound); + } + + let mut results: Vec = Vec::new(); + for (_account_id, (account, mut pairs)) in account_pairs { + // Deduplicate by (salt, key) and sort for stable output. + pairs.sort_by(|a, b| a.salt.cmp(&b.salt).then_with(|| a.key.cmp(&b.key))); + pairs.dedup(); + + if pairs.is_empty() { + continue; + } + + // Always use EncKeys as the canonical format, even for a single pair. + results.push(MachCaptureResult { + key_material: KeyMaterial::EncKeys(pairs), + matched_account: account, + }); + } + + if results.is_empty() { + return Err(KeychainError::NoKeysFound); + } + + // Sort by account_id for stable cross-account output order. + results.sort_by(|a, b| { + a.matched_account + .account_id + .cmp(&b.matched_account.account_id) + }); + + Ok(results) +} diff --git a/crates/wx-keychain/src/mach_vm/pattern.rs b/crates/wx-keychain/src/mach_vm/pattern.rs new file mode 100644 index 0000000..d89bc76 --- /dev/null +++ b/crates/wx-keychain/src/mach_vm/pattern.rs @@ -0,0 +1,199 @@ +use std::collections::HashSet; + +pub const MAX_HEX_LEN: usize = 192; +pub const MAX_PATTERN_BYTES: usize = MAX_HEX_LEN + 3; + +/// A candidate key found by scanning a memory chunk for an SQL hex literal. +/// +/// Supported payload forms (mirrors `refs/wx-decrypt/find_all_keys.py`): +/// - `x'<64 hex>'` → enc_key only +/// - `x'<96 hex>'` → enc_key + salt +/// - `x'<98..192 hex, even>'` → enc_key = first 64 hex, salt = last 32 hex +#[derive(Debug, Clone, Hash, Eq, PartialEq)] +pub struct FoundKey { + pub enc_key: [u8; 32], + pub salt: Option<[u8; 16]>, +} + +/// Scan `buf` for supported `x'<...>'` hex literal patterns. +pub fn scan_chunk(buf: &[u8]) -> Vec { + let mut seen = HashSet::new(); + let mut results = Vec::new(); + + if buf.len() < 2 { + return results; + } + + let mut i = 0; + while i + 1 < buf.len() { + if buf[i] != b'x' || buf[i + 1] != b'\'' { + i += 1; + continue; + } + + let payload_start = i + 2; + let mut payload_end = payload_start; + while payload_end < buf.len() + && payload_end - payload_start < MAX_HEX_LEN + && buf[payload_end].is_ascii_hexdigit() + { + payload_end += 1; + } + + if payload_end >= buf.len() { + break; + } + + if buf[payload_end] != b'\'' { + i += 1; + continue; + } + + let hex_len = payload_end - payload_start; + if !is_supported_hex_len(hex_len) { + i += 1; + continue; + } + + let hex_slice = &buf[payload_start..payload_end]; + if let Some(found) = decode_found_key(hex_slice) { + if seen.insert(found.clone()) { + results.push(found); + } + } + + i = payload_end + 1; + } + + results +} + +fn is_supported_hex_len(len: usize) -> bool { + len == 64 || len == 96 || (len > 96 && len <= MAX_HEX_LEN && len.is_multiple_of(2)) +} + +fn decode_found_key(hex_slice: &[u8]) -> Option { + let enc_vec = hex::decode(&hex_slice[..64]).ok()?; + let mut enc_key = [0u8; 32]; + enc_key.copy_from_slice(&enc_vec); + + let salt = match hex_slice.len() { + 64 => None, + 96 => decode_salt(&hex_slice[64..96]), + len if len > 96 && len % 2 == 0 => decode_salt(&hex_slice[len - 32..len]), + _ => None, + }; + + Some(FoundKey { enc_key, salt }) +} + +fn decode_salt(hex_slice: &[u8]) -> Option<[u8; 16]> { + let salt_vec = hex::decode(hex_slice).ok()?; + let mut salt = [0u8; 16]; + salt.copy_from_slice(&salt_vec); + Some(salt) +} + +#[cfg(test)] +mod tests { + use super::*; + + const ENC_HEX: &str = "0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef"; + const SALT_HEX: &str = "fedcba9876543210fedcba9876543210"; + + fn expected_enc_key() -> [u8; 32] { + let v = hex::decode(ENC_HEX).unwrap(); + let mut arr = [0u8; 32]; + arr.copy_from_slice(&v); + arr + } + + fn expected_salt() -> [u8; 16] { + let v = hex::decode(SALT_HEX).unwrap(); + let mut arr = [0u8; 16]; + arr.copy_from_slice(&v); + arr + } + + #[test] + fn exact_96_hex_returns_key_and_salt() { + let buf = format!("x'{}{}'", ENC_HEX, SALT_HEX).into_bytes(); + let keys = scan_chunk(&buf); + assert_eq!(keys.len(), 1); + assert_eq!(keys[0].enc_key, expected_enc_key()); + assert_eq!(keys[0].salt, Some(expected_salt())); + } + + #[test] + fn exact_64_hex_returns_key_only() { + let buf = format!("x'{}'", ENC_HEX).into_bytes(); + let keys = scan_chunk(&buf); + assert_eq!(keys.len(), 1); + assert_eq!(keys[0].enc_key, expected_enc_key()); + assert_eq!(keys[0].salt, None); + } + + #[test] + fn long_hex_uses_first_key_and_last_salt() { + let middle = "a1".repeat(20); + let buf = format!("x'{}{}{}'", ENC_HEX, middle, SALT_HEX).into_bytes(); + let keys = scan_chunk(&buf); + assert_eq!(keys.len(), 1); + assert_eq!(keys[0].enc_key, expected_enc_key()); + assert_eq!(keys[0].salt, Some(expected_salt())); + } + + #[test] + fn invalid_mid_length_hex_is_ignored() { + let payload = format!("{}{}", ENC_HEX, "ab".repeat(8)); // 80 hex + let buf = format!("x'{}'", payload).into_bytes(); + let keys = scan_chunk(&buf); + assert!(keys.is_empty()); + } + + #[test] + fn mixed_case_hex_decoded_correctly() { + let mixed_enc: String = ENC_HEX + .chars() + .enumerate() + .map(|(i, c)| { + if i % 2 == 0 { + c.to_ascii_uppercase() + } else { + c + } + }) + .collect(); + let mixed_salt: String = SALT_HEX + .chars() + .enumerate() + .map(|(i, c)| { + if i % 2 == 0 { + c.to_ascii_uppercase() + } else { + c + } + }) + .collect(); + let buf = format!("x'{}{}'", mixed_enc, mixed_salt).into_bytes(); + let keys = scan_chunk(&buf); + assert_eq!(keys.len(), 1); + assert_eq!(keys[0].enc_key, expected_enc_key()); + assert_eq!(keys[0].salt, Some(expected_salt())); + } + + #[test] + fn duplicate_patterns_are_deduplicated() { + let pattern = format!("x'{}{}'", ENC_HEX, SALT_HEX); + let buf = format!("{}__{}", pattern, pattern).into_bytes(); + let keys = scan_chunk(&buf); + assert_eq!(keys.len(), 1); + } + + #[test] + fn incomplete_pattern_at_end_is_ignored() { + let buf = format!("x'{}", ENC_HEX).into_bytes(); + let keys = scan_chunk(&buf); + assert!(keys.is_empty()); + } +} diff --git a/crates/wx-keychain/src/mach_vm/reader.rs b/crates/wx-keychain/src/mach_vm/reader.rs new file mode 100644 index 0000000..36cfcb7 --- /dev/null +++ b/crates/wx-keychain/src/mach_vm/reader.rs @@ -0,0 +1,12 @@ +use crate::error::KeychainError; + +#[derive(Debug, Clone)] +pub struct MemRegion { + pub start: u64, + pub end: u64, +} + +pub trait MemoryReader { + fn rw_regions(&self) -> Result, KeychainError>; + fn read_bytes(&self, addr: u64, len: usize) -> Result, KeychainError>; +} diff --git a/crates/wx-keychain/src/mach_vm/scanner.rs b/crates/wx-keychain/src/mach_vm/scanner.rs new file mode 100644 index 0000000..d9a5eb3 --- /dev/null +++ b/crates/wx-keychain/src/mach_vm/scanner.rs @@ -0,0 +1,423 @@ +use std::collections::HashSet; +use std::path::Path; + +use crate::error::KeychainError; +use crate::mach_vm::pattern::{scan_chunk, FoundKey, MAX_PATTERN_BYTES}; +use crate::mach_vm::reader::MemoryReader; + +/// The maximum supported pattern is `x'<192 hex>'` = 195 bytes. +/// Keep `MAX_PATTERN_BYTES - 1` bytes so chunk-boundary matches are not missed. +const OVERLAP: usize = MAX_PATTERN_BYTES - 1; + +/// Default chunk size for reading memory regions. +const CHUNK_SIZE: usize = 2 * 1024 * 1024; // 2 MiB + +/// A validated scan result: enc_key + salt matched to a specific DB file. +#[derive(Debug, Clone)] +pub struct ScanResult { + pub enc_key: [u8; 32], + pub salt: [u8; 16], + pub db_path: std::path::PathBuf, +} + +/// Scan process memory for enc_key candidates, match them against known DB salts, +/// and HMAC-validate each match. +pub struct MemoryScanner { + reader: R, +} + +impl MemoryScanner { + pub fn new(reader: R) -> Self { + Self { reader } + } + + /// Scan all RW regions for supported SQL hex literal patterns, then match + /// candidates against known DB salts via HMAC validation. + pub fn scan( + &self, + db_salts: &[([u8; 16], &Path)], + params: &wx_decrypt::CryptoParams, + ) -> Result, KeychainError> { + let regions = self.reader.rw_regions()?; + let candidates = self.scan_regions(®ions)?; + + if candidates.is_empty() { + return Ok(Vec::new()); + } + + let mut results = Vec::new(); + let mut seen = HashSet::new(); + + for found in &candidates { + self.validate_candidate(found, db_salts, params, &mut results, &mut seen); + } + + self.cross_validate_known_keys(db_salts, params, &mut results, &mut seen); + + Ok(results) + } + + fn validate_candidate( + &self, + found: &FoundKey, + db_salts: &[([u8; 16], &Path)], + params: &wx_decrypt::CryptoParams, + results: &mut Vec, + seen: &mut HashSet<([u8; 32], [u8; 16])>, + ) { + match found.salt { + Some(salt_hint) => { + for (salt, db_path) in db_salts { + if *salt != salt_hint { + continue; + } + if validate_key_for_db(&found.enc_key, salt, db_path, params) + && seen.insert((found.enc_key, *salt)) + { + results.push(ScanResult { + enc_key: found.enc_key, + salt: *salt, + db_path: db_path.to_path_buf(), + }); + } + break; + } + } + None => { + for (salt, db_path) in db_salts { + if validate_key_for_db(&found.enc_key, salt, db_path, params) + && seen.insert((found.enc_key, *salt)) + { + results.push(ScanResult { + enc_key: found.enc_key, + salt: *salt, + db_path: db_path.to_path_buf(), + }); + } + } + } + } + } + + fn cross_validate_known_keys( + &self, + db_salts: &[([u8; 16], &Path)], + params: &wx_decrypt::CryptoParams, + results: &mut Vec, + seen: &mut HashSet<([u8; 32], [u8; 16])>, + ) { + if results.is_empty() { + return; + } + + let matched_salts: HashSet<[u8; 16]> = results.iter().map(|r| r.salt).collect(); + let known_keys: HashSet<[u8; 32]> = results.iter().map(|r| r.enc_key).collect(); + + for (salt, db_path) in db_salts { + if matched_salts.contains(salt) { + continue; + } + for enc_key in &known_keys { + if validate_key_for_db(enc_key, salt, db_path, params) + && seen.insert((*enc_key, *salt)) + { + results.push(ScanResult { + enc_key: *enc_key, + salt: *salt, + db_path: db_path.to_path_buf(), + }); + break; + } + } + } + } + + /// Scan all regions, returning deduplicated candidates. + fn scan_regions( + &self, + regions: &[crate::mach_vm::MemRegion], + ) -> Result, KeychainError> { + let mut seen = HashSet::new(); + let mut all_keys = Vec::new(); + + for region in regions { + let region_len = (region.end - region.start) as usize; + if region_len == 0 { + continue; + } + + let mut offset: u64 = 0; + let mut carry: Vec = Vec::new(); + + while (offset as usize) < region_len { + let read_len = CHUNK_SIZE.min(region_len - offset as usize); + let chunk = match self.reader.read_bytes(region.start + offset, read_len) { + Ok(data) => data, + Err(_) => { + offset += read_len as u64; + carry.clear(); + continue; + } + }; + + let scan_buf = if carry.is_empty() { + chunk.clone() + } else { + let mut buf = carry.clone(); + buf.extend_from_slice(&chunk); + buf + }; + + for found in scan_chunk(&scan_buf) { + if seen.insert(found.clone()) { + all_keys.push(found); + } + } + + carry = if chunk.len() > OVERLAP { + chunk[chunk.len() - OVERLAP..].to_vec() + } else { + chunk.clone() + }; + + offset += read_len as u64; + } + } + + Ok(all_keys) + } +} + +fn validate_key_for_db( + enc_key: &[u8; 32], + salt: &[u8; 16], + db_path: &Path, + params: &wx_decrypt::CryptoParams, +) -> bool { + let first_page = match std::fs::read(db_path) { + Ok(data) if data.len() >= params.page_size => data[..params.page_size].to_vec(), + _ => return false, + }; + + wx_decrypt::validate_enc_key(&first_page, enc_key, salt, params) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::mach_vm::reader::MemRegion; + use wx_decrypt::MACOS_4_1_7_31; + + struct MockReader { + regions: Vec, + data: Vec, + } + + impl MemoryReader for MockReader { + fn rw_regions(&self) -> Result, KeychainError> { + Ok(self.regions.clone()) + } + + fn read_bytes(&self, addr: u64, len: usize) -> Result, KeychainError> { + let start = addr as usize; + let end = (start + len).min(self.data.len()); + if start >= self.data.len() { + return Err(KeychainError::Other("out of bounds".into())); + } + Ok(self.data[start..end].to_vec()) + } + } + + fn make_pattern(enc_key: &[u8; 32], salt: &[u8; 16]) -> Vec { + format!("x'{}{}'", hex::encode(enc_key), hex::encode(salt)).into_bytes() + } + + fn make_key_only_pattern(enc_key: &[u8; 32]) -> Vec { + format!("x'{}'", hex::encode(enc_key)).into_bytes() + } + + fn make_long_pattern(enc_key: &[u8; 32], middle_hex: &str, salt: &[u8; 16]) -> Vec { + format!( + "x'{}{}{}'", + hex::encode(enc_key), + middle_hex, + hex::encode(salt) + ) + .into_bytes() + } + + fn build_first_page(enc_key: &[u8; 32], salt: &[u8; 16]) -> Vec { + use aes::cipher::{block_padding::NoPadding, BlockEncryptMut, KeyIvInit}; + use hmac::{Hmac, Mac}; + use sha2::Sha512; + + let params = &MACOS_4_1_7_31; + let iv = [0x42u8; 16]; + let data_len = params.page_size - params.reserve - params.salt_size; + let plaintext = vec![0u8; data_len]; + + let mut ciphertext = plaintext; + type Aes256CbcEnc = cbc::Encryptor; + Aes256CbcEnc::new(enc_key.into(), (&iv).into()) + .encrypt_padded_mut::(&mut ciphertext, data_len) + .unwrap(); + + let mut page = vec![0u8; params.page_size]; + page[..params.salt_size].copy_from_slice(salt); + page[params.salt_size..params.salt_size + data_len].copy_from_slice(&ciphertext); + + let iv_start = params.page_size - params.reserve; + page[iv_start..iv_start + params.iv_size].copy_from_slice(&iv); + + let hmac_data_end = params.page_size - params.reserve + params.iv_size; + let mac_key = wx_decrypt::kdf::derive_mac_key(enc_key, salt, params); + let mut mac = as Mac>::new_from_slice(&mac_key).unwrap(); + mac.update(&page[params.salt_size..hmac_data_end]); + mac.update(&1u32.to_le_bytes()); + let hmac_result = mac.finalize().into_bytes(); + page[hmac_data_end..hmac_data_end + params.hmac_size] + .copy_from_slice(&hmac_result[..params.hmac_size]); + + page + } + + #[test] + fn mock_reader_finds_valid_pattern() { + let enc_key = [0xABu8; 32]; + let salt = [0x01u8; 16]; + + let dir = tempfile::tempdir().unwrap(); + let db_path = dir.path().join("test.db"); + std::fs::write(&db_path, build_first_page(&enc_key, &salt)).unwrap(); + + let mut data = vec![0u8; 1000]; + let pattern = make_pattern(&enc_key, &salt); + data[100..100 + pattern.len()].copy_from_slice(&pattern); + + let scanner = MemoryScanner::new(MockReader { + regions: vec![MemRegion { + start: 0, + end: data.len() as u64, + }], + data, + }); + let results = scanner.scan(&[(salt, &db_path)], &MACOS_4_1_7_31).unwrap(); + assert_eq!(results.len(), 1); + assert_eq!(results[0].enc_key, enc_key); + assert_eq!(results[0].salt, salt); + } + + #[test] + fn key_only_pattern_validates_against_known_dbs() { + let enc_key = [0xABu8; 32]; + let salt = [0x01u8; 16]; + + let dir = tempfile::tempdir().unwrap(); + let db_path = dir.path().join("test.db"); + std::fs::write(&db_path, build_first_page(&enc_key, &salt)).unwrap(); + + let mut data = vec![0u8; 1000]; + let pattern = make_key_only_pattern(&enc_key); + data[100..100 + pattern.len()].copy_from_slice(&pattern); + + let scanner = MemoryScanner::new(MockReader { + regions: vec![MemRegion { + start: 0, + end: data.len() as u64, + }], + data, + }); + let results = scanner.scan(&[(salt, &db_path)], &MACOS_4_1_7_31).unwrap(); + assert_eq!(results.len(), 1); + assert_eq!(results[0].salt, salt); + } + + #[test] + fn long_hex_pattern_uses_first_key_and_last_salt() { + let enc_key = [0xCDu8; 32]; + let salt = [0x02u8; 16]; + + let dir = tempfile::tempdir().unwrap(); + let db_path = dir.path().join("test.db"); + std::fs::write(&db_path, build_first_page(&enc_key, &salt)).unwrap(); + + let mut data = vec![0u8; 2048]; + let pattern = make_long_pattern(&enc_key, &"a1".repeat(20), &salt); + data[300..300 + pattern.len()].copy_from_slice(&pattern); + + let scanner = MemoryScanner::new(MockReader { + regions: vec![MemRegion { + start: 0, + end: data.len() as u64, + }], + data, + }); + let results = scanner.scan(&[(salt, &db_path)], &MACOS_4_1_7_31).unwrap(); + assert_eq!(results.len(), 1); + assert_eq!(results[0].enc_key, enc_key); + assert_eq!(results[0].salt, salt); + } + + #[test] + fn cross_validation_reuses_known_key_for_missing_salt() { + let enc_key = [0xAAu8; 32]; + let salt1 = [0x01u8; 16]; + let salt2 = [0x02u8; 16]; + + let dir = tempfile::tempdir().unwrap(); + let db1 = dir.path().join("db1.db"); + let db2 = dir.path().join("db2.db"); + std::fs::write(&db1, build_first_page(&enc_key, &salt1)).unwrap(); + std::fs::write(&db2, build_first_page(&enc_key, &salt2)).unwrap(); + + let mut data = vec![0u8; 1000]; + let pattern = make_pattern(&enc_key, &salt1); + data[100..100 + pattern.len()].copy_from_slice(&pattern); + + let scanner = MemoryScanner::new(MockReader { + regions: vec![MemRegion { + start: 0, + end: data.len() as u64, + }], + data, + }); + let results = scanner + .scan( + &[(salt1, db1.as_path()), (salt2, db2.as_path())], + &MACOS_4_1_7_31, + ) + .unwrap(); + + assert_eq!(results.len(), 2); + let salts: Vec<_> = results.iter().map(|r| r.salt).collect(); + assert!(salts.contains(&salt1)); + assert!(salts.contains(&salt2)); + } + + #[test] + fn pattern_spanning_chunks_found_via_overlap() { + let enc_key = [0xEFu8; 32]; + let salt = [0x03u8; 16]; + + let dir = tempfile::tempdir().unwrap(); + let db_path = dir.path().join("test.db"); + std::fs::write(&db_path, build_first_page(&enc_key, &salt)).unwrap(); + + let pattern = make_long_pattern(&enc_key, &"a1".repeat(20), &salt); + let split_point = CHUNK_SIZE - 100; + let total_len = split_point + pattern.len(); + let mut data = vec![0u8; total_len]; + data[split_point..split_point + pattern.len()].copy_from_slice(&pattern); + + let scanner = MemoryScanner::new(MockReader { + regions: vec![MemRegion { + start: 0, + end: data.len() as u64, + }], + data, + }); + let results = scanner.scan(&[(salt, &db_path)], &MACOS_4_1_7_31).unwrap(); + assert_eq!(results.len(), 1); + assert_eq!(results[0].salt, salt); + } +} diff --git a/crates/wx-keychain/src/nickname.rs b/crates/wx-keychain/src/nickname.rs new file mode 100644 index 0000000..cfdee7a --- /dev/null +++ b/crates/wx-keychain/src/nickname.rs @@ -0,0 +1,133 @@ +use std::path::{Path, PathBuf}; +use std::process::Command; +use std::time::{SystemTime, UNIX_EPOCH}; + +use wx_decrypt::KeyMaterial; + +use crate::error::KeychainError; + +/// Try to resolve an account's nickname from its contact.db. +/// +/// This is best-effort: returns `Ok(None)` if the nickname cannot be resolved +/// (e.g. contact.db doesn't exist, decryption fails, no matching record). +/// Errors are only returned for unexpected I/O failures. +pub fn resolve_nickname( + data_dir: &Path, + key_material: &KeyMaterial, + base_wxid: &str, +) -> Result, KeychainError> { + let contact_db = data_dir.join("db_storage/contact/contact.db"); + if !contact_db.exists() { + return Ok(None); + } + + let params = &wx_decrypt::MACOS_4_1_7_31; + let tmp_db = unique_temp_db_path(); + + // Decrypt contact.db to temp file + let decrypt_result = match key_material { + KeyMaterial::RawKey(key) => wx_decrypt::decrypt_db(&contact_db, &tmp_db, key, params), + KeyMaterial::EncKey { key, salt } => { + wx_decrypt::decrypt_db_direct(&contact_db, &tmp_db, key, salt, params) + } + KeyMaterial::EncKeys(pairs) => { + match wx_decrypt::read_main_db_salt_for_path(&contact_db) { + Ok(db_salt) => { + match pairs.iter().find(|p| p.salt == db_salt) { + Some(pair) => wx_decrypt::decrypt_db_direct( + &contact_db, + &tmp_db, + &pair.key, + &pair.salt, + params, + ), + None => return Ok(None), // best-effort: no matching key + } + } + Err(_) => return Ok(None), // best-effort: can't read salt + } + } + }; + + match decrypt_result { + Ok(()) => {} + Err(wx_decrypt::DecryptError::AlreadyDecrypted) => { + if std::fs::copy(&contact_db, &tmp_db).is_err() { + return Ok(None); + } + } + Err(_) => return Ok(None), + } + + // Query nickname using sqlite3 CLI + let result = query_nickname_from_db(&tmp_db, base_wxid); + + // Clean up temp files + let _ = std::fs::remove_file(&tmp_db); + let _ = std::fs::remove_file(with_suffix(&tmp_db, "-wal")); + let _ = std::fs::remove_file(with_suffix(&tmp_db, "-shm")); + + result +} + +fn query_nickname_from_db( + db_path: &Path, + base_wxid: &str, +) -> Result, KeychainError> { + let query = format!( + "SELECT COALESCE(\ + NULLIF(remark, ''),\ + NULLIF(nick_name, ''),\ + NULLIF(alias, '')\ + ) FROM contact WHERE username = '{}' LIMIT 1;", + base_wxid.replace('\'', "''") + ); + + let output = Command::new("sqlite3") + .args([db_path.to_str().unwrap_or(""), &query]) + .output(); + + match output { + Ok(out) if out.status.success() => { + let name = String::from_utf8_lossy(&out.stdout).trim().to_string(); + if name.is_empty() { + Ok(None) + } else { + Ok(Some(name)) + } + } + _ => Ok(None), + } +} + +fn unique_temp_db_path() -> PathBuf { + let nanos = SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|d| d.as_nanos()) + .unwrap_or(0); + let path = wx_paths::AppPaths::nickname_temp_db(std::process::id(), nanos); + if let Some(parent) = path.parent() { + let _ = wx_paths::AppPaths::ensure_dir(parent); + } + path +} + +fn with_suffix(path: &Path, suffix: &str) -> PathBuf { + PathBuf::from(format!("{}{}", path.display(), suffix)) +} + +// ---- Test coverage notes ---- +// +// `resolve_nickname()` EncKeys branch is not directly unit-tested because: +// 1. It requires constructing an encrypted contact.db with a valid `contact` table +// schema, decrypting it, then querying via sqlite3 CLI — high integration cost. +// 2. The EncKeys branch only calls `read_main_db_salt_for_path()` (tested in +// wx-decrypt db.rs) and `decrypt_db_direct()` (tested in wx-decrypt db.rs). +// 3. The salt-matching + best-effort `Ok(None)` fallback is a trivial code path. +// +// Indirect coverage: +// - `read_main_db_salt_for_path()`: 3 tests in wx-decrypt/src/db.rs +// - `decrypt_db_direct()`: 2 tests in wx-decrypt/src/db.rs +// - EncKeys salt matching: tested in wx-context/src/cache.rs +// (`enc_keys_decrypt_db_selects_matching_pair`, `enc_keys_decrypt_db_no_match_returns_error`) +// - E2E coverage via VM test scenario 4 (key scan + decrypt with EncKeys) diff --git a/crates/wx-keychain/src/process.rs b/crates/wx-keychain/src/process.rs new file mode 100644 index 0000000..4ccda0c --- /dev/null +++ b/crates/wx-keychain/src/process.rs @@ -0,0 +1,466 @@ +use std::path::{Path, PathBuf}; +use std::process::Command; + +use crate::error::KeychainError; + +pub const SUPPORTED_VERSION: &str = "4.1.8.21"; + +/// Version prefixes accepted for LLDB key extraction. +/// Encryption params (PBKDF2-HMAC-SHA512, 256K iterations) are identical across these versions. +const EXTRACTION_VERSION_PREFIXES: &[&str] = &["4.1.7", "4.1.8"]; + +/// Check whether a version string is compatible with our LLDB key extraction. +fn is_extraction_compatible(version: &str) -> bool { + EXTRACTION_VERSION_PREFIXES + .iter() + .any(|prefix| version == *prefix || version.starts_with(&format!("{prefix}."))) +} + +#[derive(Debug, Clone)] +pub struct AccountDirInfo { + /// Directory name, e.g. "wxid_example123abc_ab12" + pub account_id: String, + /// Normalized base wxid, e.g. "wxid_example123abc" + pub base_wxid: String, + /// Full path to the account data directory + pub data_dir: PathBuf, + /// Path to message_0.db + pub message_db_path: PathBuf, +} + +#[derive(Debug, Clone, Copy)] +pub enum DetectionSource { + RunningProcess, + LoginKeyInfoMtime, +} + +impl std::fmt::Display for DetectionSource { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + DetectionSource::RunningProcess => write!(f, "running-process"), + DetectionSource::LoginKeyInfoMtime => write!(f, "login-key-info-mtime"), + } + } +} + +pub struct ActiveAccount { + pub info: AccountDirInfo, + pub source: DetectionSource, +} + +/// Extract base wxid from an account directory name (conservative). +/// +/// Only strips suffix for `wxid_*` prefix accounts. For arbitrary strings +/// (e.g. user CLI input), this is safe. For confirmed account directories, +/// use `extract_base_wxid_for_account_dir()` instead. +pub fn extract_base_wxid(account_id: &str) -> String { + crate::account_id::canonical_base(account_id) +} + +/// Extract base wxid from a confirmed account directory name (aggressive). +/// +/// Remains conservative for non-`wxid_` inputs unless the caller can provide +/// an independent confirmation signal via `extract_base_wxid_for_account_dir_under_root()`. +pub fn extract_base_wxid_for_account_dir(account_id: &str) -> String { + crate::account_id::canonical_base_for_account_dir(account_id) +} + +/// Extract base wxid from a confirmed account directory using sibling +/// `all_users/login/` directories as the confirmation signal. +pub fn extract_base_wxid_for_account_dir_under_root( + xwechat_root: &Path, + account_id: &str, +) -> String { + let account = crate::account_id::AccountId::parse(account_id); + let confirmed = read_login_names(xwechat_root) + .into_iter() + .find(|login| account.alias_candidate() == Some(login.as_str())); + crate::account_id::canonical_base_for_account_dir_with_confirmed_base( + account_id, + confirmed.as_deref(), + ) +} + +/// Find a running WeChat process PID and validate its version. +/// +/// Uses `pgrep` to find the PID and checks the installed WeChat version. +/// Does NOT use `lsof` — suitable for commands that only need PID + version. +pub fn find_wechat_pid() -> Result<(u32, String), KeychainError> { + let pgrep_output = Command::new("pgrep").args(["-x", "WeChat"]).output()?; + + if !pgrep_output.status.success() { + return Err(KeychainError::WeChatNotRunning); + } + + let pid_str = String::from_utf8_lossy(&pgrep_output.stdout) + .lines() + .next() + .unwrap_or("") + .trim() + .to_string(); + + let pid: u32 = pid_str + .parse() + .map_err(|_| KeychainError::WeChatNotRunning)?; + + let version = ensure_supported_wechat_version()?; + + Ok((pid, version)) +} + +/// Ensure the installed WeChat version matches the supported target. +pub fn ensure_supported_wechat_version() -> Result { + let version = get_wechat_version()?; + if !is_extraction_compatible(&version) { + return Err(KeychainError::UnsupportedVersion { version }); + } + Ok(version) +} + +/// Get WeChat version from the application bundle. +fn get_wechat_version() -> Result { + let output = Command::new("defaults") + .args([ + "read", + "/Applications/WeChat.app/Contents/Info.plist", + "CFBundleShortVersionString", + ]) + .output()?; + + if !output.status.success() { + return Err(KeychainError::Other( + "failed to read WeChat version from Info.plist".into(), + )); + } + + Ok(String::from_utf8_lossy(&output.stdout).trim().to_string()) +} + +/// Shared config directory resolved by `AppPaths`. +pub fn config_dir() -> Result { + let ap = wx_paths::AppPaths::new() + .map_err(|e| KeychainError::Other(e.to_string()))?; + Ok(ap.config_dir()) +} + +/// Default xwechat_files base path. +fn default_xwechat_files_base() -> Result { + let ap = wx_paths::AppPaths::new() + .map_err(|e| KeychainError::Other(e.to_string()))?; + Ok(ap.home().join("Library/Containers/com.tencent.xinWeChat/Data/Documents/xwechat_files")) +} + +/// Detect account directories from the filesystem (without WeChat running). +/// +/// Scans `~/Library/Containers/com.tencent.xinWeChat/Data/Documents/xwechat_files/` +/// for subdirectories containing `db_storage/message/message_0.db`. +pub fn find_account_dirs() -> Result, KeychainError> { + let base = default_xwechat_files_base()?; + if !base.exists() { + return Ok(vec![]); + } + find_account_dirs_under(&base) +} + +/// Detect account directories under a specified root path. +/// +/// Scans `base` for subdirectories containing `db_storage/message/message_0.db`. +pub fn find_account_dirs_under(base: &Path) -> Result, KeychainError> { + let mut accounts = Vec::new(); + for entry in std::fs::read_dir(base)? { + let entry = entry?; + if entry.file_type()?.is_dir() { + let name = entry.file_name().to_string_lossy().to_string(); + let db_path = entry.path().join("db_storage/message/message_0.db"); + if db_path.exists() { + accounts.push(AccountDirInfo { + base_wxid: extract_base_wxid_for_account_dir_under_root(base, &name), + account_id: name, + data_dir: entry.path(), + message_db_path: db_path, + }); + } + } + } + Ok(accounts) +} + +/// Check if a path is an xwechat_files root directory (contains `all_users/` subdirectory). +pub fn is_xwechat_files_root(path: &Path) -> bool { + path.join("all_users").is_dir() +} + +/// Check whether WeChat is currently running via pgrep. +/// +/// Unlike `find_wechat_pid()`, this does NOT perform version validation, +/// so it won't block commands (sessions, query, etc.) that don't depend on version. +fn is_wechat_running() -> bool { + Command::new("pgrep") + .args(["-x", "WeChat"]) + .output() + .is_ok_and(|o| o.status.success()) +} + +/// Detect the currently active WeChat account. +/// +/// Uses mtime of `all_users/login/*/key_info.db-wal` to identify the active account. +/// `pgrep` is only used to determine `DetectionSource` (running-process vs login-key-info-mtime). +pub fn detect_active_account(accounts: &[AccountDirInfo]) -> Result { + if accounts.is_empty() { + return Err(KeychainError::AccountDetectionFailed { + reason: "no account directories provided".into(), + candidates: String::new(), + }); + } + + let wechat_running = is_wechat_running(); + let source = if wechat_running { + DetectionSource::RunningProcess + } else { + DetectionSource::LoginKeyInfoMtime + }; + + // Mtime strategy: find most recent login key_info.db + let xwechat_root = accounts[0] + .data_dir + .parent() + .ok_or_else(|| KeychainError::Other("cannot determine xwechat_files root".into()))?; + + let login_dir = xwechat_root.join("all_users/login"); + if login_dir.is_dir() { + if let Ok(best_login) = find_most_recent_login(&login_dir) { + // Match login name to account directories using alias-aware matching + let matches: Vec<&AccountDirInfo> = accounts + .iter() + .filter(|a| { + let id = crate::account_id::AccountId::parse(&a.account_id); + id.matches(&best_login) + }) + .collect(); + + match matches.len() { + 1 => { + return Ok(ActiveAccount { + info: matches[0].clone(), + source, + }); + } + n if n > 1 => { + // Tiebreak by message_0.db-wal mtime + if let Some(best) = tiebreak_by_wal_mtime(&matches) { + return Ok(ActiveAccount { + info: best.clone(), + source, + }); + } + // Still ambiguous + let candidates = matches + .iter() + .map(|a| format!(" - {}", a.account_id)) + .collect::>() + .join("\n"); + return Err(KeychainError::AccountDetectionFailed { + reason: format!("multiple account directories match login '{best_login}'"), + candidates, + }); + } + _ => {} // no match, fall through + } + } + } + + // If only one account exists, use it + if accounts.len() == 1 { + return Ok(ActiveAccount { + info: accounts[0].clone(), + source, + }); + } + + // All strategies failed + let candidates = accounts + .iter() + .map(|a| format!(" - {}", a.account_id)) + .collect::>() + .join("\n"); + Err(KeychainError::AccountDetectionFailed { + reason: "login mtime detection failed".into(), + candidates, + }) +} + +/// Find the login subdirectory with the most recent key_info.db-wal mtime. +fn find_most_recent_login(login_dir: &Path) -> Result { + let mut best_name: Option = None; + let mut best_mtime: Option = None; + + for entry in std::fs::read_dir(login_dir)? { + let entry = entry?; + if !entry.file_type()?.is_dir() { + continue; + } + let name = entry.file_name().to_string_lossy().to_string(); + + let wal_path = entry.path().join("key_info.db-wal"); + if let Ok(meta) = std::fs::metadata(&wal_path) { + if let Ok(mtime) = meta.modified() { + if best_mtime.is_none_or(|t| mtime > t) { + best_mtime = Some(mtime); + best_name = Some(name); + } + } + } + } + + best_name.ok_or_else(|| { + KeychainError::Other("no login directories with key_info.db-wal found".into()) + }) +} + +fn read_login_names(xwechat_root: &Path) -> Vec { + let login_dir = xwechat_root.join("all_users/login"); + let Ok(entries) = std::fs::read_dir(&login_dir) else { + return Vec::new(); + }; + + entries + .filter_map(|entry| entry.ok()) + .filter(|entry| entry.file_type().map(|ft| ft.is_dir()).unwrap_or(false)) + .map(|entry| entry.file_name().to_string_lossy().to_string()) + .collect() +} + +/// Tiebreak multiple matching accounts by message_0.db-wal mtime. +fn tiebreak_by_wal_mtime<'a>(accounts: &[&'a AccountDirInfo]) -> Option<&'a AccountDirInfo> { + let mut best: Option<&AccountDirInfo> = None; + let mut best_mtime: Option = None; + + for account in accounts { + let wal_path = account.data_dir.join("db_storage/message/message_0.db-wal"); + if let Ok(meta) = std::fs::metadata(&wal_path) { + if let Ok(mtime) = meta.modified() { + if best_mtime.is_none_or(|t| mtime > t) { + best_mtime = Some(mtime); + best = Some(account); + } + } + } + } + + best +} + + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_extract_base_wxid_with_suffix() { + assert_eq!( + extract_base_wxid("wxid_example123abc_ab12"), + "wxid_example123abc" + ); + } + + #[test] + fn test_extract_base_wxid_no_suffix() { + assert_eq!(extract_base_wxid("wxid_xxx"), "wxid_xxx"); + } + + #[test] + fn test_extract_base_wxid_not_wxid() { + // Conservative: non-wxid pattern is NOT stripped + assert_eq!(extract_base_wxid("not_a_wxid"), "not_a_wxid"); + } + + #[test] + fn test_extract_base_wxid_bare_wxid() { + assert_eq!(extract_base_wxid("wxid_x"), "wxid_x"); + } + + #[test] + fn test_extract_base_wxid_multiple_underscores() { + assert_eq!( + extract_base_wxid("wxid_foobar456def_c3e7"), + "wxid_foobar456def" + ); + } + + #[test] + fn test_extract_base_wxid_legacy_account_conservative() { + // Conservative: non-wxid not stripped + assert_eq!(extract_base_wxid("testuser001_1662"), "testuser001_1662"); + } + + #[test] + fn test_extract_base_wxid_legacy_account_confirmed_dir() { + // Confirmed dir without extra signal: still conservative + assert_eq!( + extract_base_wxid_for_account_dir("testuser001_1662"), + "testuser001_1662" + ); + } + + #[test] + fn test_find_account_dirs_under_keeps_raw_base_without_login_signal() { + let tmp = tempfile::tempdir().unwrap(); + let account_dir = tmp.path().join("not_a_wxid"); + let db_dir = account_dir.join("db_storage/message"); + std::fs::create_dir_all(&db_dir).unwrap(); + std::fs::write(db_dir.join("message_0.db"), b"fake").unwrap(); + + let accounts = find_account_dirs_under(tmp.path()).unwrap(); + assert_eq!(accounts.len(), 1); + assert_eq!(accounts[0].account_id, "not_a_wxid"); + assert_eq!(accounts[0].base_wxid, "not_a_wxid"); + } + + #[test] + fn test_find_account_dirs_under_uses_login_dir_as_confirmation() { + let tmp = tempfile::tempdir().unwrap(); + let account_dir = tmp.path().join("testuser001_1662"); + let db_dir = account_dir.join("db_storage/message"); + std::fs::create_dir_all(&db_dir).unwrap(); + std::fs::write(db_dir.join("message_0.db"), b"fake").unwrap(); + std::fs::create_dir_all(tmp.path().join("all_users/login/testuser001")).unwrap(); + + let accounts = find_account_dirs_under(tmp.path()).unwrap(); + assert_eq!(accounts.len(), 1); + assert_eq!(accounts[0].account_id, "testuser001_1662"); + assert_eq!(accounts[0].base_wxid, "testuser001"); + } + + #[test] + fn test_extract_base_wxid_wxid_test_not_stripped() { + assert_eq!(extract_base_wxid("wxid_test"), "wxid_test"); + } + + #[test] + fn test_is_xwechat_files_root_with_all_users() { + let tmp = std::env::temp_dir().join("test_xwechat_root"); + let _ = std::fs::create_dir_all(tmp.join("all_users")); + assert!(is_xwechat_files_root(&tmp)); + let _ = std::fs::remove_dir_all(&tmp); + } + + #[test] + fn test_is_xwechat_files_root_without_all_users() { + let tmp = std::env::temp_dir().join("test_xwechat_no_root"); + let _ = std::fs::create_dir_all(&tmp); + assert!(!is_xwechat_files_root(&tmp)); + let _ = std::fs::remove_dir_all(&tmp); + } + + #[test] + fn test_detection_source_display_all_variants() { + assert_eq!( + format!("{}", DetectionSource::RunningProcess), + "running-process" + ); + assert_eq!( + format!("{}", DetectionSource::LoginKeyInfoMtime), + "login-key-info-mtime" + ); + } +} diff --git a/crates/wx-keychain/src/script.rs b/crates/wx-keychain/src/script.rs new file mode 100644 index 0000000..4bb0eff --- /dev/null +++ b/crates/wx-keychain/src/script.rs @@ -0,0 +1,55 @@ +/// Embedded LLDB Python script for capturing PBKDF2 calls during WeChat startup. +/// +/// This is written to a temp file at runtime and imported into LLDB. +/// +/// Hooks `CCKeyDerivationPBKDF` and prints password/salt for rounds=256000 calls. +pub const CAPTURE_KEY_SCRIPT: &str = r#"import lldb +import binascii + +keys_found = [] +call_count = 0 +module_name = __name__ + +def pbkdf_callback(frame, bp_loc, dict): + global call_count, keys_found + call_count += 1 + process = frame.GetThread().GetProcess() + gpr = frame.GetRegisters()[0] + + pwd_ptr = gpr.GetChildMemberWithName("x1").GetValueAsUnsigned() + pwd_len = gpr.GetChildMemberWithName("x2").GetValueAsUnsigned() + salt_ptr = gpr.GetChildMemberWithName("x3").GetValueAsUnsigned() + salt_len = gpr.GetChildMemberWithName("x4").GetValueAsUnsigned() + prf = gpr.GetChildMemberWithName("x5").GetValueAsUnsigned() + rounds = gpr.GetChildMemberWithName("x6").GetValueAsUnsigned() + + error = lldb.SBError() + pwd_hex, salt_hex = "", "" + if 0 < pwd_len < 1024: + d = process.ReadMemory(pwd_ptr, pwd_len, error) + if error.Success(): pwd_hex = binascii.hexlify(d).decode() + if 0 < salt_len < 1024: + d = process.ReadMemory(salt_ptr, salt_len, error) + if error.Success(): salt_hex = binascii.hexlify(d).decode() + + prf_names = {3: "SHA1", 4: "SHA256", 5: "SHA512"} + print(f"[PBKDF2 #{call_count}] PRF={prf_names.get(prf, prf)} rounds={rounds} " + f"pwdLen={pwd_len} saltLen={salt_len}", flush=True) + print(f" Password: {pwd_hex}", flush=True) + print(f" Salt: {salt_hex}", flush=True) + return False + +def setup(debugger, command, result, internal_dict): + target = debugger.GetSelectedTarget() + bp = target.BreakpointCreateByName("CCKeyDerivationPBKDF") + bp.SetScriptCallbackFunction(f"{module_name}.pbkdf_callback") + bp.SetAutoContinue(True) + print(f"Breakpoint on CCKeyDerivationPBKDF (id={bp.GetID()})", flush=True) + print("Resuming process...", flush=True) + target.GetProcess().Continue() + +def __lldb_init_module(debugger, internal_dict): + debugger.HandleCommand( + f'command script add -f {module_name}.setup capture_keys') + print("Run: capture_keys", flush=True) +"#; diff --git a/crates/wx-keychain/src/store.rs b/crates/wx-keychain/src/store.rs new file mode 100644 index 0000000..93794a6 --- /dev/null +++ b/crates/wx-keychain/src/store.rs @@ -0,0 +1,816 @@ +use std::collections::HashMap; +use std::fs; +use std::path::{Path, PathBuf}; + +use chrono::{DateTime, Utc}; +use serde::{Deserialize, Serialize}; + +use crate::error::KeychainError; +use wx_decrypt::{EncKeyPair, KeyMaterial}; + +/// Serializable enc_key + salt pair for the new per-DB format. +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] +pub struct EncKeyEntry { + pub enc_key: String, + pub salt: String, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct AccountKey { + pub account_id: String, + pub data_key: String, + pub extracted_at: DateTime, + pub wechat_version: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub nickname: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub base_wxid: Option, + /// V2 image AES key (16-byte hex string, independent of data_key). + #[serde(default, skip_serializing_if = "Option::is_none")] + pub image_aes_key: Option, + /// Legacy single pre-derived encryption key (64-char hex = 32 bytes). + #[serde(default, skip_serializing_if = "Option::is_none")] + pub enc_key: Option, + /// Legacy DB salt associated with the enc_key (32-char hex = 16 bytes). + #[serde(default, skip_serializing_if = "Option::is_none")] + pub enc_key_salt: Option, + /// Per-DB enc_key + salt pairs (new canonical format). + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub enc_keys: Vec, +} + +impl AccountKey { + pub fn display_name(&self) -> String { + match &self.nickname { + Some(nick) => format!("{} ({})", self.account_id, nick), + None => format!("{} (昵称未知)", self.account_id), + } + } +} + +#[derive(Debug, Default, Serialize, Deserialize)] +pub struct KeyStore { + #[serde(default)] + pub accounts: HashMap, +} + +impl KeyStore { + /// Default path: `/keys.toml` + pub fn default_path() -> Result { + let ap = wx_paths::AppPaths::new() + .map_err(|e| KeychainError::Other(e.to_string()))?; + Ok(ap.keys_file()) + } + + /// Load from the default path, creating an empty store if the file doesn't exist. + pub fn load_default() -> Result { + let ap = wx_paths::AppPaths::new() + .map_err(|e| KeychainError::Other(e.to_string()))?; + ap.migrate_config() + .map_err(|e| KeychainError::Other(format!("config migration failed: {}", e)))?; + let path = ap.keys_file(); + Self::load(&path) + } + + /// Load from a specific path. + pub fn load(path: &Path) -> Result { + if !path.exists() { + return Ok(Self::default()); + } + let content = fs::read_to_string(path)?; + toml::from_str(&content) + .map_err(|e| KeychainError::Store(format!("failed to parse {}: {}", path.display(), e))) + } + + /// Save to the default path. + pub fn save_default(&self) -> Result<(), KeychainError> { + let path = Self::default_path()?; + self.save(&path) + } + + /// Save to a specific path, creating parent directories as needed. + /// + /// Uses atomic write (write to `.tmp` sibling, then `rename`) to prevent + /// partial TOML files on crash. + pub fn save(&self, path: &Path) -> Result<(), KeychainError> { + let parent_existed = path.parent().is_none_or(|p| p.exists()); + if let Some(parent) = path.parent() { + fs::create_dir_all(parent)?; + } + let content = toml::to_string_pretty(self) + .map_err(|e| KeychainError::Store(format!("failed to serialize: {}", e)))?; + + let tmp_path = path.with_extension("toml.tmp"); + fs::write(&tmp_path, &content)?; + fs::rename(&tmp_path, path).inspect_err(|_| { + // Clean up the temp file on rename failure. + let _ = fs::remove_file(&tmp_path); + })?; + + wx_paths::sudo::chown_to_sudo_user(path); + if !parent_existed { + if let Some(parent) = path.parent() { + wx_paths::sudo::chown_to_sudo_user(parent); + } + } + Ok(()) + } + + /// Insert or update a key for an account. + /// + /// If `hex_key` differs from the existing `data_key`, derived key fields + /// (`enc_key`, `enc_key_salt`, `enc_keys`) are cleared because they were + /// derived from the old key and are no longer valid. + pub fn set( + &mut self, + account_id: &str, + hex_key: &str, + version: &str, + nickname: Option, + base_wxid: Option, + ) { + let existing = self.accounts.get(account_id); + let existing_image_key = existing.and_then(|k| k.image_aes_key.clone()); + + let data_key_changed = existing.map(|k| k.data_key != hex_key).unwrap_or(false); + + let (existing_enc_key, existing_enc_key_salt, existing_enc_keys) = if data_key_changed { + // data_key changed — derived keys are stale, clear them. + (None, None, Vec::new()) + } else { + ( + existing.and_then(|k| k.enc_key.clone()), + existing.and_then(|k| k.enc_key_salt.clone()), + existing.map(|k| k.enc_keys.clone()).unwrap_or_default(), + ) + }; + + self.accounts.insert( + account_id.to_string(), + AccountKey { + account_id: account_id.to_string(), + data_key: hex_key.to_string(), + extracted_at: Utc::now(), + wechat_version: version.to_string(), + nickname, + base_wxid, + image_aes_key: existing_image_key, + enc_key: existing_enc_key, + enc_key_salt: existing_enc_key_salt, + enc_keys: existing_enc_keys, + }, + ); + } + + /// Set or update the V2 image AES key for an account. + pub fn set_image_key(&mut self, account_id: &str, image_key_hex: &str) { + if let Some(entry) = self.accounts.get_mut(account_id) { + entry.image_aes_key = Some(image_key_hex.to_string()); + } else { + // Create a minimal entry if account doesn't exist yet. + self.accounts.insert( + account_id.to_string(), + AccountKey { + account_id: account_id.to_string(), + data_key: String::new(), + extracted_at: Utc::now(), + wechat_version: String::new(), + nickname: None, + base_wxid: None, + image_aes_key: Some(image_key_hex.to_string()), + enc_key: None, + enc_key_salt: None, + enc_keys: Vec::new(), + }, + ); + } + } + + /// Set or update the pre-derived enc_key and salt from Mach VM scan. + /// + /// Preserves existing `data_key` and `image_aes_key` if the account already exists. + pub fn set_enc_key( + &mut self, + account_id: &str, + enc_key_hex: &str, + salt_hex: &str, + version: &str, + nickname: Option, + base_wxid: Option, + ) { + if let Some(entry) = self.accounts.get_mut(account_id) { + entry.enc_key = Some(enc_key_hex.to_string()); + entry.enc_key_salt = Some(salt_hex.to_string()); + entry.extracted_at = Utc::now(); + entry.wechat_version = version.to_string(); + if nickname.is_some() { + entry.nickname = nickname; + } + if base_wxid.is_some() { + entry.base_wxid = base_wxid; + } + } else { + self.accounts.insert( + account_id.to_string(), + AccountKey { + account_id: account_id.to_string(), + data_key: String::new(), + extracted_at: Utc::now(), + wechat_version: version.to_string(), + nickname, + base_wxid, + image_aes_key: None, + enc_key: Some(enc_key_hex.to_string()), + enc_key_salt: Some(salt_hex.to_string()), + enc_keys: Vec::new(), + }, + ); + } + } + + /// Set or update multiple per-DB enc_key + salt pairs from Mach VM scan. + /// + /// Deduplicates and sorts entries by `(salt, enc_key)` for stable output. + /// Clears legacy `enc_key` / `enc_key_salt` fields, making `enc_keys` the canonical format. + /// Preserves existing `data_key` and `image_aes_key`. + pub fn set_enc_keys( + &mut self, + account_id: &str, + pairs: &[EncKeyPair], + version: &str, + nickname: Option, + base_wxid: Option, + ) { + // Deduplicate and sort + let mut entries: Vec = pairs + .iter() + .map(|p| EncKeyEntry { + enc_key: hex::encode(p.key), + salt: hex::encode(p.salt), + }) + .collect(); + entries.sort_by(|a, b| a.salt.cmp(&b.salt).then_with(|| a.enc_key.cmp(&b.enc_key))); + entries.dedup(); + + if let Some(entry) = self.accounts.get_mut(account_id) { + entry.enc_keys = entries; + entry.enc_key = None; + entry.enc_key_salt = None; + entry.extracted_at = Utc::now(); + entry.wechat_version = version.to_string(); + if nickname.is_some() { + entry.nickname = nickname; + } + if base_wxid.is_some() { + entry.base_wxid = base_wxid; + } + } else { + self.accounts.insert( + account_id.to_string(), + AccountKey { + account_id: account_id.to_string(), + data_key: String::new(), + extracted_at: Utc::now(), + wechat_version: version.to_string(), + nickname, + base_wxid, + image_aes_key: None, + enc_key: None, + enc_key_salt: None, + enc_keys: entries, + }, + ); + } + } + + /// Resolve the best available key material for an account. + /// + /// Priority: `enc_keys` (new per-DB format) > `enc_key` + `enc_key_salt` (legacy) + /// > `data_key` (raw key). + /// + /// **Note:** `wx-context` overrides this priority — it always prefers `RawKey` + /// when `data_key` is present, because `RawKey` can decrypt any DB regardless of + /// salt, whereas `EncKeys` only cover DBs whose salt was cached in memory at scan + /// time. Direct callers of this method should be aware of this limitation. + pub fn resolve_key_material(&self, account_id: &str) -> Option { + let entry = self.accounts.get(account_id)?; + + // Prefer new per-DB enc_keys format. + if !entry.enc_keys.is_empty() { + let pairs: Vec = entry + .enc_keys + .iter() + .filter_map(|e| { + let key_bytes = hex::decode(&e.enc_key).ok()?; + let salt_bytes = hex::decode(&e.salt).ok()?; + if key_bytes.len() == 32 && salt_bytes.len() == 16 { + let mut key = [0u8; 32]; + let mut salt = [0u8; 16]; + key.copy_from_slice(&key_bytes); + salt.copy_from_slice(&salt_bytes); + Some(EncKeyPair { key, salt }) + } else { + None + } + }) + .collect(); + if !pairs.is_empty() { + return Some(KeyMaterial::EncKeys(pairs)); + } + } + + // Legacy single enc_key path. + if let (Some(ek), Some(es)) = (&entry.enc_key, &entry.enc_key_salt) { + if !ek.is_empty() && !es.is_empty() { + if let (Ok(key_bytes), Ok(salt_bytes)) = (hex::decode(ek), hex::decode(es)) { + if key_bytes.len() == 32 && salt_bytes.len() == 16 { + let mut key = [0u8; 32]; + let mut salt = [0u8; 16]; + key.copy_from_slice(&key_bytes); + salt.copy_from_slice(&salt_bytes); + return Some(KeyMaterial::EncKey { key, salt }); + } + } + } + } + + // Fall back to raw data_key. + if !entry.data_key.is_empty() { + if let Ok(key_bytes) = hex::decode(&entry.data_key) { + if key_bytes.len() == 32 { + let mut key = [0u8; 32]; + key.copy_from_slice(&key_bytes); + return Some(KeyMaterial::RawKey(key)); + } + } + } + + None + } + + /// Get the hex key for an account. + pub fn get(&self, account_id: &str) -> Option<&AccountKey> { + self.accounts.get(account_id) + } +} + + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_old_keys_toml_without_nickname_base_wxid() { + let toml_str = r#" +[accounts.wxid_test_1234] +account_id = "wxid_test_1234" +data_key = "aabbccdd" +extracted_at = "2025-01-01T00:00:00Z" +wechat_version = "4.1.7.31" +"#; + let store: KeyStore = toml::from_str(toml_str).unwrap(); + let key = store.get("wxid_test_1234").unwrap(); + assert_eq!(key.nickname, None); + assert_eq!(key.base_wxid, None); + assert_eq!(key.display_name(), "wxid_test_1234 (昵称未知)"); + } + + #[test] + fn test_keys_toml_with_nickname() { + let toml_str = r#" +[accounts.wxid_test_1234] +account_id = "wxid_test_1234" +data_key = "aabbccdd" +extracted_at = "2025-01-01T00:00:00Z" +wechat_version = "4.1.7.31" +nickname = "TestUser" +base_wxid = "wxid_test" +"#; + let store: KeyStore = toml::from_str(toml_str).unwrap(); + let key = store.get("wxid_test_1234").unwrap(); + assert_eq!(key.nickname.as_deref(), Some("TestUser")); + assert_eq!(key.base_wxid.as_deref(), Some("wxid_test")); + assert_eq!(key.display_name(), "wxid_test_1234 (TestUser)"); + } + + #[test] + fn test_image_aes_key_roundtrip() { + let mut store = KeyStore::default(); + store.set("wxid_abc", "aabbccdd", "4.1.7.31", None, None); + assert_eq!(store.get("wxid_abc").unwrap().image_aes_key, None); + + store.set_image_key("wxid_abc", "6162636465666768696a6b6c6d6e6f70"); + let serialized = toml::to_string_pretty(&store).unwrap(); + let deserialized: KeyStore = toml::from_str(&serialized).unwrap(); + assert_eq!( + deserialized + .get("wxid_abc") + .unwrap() + .image_aes_key + .as_deref(), + Some("6162636465666768696a6b6c6d6e6f70"), + ); + // data_key preserved + assert_eq!(deserialized.get("wxid_abc").unwrap().data_key, "aabbccdd"); + } + + #[test] + fn test_set_preserves_existing_image_key() { + let mut store = KeyStore::default(); + store.set("wxid_abc", "aabbccdd", "4.1.7.31", None, None); + store.set_image_key("wxid_abc", "deadbeef12345678"); + + // Re-setting data key should preserve existing image key + store.set( + "wxid_abc", + "newdatakey", + "4.1.7.33", + Some("Nick".into()), + None, + ); + assert_eq!( + store.get("wxid_abc").unwrap().image_aes_key.as_deref(), + Some("deadbeef12345678"), + ); + assert_eq!(store.get("wxid_abc").unwrap().data_key, "newdatakey"); + } + + #[test] + fn test_old_toml_without_image_key() { + // Backward compatibility: old TOML without image_aes_key field + let toml_str = r#" +[accounts.wxid_old] +account_id = "wxid_old" +data_key = "oldkey" +extracted_at = "2025-01-01T00:00:00Z" +wechat_version = "4.1.7.31" +"#; + let store: KeyStore = toml::from_str(toml_str).unwrap(); + assert_eq!(store.get("wxid_old").unwrap().image_aes_key, None); + } + + #[test] + fn test_old_toml_without_enc_key_fields() { + let toml_str = r#" +[accounts.wxid_old] +account_id = "wxid_old" +data_key = "aabbccddaabbccddaabbccddaabbccddaabbccddaabbccddaabbccddaabbccdd" +extracted_at = "2025-01-01T00:00:00Z" +wechat_version = "4.1.7.31" +"#; + let store: KeyStore = toml::from_str(toml_str).unwrap(); + let key = store.get("wxid_old").unwrap(); + assert_eq!(key.enc_key, None); + assert_eq!(key.enc_key_salt, None); + } + + #[test] + fn test_set_clears_enc_key_when_data_key_changes() { + let mut store = KeyStore::default(); + store.set_enc_key( + "wxid_abc", + "aa".repeat(32).as_str(), + "bb".repeat(16).as_str(), + "4.1.7.31", + None, + None, + ); + + // data_key was empty; setting a new data_key should clear derived fields + store.set( + "wxid_abc", + "cc".repeat(32).as_str(), + "4.1.8.0", + Some("Nick".into()), + None, + ); + let entry = store.get("wxid_abc").unwrap(); + assert_eq!(entry.enc_key, None, "enc_key cleared on data_key change"); + assert_eq!( + entry.enc_key_salt, None, + "enc_key_salt cleared on data_key change" + ); + assert_eq!(entry.data_key, "cc".repeat(32)); + } + + #[test] + fn test_set_enc_key_preserves_data_key_and_image_key() { + let mut store = KeyStore::default(); + let data_key = "dd".repeat(32); + store.set("wxid_abc", &data_key, "4.1.7.31", None, None); + store.set_image_key("wxid_abc", "deadbeef12345678"); + + store.set_enc_key( + "wxid_abc", + "ee".repeat(32).as_str(), + "ff".repeat(16).as_str(), + "4.1.8.0", + Some("Nick".into()), + None, + ); + + let entry = store.get("wxid_abc").unwrap(); + assert_eq!(entry.data_key, data_key); + assert_eq!(entry.image_aes_key.as_deref(), Some("deadbeef12345678")); + assert_eq!(entry.enc_key.as_deref(), Some(&*"ee".repeat(32))); + assert_eq!(entry.enc_key_salt.as_deref(), Some(&*"ff".repeat(16))); + } + + #[test] + fn test_resolve_key_material_prefers_enc_key() { + let mut store = KeyStore::default(); + let data_key = "ab".repeat(32); + let enc_key = "cd".repeat(32); + let enc_salt = "ef".repeat(16); + + store.set("wxid_abc", &data_key, "4.1.7.31", None, None); + store.set_enc_key("wxid_abc", &enc_key, &enc_salt, "4.1.7.31", None, None); + + let km = store.resolve_key_material("wxid_abc").unwrap(); + match km { + KeyMaterial::EncKey { key, salt } => { + assert_eq!(hex::encode(key), enc_key); + assert_eq!(hex::encode(salt), enc_salt); + } + _ => panic!("expected EncKey variant"), + } + } + + #[test] + fn test_resolve_key_material_falls_back_to_raw_key() { + let mut store = KeyStore::default(); + let data_key = "ab".repeat(32); + store.set("wxid_abc", &data_key, "4.1.7.31", None, None); + + let km = store.resolve_key_material("wxid_abc").unwrap(); + match km { + KeyMaterial::RawKey(key) => { + assert_eq!(hex::encode(key), data_key); + } + _ => panic!("expected RawKey variant"), + } + } + + #[test] + fn test_resolve_key_material_none_when_no_keys() { + let mut store = KeyStore::default(); + // Create entry with empty data_key and no enc_key + store.set_image_key("wxid_abc", "deadbeef12345678"); + assert!(store.resolve_key_material("wxid_abc").is_none()); + } + + #[test] + fn test_resolve_key_material_nonexistent_account() { + let store = KeyStore::default(); + assert!(store.resolve_key_material("wxid_nope").is_none()); + } + + #[test] + fn test_enc_key_roundtrip_serialization() { + let mut store = KeyStore::default(); + let enc_key = "ab".repeat(32); + let enc_salt = "cd".repeat(16); + store.set_enc_key( + "wxid_abc", + &enc_key, + &enc_salt, + "4.1.7.31", + Some("Test".into()), + None, + ); + + let serialized = toml::to_string_pretty(&store).unwrap(); + let deserialized: KeyStore = toml::from_str(&serialized).unwrap(); + + let entry = deserialized.get("wxid_abc").unwrap(); + assert_eq!(entry.enc_key.as_deref(), Some(&*enc_key)); + assert_eq!(entry.enc_key_salt.as_deref(), Some(&*enc_salt)); + assert_eq!(entry.nickname.as_deref(), Some("Test")); + + // resolve_key_material should work on deserialized store + let km = deserialized.resolve_key_material("wxid_abc").unwrap(); + assert!(matches!(km, KeyMaterial::EncKey { .. })); + } + + #[test] + fn test_set_enc_keys_dedup_and_sort() { + let mut store = KeyStore::default(); + let pairs = vec![ + EncKeyPair { + key: [0xBBu8; 32], + salt: [0x02u8; 16], + }, + EncKeyPair { + key: [0xAAu8; 32], + salt: [0x01u8; 16], + }, + EncKeyPair { + key: [0xAAu8; 32], + salt: [0x01u8; 16], + }, // duplicate + ]; + store.set_enc_keys("wxid_test", &pairs, "4.1.8.0", None, None); + + let entry = store.get("wxid_test").unwrap(); + assert_eq!(entry.enc_keys.len(), 2, "duplicates should be removed"); + assert_eq!( + entry.enc_keys[0].salt, + hex::encode([0x01u8; 16]), + "sorted by salt" + ); + assert_eq!(entry.enc_keys[1].salt, hex::encode([0x02u8; 16])); + assert!(entry.enc_key.is_none(), "legacy enc_key cleared"); + assert!(entry.enc_key_salt.is_none(), "legacy enc_key_salt cleared"); + } + + #[test] + fn test_set_enc_keys_preserves_data_key_and_image_key() { + let mut store = KeyStore::default(); + let data_key = "dd".repeat(32); + store.set("wxid_test", &data_key, "4.1.7.31", None, None); + store.set_image_key("wxid_test", "deadbeef12345678"); + + let pairs = vec![EncKeyPair { + key: [0xAAu8; 32], + salt: [0x01u8; 16], + }]; + store.set_enc_keys("wxid_test", &pairs, "4.1.8.0", Some("Nick".into()), None); + + let entry = store.get("wxid_test").unwrap(); + assert_eq!(entry.data_key, data_key); + assert_eq!(entry.image_aes_key.as_deref(), Some("deadbeef12345678")); + assert_eq!(entry.enc_keys.len(), 1); + } + + #[test] + fn test_resolve_key_material_enc_keys_preferred() { + let mut store = KeyStore::default(); + let data_key = "ab".repeat(32); + store.set("wxid_test", &data_key, "4.1.7.31", None, None); + + let pairs = vec![ + EncKeyPair { + key: [0xAAu8; 32], + salt: [0x01u8; 16], + }, + EncKeyPair { + key: [0xBBu8; 32], + salt: [0x02u8; 16], + }, + ]; + store.set_enc_keys("wxid_test", &pairs, "4.1.8.0", None, None); + + let km = store.resolve_key_material("wxid_test").unwrap(); + match km { + KeyMaterial::EncKeys(ref p) => { + assert_eq!(p.len(), 2); + assert_eq!(p[0].salt, [0x01u8; 16]); + assert_eq!(p[1].salt, [0x02u8; 16]); + } + _ => panic!("expected EncKeys variant"), + } + } + + #[test] + fn test_resolve_key_material_legacy_enc_key_still_works() { + let mut store = KeyStore::default(); + let enc_key = "cd".repeat(32); + let enc_salt = "ef".repeat(16); + store.set_enc_key("wxid_test", &enc_key, &enc_salt, "4.1.7.31", None, None); + + let km = store.resolve_key_material("wxid_test").unwrap(); + assert!(matches!(km, KeyMaterial::EncKey { .. })); + } + + #[test] + fn test_enc_keys_roundtrip_serialization() { + let mut store = KeyStore::default(); + let pairs = vec![ + EncKeyPair { + key: [0xAAu8; 32], + salt: [0x01u8; 16], + }, + EncKeyPair { + key: [0xBBu8; 32], + salt: [0x02u8; 16], + }, + ]; + store.set_enc_keys("wxid_test", &pairs, "4.1.8.0", Some("Test".into()), None); + + let serialized = toml::to_string_pretty(&store).unwrap(); + let deserialized: KeyStore = toml::from_str(&serialized).unwrap(); + + let km = deserialized.resolve_key_material("wxid_test").unwrap(); + match km { + KeyMaterial::EncKeys(ref p) => assert_eq!(p.len(), 2), + _ => panic!("expected EncKeys variant after roundtrip"), + } + } + + #[test] + fn set_with_different_data_key_clears_enc_keys() { + let mut store = KeyStore::default(); + let original_key = "aa".repeat(32); + store.set("wxid_abc", &original_key, "4.1.7.31", None, None); + + // Populate derived fields via set_enc_keys + let pairs = vec![EncKeyPair { + key: [0xCCu8; 32], + salt: [0x01u8; 16], + }]; + store.set_enc_keys("wxid_abc", &pairs, "4.1.7.31", None, None); + // Also set legacy enc_key manually + { + let entry = store.accounts.get_mut("wxid_abc").unwrap(); + entry.enc_key = Some("dd".repeat(32)); + entry.enc_key_salt = Some("ee".repeat(16)); + } + + // Now set() with a DIFFERENT data_key + let new_key = "ff".repeat(32); + store.set("wxid_abc", &new_key, "4.1.8.0", None, None); + + let entry = store.get("wxid_abc").unwrap(); + assert_eq!(entry.data_key, new_key); + assert_eq!(entry.enc_key, None, "legacy enc_key should be cleared"); + assert_eq!( + entry.enc_key_salt, None, + "legacy enc_key_salt should be cleared" + ); + assert!(entry.enc_keys.is_empty(), "enc_keys should be cleared"); + } + + #[test] + fn set_with_same_data_key_preserves_enc_keys() { + let mut store = KeyStore::default(); + let data_key = "aa".repeat(32); + store.set("wxid_abc", &data_key, "4.1.7.31", None, None); + + // Populate derived fields + let pairs = vec![EncKeyPair { + key: [0xCCu8; 32], + salt: [0x01u8; 16], + }]; + store.set_enc_keys("wxid_abc", &pairs, "4.1.7.31", None, None); + // Also set legacy enc_key manually + { + let entry = store.accounts.get_mut("wxid_abc").unwrap(); + entry.enc_key = Some("dd".repeat(32)); + entry.enc_key_salt = Some("ee".repeat(16)); + } + + // Now set() with the SAME data_key + store.set("wxid_abc", &data_key, "4.1.8.0", Some("Nick".into()), None); + + let entry = store.get("wxid_abc").unwrap(); + assert_eq!(entry.data_key, data_key); + assert_eq!( + entry.enc_key.as_deref(), + Some(&*"dd".repeat(32)), + "legacy enc_key should be preserved" + ); + assert_eq!( + entry.enc_key_salt.as_deref(), + Some(&*"ee".repeat(16)), + "legacy enc_key_salt should be preserved" + ); + assert_eq!(entry.enc_keys.len(), 1, "enc_keys should be preserved"); + assert_eq!(entry.enc_keys[0].enc_key, hex::encode([0xCCu8; 32])); + } + + #[test] + fn test_set_enc_keys_overwrites_stale_base_wxid() { + let mut store = KeyStore::default(); + // Initial entry with stale base_wxid + store.set( + "testuser001_1662", + &"aa".repeat(32), + "4.1.7.31", + None, + Some("testuser001_1662".into()), // stale: same as directory name + ); + + let entry = store.get("testuser001_1662").unwrap(); + assert_eq!(entry.base_wxid.as_deref(), Some("testuser001_1662")); + + // Writeback with canonical base_wxid + let pairs = vec![EncKeyPair { + key: [0xAAu8; 32], + salt: [0x01u8; 16], + }]; + store.set_enc_keys( + "testuser001_1662", + &pairs, + "4.1.8.0", + None, + Some("testuser001".into()), // canonical value + ); + + let entry = store.get("testuser001_1662").unwrap(); + assert_eq!( + entry.base_wxid.as_deref(), + Some("testuser001"), + "stale base_wxid should be overwritten with canonical value" + ); + } +} diff --git a/crates/wx-media/Cargo.toml b/crates/wx-media/Cargo.toml new file mode 100644 index 0000000..bbfdd39 --- /dev/null +++ b/crates/wx-media/Cargo.toml @@ -0,0 +1,24 @@ +[package] +name = "wx-media" +version.workspace = true +edition.workspace = true + +[features] +default = [] +audio = ["dep:silk-rs"] + +[dependencies] +aes = "0.8" +ecb = "0.1" +cipher = { version = "0.4", features = ["block-padding"] } +rusqlite = { version = "0.32", features = ["bundled"] } +md5 = "0.7" +base64 = "0.22" +thiserror = "2" +serde = { version = "1", features = ["derive"] } +hex = "0.4" +silk-rs = { version = "0.2", optional = true } +wx-keychain = { path = "../wx-keychain" } + +[dev-dependencies] +tempfile = "3" diff --git a/crates/wx-media/src/audio_transcode.rs b/crates/wx-media/src/audio_transcode.rs new file mode 100644 index 0000000..c67289c --- /dev/null +++ b/crates/wx-media/src/audio_transcode.rs @@ -0,0 +1,125 @@ +use crate::error::MediaError; +use crate::types::TranscodeAudioResult; + +/// Strip the WeChat `\x02` prefix from SILK data if present. +/// WeChat prepends a `\x02` byte before the standard `#!SILK_V3` header. +#[cfg(feature = "audio")] +fn strip_wechat_silk_prefix(data: &[u8]) -> &[u8] { + if data.first() == Some(&0x02) { + &data[1..] + } else { + data + } +} + +#[cfg(feature = "audio")] +fn decode_silk_to_pcm(data: &[u8]) -> Result, MediaError> { + let silk_data = strip_wechat_silk_prefix(data); + silk_rs::decode_silk(silk_data, 24000).map_err(|e| MediaError::SilkDecodeFailed { + reason: e.to_string(), + }) +} + +#[cfg(feature = "audio")] +pub fn transcode_silk_to_ogg_opus(data: &[u8]) -> Result { + if !crate::ffmpeg::ffmpeg_available() { + return Err(MediaError::FfmpegNotFound); + } + + let pcm = decode_silk_to_pcm(data)?; + let ogg = crate::ffmpeg::run_ffmpeg( + &pcm, + &[ + "-hide_banner", + "-loglevel", + "error", + "-f", + "s16le", + "-ar", + "24000", + "-ac", + "1", + "-i", + "pipe:0", + "-c:a", + "libopus", + "-b:a", + "32k", + "-f", + "ogg", + "pipe:1", + ], + )?; + + Ok(TranscodeAudioResult { + data: ogg, + ext: "ogg", + mime: "audio/ogg", + transcoded: true, + }) +} + +/// Transcode SILK audio to MP3. +/// +/// Requires the `audio` cargo feature (for SILK decoding) and ffmpeg (for MP3 encoding). +/// +/// Returns: +/// - `AudioFeatureDisabled` error if the `audio` feature is not enabled. +/// - `TranscodeAudioResult { transcoded: false, ext: "silk" }` if ffmpeg is not available +/// (returns original SILK data unchanged). +/// - `TranscodeAudioResult { transcoded: true, ext: "mp3" }` on full success. +#[cfg(feature = "audio")] +pub fn transcode_silk_to_mp3(data: &[u8]) -> Result { + if !crate::ffmpeg::ffmpeg_available() { + // ffmpeg not available → return original SILK data unchanged + return Ok(TranscodeAudioResult { + data: data.to_vec(), + ext: "silk", + mime: "audio/x-silk", + transcoded: false, + }); + } + + let pcm = decode_silk_to_pcm(data)?; + + // PCM → MP3 via ffmpeg + let mp3 = crate::ffmpeg::run_ffmpeg( + &pcm, + &[ + "-hide_banner", + "-loglevel", + "error", + "-f", + "s16le", + "-ar", + "24000", + "-ac", + "1", + "-i", + "pipe:0", + "-b:a", + "64k", + "-f", + "mp3", + "pipe:1", + ], + )?; + + Ok(TranscodeAudioResult { + data: mp3, + ext: "mp3", + mime: "audio/mpeg", + transcoded: true, + }) +} + +/// Stub when `audio` feature is not enabled. +#[cfg(not(feature = "audio"))] +pub fn transcode_silk_to_mp3(_data: &[u8]) -> Result { + Err(MediaError::AudioFeatureDisabled) +} + +#[cfg(not(feature = "audio"))] +pub fn transcode_silk_to_ogg_opus(_data: &[u8]) -> Result { + Err(MediaError::AudioFeatureDisabled) +} diff --git a/crates/wx-media/src/dat.rs b/crates/wx-media/src/dat.rs new file mode 100644 index 0000000..c81d4ee --- /dev/null +++ b/crates/wx-media/src/dat.rs @@ -0,0 +1,417 @@ +use std::collections::HashMap; +use std::path::Path; + +use crate::error::MediaError; +use crate::types::{DatDecryptOptions, DatFormat, DecodedImage, ImageType}; + +/// V1 signature: `07 08 V1 08 07` +const V1_MAGIC: &[u8; 6] = b"\x07\x08V1\x08\x07"; +/// V2 signature: `07 08 V2 08 07` +const V2_MAGIC: &[u8; 6] = b"\x07\x08V2\x08\x07"; +/// V1 fixed AES key: `cfcd208495d565ef` (md5("0")[:16]) +const V1_FIXED_KEY: &[u8; 16] = b"cfcd208495d565ef"; +/// Dat file header size: 6B signature + 4B aes_size + 4B xor_size + 1B padding = 15 +const HEADER_SIZE: usize = 15; + +/// Known image magic bytes for XOR key detection (ordered by header length, descending). +const IMAGE_MAGICS: &[(&[u8], ImageType)] = &[ + (&[0x77, 0x78, 0x67, 0x66], ImageType::Wxgf), + (&[0x89, 0x50, 0x4E, 0x47], ImageType::Png), + (&[0x47, 0x49, 0x46, 0x38], ImageType::Gif), + (&[0x49, 0x49, 0x2A, 0x00], ImageType::Tif), + (&[0x52, 0x49, 0x46, 0x46], ImageType::Webp), // RIFF header + (&[0xFF, 0xD8, 0xFF], ImageType::Jpg), +]; + +/// Detect the dat encryption format from the first 6 bytes. +/// Returns `None` if it's an XOR-only file (no V1/V2 signature). +pub fn detect_dat_format(data: &[u8]) -> Option { + if data.len() < 6 { + return None; + } + if &data[..6] == V1_MAGIC { + Some(DatFormat::V1) + } else if &data[..6] == V2_MAGIC { + Some(DatFormat::V2) + } else { + None + } +} + +/// Detect image type from decrypted data header. +pub fn detect_image_type(data: &[u8]) -> ImageType { + if data.len() >= 4 && data[..4] == [0x77, 0x78, 0x67, 0x66] { + return ImageType::Wxgf; + } + // WEBP: RIFF....WEBP + if data.len() >= 12 && data[..4] == [0x52, 0x49, 0x46, 0x46] && data[8..12] == *b"WEBP" { + return ImageType::Webp; + } + for &(magic, img_type) in IMAGE_MAGICS { + if img_type == ImageType::Webp { + continue; // handled above with full check + } + if data.len() >= magic.len() && data[..magic.len()] == *magic { + return img_type; + } + } + // BMP: 2-byte magic 42 4D — only check if nothing else matched + if data.len() >= 2 && data[..2] == [0x42, 0x4D] { + return ImageType::Bmp; + } + ImageType::Unknown +} + +/// Decrypt a `.dat` file (in-memory). Unified entry for XOR, V1, V2 formats. +pub fn decrypt_dat(data: &[u8], opts: &DatDecryptOptions) -> Result { + if data.len() < 3 { + return Err(MediaError::InvalidFormat { + reason: format!("data too short: {} bytes", data.len()), + }); + } + + match detect_dat_format(data) { + Some(DatFormat::V1) => decrypt_v1_v2(data, V1_FIXED_KEY, opts.xor_key, DatFormat::V1), + Some(DatFormat::V2) => { + let key = opts.v2_aes_key.as_ref().ok_or(MediaError::MissingV2Key)?; + decrypt_v1_v2(data, key, opts.xor_key, DatFormat::V2) + } + None => decrypt_xor(data), + _ => unreachable!(), + } +} + +/// XOR decryption: detect single-byte key via known image magic. +fn decrypt_xor(data: &[u8]) -> Result { + for &(magic, _) in IMAGE_MAGICS { + if data.len() < magic.len() { + continue; + } + let key = data[0] ^ magic[0]; + let matched = magic.iter().enumerate().all(|(i, &m)| data[i] ^ key == m); + if matched { + let decrypted: Vec = data.iter().map(|b| b ^ key).collect(); + let img_type = detect_image_type(&decrypted); + return Ok(DecodedImage { + data: decrypted, + format: DatFormat::Xor, + ext: img_type.ext().to_string(), + }); + } + } + Err(MediaError::XorKeyDetectionFailed) +} + +/// V1/V2 decryption: AES-128-ECB header + optional raw middle + optional XOR tail. +fn decrypt_v1_v2( + data: &[u8], + aes_key: &[u8; 16], + xor_key: Option, + format: DatFormat, +) -> Result { + if data.len() < HEADER_SIZE { + return Err(MediaError::InvalidFormat { + reason: format!( + "V1/V2 header requires {} bytes, got {}", + HEADER_SIZE, + data.len() + ), + }); + } + + let aes_size = u32::from_le_bytes(data[6..10].try_into().unwrap()) as usize; + let xor_size = u32::from_le_bytes(data[10..14].try_into().unwrap()) as usize; + + // AES alignment: round up to next multiple of 16 (+ 16 for PKCS7 padding block) + let aligned_aes_size = (aes_size / 16 + 1) * 16; + + let payload = &data[HEADER_SIZE..]; + + if aligned_aes_size > payload.len() { + return Err(MediaError::InvalidFormat { + reason: format!( + "AES section ({} aligned) exceeds payload ({})", + aligned_aes_size, + payload.len() + ), + }); + } + + // Decrypt AES-ECB section + let aes_ciphertext = &payload[..aligned_aes_size]; + let aes_plaintext = aes_ecb_decrypt(aes_ciphertext, aes_key)?; + // Truncate to original aes_size (remove PKCS7 padding overshoot) + let aes_out = if aes_plaintext.len() > aes_size { + &aes_plaintext[..aes_size] + } else { + &aes_plaintext + }; + + // Raw middle section + let raw_start = aligned_aes_size; + let raw_end = payload.len().saturating_sub(xor_size); + let raw_data = if raw_start < raw_end { + &payload[raw_start..raw_end] + } else { + &[] + }; + + // XOR tail section + let xor_data = &payload[payload.len().saturating_sub(xor_size)..]; + let xor_decrypted: Vec = if xor_size > 0 { + let k = xor_key.unwrap_or(0x37); // default xor key + xor_data.iter().map(|b| b ^ k).collect() + } else { + Vec::new() + }; + + let mut result = Vec::with_capacity(aes_out.len() + raw_data.len() + xor_decrypted.len()); + result.extend_from_slice(aes_out); + result.extend_from_slice(raw_data); + result.extend_from_slice(&xor_decrypted); + + let img_type = detect_image_type(&result); + + Ok(DecodedImage { + data: result, + format, + ext: img_type.ext().to_string(), + }) +} + +/// AES-128-ECB decrypt with PKCS7 unpadding. +fn aes_ecb_decrypt(ciphertext: &[u8], key: &[u8; 16]) -> Result, MediaError> { + use aes::cipher::{BlockDecrypt, KeyInit}; + use aes::Aes128; + + if ciphertext.is_empty() { + return Ok(Vec::new()); + } + if !ciphertext.len().is_multiple_of(16) { + return Err(MediaError::AesDecryptFailed { + reason: format!( + "ciphertext length {} is not a multiple of 16", + ciphertext.len() + ), + }); + } + + let cipher = Aes128::new(key.into()); + let mut decrypted = ciphertext.to_vec(); + + for chunk in decrypted.chunks_exact_mut(16) { + cipher.decrypt_block(chunk.into()); + } + + // PKCS7 unpadding + let padding = decrypted[decrypted.len() - 1] as usize; + if padding == 0 || padding > 16 { + return Err(MediaError::AesDecryptFailed { + reason: format!( + "invalid PKCS7 padding byte: {}", + decrypted[decrypted.len() - 1] + ), + }); + } + let valid = decrypted[decrypted.len() - padding..] + .iter() + .all(|&b| b as usize == padding); + if !valid { + return Err(MediaError::AesDecryptFailed { + reason: "PKCS7 padding validation failed (wrong key?)".into(), + }); + } + decrypted.truncate(decrypted.len() - padding); + + Ok(decrypted) +} + +/// Auto-detect the XOR key by scanning `_t.dat` thumbnail files in `attach_dir`. +/// +/// JPEG files end with `FF D9`. In XOR-encrypted `.dat` files, the last 2 bytes +/// are `(0xFF ^ xor_key)` and `(0xD9 ^ xor_key)`. We derive the key from multiple +/// thumbnails and use majority voting for robustness. +/// +/// Returns `None` if no thumbnails are found or votes are inconsistent. +pub fn detect_xor_key(attach_dir: &Path) -> Option { + let mut votes: HashMap = HashMap::new(); + + for entry in walkdir_thumbnails(attach_dir) { + let path = entry.path(); + let data = match std::fs::read(&path) { + Ok(d) if d.len() >= 2 => d, + _ => continue, + }; + + // For V1/V2 files, extract XOR key from the XOR tail section + if data.len() >= HEADER_SIZE + 2 && (&data[..6] == V1_MAGIC || &data[..6] == V2_MAGIC) { + let xor_size = u32::from_le_bytes(data[10..14].try_into().unwrap()) as usize; + if xor_size >= 2 { + // Last 2 bytes of file = last 2 bytes of XOR tail + let tail_penultimate = data[data.len() - 2]; + let tail_last = data[data.len() - 1]; + let candidate = tail_penultimate ^ 0xFF; + if tail_last ^ 0xD9 == candidate { + *votes.entry(candidate).or_insert(0) += 1; + } + } + continue; + } + + let tail_penultimate = data[data.len() - 2]; + let tail_last = data[data.len() - 1]; + + // Derive XOR key from penultimate byte (should be 0xFF ^ key) + let candidate = tail_penultimate ^ 0xFF; + // Verify with last byte (should be 0xD9 ^ key) + if tail_last ^ 0xD9 == candidate { + *votes.entry(candidate).or_insert(0) += 1; + } + } + + // Return the most-voted key + votes + .into_iter() + .max_by_key(|&(_, count)| count) + .map(|(key, _)| key) +} + +/// Walk directory recursively collecting `_t.dat` thumbnail files. +fn walkdir_thumbnails(dir: &Path) -> Vec { + let mut results = Vec::new(); + if let Ok(entries) = std::fs::read_dir(dir) { + for entry in entries.flatten() { + if let Ok(ft) = entry.file_type() { + if ft.is_dir() { + results.extend(walkdir_thumbnails(&entry.path())); + } else if ft.is_file() { + let name = entry.file_name().to_string_lossy().to_string(); + if name.ends_with("_t.dat") { + results.push(entry); + } + } + } + } + } + results +} + +#[cfg(test)] +mod tests { + use super::*; + use std::fs; + + #[test] + fn test_detect_xor_key_from_thumbnails() { + let dir = std::env::temp_dir().join("wechat_xor_test"); + let _ = fs::remove_dir_all(&dir); + fs::create_dir_all(&dir).unwrap(); + + // XOR key = 0xa5 + // JPEG ends with FF D9 → encrypted tail = (0xFF ^ 0xa5, 0xD9 ^ 0xa5) = (0x5a, 0x7c) + let xor_key: u8 = 0xa5; + let tail = [0xFF ^ xor_key, 0xD9 ^ xor_key]; + + // Create 3 fake _t.dat thumbnails with consistent tail + for i in 0..3 { + let mut data = vec![0x00; 100]; + data.extend_from_slice(&tail); + fs::write(dir.join(format!("img{i}_t.dat")), &data).unwrap(); + } + + let result = detect_xor_key(&dir); + assert_eq!(result, Some(0xa5)); + + let _ = fs::remove_dir_all(&dir); + } + + #[test] + fn test_detect_xor_key_empty_dir() { + let dir = std::env::temp_dir().join("wechat_xor_test_empty"); + let _ = fs::remove_dir_all(&dir); + fs::create_dir_all(&dir).unwrap(); + + assert_eq!(detect_xor_key(&dir), None); + + let _ = fs::remove_dir_all(&dir); + } + + #[test] + fn test_detect_xor_key_from_v2_thumbnails() { + let dir = std::env::temp_dir().join("wechat_xor_test_v2"); + let _ = fs::remove_dir_all(&dir); + fs::create_dir_all(&dir).unwrap(); + + // V2 thumbnail with xor_size=50, xor_key=0xa5 + let xor_key: u8 = 0xa5; + let xor_size: u32 = 50; + let mut data = Vec::new(); + data.extend_from_slice(V2_MAGIC); // 6 bytes + data.extend_from_slice(&1024u32.to_le_bytes()); // aes_size + data.extend_from_slice(&xor_size.to_le_bytes()); // xor_size + data.push(0x00); // padding → 15 bytes header + // Payload: some AES data + raw middle + XOR tail + data.extend_from_slice(&[0x00; 100]); // filler + // XOR tail ends with JPEG EOI encrypted + let tail_len = xor_size as usize; + let filler_tail = vec![0x00; tail_len - 2]; + data.extend_from_slice(&filler_tail); + data.push(0xFF ^ xor_key); // penultimate + data.push(0xD9 ^ xor_key); // last + + fs::write(dir.join("img0_t.dat"), &data).unwrap(); + + assert_eq!(detect_xor_key(&dir), Some(0xa5)); + + let _ = fs::remove_dir_all(&dir); + } + + #[test] + fn test_detect_xor_key_v2_no_xor_section() { + let dir = std::env::temp_dir().join("wechat_xor_test_v2_noxor"); + let _ = fs::remove_dir_all(&dir); + fs::create_dir_all(&dir).unwrap(); + + // V2 thumbnail with xor_size=0 (no XOR tail — cannot derive key) + let mut data = Vec::new(); + data.extend_from_slice(V2_MAGIC); + data.extend_from_slice(&1024u32.to_le_bytes()); // aes_size + data.extend_from_slice(&0u32.to_le_bytes()); // xor_size = 0 + data.push(0x00); + data.extend_from_slice(&[0x00; 100]); + fs::write(dir.join("img0_t.dat"), &data).unwrap(); + + assert_eq!(detect_xor_key(&dir), None); + + let _ = fs::remove_dir_all(&dir); + } + + #[test] + fn test_detect_xor_key_majority_voting() { + let dir = std::env::temp_dir().join("wechat_xor_test_vote"); + let _ = fs::remove_dir_all(&dir); + fs::create_dir_all(&dir).unwrap(); + + let xor_key: u8 = 0xa5; + let good_tail = [0xFF ^ xor_key, 0xD9 ^ xor_key]; + + // 3 files with key 0xa5 + for i in 0..3 { + let mut data = vec![0x00; 50]; + data.extend_from_slice(&good_tail); + fs::write(dir.join(format!("good{i}_t.dat")), &data).unwrap(); + } + + // 1 file with different key 0x37 + let other_key: u8 = 0x37; + let bad_tail = [0xFF ^ other_key, 0xD9 ^ other_key]; + let mut data = vec![0x00; 50]; + data.extend_from_slice(&bad_tail); + fs::write(dir.join("bad0_t.dat"), &data).unwrap(); + + // Majority should win + assert_eq!(detect_xor_key(&dir), Some(0xa5)); + + let _ = fs::remove_dir_all(&dir); + } +} diff --git a/crates/wx-media/src/error.rs b/crates/wx-media/src/error.rs new file mode 100644 index 0000000..7a52095 --- /dev/null +++ b/crates/wx-media/src/error.rs @@ -0,0 +1,61 @@ +use std::path::PathBuf; + +#[derive(Debug, thiserror::Error)] +pub enum MediaError { + #[error("I/O error: {0}")] + Io(#[from] std::io::Error), + + #[error("SQLite error: {0}")] + Sqlite(#[from] rusqlite::Error), + + #[error("invalid dat format: {reason}")] + InvalidFormat { reason: String }, + + #[error("V2 AES key required but not provided")] + MissingV2Key, + + #[error("AES decryption failed: {reason}")] + AesDecryptFailed { reason: String }, + + #[error("XOR key detection failed: no known image magic matched")] + XorKeyDetectionFailed, + + #[error("resource not found: {0}")] + NotFound(String), + + #[error("media lookup miss: {0}")] + LookupMiss(String), + + #[error("media schema missing: {0}")] + SchemaMissing(String), + + #[error("packed_info parse failed for local_id {local_id}: {reason}")] + PackedInfoParseFailed { local_id: i64, reason: String }, + + #[error("no dat files found for md5 {md5} under {path}")] + NoDatFiles { md5: String, path: PathBuf }, + + #[error("no media databases found in {0}")] + NoMediaDbs(PathBuf), + + #[error("invalid or unsupported WXGF container")] + InvalidWxgf, + + #[error("ffmpeg not found")] + FfmpegNotFound, + + #[error("ffmpeg failed (exit {status}): {stderr}")] + FfmpegFailed { status: i32, stderr: String }, + + #[error("SILK decode failed: {reason}")] + SilkDecodeFailed { reason: String }, + + #[error("audio feature not enabled (build with --features audio)")] + AudioFeatureDisabled, +} + +impl MediaError { + pub fn ffmpeg_install_hint() -> &'static str { + "install ffmpeg and ensure it is on PATH, or set FFMPEG_PATH to the ffmpeg binary" + } +} diff --git a/crates/wx-media/src/fallback.rs b/crates/wx-media/src/fallback.rs new file mode 100644 index 0000000..b4ebd64 --- /dev/null +++ b/crates/wx-media/src/fallback.rs @@ -0,0 +1,89 @@ +use std::path::{Path, PathBuf}; + +/// Find a video file by MD5 hash, scanning `video_dir/{YYYY-MM}/{md5}.mp4`. +/// +/// Strategy: try `month_hint` first (fast path), then scan all YYYY-MM subdirectories. +pub fn find_video_by_md5(video_dir: &Path, md5: &str, month_hint: &str) -> Option { + let target = format!("{md5}.mp4"); + + // Fast path: check hint month first + let hint_path = video_dir.join(month_hint).join(&target); + if hint_path.is_file() { + return Some(hint_path); + } + + // Scan all YYYY-MM subdirectories (skip hint month, already tried) + let entries = std::fs::read_dir(video_dir).ok()?; + for entry in entries.flatten() { + let name = entry.file_name(); + let name_str = name.to_string_lossy(); + if !is_month_dir(&name_str) || name_str.as_ref() == month_hint { + continue; + } + let candidate = entry.path().join(&target); + if candidate.is_file() { + return Some(candidate); + } + } + + None +} + +/// Find a file by title in `file_dir/{month_hint}/{title}`. +/// +/// Only checks the specified month (no broad scan to avoid same-name ambiguity). +/// Sanitizes `title` to prevent path traversal. +pub fn find_file_by_name(file_dir: &Path, title: &str, month_hint: &str) -> Option { + // Sanitize: extract just the filename component to prevent path traversal + let safe_basename = Path::new(title).file_name()?; + let candidate = file_dir.join(month_hint).join(safe_basename); + if candidate.is_file() { + Some(candidate) + } else { + None + } +} + +/// Check if a directory name matches the `YYYY-MM` pattern. +fn is_month_dir(name: &str) -> bool { + if name.len() != 7 { + return false; + } + let bytes = name.as_bytes(); + // YYYY must be digits + if !bytes[..4].iter().all(|b| b.is_ascii_digit()) { + return false; + } + // Separator must be '-' + if bytes[4] != b'-' { + return false; + } + // MM must be 01-12 + let month: u8 = match name[5..7].parse() { + Ok(m) => m, + Err(_) => return false, + }; + (1..=12).contains(&month) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_is_month_dir_valid() { + assert!(is_month_dir("2024-01")); + assert!(is_month_dir("2024-12")); + assert!(is_month_dir("1999-06")); + } + + #[test] + fn test_is_month_dir_invalid() { + assert!(!is_month_dir("2024-00")); + assert!(!is_month_dir("2024-13")); + assert!(!is_month_dir("other")); + assert!(!is_month_dir("24-01")); + assert!(!is_month_dir("2024/01")); + assert!(!is_month_dir("")); + } +} diff --git a/crates/wx-media/src/ffmpeg.rs b/crates/wx-media/src/ffmpeg.rs new file mode 100644 index 0000000..f2ea33c --- /dev/null +++ b/crates/wx-media/src/ffmpeg.rs @@ -0,0 +1,135 @@ +use std::io::Write; +use std::process::{Command, Output, Stdio}; +use std::sync::atomic::{AtomicBool, Ordering}; + +use crate::error::MediaError; + +fn ffmpeg_bin() -> String { + std::env::var("FFMPEG_PATH").unwrap_or_else(|_| "ffmpeg".to_string()) +} + +fn ffprobe_bin() -> String { + std::env::var("FFPROBE_PATH").unwrap_or_else(|_| "ffprobe".to_string()) +} + +static FFMPEG_CACHED: AtomicBool = AtomicBool::new(false); +static FFMPEG_VALUE: AtomicBool = AtomicBool::new(false); + +static FFPROBE_CACHED: AtomicBool = AtomicBool::new(false); +static FFPROBE_VALUE: AtomicBool = AtomicBool::new(false); + +/// Check whether ffmpeg is available on the system (result cached after first call). +pub fn ffmpeg_available() -> bool { + if FFMPEG_CACHED.load(Ordering::Acquire) { + return FFMPEG_VALUE.load(Ordering::Acquire); + } + let available = Command::new(ffmpeg_bin()) + .arg("-version") + .stdout(Stdio::null()) + .stderr(Stdio::null()) + .status() + .map(|s| s.success()) + .unwrap_or(false); + FFMPEG_VALUE.store(available, Ordering::Release); + FFMPEG_CACHED.store(true, Ordering::Release); + available +} + +/// Check whether ffprobe is available on the system (result cached after first call). +pub fn ffprobe_available() -> bool { + if FFPROBE_CACHED.load(Ordering::Acquire) { + return FFPROBE_VALUE.load(Ordering::Acquire); + } + let available = Command::new(ffprobe_bin()) + .arg("-version") + .stdout(Stdio::null()) + .stderr(Stdio::null()) + .status() + .map(|s| s.success()) + .unwrap_or(false); + FFPROBE_VALUE.store(available, Ordering::Release); + FFPROBE_CACHED.store(true, Ordering::Release); + available +} + +/// Reset ffmpeg/ffprobe availability caches (for testing). +pub fn reset_ffmpeg_cache() { + FFMPEG_CACHED.store(false, Ordering::Release); + FFPROBE_CACHED.store(false, Ordering::Release); +} + +fn run_command_with_piped_input(bin: String, input: &[u8], args: &[&str]) -> Result { + let mut child = Command::new(&bin) + .args(args) + .stdin(Stdio::piped()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .spawn() + .map_err(|_| MediaError::FfmpegNotFound)?; + + std::thread::scope(|scope| { + let writer = child.stdin.take().map(|mut stdin| { + scope.spawn(move || { + // Ignore broken pipe errors because some ffmpeg/ffprobe invocations + // may stop reading stdin once they have enough input. + let _ = stdin.write_all(input); + }) + }); + + let output = child.wait_with_output().map_err(|e| MediaError::FfmpegFailed { + status: -1, + stderr: e.to_string(), + })?; + + if let Some(writer) = writer { + let _ = writer.join(); + } + + Ok(output) + }) +} + +/// Run ffmpeg with the given input bytes piped to stdin, using the provided arguments. +/// Returns the stdout output on success. +pub fn run_ffmpeg(input: &[u8], args: &[&str]) -> Result, MediaError> { + if !ffmpeg_available() { + return Err(MediaError::FfmpegNotFound); + } + + let output = run_command_with_piped_input(ffmpeg_bin(), input, args)?; + + if !output.status.success() { + return Err(MediaError::FfmpegFailed { + status: output.status.code().unwrap_or(-1), + stderr: String::from_utf8_lossy(&output.stderr).to_string(), + }); + } + + if output.stdout.is_empty() { + return Err(MediaError::FfmpegFailed { + status: 0, + stderr: "ffmpeg produced no output".to_string(), + }); + } + + Ok(output.stdout) +} + +/// Run ffprobe with the given input bytes piped to stdin, using the provided arguments. +/// Returns the stdout output as a string on success. +pub fn run_ffprobe(input: &[u8], args: &[&str]) -> Result { + if !ffprobe_available() { + return Err(MediaError::FfmpegNotFound); + } + + let output = run_command_with_piped_input(ffprobe_bin(), input, args)?; + + if !output.status.success() { + return Err(MediaError::FfmpegFailed { + status: output.status.code().unwrap_or(-1), + stderr: String::from_utf8_lossy(&output.stderr).to_string(), + }); + } + + Ok(String::from_utf8_lossy(&output.stdout).trim().to_string()) +} diff --git a/crates/wx-media/src/hardlink.rs b/crates/wx-media/src/hardlink.rs new file mode 100644 index 0000000..5aa3580 --- /dev/null +++ b/crates/wx-media/src/hardlink.rs @@ -0,0 +1,107 @@ +use std::path::Path; + +use crate::error::MediaError; +use crate::types::HardlinkEntry; + +/// Query `hardlink.db` for image/video/file entries by MD5 or file_name prefix. +/// +/// Automatically falls back from v3 to v4 table variants. +/// For `"image"` type, results are sorted with `_h.dat` (high quality) first. +pub fn query_hardlink( + db_path: &Path, + media_type: &str, + key: &str, +) -> Result, MediaError> { + let conn = + rusqlite::Connection::open_with_flags(db_path, rusqlite::OpenFlags::SQLITE_OPEN_READ_ONLY)?; + query_hardlink_with_conn(&conn, media_type, key) +} + +pub fn query_hardlink_with_conn( + conn: &rusqlite::Connection, + media_type: &str, + key: &str, +) -> Result, MediaError> { + let table = resolve_table(&conn, media_type)?; + + let query = format!( + "SELECT f.md5, f.file_name, f.file_size, f.modify_time, + IFNULL(d1.username, ''), IFNULL(d2.username, '') + FROM {} f + LEFT JOIN dir2id d1 ON d1.rowid = f.dir1 + LEFT JOIN dir2id d2 ON d2.rowid = f.dir2 + WHERE f.md5 = ?1 OR f.file_name LIKE ?2 || '%'", + table + ); + + let mut stmt = conn.prepare(&query)?; + let rows = stmt.query_map(rusqlite::params![key, key], |row| { + Ok(HardlinkEntry { + media_type: media_type.to_string(), + md5: row.get(0)?, + file_name: row.get(1)?, + file_size: row.get(2)?, + modify_time: row.get(3)?, + dir1: row.get(4)?, + dir2: row.get(5)?, + }) + })?; + + let mut entries: Vec = rows.filter_map(|r| r.ok()).collect(); + + if entries.is_empty() { + return Err(MediaError::LookupMiss(format!( + "no {} entries found for key '{}'", + media_type, key, + ))); + } + + // For images, sort _h.dat (high quality) first + if media_type == "image" { + entries.sort_by(|a, b| { + let a_h = a.file_name.contains("_h."); + let b_h = b.file_name.contains("_h."); + b_h.cmp(&a_h) // true (has _h) sorts before false + }); + } + + Ok(entries) +} + +/// Resolve the correct table name, falling back from v3 to v4. +fn resolve_table(conn: &rusqlite::Connection, media_type: &str) -> Result { + let prefix = match media_type { + "image" => "image", + "video" => "video", + "file" => "file", + other => { + return Err(MediaError::InvalidFormat { + reason: format!("unsupported media type: {}", other), + }) + } + }; + + let v3 = format!("{}_hardlink_info_v3", prefix); + if table_exists(conn, &v3) { + return Ok(v3); + } + + let v4 = format!("{}_hardlink_info_v4", prefix); + if table_exists(conn, &v4) { + return Ok(v4); + } + + Err(MediaError::SchemaMissing(format!( + "no hardlink table found for type '{}' (tried {} and {})", + media_type, v3, v4 + ))) +} + +fn table_exists(conn: &rusqlite::Connection, table: &str) -> bool { + conn.query_row( + "SELECT name FROM sqlite_master WHERE type='table' AND name=?", + [table], + |_| Ok(()), + ) + .is_ok() +} diff --git a/crates/wx-media/src/image_resolver.rs b/crates/wx-media/src/image_resolver.rs new file mode 100644 index 0000000..0386b13 --- /dev/null +++ b/crates/wx-media/src/image_resolver.rs @@ -0,0 +1,140 @@ +use std::path::{Path, PathBuf}; + +use crate::error::MediaError; +use crate::resource; +use crate::types::MediaLookupResult; + +/// Resolve an image file from `local_id` through the full lookup chain: +/// +/// `local_id` → `message_resource.db(packed_info)` → MD5 → `attach//*/Img/*.dat` +/// +/// # Arguments +/// - `resource_db`: Path to `message_resource.db` +/// - `local_id`: Message local ID +/// - `username`: Chat partner username (for hashing to directory name) +/// - `attach_dir`: Root `attach/` directory +pub fn resolve_image( + resource_db: &Path, + local_id: i64, + username: &str, + attach_dir: &Path, +) -> Result { + // Step 1: Get packed_info from DB + let packed_info = resource::get_packed_info(resource_db, local_id)?; + + // Step 2: Extract MD5 + let file_md5 = resource::extract_md5_from_packed_info(&packed_info).ok_or_else(|| { + MediaError::PackedInfoParseFailed { + local_id, + reason: "no MD5 found in packed_info blob".into(), + } + })?; + + // Step 3: Locate .dat files + let username_hash = format!("{:x}", md5::compute(username.as_bytes())); + let search_base = attach_dir.join(&username_hash); + + let candidates = find_dat_files(&search_base, &file_md5); + + if candidates.is_empty() { + return Err(MediaError::NoDatFiles { + md5: file_md5, + path: search_base, + }); + } + + // Step 4: Recommend best file (original > _h > _t) + let recommended = pick_recommended(&candidates, &file_md5); + + Ok(MediaLookupResult { + file_md5, + candidates, + recommended, + }) +} + +/// Resolve image .dat files by MD5 directly (without local_id). +/// Used when caller already has MD5 from MessageContent::Image. +pub fn resolve_image_by_md5( + username: &str, + attach_dir: &Path, + file_md5: &str, +) -> Result { + let username_hash = format!("{:x}", md5::compute(username.as_bytes())); + let search_base = attach_dir.join(&username_hash); + let candidates = find_dat_files(&search_base, file_md5); + + if candidates.is_empty() { + return Err(MediaError::NoDatFiles { + md5: file_md5.to_string(), + path: search_base, + }); + } + + let recommended = pick_recommended(&candidates, file_md5); + + Ok(MediaLookupResult { + file_md5: file_md5.to_string(), + candidates, + recommended, + }) +} + +/// Find all `.dat` files matching `*.dat` under `/*/Img/`. +fn find_dat_files(base: &Path, file_md5: &str) -> Vec { + let mut results = Vec::new(); + + let entries = match std::fs::read_dir(base) { + Ok(e) => e, + Err(_) => return results, + }; + + for entry in entries.flatten() { + if !entry.file_type().is_ok_and(|t| t.is_dir()) { + continue; + } + let img_dir = entry.path().join("Img"); + if !img_dir.is_dir() { + continue; + } + let inner = match std::fs::read_dir(&img_dir) { + Ok(e) => e, + Err(_) => continue, + }; + for file_entry in inner.flatten() { + let name = file_entry.file_name(); + let name_str = name.to_string_lossy(); + if name_str.starts_with(file_md5) && name_str.ends_with(".dat") { + results.push(file_entry.path()); + } + } + } + + results.sort(); + results +} + +/// Pick the recommended file: prefer `_h.dat` > exact `md5.dat` > `_t.dat`. +fn pick_recommended(candidates: &[PathBuf], file_md5: &str) -> Option { + let exact = format!("{}.dat", file_md5); + let hd = format!("{}_h.dat", file_md5); + + // Prefer _h (high quality/original download) + if let Some(p) = candidates + .iter() + .find(|p| p.file_name().is_some_and(|n| n.to_string_lossy() == hd)) + { + return Some(p.clone()); + } + + // Then exact match + if let Some(p) = candidates + .iter() + .find(|p| p.file_name().is_some_and(|n| n.to_string_lossy() == exact)) + { + return Some(p.clone()); + } + + // Fallback to first available + candidates.first().cloned() +} diff --git a/crates/wx-media/src/image_transcode.rs b/crates/wx-media/src/image_transcode.rs new file mode 100644 index 0000000..1691a93 --- /dev/null +++ b/crates/wx-media/src/image_transcode.rs @@ -0,0 +1,118 @@ +use crate::error::MediaError; +use crate::ffmpeg::{ffmpeg_available, ffprobe_available, run_ffmpeg, run_ffprobe}; +use crate::types::TranscodeImageResult; +use crate::wxgf::{parse_wxgf, WxgfContent}; + +/// Transcode a WXGF container to a standard image format. +/// +/// - Embedded JPG/PNG: returned directly, `transcoded: true`. +/// - HEVC + ffmpeg available: single frame → PNG, multi frame → GIF, `transcoded: true`. +/// - HEVC + no ffmpeg: returns raw HEVC bytes, `transcoded: false`. +pub fn transcode_wxgf(data: &[u8]) -> Result { + let content = parse_wxgf(data)?; + + match content { + WxgfContent::EmbeddedImage { data, ext } => Ok(TranscodeImageResult { + data, + ext, + transcoded: true, + }), + WxgfContent::Hevc(hevc) => { + if !ffmpeg_available() { + return Ok(TranscodeImageResult { + data: hevc, + ext: "hevc", + transcoded: false, + }); + } + + let frame_count = count_hevc_frames(&hevc); + + if frame_count > 1 { + // Multi-frame → GIF + let gif = run_ffmpeg( + &hevc, + &[ + "-hide_banner", + "-loglevel", + "error", + "-f", + "hevc", + "-i", + "pipe:0", + "-filter_complex", + "[0:v]split[s0][s1];[s0]palettegen[p];[s1][p]paletteuse", + "-loop", + "0", + "-f", + "gif", + "pipe:1", + ], + )?; + Ok(TranscodeImageResult { + data: gif, + ext: "gif", + transcoded: true, + }) + } else { + // Single frame → PNG + let png = run_ffmpeg( + &hevc, + &[ + "-hide_banner", + "-loglevel", + "error", + "-f", + "hevc", + "-i", + "pipe:0", + "-frames:v", + "1", + "-f", + "image2pipe", + "-vcodec", + "png", + "pipe:1", + ], + )?; + Ok(TranscodeImageResult { + data: png, + ext: "png", + transcoded: true, + }) + } + } + } +} + +/// Count frames in an HEVC bitstream using ffprobe. +/// Returns 1 as default if ffprobe is unavailable or parsing fails. +fn count_hevc_frames(hevc: &[u8]) -> usize { + if !ffprobe_available() { + return 1; + } + + let result = run_ffprobe( + hevc, + &[ + "-v", + "error", + "-count_frames", + "-select_streams", + "v:0", + "-show_entries", + "stream=nb_read_frames", + "-of", + "default=nw=1:nk=1", + "-f", + "hevc", + "-i", + "pipe:0", + ], + ); + + match result { + Ok(s) => s.trim().parse::().unwrap_or(1), + Err(_) => 1, + } +} diff --git a/crates/wx-media/src/isaac64.rs b/crates/wx-media/src/isaac64.rs new file mode 100644 index 0000000..3a64d7f --- /dev/null +++ b/crates/wx-media/src/isaac64.rs @@ -0,0 +1,165 @@ +/// ISAAC-64: A fast cryptographic PRNG used for WeChat Channels video decryption. +/// +/// Ported from the CipherTalk TypeScript implementation which matches +/// WeChat's standard ISAAC-64 with big-endian keystream and reverse-index consumption. +pub struct Isaac64 { + mm: [u64; 256], + randrsl: [u64; 256], + aa: u64, + bb: u64, + cc: u64, + randcnt: usize, +} + +impl Isaac64 { + pub fn new(seed: u64) -> Self { + let mut rng = Isaac64 { + mm: [0u64; 256], + randrsl: [0u64; 256], + aa: 0, + bb: 0, + cc: 0, + randcnt: 0, + }; + rng.randrsl[0] = seed; + rng.init(); + rng + } + + pub fn next_u64(&mut self) -> u64 { + if self.randcnt == 0 { + self.generate(); + self.randcnt = 256; + } + self.randcnt -= 1; + self.randrsl[self.randcnt] + } + + /// Generate keystream bytes. Each u64 is written as big-endian. + /// Trailing bytes (when `len` is not a multiple of 8) take the BE prefix. + pub fn keystream(&mut self, len: usize) -> Vec { + let mut buf = Vec::with_capacity(len); + let full_blocks = len / 8; + + for _ in 0..full_blocks { + buf.extend_from_slice(&self.next_u64().to_be_bytes()); + } + + let remaining = len % 8; + if remaining > 0 { + let last = self.next_u64().to_be_bytes(); + buf.extend_from_slice(&last[..remaining]); + } + + buf + } + + fn init(&mut self) { + const GOLDEN: u64 = 0x9e3779b97f4a7c15; + let (mut a, mut b, mut c, mut d) = (GOLDEN, GOLDEN, GOLDEN, GOLDEN); + let (mut e, mut f, mut g, mut h) = (GOLDEN, GOLDEN, GOLDEN, GOLDEN); + + macro_rules! mix { + () => { + a = a.wrapping_sub(e); + f ^= h >> 9; + h = h.wrapping_add(a); + b = b.wrapping_sub(f); + g ^= a << 9; + a = a.wrapping_add(b); + c = c.wrapping_sub(g); + h ^= b >> 23; + b = b.wrapping_add(c); + d = d.wrapping_sub(h); + a ^= c << 15; + c = c.wrapping_add(d); + e = e.wrapping_sub(a); + b ^= d >> 14; + d = d.wrapping_add(e); + f = f.wrapping_sub(b); + c ^= e << 20; + e = e.wrapping_add(f); + g = g.wrapping_sub(c); + d ^= f >> 17; + f = f.wrapping_add(g); + h = h.wrapping_sub(d); + e ^= g << 14; + g = g.wrapping_add(h); + }; + } + + // 4 rounds of mixing + for _ in 0..4 { + mix!(); + } + + // First pass: mix in seed material from randrsl + for i in (0..256).step_by(8) { + a = a.wrapping_add(self.randrsl[i]); + b = b.wrapping_add(self.randrsl[i + 1]); + c = c.wrapping_add(self.randrsl[i + 2]); + d = d.wrapping_add(self.randrsl[i + 3]); + e = e.wrapping_add(self.randrsl[i + 4]); + f = f.wrapping_add(self.randrsl[i + 5]); + g = g.wrapping_add(self.randrsl[i + 6]); + h = h.wrapping_add(self.randrsl[i + 7]); + mix!(); + self.mm[i] = a; + self.mm[i + 1] = b; + self.mm[i + 2] = c; + self.mm[i + 3] = d; + self.mm[i + 4] = e; + self.mm[i + 5] = f; + self.mm[i + 6] = g; + self.mm[i + 7] = h; + } + + // Second pass: mix in mm values + for i in (0..256).step_by(8) { + a = a.wrapping_add(self.mm[i]); + b = b.wrapping_add(self.mm[i + 1]); + c = c.wrapping_add(self.mm[i + 2]); + d = d.wrapping_add(self.mm[i + 3]); + e = e.wrapping_add(self.mm[i + 4]); + f = f.wrapping_add(self.mm[i + 5]); + g = g.wrapping_add(self.mm[i + 6]); + h = h.wrapping_add(self.mm[i + 7]); + mix!(); + self.mm[i] = a; + self.mm[i + 1] = b; + self.mm[i + 2] = c; + self.mm[i + 3] = d; + self.mm[i + 4] = e; + self.mm[i + 5] = f; + self.mm[i + 6] = g; + self.mm[i + 7] = h; + } + + // Generate first batch + self.generate(); + self.randcnt = 256; + } + + fn generate(&mut self) { + self.cc = self.cc.wrapping_add(1); + self.bb = self.bb.wrapping_add(self.cc); + + for i in 0..256 { + let x = self.mm[i]; + match i & 3 { + 0 => self.aa ^= !(self.aa << 21), + 1 => self.aa ^= self.aa >> 5, + 2 => self.aa ^= self.aa << 12, + 3 => self.aa ^= self.aa >> 33, + _ => unreachable!(), + } + self.aa = self.mm[(i + 128) & 255].wrapping_add(self.aa); + let y = self.mm[((x >> 3) as usize) & 255] + .wrapping_add(self.aa) + .wrapping_add(self.bb); + self.mm[i] = y; + self.bb = self.mm[((y >> 11) as usize) & 255].wrapping_add(x); + self.randrsl[i] = self.bb; + } + } +} diff --git a/crates/wx-media/src/key.rs b/crates/wx-media/src/key.rs new file mode 100644 index 0000000..241562d --- /dev/null +++ b/crates/wx-media/src/key.rs @@ -0,0 +1,221 @@ +use std::path::Path; + +use base64::Engine; + +use crate::error::MediaError; + +/// Derive V2 AES key from UIN and WXID. +/// +/// Formula: `MD5(format!("{uin}{wxid}")).hex()[:16].as_bytes()` → 16 ASCII bytes. +/// +/// The V1 fixed key `cfcd208495d565ef` is a special case: `MD5("0")[:16]`. +pub fn derive_v2_aes_key(uin: &str, wxid: &str) -> [u8; 16] { + let input = format!("{uin}{wxid}"); + let digest = md5::compute(input.as_bytes()); + let hex_str = format!("{digest:x}"); // 32-char lowercase hex + let mut key = [0u8; 16]; + key.copy_from_slice(&hex_str.as_bytes()[..16]); + key +} + +/// Extract canonical base account ID from a directory name. +/// +/// Example: `wxid_example123abc_ab12` → `"wxid_example123abc"` +/// Example: `testuser001_1662` → `"testuser001_1662"` without extra signal +/// Example: `wxid_test` → `"wxid_test"` (not stripped: would leave bare `wxid`) +pub fn extract_wxid(dir_name: &str) -> String { + wx_keychain::account_id::canonical_base_for_account_dir(dir_name) +} + +fn extract_wxid_from_data_dir(data_dir: &Path) -> Result { + let dir_name = data_dir + .file_name() + .and_then(|n| n.to_str()) + .ok_or_else(|| MediaError::InvalidFormat { + reason: "cannot determine directory name from data_dir".into(), + })?; + + Ok(data_dir.parent().map_or_else( + || extract_wxid(dir_name), + |root| { + wx_keychain::process::extract_base_wxid_for_account_dir_under_root(root, dir_name) + }, + )) +} + +/// Read UIN from `config.ini` for the given WeChat account data directory. +/// +/// Searches for `last_uin=` in config.ini under the ilink directory. +/// +/// On macOS, the directory layout separates account data from shared app_data: +/// ```text +/// Documents/ +/// ├── app_data/radium/ilink//kvcomm/config.ini +/// └── xwechat_files// ← data_dir +/// ``` +/// +/// This function tries two locations: +/// 1. `/app_data/radium/ilink/` (for tests / alternative layouts) +/// 2. `/../../app_data/radium/ilink/` (standard macOS layout) +/// +/// When multiple accounts exist under ilink, the `account_suffix` (e.g. `ab12` +/// from `wxid_xxx_ab12`) is used to match the correct subdirectory. +pub fn read_uin(data_dir: &Path) -> Result { + let suffix = + extract_account_suffix(data_dir.file_name().and_then(|n| n.to_str()).unwrap_or("")); + + // Candidate ilink directories (in priority order) + let candidates = [ + data_dir.join("app_data").join("radium").join("ilink"), + data_dir + .join("..") + .join("..") + .join("app_data") + .join("radium") + .join("ilink"), + ]; + + for ilink_dir in &candidates { + if !ilink_dir.is_dir() { + continue; + } + if let Ok(uin) = read_uin_from_ilink(ilink_dir, suffix.as_deref()) { + return Ok(uin); + } + } + + Err(MediaError::NotFound(format!( + "no config.ini with last_uin found (searched {} candidate paths)", + candidates.len() + ))) +} + +/// Extract the 4-char account hash suffix from a directory name. +/// +/// Delegates to `wx_keychain::AccountId::parse` for consistent suffix detection. +/// +/// `wxid_example123abc_ab12` → `Some("ab12")` +fn extract_account_suffix(dir_name: &str) -> Option { + wx_keychain::AccountId::parse(dir_name) + .suffix() + .map(|s| s.to_string()) +} + +/// Scan an ilink directory for config.ini with last_uin. +/// +/// If `account_suffix` is provided, prefer subdirectories whose name starts with it. +fn read_uin_from_ilink( + ilink_dir: &Path, + account_suffix: Option<&str>, +) -> Result { + let mut entries: Vec<_> = std::fs::read_dir(ilink_dir)? + .filter_map(|e| e.ok()) + .filter(|e| e.file_type().map(|ft| ft.is_dir()).unwrap_or(false)) + .collect(); + + // Sort: matching-suffix directories first + if let Some(suffix) = account_suffix { + entries.sort_by_key(|e| { + let name = e.file_name().to_string_lossy().to_string(); + if name.starts_with(suffix) { + 0 + } else { + 1 + } + }); + } + + for entry in entries { + let config_path = entry.path().join("kvcomm").join("config.ini"); + if !config_path.is_file() { + continue; + } + let content = std::fs::read_to_string(&config_path)?; + if let Some(uin) = parse_uin_from_config(&content) { + return Ok(uin); + } + } + + Err(MediaError::NotFound( + "no config.ini with last_uin found in ilink directory".into(), + )) +} + +/// Parse `last_uin=` from config.ini content and return the decoded UIN. +fn parse_uin_from_config(content: &str) -> Option { + for line in content.lines() { + let line = line.trim(); + if let Some(value) = line.strip_prefix("last_uin=") { + let value = value.trim(); + if value.is_empty() { + continue; + } + let decoded = base64::engine::general_purpose::STANDARD + .decode(value) + .ok()?; + let uin = String::from_utf8(decoded).ok()?; + if !uin.is_empty() && uin.chars().all(|c| c.is_ascii_digit()) { + return Some(uin); + } + } + } + None +} + +/// Derive V2 AES key automatically from a WeChat account data directory. +/// +/// Combines: read UIN from config.ini + extract WXID from directory name + MD5 derivation. +pub fn derive_v2_key_from_dir(data_dir: &Path) -> Result<[u8; 16], MediaError> { + let wxid = extract_wxid_from_data_dir(data_dir)?; + let uin = read_uin(data_dir)?; + + Ok(derive_v2_aes_key(&uin, &wxid)) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_parse_uin_from_config_valid() { + let config = "[General]\nlast_uin=MTIzNDU2Nzg5MA==\n"; + assert_eq!( + parse_uin_from_config(config), + Some("1234567890".to_string()) + ); + } + + #[test] + fn test_parse_uin_from_config_missing() { + assert_eq!(parse_uin_from_config("[General]\nsome_key=value\n"), None); + } + + #[test] + fn test_parse_uin_from_config_empty_value() { + assert_eq!(parse_uin_from_config("last_uin=\n"), None); + } + + #[test] + fn test_parse_uin_from_config_non_numeric() { + assert_eq!(parse_uin_from_config("last_uin=YWJj\n"), None); // "abc" + } + + #[test] + fn test_extract_account_suffix() { + assert_eq!( + extract_account_suffix("wxid_test_ab12"), + Some("ab12".to_string()) + ); + assert_eq!( + extract_account_suffix("wxid_test_c3e7"), + Some("c3e7".to_string()) + ); + assert_eq!( + extract_account_suffix("wxid_test"), + Some("test".to_string()) + ); + assert_eq!(extract_account_suffix("wxid_test_abcde"), None); + // "no_underscore" has rfind('_') at the `_` after `no`, tail = "underscore" (10 chars, not 4) + assert_eq!(extract_account_suffix("no_underscore"), None); + } +} diff --git a/crates/wx-media/src/lib.rs b/crates/wx-media/src/lib.rs new file mode 100644 index 0000000..73d79a0 --- /dev/null +++ b/crates/wx-media/src/lib.rs @@ -0,0 +1,106 @@ +//! WeChat media decryption, transcoding, and resource resolution. +//! +//! # Decrypt a `.dat` image file +//! +//! ```no_run +//! use wx_media::{decrypt_dat, DatDecryptOptions}; +//! +//! let data = std::fs::read("image.dat").unwrap(); +//! // XOR / V1 files need no extra key: +//! let result = decrypt_dat(&data, &DatDecryptOptions::default()).unwrap(); +//! std::fs::write(format!("output.{}", result.ext), &result.data).unwrap(); +//! +//! // V2 files require a 16-byte AES key: +//! let opts = DatDecryptOptions { +//! v2_aes_key: Some(*b"abcdefghijklmnop"), +//! xor_key: None, +//! }; +//! let result = decrypt_dat(&data, &opts).unwrap(); +//! ``` +//! +//! # Transcode a WXGF image to PNG/GIF +//! +//! ```no_run +//! use wx_media::transcode_wxgf; +//! +//! let wxgf_data = std::fs::read("image.wxgf").unwrap(); +//! let result = transcode_wxgf(&wxgf_data).unwrap(); +//! std::fs::write(format!("output.{}", result.ext), &result.data).unwrap(); +//! // result.transcoded == true if ffmpeg converted to PNG/GIF +//! ``` +//! +//! # Decrypt a WeChat Channels video +//! +//! ```no_run +//! use wx_media::decrypt_video; +//! +//! let encrypted = std::fs::read("video.enc").unwrap(); +//! let result = decrypt_video(&encrypted, 12345); // seed from decode_key +//! std::fs::write("output.mp4", &result.data).unwrap(); +//! // result.is_valid_mp4 == true if "ftyp" signature found +//! ``` +//! +//! # Resolve an image by `local_id` +//! +//! ```no_run +//! use std::path::Path; +//! use wx_media::resolve_image; +//! +//! let lookup = resolve_image( +//! Path::new("db_storage/message/message_resource.db"), +//! 12345, // local_id from Message table +//! "wxid_alice", // chat partner username +//! Path::new("msg/attach"), // attach base directory +//! ).unwrap(); +//! // lookup.recommended is the best-match .dat file path +//! ``` +//! +//! # Extract a voice BLOB +//! +//! ```no_run +//! use std::path::Path; +//! use wx_media::extract_voice; +//! +//! let blob = extract_voice(Path::new("db_storage/media"), "123456789").unwrap(); +//! std::fs::write("voice.silk", &blob.data).unwrap(); +//! ``` + +pub mod audio_transcode; +mod dat; +mod error; +mod fallback; +pub mod ffmpeg; +mod hardlink; +mod image_resolver; +pub mod image_transcode; +pub mod isaac64; +pub mod key; +mod resource; +mod types; +pub mod video_decrypt; +mod voice; +pub mod wxgf; + +pub use audio_transcode::{transcode_silk_to_mp3, transcode_silk_to_ogg_opus}; +pub use dat::{decrypt_dat, detect_dat_format, detect_image_type, detect_xor_key}; +pub use error::MediaError; +pub use fallback::{find_file_by_name, find_video_by_md5}; +pub use ffmpeg::{ffmpeg_available, ffprobe_available, reset_ffmpeg_cache, run_ffmpeg, run_ffprobe}; +pub use hardlink::{query_hardlink, query_hardlink_with_conn}; +pub use image_resolver::{resolve_image, resolve_image_by_md5}; +pub use image_transcode::transcode_wxgf; +pub use isaac64::Isaac64; +pub use key::{derive_v2_aes_key, derive_v2_key_from_dir, extract_wxid, read_uin}; +pub use resource::extract_md5_from_packed_info; +pub use types::{ + DatDecryptOptions, DatFormat, DecodedImage, DecryptVideoResult, HardlinkEntry, ImageType, + MediaLookupResult, TranscodeAudioResult, TranscodeImageResult, VoiceBlob, +}; +pub use video_decrypt::{decrypt_video, decrypt_video_with_keystream}; +pub use voice::{extract_voice, extract_voice_with_conn, extract_voice_with_conn_hint, find_media_dbs}; +pub use wxgf::{parse_wxgf, WxgfContent}; + +/// Compute MD5 hash of bytes, returning the `md5::Digest` (displays as hex). +pub fn md5_hash(data: &[u8]) -> md5::Digest { + md5::compute(data) +} diff --git a/crates/wx-media/src/resource.rs b/crates/wx-media/src/resource.rs new file mode 100644 index 0000000..e73055a --- /dev/null +++ b/crates/wx-media/src/resource.rs @@ -0,0 +1,73 @@ +use crate::error::MediaError; + +/// Protobuf marker preceding the 32-byte MD5 hex string in `packed_info`. +const PROTOBUF_MARKER: &[u8] = b"\x12\x22\x0a\x20"; + +/// Extract the 32-char hex MD5 from a `packed_info` protobuf blob. +/// +/// Strategy: +/// 1. Primary: locate protobuf marker `\x12\x22\x0a\x20`, read next 32 bytes as hex. +/// 2. Fallback: scan for 32 contiguous lowercase hex characters. +pub fn extract_md5_from_packed_info(blob: &[u8]) -> Option { + if blob.is_empty() { + return None; + } + + // Primary: protobuf marker + if let Some(idx) = find_subsequence(blob, PROTOBUF_MARKER) { + let start = idx + PROTOBUF_MARKER.len(); + if start + 32 <= blob.len() { + if let Ok(s) = std::str::from_utf8(&blob[start..start + 32]) { + if is_hex_string(s) { + return Some(s.to_string()); + } + } + } + } + + // Fallback: scan for 32 contiguous hex chars + let hex_chars: &[u8] = b"0123456789abcdef"; + let mut i = 0; + while i + 32 <= blob.len() { + if hex_chars.contains(&blob[i]) { + let candidate = &blob[i..i + 32]; + if candidate.iter().all(|b| hex_chars.contains(b)) { + if let Ok(s) = std::str::from_utf8(candidate) { + return Some(s.to_string()); + } + } + i += 32; + } else { + i += 1; + } + } + + None +} + +/// Look up `packed_info` for a given `local_id` in `message_resource.db`. +pub fn get_packed_info(db_path: &std::path::Path, local_id: i64) -> Result, MediaError> { + let conn = + rusqlite::Connection::open_with_flags(db_path, rusqlite::OpenFlags::SQLITE_OPEN_READ_ONLY)?; + let mut stmt = + conn.prepare("SELECT packed_info FROM MessageResourceInfo WHERE local_id = ?")?; + let blob: Vec = stmt + .query_row(rusqlite::params![local_id], |row| row.get(0)) + .map_err(|e| match e { + rusqlite::Error::QueryReturnedNoRows => { + MediaError::NotFound(format!("local_id {} not in MessageResourceInfo", local_id)) + } + other => MediaError::Sqlite(other), + })?; + Ok(blob) +} + +fn find_subsequence(haystack: &[u8], needle: &[u8]) -> Option { + haystack.windows(needle.len()).position(|w| w == needle) +} + +fn is_hex_string(s: &str) -> bool { + s.len() == 32 + && s.chars() + .all(|c| c.is_ascii_hexdigit() && !c.is_ascii_uppercase()) +} diff --git a/crates/wx-media/src/types.rs b/crates/wx-media/src/types.rs new file mode 100644 index 0000000..22decaf --- /dev/null +++ b/crates/wx-media/src/types.rs @@ -0,0 +1,133 @@ +/// Detected `.dat` file encryption format. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum DatFormat { + /// Single-byte XOR encryption (legacy). + Xor, + /// V1: AES-128-ECB with fixed key `cfcd208495d565ef`. + V1, + /// V2: AES-128-ECB with per-account key + XOR tail. + V2, +} + +/// Options for decrypting a `.dat` file. +#[derive(Debug, Clone, Default)] +pub struct DatDecryptOptions { + /// AES key for V2 format (16 bytes, ASCII alphanumeric). + /// Required for V2, ignored for XOR/V1. + pub v2_aes_key: Option<[u8; 16]>, + /// XOR key for V2 tail section. If `None`, uses auto-detected or default. + pub xor_key: Option, +} + +/// Result of decrypting a `.dat` file. +#[derive(Debug)] +pub struct DecodedImage { + /// Decrypted image bytes. + pub data: Vec, + /// Detected encryption format. + pub format: DatFormat, + /// Detected image file extension (jpg, png, gif, bmp, webp, tif, wxgf, bin). + pub ext: String, +} + +/// Detected image type from magic bytes. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum ImageType { + Jpg, + Png, + Gif, + Bmp, + Webp, + Tif, + Wxgf, + Unknown, +} + +impl ImageType { + pub fn ext(&self) -> &'static str { + match self { + Self::Jpg => "jpg", + Self::Png => "png", + Self::Gif => "gif", + Self::Bmp => "bmp", + Self::Webp => "webp", + Self::Tif => "tif", + Self::Wxgf => "wxgf", + Self::Unknown => "bin", + } + } +} + +/// Result of resolving a media file path from message metadata. +#[derive(Debug, Clone)] +pub struct MediaLookupResult { + /// The md5 extracted from packed_info. + pub file_md5: String, + /// All candidate `.dat` file paths found. + pub candidates: Vec, + /// Recommended file (original > `_h` > `_t`). + pub recommended: Option, +} + +/// Hardlink query result for image/video/file. +#[derive(Debug, Clone, serde::Serialize)] +pub struct HardlinkEntry { + /// Media type: "image", "video", "file". + pub media_type: String, + /// MD5 key. + pub md5: String, + /// File name. + pub file_name: String, + /// File size in bytes. + pub file_size: i64, + /// Last modification time (unix seconds). + pub modify_time: i64, + /// First directory segment. + pub dir1: String, + /// Second directory segment. + pub dir2: String, +} + +/// Result of decrypting a WeChat Channels video. +#[derive(Debug)] +pub struct DecryptVideoResult { + /// Decrypted video bytes. + pub data: Vec, + /// Whether the decrypted data contains an MP4 "ftyp" signature in the first 32 bytes. + pub is_valid_mp4: bool, +} + +/// Result of transcoding a WXGF image. +#[derive(Debug)] +pub struct TranscodeImageResult { + /// Output image bytes. + pub data: Vec, + /// File extension: "png", "gif", "jpg", or "hevc" (fallback). + pub ext: &'static str, + /// `true` if data is in a standard viewable format; `false` if raw HEVC (ffmpeg missing). + pub transcoded: bool, +} + +/// Result of transcoding SILK audio. +#[derive(Debug)] +pub struct TranscodeAudioResult { + /// Output audio bytes. + pub data: Vec, + /// File extension: "ogg", "mp3", or "silk" (fallback when ffmpeg absent). + pub ext: &'static str, + /// Response MIME type for the output payload. + pub mime: &'static str, + /// `true` if output was transcoded; `false` if original SILK was returned unchanged. + pub transcoded: bool, +} + +/// Extracted voice blob from `media_N.db`. +#[derive(Debug)] +pub struct VoiceBlob { + /// Server ID used to locate this voice. + pub svr_id: String, + /// Optional `VoiceInfo.chat_name_id` captured from indexed schemas. + pub chat_name_id: Option, + /// Raw SILK audio bytes. + pub data: Vec, +} diff --git a/crates/wx-media/src/video_decrypt.rs b/crates/wx-media/src/video_decrypt.rs new file mode 100644 index 0000000..2833862 --- /dev/null +++ b/crates/wx-media/src/video_decrypt.rs @@ -0,0 +1,37 @@ +use crate::isaac64::Isaac64; +use crate::types::DecryptVideoResult; + +const MAX_DECRYPT_LEN: usize = 131072; // 128 KB + +/// Decrypt a WeChat Channels encrypted video using Isaac64 keystream XOR. +/// +/// Only the first 128KB of the ciphertext is encrypted; the rest is plaintext. +/// Returns the decrypted data and whether it appears to be a valid MP4 (contains "ftyp"). +pub fn decrypt_video(ciphertext: &[u8], seed: u64) -> DecryptVideoResult { + let mut isaac = Isaac64::new(seed); + let decrypt_len = ciphertext.len().min(MAX_DECRYPT_LEN); + let keystream = isaac.keystream(decrypt_len); + decrypt_video_with_keystream(ciphertext, &keystream) +} + +/// Decrypt using a pre-generated keystream (useful for testing). +pub fn decrypt_video_with_keystream(ciphertext: &[u8], keystream: &[u8]) -> DecryptVideoResult { + let decrypt_len = ciphertext.len().min(keystream.len()); + let mut data = Vec::with_capacity(ciphertext.len()); + + // XOR the encrypted prefix + for i in 0..decrypt_len { + data.push(ciphertext[i] ^ keystream[i]); + } + + // Append unencrypted tail + if decrypt_len < ciphertext.len() { + data.extend_from_slice(&ciphertext[decrypt_len..]); + } + + // Check for MP4 signature in first 32 bytes + let check_len = data.len().min(32); + let is_valid_mp4 = data[..check_len].windows(4).any(|w| w == b"ftyp"); + + DecryptVideoResult { data, is_valid_mp4 } +} diff --git a/crates/wx-media/src/voice.rs b/crates/wx-media/src/voice.rs new file mode 100644 index 0000000..bd24733 --- /dev/null +++ b/crates/wx-media/src/voice.rs @@ -0,0 +1,177 @@ +use std::path::Path; + +use crate::error::MediaError; +use crate::types::VoiceBlob; + +fn voice_table_exists(conn: &rusqlite::Connection) -> Result { + Ok(conn.query_row( + "SELECT COUNT(*) FROM sqlite_master WHERE type='table' AND name='VoiceInfo'", + [], + |row| row.get::<_, i64>(0), + )? > 0) +} + +fn voice_table_has_chat_name_id(conn: &rusqlite::Connection) -> Result { + let mut stmt = conn.prepare("PRAGMA table_info([VoiceInfo])")?; + let mut rows = stmt.query([])?; + while let Some(row) = rows.next()? { + if row.get::<_, String>(1)? == "chat_name_id" { + return Ok(true); + } + } + Ok(false) +} + +fn query_voice_rows( + conn: &rusqlite::Connection, + sql: &str, + params: impl rusqlite::Params, + svr_id: &str, +) -> Result, MediaError> { + let mut stmt = conn.prepare(sql)?; + let mut rows = stmt.query(params)?; + + while let Some(row) = rows.next()? { + let chat_name_id = row.get::<_, Option>(0)?; + let data: Vec = row.get(1)?; + if !data.is_empty() { + return Ok(Some(VoiceBlob { + svr_id: svr_id.to_string(), + chat_name_id, + data, + })); + } + } + + Ok(None) +} + +/// Extract a voice BLOB from `media_*.db` files by `svr_id`. +/// +/// Scans all `media*.db` files in the given directory (supports `media.db`, +/// `media_0.db`, `media_1.db`, etc.). Returns the first non-empty match. +/// +/// Error classification: +/// - [`MediaError::NoMediaDbs`] — directory missing or no media DBs found +/// - [`MediaError::Sqlite`] — all accessible DBs produced SQLite errors +/// - [`MediaError::LookupMiss`] — query succeeded but svr_id not found +pub fn extract_voice(media_dir: &Path, svr_id: &str) -> Result { + let db_paths = find_media_dbs(media_dir)?; + + let mut first_sqlite_err: Option = None; + let mut any_queried = false; + + for db_path in &db_paths { + let conn = match rusqlite::Connection::open_with_flags( + db_path, + rusqlite::OpenFlags::SQLITE_OPEN_READ_ONLY, + ) { + Ok(c) => c, + Err(e) => { + if first_sqlite_err.is_none() { + first_sqlite_err = Some(e); + } + continue; + } + }; + + match extract_voice_with_conn(&conn, svr_id) { + Ok(blob) => return Ok(blob), + Err(MediaError::LookupMiss(_)) => { + any_queried = true; + } + Err(MediaError::Sqlite(err)) => { + if first_sqlite_err.is_none() { + first_sqlite_err = Some(err); + } + } + Err(other) => return Err(other), + } + } + + // If we never successfully queried any DB, report the SQLite error + if !any_queried { + if let Some(e) = first_sqlite_err { + return Err(MediaError::Sqlite(e)); + } + } + + Err(MediaError::LookupMiss(format!( + "voice not found for svr_id {svr_id}" + ))) +} + +pub fn extract_voice_with_conn( + conn: &rusqlite::Connection, + svr_id: &str, +) -> Result { + extract_voice_with_conn_hint(conn, svr_id, None) +} + +pub fn extract_voice_with_conn_hint( + conn: &rusqlite::Connection, + svr_id: &str, + chat_name_id_hint: Option, +) -> Result { + if !voice_table_exists(conn)? { + return Err(MediaError::SchemaMissing("VoiceInfo table missing".into())); + } + + let has_chat_name_id = voice_table_has_chat_name_id(conn)?; + + if has_chat_name_id { + if let Some(chat_name_id) = chat_name_id_hint { + if let Some(blob) = query_voice_rows( + conn, + "SELECT chat_name_id, voice_data FROM VoiceInfo WHERE chat_name_id = ? AND svr_id = ?", + rusqlite::params![chat_name_id, svr_id], + svr_id, + )? { + return Ok(blob); + } + } + + if let Some(blob) = query_voice_rows( + conn, + "SELECT chat_name_id, voice_data FROM VoiceInfo WHERE svr_id = ?", + rusqlite::params![svr_id], + svr_id, + )? { + return Ok(blob); + } + } else if let Some(blob) = query_voice_rows( + conn, + "SELECT NULL, voice_data FROM VoiceInfo WHERE svr_id = ?", + rusqlite::params![svr_id], + svr_id, + )? { + return Ok(blob); + } + + Err(MediaError::LookupMiss(format!( + "voice not found for svr_id {svr_id}" + ))) +} + +/// Find all `media*.db` files in the directory, sorted by name. +pub fn find_media_dbs(dir: &Path) -> Result, MediaError> { + let entries = std::fs::read_dir(dir).map_err(|_| MediaError::NoMediaDbs(dir.to_path_buf()))?; + + let mut paths: Vec = entries + .filter_map(|e| e.ok()) + .filter(|e| { + let name = e.file_name(); + let name = name.to_string_lossy(); + // Match: media.db, media_0.db, media_1.db, media_12.db, etc. + (name == "media.db" || name.starts_with("media_")) && name.ends_with(".db") + }) + .map(|e| e.path()) + .collect(); + + if paths.is_empty() { + return Err(MediaError::NoMediaDbs(dir.to_path_buf())); + } + + paths.sort(); + Ok(paths) +} diff --git a/crates/wx-media/src/wxgf.rs b/crates/wx-media/src/wxgf.rs new file mode 100644 index 0000000..196c689 --- /dev/null +++ b/crates/wx-media/src/wxgf.rs @@ -0,0 +1,61 @@ +use crate::error::MediaError; + +/// Result of parsing a WXGF container. +#[derive(Debug)] +pub enum WxgfContent { + /// Embedded standard image (JPG/PNG) found within the WXGF container. + EmbeddedImage { data: Vec, ext: &'static str }, + /// HEVC bitstream extracted from the WXGF container (needs ffmpeg to transcode). + Hevc(Vec), +} + +const WXGF_MAGIC: &[u8; 4] = b"wxgf"; +const HEVC_START_CODE_4: &[u8; 4] = &[0x00, 0x00, 0x00, 0x01]; +const HEVC_START_CODE_3: &[u8; 3] = &[0x00, 0x00, 0x01]; +const JPG_MAGIC: &[u8; 3] = &[0xFF, 0xD8, 0xFF]; +const PNG_MAGIC: &[u8; 4] = &[0x89, 0x50, 0x4E, 0x47]; + +/// Scan the first `limit` bytes (after WXGF magic) for embedded JPG/PNG. +const EMBEDDED_SCAN_LIMIT: usize = 4096; + +/// Parse a WXGF container, extracting the inner content. +/// +/// Returns `EmbeddedImage` if a standard image (JPG/PNG) is found within the container, +/// or `Hevc` if an HEVC bitstream is found. Returns an error if the data is not a valid +/// WXGF container or contains no recognizable content. +pub fn parse_wxgf(data: &[u8]) -> Result { + if data.len() < 4 || &data[..4] != WXGF_MAGIC { + return Err(MediaError::InvalidWxgf); + } + + // Check for embedded JPG/PNG in the first 4KB + let scan_end = data.len().min(EMBEDDED_SCAN_LIMIT); + for i in 4..scan_end { + if i + 3 <= data.len() && &data[i..i + 3] == JPG_MAGIC { + return Ok(WxgfContent::EmbeddedImage { + data: data[i..].to_vec(), + ext: "jpg", + }); + } + if i + 4 <= data.len() && &data[i..i + 4] == PNG_MAGIC { + return Ok(WxgfContent::EmbeddedImage { + data: data[i..].to_vec(), + ext: "png", + }); + } + } + + // Search for HEVC start code (4-byte first, then 3-byte fallback) + if let Some(pos) = find_subsequence(data, HEVC_START_CODE_4) { + return Ok(WxgfContent::Hevc(data[pos..].to_vec())); + } + if let Some(pos) = find_subsequence(data, HEVC_START_CODE_3) { + return Ok(WxgfContent::Hevc(data[pos..].to_vec())); + } + + Err(MediaError::InvalidWxgf) +} + +fn find_subsequence(haystack: &[u8], needle: &[u8]) -> Option { + haystack.windows(needle.len()).position(|w| w == needle) +} diff --git a/crates/wx-media/tests/audio-transcode-missing-ffmpeg.rs b/crates/wx-media/tests/audio-transcode-missing-ffmpeg.rs new file mode 100644 index 0000000..95d1148 --- /dev/null +++ b/crates/wx-media/tests/audio-transcode-missing-ffmpeg.rs @@ -0,0 +1,29 @@ +#[cfg(feature = "audio")] +use wx_media::MediaError; + +#[cfg(feature = "audio")] +fn silent_pcm_frame() -> Vec { + vec![0_u8; 24_000 / 1_000 * 40 * 2] +} + +#[cfg(feature = "audio")] +fn sample_silk() -> Vec { + silk_rs::encode_silk(silent_pcm_frame(), 24_000, 24_000, true).unwrap() +} + +#[cfg(feature = "audio")] +#[test] +fn audio_transcode_handles_missing_ffmpeg_with_explicit_results() { + unsafe { + std::env::set_var("FFMPEG_PATH", "/definitely-missing-ffmpeg"); + } + wx_media::reset_ffmpeg_cache(); + + let ogg = wx_media::transcode_silk_to_ogg_opus(&sample_silk()); + assert!(matches!(ogg, Err(MediaError::FfmpegNotFound))); + + let mp3 = wx_media::transcode_silk_to_mp3(&sample_silk()).unwrap(); + assert_eq!(mp3.ext, "silk"); + assert_eq!(mp3.mime, "audio/x-silk"); + assert!(!mp3.transcoded); +} diff --git a/crates/wx-media/tests/audio-transcode.rs b/crates/wx-media/tests/audio-transcode.rs new file mode 100644 index 0000000..651d392 --- /dev/null +++ b/crates/wx-media/tests/audio-transcode.rs @@ -0,0 +1,135 @@ +use wx_media::MediaError; + +#[cfg(feature = "audio")] +fn silent_pcm_frame() -> Vec { + vec![0_u8; 24_000 / 1_000 * 40 * 2] +} + +#[cfg(feature = "audio")] +fn sample_silk() -> Vec { + silk_rs::encode_silk(silent_pcm_frame(), 24_000, 24_000, true).unwrap() +} + +#[cfg(feature = "audio")] +fn long_sample_silk() -> Vec { + silk_rs::encode_silk(vec![0_u8; silent_pcm_frame().len() * 250], 24_000, 24_000, true) + .unwrap() +} + +#[cfg(feature = "audio")] +#[test] +fn audio_transcode_ogg_returns_ogg_when_ffmpeg_is_available() { + if !wx_media::ffmpeg_available() { + return; + } + + let result = wx_media::transcode_silk_to_ogg_opus(&sample_silk()).unwrap(); + assert_eq!(result.ext, "ogg"); + assert_eq!(result.mime, "audio/ogg"); + assert!(result.transcoded); + assert!(!result.data.is_empty()); +} + +#[cfg(feature = "audio")] +#[test] +fn audio_transcode_ogg_errors_when_ffmpeg_is_missing() { + if wx_media::ffmpeg_available() { + return; + } + + let result = wx_media::transcode_silk_to_ogg_opus(&sample_silk()); + assert!(matches!(result, Err(MediaError::FfmpegNotFound))); +} + +#[cfg(feature = "audio")] +#[test] +fn audio_transcode_mp3_keeps_existing_compatibility() { + let result = wx_media::transcode_silk_to_mp3(&sample_silk()).unwrap(); + if wx_media::ffmpeg_available() { + assert_eq!(result.ext, "mp3"); + assert_eq!(result.mime, "audio/mpeg"); + assert!(result.transcoded); + assert!(!result.data.is_empty()); + } else { + assert_eq!(result.ext, "silk"); + assert_eq!(result.mime, "audio/x-silk"); + assert!(!result.transcoded); + } +} + +#[cfg(feature = "audio")] +#[test] +fn audio_transcode_mp3_handles_long_audio_without_hanging() { + if !wx_media::ffmpeg_available() { + return; + } + + let output = run_long_audio_child(std::time::Duration::from_secs(5)); + assert!( + output.status.success(), + "child failed: stdout={}\nstderr={}", + String::from_utf8_lossy(&output.stdout), + String::from_utf8_lossy(&output.stderr) + ); + assert!(String::from_utf8_lossy(&output.stdout).contains("long-audio-ok")); +} + +#[cfg(feature = "audio")] +#[test] +fn audio_transcode_long_audio_child_mode() { + if std::env::var_os("WECHAT_MEDIA_LONG_AUDIO_CHILD").is_none() { + return; + } + + let result = wx_media::transcode_silk_to_mp3(&long_sample_silk()).unwrap(); + assert_eq!(result.ext, "mp3"); + assert_eq!(result.mime, "audio/mpeg"); + assert!(result.transcoded); + assert!(!result.data.is_empty()); + println!("long-audio-ok {}", result.data.len()); +} + +#[cfg(feature = "audio")] +fn run_long_audio_child(timeout: std::time::Duration) -> std::process::Output { + let current_exe = std::env::current_exe().unwrap(); + let mut child = std::process::Command::new(current_exe) + .arg("--exact") + .arg("audio_transcode_long_audio_child_mode") + .arg("--nocapture") + .arg("--test-threads=1") + .env("WECHAT_MEDIA_LONG_AUDIO_CHILD", "1") + .stdout(std::process::Stdio::piped()) + .stderr(std::process::Stdio::piped()) + .spawn() + .unwrap(); + + let start = std::time::Instant::now(); + loop { + if child.try_wait().unwrap().is_some() { + return child.wait_with_output().unwrap(); + } + + if start.elapsed() >= timeout { + let _ = child.kill(); + let output = child.wait_with_output().unwrap(); + panic!( + "child timed out after {:?}: stdout={}\nstderr={}", + timeout, + String::from_utf8_lossy(&output.stdout), + String::from_utf8_lossy(&output.stderr) + ); + } + + std::thread::sleep(std::time::Duration::from_millis(25)); + } +} + +#[cfg(not(feature = "audio"))] +#[test] +fn audio_transcode_feature_disabled_returns_error() { + let result = wx_media::transcode_silk_to_mp3(b"ignored"); + assert!(matches!(result, Err(MediaError::AudioFeatureDisabled))); + + let result = wx_media::transcode_silk_to_ogg_opus(b"ignored"); + assert!(matches!(result, Err(MediaError::AudioFeatureDisabled))); +} diff --git a/crates/wx-media/tests/dat-decrypt.rs b/crates/wx-media/tests/dat-decrypt.rs new file mode 100644 index 0000000..35a8024 --- /dev/null +++ b/crates/wx-media/tests/dat-decrypt.rs @@ -0,0 +1,269 @@ +use wx_media::{DatDecryptOptions, DatFormat, MediaError}; + +// ── Helpers ────────────────────────────────────────────────────────── + +/// Build a simple XOR-encrypted `.dat` from a known JPEG header. +fn make_xor_dat(key: u8) -> Vec { + // Minimal JPEG: FF D8 FF E0 + padding + let plain = [ + 0xFFu8, 0xD8, 0xFF, 0xE0, 0x00, 0x10, 0x4A, 0x46, 0x49, 0x46, 0x00, + ]; + plain.iter().map(|b| b ^ key).collect() +} + +/// Build a V1-encrypted `.dat` file. +/// Header: 07 08 V1 08 07 (6B) + aes_size LE (4B) + xor_size LE (4B) + 0x01 (1B) = 15B +/// Then AES-ECB encrypted payload with fixed key, then raw, then XOR tail. +fn make_v1_dat() -> Vec { + use aes::cipher::{BlockEncrypt, KeyInit}; + use aes::Aes128; + + let key = b"cfcd208495d565ef"; // md5("0")[:16] + let cipher = Aes128::new(key.into()); + + // Plaintext: JPEG header (16 bytes = 1 AES block) with PKCS7 padding + let mut block1 = [ + 0xFFu8, 0xD8, 0xFF, 0xE0, 0x00, 0x10, 0x4A, 0x46, 0x49, 0x46, 0x00, 0x01, 0x01, 0x00, 0x00, + 0x01, + ]; + let mut block2 = [16u8; 16]; // full PKCS7 padding block + + cipher.encrypt_block((&mut block1).into()); + cipher.encrypt_block((&mut block2).into()); + + let aes_size: u32 = 16; // original plaintext size + let xor_size: u32 = 0; + + let mut dat = Vec::new(); + dat.extend_from_slice(b"\x07\x08V1\x08\x07"); // 6B signature + dat.extend_from_slice(&aes_size.to_le_bytes()); // 4B aes_size + dat.extend_from_slice(&xor_size.to_le_bytes()); // 4B xor_size + dat.push(0x01); // 1B padding + dat.extend_from_slice(&block1); // AES ciphertext block 1 + dat.extend_from_slice(&block2); // AES ciphertext block 2 (PKCS7 padding) + dat +} + +/// Build a V2-encrypted `.dat` file with known AES key and XOR tail. +fn make_v2_dat(aes_key: &[u8; 16], xor_key: u8) -> Vec { + use aes::cipher::{BlockEncrypt, KeyInit}; + use aes::Aes128; + + let cipher = Aes128::new(aes_key.into()); + + // Plaintext: PNG header (16 bytes = 1 block) + let mut block1 = [ + 0x89u8, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A, 0x00, 0x00, 0x00, 0x0D, 0x49, 0x48, 0x44, + 0x52, + ]; + let mut block2 = [16u8; 16]; // PKCS7 padding block + + cipher.encrypt_block((&mut block1).into()); + cipher.encrypt_block((&mut block2).into()); + + let aes_size: u32 = 16; + // Tail: 4 bytes XOR-encrypted + let xor_plain = [0x00u8, 0x00, 0x00, 0x00]; + let xor_enc: Vec = xor_plain.iter().map(|b| b ^ xor_key).collect(); + let xor_size: u32 = xor_enc.len() as u32; + + // Middle raw section: 8 bytes of unencrypted data + let raw_data = [0xAA, 0xBB, 0xCC, 0xDD, 0x11, 0x22, 0x33, 0x44]; + + let mut dat = Vec::new(); + dat.extend_from_slice(b"\x07\x08V2\x08\x07"); // 6B signature + dat.extend_from_slice(&aes_size.to_le_bytes()); // 4B aes_size + dat.extend_from_slice(&xor_size.to_le_bytes()); // 4B xor_size + dat.push(0x01); // 1B padding + dat.extend_from_slice(&block1); // AES ciphertext + dat.extend_from_slice(&block2); // PKCS7 padding ciphertext + dat.extend_from_slice(&raw_data); // unencrypted middle + dat.extend_from_slice(&xor_enc); // XOR tail + dat +} + +// ── XOR tests ──────────────────────────────────────────────────────── + +#[test] +fn xor_decrypt_jpg() { + let dat = make_xor_dat(0xAB); + let result = wx_media::decrypt_dat(&dat, &DatDecryptOptions::default()).unwrap(); + assert_eq!(result.format, DatFormat::Xor); + assert_eq!(result.ext, "jpg"); + assert_eq!(&result.data[..3], &[0xFF, 0xD8, 0xFF]); +} + +#[test] +fn xor_decrypt_png() { + let plain = [0x89u8, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A]; + let key = 0x55u8; + let dat: Vec = plain.iter().map(|b| b ^ key).collect(); + let result = wx_media::decrypt_dat(&dat, &DatDecryptOptions::default()).unwrap(); + assert_eq!(result.format, DatFormat::Xor); + assert_eq!(result.ext, "png"); +} + +#[test] +fn xor_detect_gif() { + let plain = [0x47u8, 0x49, 0x46, 0x38, 0x39, 0x61]; // GIF89a + let key = 0x12u8; + let dat: Vec = plain.iter().map(|b| b ^ key).collect(); + let result = wx_media::decrypt_dat(&dat, &DatDecryptOptions::default()).unwrap(); + assert_eq!(result.ext, "gif"); +} + +#[test] +fn xor_detect_wxgf() { + let plain = [0x77u8, 0x78, 0x67, 0x66, 0x00, 0x01]; + let key = 0x24u8; + let dat: Vec = plain.iter().map(|b| b ^ key).collect(); + let result = wx_media::decrypt_dat(&dat, &DatDecryptOptions::default()).unwrap(); + assert_eq!(result.ext, "wxgf"); + assert_eq!(&result.data[..4], b"wxgf"); +} + +#[test] +fn xor_header_too_short() { + let dat = vec![0xAB, 0xCD]; // Only 2 bytes, not enough for reliable detection + let result = wx_media::decrypt_dat(&dat, &DatDecryptOptions::default()); + // Should fail since no 3+ byte magic matches + assert!(result.is_err()); +} + +// ── V1 tests ───────────────────────────────────────────────────────── + +#[test] +fn v1_decrypt_success() { + let dat = make_v1_dat(); + let result = wx_media::decrypt_dat(&dat, &DatDecryptOptions::default()).unwrap(); + assert_eq!(result.format, DatFormat::V1); + assert_eq!(result.ext, "jpg"); + assert_eq!(&result.data[..3], &[0xFF, 0xD8, 0xFF]); +} + +// ── V2 tests ───────────────────────────────────────────────────────── + +#[test] +fn v2_decrypt_success() { + let aes_key = b"abcdefghijklmnop"; + let xor_key = 0x37u8; + let dat = make_v2_dat(aes_key, xor_key); + + let opts = DatDecryptOptions { + v2_aes_key: Some(*aes_key), + xor_key: Some(xor_key), + }; + let result = wx_media::decrypt_dat(&dat, &opts).unwrap(); + assert_eq!(result.format, DatFormat::V2); + assert_eq!(result.ext, "png"); + assert_eq!(&result.data[..4], &[0x89, 0x50, 0x4E, 0x47]); + // Verify raw middle section preserved + assert_eq!( + &result.data[16..24], + &[0xAA, 0xBB, 0xCC, 0xDD, 0x11, 0x22, 0x33, 0x44] + ); + // Verify XOR tail decrypted + assert_eq!(&result.data[24..28], &[0x00, 0x00, 0x00, 0x00]); +} + +#[test] +fn v2_missing_key() { + let aes_key = b"abcdefghijklmnop"; + let dat = make_v2_dat(aes_key, 0x37); + let result = wx_media::decrypt_dat(&dat, &DatDecryptOptions::default()); + assert!(matches!(result, Err(MediaError::MissingV2Key))); +} + +#[test] +fn v2_wrong_key() { + let aes_key = b"abcdefghijklmnop"; + let dat = make_v2_dat(aes_key, 0x37); + let wrong_key = b"0000000000000000"; + let opts = DatDecryptOptions { + v2_aes_key: Some(*wrong_key), + xor_key: Some(0x37), + }; + // Wrong key → PKCS7 validation fails → AesDecryptFailed error + let result = wx_media::decrypt_dat(&dat, &opts); + assert!(matches!(result, Err(MediaError::AesDecryptFailed { .. }))); +} + +// ── Edge cases ─────────────────────────────────────────────────────── + +#[test] +fn empty_input() { + let result = wx_media::decrypt_dat(&[], &DatDecryptOptions::default()); + assert!(result.is_err()); +} + +#[test] +fn truncated_v2_header() { + // V2 signature but truncated before aes_size field + let dat = b"\x07\x08V2\x08\x07\x00\x04"; + let opts = DatDecryptOptions { + v2_aes_key: Some(*b"abcdefghijklmnop"), + xor_key: Some(0x37), + }; + let result = wx_media::decrypt_dat(dat, &opts); + assert!(result.is_err()); +} + +// ── detect_dat_format tests ────────────────────────────────────────── + +#[test] +fn detect_format_xor() { + let dat = make_xor_dat(0xAB); + assert_eq!(wx_media::detect_dat_format(&dat), None); + // XOR format is detected by elimination (no V1/V2 signature), returns None for the enum +} + +#[test] +fn detect_format_v1() { + let dat = make_v1_dat(); + assert_eq!(wx_media::detect_dat_format(&dat), Some(DatFormat::V1)); +} + +#[test] +fn detect_format_v2() { + let aes_key = b"abcdefghijklmnop"; + let dat = make_v2_dat(aes_key, 0x37); + assert_eq!(wx_media::detect_dat_format(&dat), Some(DatFormat::V2)); +} + +// ── detect_image_type tests ────────────────────────────────────────── + +#[test] +fn detect_image_types() { + use wx_media::ImageType; + + assert_eq!( + wx_media::detect_image_type(&[0xFF, 0xD8, 0xFF]), + ImageType::Jpg + ); + assert_eq!( + wx_media::detect_image_type(&[0x89, 0x50, 0x4E, 0x47]), + ImageType::Png + ); + assert_eq!( + wx_media::detect_image_type(&[0x47, 0x49, 0x46, 0x38]), + ImageType::Gif + ); + assert_eq!( + wx_media::detect_image_type(&[ + 0x52, 0x49, 0x46, 0x46, 0x00, 0x00, 0x00, 0x00, 0x57, 0x45, 0x42, 0x50 + ]), + ImageType::Webp + ); + assert_eq!( + wx_media::detect_image_type(&[0x49, 0x49, 0x2A, 0x00]), + ImageType::Tif + ); + assert_eq!( + wx_media::detect_image_type(&[0x77, 0x78, 0x67, 0x66]), + ImageType::Wxgf + ); + assert_eq!( + wx_media::detect_image_type(&[0x00, 0x00, 0x00, 0x00]), + ImageType::Unknown + ); +} diff --git a/crates/wx-media/tests/fallback.rs b/crates/wx-media/tests/fallback.rs new file mode 100644 index 0000000..571dcaa --- /dev/null +++ b/crates/wx-media/tests/fallback.rs @@ -0,0 +1,137 @@ +use std::fs; +use tempfile::TempDir; + +// --- find_video_by_md5 tests --- + +#[test] +fn test_video_found_in_hint_month() { + let dir = TempDir::new().unwrap(); + let month_dir = dir.path().join("2024-03"); + fs::create_dir(&month_dir).unwrap(); + fs::write(month_dir.join("abc123.mp4"), b"video").unwrap(); + + let result = wx_media::find_video_by_md5(dir.path(), "abc123", "2024-03"); + assert_eq!(result, Some(month_dir.join("abc123.mp4"))); +} + +#[test] +fn test_video_found_in_other_month() { + let dir = TempDir::new().unwrap(); + let month_dir = dir.path().join("2024-05"); + fs::create_dir(&month_dir).unwrap(); + fs::write(month_dir.join("abc123.mp4"), b"video").unwrap(); + + let result = wx_media::find_video_by_md5(dir.path(), "abc123", "2024-03"); + assert_eq!(result, Some(month_dir.join("abc123.mp4"))); +} + +#[test] +fn test_video_not_found() { + let dir = TempDir::new().unwrap(); + let result = wx_media::find_video_by_md5(dir.path(), "abc123", "2024-03"); + assert_eq!(result, None); +} + +#[test] +fn test_video_ignores_non_month_dirs() { + let dir = TempDir::new().unwrap(); + let other_dir = dir.path().join("other"); + fs::create_dir(&other_dir).unwrap(); + fs::write(other_dir.join("abc123.mp4"), b"video").unwrap(); + + let result = wx_media::find_video_by_md5(dir.path(), "abc123", "2024-03"); + assert_eq!(result, None); +} + +// --- find_file_by_name tests --- + +#[test] +fn test_file_found_in_month() { + let dir = TempDir::new().unwrap(); + let month_dir = dir.path().join("2024-03"); + fs::create_dir(&month_dir).unwrap(); + fs::write(month_dir.join("report.pdf"), b"data").unwrap(); + + let result = wx_media::find_file_by_name(dir.path(), "report.pdf", "2024-03"); + assert_eq!(result, Some(month_dir.join("report.pdf"))); +} + +#[test] +fn test_file_not_found_wrong_month() { + let dir = TempDir::new().unwrap(); + let month_dir = dir.path().join("2024-05"); + fs::create_dir(&month_dir).unwrap(); + fs::write(month_dir.join("report.pdf"), b"data").unwrap(); + + let result = wx_media::find_file_by_name(dir.path(), "report.pdf", "2024-03"); + assert_eq!(result, None); +} + +#[test] +fn test_file_not_found_empty() { + let dir = TempDir::new().unwrap(); + let result = wx_media::find_file_by_name(dir.path(), "report.pdf", "2024-03"); + assert_eq!(result, None); +} + +// --- Edge case tests --- + +#[test] +fn test_video_hint_month_preferred() { + let dir = TempDir::new().unwrap(); + // Create same md5 file in two months + let hint_dir = dir.path().join("2024-03"); + let other_dir = dir.path().join("2024-05"); + fs::create_dir(&hint_dir).unwrap(); + fs::create_dir(&other_dir).unwrap(); + fs::write(hint_dir.join("abc123.mp4"), b"video-hint").unwrap(); + fs::write(other_dir.join("abc123.mp4"), b"video-other").unwrap(); + + let result = wx_media::find_video_by_md5(dir.path(), "abc123", "2024-03"); + assert_eq!(result, Some(hint_dir.join("abc123.mp4"))); +} + +#[test] +fn test_file_with_special_chars() { + let dir = TempDir::new().unwrap(); + let month_dir = dir.path().join("2024-03"); + fs::create_dir(&month_dir).unwrap(); + + // CJK characters and spaces in filename + let name = "\u{4F1A}\u{8BAE}\u{8BB0}\u{5F55} 2024.pdf"; + fs::write(month_dir.join(name), b"data").unwrap(); + + let result = wx_media::find_file_by_name(dir.path(), name, "2024-03"); + assert_eq!(result, Some(month_dir.join(name))); +} + +#[test] +fn test_file_path_traversal_blocked() { + let dir = TempDir::new().unwrap(); + let month_dir = dir.path().join("2024-03"); + fs::create_dir(&month_dir).unwrap(); + + // Create a file that a naive path join would reach via traversal + let escape_target = dir.path().join("passwd"); + fs::write(&escape_target, b"secret").unwrap(); + + // Path traversal attempt: "../passwd" relative to file_dir/2024-03/ would reach file_dir/passwd + let result = wx_media::find_file_by_name(dir.path(), "../passwd", "2024-03"); + assert_eq!(result, None, "path traversal should be blocked"); + + // Also test deeper traversal + let result = wx_media::find_file_by_name(dir.path(), "../../etc/passwd", "2024-03"); + assert_eq!(result, None, "deep path traversal should be blocked"); +} + +#[test] +fn test_video_skips_non_mp4() { + let dir = TempDir::new().unwrap(); + let month_dir = dir.path().join("2024-03"); + fs::create_dir(&month_dir).unwrap(); + // Create .avi instead of .mp4 + fs::write(month_dir.join("abc123.avi"), b"video").unwrap(); + + let result = wx_media::find_video_by_md5(dir.path(), "abc123", "2024-03"); + assert_eq!(result, None); +} diff --git a/crates/wx-media/tests/ffmpeg-pipe.rs b/crates/wx-media/tests/ffmpeg-pipe.rs new file mode 100644 index 0000000..7fab5c3 --- /dev/null +++ b/crates/wx-media/tests/ffmpeg-pipe.rs @@ -0,0 +1,127 @@ +use std::process::{Command, Output, Stdio}; +use std::sync::Mutex; +use std::time::{Duration, Instant}; + +static CHILD_LOCK: Mutex<()> = Mutex::new(()); + +#[test] +fn run_ffmpeg_drains_stdout_while_streaming_stdin() { + let _guard = CHILD_LOCK.lock().unwrap(); + + let output = run_child("ffmpeg", Duration::from_secs(3)); + assert!( + output.status.success(), + "child failed: stdout={}\nstderr={}", + String::from_utf8_lossy(&output.stdout), + String::from_utf8_lossy(&output.stderr) + ); + assert!(String::from_utf8_lossy(&output.stdout).contains("ffmpeg-ok")); +} + +#[test] +fn run_ffprobe_drains_stdout_while_streaming_stdin() { + let _guard = CHILD_LOCK.lock().unwrap(); + + let output = run_child("ffprobe", Duration::from_secs(3)); + assert!( + output.status.success(), + "child failed: stdout={}\nstderr={}", + String::from_utf8_lossy(&output.stdout), + String::from_utf8_lossy(&output.stderr) + ); + assert!(String::from_utf8_lossy(&output.stdout).contains("ffprobe-ok")); +} + +#[test] +fn ffmpeg_pipe_child_mode() { + let Ok(mode) = std::env::var("WECHAT_MEDIA_PIPE_CHILD_MODE") else { + return; + }; + + let temp = tempfile::TempDir::new().unwrap(); + let tool_path = temp.path().join("fake-ffmpeg.py"); + std::fs::write(&tool_path, fake_ffmpeg_script()).unwrap(); + + #[cfg(unix)] + { + use std::os::unix::fs::PermissionsExt; + let mut perms = std::fs::metadata(&tool_path).unwrap().permissions(); + perms.set_mode(0o755); + std::fs::set_permissions(&tool_path, perms).unwrap(); + } + + unsafe { + std::env::set_var("FFMPEG_PATH", &tool_path); + std::env::set_var("FFPROBE_PATH", &tool_path); + } + wx_media::reset_ffmpeg_cache(); + + let input = vec![b'i'; 1024 * 1024]; + match mode.as_str() { + "ffmpeg" => { + let output = wx_media::run_ffmpeg(&input, &["-i", "pipe:0", "pipe:1"]).unwrap(); + assert_eq!(output.len(), 262_144); + println!("ffmpeg-ok {}", output.len()); + } + "ffprobe" => { + let output = wx_media::run_ffprobe(&input, &["-i", "pipe:0"]).unwrap(); + assert_eq!(output.len(), 262_144); + println!("ffprobe-ok {}", output.len()); + } + other => panic!("unexpected child mode: {other}"), + } +} + +fn run_child(mode: &str, timeout: Duration) -> Output { + let current_exe = std::env::current_exe().unwrap(); + let mut child = Command::new(current_exe) + .arg("--exact") + .arg("ffmpeg_pipe_child_mode") + .arg("--nocapture") + .arg("--test-threads=1") + .env("WECHAT_MEDIA_PIPE_CHILD_MODE", mode) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .spawn() + .unwrap(); + + let start = Instant::now(); + loop { + if child.try_wait().unwrap().is_some() { + return child.wait_with_output().unwrap(); + } + + if start.elapsed() >= timeout { + let _ = child.kill(); + let output = child.wait_with_output().unwrap(); + panic!( + "child timed out after {:?}: stdout={}\nstderr={}", + timeout, + String::from_utf8_lossy(&output.stdout), + String::from_utf8_lossy(&output.stderr) + ); + } + + std::thread::sleep(Duration::from_millis(25)); + } +} + +fn fake_ffmpeg_script() -> &'static str { + r#"#!/usr/bin/env python3 +import os +import sys + +if "-version" in sys.argv: + sys.stdout.write("fake ffmpeg 1.0\n") + sys.exit(0) + +chunk = b"x" * 4096 +remaining = 262144 +while remaining > 0: + piece = chunk if remaining >= len(chunk) else b"x" * remaining + os.write(1, piece) + remaining -= len(piece) + +sys.stdin.buffer.read() +"# +} diff --git a/crates/wx-media/tests/hardlink.rs b/crates/wx-media/tests/hardlink.rs new file mode 100644 index 0000000..ada74dc --- /dev/null +++ b/crates/wx-media/tests/hardlink.rs @@ -0,0 +1,242 @@ +use std::path::Path; +use tempfile::TempDir; +use wx_media::{self, MediaError}; + +/// Create a hardlink.db with v3-style tables. +fn create_hardlink_db_v3(path: &Path, entries: &[(&str, &str, &str, i64, i64, i64, i64)]) { + // entries: (table_prefix, md5, file_name, file_size, modify_time, dir1_rowid, dir2_rowid) + let conn = rusqlite::Connection::open(path).unwrap(); + conn.execute_batch( + "CREATE TABLE dir2id (rowid INTEGER PRIMARY KEY, username TEXT); + CREATE TABLE image_hardlink_info_v3 ( + md5 TEXT, file_name TEXT, file_size INTEGER, modify_time INTEGER, + dir1 INTEGER, dir2 INTEGER + ); + CREATE TABLE video_hardlink_info_v3 ( + md5 TEXT, file_name TEXT, file_size INTEGER, modify_time INTEGER, + dir1 INTEGER, dir2 INTEGER + ); + CREATE TABLE file_hardlink_info_v3 ( + md5 TEXT, file_name TEXT, file_size INTEGER, modify_time INTEGER, + dir1 INTEGER, dir2 INTEGER + );", + ) + .unwrap(); + + // Insert dir2id entries + conn.execute( + "INSERT INTO dir2id (rowid, username) VALUES (1, 'wxid_alice')", + [], + ) + .unwrap(); + conn.execute( + "INSERT INTO dir2id (rowid, username) VALUES (2, '2026-03')", + [], + ) + .unwrap(); + + for &(table_prefix, md5, file_name, file_size, modify_time, dir1, dir2) in entries { + let table = format!("{}_hardlink_info_v3", table_prefix); + conn.execute( + &format!( + "INSERT INTO {} (md5, file_name, file_size, modify_time, dir1, dir2) VALUES (?, ?, ?, ?, ?, ?)", + table + ), + rusqlite::params![md5, file_name, file_size, modify_time, dir1, dir2], + ) + .unwrap(); + } +} + +/// Create a hardlink.db with only v4-style tables. +fn create_hardlink_db_v4(path: &Path) { + let conn = rusqlite::Connection::open(path).unwrap(); + conn.execute_batch( + "CREATE TABLE dir2id (rowid INTEGER PRIMARY KEY, username TEXT); + CREATE TABLE image_hardlink_info_v4 ( + md5 TEXT, file_name TEXT, file_size INTEGER, modify_time INTEGER, + dir1 INTEGER, dir2 INTEGER + ); + CREATE TABLE video_hardlink_info_v4 ( + md5 TEXT, file_name TEXT, file_size INTEGER, modify_time INTEGER, + dir1 INTEGER, dir2 INTEGER + ); + CREATE TABLE file_hardlink_info_v4 ( + md5 TEXT, file_name TEXT, file_size INTEGER, modify_time INTEGER, + dir1 INTEGER, dir2 INTEGER + );", + ) + .unwrap(); + + conn.execute( + "INSERT INTO dir2id (rowid, username) VALUES (1, 'wxid_bob')", + [], + ) + .unwrap(); + conn.execute( + "INSERT INTO dir2id (rowid, username) VALUES (2, '2026-01')", + [], + ) + .unwrap(); + conn.execute( + "INSERT INTO image_hardlink_info_v4 (md5, file_name, file_size, modify_time, dir1, dir2) \ + VALUES ('abc123', 'abc123_h.dat', 4096, 1709000000, 1, 2)", + [], + ) + .unwrap(); +} + +// ── Tests ──────────────────────────────────────────────────────────── + +#[test] +fn query_image_v3_by_md5() { + let tmp = TempDir::new().unwrap(); + let db_path = tmp.path().join("hardlink.db"); + create_hardlink_db_v3( + &db_path, + &[ + ("image", "aabbcc", "aabbcc_t.dat", 1024, 1709000000, 1, 2), + ("image", "aabbcc", "aabbcc_h.dat", 8192, 1709000001, 1, 2), + ], + ); + + let results = wx_media::query_hardlink(&db_path, "image", "aabbcc").unwrap(); + assert_eq!(results.len(), 2); + // Should prefer _h.dat (high quality) — returned first + assert_eq!(results[0].file_name, "aabbcc_h.dat"); + assert_eq!(results[0].dir1, "wxid_alice"); + assert_eq!(results[0].dir2, "2026-03"); +} + +#[test] +fn query_video_v3() { + let tmp = TempDir::new().unwrap(); + let db_path = tmp.path().join("hardlink.db"); + create_hardlink_db_v3( + &db_path, + &[("video", "vid001", "vid001.mp4", 102400, 1709000000, 1, 2)], + ); + + let results = wx_media::query_hardlink(&db_path, "video", "vid001").unwrap(); + assert_eq!(results.len(), 1); + assert_eq!(results[0].media_type, "video"); + assert_eq!(results[0].file_size, 102400); +} + +#[test] +fn query_file_by_name_prefix() { + let tmp = TempDir::new().unwrap(); + let db_path = tmp.path().join("hardlink.db"); + create_hardlink_db_v3( + &db_path, + &[( + "file", + "doc999", + "doc999_report.pdf", + 51200, + 1709000000, + 1, + 2, + )], + ); + + let results = wx_media::query_hardlink(&db_path, "file", "doc999").unwrap(); + assert_eq!(results.len(), 1); + assert_eq!(results[0].file_name, "doc999_report.pdf"); +} + +#[test] +fn query_v4_fallback() { + let tmp = TempDir::new().unwrap(); + let db_path = tmp.path().join("hardlink.db"); + create_hardlink_db_v4(&db_path); + + let results = wx_media::query_hardlink(&db_path, "image", "abc123").unwrap(); + assert_eq!(results.len(), 1); + assert_eq!(results[0].dir1, "wxid_bob"); +} + +#[test] +fn query_not_found() { + let tmp = TempDir::new().unwrap(); + let db_path = tmp.path().join("hardlink.db"); + create_hardlink_db_v3(&db_path, &[]); + + let result = wx_media::query_hardlink(&db_path, "image", "nonexistent"); + assert!(matches!(result, Err(MediaError::LookupMiss(_)))); +} + +#[test] +fn query_image_prefers_h_dat() { + let tmp = TempDir::new().unwrap(); + let db_path = tmp.path().join("hardlink.db"); + create_hardlink_db_v3( + &db_path, + &[ + ("image", "xyz", "xyz_t.dat", 512, 1709000000, 1, 2), + ("image", "xyz", "xyz_h.dat", 4096, 1709000001, 1, 2), + ("image", "xyz", "xyz.dat", 2048, 1709000002, 1, 2), + ], + ); + + let results = wx_media::query_hardlink(&db_path, "image", "xyz").unwrap(); + // _h.dat should be first + assert_eq!(results[0].file_name, "xyz_h.dat"); +} + +#[test] +fn hardlink_query_with_conn_prefers_v3_and_h_image() { + let tmp = TempDir::new().unwrap(); + let db_path = tmp.path().join("hardlink.db"); + create_hardlink_db_v3( + &db_path, + &[ + ("image", "img001", "img001.dat", 2048, 1709000002, 1, 2), + ("image", "img001", "img001_h.dat", 4096, 1709000001, 1, 2), + ], + ); + let conn = rusqlite::Connection::open(&db_path).unwrap(); + + let results = wx_media::query_hardlink_with_conn(&conn, "image", "img001").unwrap(); + assert_eq!(results.len(), 2); + assert_eq!(results[0].file_name, "img001_h.dat"); +} + +#[test] +fn hardlink_query_with_conn_matches_video_and_file_prefixes() { + let tmp = TempDir::new().unwrap(); + let db_path = tmp.path().join("hardlink.db"); + create_hardlink_db_v3( + &db_path, + &[ + ("video", "vid002", "vid002.mp4", 102400, 1709000000, 1, 2), + ( + "file", + "doc123", + "doc123_notes.txt", + 51200, + 1709000100, + 1, + 2, + ), + ], + ); + let conn = rusqlite::Connection::open(&db_path).unwrap(); + + let video = wx_media::query_hardlink_with_conn(&conn, "video", "vid002").unwrap(); + assert_eq!(video[0].file_name, "vid002.mp4"); + + let file = wx_media::query_hardlink_with_conn(&conn, "file", "doc123").unwrap(); + assert_eq!(file[0].file_name, "doc123_notes.txt"); +} + +#[test] +fn hardlink_query_with_conn_rejects_unsupported_media_type() { + let tmp = TempDir::new().unwrap(); + let db_path = tmp.path().join("hardlink.db"); + create_hardlink_db_v3(&db_path, &[]); + let conn = rusqlite::Connection::open(&db_path).unwrap(); + + let result = wx_media::query_hardlink_with_conn(&conn, "voice", "abc"); + assert!(matches!(result, Err(MediaError::InvalidFormat { .. }))); +} diff --git a/crates/wx-media/tests/image-resolver.rs b/crates/wx-media/tests/image-resolver.rs new file mode 100644 index 0000000..bdd7e92 --- /dev/null +++ b/crates/wx-media/tests/image-resolver.rs @@ -0,0 +1,185 @@ +use std::fs; +use std::path::Path; +use tempfile::TempDir; +use wx_media::MediaError; + +/// Create a minimal `message_resource.db` with a `MessageResourceInfo` table. +fn create_message_resource_db(path: &Path, rows: &[(i64, &[u8])]) { + let conn = rusqlite::Connection::open(path).unwrap(); + conn.execute_batch( + "CREATE TABLE MessageResourceInfo ( + local_id INTEGER PRIMARY KEY, + packed_info BLOB + );", + ) + .unwrap(); + let mut stmt = conn + .prepare("INSERT INTO MessageResourceInfo (local_id, packed_info) VALUES (?, ?)") + .unwrap(); + for &(local_id, blob) in rows { + stmt.execute(rusqlite::params![local_id, blob]).unwrap(); + } +} + +/// Build a packed_info blob containing a protobuf-style MD5 marker. +/// Format: prefix + `\x12\x22\x0a\x20` + 32-byte hex MD5 +fn make_packed_info(md5_hex: &str) -> Vec { + let mut blob = vec![0x0A, 0x10]; // some prefix bytes + blob.extend_from_slice(b"\x12\x22\x0a\x20"); + blob.extend_from_slice(md5_hex.as_bytes()); + blob.extend_from_slice(&[0x18, 0x01]); // trailing protobuf + blob +} + +/// Create a fake attach directory structure with .dat files. +fn create_attach_dir( + base: &Path, + username_hash: &str, + month: &str, + file_md5: &str, + suffixes: &[&str], +) { + let img_dir = base + .join("msg") + .join("attach") + .join(username_hash) + .join(month) + .join("Img"); + fs::create_dir_all(&img_dir).unwrap(); + for suffix in suffixes { + let name = format!("{}{}.dat", file_md5, suffix); + fs::write(img_dir.join(name), b"fake dat content").unwrap(); + } +} + +// ── packed_info parsing tests ──────────────────────────────────────── + +#[test] +fn extract_md5_protobuf_marker() { + let md5 = "d41d8cd98f00b204e9800998ecf8427e"; + let blob = make_packed_info(md5); + let result = wx_media::extract_md5_from_packed_info(&blob).unwrap(); + assert_eq!(result, md5); +} + +#[test] +fn extract_md5_fallback_hex_scan() { + // No protobuf marker, but contains 32 contiguous hex chars + let md5 = "abcdef0123456789abcdef0123456789"; + let mut blob = vec![0x00, 0x01, 0x02]; + blob.extend_from_slice(md5.as_bytes()); + blob.extend_from_slice(&[0xFF, 0xFE]); + let result = wx_media::extract_md5_from_packed_info(&blob).unwrap(); + assert_eq!(result, md5); +} + +#[test] +fn extract_md5_no_match() { + let blob = vec![0x00; 10]; + let result = wx_media::extract_md5_from_packed_info(&blob); + assert!(result.is_none()); +} + +#[test] +fn extract_md5_empty() { + assert!(wx_media::extract_md5_from_packed_info(&[]).is_none()); +} + +// ── resolve_image_by_md5 tests ─────────────────────────────────────── + +#[test] +fn resolve_image_by_md5_success() { + let tmp = TempDir::new().unwrap(); + let base = tmp.path(); + + let username = "testuser"; + let username_hash = format!("{:x}", md5::compute(username.as_bytes())); + let file_md5 = "d41d8cd98f00b204e9800998ecf8427e"; + + // Create attach dir with dat files (reuse helper) + create_attach_dir(base, &username_hash, "2026-03", file_md5, &["", "_t", "_h"]); + + let result = + wx_media::resolve_image_by_md5(username, &base.join("msg").join("attach"), file_md5) + .unwrap(); + + assert_eq!(result.file_md5, file_md5); + assert_eq!(result.candidates.len(), 3); + let rec = result.recommended.unwrap(); + let name = rec.file_name().unwrap().to_str().unwrap(); + assert_eq!(name, format!("{}_h.dat", file_md5)); +} + +#[test] +fn resolve_image_by_md5_no_dat_files() { + let tmp = TempDir::new().unwrap(); + let attach_dir = tmp.path().join("msg").join("attach"); + fs::create_dir_all(&attach_dir).unwrap(); + + let result = + wx_media::resolve_image_by_md5("user", &attach_dir, "deadbeef12345678deadbeef12345678"); + assert!(matches!(result, Err(MediaError::NoDatFiles { .. }))); +} + +// ── image resolver integration tests ───────────────────────────────── + +#[test] +fn resolve_image_path_success() { + let tmp = TempDir::new().unwrap(); + let base = tmp.path(); + + let username = "testuser"; + let username_hash = format!("{:x}", md5::compute(username.as_bytes())); + let file_md5 = "d41d8cd98f00b204e9800998ecf8427e"; + + // Create message_resource.db + let db_path = base.join("db_storage").join("message"); + fs::create_dir_all(&db_path).unwrap(); + let resource_db = db_path.join("message_resource.db"); + let packed_info = make_packed_info(file_md5); + create_message_resource_db(&resource_db, &[(42, &packed_info)]); + + // Create attach dir with dat files + create_attach_dir(base, &username_hash, "2026-03", file_md5, &["", "_t", "_h"]); + + let result = wx_media::resolve_image( + &resource_db, + 42, // local_id + username, + &base.join("msg").join("attach"), + ) + .unwrap(); + + assert_eq!(result.file_md5, file_md5); + assert_eq!(result.candidates.len(), 3); + // Recommended should prefer _h over the plain md5.dat variant. + let rec = result.recommended.unwrap(); + let name = rec.file_name().unwrap().to_str().unwrap(); + assert_eq!(name, format!("{}_h.dat", file_md5)); +} + +#[test] +fn resolve_image_local_id_not_found() { + let tmp = TempDir::new().unwrap(); + let db_path = tmp.path().join("message_resource.db"); + create_message_resource_db(&db_path, &[]); + + let result = wx_media::resolve_image(&db_path, 999, "user", tmp.path()); + assert!(matches!(result, Err(MediaError::NotFound(_)))); +} + +#[test] +fn resolve_image_no_dat_files() { + let tmp = TempDir::new().unwrap(); + let db_path = tmp.path().join("message_resource.db"); + let file_md5 = "d41d8cd98f00b204e9800998ecf8427e"; + let packed_info = make_packed_info(file_md5); + create_message_resource_db(&db_path, &[(1, &packed_info)]); + + // Don't create any attach dirs + let attach_dir = tmp.path().join("msg").join("attach"); + fs::create_dir_all(&attach_dir).unwrap(); + + let result = wx_media::resolve_image(&db_path, 1, "user", &attach_dir); + assert!(matches!(result, Err(MediaError::NoDatFiles { .. }))); +} diff --git a/crates/wx-media/tests/isaac64.rs b/crates/wx-media/tests/isaac64.rs new file mode 100644 index 0000000..2d6ad5a --- /dev/null +++ b/crates/wx-media/tests/isaac64.rs @@ -0,0 +1,72 @@ +use wx_media::Isaac64; + +// Test vectors generated from correct BigInt-safe ISAAC-64 implementation. +// NOTE: The CipherTalk TS reference has a Number() precision bug in +// `mm[Number(x >> 3n) & 255]` that loses precision for 64-bit values > 2^53. +// Our Rust implementation matches the standard ISAAC-64 algorithm and the +// WeChat WASM module's correct 64-bit behavior. + +#[test] +fn test_seed_0_next_u64() { + let mut rng = Isaac64::new(0); + assert_eq!(rng.next_u64(), 0x673f4a26e311355a); + assert_eq!(rng.next_u64(), 0x80b58ce58cf970fd); + assert_eq!(rng.next_u64(), 0xdfdbe39c37e83fb2); + assert_eq!(rng.next_u64(), 0x7b81201809b6bbf7); + assert_eq!(rng.next_u64(), 0x7fa9c030f3ff9cfc); +} + +#[test] +fn test_seed_12345_next_u64() { + let mut rng = Isaac64::new(12345); + assert_eq!(rng.next_u64(), 0x22286913bb698089); + assert_eq!(rng.next_u64(), 0x276beba6b7d70db1); + assert_eq!(rng.next_u64(), 0x7c228f4bc32b9af1); + assert_eq!(rng.next_u64(), 0x5f46814fb1b21e59); + assert_eq!(rng.next_u64(), 0x977adb409ed5f786); +} + +#[test] +fn test_keystream_16_bytes() { + let mut rng = Isaac64::new(0); + let ks = rng.keystream(16); + assert_eq!(ks.len(), 16); + assert_eq!(hex::encode(&ks), "673f4a26e311355a80b58ce58cf970fd"); +} + +#[test] +fn test_keystream_9_bytes_non_aligned() { + let mut rng = Isaac64::new(0); + let ks = rng.keystream(9); + assert_eq!(ks.len(), 9); + // First 8 bytes: BE of next_u64[0] = 0x673f4a26e311355a + // Next 1 byte: first BE byte of next_u64[1] = 0x80 + assert_eq!(hex::encode(&ks), "673f4a26e311355a80"); +} + +#[test] +fn test_keystream_0_bytes() { + let mut rng = Isaac64::new(0); + let ks = rng.keystream(0); + assert!(ks.is_empty()); +} + +#[test] +fn test_keystream_1_byte() { + let mut rng = Isaac64::new(0); + let ks = rng.keystream(1); + assert_eq!(ks.len(), 1); + // First BE byte of 0x673f4a26e311355a = 0x67 + assert_eq!(ks[0], 0x67); +} + +#[test] +fn test_generate_exhausts_and_refills() { + let mut rng = Isaac64::new(0); + for _ in 0..256 { + rng.next_u64(); + } + // The 257th call triggers generate() and still works + let val = rng.next_u64(); + assert_ne!(val, 0); +} diff --git a/crates/wx-media/tests/key-derivation.rs b/crates/wx-media/tests/key-derivation.rs new file mode 100644 index 0000000..59c6591 --- /dev/null +++ b/crates/wx-media/tests/key-derivation.rs @@ -0,0 +1,164 @@ +use wx_media::key::{derive_v2_key_from_dir, read_uin}; +use wx_media::{derive_v2_aes_key, extract_wxid}; + +// ── V2 key derivation formula ────────────────────────────────────── + +#[test] +fn derive_v2_key_known_account() { + // Verified: MD5("1234567890wxid_example123abc") → hex[:16] = d5b763f94307be38 + let key = derive_v2_aes_key("1234567890", "wxid_example123abc"); + assert_eq!(&key, b"d5b763f94307be38"); +} + +#[test] +fn v1_fixed_key_same_formula() { + // V1 fixed key = MD5("0")[:16] = cfcd208495d565ef + let key = derive_v2_aes_key("", "0"); + assert_eq!(&key, b"cfcd208495d565ef"); +} + +// ── WXID suffix stripping (shared via wx-keychain AccountId) ─── + +#[test] +fn extract_wxid_strips_suffix() { + assert_eq!( + extract_wxid("wxid_example123abc_ab12"), + "wxid_example123abc" + ); +} + +#[test] +fn extract_wxid_no_suffix() { + assert_eq!(extract_wxid("wxid_example123abc"), "wxid_example123abc"); +} + +#[test] +fn extract_wxid_long_suffix_not_stripped() { + assert_eq!(extract_wxid("wxid_test_abcde"), "wxid_test_abcde"); +} + +#[test] +fn extract_wxid_non_alnum_suffix_not_stripped() { + assert_eq!(extract_wxid("wxid_test_ab-c"), "wxid_test_ab-c"); +} + +#[test] +fn extract_wxid_bare_id_not_shortened() { + // Regression: wxid_test must NOT be shortened to "wxid" + assert_eq!(extract_wxid("wxid_test"), "wxid_test"); +} + +#[test] +fn extract_wxid_legacy_non_wxid_without_signal_stays_raw() { + assert_eq!(extract_wxid("testuser001_1662"), "testuser001_1662"); +} + +// ── Base64 UIN parsing ───────────────────────────────────────────── + +#[test] +fn read_uin_from_dir() { + let tmp = tempfile::tempdir().unwrap(); + let data_dir = tmp.path().join("wxid_test1234_ab12"); + let config_dir = data_dir.join("app_data/radium/ilink/somehash/kvcomm"); + std::fs::create_dir_all(&config_dir).unwrap(); + std::fs::write( + config_dir.join("config.ini"), + "[General]\nlast_uin=MTIzNDU2Nzg5MA==\n", + ) + .unwrap(); + + let uin = read_uin(&data_dir).unwrap(); + assert_eq!(uin, "1234567890"); +} + +#[test] +fn read_uin_no_ilink_dir() { + let tmp = tempfile::tempdir().unwrap(); + let data_dir = tmp.path().join("wxid_test_ab12"); + std::fs::create_dir_all(&data_dir).unwrap(); + assert!(read_uin(&data_dir).is_err()); +} + +#[test] +fn read_uin_documents_level_path() { + // Simulate macOS layout: Documents/app_data/ + Documents/xwechat_files// + let tmp = tempfile::tempdir().unwrap(); + let docs = tmp.path().join("Documents"); + let data_dir = docs.join("xwechat_files/wxid_example123abc_ab12"); + std::fs::create_dir_all(&data_dir).unwrap(); + + let config_dir = docs.join("app_data/radium/ilink/ab12000000000000/kvcomm"); + std::fs::create_dir_all(&config_dir).unwrap(); + std::fs::write( + config_dir.join("config.ini"), + "[default]\nlast_uin=MTIzNDU2Nzg5MA==\n", + ) + .unwrap(); + + let uin = read_uin(&data_dir).unwrap(); + assert_eq!(uin, "1234567890"); +} + +#[test] +fn read_uin_suffix_disambiguation() { + // Two accounts under ilink — suffix selects the right one + let tmp = tempfile::tempdir().unwrap(); + let docs = tmp.path().join("Documents"); + let data_dir = docs.join("xwechat_files/wxid_example123abc_ab12"); + std::fs::create_dir_all(&data_dir).unwrap(); + + // ab12 account → UIN 1234567890 + let ab12_dir = docs.join("app_data/radium/ilink/ab12000000000000/kvcomm"); + std::fs::create_dir_all(&ab12_dir).unwrap(); + std::fs::write(ab12_dir.join("config.ini"), "last_uin=MTIzNDU2Nzg5MA==\n").unwrap(); + + // c3e7 account → UIN 9999999999 + let c3e7_dir = docs.join("app_data/radium/ilink/c3e7000000000000/kvcomm"); + std::fs::create_dir_all(&c3e7_dir).unwrap(); + std::fs::write( + c3e7_dir.join("config.ini"), + "last_uin=OTk5OTk5OTk5OQ==\n", // base64("9999999999") + ) + .unwrap(); + + let uin = read_uin(&data_dir).unwrap(); + assert_eq!(uin, "1234567890"); // ab12 suffix → picks ab12 config +} + +// ── Full derive_v2_key_from_dir ──────────────────────────────────── + +#[test] +fn derive_v2_key_from_dir_integration() { + let tmp = tempfile::tempdir().unwrap(); + let data_dir = tmp.path().join("wxid_example123abc_ab12"); + let config_dir = data_dir.join("app_data/radium/ilink/somehash/kvcomm"); + std::fs::create_dir_all(&config_dir).unwrap(); + std::fs::write( + config_dir.join("config.ini"), + "[General]\nlast_uin=MTIzNDU2Nzg5MA==\n", + ) + .unwrap(); + + let key = derive_v2_key_from_dir(&data_dir).unwrap(); + assert_eq!(&key, b"d5b763f94307be38"); +} + +#[test] +fn derive_v2_key_from_dir_legacy_account_with_login_signal() { + let tmp = tempfile::tempdir().unwrap(); + let docs = tmp.path().join("Documents"); + let data_dir = docs.join("xwechat_files/testuser001_1662"); + let config_dir = docs.join("app_data/radium/ilink/1662000000000000/kvcomm"); + std::fs::create_dir_all(&config_dir).unwrap(); + std::fs::create_dir_all(docs.join("xwechat_files/all_users/login/testuser001")).unwrap(); + std::fs::create_dir_all(&data_dir).unwrap(); + std::fs::write( + config_dir.join("config.ini"), + "[General]\nlast_uin=MTIzNDU2Nzg5MA==\n", + ) + .unwrap(); + + let key = derive_v2_key_from_dir(&data_dir).unwrap(); + let expected = derive_v2_aes_key("1234567890", "testuser001"); + assert_eq!(key, expected); +} diff --git a/crates/wx-media/tests/video-decrypt.rs b/crates/wx-media/tests/video-decrypt.rs new file mode 100644 index 0000000..64e9f19 --- /dev/null +++ b/crates/wx-media/tests/video-decrypt.rs @@ -0,0 +1,72 @@ +use wx_media::{decrypt_video, decrypt_video_with_keystream}; + +#[test] +fn test_zero_keystream_returns_input_unchanged() { + let ciphertext = b"hello world, this is a test video file"; + let keystream = vec![0u8; ciphertext.len()]; + let result = decrypt_video_with_keystream(ciphertext, &keystream); + assert_eq!(result.data, ciphertext); + assert!(!result.is_valid_mp4); +} + +#[test] +fn test_keystream_shorter_than_ciphertext() { + let ciphertext = vec![0xAAu8; 100]; + let keystream = vec![0xBBu8; 30]; + let result = decrypt_video_with_keystream(&ciphertext, &keystream); + assert_eq!(result.data.len(), 100); + // First 30 bytes: 0xAA ^ 0xBB = 0x11 + for &b in &result.data[..30] { + assert_eq!(b, 0x11); + } + // Remaining 70 bytes: unchanged 0xAA + for &b in &result.data[30..] { + assert_eq!(b, 0xAA); + } +} + +#[test] +fn test_ftyp_detection_valid() { + // Construct mock data that contains "ftyp" at offset 4 (standard MP4) + let mut plaintext = vec![0u8; 64]; + plaintext[4..8].copy_from_slice(b"ftyp"); + let result = decrypt_video_with_keystream(&plaintext, &[0u8; 64]); + assert!(result.is_valid_mp4); +} + +#[test] +fn test_ftyp_detection_invalid() { + let plaintext = vec![0u8; 64]; + let result = decrypt_video_with_keystream(&plaintext, &[0u8; 64]); + assert!(!result.is_valid_mp4); +} + +#[test] +fn test_ftyp_detected_after_xor() { + // "ftyp" XORed with key at offset 4 + let key = vec![0x42u8; 32]; + let mut ciphertext = vec![0u8; 32]; + // Set bytes 4..8 so that after XOR with 0x42 they become "ftyp" + ciphertext[4] = b'f' ^ 0x42; + ciphertext[5] = b't' ^ 0x42; + ciphertext[6] = b'y' ^ 0x42; + ciphertext[7] = b'p' ^ 0x42; + let result = decrypt_video_with_keystream(&ciphertext, &key); + assert!(result.is_valid_mp4); + assert_eq!(&result.data[4..8], b"ftyp"); +} + +#[test] +fn test_decrypt_video_with_seed() { + // Just verify it runs without panic and produces output of correct length + let ciphertext = vec![0u8; 256]; + let result = decrypt_video(&ciphertext, 42); + assert_eq!(result.data.len(), 256); +} + +#[test] +fn test_empty_ciphertext() { + let result = decrypt_video(b"", 0); + assert!(result.data.is_empty()); + assert!(!result.is_valid_mp4); +} diff --git a/crates/wx-media/tests/voice.rs b/crates/wx-media/tests/voice.rs new file mode 100644 index 0000000..d0a637f --- /dev/null +++ b/crates/wx-media/tests/voice.rs @@ -0,0 +1,216 @@ +use std::fs; +use std::path::Path; +use tempfile::TempDir; +use wx_media::MediaError; + +/// Create a media_N.db with VoiceInfo table. +fn create_media_db(path: &Path, rows: &[(&str, &[u8])]) { + let conn = rusqlite::Connection::open(path).unwrap(); + conn.execute_batch( + "CREATE TABLE VoiceInfo ( + svr_id TEXT, + voice_data BLOB + );", + ) + .unwrap(); + let mut stmt = conn + .prepare("INSERT INTO VoiceInfo (svr_id, voice_data) VALUES (?, ?)") + .unwrap(); + for &(svr_id, data) in rows { + stmt.execute(rusqlite::params![svr_id, data]).unwrap(); + } +} + +fn create_indexed_media_db(path: &Path, rows: &[(i64, i64, i64, &str, &[u8])]) { + let conn = rusqlite::Connection::open(path).unwrap(); + conn.execute_batch( + "CREATE TABLE VoiceInfo ( + chat_name_id INTEGER, + create_time INTEGER, + local_id INTEGER, + svr_id TEXT, + voice_data BLOB, + data_index TEXT DEFAULT '0' + ); + CREATE INDEX VoiceInfo_INDEX ON VoiceInfo(chat_name_id, svr_id);", + ) + .unwrap(); + let mut stmt = conn + .prepare( + "INSERT INTO VoiceInfo (chat_name_id, create_time, local_id, svr_id, voice_data) + VALUES (?, ?, ?, ?, ?)", + ) + .unwrap(); + for &(chat_name_id, create_time, local_id, svr_id, data) in rows { + stmt.execute(rusqlite::params![ + chat_name_id, + create_time, + local_id, + svr_id, + data + ]) + .unwrap(); + } +} + +#[test] +fn extract_voice_single_db() { + let tmp = TempDir::new().unwrap(); + let media_dir = tmp.path().join("media"); + fs::create_dir_all(&media_dir).unwrap(); + + let silk_data = b"\x02\x23\x21SILK_V3"; + create_media_db(&media_dir.join("media_0.db"), &[("srv_001", silk_data)]); + + let blob = wx_media::extract_voice(&media_dir, "srv_001").unwrap(); + assert_eq!(blob.svr_id, "srv_001"); + assert_eq!(blob.data, silk_data); +} + +#[test] +fn extract_voice_multi_db_scan() { + let tmp = TempDir::new().unwrap(); + let media_dir = tmp.path().join("media"); + fs::create_dir_all(&media_dir).unwrap(); + + // Voice is in media_1.db, not media_0.db + create_media_db(&media_dir.join("media_0.db"), &[]); + create_media_db( + &media_dir.join("media_1.db"), + &[("srv_002", b"silk_audio_bytes")], + ); + + let blob = wx_media::extract_voice(&media_dir, "srv_002").unwrap(); + assert_eq!(blob.svr_id, "srv_002"); + assert_eq!(blob.data, b"silk_audio_bytes"); +} + +#[test] +fn extract_voice_not_found() { + let tmp = TempDir::new().unwrap(); + let media_dir = tmp.path().join("media"); + fs::create_dir_all(&media_dir).unwrap(); + create_media_db(&media_dir.join("media_0.db"), &[]); + + let result = wx_media::extract_voice(&media_dir, "nonexistent"); + assert!(matches!(result, Err(MediaError::LookupMiss(_)))); +} + +#[test] +fn extract_voice_skips_empty_blob() { + let tmp = TempDir::new().unwrap(); + let media_dir = tmp.path().join("media"); + fs::create_dir_all(&media_dir).unwrap(); + + // First db has empty blob, second has real data + create_media_db( + &media_dir.join("media_0.db"), + &[("srv_003", b"")], // empty blob + ); + create_media_db(&media_dir.join("media_1.db"), &[("srv_003", b"real_data")]); + + let blob = wx_media::extract_voice(&media_dir, "srv_003").unwrap(); + assert_eq!(blob.data, b"real_data"); +} + +#[test] +fn extract_voice_no_media_dbs() { + let tmp = TempDir::new().unwrap(); + let media_dir = tmp.path().join("media"); + fs::create_dir_all(&media_dir).unwrap(); + // No media_*.db files + + let result = wx_media::extract_voice(&media_dir, "srv_001"); + assert!(matches!(result, Err(MediaError::NoMediaDbs(_)))); +} + +#[test] +fn extract_voice_compat_media_db() { + // Test `media.db` (no numeric suffix) is also scanned + let tmp = TempDir::new().unwrap(); + let media_dir = tmp.path().join("media"); + fs::create_dir_all(&media_dir).unwrap(); + + create_media_db(&media_dir.join("media.db"), &[("srv_004", b"compat_voice")]); + + let blob = wx_media::extract_voice(&media_dir, "srv_004").unwrap(); + assert_eq!(blob.data, b"compat_voice"); +} + +#[test] +fn voice_query_with_conn_returns_non_empty_blob() { + let tmp = TempDir::new().unwrap(); + let db_path = tmp.path().join("media.db"); + create_media_db(&db_path, &[("srv_conn_1", b"blob_one")]); + let conn = rusqlite::Connection::open(&db_path).unwrap(); + + let blob = wx_media::extract_voice_with_conn(&conn, "srv_conn_1").unwrap(); + assert_eq!(blob.svr_id, "srv_conn_1"); + assert_eq!(blob.data, b"blob_one"); + assert_eq!(blob.chat_name_id, None); +} + +#[test] +fn voice_query_with_conn_skips_empty_blob_rows() { + let tmp = TempDir::new().unwrap(); + let db_path = tmp.path().join("media.db"); + create_media_db( + &db_path, + &[("srv_conn_2", b""), ("srv_conn_2", b"blob_two")], + ); + let conn = rusqlite::Connection::open(&db_path).unwrap(); + + let blob = wx_media::extract_voice_with_conn(&conn, "srv_conn_2").unwrap(); + assert_eq!(blob.data, b"blob_two"); + assert_eq!(blob.chat_name_id, None); +} + +#[test] +fn voice_query_with_conn_returns_stable_not_found_errors() { + let tmp = TempDir::new().unwrap(); + let db_path = tmp.path().join("media.db"); + create_media_db(&db_path, &[]); + let conn = rusqlite::Connection::open(&db_path).unwrap(); + + let result = wx_media::extract_voice_with_conn(&conn, "missing"); + assert!(matches!(result, Err(MediaError::LookupMiss(s)) if s.contains("svr_id missing"))); + + let no_table = rusqlite::Connection::open(tmp.path().join("no-table.db")).unwrap(); + let result = wx_media::extract_voice_with_conn(&no_table, "missing"); + assert!(matches!(result, Err(MediaError::SchemaMissing(s)) if s == "VoiceInfo table missing")); +} + +#[test] +fn voice_query_with_conn_hint_returns_chat_name_id_from_indexed_lookup() { + let tmp = TempDir::new().unwrap(); + let db_path = tmp.path().join("media.db"); + create_indexed_media_db( + &db_path, + &[(55, 1000, 1, "srv_hint_1", b"hinted_blob")], + ); + let conn = rusqlite::Connection::open(&db_path).unwrap(); + + let blob = wx_media::extract_voice_with_conn_hint(&conn, "srv_hint_1", Some(55)).unwrap(); + assert_eq!(blob.svr_id, "srv_hint_1"); + assert_eq!(blob.data, b"hinted_blob"); + assert_eq!(blob.chat_name_id, Some(55)); +} + +#[test] +fn voice_query_with_conn_hint_falls_back_to_scan_when_hint_is_wrong() { + let tmp = TempDir::new().unwrap(); + let db_path = tmp.path().join("media.db"); + create_indexed_media_db( + &db_path, + &[ + (99, 1000, 1, "srv_hint_2", b""), + (77, 1001, 2, "srv_hint_2", b"fallback_blob"), + ], + ); + let conn = rusqlite::Connection::open(&db_path).unwrap(); + + let blob = wx_media::extract_voice_with_conn_hint(&conn, "srv_hint_2", Some(55)).unwrap(); + assert_eq!(blob.svr_id, "srv_hint_2"); + assert_eq!(blob.data, b"fallback_blob"); + assert_eq!(blob.chat_name_id, Some(77)); +} diff --git a/crates/wx-media/tests/wxgf-transcode.rs b/crates/wx-media/tests/wxgf-transcode.rs new file mode 100644 index 0000000..4b63e68 --- /dev/null +++ b/crates/wx-media/tests/wxgf-transcode.rs @@ -0,0 +1,94 @@ +use std::process::Command; +use std::sync::Mutex; + +static FFMPEG_ENV_LOCK: Mutex<()> = Mutex::new(()); + +#[test] +fn transcode_wxgf_respects_embedded_png_and_missing_ffmpeg_hevc_fallback() { + let _guard = FFMPEG_ENV_LOCK.lock().unwrap(); + unsafe { + std::env::set_var("FFMPEG_PATH", "/definitely-missing-ffmpeg"); + } + wx_media::reset_ffmpeg_cache(); + + let mut embedded_png = b"wxgfmetadata".to_vec(); + embedded_png.extend_from_slice(&sample_png()); + let embedded = wx_media::transcode_wxgf(&embedded_png).unwrap(); + assert_eq!(embedded.ext, "png"); + assert!(embedded.transcoded); + assert_eq!(&embedded.data[..8], b"\x89PNG\r\n\x1a\n"); + + let mut hevc = b"wxgfmetadata".to_vec(); + hevc.extend_from_slice(&[0x00, 0x00, 0x00, 0x01, 0x26, 0x01, 0x02, 0x03, 0x04]); + let hevc_result = wx_media::transcode_wxgf(&hevc).unwrap(); + assert_eq!(hevc_result.ext, "hevc"); + assert!(!hevc_result.transcoded); + assert_eq!( + hevc_result.data, + vec![0x00, 0x00, 0x00, 0x01, 0x26, 0x01, 0x02, 0x03, 0x04] + ); + + unsafe { + std::env::remove_var("FFMPEG_PATH"); + } +} + +#[test] +fn transcode_wxgf_hevc_returns_png_when_ffmpeg_is_available() { + let _guard = FFMPEG_ENV_LOCK.lock().unwrap(); + unsafe { + std::env::remove_var("FFMPEG_PATH"); + } + wx_media::reset_ffmpeg_cache(); + + if !wx_media::ffmpeg_available() { + return; + } + + let mut wxgf = b"wxgfmetadata".to_vec(); + wxgf.extend_from_slice(&sample_valid_hevc()); + + let result = wx_media::transcode_wxgf(&wxgf).unwrap(); + assert_eq!(result.ext, "png"); + assert!(result.transcoded); + assert_eq!(&result.data[..8], b"\x89PNG\r\n\x1a\n"); +} + +fn sample_png() -> Vec { + vec![ + 0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A, 0x00, 0x00, 0x00, 0x0D, 0x49, 0x48, 0x44, + 0x52, 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01, 0x08, 0x06, 0x00, 0x00, 0x00, 0x1F, + 0x15, 0xC4, 0x89, 0x00, 0x00, 0x00, 0x0D, 0x49, 0x44, 0x41, 0x54, 0x78, 0x9C, 0x63, 0xF8, + 0xCF, 0xC0, 0xF0, 0x1F, 0x00, 0x05, 0x00, 0x01, 0xFF, 0x89, 0x99, 0x3D, 0x1D, 0x00, 0x00, + 0x00, 0x00, 0x49, 0x45, 0x4E, 0x44, 0xAE, 0x42, 0x60, 0x82, + ] +} + +fn sample_valid_hevc() -> Vec { + let temp = tempfile::TempDir::new().unwrap(); + let output = temp.path().join("frame.hevc"); + let ffmpeg = std::env::var("FFMPEG_PATH").unwrap_or_else(|_| "ffmpeg".to_string()); + let status = Command::new(ffmpeg) + .args([ + "-hide_banner", + "-loglevel", + "error", + "-f", + "lavfi", + "-i", + "color=c=red:s=64x64:d=0.04:r=1", + "-frames:v", + "1", + "-c:v", + "libx265", + "-x265-params", + "log-level=error", + "-f", + "hevc", + output.to_str().unwrap(), + ]) + .status() + .unwrap(); + assert!(status.success()); + std::fs::read(output).unwrap() +} diff --git a/crates/wx-media/tests/wxgf.rs b/crates/wx-media/tests/wxgf.rs new file mode 100644 index 0000000..a26c2cf --- /dev/null +++ b/crates/wx-media/tests/wxgf.rs @@ -0,0 +1,115 @@ +use wx_media::{parse_wxgf, ImageType, MediaError, WxgfContent}; + +#[test] +fn test_non_wxgf_magic_returns_error() { + let data = b"not wxgf data at all"; + let err = parse_wxgf(data).unwrap_err(); + assert!(matches!(err, MediaError::InvalidWxgf)); +} + +#[test] +fn test_too_short_returns_error() { + let data = b"wxg"; // Only 3 bytes + let err = parse_wxgf(data).unwrap_err(); + assert!(matches!(err, MediaError::InvalidWxgf)); +} + +#[test] +fn test_wxgf_with_hevc_4byte_start_code() { + let mut data = b"wxgf".to_vec(); + // Some padding bytes before the HEVC start code + data.extend_from_slice(&[0x00, 0x01, 0x02, 0x03]); + // 4-byte HEVC start code + data.extend_from_slice(&[0x00, 0x00, 0x00, 0x01]); + // Mock HEVC NALU data + data.extend_from_slice(&[0x40, 0x01, 0xFF, 0xFF]); + + match parse_wxgf(&data).unwrap() { + WxgfContent::Hevc(hevc_data) => { + assert_eq!(&hevc_data[..4], &[0x00, 0x00, 0x00, 0x01]); + assert_eq!(hevc_data.len(), 8); // start code + NALU data + } + WxgfContent::EmbeddedImage { .. } => panic!("expected Hevc"), + } +} + +#[test] +fn test_wxgf_with_hevc_3byte_start_code() { + let mut data = b"wxgf".to_vec(); + data.extend_from_slice(&[0x01, 0x02, 0x03]); + // 3-byte start code (no 4-byte found) + data.extend_from_slice(&[0x00, 0x00, 0x01]); + data.extend_from_slice(&[0x26, 0x01]); + + match parse_wxgf(&data).unwrap() { + WxgfContent::Hevc(hevc_data) => { + assert_eq!(&hevc_data[..3], &[0x00, 0x00, 0x01]); + } + WxgfContent::EmbeddedImage { .. } => panic!("expected Hevc"), + } +} + +#[test] +fn test_wxgf_with_embedded_jpg() { + let mut data = b"wxgf".to_vec(); + // Some header bytes + data.extend_from_slice(&[0x00; 10]); + // JPG magic at offset 14 + data.extend_from_slice(&[0xFF, 0xD8, 0xFF, 0xE0]); + // JPG data + data.extend_from_slice(&[0x01, 0x02, 0x03]); + + match parse_wxgf(&data).unwrap() { + WxgfContent::EmbeddedImage { data: img, ext } => { + assert_eq!(ext, "jpg"); + assert_eq!(&img[..3], &[0xFF, 0xD8, 0xFF]); + } + WxgfContent::Hevc(_) => panic!("expected EmbeddedImage"), + } +} + +#[test] +fn test_wxgf_with_embedded_png() { + let mut data = b"wxgf".to_vec(); + data.extend_from_slice(&[0x00; 10]); + // PNG magic + data.extend_from_slice(&[0x89, 0x50, 0x4E, 0x47]); + data.extend_from_slice(&[0x0D, 0x0A, 0x1A, 0x0A]); + + match parse_wxgf(&data).unwrap() { + WxgfContent::EmbeddedImage { data: img, ext } => { + assert_eq!(ext, "png"); + assert_eq!(&img[..4], &[0x89, 0x50, 0x4E, 0x47]); + } + WxgfContent::Hevc(_) => panic!("expected EmbeddedImage"), + } +} + +#[test] +fn test_embedded_image_takes_priority_over_hevc() { + let mut data = b"wxgf".to_vec(); + data.extend_from_slice(&[0x00; 10]); + // JPG magic appears first + data.extend_from_slice(&[0xFF, 0xD8, 0xFF, 0xE0]); + data.extend_from_slice(&[0x00; 20]); + // HEVC start code appears later + data.extend_from_slice(&[0x00, 0x00, 0x00, 0x01]); + + match parse_wxgf(&data).unwrap() { + WxgfContent::EmbeddedImage { ext, .. } => assert_eq!(ext, "jpg"), + WxgfContent::Hevc(_) => panic!("embedded image should take priority"), + } +} + +#[test] +fn test_no_start_code_no_embedded_returns_error() { + let mut data = b"wxgf".to_vec(); + data.extend_from_slice(&[0x01, 0x02, 0x03, 0x04, 0x05]); + let err = parse_wxgf(&data).unwrap_err(); + assert!(matches!(err, MediaError::InvalidWxgf)); +} + +#[test] +fn test_image_type_wxgf_ext() { + assert_eq!(ImageType::Wxgf.ext(), "wxgf"); +} diff --git a/crates/wx-monitor/Cargo.toml b/crates/wx-monitor/Cargo.toml new file mode 100644 index 0000000..2970482 --- /dev/null +++ b/crates/wx-monitor/Cargo.toml @@ -0,0 +1,24 @@ +[package] +name = "wx-monitor" +version.workspace = true +edition.workspace = true + +[dependencies] +wx-decrypt = { path = "../wx-decrypt" } +wx-db = { path = "../wx-db" } +notify = "8" +tokio = { version = "1", features = ["rt", "sync", "time"] } +futures-core = "0.3" +tempfile = "3" +tracing = "0.1" +thiserror = "2" +serde = { version = "1", features = ["derive"] } + +[dev-dependencies] +aes = "0.8" +cbc = "0.1" +hmac = "0.12" +sha2 = "0.10" +pbkdf2 = { version = "0.12", features = ["sha2"] } +rusqlite = { version = "0.32", features = ["bundled"] } +tokio = { version = "1", features = ["rt-multi-thread", "macros", "time"] } diff --git a/crates/wx-monitor/src/cache.rs b/crates/wx-monitor/src/cache.rs new file mode 100644 index 0000000..fb1db78 --- /dev/null +++ b/crates/wx-monitor/src/cache.rs @@ -0,0 +1,472 @@ +use std::path::{Path, PathBuf}; +use std::time::SystemTime; + +use wx_decrypt::{CryptoParams, KeyMaterial}; + +use crate::error::MonitorError; + +/// Outcome of a cache update cycle. +#[derive(Debug, PartialEq, Eq)] +pub(crate) enum UpdateKind { + /// The main .db file changed; a full re-decrypt was performed. + FullDecrypt, + /// Only the WAL file changed; patched in-place. + WalPatched, + /// Neither file changed. + NoChange, +} + +/// RAII cache that manages decrypted copies in a temp directory. +/// +/// Mirrors the `db_storage/` layout so that `WechatDb::open()` can operate +/// on the decrypted root. +pub(crate) struct DecryptCache { + dir: tempfile::TempDir, + encrypted_session_dir: PathBuf, + key_material: KeyMaterial, + params: &'static CryptoParams, + last_db_mtime: Option, + last_wal_mtime: Option, +} + +/// SQLite database file header (first 16 bytes). +const SQLITE_HEADER: &[u8; 16] = b"SQLite format 3\0"; + +impl DecryptCache { + pub fn new( + encrypted_session_dir: PathBuf, + key_material: KeyMaterial, + params: &'static CryptoParams, + ) -> Result { + let dir = tempfile::TempDir::new()?; + let root = dir.path(); + + // Create directory layout matching db_storage/ + std::fs::create_dir_all(root.join("session"))?; + std::fs::create_dir_all(root.join("contact"))?; + std::fs::create_dir_all(root.join("message"))?; + + // Write placeholder SQLite files so WechatDb::open() succeeds + let minimal = make_minimal_sqlite(); + std::fs::write(root.join("contact").join("contact.db"), &minimal)?; + std::fs::write(root.join("message").join("message_0.db"), &minimal)?; + + Ok(Self { + dir, + encrypted_session_dir, + key_material, + params, + last_db_mtime: None, + last_wal_mtime: None, + }) + } + + /// Root path for `WechatDb::open()`. + pub fn decrypted_root(&self) -> &Path { + self.dir.path() + } + + /// Perform the initial full decrypt of session.db (and WAL if present). + pub fn initial_decrypt(&mut self) -> Result<(), MonitorError> { + let enc_db = self.encrypted_session_dir.join("session.db"); + let dec_db = self.dir.path().join("session").join("session.db"); + + self.do_decrypt_db(&enc_db, &dec_db)?; + self.last_db_mtime = file_mtime(&enc_db); + + // Patch WAL if it exists + let wal = self.encrypted_session_dir.join("session.db-wal"); + if wal.exists() { + let _ = self.do_decrypt_wal(&wal, &dec_db); + self.last_wal_mtime = file_mtime(&wal); + } + + Ok(()) + } + + /// Check for changes and update the decrypted cache accordingly. + pub fn update(&mut self) -> Result { + let enc_db = self.encrypted_session_dir.join("session.db"); + let dec_db = self.dir.path().join("session").join("session.db"); + let wal = self.encrypted_session_dir.join("session.db-wal"); + + let db_mtime = file_mtime(&enc_db); + let wal_mtime = if wal.exists() { file_mtime(&wal) } else { None }; + + if db_mtime != self.last_db_mtime { + // Full re-decrypt (DB changed, WAL may also have changed) + self.do_decrypt_db(&enc_db, &dec_db)?; + if wal.exists() { + let _ = self.do_decrypt_wal(&wal, &dec_db); + } + self.last_db_mtime = db_mtime; + self.last_wal_mtime = wal_mtime; + Ok(UpdateKind::FullDecrypt) + } else if wal_mtime != self.last_wal_mtime { + // WAL-only patch + if wal.exists() { + self.do_decrypt_wal(&wal, &dec_db)?; + } + self.last_wal_mtime = wal_mtime; + Ok(UpdateKind::WalPatched) + } else { + Ok(UpdateKind::NoChange) + } + } + + fn do_decrypt_db(&self, src: &Path, dst: &Path) -> Result<(), wx_decrypt::DecryptError> { + wx_decrypt::dispatch_decrypt_db(src, dst, &self.key_material, self.params) + } + + fn do_decrypt_wal( + &self, + wal: &Path, + dst: &Path, + ) -> Result { + wx_decrypt::dispatch_decrypt_wal(wal, dst, &self.key_material, self.params) + } +} + +/// Build a minimal valid SQLite database (single 4096-byte page with header). +fn make_minimal_sqlite() -> Vec { + let mut page = vec![0u8; 4096]; + // SQLite header + page[..16].copy_from_slice(SQLITE_HEADER); + // Page size = 4096 (big-endian at offset 16) + page[16] = 0x10; + page[17] = 0x00; + // File format write version = 1 + page[18] = 1; + // File format read version = 1 + page[19] = 1; + // Reserved space per page = 0 + page[20] = 0; + // Max embedded payload fraction = 64 + page[21] = 64; + // Min embedded payload fraction = 32 + page[22] = 32; + // Leaf payload fraction = 32 + page[23] = 32; + // Page count = 1 (big-endian at offset 28) + page[31] = 1; + // Schema format number = 4 (offset 44) + page[47] = 4; + // Text encoding = UTF-8 = 1 (offset 56) + page[59] = 1; + page +} + +/// Get file mtime as `SystemTime`, or `None` if unavailable. +fn file_mtime(path: &Path) -> Option { + std::fs::metadata(path).ok().and_then(|m| m.modified().ok()) +} + +#[cfg(test)] +mod tests { + use super::*; + use wx_decrypt::{KeyMaterial, MACOS_4_1_7_31}; + + #[test] + fn new_creates_correct_layout() { + let enc_dir = tempfile::TempDir::new().unwrap(); + let cache = DecryptCache::new( + enc_dir.path().to_path_buf(), + KeyMaterial::RawKey([0u8; 32]), + &MACOS_4_1_7_31, + ) + .unwrap(); + + let root = cache.decrypted_root(); + assert!(root.join("session").is_dir()); + assert!(root.join("contact").is_dir()); + assert!(root.join("message").is_dir()); + assert!(root.join("contact").join("contact.db").is_file()); + assert!(root.join("message").join("message_0.db").is_file()); + + // Placeholder files should be valid SQLite + let data = std::fs::read(root.join("contact").join("contact.db")).unwrap(); + assert_eq!(&data[..16], SQLITE_HEADER); + assert_eq!(data.len(), 4096); + } + + #[test] + fn update_returns_no_change_when_nothing_changed() { + let enc_dir = tempfile::TempDir::new().unwrap(); + let session_dir = enc_dir.path().join("session_dir"); + std::fs::create_dir_all(&session_dir).unwrap(); + + // Create a minimal encrypted session.db + let raw_key = [0xABu8; 32]; + build_encrypted_session_db(&session_dir.join("session.db"), &raw_key); + + let mut cache = DecryptCache::new( + session_dir.clone(), + KeyMaterial::RawKey(raw_key), + &MACOS_4_1_7_31, + ) + .unwrap(); + cache.initial_decrypt().unwrap(); + + // Verify the decrypted file exists and is valid SQLite + let dec_db = cache.decrypted_root().join("session").join("session.db"); + assert!(dec_db.is_file()); + let data = std::fs::read(&dec_db).unwrap(); + assert_eq!(&data[..16], SQLITE_HEADER); + + // No changes → NoChange + assert_eq!(cache.update().unwrap(), UpdateKind::NoChange); + } + + #[test] + fn update_detects_full_decrypt_on_db_change() { + let session_dir = tempfile::TempDir::new().unwrap(); + let raw_key = [0xABu8; 32]; + build_encrypted_session_db(&session_dir.path().join("session.db"), &raw_key); + + let mut cache = DecryptCache::new( + session_dir.path().to_path_buf(), + KeyMaterial::RawKey(raw_key), + &MACOS_4_1_7_31, + ) + .unwrap(); + cache.initial_decrypt().unwrap(); + + // Wait to ensure mtime difference + std::thread::sleep(std::time::Duration::from_millis(1100)); + + // Re-write the encrypted DB (simulates DB checkpoint) + build_encrypted_session_db(&session_dir.path().join("session.db"), &raw_key); + + assert_eq!(cache.update().unwrap(), UpdateKind::FullDecrypt); + } + + #[test] + fn update_detects_wal_patched_on_wal_change() { + let session_dir = tempfile::TempDir::new().unwrap(); + let raw_key = [0xABu8; 32]; + build_encrypted_session_db(&session_dir.path().join("session.db"), &raw_key); + + let mut cache = DecryptCache::new( + session_dir.path().to_path_buf(), + KeyMaterial::RawKey(raw_key), + &MACOS_4_1_7_31, + ) + .unwrap(); + cache.initial_decrypt().unwrap(); + assert_eq!(cache.update().unwrap(), UpdateKind::NoChange); + + // Wait to ensure mtime difference + std::thread::sleep(std::time::Duration::from_millis(1100)); + + // Create a minimal valid WAL file (header only, no frames). + // decrypt_wal will return Ok(0) — no frames patched — but mtime changed. + let wal_path = session_dir.path().join("session.db-wal"); + let mut wal_header = [0u8; 32]; + wal_header[0..4].copy_from_slice(&0x377f_0682u32.to_be_bytes()); // WAL_MAGIC_BE + std::fs::write(&wal_path, wal_header).unwrap(); + + assert_eq!(cache.update().unwrap(), UpdateKind::WalPatched); + + // Subsequent call with no changes → NoChange + assert_eq!(cache.update().unwrap(), UpdateKind::NoChange); + } + + // ---- test helper: build encrypted session.db ---- + + /// Build a minimal encrypted session.db that `decrypt_db` can process. + fn build_encrypted_session_db(path: &Path, raw_key: &[u8; 32]) -> PathBuf { + use aes::cipher::{BlockEncryptMut, KeyIvInit}; + use hmac::{Hmac, Mac}; + use sha2::Sha512; + + let params = &MACOS_4_1_7_31; + let salt: [u8; 16] = [0x01; 16]; + let iv: [u8; 16] = [0x42; 16]; + + // Derive keys + let enc_key = derive_enc_key(raw_key, &salt, params); + let mac_key = derive_mac_key(&enc_key, &salt); + + // Build plaintext page 0: a minimal SQLite page (without the 16-byte header, + // since decrypt_db will prepend it) + let data_size = params.page_size - params.reserve - params.salt_size; // 4000 + let mut plaintext = vec![0u8; data_size]; + // Page size at offset 0 (which maps to file offset 16): 0x10 0x00 = 4096 + plaintext[0] = 0x10; + plaintext[1] = 0x00; + // Write/read format versions + plaintext[2] = 1; + plaintext[3] = 1; + // Max embedded payload fraction = 64 + plaintext[5] = 64; + // Min embedded payload fraction = 32 + plaintext[6] = 32; + // Leaf payload fraction = 32 + plaintext[7] = 32; + // Page count = 1 at offset 12 (file offset 28) + plaintext[15] = 1; + // Schema format = 4 at offset 28 (file offset 44) + plaintext[31] = 4; + // Text encoding = UTF-8 at offset 40 (file offset 56) + plaintext[43] = 1; + + // Pad to AES block size (already aligned: 4000 is divisible by 16) + + // Encrypt + type Aes256CbcEnc = cbc::Encryptor; + let mut ciphertext = plaintext.clone(); + let encryptor = Aes256CbcEnc::new((&enc_key).into(), (&iv).into()); + encryptor + .encrypt_padded_mut::(&mut ciphertext, data_size) + .unwrap(); + + // Assemble page: salt + ciphertext + IV + HMAC + let mut page = Vec::with_capacity(params.page_size); + page.extend_from_slice(&salt); + page.extend_from_slice(&ciphertext); + // Reserve area: IV(16) + HMAC(64) + page.extend_from_slice(&iv); + page.resize(params.page_size, 0); // zero-fill HMAC area + + // Compute HMAC over: page[salt_size..page_size - reserve + iv_size] + page_num(1, LE) + let hmac_data_end = params.page_size - params.reserve + params.iv_size; + let mut mac = as Mac>::new_from_slice(&mac_key).unwrap(); + mac.update(&page[params.salt_size..hmac_data_end]); + mac.update(&1u32.to_le_bytes()); // page number is 1-indexed + let hmac_result = mac.finalize().into_bytes(); + + // Place HMAC + let hmac_start = params.page_size - params.reserve + params.iv_size; + page[hmac_start..hmac_start + params.hmac_size] + .copy_from_slice(&hmac_result[..params.hmac_size]); + + std::fs::write(path, &page).unwrap(); + path.to_path_buf() + } + + fn derive_enc_key(raw_key: &[u8; 32], salt: &[u8; 16], params: &CryptoParams) -> [u8; 32] { + let mut key = [0u8; 32]; + pbkdf2::pbkdf2_hmac::(raw_key, salt, params.kdf_iter, &mut key); + key + } + + fn derive_mac_key(enc_key: &[u8; 32], salt: &[u8; 16]) -> [u8; 32] { + let mut mac_salt = [0u8; 16]; + for (i, b) in salt.iter().enumerate() { + mac_salt[i] = b ^ 0x3a; + } + let mut key = [0u8; 32]; + pbkdf2::pbkdf2_hmac::(enc_key, &mac_salt, 2, &mut key); + key + } + + /// Derive enc_key + salt from a raw key, for building EncKey test fixtures. + fn derive_enc_key_and_salt(raw_key: &[u8; 32]) -> ([u8; 32], [u8; 16]) { + let salt: [u8; 16] = [0x01; 16]; + let enc_key = derive_enc_key(raw_key, &salt, &MACOS_4_1_7_31); + (enc_key, salt) + } + + #[test] + fn enc_key_initial_decrypt_succeeds() { + let session_dir = tempfile::TempDir::new().unwrap(); + let raw_key = [0xABu8; 32]; + build_encrypted_session_db(&session_dir.path().join("session.db"), &raw_key); + + let (enc_key, salt) = derive_enc_key_and_salt(&raw_key); + let key_material = KeyMaterial::EncKey { key: enc_key, salt }; + + let mut cache = DecryptCache::new( + session_dir.path().to_path_buf(), + key_material, + &MACOS_4_1_7_31, + ) + .unwrap(); + cache.initial_decrypt().unwrap(); + + // Verify decrypted file is valid SQLite + let dec_db = cache.decrypted_root().join("session").join("session.db"); + assert!(dec_db.is_file()); + let data = std::fs::read(&dec_db).unwrap(); + assert_eq!(&data[..16], SQLITE_HEADER); + } + + #[test] + fn enc_key_update_detects_changes() { + let session_dir = tempfile::TempDir::new().unwrap(); + let raw_key = [0xABu8; 32]; + build_encrypted_session_db(&session_dir.path().join("session.db"), &raw_key); + + let (enc_key, salt) = derive_enc_key_and_salt(&raw_key); + let key_material = KeyMaterial::EncKey { key: enc_key, salt }; + + let mut cache = DecryptCache::new( + session_dir.path().to_path_buf(), + key_material, + &MACOS_4_1_7_31, + ) + .unwrap(); + cache.initial_decrypt().unwrap(); + assert_eq!(cache.update().unwrap(), UpdateKind::NoChange); + + // Wait for mtime difference then re-write + std::thread::sleep(std::time::Duration::from_millis(1100)); + build_encrypted_session_db(&session_dir.path().join("session.db"), &raw_key); + assert_eq!(cache.update().unwrap(), UpdateKind::FullDecrypt); + } + + #[test] + fn enc_keys_initial_decrypt_succeeds() { + use wx_decrypt::EncKeyPair; + + let session_dir = tempfile::TempDir::new().unwrap(); + let raw_key = [0xABu8; 32]; + build_encrypted_session_db(&session_dir.path().join("session.db"), &raw_key); + + let (enc_key, salt) = derive_enc_key_and_salt(&raw_key); + // Wrap in EncKeys (single pair) — the canonical format from capture_key_mach + let key_material = KeyMaterial::EncKeys(vec![EncKeyPair { key: enc_key, salt }]); + + let mut cache = DecryptCache::new( + session_dir.path().to_path_buf(), + key_material, + &MACOS_4_1_7_31, + ) + .unwrap(); + cache.initial_decrypt().unwrap(); + + let dec_db = cache.decrypted_root().join("session").join("session.db"); + assert!(dec_db.is_file()); + let data = std::fs::read(&dec_db).unwrap(); + assert_eq!( + &data[..16], + SQLITE_HEADER, + "EncKeys path should produce valid SQLite" + ); + } + + #[test] + fn enc_keys_update_detects_changes() { + use wx_decrypt::EncKeyPair; + + let session_dir = tempfile::TempDir::new().unwrap(); + let raw_key = [0xABu8; 32]; + build_encrypted_session_db(&session_dir.path().join("session.db"), &raw_key); + + let (enc_key, salt) = derive_enc_key_and_salt(&raw_key); + let key_material = KeyMaterial::EncKeys(vec![EncKeyPair { key: enc_key, salt }]); + + let mut cache = DecryptCache::new( + session_dir.path().to_path_buf(), + key_material, + &MACOS_4_1_7_31, + ) + .unwrap(); + cache.initial_decrypt().unwrap(); + assert_eq!(cache.update().unwrap(), UpdateKind::NoChange); + + std::thread::sleep(std::time::Duration::from_millis(1100)); + build_encrypted_session_db(&session_dir.path().join("session.db"), &raw_key); + assert_eq!(cache.update().unwrap(), UpdateKind::FullDecrypt); + } +} diff --git a/crates/wx-monitor/src/error.rs b/crates/wx-monitor/src/error.rs new file mode 100644 index 0000000..6072d09 --- /dev/null +++ b/crates/wx-monitor/src/error.rs @@ -0,0 +1,16 @@ +use thiserror::Error; + +/// Errors that can occur during monitoring. +#[derive(Debug, Error)] +pub enum MonitorError { + #[error("watcher error: {0}")] + Watcher(#[from] notify::Error), + #[error("decrypt error: {0}")] + Decrypt(#[from] wx_decrypt::DecryptError), + #[error("database error: {0}")] + Db(#[from] wx_db::DbError), + #[error("io error: {0}")] + Io(#[from] std::io::Error), + #[error("data directory not found: {0}")] + DataDirNotFound(String), +} diff --git a/crates/wx-monitor/src/event.rs b/crates/wx-monitor/src/event.rs new file mode 100644 index 0000000..3761234 --- /dev/null +++ b/crates/wx-monitor/src/event.rs @@ -0,0 +1,49 @@ +use std::path::PathBuf; + +use serde::Serialize; + +/// Internal event from the file watcher layer. +#[derive(Debug)] +#[allow(dead_code)] // fields read in tests; monitor loop only checks event arrival +pub(crate) struct FileEvent { + pub path: PathBuf, + pub kind: FileEventKind, +} + +/// Kind of file change detected. +#[derive(Debug)] +pub(crate) enum FileEventKind { + Modified, + Created, +} + +/// A detected change in the WeChat session list. +#[derive(Debug, Clone, Serialize)] +pub struct SessionEvent { + /// The wxid or chatroom username that changed. + pub username: String, + /// The sort_timestamp from session.db. + pub sort_timestamp: i64, + /// Unix timestamp (seconds) when the change was detected. + pub detected_at: i64, + /// Whether this is an incremental update or a full reset. + pub kind: SessionEventKind, + /// Summary text of the last message. + pub summary: String, + /// The message type of the last message. + #[serde(skip_serializing_if = "Option::is_none")] + pub last_msg_type: Option, + /// The wxid of the last message sender. + #[serde(skip_serializing_if = "Option::is_none")] + pub last_msg_sender: Option, + /// The display name of the last message sender. + #[serde(skip_serializing_if = "Option::is_none")] + pub last_sender_display_name: Option, +} + +/// Kind of session change. +#[derive(Debug, Clone, Serialize)] +pub enum SessionEventKind { + /// Session was updated (new or changed sort_timestamp). + Updated, +} diff --git a/crates/wx-monitor/src/lib.rs b/crates/wx-monitor/src/lib.rs new file mode 100644 index 0000000..0405d4a --- /dev/null +++ b/crates/wx-monitor/src/lib.rs @@ -0,0 +1,12 @@ +mod cache; +mod error; +mod event; +mod monitor; +mod tracker; +mod watcher; + +pub use error::MonitorError; +pub use event::{SessionEvent, SessionEventKind}; +pub use monitor::{ + resolve_watch_mode, MonitorConfig, MonitorStream, ResolvedWatcher, WatchMode, WechatMonitor, +}; diff --git a/crates/wx-monitor/src/monitor.rs b/crates/wx-monitor/src/monitor.rs new file mode 100644 index 0000000..a705eb3 --- /dev/null +++ b/crates/wx-monitor/src/monitor.rs @@ -0,0 +1,382 @@ +use std::path::PathBuf; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::sync::Arc; +use std::time::Duration; + +use futures_core::Stream; +use wx_db::{SessionQuery, WechatDb}; +use wx_decrypt::{CryptoParams, KeyMaterial}; + +use crate::cache::{DecryptCache, UpdateKind}; +use crate::error::MonitorError; +use crate::event::SessionEvent; +use crate::tracker::SessionTracker; +use crate::watcher::{FileWatcher, NotifyWatcher, PollingWatcher}; + +/// Watcher mode selection. +#[derive(Debug, Clone, Default)] +pub enum WatchMode { + /// Automatic: polling on macOS, fsnotify on other platforms. + #[default] + Auto, + /// Force polling. + Poll, + /// Force fsnotify (opt-in on macOS). + Fsnotify, +} + +/// Which watcher implementation to use (resolved from WatchMode). +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum ResolvedWatcher { + Polling, + Notify, +} + +/// Resolve a `WatchMode` to the concrete watcher implementation. +pub fn resolve_watch_mode(mode: &WatchMode) -> ResolvedWatcher { + match mode { + WatchMode::Poll => ResolvedWatcher::Polling, + WatchMode::Fsnotify => ResolvedWatcher::Notify, + WatchMode::Auto => { + #[cfg(target_os = "macos")] + { + ResolvedWatcher::Polling + } + #[cfg(not(target_os = "macos"))] + { + ResolvedWatcher::Notify + } + } + } +} + +/// Configuration for the monitor. +pub struct MonitorConfig { + pub encrypted_session_dir: PathBuf, + pub key_material: KeyMaterial, + pub params: &'static CryptoParams, + /// Watcher mode selection. Default: Auto (polling on macOS, fsnotify elsewhere). + pub watch_mode: WatchMode, + pub poll_interval: Duration, + pub channel_capacity: usize, + /// Raw key for direct encrypted open (bypasses DecryptCache). + pub raw_key: Option<[u8; 32]>, + /// Full db_storage root path (required when raw_key is set). + pub encrypted_root: Option, +} + +/// The main monitor handle. +/// +/// Detects changes to an encrypted WeChat session.db, decrypts incrementally, +/// and emits [`SessionEvent`]s. +pub struct WechatMonitor { + receiver: Option>, + shutdown_tx: Option>, + shutdown_flag: Option>, + _task: Option>, +} + +/// Returns true if a `notify::Error` indicates a backend-unavailable condition +/// (should fall back to polling). Returns false for path/permission/config errors +/// (should propagate as Err). +fn is_backend_unavailable(err: ¬ify::Error) -> bool { + match &err.kind { + notify::ErrorKind::Generic(_) => true, + notify::ErrorKind::Io(io_err) => !matches!( + io_err.kind(), + std::io::ErrorKind::NotFound | std::io::ErrorKind::PermissionDenied + ), + notify::ErrorKind::PathNotFound => false, + notify::ErrorKind::WatchNotFound => false, + notify::ErrorKind::InvalidConfig(_) => false, + notify::ErrorKind::MaxFilesWatch => true, + } +} + +impl WechatMonitor { + /// Start monitoring the encrypted session directory. + pub fn start(config: MonitorConfig) -> Result { + // 1 & 2. Open DB — direct encrypted or decrypt+cache + let (db, cache) = if let (Some(raw_key), Some(ref encrypted_root)) = + (config.raw_key, &config.encrypted_root) + { + let db = WechatDb::open_encrypted(encrypted_root, raw_key)?; + (db, None) + } else { + let mut cache = DecryptCache::new( + config.encrypted_session_dir.clone(), + config.key_material.clone(), + config.params, + )?; + cache.initial_decrypt()?; + let db = WechatDb::open(cache.decrypted_root())?; + (db, Some(cache)) + }; + + let initial_sessions = db + .query_sessions(&SessionQuery::new().limit(10_000)) + .map_err(MonitorError::Db)?; + + // 3. Initialize tracker with initial snapshot + let mut tracker = SessionTracker::new(); + tracker.diff(&initial_sessions.items); + + // 4. Shutdown channels + let (shutdown_tx, shutdown_rx) = tokio::sync::watch::channel(false); + let shutdown_flag = Arc::new(AtomicBool::new(false)); + + // 5. File watcher + event channel + let (file_tx, file_rx) = std::sync::mpsc::channel(); + + // Watcher stored as Box and moved into the task + let resolved = resolve_watch_mode(&config.watch_mode); + tracing::info!(watch_mode = ?config.watch_mode, resolved = ?resolved, "watcher mode selected"); + + let session_db = config.encrypted_session_dir.join("session.db"); + let session_wal = config.encrypted_session_dir.join("session.db-wal"); + let poll_interval = config.poll_interval; + + let make_polling = + |file_tx: std::sync::mpsc::Sender<_>| -> Result, MonitorError> { + let mut pw = PollingWatcher::new(poll_interval, file_tx, shutdown_flag.clone()); + pw.watch(&session_db)?; + pw.watch(&session_wal)?; + Ok(Box::new(pw)) + }; + + let watcher: Box = match resolved { + ResolvedWatcher::Polling => make_polling(file_tx)?, + ResolvedWatcher::Notify => match NotifyWatcher::new(file_tx.clone()) { + Ok(mut nw) => match nw.watch(&config.encrypted_session_dir) { + Ok(()) => Box::new(nw), + Err(ref e) if is_backend_unavailable_monitor(e) => { + tracing::warn!(error = %e, "notify watcher failed to watch directory, falling back to polling"); + make_polling(file_tx)? + } + Err(e) => return Err(e), + }, + Err(ref e) if is_backend_unavailable_monitor(e) => { + tracing::warn!(error = %e, "notify watcher init failed, falling back to polling"); + make_polling(file_tx)? + } + Err(e) => return Err(e), + }, + }; + + // 6. Output channel + let (event_tx, event_rx) = tokio::sync::mpsc::channel(config.channel_capacity); + + // 7. Spawn monitor loop — watcher is moved in to keep it alive + tracing::info!( + encrypted_session_dir = %config.encrypted_session_dir.display(), + poll_interval_ms = poll_interval.as_millis() as u64, + resolved_watcher = ?resolved, + "monitor loop starting" + ); + + let task = tokio::task::spawn_blocking(move || { + let mut db = db; + let mut cache = cache; + let shutdown_rx = shutdown_rx; + let _watcher = watcher; // held alive for the duration of the loop + + loop { + let got_event = file_rx.recv_timeout(Duration::from_millis(500)); + + if shutdown_rx.has_changed().unwrap_or(true) { + break; + } + + if got_event.is_ok() { + tracing::debug!("file event received"); + + if let Some(ref mut c) = cache { + // Decrypt-cache mode + match c.update() { + Ok(UpdateKind::WalPatched) => { + tracing::debug!(update = "WalPatched", "cache update result"); + if let Err(e) = db.reopen_sessions() { + tracing::warn!(error = %e, "reopen_sessions failed"); + continue; + } + } + Ok(UpdateKind::FullDecrypt) => { + tracing::debug!(update = "FullDecrypt", "cache update result"); + match WechatDb::open(c.decrypted_root()) { + Ok(new_db) => db = new_db, + Err(e) => { + tracing::warn!(error = %e, "WechatDb::open failed"); + continue; + } + } + } + Ok(UpdateKind::NoChange) => { + tracing::debug!(update = "NoChange", "cache update result"); + continue; + } + Err(e) => { + tracing::warn!(error = %e, "cache update failed"); + continue; + } + } + } else { + // Direct encrypted mode — just reopen session connection + tracing::debug!("direct mode: reopening session connection"); + if let Err(e) = db.reopen_sessions() { + tracing::warn!(error = %e, "reopen_sessions failed"); + continue; + } + } + + // Query sessions and emit events + let sessions = match db.query_sessions(&SessionQuery::new().limit(10_000)) { + Ok(r) => r, + Err(e) => { + tracing::warn!(error = %e, "query_sessions failed"); + continue; + } + }; + let events = tracker.diff(&sessions.items); + tracing::debug!(count = events.len(), "session events emitted"); + for ev in events { + if event_tx.blocking_send(ev).is_err() { + return; + } + } + } + } + }); + + Ok(Self { + receiver: Some(event_rx), + shutdown_tx: Some(shutdown_tx), + shutdown_flag: Some(shutdown_flag), + _task: Some(task), + }) + } + + /// Take the mpsc receiver, allowing an external task to consume events directly. + /// + /// After calling this, [`recv()`](Self::recv) will always return `None`. + pub fn take_receiver(&mut self) -> Option> { + self.receiver.take() + } + + /// Receive the next session event. + pub async fn recv(&mut self) -> Option { + self.receiver.as_mut()?.recv().await + } + + /// Signal the monitor to stop. + pub fn stop(&self) { + if let Some(tx) = &self.shutdown_tx { + let _ = tx.send(true); + } + if let Some(flag) = &self.shutdown_flag { + flag.store(true, Ordering::Relaxed); + } + } + + /// Convert into a `Stream` of session events, consuming the monitor handle. + pub fn into_stream(mut self) -> MonitorStream { + MonitorStream { + receiver: self.receiver.take(), + shutdown_tx: self.shutdown_tx.take(), + shutdown_flag: self.shutdown_flag.take(), + _task: self._task.take(), + } + } +} + +impl Drop for WechatMonitor { + fn drop(&mut self) { + if let Some(tx) = self.shutdown_tx.take() { + let _ = tx.send(true); + } + if let Some(flag) = self.shutdown_flag.take() { + flag.store(true, Ordering::Relaxed); + } + } +} + +/// A `Stream` of `SessionEvent`s. +/// +/// Created by [`WechatMonitor::into_stream()`]. Sends shutdown signal on drop. +pub struct MonitorStream { + receiver: Option>, + shutdown_tx: Option>, + shutdown_flag: Option>, + _task: Option>, +} + +impl Stream for MonitorStream { + type Item = SessionEvent; + + fn poll_next( + mut self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + match &mut self.receiver { + Some(rx) => rx.poll_recv(cx), + None => std::task::Poll::Ready(None), + } + } +} + +impl Drop for MonitorStream { + fn drop(&mut self) { + if let Some(tx) = self.shutdown_tx.take() { + let _ = tx.send(true); + } + if let Some(flag) = self.shutdown_flag.take() { + flag.store(true, Ordering::Relaxed); + } + } +} + +/// Check if a MonitorError wraps a backend-unavailable notify error. +fn is_backend_unavailable_monitor(err: &MonitorError) -> bool { + match err { + MonitorError::Watcher(notify_err) => is_backend_unavailable(notify_err), + _ => false, + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn watch_mode_poll_resolves_to_polling() { + assert_eq!( + resolve_watch_mode(&WatchMode::Poll), + ResolvedWatcher::Polling + ); + } + + #[test] + fn watch_mode_fsnotify_resolves_to_notify() { + assert_eq!( + resolve_watch_mode(&WatchMode::Fsnotify), + ResolvedWatcher::Notify + ); + } + + #[cfg(target_os = "macos")] + #[test] + fn watch_mode_auto_resolves_to_polling_on_macos() { + assert_eq!( + resolve_watch_mode(&WatchMode::Auto), + ResolvedWatcher::Polling + ); + } + + #[cfg(not(target_os = "macos"))] + #[test] + fn watch_mode_auto_resolves_to_notify_on_non_macos() { + assert_eq!( + resolve_watch_mode(&WatchMode::Auto), + ResolvedWatcher::Notify + ); + } +} diff --git a/crates/wx-monitor/src/tracker.rs b/crates/wx-monitor/src/tracker.rs new file mode 100644 index 0000000..bb8c6f4 --- /dev/null +++ b/crates/wx-monitor/src/tracker.rs @@ -0,0 +1,130 @@ +use std::collections::HashMap; + +use wx_db::Session; + +use crate::event::{SessionEvent, SessionEventKind}; + +/// Pure diff logic for detecting session changes. +/// +/// Maintains a snapshot of `(username → sort_timestamp)` and compares +/// incoming session lists to emit change events. +pub(crate) struct SessionTracker { + snapshot: HashMap, +} + +impl SessionTracker { + pub fn new() -> Self { + Self { + snapshot: HashMap::new(), + } + } + + /// Compare current sessions against the stored snapshot. + /// + /// Emits `Updated` for sessions that are new or have a newer `sort_timestamp`. + pub fn diff(&mut self, sessions: &[Session]) -> Vec { + let now = now_secs(); + let mut events = Vec::new(); + + for s in sessions { + let is_new_or_updated = match self.snapshot.get(&s.username) { + None => true, + Some(&old_ts) => s.sort_timestamp > old_ts, + }; + + if is_new_or_updated { + events.push(SessionEvent { + username: s.username.clone(), + sort_timestamp: s.sort_timestamp, + detected_at: now, + kind: SessionEventKind::Updated, + summary: s.summary.clone(), + last_msg_type: s.last_msg_type, + last_msg_sender: s.last_msg_sender.clone(), + last_sender_display_name: s.last_sender_display_name.clone(), + }); + } + } + + // Update snapshot with all entries + for s in sessions { + self.snapshot.insert(s.username.clone(), s.sort_timestamp); + } + + events + } +} + +fn now_secs() -> i64 { + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_secs() as i64 +} + +#[cfg(test)] +mod tests { + use super::*; + + fn session(username: &str, ts: i64) -> Session { + Session { + username: username.to_string(), + summary: String::new(), + sort_timestamp: ts, + last_msg_type: None, + last_msg_sender: None, + last_sender_display_name: None, + } + } + + #[test] + fn diff_new_session_emits_updated() { + let mut tracker = SessionTracker::new(); + let events = tracker.diff(&[session("alice", 100)]); + assert_eq!(events.len(), 1); + assert_eq!(events[0].username, "alice"); + assert!(matches!(events[0].kind, SessionEventKind::Updated)); + } + + #[test] + fn diff_updated_timestamp_emits_updated() { + let mut tracker = SessionTracker::new(); + tracker.diff(&[session("alice", 100)]); + + let events = tracker.diff(&[session("alice", 200)]); + assert_eq!(events.len(), 1); + assert_eq!(events[0].username, "alice"); + assert_eq!(events[0].sort_timestamp, 200); + } + + #[test] + fn diff_unchanged_emits_nothing() { + let mut tracker = SessionTracker::new(); + tracker.diff(&[session("alice", 100)]); + + let events = tracker.diff(&[session("alice", 100)]); + assert!(events.is_empty()); + } + + #[test] + fn diff_carries_content_fields() { + let mut tracker = SessionTracker::new(); + let s = Session { + username: "wxid_alice".to_string(), + summary: "hello world".to_string(), + sort_timestamp: 100, + last_msg_type: Some(1), + last_msg_sender: Some("wxid_sender".to_string()), + last_sender_display_name: Some("Sender Name".to_string()), + }; + let events = tracker.diff(&[s]); + assert_eq!(events.len(), 1); + assert_eq!(events[0].summary, "hello world"); + assert_eq!(events[0].last_msg_type, Some(1)); + assert_eq!(events[0].last_msg_sender.as_deref(), Some("wxid_sender")); + assert_eq!( + events[0].last_sender_display_name.as_deref(), + Some("Sender Name") + ); + } +} diff --git a/crates/wx-monitor/src/watcher.rs b/crates/wx-monitor/src/watcher.rs new file mode 100644 index 0000000..a874521 --- /dev/null +++ b/crates/wx-monitor/src/watcher.rs @@ -0,0 +1,261 @@ +use std::collections::HashMap; +use std::path::{Path, PathBuf}; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::sync::mpsc::Sender; +use std::sync::{Arc, Mutex}; +use std::time::{Duration, SystemTime}; + +use notify::{EventKind, RecommendedWatcher, RecursiveMode, Watcher}; + +use crate::error::MonitorError; +use crate::event::{FileEvent, FileEventKind}; + +pub(crate) trait FileWatcher: Send + 'static { + fn watch(&mut self, path: &Path) -> Result<(), MonitorError>; +} + +// --------------------------------------------------------------------------- +// NotifyWatcher +// --------------------------------------------------------------------------- + +pub(crate) struct NotifyWatcher { + inner: RecommendedWatcher, +} + +impl NotifyWatcher { + pub fn new(tx: Sender) -> Result { + let watcher = RecommendedWatcher::new( + move |res: Result| match res { + Ok(event) => { + let kind = match event.kind { + EventKind::Create(_) => FileEventKind::Created, + EventKind::Modify(_) => FileEventKind::Modified, + EventKind::Remove(_) => { + tracing::debug!(event_kind = ?event.kind, paths = ?event.paths, "remove event, treating as Modified"); + FileEventKind::Modified + } + _ => { + tracing::debug!(event_kind = ?event.kind, paths = ?event.paths, "ignoring event"); + return; + } + }; + tracing::debug!(event_kind = ?event.kind, paths = ?event.paths, "file event received"); + for path in event.paths { + let _ = tx.send(FileEvent { + path, + kind: match kind { + FileEventKind::Created => FileEventKind::Created, + FileEventKind::Modified => FileEventKind::Modified, + }, + }); + } + } + Err(e) => { + tracing::warn!(error = %e, "notify watcher error"); + } + }, + notify::Config::default(), + )?; + Ok(Self { inner: watcher }) + } +} + +impl FileWatcher for NotifyWatcher { + fn watch(&mut self, path: &Path) -> Result<(), MonitorError> { + tracing::info!(path = %path.display(), "notify watcher: watching path"); + self.inner.watch(path, RecursiveMode::NonRecursive)?; + Ok(()) + } +} + +// --------------------------------------------------------------------------- +// PollingWatcher +// --------------------------------------------------------------------------- + +pub(crate) struct PollingWatcher { + interval: Duration, + tx: Sender, + shutdown: Arc, + paths: Arc>>>, + thread: Option>, +} + +impl PollingWatcher { + pub fn new(interval: Duration, tx: Sender, shutdown: Arc) -> Self { + Self { + interval, + tx, + shutdown, + paths: Arc::new(Mutex::new(HashMap::new())), + thread: None, + } + } + + fn ensure_thread(&mut self) { + if self.thread.is_some() { + return; + } + + let interval = self.interval; + let tx = self.tx.clone(); + let shutdown = self.shutdown.clone(); + let paths = self.paths.clone(); + + let handle = std::thread::spawn(move || loop { + if shutdown.load(Ordering::Relaxed) { + break; + } + std::thread::sleep(interval); + if shutdown.load(Ordering::Relaxed) { + break; + } + + tracing::trace!("polling cycle"); + let mut tracked = paths.lock().unwrap(); + for (path, prev_mtime) in tracked.iter_mut() { + let current = std::fs::metadata(path).ok().and_then(|m| m.modified().ok()); + + match (prev_mtime.as_ref(), current) { + (None, Some(mtime)) => { + *prev_mtime = Some(mtime); + tracing::debug!(path = %path.display(), "file created"); + let _ = tx.send(FileEvent { + path: path.clone(), + kind: FileEventKind::Created, + }); + } + (Some(old), Some(new)) if new > *old => { + tracing::debug!(path = %path.display(), old_mtime = ?old, new_mtime = ?new, "file modified"); + *prev_mtime = Some(new); + let _ = tx.send(FileEvent { + path: path.clone(), + kind: FileEventKind::Modified, + }); + } + (Some(_), None) => { + *prev_mtime = None; + } + _ => {} + } + } + }); + + self.thread = Some(handle); + } +} + +impl FileWatcher for PollingWatcher { + fn watch(&mut self, path: &Path) -> Result<(), MonitorError> { + let initial_mtime = std::fs::metadata(path).ok().and_then(|m| m.modified().ok()); + tracing::info!(path = %path.display(), initial_mtime = ?initial_mtime, "polling watcher: watching path"); + + self.paths + .lock() + .unwrap() + .insert(path.to_path_buf(), initial_mtime); + + self.ensure_thread(); + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::mpsc; + + #[test] + fn polling_watcher_detects_modification() { + let dir = tempfile::TempDir::new().unwrap(); + let file = dir.path().join("test.db"); + std::fs::write(&file, b"initial").unwrap(); + + std::thread::sleep(Duration::from_millis(50)); + + let (tx, rx) = mpsc::channel(); + let shutdown = Arc::new(AtomicBool::new(false)); + + let mut watcher = PollingWatcher::new(Duration::from_millis(100), tx, shutdown.clone()); + watcher.watch(&file).unwrap(); + + std::thread::sleep(Duration::from_millis(50)); + std::fs::write(&file, b"modified content").unwrap(); + + let event = rx.recv_timeout(Duration::from_secs(3)).unwrap(); + assert_eq!(event.path, file); + assert!(matches!(event.kind, FileEventKind::Modified)); + + shutdown.store(true, Ordering::Relaxed); + } + + #[test] + fn polling_watcher_detects_creation() { + let dir = tempfile::TempDir::new().unwrap(); + let file = dir.path().join("new.db-wal"); + + let (tx, rx) = mpsc::channel(); + let shutdown = Arc::new(AtomicBool::new(false)); + + let mut watcher = PollingWatcher::new(Duration::from_millis(100), tx, shutdown.clone()); + watcher.watch(&file).unwrap(); + + std::thread::sleep(Duration::from_millis(50)); + std::fs::write(&file, b"wal data").unwrap(); + + let event = rx.recv_timeout(Duration::from_secs(3)).unwrap(); + assert_eq!(event.path, file); + assert!(matches!(event.kind, FileEventKind::Created)); + + shutdown.store(true, Ordering::Relaxed); + } + + #[test] + fn polling_watcher_tracks_paths_added_after_thread_start() { + let dir = tempfile::TempDir::new().unwrap(); + let file1 = dir.path().join("session.db"); + let file2 = dir.path().join("session.db-wal"); + std::fs::write(&file1, b"db initial").unwrap(); + + let (tx, rx) = mpsc::channel(); + let shutdown = Arc::new(AtomicBool::new(false)); + + let mut watcher = PollingWatcher::new(Duration::from_millis(100), tx, shutdown.clone()); + + // First watch starts the thread + watcher.watch(&file1).unwrap(); + // Second watch adds path after thread is running + watcher.watch(&file2).unwrap(); + + std::thread::sleep(Duration::from_millis(50)); + + // Modify the second file (added after thread start) + std::fs::write(&file2, b"wal data").unwrap(); + + let event = rx.recv_timeout(Duration::from_secs(3)).unwrap(); + assert_eq!(event.path, file2); + assert!(matches!(event.kind, FileEventKind::Created)); + + shutdown.store(true, Ordering::Relaxed); + } + + #[test] + #[ignore] // FSEvents on macOS temp dirs can be slow/flaky in CI + fn notify_watcher_detects_modification() { + let dir = tempfile::TempDir::new().unwrap(); + let file = dir.path().join("test.db"); + std::fs::write(&file, b"initial").unwrap(); + + let (tx, rx) = mpsc::channel(); + + let mut watcher = NotifyWatcher::new(tx).unwrap(); + watcher.watch(dir.path()).unwrap(); + + std::thread::sleep(Duration::from_millis(200)); + std::fs::write(&file, b"modified").unwrap(); + + let event = rx.recv_timeout(Duration::from_secs(3)).unwrap(); + let expected = file.canonicalize().unwrap(); + let actual = event.path.canonicalize().unwrap(); + assert_eq!(actual, expected); + } +} diff --git a/crates/wx-monitor/tests/integration.rs b/crates/wx-monitor/tests/integration.rs new file mode 100644 index 0000000..0fcf94a --- /dev/null +++ b/crates/wx-monitor/tests/integration.rs @@ -0,0 +1,296 @@ +use std::path::{Path, PathBuf}; +use std::time::Duration; + +use rusqlite::{params, Connection}; +use wx_decrypt::{KeyMaterial, MACOS_4_1_7_31}; +use wx_monitor::{MonitorConfig, WechatMonitor}; + +// ---- crypto helpers (standalone, matching wx-decrypt internals) ---- + +fn derive_enc_key(raw_key: &[u8; 32], salt: &[u8; 16]) -> [u8; 32] { + let mut key = [0u8; 32]; + pbkdf2::pbkdf2_hmac::(raw_key, salt, MACOS_4_1_7_31.kdf_iter, &mut key); + key +} + +fn derive_mac_key(enc_key: &[u8; 32], salt: &[u8; 16]) -> [u8; 32] { + let mut mac_salt = [0u8; 16]; + for (i, b) in salt.iter().enumerate() { + mac_salt[i] = b ^ 0x3a; + } + let mut key = [0u8; 32]; + pbkdf2::pbkdf2_hmac::(enc_key, &mac_salt, 2, &mut key); + key +} + +/// Encrypt a single page from a plaintext SQLite file. +/// +/// - `page_data`: raw 4096 bytes from the plaintext SQLite file +/// - `page_num`: 0-indexed page number +/// - `salt`: 16-byte salt (required for page 0, ignored for others) +fn encrypt_page( + page_data: &[u8], + enc_key: &[u8; 32], + mac_key: &[u8; 32], + page_num: u32, + salt: &[u8; 16], +) -> Vec { + use aes::cipher::{BlockEncryptMut, KeyIvInit}; + use hmac::{Hmac, Mac}; + use sha2::Sha512; + + let params = &MACOS_4_1_7_31; + let iv: [u8; 16] = [0x42; 16]; + let offset = if page_num == 0 { 16 } else { 0 }; + let data_size = params.page_size - params.reserve - offset; // 4000 for page 0, 4016 for others + + // Extract plaintext from the original page (skip SQLite header for page 0, skip reserved area) + let plaintext = &page_data[offset..offset + data_size]; + + // Encrypt with AES-256-CBC + type Aes256CbcEnc = cbc::Encryptor; + let mut ciphertext = plaintext.to_vec(); + let encryptor = Aes256CbcEnc::new(enc_key.into(), (&iv).into()); + encryptor + .encrypt_padded_mut::(&mut ciphertext, data_size) + .unwrap(); + + // Assemble encrypted page + let mut page = Vec::with_capacity(params.page_size); + if page_num == 0 { + page.extend_from_slice(salt); + } + page.extend_from_slice(&ciphertext); + // Reserve area: IV + HMAC + page.extend_from_slice(&iv); + page.resize(params.page_size, 0); // zero-fill HMAC placeholder + + // Compute HMAC + let hmac_data_end = params.page_size - params.reserve + params.iv_size; + let mut mac = as Mac>::new_from_slice(mac_key).unwrap(); + mac.update(&page[offset..hmac_data_end]); + mac.update(&(page_num + 1).to_le_bytes()); // 1-indexed + let hmac_result = mac.finalize().into_bytes(); + + let hmac_start = params.page_size - params.reserve + params.iv_size; + page[hmac_start..hmac_start + params.hmac_size] + .copy_from_slice(&hmac_result[..params.hmac_size]); + + page +} + +/// Create a valid SQLite session.db with reserved_page_size=80, then encrypt it. +/// +/// Returns the path to the encrypted file. +fn create_encrypted_session_db( + dir: &Path, + raw_key: &[u8; 32], + sessions: &[(&str, i64, &str)], +) -> PathBuf { + let salt: [u8; 16] = [0x01; 16]; + + // 1. Create a valid SQLite DB with reserved_page_size=80 + let plain_path = dir.join("session_plain.db"); + { + let conn = Connection::open(&plain_path).unwrap(); + conn.execute_batch("PRAGMA page_size = 4096;").unwrap(); + + // Set reserved bytes via sqlite3_file_control + unsafe { + let mut reserve: i32 = 80; + let rc = rusqlite::ffi::sqlite3_file_control( + conn.handle(), + c"main".as_ptr(), + 38, // SQLITE_FCNTL_RESERVE_BYTES + &mut reserve as *mut _ as *mut std::ffi::c_void, + ); + assert_eq!(rc, 0, "sqlite3_file_control failed"); + } + + conn.execute_batch( + "CREATE TABLE SessionTable ( + username TEXT, + sort_timestamp INTEGER, + summary TEXT, + last_msg_type INTEGER, + last_msg_sender TEXT, + last_sender_display_name TEXT + );", + ) + .unwrap(); + + for (username, ts, summary) in sessions { + conn.execute( + "INSERT INTO SessionTable VALUES (?1, ?2, ?3, NULL, NULL, NULL)", + params![username, ts, summary], + ) + .unwrap(); + } + } + + // 2. Read raw bytes and verify reserved=80 + let plain_data = std::fs::read(&plain_path).unwrap(); + assert_eq!( + plain_data[20], 80, + "reserved_page_size should be 80, got {}", + plain_data[20] + ); + let page_count = plain_data.len() / 4096; + assert!(page_count >= 1, "expected at least 1 page"); + + // 3. Derive keys + let enc_key = derive_enc_key(raw_key, &salt); + let mac_key = derive_mac_key(&enc_key, &salt); + + // 4. Encrypt each page + let enc_path = dir.join("session.db"); + let mut enc_data = Vec::with_capacity(plain_data.len()); + for i in 0..page_count { + let start = i * 4096; + let page = &plain_data[start..start + 4096]; + let encrypted = encrypt_page(page, &enc_key, &mac_key, i as u32, &salt); + enc_data.extend_from_slice(&encrypted); + } + std::fs::write(&enc_path, &enc_data).unwrap(); + + // Clean up plain file + let _ = std::fs::remove_file(&plain_path); + + enc_path +} + +// ---- integration test ---- + +// Slow by design: exercises real polling + PBKDF2(256k) decrypt flow and +// routinely takes tens of seconds in debug builds. Keep it opt-in unless +// explicitly validating monitor/decrypt integration. +#[tokio::test] +#[ignore = "slow integration test; runs real PBKDF2/decrypt path"] +async fn monitor_detects_session_change() { + let dir = tempfile::TempDir::new().unwrap(); + let session_dir = dir.path().to_path_buf(); + + let raw_key: [u8; 32] = [0xAB; 32]; + + // Create initial encrypted session.db with one session + create_encrypted_session_db( + &session_dir, + &raw_key, + &[("wxid_alice", 1000, "hello from alice")], + ); + + // Start monitor with polling (200ms interval) + let mut monitor = WechatMonitor::start(MonitorConfig { + encrypted_session_dir: session_dir.clone(), + key_material: KeyMaterial::RawKey(raw_key), + params: &MACOS_4_1_7_31, + watch_mode: wx_monitor::WatchMode::Poll, + poll_interval: Duration::from_millis(200), + channel_capacity: 100, + raw_key: None, + encrypted_root: None, + }) + .expect("monitor should start"); + + // Wait for initial setup to stabilize + tokio::time::sleep(Duration::from_millis(500)).await; + + // Overwrite with updated data (new session added) + // Need mtime to change, so sleep briefly + std::thread::sleep(Duration::from_millis(1100)); + create_encrypted_session_db( + &session_dir, + &raw_key, + &[ + ("wxid_alice", 1000, "hello from alice"), + ("wxid_bob", 2000, "hello from bob"), + ], + ); + + // Wait for event (up to 10 seconds, accounting for PBKDF2 overhead) + // Plan specifies 5s; PBKDF2 256k iterations takes ~1s release / ~9s debug per call. + // update() triggers a second PBKDF2, so debug needs ~20s total. + let event = tokio::time::timeout(Duration::from_secs(25), monitor.recv()) + .await + .expect("should receive event within timeout") + .expect("event should not be None"); + + // Should be an Updated event for wxid_bob (the new session) + assert_eq!(event.username, "wxid_bob"); + assert!(matches!( + event.kind, + wx_monitor::SessionEventKind::Updated + )); + + // Stop monitor and assert clean exit + monitor.stop(); + + // Drain any remaining events, then recv must return None (task exited, channel closed). + // The monitor loop checks shutdown every 500ms, so 5s is generous. + loop { + let result = tokio::time::timeout(Duration::from_secs(1), monitor.recv()) + .await + .expect("monitor task should exit within timeout after stop()"); + match result { + Some(_) => continue, // drain buffered events + None => break, // channel closed — task exited cleanly + } + } +} + +// Slow by design: re-encrypts the same DB and waits long enough to prove the +// monitor does not flood reset events after a full decrypt. Skip by default +// because the PBKDF2/debug path makes this a tens-of-seconds test. +#[tokio::test] +#[ignore = "slow integration test; runs real PBKDF2/decrypt path"] +async fn full_decrypt_does_not_flood_resets() { + let dir = tempfile::TempDir::new().unwrap(); + let session_dir = dir.path().to_path_buf(); + + let raw_key: [u8; 32] = [0xAB; 32]; + + // Create initial encrypted session.db with 5 sessions + let sessions: Vec<(&str, i64, &str)> = vec![ + ("wxid_alice", 1000, "hello alice"), + ("wxid_bob", 2000, "hello bob"), + ("wxid_charlie", 3000, "hello charlie"), + ("wxid_dave", 4000, "hello dave"), + ("wxid_eve", 5000, "hello eve"), + ]; + create_encrypted_session_db(&session_dir, &raw_key, &sessions); + + // Start monitor with polling (200ms interval) + let mut monitor = WechatMonitor::start(MonitorConfig { + encrypted_session_dir: session_dir.clone(), + key_material: KeyMaterial::RawKey(raw_key), + params: &MACOS_4_1_7_31, + watch_mode: wx_monitor::WatchMode::Poll, + poll_interval: Duration::from_millis(200), + channel_capacity: 100, + raw_key: None, + encrypted_root: None, + }) + .expect("monitor should start"); + + // Wait for initial setup to fully stabilize + tokio::time::sleep(Duration::from_secs(2)).await; + + // Drain any buffered events from initialization + while let Ok(Some(_)) = tokio::time::timeout(Duration::from_millis(100), monitor.recv()).await { + } + + // Re-encrypt the same session.db with identical data (simulates WAL checkpoint) + std::thread::sleep(Duration::from_millis(1100)); + create_encrypted_session_db(&session_dir, &raw_key, &sessions); + + // Wait long enough for FullDecrypt to complete (PBKDF2 ~20s in debug mode) + // then verify no events arrived. With the old reset() code, 5 events would + // arrive after PBKDF2 completes. With diff(), zero events are produced. + let result = tokio::time::timeout(Duration::from_secs(30), monitor.recv()).await; + assert!( + result.is_err(), + "expected no events after re-encrypting identical data, but got one — FullDecrypt is still flooding" + ); + + monitor.stop(); +} diff --git a/crates/wx-paths/Cargo.toml b/crates/wx-paths/Cargo.toml new file mode 100644 index 0000000..4d40ae2 --- /dev/null +++ b/crates/wx-paths/Cargo.toml @@ -0,0 +1,15 @@ +[package] +name = "wx-paths" +version.workspace = true +edition.workspace = true + +[dependencies] +dirs = "6" +thiserror = "2" +serde = { version = "1", features = ["derive"] } + +[target.'cfg(unix)'.dependencies] +libc = "0.2" + +[dev-dependencies] +tempfile = "3" diff --git a/crates/wx-paths/src/lib.rs b/crates/wx-paths/src/lib.rs new file mode 100644 index 0000000..6ab85e7 --- /dev/null +++ b/crates/wx-paths/src/lib.rs @@ -0,0 +1,484 @@ +mod migration; +mod platform; +pub mod sudo; + +use std::io; +use std::path::{Path, PathBuf}; + +use platform::PlatformBaseDirs; + +/// Centralized path resolution for wx-cli. +/// +/// All implicit/system-managed paths (config, cache, state, logs, temp) are +/// resolved through this struct. User-specified output paths (decrypt --output, +/// export, media) are out of scope. +#[derive(Clone, Debug)] +pub struct AppPaths { + home: PathBuf, + config_root: PathBuf, + cache_root: PathBuf, + state_root: PathBuf, + logs_root: PathBuf, + runtime_root_override: Option, +} + +#[derive(Debug, serde::Serialize)] +pub struct PathsSummary { + pub platform: &'static str, + pub config_dir: PathBuf, + pub keys_file: PathBuf, + pub settings_file: PathBuf, + pub cache_root: PathBuf, + pub state_root: PathBuf, + pub logs_dir: PathBuf, + pub server_state_dir: PathBuf, + pub server_stdout_log: PathBuf, + pub server_stderr_log: PathBuf, + pub temp_root: PathBuf, +} + +#[derive(Debug, thiserror::Error)] +pub enum PathsError { + #[error("cannot determine home directory")] + NoHome, + #[error("cannot determine config directory")] + NoConfig, + #[error("cannot determine cache directory")] + NoCache, + #[error("cannot determine state directory")] + NoState, + #[error("cannot determine log directory")] + NoLog, +} + +impl AppPaths { + /// Create a new `AppPaths` by resolving the real user's home directory. + /// + /// Under `sudo`, resolves the original user's home via `SUDO_USER` + `getpwnam`. + pub fn new() -> Result { + let home = sudo::resolve_real_home()?; + let dirs = PlatformBaseDirs::resolve(&home)?; + + Ok(Self { + home, + config_root: dirs.config_root, + cache_root: dirs.cache_root, + state_root: dirs.state_root, + logs_root: dirs.logs_root, + runtime_root_override: None, + }) + } + + /// Create `AppPaths` with a runtime root override for server commands. + /// + /// When specified, ALL server runtime files (state, config, lock, logs) + /// are placed under the given root instead of their platform-default locations. + pub fn with_runtime_root(root: PathBuf) -> Result { + let mut ap = Self::new()?; + ap.runtime_root_override = Some(root); + Ok(ap) + } + + /// The resolved home directory. + pub fn home(&self) -> &Path { + &self.home + } + + // ── Config ── + + /// Config directory. + /// + /// macOS: `~/Library/Application Support/wx-cli/config/` + /// Linux: `~/.config/wx-cli/` + pub fn config_dir(&self) -> PathBuf { + self.config_root.clone() + } + + /// `/keys.toml` + pub fn keys_file(&self) -> PathBuf { + self.config_root.join("keys.toml") + } + + /// `/settings.toml` + pub fn settings_file(&self) -> PathBuf { + self.config_root.join("settings.toml") + } + + // ── Cache ── + + /// Cache root: `~/Library/Caches/wx-cli/` (macOS) + pub fn cache_root(&self) -> &Path { + &self.cache_root + } + + /// `//` + pub fn account_cache_dir(&self, id: &str) -> PathBuf { + self.cache_root.join(id) + } + + /// `//db_storage/` + pub fn account_db_cache_dir(&self, id: &str) -> PathBuf { + self.cache_root.join(id).join("db_storage") + } + + // ── State ── + + /// State root: `~/Library/Application Support/wx-cli/state/` (macOS) + pub fn state_root(&self) -> &Path { + &self.state_root + } + + /// Server state directory. + /// + /// Default: `/server/` + /// With `--runtime-root`: `/` + pub fn server_state_dir(&self) -> PathBuf { + match &self.runtime_root_override { + Some(root) => root.clone(), + None => self.state_root.join("server"), + } + } + + /// Server lock file: `/manager.lock` + pub fn server_lock_file(&self) -> PathBuf { + self.server_state_dir().join("manager.lock") + } + + /// Server config file: `/config.json` + pub fn server_config_file(&self) -> PathBuf { + self.server_state_dir().join("config.json") + } + + /// Server state file: `/state.json` + pub fn server_state_file(&self) -> PathBuf { + self.server_state_dir().join("state.json") + } + + // ── Logs ── + + /// Logs directory: `~/Library/Logs/wx-cli/` (macOS) + pub fn logs_dir(&self) -> &Path { + &self.logs_root + } + + /// Server stdout log. + /// + /// Default: `/server/stdout.log` + /// With `--runtime-root`: `/stdout.log` + pub fn server_stdout_log(&self) -> PathBuf { + match &self.runtime_root_override { + Some(root) => root.join("stdout.log"), + None => self.logs_root.join("server").join("stdout.log"), + } + } + + /// Server stderr log. + /// + /// Default: `/server/stderr.log` + /// With `--runtime-root`: `/stderr.log` + pub fn server_stderr_log(&self) -> PathBuf { + match &self.runtime_root_override { + Some(root) => root.join("stderr.log"), + None => self.logs_root.join("server").join("stderr.log"), + } + } + + // ── Server directories ── + + /// Ensure server state and log directories exist. + pub fn ensure_server_dirs(&self) -> io::Result<()> { + std::fs::create_dir_all(self.server_state_dir())?; + match &self.runtime_root_override { + Some(_) => {} // logs go in the same dir, already created + None => { + std::fs::create_dir_all(self.logs_root.join("server"))?; + } + } + Ok(()) + } + + // ── Temp (associated fns — system-level, no &self) ── + + /// Temp root: `std::env::temp_dir()/wx-cli/` + pub fn temp_root() -> PathBuf { + std::env::temp_dir().join("wx-cli") + } + + /// `/lldb/wechat_capture_key.py` + pub fn lldb_script_file() -> PathBuf { + Self::temp_root().join("lldb").join("wechat_capture_key.py") + } + + /// `/lldb/wechat_lldb_output.txt` + pub fn lldb_output_file() -> PathBuf { + Self::temp_root().join("lldb").join("wechat_lldb_output.txt") + } + + /// `/nickname/_.db` + pub fn nickname_temp_db(pid: u32, nanos: u128) -> PathBuf { + Self::temp_root() + .join("nickname") + .join(format!("{pid}_{nanos}.db")) + } + + // ── Utility ── + + /// Create directory and all parents, returning the path on success. + pub fn ensure_dir(path: &Path) -> io::Result<&Path> { + std::fs::create_dir_all(path)?; + Ok(path) + } + + /// One-time config migration from legacy `~/.config/wechat-utils/`. + /// + /// Migrates `keys.toml` and `settings.toml` to the new platform-correct + /// config directory. Idempotent via sentinel file. Called by + /// `KeyStore::load_default()` and `Settings::load_default()`. + pub fn migrate_config(&self) -> Result<(), io::Error> { + let _ = migration::ensure_config_migrated(&self.home, &self.config_root)?; + Ok(()) + } + + /// Current platform identifier. + pub fn platform() -> &'static str { + #[cfg(target_os = "macos")] + { "macos" } + #[cfg(target_os = "linux")] + { "linux" } + #[cfg(target_os = "windows")] + { "windows" } + #[cfg(not(any(target_os = "macos", target_os = "linux", target_os = "windows")))] + { "unknown" } + } + + /// Build a summary of all paths. + pub fn summary(&self) -> PathsSummary { + PathsSummary { + platform: Self::platform(), + config_dir: self.config_dir(), + keys_file: self.keys_file(), + settings_file: self.settings_file(), + cache_root: self.cache_root.clone(), + state_root: self.state_root.to_path_buf(), + logs_dir: self.logs_root.clone(), + server_state_dir: self.server_state_dir(), + server_stdout_log: self.server_stdout_log(), + server_stderr_log: self.server_stderr_log(), + temp_root: Self::temp_root(), + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn app_paths_new_returns_sensible_paths() { + let ap = AppPaths::new().expect("AppPaths::new() should succeed in test env"); + assert!(ap.home().is_absolute()); + assert!(ap.config_dir().is_absolute()); + assert!(ap.keys_file().is_absolute()); + assert!(ap.cache_root().is_absolute()); + assert!(ap.state_root().is_absolute()); + assert!(ap.logs_dir().is_absolute()); + } + + #[test] + fn config_dir_uses_wx_cli_namespace() { + let ap = AppPaths::new().unwrap(); + let config = ap.config_dir(); + assert!( + config.to_str().unwrap().contains("wx-cli"), + "config_dir should contain wx-cli: {:?}", + config + ); + } + + #[test] + fn temp_root_uses_wx_cli_namespace() { + let temp = AppPaths::temp_root(); + assert!( + temp.to_str().unwrap().ends_with("wx-cli"), + "temp_root should end with wx-cli: {:?}", + temp + ); + } + + #[test] + fn lldb_files_under_temp_root() { + let script = AppPaths::lldb_script_file(); + let output = AppPaths::lldb_output_file(); + assert!(script.starts_with(AppPaths::temp_root())); + assert!(output.starts_with(AppPaths::temp_root())); + assert!(script.to_str().unwrap().ends_with("wechat_capture_key.py")); + assert!(output.to_str().unwrap().ends_with("wechat_lldb_output.txt")); + } + + #[test] + fn nickname_temp_db_under_temp_root() { + let db = AppPaths::nickname_temp_db(1234, 9999); + assert!(db.starts_with(AppPaths::temp_root())); + assert!(db.to_str().unwrap().contains("nickname")); + assert!(db.to_str().unwrap().ends_with("1234_9999.db")); + } + + #[test] + fn keys_file_ends_with_keys_toml() { + let ap = AppPaths::new().unwrap(); + assert!(ap.keys_file().ends_with("keys.toml")); + } + + #[test] + fn settings_file_ends_with_settings_toml() { + let ap = AppPaths::new().unwrap(); + assert!(ap.settings_file().ends_with("settings.toml")); + } + + #[test] + fn account_cache_dir_contains_account_id() { + let ap = AppPaths::new().unwrap(); + let dir = ap.account_cache_dir("wxid_test_ab12"); + assert!(dir.ends_with("wxid_test_ab12")); + } + + #[test] + fn account_db_cache_dir_contains_db_storage() { + let ap = AppPaths::new().unwrap(); + let dir = ap.account_db_cache_dir("wxid_test_ab12"); + assert!(dir.ends_with("wxid_test_ab12/db_storage")); + } + + #[test] + fn server_state_dir_under_state_root() { + let ap = AppPaths::new().unwrap(); + assert!(ap.server_state_dir().starts_with(ap.state_root())); + assert!(ap.server_state_dir().ends_with("server")); + } + + #[test] + fn server_lock_file_under_server_state() { + let ap = AppPaths::new().unwrap(); + assert!(ap.server_lock_file().starts_with(&ap.server_state_dir())); + assert!(ap.server_lock_file().ends_with("manager.lock")); + } + + #[test] + fn server_config_file_under_server_state() { + let ap = AppPaths::new().unwrap(); + assert!(ap.server_config_file().starts_with(&ap.server_state_dir())); + assert!(ap.server_config_file().ends_with("config.json")); + } + + #[test] + fn server_state_file_under_server_state() { + let ap = AppPaths::new().unwrap(); + assert!(ap.server_state_file().starts_with(&ap.server_state_dir())); + assert!(ap.server_state_file().ends_with("state.json")); + } + + #[test] + fn server_logs_under_logs_dir_default() { + let ap = AppPaths::new().unwrap(); + assert!(ap.server_stdout_log().to_str().unwrap().contains("server")); + assert!(ap.server_stderr_log().to_str().unwrap().contains("server")); + assert!(ap.server_stdout_log().ends_with("stdout.log")); + assert!(ap.server_stderr_log().ends_with("stderr.log")); + } + + #[test] + fn with_runtime_root_overrides_server_paths() { + let root = PathBuf::from("/tmp/test-runtime"); + let ap = AppPaths::with_runtime_root(root.clone()).unwrap(); + assert_eq!(ap.server_state_dir(), root); + assert_eq!(ap.server_lock_file(), root.join("manager.lock")); + assert_eq!(ap.server_config_file(), root.join("config.json")); + assert_eq!(ap.server_state_file(), root.join("state.json")); + assert_eq!(ap.server_stdout_log(), root.join("stdout.log")); + assert_eq!(ap.server_stderr_log(), root.join("stderr.log")); + } + + #[test] + fn platform_is_known() { + let p = AppPaths::platform(); + assert!( + ["macos", "linux", "windows", "unknown"].contains(&p), + "unexpected platform: {}", + p + ); + } + + #[test] + fn summary_has_all_fields() { + let ap = AppPaths::new().unwrap(); + let s = ap.summary(); + assert_eq!(s.platform, AppPaths::platform()); + assert_eq!(s.config_dir, ap.config_dir()); + assert_eq!(s.keys_file, ap.keys_file()); + assert_eq!(s.settings_file, ap.settings_file()); + assert_eq!(s.cache_root, ap.cache_root().to_path_buf()); + assert_eq!(s.state_root, ap.state_root().to_path_buf()); + assert_eq!(s.logs_dir, ap.logs_dir().to_path_buf()); + assert_eq!(s.server_state_dir, ap.server_state_dir()); + assert_eq!(s.temp_root, AppPaths::temp_root()); + } + + #[test] + fn ensure_dir_creates_and_returns() { + let tmp = std::env::temp_dir().join("wx_paths_test_ensure"); + let _ = std::fs::remove_dir_all(&tmp); + let result = AppPaths::ensure_dir(&tmp); + assert!(result.is_ok()); + assert!(tmp.is_dir()); + let _ = std::fs::remove_dir_all(&tmp); + } + + #[cfg(target_os = "macos")] + mod macos_tests { + use super::*; + + #[test] + fn macos_config_under_application_support() { + let ap = AppPaths::new().unwrap(); + let config = ap.config_dir(); + assert!( + config.to_str().unwrap().contains("Application Support/wx-cli/config"), + "macOS config should be under Application Support: {:?}", + config + ); + } + + #[test] + fn macos_cache_under_library_caches() { + let ap = AppPaths::new().unwrap(); + let cache = ap.cache_root(); + assert!( + cache.to_str().unwrap().contains("Library/Caches/wx-cli"), + "macOS cache should be under Library/Caches: {:?}", + cache + ); + } + + #[test] + fn macos_state_under_application_support() { + let ap = AppPaths::new().unwrap(); + let state = ap.state_root(); + assert!( + state.to_str().unwrap().contains("Application Support/wx-cli/state"), + "macOS state should be under Application Support: {:?}", + state + ); + } + + #[test] + fn macos_logs_under_library_logs() { + let ap = AppPaths::new().unwrap(); + let logs = ap.logs_dir(); + assert!( + logs.to_str().unwrap().contains("Library/Logs/wx-cli"), + "macOS logs should be under Library/Logs: {:?}", + logs + ); + } + } +} diff --git a/crates/wx-paths/src/migration.rs b/crates/wx-paths/src/migration.rs new file mode 100644 index 0000000..2892a6e --- /dev/null +++ b/crates/wx-paths/src/migration.rs @@ -0,0 +1,178 @@ +use std::path::Path; + +use crate::sudo::chown_to_sudo_user; + +const OLD_CONFIG_DIR_NAME: &str = "wechat-utils"; +const SENTINEL_FILE: &str = ".migrated"; + +pub(crate) enum MigrationOutcome { + AlreadyMigrated, + NoLegacyConfig, + Migrated, +} + +/// Migrate config files from legacy `~/.config/wechat-utils/` to the new +/// platform-correct config directory. +/// +/// Strategy: **move then delete originals**. Files are first copied to the +/// new location, then originals are deleted. If copy fails, originals are +/// left untouched. If delete fails, both copies exist but new location wins. +/// +/// Sentinel is only written when all operations succeed. If any step +/// fails, returns an error and the sentinel is NOT written, allowing retry. +pub(crate) fn ensure_config_migrated( + home: &Path, + new_config_root: &Path, +) -> Result { + let sentinel = new_config_root.join(SENTINEL_FILE); + if sentinel.exists() { + return Ok(MigrationOutcome::AlreadyMigrated); + } + + let old_config = home.join(".config").join(OLD_CONFIG_DIR_NAME); + if !old_config.exists() { + std::fs::create_dir_all(new_config_root)?; + chown_config_tree(home, new_config_root); + std::fs::File::create(&sentinel)?; + chown_to_sudo_user(&sentinel); + return Ok(MigrationOutcome::NoLegacyConfig); + } + + std::fs::create_dir_all(new_config_root)?; + chown_config_tree(home, new_config_root); + + for file in &["keys.toml", "settings.toml"] { + let src = old_config.join(file); + let dst = new_config_root.join(file); + if src.exists() && !dst.exists() { + std::fs::copy(&src, &dst)?; + chown_to_sudo_user(&dst); + // Delete original after successful copy + std::fs::remove_file(&src)?; + } + } + + std::fs::File::create(&sentinel)?; + chown_to_sudo_user(&sentinel); + Ok(MigrationOutcome::Migrated) +} + +/// Chown all path segments from the new config root up to (but not including) +/// the home directory. This ensures intermediate directories created under +/// sudo (e.g. ~/Library/Application Support/wx-cli/) are owned by the +/// real user. +fn chown_config_tree(home: &Path, config_root: &Path) { + let mut current = config_root.to_path_buf(); + while current.starts_with(home) && current != home { + chown_to_sudo_user(¤t); + match current.parent() { + Some(parent) => current = parent.to_path_buf(), + None => break, + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::fs; + + #[test] + fn migration_no_legacy_config() { + let tmp = tempfile::tempdir().unwrap(); + let home = tmp.path(); + let new_config = home.join("new_config"); + + let result = ensure_config_migrated(home, &new_config).unwrap(); + assert!(matches!(result, MigrationOutcome::NoLegacyConfig)); + assert!(new_config.join(SENTINEL_FILE).exists()); + } + + #[test] + fn migration_already_migrated() { + let tmp = tempfile::tempdir().unwrap(); + let home = tmp.path(); + let new_config = home.join("new_config"); + fs::create_dir_all(&new_config).unwrap(); + fs::File::create(new_config.join(SENTINEL_FILE)).unwrap(); + + let result = ensure_config_migrated(home, &new_config).unwrap(); + assert!(matches!(result, MigrationOutcome::AlreadyMigrated)); + } + + #[test] + fn migration_moves_files() { + let tmp = tempfile::tempdir().unwrap(); + let home = tmp.path(); + + // Create legacy config + let old_config = home.join(".config").join(OLD_CONFIG_DIR_NAME); + fs::create_dir_all(&old_config).unwrap(); + fs::write(old_config.join("keys.toml"), "test-keys").unwrap(); + fs::write(old_config.join("settings.toml"), "test-settings").unwrap(); + + let new_config = home.join("new_config"); + let result = ensure_config_migrated(home, &new_config).unwrap(); + + match result { + MigrationOutcome::Migrated => {} + _ => panic!("expected Migrated"), + } + + // New files exist + assert_eq!(fs::read_to_string(new_config.join("keys.toml")).unwrap(), "test-keys"); + assert_eq!(fs::read_to_string(new_config.join("settings.toml")).unwrap(), "test-settings"); + + // Old files deleted + assert!(!old_config.join("keys.toml").exists()); + assert!(!old_config.join("settings.toml").exists()); + + // Sentinel exists + assert!(new_config.join(SENTINEL_FILE).exists()); + } + + #[test] + fn migration_skips_existing_dst() { + let tmp = tempfile::tempdir().unwrap(); + let home = tmp.path(); + + // Create legacy config + let old_config = home.join(".config").join(OLD_CONFIG_DIR_NAME); + fs::create_dir_all(&old_config).unwrap(); + fs::write(old_config.join("keys.toml"), "old-keys").unwrap(); + + // Create new config with existing file + let new_config = home.join("new_config"); + fs::create_dir_all(&new_config).unwrap(); + fs::write(new_config.join("keys.toml"), "new-keys").unwrap(); + + let result = ensure_config_migrated(home, &new_config).unwrap(); + match result { + MigrationOutcome::Migrated => {} + _ => panic!("expected Migrated"), + } + + // New file unchanged + assert_eq!(fs::read_to_string(new_config.join("keys.toml")).unwrap(), "new-keys"); + // Old file still exists (wasn't migrated because dst exists) + assert!(old_config.join("keys.toml").exists()); + } + + #[test] + fn migration_is_idempotent() { + let tmp = tempfile::tempdir().unwrap(); + let home = tmp.path(); + + let old_config = home.join(".config").join(OLD_CONFIG_DIR_NAME); + fs::create_dir_all(&old_config).unwrap(); + fs::write(old_config.join("keys.toml"), "test").unwrap(); + + let new_config = home.join("new_config"); + + // First run + ensure_config_migrated(home, &new_config).unwrap(); + // Second run + let result = ensure_config_migrated(home, &new_config).unwrap(); + assert!(matches!(result, MigrationOutcome::AlreadyMigrated)); + } +} diff --git a/crates/wx-paths/src/platform.rs b/crates/wx-paths/src/platform.rs new file mode 100644 index 0000000..72cdf91 --- /dev/null +++ b/crates/wx-paths/src/platform.rs @@ -0,0 +1,91 @@ +use std::path::{Path, PathBuf}; + +use crate::PathsError; + +pub(crate) struct PlatformBaseDirs { + pub config_root: PathBuf, + pub cache_root: PathBuf, + pub state_root: PathBuf, + pub logs_root: PathBuf, +} + +impl PlatformBaseDirs { + pub(crate) fn resolve(home: &Path) -> Result { + #[cfg(target_os = "macos")] + { + Ok(Self { + config_root: home.join("Library/Application Support/wx-cli/config"), + cache_root: home.join("Library/Caches/wx-cli"), + state_root: home.join("Library/Application Support/wx-cli/state"), + logs_root: home.join("Library/Logs/wx-cli"), + }) + } + + #[cfg(target_os = "linux")] + { + let config_root = dirs::config_dir() + .ok_or(PathsError::NoConfig)? + .join("wx-cli"); + let cache_root = dirs::cache_dir() + .ok_or(PathsError::NoCache)? + .join("wx-cli"); + let state_root = dirs::state_dir() + .or_else(dirs::data_local_dir) + .or_else(dirs::data_dir) + .ok_or(PathsError::NoState)? + .join("wx-cli"); + let logs_root = state_root.join("logs"); + Ok(Self { + config_root, + cache_root, + state_root, + logs_root, + }) + } + + #[cfg(target_os = "windows")] + { + let config_root = dirs::config_dir() + .ok_or(PathsError::NoConfig)? + .join("wx-cli"); + let cache_root = dirs::cache_dir() + .ok_or(PathsError::NoCache)? + .join("wx-cli"); + let local_data = dirs::data_local_dir() + .or_else(dirs::data_dir) + .ok_or(PathsError::NoState)? + .join("wx-cli"); + let state_root = local_data.join("state"); + let logs_root = local_data.join("logs"); + Ok(Self { + config_root, + cache_root, + state_root, + logs_root, + }) + } + + #[cfg(not(any(target_os = "macos", target_os = "linux", target_os = "windows")))] + { + // Fallback: use XDG-like conventions + let config_root = dirs::config_dir() + .ok_or(PathsError::NoConfig)? + .join("wx-cli"); + let cache_root = dirs::cache_dir() + .ok_or(PathsError::NoCache)? + .join("wx-cli"); + let state_root = dirs::state_dir() + .or_else(dirs::data_local_dir) + .or_else(dirs::data_dir) + .ok_or(PathsError::NoState)? + .join("wx-cli"); + let logs_root = state_root.join("logs"); + Ok(Self { + config_root, + cache_root, + state_root, + logs_root, + }) + } + } +} diff --git a/crates/wx-paths/src/sudo.rs b/crates/wx-paths/src/sudo.rs new file mode 100644 index 0000000..9ef6278 --- /dev/null +++ b/crates/wx-paths/src/sudo.rs @@ -0,0 +1,65 @@ +use std::path::{Path, PathBuf}; + +use crate::PathsError; + +/// Resolve the real user's home directory, even under `sudo -H`. +/// +/// When `SUDO_USER` is set, uses `libc::getpwnam()` to look up the original +/// user's home directory. Otherwise falls back to `dirs::home_dir()`. +pub(crate) fn resolve_real_home() -> Result { + #[cfg(unix)] + { + if let Ok(sudo_user) = std::env::var("SUDO_USER") { + if !sudo_user.is_empty() { + if let Some(home) = home_from_getpwnam(&sudo_user) { + return Ok(home); + } + } + } + } + dirs::home_dir().ok_or(PathsError::NoHome) +} + +#[cfg(unix)] +fn home_from_getpwnam(username: &str) -> Option { + use std::ffi::{CStr, CString}; + let c_user = CString::new(username).ok()?; + // Safety: getpwnam returns a pointer to a static struct or null. + let pw = unsafe { libc::getpwnam(c_user.as_ptr()) }; + if pw.is_null() { + return None; + } + // Safety: pw_dir is a valid C string if pw is non-null. + let home_cstr = unsafe { CStr::from_ptr((*pw).pw_dir) }; + let home_str = home_cstr.to_str().ok()?; + Some(PathBuf::from(home_str)) +} + +/// Best-effort chown to the real user when running under sudo. +/// +/// Reads `SUDO_UID` and `SUDO_GID` environment variables and uses `libc::chown()` +/// to restore ownership. Silently does nothing if not running under sudo or if +/// the chown fails (e.g. on non-Unix platforms). +pub fn chown_to_sudo_user(path: &Path) { + #[cfg(unix)] + { + use std::ffi::CString; + let uid: u32 = match std::env::var("SUDO_UID").ok().and_then(|v| v.parse().ok()) { + Some(uid) => uid, + None => return, + }; + let gid: u32 = std::env::var("SUDO_GID") + .ok() + .and_then(|v| v.parse().ok()) + .unwrap_or(uid); + if let Some(c_path) = path.to_str().and_then(|s| CString::new(s).ok()) { + unsafe { + libc::chown(c_path.as_ptr(), uid, gid); + } + } + } + #[cfg(not(unix))] + { + let _ = path; + } +}