Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 20 additions & 1 deletion src/wechat_decrypt_tool/ai/agent_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion src/wechat_decrypt_tool/ai/deep_runtime.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
20 changes: 13 additions & 7 deletions src/wechat_decrypt_tool/ai/service.py
Original file line number Diff line number Diff line change
Expand Up @@ -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():
Expand All @@ -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
Expand Down
86 changes: 86 additions & 0 deletions tests/test_ai_services.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import asyncio
import io
import json
import sqlite3
import sys
import time
from pathlib import Path
Expand Down Expand Up @@ -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
Expand Down
Loading