from __future__ import annotations import asyncio import logging import os from typing import Optional from fastapi import APIRouter from woc_bridge.config import _state, _require_db_reader from woc_bridge.db.coordinator import ( _verify_db_key, _read_db_salt, _invalidate_db_state_cache, ) from woc_bridge.models import ( DbDecryptRequest, DbDecryptResponse, DbKeyStatusResponse, DbInitRequest, DbInitResponse, DbInitStatusResponse, BridgeError, ) logger = logging.getLogger("woc-bridge") router = APIRouter() @router.post("/api/db/decrypt", response_model=DbDecryptResponse) async def decrypt_db_with_key(request: DbDecryptRequest) -> DbDecryptResponse: """手动传入 SQLCipher 密钥触发解密验证与缓存。 外部系统(如 ForcePilot)或管理员通过本接口注入 64 位十六进制密钥, bridge 验证 key 是否匹配当前 DB,验证通过后按 salt 缓存到 KeyCache, 后续所有 DB 查询接口立即恢复可用。 请求体除 key 外,可选传 salt 字段指定该 key 对应的 salt;不传时 bridge 自动读取当前 DB 的 salt 进行匹配。 Args: request: 含 key 字段(64 位十六进制字符串),可选 salt 字段 Returns: DbDecryptResponse:含 success / verified / key_mode / error """ key = request.key.strip() # 格式校验 if len(key) != 64 or not all(c in "0123456789abcdefABCDEF" for c in key): raise BridgeError( code="INVALID_PARAMS", message="key 必须是 64 位十六进制字符串", ) db_reader = _require_db_reader() # 检查 DB 是否存在 db_status = await asyncio.to_thread(db_reader.check_db_status) if db_status == "not_found": raise BridgeError( code="DB_NOT_FOUND", message="未找到微信 DB,无法验证 key", ) if db_status == "unreadable": raise BridgeError( code="DB_NOT_FOUND", message="DB 文件存在但不可读(权限问题),无法验证 key", ) # 明文 DB 无需 key,直接返回提示 if db_status == "ok": return DbDecryptResponse( success=True, verified=False, key_mode=None, error="db_not_encrypted", ) # 获取当前 DB 的 salt(用户未指定时自动识别) salt_hex = getattr(request, "salt", None) if not salt_hex: salt_hex = await asyncio.to_thread(_read_db_salt, db_reader) if not salt_hex: raise BridgeError( code="DB_ENCRYPTED", message="无法读取当前 DB 的 salt", ) # 验证 key(_verify_db_key 已内置 key_mode 获取,无需再读一次 page1) try: verified, key_mode = await asyncio.to_thread(_verify_db_key, db_reader, key) except Exception: verified, key_mode = False, None if verified: # 缓存 key 并注入 DbReader(按 salt 存储) # db_reader.set_key 已转发到 key_cache.set_key,避免双份同步 db_reader.set_key(key, salt_hex=salt_hex) _invalidate_db_state_cache() # 补充元信息(source/verified 在 set_key 默认为 auto_extract/True,这里覆盖为 api) _state.key_cache.set_meta(source="api", verified=True) logger.info( "db/decrypt: key=%s...%s salt=%s → 验证成功, key_mode=%s", key[:4], key[-4:], salt_hex, key_mode, ) return DbDecryptResponse( success=True, verified=True, key_mode=key_mode, error=None, ) else: logger.warning( "db/decrypt: key=%s...%s salt=%s → 验证失败(key 不匹配)", key[:4], key[-4:], salt_hex, ) return DbDecryptResponse( success=True, verified=False, key_mode=None, error="key_mismatch", ) # --------------------------------------------------------------------------- # 路由:GET /api/db/key/status(key 缓存状态查询) # --------------------------------------------------------------------------- @router.get("/api/db/key/status", response_model=DbKeyStatusResponse) async def get_db_key_status() -> DbKeyStatusResponse: """返回当前 DB 解密密钥缓存状态。 供调用方判断是否需要注入 key 或等待自动提取。 Returns: DbKeyStatusResponse:含 cached / source / verified / key_prefix """ return DbKeyStatusResponse( cached=_state.key_cache.cached, source=_state.key_cache.source, verified=_state.key_cache.verified, key_prefix=_state.key_cache.key_prefix(), ) # --------------------------------------------------------------------------- # 路由:POST /api/db/init(显式触发密钥提取) # --------------------------------------------------------------------------- @router.post("/api/db/init", response_model=DbInitResponse) async def init_db(request: DbInitRequest = DbInitRequest()) -> DbInitResponse: """显式触发 DB 初始化(密钥提取)。 后台线程执行内存扫描,避免阻塞 /api/status 等接口。调用方可通过 GET /api/db/init/status 轮询进度,或观察 /api/status 中的 init_in_progress / init_progress_pct / init_message 字段。 Args: request: 可选 pid / db_dir / force Returns: DbInitResponse:请求受理状态 """ _require_db_reader() if _state.init_state.state == "running": return DbInitResponse( success=True, state="in_progress", message=_state.init_state.message or "初始化中,请稍后", key_count=None, ) if _state.db_reader.has_keys() and not request.force: key_count = len(_state.db_reader.get_keys()) return DbInitResponse( success=True, state="already_done", message=f"已初始化({key_count} 个密钥),使用 force=true 可强制重新提取", key_count=key_count, ) def _init_task(pid: Optional[int], db_dir: Optional[str], force: bool) -> None: from woc_bridge.db.key_extractor import KeyExtractor, _get_wechat_pids try: _state.init_state.set_progress(10.0, "正在检测微信进程") if pid: pids = [(pid, 0)] else: pids = _get_wechat_pids() if not pids: _state.init_state.set_failed("微信进程未运行") return _state.init_state.set_progress(20.0, "正在定位数据库目录") if db_dir is None: db_path = _state.db_reader._find_db_path() if db_path is None: _state.init_state.set_failed("未找到微信数据库目录") return db_dir = os.path.dirname(db_path) scan_dir = "/config/xwechat_files" if os.path.isdir("/config/xwechat_files") else db_dir _state.init_state.set_progress(30.0, "正在扫描内存提取密钥") extractor = KeyExtractor(pid or 0, db_dir) key_map = extractor.extract_all_keys(scan_dir) if not key_map: _state.init_state.set_failed("未提取到任何有效密钥") return _state.init_state.set_progress(80.0, "正在保存密钥到文件") wechat_pid = pids[0][0] # 读取真实 start_time,用于进程重启后 key 缓存失效检测 start_time = 0.0 try: with open(f"/proc/{wechat_pid}/stat", "r") as f: stat_fields = f.read().split() start_time = float(stat_fields[21]) if len(stat_fields) > 21 else 0.0 except (OSError, IndexError, ValueError): start_time = 0.0 _state.db_reader.set_keys( key_map, source="auto_extract", pid=wechat_pid, start_time=start_time, ) _invalidate_db_state_cache() _state.init_state.set_success( len(key_map), f"成功提取 {len(key_map)} 个密钥", ) except Exception as e: _state.init_state.set_failed(str(e)) _state.init_state.start(_init_task, args=(request.pid, request.db_dir, request.force)) return DbInitResponse( success=True, state="started", message="初始化已启动,可通过 /api/db/init/status 查询进度", key_count=None, ) # --------------------------------------------------------------------------- # 路由:GET /api/db/init/status # --------------------------------------------------------------------------- @router.get("/api/db/init/status", response_model=DbInitStatusResponse) async def get_init_status() -> DbInitStatusResponse: """查询 DB 初始化后台任务状态。""" return DbInitStatusResponse( state=_state.init_state.state, progress_pct=_state.init_state.progress_pct, message=_state.init_state.message, key_count=_state.init_state.key_count, error=_state.init_state.error, )