diff --git a/src/wechat_decrypt_tool/ai/agent_service.py b/src/wechat_decrypt_tool/ai/agent_service.py index a81f36ca..61394bc9 100644 --- a/src/wechat_decrypt_tool/ai/agent_service.py +++ b/src/wechat_decrypt_tool/ai/agent_service.py @@ -42,6 +42,7 @@ def __init__(self, ai, tools=None, model=None): self.readers = {} self.index_workers = {} self.reference_contacts = {} + self._draft_saved_at = {} from .deep_subtasks import DeepSubtasks self.subtasks = DeepSubtasks(self) @@ -68,14 +69,20 @@ def run(self, id, account=None): return record @serialized - def update(self, id, **fields): + def update(self, id, *, transient=False, emit=True, **fields): + with self.store.lock if transient else self.store.connection(): + return self._update(id, transient=transient, emit=emit, **fields) + + def _update(self, id, *, transient, emit, **fields): record = self.run(id) values = fields.pop('evidence',None) if values is not None and not isinstance(values,Evidence): record['evidence'].replace(values) record.pop('evidence',None) + if all(record.get(key) == value for key, value in fields.items()) and not self.store.has_live_record('agent_run', id): + return self.run(id) record.update(fields, updated_at=time.time()) - self.store.put('agent_run', record) + self.store.put('agent_run', record, transient=transient) if 'version' in fields or 'status' in fields: diagnostic_event('agent.run.state', run_id=id, thread_id=record['thread_id'], version=record['version'], status=record['status']) event = {'type': 'run_patch', 'run_id': id, 'thread_id': record['thread_id'], 'status': record['status'], @@ -104,9 +111,24 @@ def update(self, id, **fields): } if patch: event['patch'] = patch - self.store.event(record['account'], 'agent', event) + if emit and (patch or any(key in fields for key in ('context_budget', 'analysis', 'status', 'version'))): + self.store.event(record['account'], 'agent', event, transient=transient, + unique_key=f'run_patch:{id}' if transient else None) return self.run(id) + @serialized + def stream_answer(self, id, text, stage): + """实时正文只推送内存事件;每两秒及终态保存可恢复草稿。""" + run = self.guard(id) + now = time.monotonic() + key = (id, run['version']) + saved_at = self._draft_saved_at.setdefault(key, now) + persist = now - saved_at >= 2.0 + self.timeline_item(id, 'answer', text, item_id='answer:' + id, status='running', + persist=persist, transient_event=True, run_fields={'answer': text, 'stage': stage}) + if persist: + self._draft_saved_at[key] = now + def guard(self, id): run = self.run(id) if self.stopping or run['status'] not in ACTIVE: @@ -190,10 +212,16 @@ def profile(self, run, vision=False): return snapshot | {'api_key': live.get('api_key', '')} @observed('agent.finish', id_field='run_id') + @serialized def finish(self, id, status, error='', error_info=None): + with self.store.connection(): + return self._finish(id, status, error, error_info) + + def _finish(self, id, status, error, error_info): run = self.store.get('agent_run', id) if not run: return + self._draft_saved_at = {key: value for key, value in self._draft_saved_at.items() if key[0] != id} if status == 'completed' and run.get('delegation_partial'): status = 'interrupted' error = '部分子任务未完成,当前为阶段结果,可继续分析。' diff --git a/src/wechat_decrypt_tool/ai/agent_timeline.py b/src/wechat_decrypt_tool/ai/agent_timeline.py index 1ae4773c..2b50e6c6 100644 --- a/src/wechat_decrypt_tool/ai/agent_timeline.py +++ b/src/wechat_decrypt_tool/ai/agent_timeline.py @@ -2,15 +2,33 @@ import time import uuid import re +from contextlib import nullcontext +from .deep_synchronization import serialized class AgentTimeline: - def timeline_item(self, id, kind, text, *, item_id=None, status='completed', **fields): + @serialized + def timeline_item(self, id, kind, text, *, item_id=None, status='completed', persist=True, + transient_event=False, run_fields=None, **fields): + # 业务时间线、任务快照和持久事件在同一事务提交,避免一个动作多次刷盘。 + with self.store.connection() if persist else nullcontext(): + return self._timeline_item(id, kind, text, item_id=item_id, status=status, persist=persist, + transient_event=transient_event, run_fields=run_fields, **fields) + + def _timeline_item(self, id, kind, text, *, item_id, status, persist, transient_event, run_fields, **fields): run = self.run(id) items = run.get('timeline') or [dict(x, seq=i+1, revision=1, kind='status') for i,x in enumerate(run.get('activity', []))] now = time.time() item = next((x for x in items if x['id'] == item_id), None) previous_text = item.get('text', '') if item else '' + if (item and item.get('text') == text and item.get('status') == status + and all(item.get(k) == v for k, v in fields.items()) + and all(run.get(k) == v for k, v in (run_fields or {}).items())): + if persist and self.store.has_live_record('agent_run', id): + self.update(id, emit=False, **(run_fields or {})) + if hasattr(self, 'workspace'): + self.workspace.put(id, run['version'], f'timeline:{item["seq"]:012d}', 'timeline', item) + return item['id'] if item is None: item = dict(id=item_id or uuid.uuid4().hex, seq=max(run.get('timeline_seq',0),max((x.get('seq',0) for x in items),default=0)) + 1, revision=0, kind=kind, started_at=now, input_version=run['version']) @@ -19,12 +37,18 @@ def timeline_item(self, id, kind, text, *, item_id=None, status='completed', **f if status not in ('running', 'received'): item['finished_at'] = now # 普通过程保留最近 200 步;压缩节点是持久的对话分隔,不能随步骤淘汰。 - if hasattr(self,'workspace'): + if persist and hasattr(self,'workspace'): self.workspace.put(id,run['version'],f'timeline:{item["seq"]:012d}','timeline',item) retained = [x for i, x in enumerate(items) if i >= len(items) - 200 or (x.get('kind') == 'notice' and x.get('context_job', {}).get('before') is not None)] - self.update(id, timeline=retained,timeline_seq=max(x.get('seq',0) for x in items)) + self.update(id, transient=not persist, emit=False, timeline=retained, + timeline_seq=max(x.get('seq',0) for x in items), **(run_fields or {})) event = {'type':'timeline_item', 'run_id':id, 'thread_id':run['thread_id'], 'version':run['version'], 'timeline_item':item} + if run_fields: + current = self.run(id) + event['updated_at'] = current['updated_at'] + # 正文仍从时间线合并,阶段信息同一个事件送达,避免额外写入 run_patch。 + event['patch'] = {key: current[key] for key in run_fields if key != 'answer'} if kind in ('answer', 'progress'): # 新标记与已校验身份同时送达,避免正文先出现、来源等待慢速快照。 # 只发送本次新增标记,完整映射仍由运行快照和历史保存负责。 @@ -35,7 +59,8 @@ def timeline_item(self, id, kind, text, *, item_id=None, status='completed', **f from .agent_references import cited_references event['citations'] = self.citations({**run, 'answer': added, 'timeline': [], 'answer_context': {}}) event['references'] = cited_references(added, run.get('references', {})) - self.store.event(run['account'], 'agent', event) + self.store.event(run['account'], 'agent', event, transient=transient_event, + unique_key=f'timeline:{id}:{item["id"]}' if transient_event else None) return item['id'] def close_activity(self, id, status='completed'): diff --git a/src/wechat_decrypt_tool/ai/deep_runtime.py b/src/wechat_decrypt_tool/ai/deep_runtime.py index 57ad6ca8..8675fa43 100644 --- a/src/wechat_decrypt_tool/ai/deep_runtime.py +++ b/src/wechat_decrypt_tool/ai/deep_runtime.py @@ -802,8 +802,7 @@ async def execute(self, id): text = delta.get('text', '') if delta.get('type') == 'text-delta' else '' partial += text if text and time.monotonic() - last_emit >= .08: - self.update(id, answer=partial, stage='提取局部事实' if run.get('parent_run_id') and run.get('subtask_plan_version') == REVISION else '正在回答') - self.timeline_item(id, 'answer', partial, item_id='answer:' + id, status='running') + self.stream_answer(id, partial, '提取局部事实' if run.get('parent_run_id') and run.get('subtask_plan_version') == REVISION else '正在回答') last_emit = time.monotonic() final = await agent.aget_state(config) answer = next((m for m in reversed(final.values.get('messages', [])) if isinstance(m, AIMessage) and not m.tool_calls), None) diff --git a/src/wechat_decrypt_tool/ai/lifecycle.py b/src/wechat_decrypt_tool/ai/lifecycle.py index 115ae1f8..db3469e4 100644 --- a/src/wechat_decrypt_tool/ai/lifecycle.py +++ b/src/wechat_decrypt_tool/ai/lifecycle.py @@ -92,3 +92,6 @@ async def stop_services(): await factory().stop() except Exception as error: event('lifecycle.service.stop_failed', level=logging.ERROR, component=name, error=error) + # Agent 与摘要共享业务库,全部工作线程退出后才关闭持久连接。 + for store in (get_ai_service().store, get_local_search().store): + store.close() diff --git a/src/wechat_decrypt_tool/ai/runtime_check.py b/src/wechat_decrypt_tool/ai/runtime_check.py index d42d0c25..1585fc94 100644 --- a/src/wechat_decrypt_tool/ai/runtime_check.py +++ b/src/wechat_decrypt_tool/ai/runtime_check.py @@ -4,6 +4,7 @@ import platform import sys import tempfile +from contextlib import closing from pathlib import Path @@ -40,26 +41,28 @@ def bind_tools(self, tools, **kwargs): from .providers import ModelService from .service import AIService from .agent_service import AgentService - store = AIStore(root / 'application') - store.put('profile', {'model': 'synthetic', 'name': '冻结验收', 'protocol': 'openai', - 'base_url': 'http://127.0.0.1:1/v1', 'api_key': 'synthetic-unused-key', - 'model_overrides': {'context_window': 32768}}, id='synthetic') - store.put('defaults', {'text': 'synthetic'}, id='global') - models = ModelService(store) - models.client = lambda profile: RuntimeModel(messages=iter([AIMessage(content='运行正常')])) - class NoQueries: - async def conversations(self, account): - raise AssertionError('冻结问候检查不能读取聊天目录') - service = AgentService(AIService(store, models), tools=NoQueries()) - thread = await service.create_thread('synthetic', '', '运行检查') - run = await service.submit(thread['id'], 'synthetic', {'text': '你好', 'request_id': 'smoke'}) - await service.workers[run['id']] - result = service.public_run(run['id'], 'synthetic') - assert result['status'] == 'completed', result.get('error') - assert result['answer'] == '运行正常' - assert result['used']['models'] == 1 and result['used']['tools'] == 0 - assert result['engine'] == 'deepagents' and result['engine_version'] == 3 - await service.stop() + with closing(AIStore(root / 'application')) as store: + store.put('profile', {'model': 'synthetic', 'name': '冻结验收', 'protocol': 'openai', + 'base_url': 'http://127.0.0.1:1/v1', 'api_key': 'synthetic-unused-key', + 'model_overrides': {'context_window': 32768}}, id='synthetic') + store.put('defaults', {'text': 'synthetic'}, id='global') + models = ModelService(store) + models.client = lambda profile: RuntimeModel(messages=iter([AIMessage(content='运行正常')])) + class NoQueries: + async def conversations(self, account): + raise AssertionError('冻结问候检查不能读取聊天目录') + service = AgentService(AIService(store, models), tools=NoQueries()) + try: + thread = await service.create_thread('synthetic', '', '运行检查') + run = await service.submit(thread['id'], 'synthetic', {'text': '你好', 'request_id': 'smoke'}) + await service.workers[run['id']] + result = service.public_run(run['id'], 'synthetic') + assert result['status'] == 'completed', result.get('error') + assert result['answer'] == '运行正常' + assert result['used']['models'] == 1 and result['used']['tools'] == 0 + assert result['engine'] == 'deepagents' and result['engine_version'] == 3 + finally: + await service.stop() def check_runtime(model_root=None): @@ -90,9 +93,10 @@ def check_runtime(model_root=None): assert 'cl100k_base' in tiktoken.list_encoding_names() with tempfile.TemporaryDirectory(prefix='wechat-ai-runtime-') as directory: root = Path(directory) - store = AIStore(root / 'business') - store.put('runtime_check', {'ok': True}, id='check', account='synthetic') - assert store.get('runtime_check', 'check')['ok'] + # 自检持有的业务连接必须在删除临时目录前释放,Windows 不允许删除打开的库。 + with closing(AIStore(root / 'business')) as store: + store.put('runtime_check', {'ok': True}, id='check', account='synthetic') + assert store.get('runtime_check', 'check')['ok'] asyncio.run(_checkpoint(root)) report['application_graph'] = 'deepagents-v3-one-call-no-query' index = SemanticIndex(root / 'vectors.sqlite3') diff --git a/src/wechat_decrypt_tool/ai/storage.py b/src/wechat_decrypt_tool/ai/storage.py index 7c49fca8..49e88c4b 100644 --- a/src/wechat_decrypt_tool/ai/storage.py +++ b/src/wechat_decrypt_tool/ai/storage.py @@ -1,6 +1,8 @@ from __future__ import annotations from .diagnostics import observed, event as diagnostic_event import logging +import copy +from collections import OrderedDict import json import os @@ -44,45 +46,133 @@ def __init__(self, root: Path | None = None): self.root.mkdir(parents=True, exist_ok=True) self.path = self.root / "ai.sqlite3" self.lock = threading.RLock() + self._db = None + self._depth = 0 + self._live_records = {} + self._live_events = OrderedDict() + self._live_event_bytes = 0 + self._last_event_id = 0 self.revoked_accounts = set() - # SSE 订阅者通过条件变量等待新事件;SQLite 仍是断线重放的权威来源。 + # SSE 订阅者通过条件变量等待新事件;持久事件及内存快照共同支持断线重放。 # 使用按账号修订号,避免其他账号的高频事件无谓唤醒当前连接。 self._event_condition = threading.Condition() self._event_revisions = {} with self.connection() as db: db.executescript(SCHEMA_SQL) + self._reserve_event_ids() + + def _reserve_event_ids(self): + """低频预留事件编号,内存事件无需写盘也能跨重启保持游标递增。""" + with self.connection() as db: + sequence = db.execute("SELECT coalesce(max(seq),0) FROM sqlite_sequence WHERE name='events'").fetchone()[0] + maximum = db.execute('SELECT coalesce(max(id),0) FROM events').fetchone()[0] + self._next_event_id = max(sequence, maximum, self._last_event_id) + 1 + self._event_id_limit = self._next_event_id + 1_000_000 + if not db.execute("UPDATE sqlite_sequence SET seq=? WHERE name='events'", (self._event_id_limit,)).rowcount: + db.execute("INSERT INTO sqlite_sequence(name,seq) VALUES('events',?)", (self._event_id_limit,)) + + def _allocate_event_id(self): + id = max(self._next_event_id, self.latest_event_id() + 1) + if id >= self._event_id_limit: + self._reserve_event_ids() + id = self._next_event_id + self._next_event_id = id + 1 + self._last_event_id = id + return id @contextmanager def connection(self): with self.lock: - db = sqlite3.connect(self.path, timeout=30) - db.row_factory = sqlite3.Row - db.execute("PRAGMA journal_mode=WAL") + if self._db is None: + self._db = sqlite3.connect(self.path, timeout=30, check_same_thread=False) + self._db.row_factory = sqlite3.Row + self._db.execute("PRAGMA journal_mode=WAL") + db = self._db + depth = self._depth + live_records = self._live_records.copy() + live_events = self._live_events.copy() + live_bytes = self._live_event_bytes + event_range = (getattr(self, '_next_event_id', None), getattr(self, '_event_id_limit', None)) + if depth: + db.execute(f'SAVEPOINT nested_{depth}') + else: + db.execute('BEGIN') + self._depth += 1 try: - with db: + if depth: yield db - except Exception as error: + db.execute(f'RELEASE nested_{depth}') + else: + with db: + yield db + except BaseException as error: + if depth: + db.execute(f'ROLLBACK TO nested_{depth}') + db.execute(f'RELEASE nested_{depth}') + self._live_records = live_records + self._live_events = live_events + self._live_event_bytes = live_bytes + if event_range[0] is not None: + self._next_event_id, self._event_id_limit = event_range diagnostic_event('storage.transaction.failed', level=logging.ERROR, error=error, committed=False) raise finally: - db.close() + self._depth -= 1 - def put(self, kind, body, id=None, account=""): + def close(self): + """工作线程退出后关闭连接;提交仍保持 SQLite 默认的持久性级别。""" + with self.lock: + if self._depth: + raise RuntimeError('不能在事务中关闭 AI 存储') + if self._db is not None: + self._db.close() + self._db = None + + def has_live_record(self, kind, id): + with self.lock: + return (kind, id) in self._live_records + + def put(self, kind, body, id=None, account="", transient=False): id = id or body.get("id") or uuid.uuid4().hex body = {**body, "id": id} + owner = account or body.get("account", "") + with self.lock: + if owner in self.revoked_accounts: + return body + if transient: + self._live_records[kind, id] = (copy.deepcopy(body), owner, time.time()) + return body with self.connection() as db: if (account or body.get("account", "")) in self.revoked_accounts: return body - db.execute("INSERT INTO records VALUES(?,?,?,?,?) ON CONFLICT(kind,id) DO UPDATE SET body=excluded.body, account=excluded.account, updated=excluded.updated", + db.execute("INSERT INTO records VALUES(?,?,?,?,?) ON CONFLICT(kind,id) DO UPDATE SET body=excluded.body, account=excluded.account, updated=excluded.updated WHERE records.body<>excluded.body OR records.account<>excluded.account", (kind, id, account or body.get("account", ""), json.dumps(body, ensure_ascii=False), time.time())) + self._live_records.pop((kind, id), None) return body def get(self, kind, id): + with self.lock: + if (kind, id) in self._live_records: + return copy.deepcopy(self._live_records[kind, id][0]) with self.connection() as db: row = db.execute("SELECT body FROM records WHERE kind=? AND id=?", (kind, id)).fetchone() return json.loads(row[0]) if row else None def list(self, kind, account=None, limit=None, offset=0, compact=False): + with self.lock: + live = {id: entry for (entry_kind, id), entry in self._live_records.items() + if entry_kind == kind and (account is None or entry[1] == account)} + if live: + with self.connection() as db: + rows = db.execute('SELECT id,body,updated FROM records WHERE kind=?' + + (' AND account=?' if account is not None else ''), + [kind] if account is None else [kind, account]).fetchall() + merged = {row['id']: (json.loads(row['body']), row['updated']) for row in rows} + merged.update({id: (copy.deepcopy(entry[0]), entry[2]) for id, entry in live.items()}) + bodies = [body for body, _ in sorted(merged.values(), key=lambda entry: entry[1], reverse=True)] + bodies = bodies[offset:offset + limit if limit is not None else None] + return [{k: v for k, v in body.items() if not compact or k not in + ('results', 'overview', 'models', 'cursors')} for body in bodies] column = "json_remove(body,'$.results','$.overview','$.models','$.cursors')" if compact else "body" args = [kind] if account is None else [kind, account] sql = f"SELECT {column} FROM records WHERE kind=?" + (" AND account=?" if account is not None else "") + " ORDER BY updated DESC" @@ -114,41 +204,94 @@ def recover_interrupted_usage(self): def latest_event_id(self): with self.connection() as db: - return db.execute("SELECT coalesce(max(id),0) FROM events").fetchone()[0] + return max(self._last_event_id, db.execute("SELECT coalesce(max(id),0) FROM events").fetchone()[0]) def delete(self, kind, id): with self.connection() as db: db.execute("DELETE FROM records WHERE kind=? AND id=?", (kind, id)) + self._live_records.pop((kind, id), None) + if kind in ('agent_run', 'agent_thread'): + field = 'run_id' if kind == 'agent_run' else 'thread_id' + self.discard_live_events(field, id) - def event(self, account, kind, body, unique_key=None, replace=False): + def discard_live_events(self, field, value): + """删除任务或账号时同步移除内存重放,避免已删内容再次送达。""" + with self.lock: + for key, (row, size) in list(self._live_events.items()): + if row['body'].get(field) == value: + del self._live_events[key] + self._live_event_bytes -= size + + def event(self, account, kind, body, unique_key=None, replace=False, transient=False): """写入事件供 SSE 重放。 默认行为保持不变:提供 `unique_key` 时按去重语义写入(同 key 已存在则忽略), 用于提醒等只应投递一次的事件。`replace=True` 时改为用最新快照替换旧行, 让高频进度事件每个逻辑任务只保留一行,同时因 INSERT OR REPLACE 会删除旧行、 - 新行仍获得递增的自增 id,断线重连的 EventSource 依然能收到最新状态。 + 新行仍获得递增 id,断线重连的 EventSource 依然能收到最新状态。 + `transient=True` 仅在有界内存缓存中合并展示快照,不写 SQLite;通知仍使用默认持久化。 """ + with self.lock: + if account in self.revoked_accounts: + return + if transient: + id = self._allocate_event_id() + key = (account, kind, unique_key) if unique_key is not None else id + old = self._live_events.pop(key, None) + if old: + self._live_event_bytes -= old[1] + # 累计正文快照被合并后,早先送达的引用映射也必须随最新快照保留。 + if old[0]['body'].get('version') == body.get('version'): + body = dict(body) + for field, identity in (('citations', 'source'), ('references', 'id')): + if old[0]['body'].get(field): + merged = {item[identity]: item for item in old[0]['body'][field]} + merged.update({item[identity]: item for item in body.get(field, [])}) + body[field] = list(merged.values()) + payload = json.dumps(body, ensure_ascii=False) + row = dict(id=id, account=account, kind=kind, body=json.loads(payload), + unique_key=unique_key, delivered=0, created=time.time()) + size = len(payload.encode('utf-8')) + self._live_events[key] = (row, size) + self._live_event_bytes += size + # 短暂断线重放最新快照;长断线由前端已有的重连 GET 补齐权威状态。 + while len(self._live_events) > 1 and ( + len(self._live_events) > 256 or self._live_event_bytes > 8 * 1024 * 1024): + _, (_, removed_size) = self._live_events.popitem(last=False) + self._live_event_bytes -= removed_size + self._notify_event(account) + return with self.connection() as db: if account in self.revoked_accounts: return payload = json.dumps(body, ensure_ascii=False) + id = self._allocate_event_id() if unique_key is None: cursor = db.execute( - "INSERT INTO events(account,kind,body,created) VALUES(?,?,?,?)", - (account, kind, payload, time.time())) + "INSERT INTO events(id,account,kind,body,created) VALUES(?,?,?,?,?)", + (id, account, kind, payload, time.time())) elif replace: cursor = db.execute( - "INSERT OR REPLACE INTO events(account,kind,body,unique_key,created) VALUES(?,?,?,?,?)", - (account, kind, payload, unique_key, time.time())) + "INSERT OR REPLACE INTO events(id,account,kind,body,unique_key,created) VALUES(?,?,?,?,?,?)", + (id, account, kind, payload, unique_key, time.time())) else: cursor = db.execute( - "INSERT OR IGNORE INTO events(account,kind,body,unique_key,created) VALUES(?,?,?,?,?)", - (account, kind, payload, unique_key, time.time())) + "INSERT OR IGNORE INTO events(id,account,kind,body,unique_key,created) VALUES(?,?,?,?,?,?)", + (id, account, kind, payload, unique_key, time.time())) inserted = cursor.rowcount > 0 + if inserted: + self._last_event_id = id + if replace: + old = self._live_events.pop((account, kind, unique_key), None) + if old: + self._live_event_bytes -= old[1] if inserted: - with self._event_condition: - self._event_revisions[account] = self._event_revisions.get(account, 0) + 1 - self._event_condition.notify_all() + self._notify_event(account) + + def _notify_event(self, account): + with self._event_condition: + self._event_revisions[account] = self._event_revisions.get(account, 0) + 1 + self._event_condition.notify_all() def event_revision(self, account): """返回进程内事件修订号,用于无竞态地建立 SSE 等待点。""" @@ -173,7 +316,11 @@ def events(self, after=0, account=None, pending=False): sql += " AND delivered=0 AND kind='notification'" with self.connection() as db: rows = db.execute(sql + " ORDER BY id LIMIT 100", args).fetchall() - return [{**dict(row), "body": json.loads(row["body"])} for row in rows] + result = [{**dict(row), "body": json.loads(row["body"])} for row in rows] + if not pending: + result.extend(copy.deepcopy(row) for row, _ in self._live_events.values() + if row['id'] > after and (account is None or row['account'] == account)) + return sorted(result, key=lambda row: row['id'])[:100] @observed('storage.acknowledge') def acknowledge(self, id): @@ -237,6 +384,7 @@ def prune_events(self, max_age=EVENT_RETENTION_SECONDS, batch=2000): def compact(self, minimum_bytes=COMPACT_MINIMUM_BYTES): """回收已删除行遗留的空闲页;空闲空间不多时不做全库重写。""" with self.lock: + self.close() probe = sqlite3.connect(self.path, timeout=30) try: page_size = probe.execute('PRAGMA page_size').fetchone()[0] @@ -266,6 +414,7 @@ def repair_oversized(self, max_database_bytes=MAINTENANCE_MAX_DATABASE_BYTES): """ with self.lock: try: + self.close() probe = sqlite3.connect(self.path, timeout=30) try: database_bytes = self._database_bytes(probe) @@ -357,3 +506,8 @@ def purge_account(self, account): self.revoked_accounts.add(account) db.execute("DELETE FROM records WHERE account=?", (account,)) db.execute("DELETE FROM events WHERE account=?", (account,)) + self._live_records = {key: value for key, value in self._live_records.items() if value[1] != account} + for key, (row, size) in list(self._live_events.items()): + if row['account'] == account: + del self._live_events[key] + self._live_event_bytes -= size diff --git a/src/wechat_decrypt_tool/local_search/service.py b/src/wechat_decrypt_tool/local_search/service.py index eb55a783..34997cce 100644 --- a/src/wechat_decrypt_tool/local_search/service.py +++ b/src/wechat_decrypt_tool/local_search/service.py @@ -102,13 +102,23 @@ async def _configure(self, account, values): return cfg def update(self, job, **changes): + # 纯展示计数不推进断点;索引数据和游标仍由 index.commit 原子保存。 + transient = bool(changes) and set(changes) <= {'read_count', 'embedded_count', 'read_batch_size_effective'} + dirty = self.store.has_live_record('index_job', job['id']) + if changes and all(job.get(key) == value for key, value in changes.items()) and (transient or not dirty): + return + if not changes and not dirty: + previous = self.store.get('index_job', job['id']) + if previous and {k: v for k, v in previous.items() if k != 'updated'} == {k: v for k, v in job.items() if k != 'updated'}: + return job.update(changes, updated=time.time()) if job['id'] in self.restarting: job['resume_on_start']=True if job['account'] in self.revoked: return - self.store.put('index_job', job, id=job['id'], account=job['account']) - # 事件表只保留每个任务的最新进度;完整快照以 records 表为准。 - progress = {key: value for key, value in job.items() if key not in EVENT_OMITTED_FIELDS} - self.store.event(job['account'], 'local_search_index', progress, unique_key=f'index_job:{job["id"]}', replace=True) + with self.store.connection() if not transient else self.store.lock: + self.store.put('index_job', job, id=job['id'], account=job['account'], transient=transient) + progress = {key: value for key, value in job.items() if key not in EVENT_OMITTED_FIELDS} + self.store.event(job['account'], 'local_search_index', progress, + unique_key=f'index_job:{job["id"]}', replace=True, transient=transient) def enrichment_version(self, account): """只检查本地提取缓存,不触发媒体分析或网络访问。""" diff --git a/tests/test_ai_message_pages.py b/tests/test_ai_message_pages.py index d9893cfc..323e8611 100644 --- a/tests/test_ai_message_pages.py +++ b/tests/test_ai_message_pages.py @@ -339,10 +339,10 @@ async def run(): service = LocalSearch(tmp_path / 'state', tmp_path / 'models', engine=engine) emitted = [] original_event = service.store.event - def record(account, kind, body, unique_key=None, replace=False): + def record(account, kind, body, unique_key=None, replace=False, transient=False): if kind == 'local_search_index': emitted.append(body) - return original_event(account, kind, body, unique_key, replace) + return original_event(account, kind, body, unique_key, replace, transient=transient) monkeypatch.setattr(service.store, 'event', record) root = model_dir(service.downloads.root, 'bge-small-zh') root.mkdir(parents=True) diff --git a/tests/test_ai_runtime_cleanup.py b/tests/test_ai_runtime_cleanup.py new file mode 100644 index 00000000..78d230c4 --- /dev/null +++ b/tests/test_ai_runtime_cleanup.py @@ -0,0 +1,69 @@ +"""运行自检必须主动释放持久连接,正常和异常路径均允许 Windows 清理临时目录。""" +import asyncio +import tempfile +from pathlib import Path + +import pytest + +from wechat_decrypt_tool.ai import runtime_check, storage, agent_service + + +def track_stores(monkeypatch): + stores = [] + + class TrackedStore(storage.AIStore): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + # 保持强引用,避免依赖垃圾回收碰巧释放 SQLite 文件句柄。 + stores.append(self) + + monkeypatch.setattr(storage, 'AIStore', TrackedStore) + return stores + + +def test_full_runtime_check_closes_owned_stores_before_temporary_directory_cleanup(monkeypatch): + stores = track_stores(monkeypatch) + report = runtime_check.check_runtime() + assert report['ok'] and report['application_graph'] == 'deepagents-v3-one-call-no-query' + assert {store.root.name for store in stores} == {'application', 'business'} + assert all(store._db is None and not store.root.exists() for store in stores) + + +@pytest.mark.parametrize('phase', ['setup', 'submit', 'validation']) +def test_checkpoint_failure_stops_workers_closes_store_and_preserves_original_error(monkeypatch, phase): + stores = track_stores(monkeypatch) + services = [] + original_agent = agent_service.AgentService + + class TrackedAgent(original_agent): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + services.append(self) + + async def submit(self, *args, **kwargs): + if phase == 'submit': + raise RuntimeError('模拟提交失败') + return await super().submit(*args, **kwargs) + + def public_run(self, *args, **kwargs): + if phase == 'validation': + raise RuntimeError('模拟校验失败') + return super().public_run(*args, **kwargs) + + monkeypatch.setattr(agent_service, 'AgentService', TrackedAgent) + if phase == 'setup': + original_put = storage.AIStore.put + + def failed_setup(self, kind, *args, **kwargs): + if kind == 'defaults': + raise RuntimeError('模拟配置失败') + return original_put(self, kind, *args, **kwargs) + + monkeypatch.setattr(storage.AIStore, 'put', failed_setup) + with tempfile.TemporaryDirectory(prefix='test-ai-runtime-cleanup-') as directory: + with pytest.raises(RuntimeError, match='模拟'): + asyncio.run(runtime_check._checkpoint(Path(directory))) + assert stores and all(store._db is None for store in stores) + assert all(service.stopping and all(worker.done() for worker in service.workers.values()) + for service in services) + assert not Path(directory).exists() diff --git a/tests/test_ai_write_budget.py b/tests/test_ai_write_budget.py new file mode 100644 index 00000000..2fc1438b --- /dev/null +++ b/tests/test_ai_write_budget.py @@ -0,0 +1,261 @@ +"""写盘预算回归:展示更新零写入,业务边界持久化,重连与恢复保持可用。""" +import asyncio +import json +import sqlite3 +import subprocess +import sys +from pathlib import Path +from concurrent.futures import ThreadPoolExecutor +from types import SimpleNamespace +from unittest.mock import patch + +import pytest + +from test_ai_agent import service, FakeModels +from test_ai_agent_sse import request +from test_ai_global_assistant import idle_run +from wechat_decrypt_tool.ai.storage import AIStore +from wechat_decrypt_tool.ai.service import AIService +from wechat_decrypt_tool.ai.agent_service import AgentService +from wechat_decrypt_tool.local_search.service import LocalSearch +from wechat_decrypt_tool.routers import ai_agent + + +def disk_record(store, kind, id): + # 绕过进程内快照,确认另一个连接真正看到的持久状态。 + with sqlite3.connect(store.path) as db: + row = db.execute('SELECT body FROM records WHERE kind=? AND id=?', (kind, id)).fetchone() + return json.loads(row[0]) if row else None + + +def changes(store): + with store.connection() as db: + return db.total_changes + + +def test_stream_updates_are_live_without_sql_writes_and_periodically_checkpoint(service): + async def run(): + _, task = await idle_run(service) + id = task['id'] + before = changes(service.store) + with patch('wechat_decrypt_tool.ai.agent_service.time.monotonic', return_value=100): + for i in range(100): + service.stream_answer(id, f'正文{i}', '正在回答') + assert changes(service.store) == before + assert service.run(id)['answer'] == '正文99' + assert disk_record(service.store, 'agent_run', id)['answer'] == '' + assert service.public_run(id, 'account')['answer'] == '正文99' + with patch('wechat_decrypt_tool.ai.agent_service.time.monotonic', return_value=102): + service.stream_answer(id, '草稿检查点', '正在回答') + assert disk_record(service.store, 'agent_run', id)['answer'] == '草稿检查点' + with service.store.connection() as db: + assert db.execute("SELECT count(*) FROM agent_piece WHERE run_id=? AND kind='timeline'", (id,)).fetchone()[0] == 1 + with patch('wechat_decrypt_tool.ai.agent_service.time.monotonic', return_value=103): + service.stream_answer(id, '检查点之后的文字', '正在回答') + reopened = AIStore(service.store.root) + try: + assert reopened.get('agent_run', id)['answer'] == '草稿检查点' + finally: + reopened.close() + asyncio.run(run()) + + +def test_service_restart_marks_checkpoint_interrupted_without_losing_saved_draft(service): + async def run(): + _, task = await idle_run(service) + with patch('wechat_decrypt_tool.ai.agent_service.time.monotonic', return_value=100): + service.stream_answer(task['id'], '第一段', '正在回答') + with patch('wechat_decrypt_tool.ai.agent_service.time.monotonic', return_value=102): + service.stream_answer(task['id'], '已提交的草稿', '正在回答') + store = AIStore(service.store.root) + restarted = AgentService(AIService(store, FakeModels(store)), service.tools) + try: + await restarted.start() + run = restarted.run(task['id']) + assert run['status'] == 'interrupted' and run['answer'] == '已提交的草稿' + assert restarted.public_run(task['id'], 'account')['can_resume'] + finally: + await restarted.stop() + store.close() + asyncio.run(run()) + + +@pytest.mark.parametrize('status', ['completed', 'cancelled', 'interrupted', 'failed']) +def test_final_state_flushes_latest_draft_and_history(service, status): + async def run(): + thread, task = await idle_run(service) + service.stream_answer(task['id'], '最后一段正文', '正在回答') + service.finish(task['id'], status) + saved = disk_record(service.store, 'agent_run', task['id']) + assert saved['answer'] == '最后一段正文' and saved['status'] == status + assert saved['timeline'][-1]['text'] == '最后一段正文' + assert not service.store.has_live_record('agent_run', task['id']) + if status == 'completed': + assert service.thread(thread['id'], 'account')['messages'][-1]['text'] == '最后一段正文' + asyncio.run(run()) + + +def test_memory_sse_wakes_replays_latest_and_keeps_account_boundary(service): + async def run(): + _, task = await idle_run(service) + with patch.object(ai_agent, 'get_agent_service', return_value=service), \ + patch.object(ai_agent, 'account_name', side_effect=lambda value: value): + response = await ai_agent.events(request(), 'account') + initial = await anext(response.body_iterator) + cursor = int(initial.split(': ')[1]) + waiting = asyncio.create_task(anext(response.body_iterator)) + await asyncio.sleep(0) + before = changes(service.store) + service.store.event('other', 'agent', {'text': '其他账号'}, transient=True) + service.stream_answer(task['id'], '实时文字', '正在回答') + event = await asyncio.wait_for(waiting, 0.5) + assert json.loads(event.split('data: ', 1)[1])['timeline_item']['text'] == '实时文字' + assert json.loads(event.split('data: ', 1)[1])['patch']['stage'] == '正在回答' + assert changes(service.store) == before + await response.body_iterator.aclose() + service.stream_answer(task['id'], '断线后的最新文字', '正在回答') + again = await ai_agent.events(request(cursor), 'account') + await anext(again.body_iterator) + replay = await anext(again.body_iterator) + assert '断线后的最新文字' in replay and '其他账号' not in replay + await again.body_iterator.aclose() + asyncio.run(run()) + + +def test_index_display_counters_do_not_write_or_advance_recovery_cursor(tmp_path): + search = LocalSearch(tmp_path, engine=SimpleNamespace(status={}, gpu_failed=False)) + job = dict(id='index', account='account', status='running', stage='reading', processed=10, + offset=10, embedded=3, read_count=10) + search.update(job) + before = changes(search.store) + for count in range(11, 111): + search.update(job, read_count=count, embedded_count=count) + assert changes(search.store) == before + assert search.store.get('index_job', 'index')['read_count'] == 110 + assert search.store.list('index_job', 'account')[0]['read_count'] == 110 + saved = disk_record(search.store, 'index_job', 'index') + assert saved['read_count'] == 10 and saved['offset'] == 10 + job.update(processed=110, offset=110, embedded=110) + search.update(job) + saved = disk_record(search.store, 'index_job', 'index') + assert saved['read_count'] == 110 and saved['offset'] == 110 + assert len(search.store.events(account='account')) == 1 + before = changes(search.store) + search.update(job) + assert changes(search.store) == before + search.store.close() + + +def test_storage_transaction_rolls_back_records_events_and_live_overlay(tmp_path): + store = AIStore(tmp_path) + store.put('agent_run', {'answer': '已提交'}, id='run', account='account') + store.put('agent_run', {'answer': '内存草稿'}, id='run', account='account', transient=True) + with pytest.raises(RuntimeError): + with store.connection(): + store.put('agent_run', {'answer': '失败的提交'}, id='run', account='account') + store.event('account', 'agent', {'text': '不应出现'}) + raise RuntimeError('模拟事务失败') + assert store.get('agent_run', 'run')['answer'] == '内存草稿' + assert disk_record(store, 'agent_run', 'run')['answer'] == '已提交' + assert store.events() == [] + store.close() + + +def test_connection_reuse_cross_thread_rollback_and_repair(tmp_path): + store = AIStore(tmp_path) + with store.connection() as first: + pass + with ThreadPoolExecutor(max_workers=2) as pool: + list(pool.map(lambda i: store.put('value', {'i': i}, id=str(i)), range(20))) + with store.connection() as second: + assert first is second + assert db_integrity(second) == 'ok' + assert store.repair_oversized(max_database_bytes=1) > 0 + assert len(store.list('value')) == 20 + with store.connection() as repaired: + assert repaired is not first and db_integrity(repaired) == 'ok' + store.close() + + +def db_integrity(db): + return db.execute('PRAGMA integrity_check').fetchone()[0] + + +def test_restart_event_ids_follow_transient_cursor_and_notifications_remain_durable(tmp_path): + store = AIStore(tmp_path) + store.event('account', 'agent', {'text': '展示'}, transient=True) + cursor = store.latest_event_id() + store.close() + restarted = AIStore(tmp_path) + restarted.event('account', 'notification', {'text': '通知'}, unique_key='notification') + assert restarted.latest_event_id() > cursor + assert restarted.events(cursor, 'account', pending=True)[0]['body']['text'] == '通知' + restarted.close() + + +def test_coalesced_live_event_keeps_verified_reference_mapping(tmp_path): + store = AIStore(tmp_path) + store.event('account', 'agent', {'version': 1, 'text': '带引用的文字', + 'citations': [{'source': 'a', 'text': '原文'}], 'references': [{'id': 'b'}]}, + unique_key='answer', transient=True) + store.event('account', 'agent', {'version': 1, 'text': '带引用的文字继续输出'}, + unique_key='answer', transient=True) + latest = store.events()[0]['body'] + assert latest['citations'][0]['source'] == 'a' and latest['references'][0]['id'] == 'b' + store.event('account', 'agent', {'version': 2, 'text': '新版本'}, unique_key='answer', transient=True) + assert 'citations' not in store.events()[0]['body'] + for i in range(300): + store.event('account', 'agent', {'text': str(i)}, unique_key=f'item:{i}', transient=True) + assert len(store._live_events) == 256 + assert store._live_event_bytes <= 8 * 1024 * 1024 + with store.connection() as db: + assert db.execute('SELECT count(*) FROM events').fetchone()[0] == 0 + store.close() + + +def test_abrupt_process_exit_retains_committed_checkpoint_and_notifications(tmp_path): + source = Path(__file__).resolve().parents[1] / 'src' + script = f""" +import sys, os +from pathlib import Path +sys.path.insert(0, {str(source)!r}) +from wechat_decrypt_tool.ai.storage import AIStore +store = AIStore(Path({str(tmp_path)!r})) +store.put('agent_run', {{'answer': '已提交的草稿'}}, id='run', account='account') +store.event('account', 'notification', {{'text': '未投递提醒'}}, unique_key='notice') +store.put('agent_run', {{'answer': '最后的展示片段'}}, id='run', account='account', transient=True) +os._exit(0) +""" + subprocess.run([sys.executable, '-c', script], check=True, timeout=30) + restarted = AIStore(tmp_path) + assert restarted.get('agent_run', 'run')['answer'] == '已提交的草稿' + assert restarted.events(account='account', pending=True)[0]['body']['text'] == '未投递提醒' + with restarted.connection() as db: + assert db_integrity(db) == 'ok' + restarted.close() + + +def test_identical_records_and_updates_do_not_write(service): + async def run(): + _, task = await idle_run(service) + before = changes(service.store) + service.update(task['id'], status='running') + service.store.put('agent_run', service.store.get('agent_run', task['id'])) + assert changes(service.store) == before + asyncio.run(run()) + + +def test_deleted_run_and_revoked_account_remove_live_state(service): + async def run(): + _, task = await idle_run(service) + service.stream_answer(task['id'], '即将删除', '正在回答') + service.store.delete('agent_run', task['id']) + assert service.store.get('agent_run', task['id']) is None + assert not any(e['body'].get('run_id') == task['id'] for e in service.store.events() + if e['unique_key'] and e['unique_key'].startswith('timeline:')) + service.store.put('index_job', {'read_count': 1}, id='job', account='account', transient=True) + service.store.event('account', 'agent', {'text': '即将清理'}, transient=True) + service.store.purge_account('account') + assert service.store.get('index_job', 'job') is None + assert service.store.events(account='account') == [] + asyncio.run(run())