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
3 changes: 3 additions & 0 deletions src/wechat_decrypt_tool/ai/service.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
import copy
import json
import hashlib
import threading
import time
from datetime import datetime, timedelta
from typing import TypedDict
Expand Down Expand Up @@ -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):
Expand Down
17 changes: 9 additions & 8 deletions src/wechat_decrypt_tool/routers/ai.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down
13 changes: 12 additions & 1 deletion src/wechat_decrypt_tool/routers/ai_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
6 changes: 6 additions & 0 deletions src/wechat_decrypt_tool/routers/chat.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.")
Expand Down
93 changes: 91 additions & 2 deletions tests/test_ai_agent.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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()
Expand Down
Loading