diff --git a/src/wechat_decrypt_tool/ai/agent_service.py b/src/wechat_decrypt_tool/ai/agent_service.py index b8b8c1b0..a81f36ca 100644 --- a/src/wechat_decrypt_tool/ai/agent_service.py +++ b/src/wechat_decrypt_tool/ai/agent_service.py @@ -309,10 +309,29 @@ async def stop(self): @observed('agent.cancel_account', id_field='run_id') def cancel_account(self, account): + workers = [] for run in self.store.list('agent_run', account): worker = self.workers.get(run['id']) - if worker: + if worker and not worker.done(): + workers.append(worker) + if not workers: + return + loop = workers[0].get_loop() + if not loop.is_running(): + raise RuntimeError('Agent 事件循环已停止,无法确认任务退出') + try: + current_loop = asyncio.get_running_loop() + except RuntimeError: + current_loop = None + if current_loop is loop: + raise RuntimeError('账号清理不能在 Agent 事件循环中同步等待') + + async def stop_workers(): + for worker in workers: worker.cancel() + await asyncio.gather(*workers, return_exceptions=True) + + asyncio.run_coroutine_threadsafe(stop_workers(), loop).result(timeout=30) _agent = None diff --git a/src/wechat_decrypt_tool/ai/deep_runtime.py b/src/wechat_decrypt_tool/ai/deep_runtime.py index fef138d0..57ad6ca8 100644 --- a/src/wechat_decrypt_tool/ai/deep_runtime.py +++ b/src/wechat_decrypt_tool/ai/deep_runtime.py @@ -755,7 +755,7 @@ def graph_input(self, run): async def execute(self, id): run = self.run(id) - if run.get('engine_version') != 3: + if run.get('engine_version') != 3 or run['account'] in self.ai.deleted_accounts: return priority = model_priority.set(0 if not run.get('parent_run_id') else 1) group = model_group.set(run.get('parent_run_id') or id) diff --git a/src/wechat_decrypt_tool/ai/service.py b/src/wechat_decrypt_tool/ai/service.py index cf6c22b4..07834dc2 100644 --- a/src/wechat_decrypt_tool/ai/service.py +++ b/src/wechat_decrypt_tool/ai/service.py @@ -558,17 +558,21 @@ async def reset_graph(self, id): @observed('summary.purge_account', id_field='task_id') def purge_account(self, account): import sqlite3 + was_deleted = account in self.deleted_accounts + self.deleted_accounts.add(account) + from . import agent_service + try: + if agent_service._agent is not None: + agent_service._agent.cancel_account(account) + except Exception: + if not was_deleted: + self.deleted_accounts.discard(account) + raise from ..local_search import service as local_search_service if local_search_service._service is not None: local_search_service._service.purge(account) - self.deleted_accounts.add(account) ids = [t["id"] for t in self.store.list("task", account)] agent_ids = [r['id'] for r in self.store.list('agent_run', account)] - deep_ids = [f'{account}:{r["id"]}:v{version}' for r in self.store.list('agent_run', account) - if r.get('engine_version') == 3 for version in range(1, r['version'] + 1)] - from . import agent_service - if agent_service._agent is not None: - agent_service._agent.cancel_account(account) self.store.purge_account(account) path = self.store.root / "checkpoints.sqlite3" if path.exists(): @@ -586,11 +590,13 @@ def purge_account(self, account): db.executemany(f'DELETE FROM {table} WHERE thread_id=?', [(id,) for id in agent_ids]) deep_path = self.store.root / 'deepagents_checkpoints.sqlite3' if deep_path.exists(): + # 历史清理可能已删掉 agent_run,只留下以账号开头的孤儿检查点。 + prefix = f'{account}:' with sqlite3.connect(deep_path, timeout=30) as db: tables = {x[0] for x in db.execute("SELECT name FROM sqlite_master WHERE type='table'")} for table in ('checkpoints', 'writes'): if table in tables: - db.executemany(f'DELETE FROM {table} WHERE thread_id=?', [(id,) for id in deep_ids]) + db.execute(f'DELETE FROM {table} WHERE substr(thread_id,1,?)=?', (len(prefix), prefix)) _service = None diff --git a/tests/test_ai_services.py b/tests/test_ai_services.py index d2563003..42f11f84 100644 --- a/tests/test_ai_services.py +++ b/tests/test_ai_services.py @@ -1,6 +1,7 @@ import asyncio import io import json +import sqlite3 import sys import time from pathlib import Path @@ -226,6 +227,91 @@ def test_account_isolation_and_cleanup(service): assert len(service.store.list("alert", "other")) == 1 +def test_purge_account_removes_deepagent_checkpoints_including_orphans(service, monkeypatch): + from wechat_decrypt_tool.ai import agent_service + + monkeypatch.setattr(agent_service, '_agent', None) + service.store.put('agent_run', { + 'account': 'account', 'thread_id': 'thread', 'engine_version': 3, + 'checkpoint_schema': 2, 'version': 2, + }, id='current') + service.store.put('agent_run', { + 'account': 'account', 'thread_id': 'old-thread', 'engine_version': 3, + 'version': 1, + }, id='legacy') + thread_ids = [ + 'account:thread:current:v1', 'account:thread:current:v2', + 'account:legacy:v1', 'account:orphan-thread:orphan:v1', + 'account2:thread:run:v1', 'other:thread:run:v1', + ] + path = service.store.root / 'deepagents_checkpoints.sqlite3' + with sqlite3.connect(path) as db: + for table in ('checkpoints', 'writes'): + db.execute(f'CREATE TABLE {table} (thread_id TEXT PRIMARY KEY)') + db.executemany(f'INSERT INTO {table} (thread_id) VALUES (?)', + [(thread_id,) for thread_id in thread_ids]) + + service.purge_account('account') + + with sqlite3.connect(path) as db: + for table in ('checkpoints', 'writes'): + assert db.execute(f'SELECT thread_id FROM {table}').fetchall() == [ + ('account2:thread:run:v1',), ('other:thread:run:v1',)] + + +def test_purge_account_waits_for_agent_checkpoint_writer(service, monkeypatch): + from wechat_decrypt_tool.ai import agent_service + + agent = agent_service.AgentService(service) + monkeypatch.setattr(agent_service, '_agent', agent) + service.store.put('agent_run', { + 'account': 'account', 'thread_id': 'thread', 'engine_version': 3, + 'checkpoint_schema': 2, 'version': 1, + }, id='running') + path = service.store.root / 'deepagents_checkpoints.sqlite3' + with sqlite3.connect(path) as db: + db.execute('CREATE TABLE checkpoints (thread_id TEXT PRIMARY KEY)') + db.execute('CREATE TABLE writes (thread_id TEXT PRIMARY KEY)') + + async def run(): + started = asyncio.Event() + + async def worker(): + started.set() + try: + await asyncio.Event().wait() + finally: + await asyncio.sleep(.05) + with sqlite3.connect(path) as db: + db.execute('INSERT INTO checkpoints VALUES (?)', + ('account:thread:running:v1',)) + + task = asyncio.create_task(worker()) + agent.workers['running'] = task + await started.wait() + await asyncio.to_thread(service.purge_account, 'account') + assert task.done() + with sqlite3.connect(path) as db: + assert db.execute('SELECT count(*) FROM checkpoints').fetchone()[0] == 0 + + asyncio.run(run()) + + +def test_deleted_account_does_not_start_new_agent_run(service): + from wechat_decrypt_tool.ai.agent_service import AgentService + + agent = AgentService(service) + service.store.put('agent_run', { + 'account': 'account', 'thread_id': 'thread', 'engine_version': 3, + 'checkpoint_schema': 2, 'version': 1, + }, id='late') + service.deleted_accounts.add('account') + + asyncio.run(agent.execute('late')) + + assert not (service.store.root / 'deepagents_checkpoints.sqlite3').exists() + + def test_docx_xlsx_pptx_and_pdf_parsers(): from docx import Document from openpyxl import Workbook