163 lines
7.0 KiB
Python
163 lines
7.0 KiB
Python
"""Live database reads on one daemon worker; GUI never waits for key acquisition."""
|
|
from concurrent.futures import Future, TimeoutError
|
|
import queue
|
|
import threading
|
|
import time
|
|
|
|
from runtime_paths import application_data_dir
|
|
from wxwork_db import DatabaseReadError, WXWorkDB, detect_wxwork_dir, load_keys
|
|
|
|
|
|
class LiveReplyDatabase:
|
|
def __init__(self, *, factory=None, read_timeout=0.35, retry_seconds=15.0):
|
|
self._factory = factory or self._open_database
|
|
self._read_timeout = max(0.0, float(read_timeout))
|
|
self._retry_seconds = max(0.0, float(retry_seconds))
|
|
self._jobs = queue.Queue(maxsize=8)
|
|
self._pending = {}
|
|
self._lock = threading.RLock()
|
|
self._stop = threading.Event()
|
|
self._thread = None
|
|
self._ready = False
|
|
self.last_error = "数据库等待初始化"
|
|
self._since = time.time() - 120.0
|
|
|
|
@staticmethod
|
|
def _open_database():
|
|
from wxwork_local_setup import acquire_local_keys, select_source_directory
|
|
source = detect_wxwork_dir()
|
|
if not source:
|
|
raise DatabaseReadError("未找到企业微信数据目录,可暂用聊天内容复制")
|
|
try:
|
|
keys = load_keys()
|
|
except ValueError:
|
|
keys = {}
|
|
cache = str(application_data_dir() / "wxwork_reply_cache")
|
|
database = WXWorkDB(source, keys, cache_dir=cache)
|
|
if database.health_check():
|
|
return database
|
|
database.close()
|
|
select_source_directory(source)
|
|
acquire_local_keys()
|
|
database = WXWorkDB(source, load_keys(), cache_dir=cache)
|
|
if not database.health_check():
|
|
database.close()
|
|
raise DatabaseReadError("自动解密尚未得到消息数据库,可暂用聊天内容复制")
|
|
return database
|
|
|
|
def _run(self):
|
|
database = None
|
|
retry_at = 0.0
|
|
checkpoint = None
|
|
try:
|
|
while not self._stop.is_set():
|
|
try:
|
|
key, operation, args, kwargs, future = self._jobs.get(timeout=0.25)
|
|
except queue.Empty:
|
|
continue
|
|
if not future.set_running_or_notify_cancel():
|
|
continue
|
|
if database is None and time.monotonic() < retry_at:
|
|
future.set_exception(DatabaseReadError(self.last_error))
|
|
continue
|
|
try:
|
|
if database is None:
|
|
database = self._factory()
|
|
if checkpoint is not None and isinstance(database, WXWorkDB):
|
|
database._message_cursors = {key: dict(value) for key, value in checkpoint.items()}
|
|
database._bootstrap_since = self._since
|
|
if self._stop.is_set():
|
|
raise DatabaseReadError("数据库读取已停止")
|
|
result = getattr(database, operation)(*args, **kwargs)
|
|
self._ready = True
|
|
self.last_error = ""
|
|
future.set_result((time.monotonic(), result))
|
|
except Exception as exc:
|
|
self._ready = False
|
|
self.last_error = str(exc)
|
|
future.set_exception(DatabaseReadError(str(exc)))
|
|
if database is not None:
|
|
if isinstance(database, WXWorkDB):
|
|
checkpoint = {key: dict(value) for key, value in database._message_cursors.items()}
|
|
database.close()
|
|
database = None
|
|
retry_at = time.monotonic() + self._retry_seconds
|
|
finally:
|
|
if database is not None:
|
|
database.close()
|
|
self._ready = False
|
|
|
|
def _request(self, key, operation, *args, **kwargs):
|
|
with self._lock:
|
|
if self._stop.is_set():
|
|
raise DatabaseReadError("数据库读取已停止")
|
|
if self._thread is None:
|
|
self._thread = threading.Thread(target=self._run, daemon=True, name="reply-database-reader")
|
|
self._thread.start()
|
|
# Do not retain completed snapshots for contacts that are never revisited.
|
|
for pending_key, pending_future in list(self._pending.items()):
|
|
if pending_key == "poll" or pending_key == key or not pending_future.done():
|
|
continue
|
|
try:
|
|
completed_at, _ = pending_future.result()
|
|
expired = time.monotonic() - completed_at > 2.0
|
|
except Exception:
|
|
expired = True
|
|
if expired:
|
|
self._pending.pop(pending_key, None)
|
|
future = self._pending.get(key)
|
|
# 文字快照过期时重新读;增量消息结果必须消费,不能丢弃已推进的游标。
|
|
if future is not None and future.done() and operation != "get_new_messages":
|
|
try:
|
|
completed_at, _ = future.result()
|
|
if time.monotonic() - completed_at > 2.0:
|
|
self._pending.pop(key, None)
|
|
future = None
|
|
except Exception:
|
|
self._pending.pop(key, None)
|
|
future = None
|
|
if future is None:
|
|
future = Future()
|
|
try:
|
|
self._jobs.put_nowait((key, operation, args, kwargs, future))
|
|
except queue.Full as exc:
|
|
raise DatabaseReadError("数据库正在读取其他消息") from exc
|
|
self._pending[key] = future
|
|
consumed = False
|
|
try:
|
|
_, result = future.result(timeout=self._read_timeout)
|
|
consumed = True
|
|
except TimeoutError as exc:
|
|
raise DatabaseReadError("数据库正在初始化或刷新,暂用备用读取") from exc
|
|
except Exception:
|
|
consumed = True
|
|
raise
|
|
finally:
|
|
if consumed:
|
|
with self._lock:
|
|
if self._pending.get(key) is future:
|
|
self._pending.pop(key, None)
|
|
return result
|
|
|
|
def get_new_messages(self, since_ts):
|
|
# 首次启动的时间边界固定,恢复连接时不会被其他账号的较晚时间推走。
|
|
return self._request("poll", "get_new_messages", self._since)
|
|
|
|
def get_conversation_context(self, display_name, *, account="", conv_id="", limit=100):
|
|
key = ("context", str(display_name), str(account), str(conv_id))
|
|
try:
|
|
return self._request(key, "get_conversation_context", display_name,
|
|
account=account, conv_id=conv_id, limit=limit)
|
|
except DatabaseReadError:
|
|
return None
|
|
|
|
def health_check(self):
|
|
return self._ready and not self._stop.is_set()
|
|
|
|
def close(self):
|
|
self._stop.set()
|
|
with self._lock:
|
|
for future in self._pending.values():
|
|
future.cancel()
|
|
self._pending.clear()
|