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.coordinator import run_auto_extract_sync try: _state.init_state.set_progress(10.0, "正在扫描内存提取密钥") key_map = run_auto_extract_sync(force=force) if key_map: _state.init_state.set_success( len(key_map), f"成功提取 {len(key_map)} 个密钥", ) else: _state.init_state.set_failed("未提取到任何有效密钥") 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, )