From c2858e233a9f5ca6639e1b614612c6a0c962d1ea Mon Sep 17 00:00:00 2001 From: superG Date: Sun, 27 Sep 2026 21:28:50 +0800 Subject: [PATCH] fix(ai): serialize account deletion and agent reactivation --- src/wechat_decrypt_tool/ai/service.py | 3 + src/wechat_decrypt_tool/routers/ai.py | 17 ++-- src/wechat_decrypt_tool/routers/ai_agent.py | 13 ++- src/wechat_decrypt_tool/routers/chat.py | 6 ++ tests/test_ai_agent.py | 93 ++++++++++++++++++++- 5 files changed, 121 insertions(+), 11 deletions(-) diff --git a/src/wechat_decrypt_tool/ai/service.py b/src/wechat_decrypt_tool/ai/service.py index cf6c22b4..249dda19 100644 --- a/src/wechat_decrypt_tool/ai/service.py +++ b/src/wechat_decrypt_tool/ai/service.py @@ -6,6 +6,7 @@ import copy import json import hashlib +import threading import time from datetime import datetime, timedelta from typing import TypedDict @@ -79,6 +80,8 @@ def __init__(self, store=None, models=None, reader=None): self.active_id = None self.stopping = False self.deleted_accounts = set() + # 删除账号与重新接入后的首次写入必须互斥。 + self.account_lifecycle_lock = threading.Lock() self.rule_locks = {} def check(self, id): diff --git a/src/wechat_decrypt_tool/routers/ai.py b/src/wechat_decrypt_tool/routers/ai.py index aabfa5dc..4087101c 100644 --- a/src/wechat_decrypt_tool/routers/ai.py +++ b/src/wechat_decrypt_tool/routers/ai.py @@ -240,14 +240,15 @@ def conversations(account: str): @router.post("/tasks") def create_task(body: TaskInput): service = get_ai_service() - data = body.model_dump() - data["account"] = account_name(body.account) - service.deleted_accounts.discard(data["account"]) - service.store.revoked_accounts.discard(data["account"]) - try: - return service.create_task(data) - except ProviderFailure as exc: - raise HTTPException(422, str(exc)) from None + with service.account_lifecycle_lock: + data = body.model_dump() + data["account"] = account_name(body.account) + service.deleted_accounts.discard(data["account"]) + service.store.revoked_accounts.discard(data["account"]) + try: + return service.create_task(data) + except ProviderFailure as exc: + raise HTTPException(422, str(exc)) from None @router.get("/tasks") diff --git a/src/wechat_decrypt_tool/routers/ai_agent.py b/src/wechat_decrypt_tool/routers/ai_agent.py index a6cce244..a81e4c0f 100644 --- a/src/wechat_decrypt_tool/routers/ai_agent.py +++ b/src/wechat_decrypt_tool/routers/ai_agent.py @@ -52,7 +52,18 @@ def threads(account: str, username: str = '', unassigned: bool = False): @router.post('/threads') async def create_thread(body: ThreadInput): try: - return await get_agent_service().create_thread(account_name(body.account), body.username, body.title) + service = get_agent_service() + lock = service.ai.account_lifecycle_lock + while not lock.acquire(blocking=False): + await asyncio.sleep(.05) + try: + owner = account_name(body.account) + # 删除完成后再验证账号;重新导入的同名账号可恢复写入。 + service.ai.deleted_accounts.discard(owner) + service.store.revoked_accounts.discard(owner) + return await service.create_thread(owner, body.username, body.title) + finally: + lock.release() except ValueError as exc: raise HTTPException(422, str(exc)) from None diff --git a/src/wechat_decrypt_tool/routers/chat.py b/src/wechat_decrypt_tool/routers/chat.py index c4e3cad6..3795f965 100644 --- a/src/wechat_decrypt_tool/routers/chat.py +++ b/src/wechat_decrypt_tool/routers/chat.py @@ -4494,6 +4494,12 @@ def get_chat_account_info(account: Optional[str] = None): @router.delete("/api/chat/account", summary="删除当前账号在本项目中的数据") def delete_chat_account(account: str): + from ..ai.service import get_ai_service + with get_ai_service().account_lifecycle_lock: + return _delete_chat_account(account) + + +def _delete_chat_account(account: str): requested_account_name = str(account or "").strip() if not requested_account_name: raise HTTPException(status_code=400, detail="Missing account.") diff --git a/tests/test_ai_agent.py b/tests/test_ai_agent.py index 98ba0486..6659c755 100644 --- a/tests/test_ai_agent.py +++ b/tests/test_ai_agent.py @@ -1,21 +1,23 @@ import asyncio import sys +import threading import time from pathlib import Path from unittest.mock import patch import httpx import pytest -from fastapi import FastAPI +from fastapi import FastAPI, HTTPException sys.path.insert(0, str(Path(__file__).resolve().parents[1] / 'src')) from wechat_decrypt_tool.ai.agent_service import AgentService, Revised from wechat_decrypt_tool.ai.agent_schemas import AgentControl from wechat_decrypt_tool.ai.agent_schemas import AgentAction +from wechat_decrypt_tool.ai.agent_schemas import ThreadInput from wechat_decrypt_tool.ai.providers import ModelService, ProviderFailure, model_attempt_hook from wechat_decrypt_tool.ai.service import AIService from wechat_decrypt_tool.ai.storage import AIStore -from wechat_decrypt_tool.routers import ai_agent +from wechat_decrypt_tool.routers import ai_agent, chat from langchain_core.messages import AIMessageChunk SOURCE = 'a' * 24 @@ -339,6 +341,93 @@ async def run(): asyncio.run(run()) +def test_reimported_account_can_create_and_send_agent_thread(service, tmp_path): + from wechat_decrypt_tool.ai import agent_service + + account_dir = tmp_path / 'account' + account_dir.mkdir() + + def resolve_account(account): + if not account_dir.exists(): + raise HTTPException(404, '账号不存在') + return account + + async def run(): + original = await service.create_thread('account', 'friend', '旧对话') + service.ai.purge_account('account') + assert service.store.get('agent_thread', original['id']) is None + assert 'account' in service.store.revoked_accounts + account_dir.rmdir() + account_dir.mkdir() # 重新导入同名账号后,目录再次出现。 + + app = FastAPI(); app.include_router(ai_agent.router) + transport = httpx.ASGITransport(app=app, client=('127.0.0.1', 100)) + async with httpx.AsyncClient(transport=transport, base_url='http://localhost') as client: + created = await client.post('/api/ai/agent/threads', + json={'account': 'account', 'username': 'friend'}) + assert created.status_code == 200, created.text + thread_id = created.json()['id'] + assert service.store.get('agent_thread', thread_id) is not None + sent = await client.post(f'/api/ai/agent/threads/{thread_id}/messages', + params={'account': 'account'}, + json={'text': '找报价', 'request_id': 'first'}) + assert sent.status_code == 200, sent.text + await service.workers[sent.json()['id']] + assert service.run(sent.json()['id'])['status'] == 'completed' + + with patch.object(agent_service, '_agent', service), \ + patch.object(ai_agent, 'get_agent_service', return_value=service), \ + patch.object(ai_agent, 'account_name', side_effect=resolve_account): + asyncio.run(run()) + + +def test_create_thread_waits_until_account_deletion_finishes(service, tmp_path): + from wechat_decrypt_tool.ai import agent_service + from wechat_decrypt_tool.ai import service as ai_service_module + + account_dir = tmp_path / 'account' + account_dir.mkdir() + purged = threading.Event() + finish_delete = threading.Event() + + def delete_account(_account): + service.ai.purge_account('account') + purged.set() + if not finish_delete.wait(5): + raise RuntimeError('删除流程未被释放') + account_dir.rmdir() + return {'status': 'success'} + + def resolve_account(account): + if not account_dir.exists(): + raise HTTPException(404, '账号不存在') + return account + + async def run(): + deleting = asyncio.create_task(asyncio.to_thread(chat.delete_chat_account, 'account')) + try: + assert await asyncio.to_thread(purged.wait, 5) + creating = asyncio.create_task(ai_agent.create_thread( + ThreadInput(account='account', username='friend'))) + await asyncio.sleep(.05) + assert not creating.done() + assert 'account' in service.store.revoked_accounts + finally: + finish_delete.set() + await deleting + with pytest.raises(HTTPException) as error: + await creating + assert error.value.status_code == 404 + assert service.store.list('agent_thread', 'account') == [] + + with patch.object(agent_service, '_agent', service), \ + patch.object(ai_service_module, 'get_ai_service', return_value=service.ai), \ + patch.object(chat, '_delete_chat_account', side_effect=delete_account), \ + patch.object(ai_agent, 'get_agent_service', return_value=service), \ + patch.object(ai_agent, 'account_name', side_effect=resolve_account): + asyncio.run(run()) + + def test_scope_shrink_during_run_does_not_reapply_previous_expansion(service): async def run(): service.model.waiting = asyncio.Event()